@@ -6,7 +6,6 @@ from django.contrib import admin
from django.contrib.admin import BooleanFieldListFilter, DateFieldListFilter
from django.contrib.admin.models import LogEntry
from django.contrib.auth.admin import UserAdmin
-from django.contrib.contenttypes.models import ContentType
from django.db.models import Count, QuerySet
from django.http import HttpRequest, HttpResponse
from django.utils import timezone
@@ -401,15 +400,13 @@ class CompanyIPWhitelistAdmin(admin.ModelAdmin):
inlines = [CompanyIPInline]
def log_change(self, request: HttpRequest, object: Any, message: list[dict[str, any]]) -> LogEntry:
- ct = ContentType.objects.get_for_model(object, for_concrete_model=False)
return [
- LogEntry.objects.log_action(
- request.user.pk,
- ct.pk,
- object.pk,
- str(object),
- 1,
- [message_entry],
+ LogEntry.objects.log_actions(
+ user_id=request.user.pk,
+ queryset=[object],
+ action_flag=1,
+ change_message=[message_entry],
+ single_object=True,
)
for message_entry in message
if not message_entry.get('changed', None)
@@ -424,5 +421,5 @@ class LogEntryAdmin(admin.ModelAdmin):
def log_change(self, *args, **kwargs) -> LogEntry:
return None
- def log_deletion(self, *args, **kwargs) -> LogEntry:
- return None
+ def log_deletions(self, *args, **kwargs) -> list[LogEntry]:
+ return []
@@ -9,6 +9,3 @@ class AuthenticationConfig(AppConfig):
default_auto_field = 'django.db.models.BigAutoField'
name = 'authentication'
verbose_name = 'Пользователи'
-
- def ready(self):
- from .signals import invalidate_user_cache
@@ -1,20 +0,0 @@
-from cacheops import cache
-from cacheops.getset import dnfs_to_conj_keys
-
-from authentication.models import BusinessAccount, CustomUserModel
-
-from django.db.models.signals import post_save, post_delete
-from django.dispatch import receiver
-
-
-@receiver([post_save, post_delete], sender=BusinessAccount)
-def invalidate_user_cache(sender, instance, signal, **kwargs):
- cache_keys = cache.conn.smembers(dnfs_to_conj_keys(
- '',
- {'authentication_customusermodel': [{'uid': instance.user_id}]}
- )[0])
- for key in cache_keys:
- data = cache.get(key.decode())
- if isinstance(data, list) and isinstance((user := data[0]), CustomUserModel):
- user.business_account = instance if signal == post_save else None
- cache.set(key.decode(), [user])
\ No newline at end of file
@@ -40,6 +40,7 @@ SYSTEM_APPS = [
'django.contrib.contenttypes',
'django.contrib.sessions',
'django.contrib.messages',
+ 'django.contrib.postgres',
'django.contrib.staticfiles',
]
@@ -54,12 +55,12 @@ EXTERNAL_APPS = [
'drf_social_oauth2',
'dj_rest_auth',
'dj_rest_auth.registration',
- 'django_minio_backend',
+ 'django_minio_backend.apps.DjangoMinioBackendConfig',
'drf_spectacular',
'drf_spectacular_sidecar',
'ordered_model',
'import_export',
- 'cacheops',
+ 'cachalot',
'django_celery_beat',
]
@@ -212,20 +213,31 @@ JWT_SETTINGS = {
DEFAULT_AUTO_FIELD = 'django.db.models.BigAutoField'
-STORAGES = {
- 'default': {'BACKEND': 'django_minio_backend.models.MinioBackend'},
- 'staticfiles': {'BACKEND': 'django.contrib.staticfiles.storage.StaticFilesStorage'},
-}
+# MinIO
+MINIO_BUCKET_CHECK_ON_SAVE = env.bool(
+ 'MINIO_BUCKET_CHECK_ON_SAVE',
+ default=False,
+)
-MINIO_BUCKET_CHECK_ON_SAVE = env.bool('MINIO_BUCKET_CHECK_ON_SAVE', default=False)
MINIO_ENDPOINT = env.str('MINIO_ENDPOINT')
MINIO_USE_HTTPS = env.bool('MINIO_USE_HTTPS', default=False)
-MINIO_EXTERNAL_ENDPOINT = env.str('MINIO_EXTERNAL_ENDPOINT', default='localhost:9000')
-MINIO_EXTERNAL_ENDPOINT_USE_HTTPS = env.bool('MINIO_EXTERNAL_ENDPOINT_USE_HTTPS', default=False)
+
+MINIO_EXTERNAL_ENDPOINT = env.str(
+ 'MINIO_EXTERNAL_ENDPOINT',
+ default='localhost:9000',
+)
+MINIO_EXTERNAL_ENDPOINT_USE_HTTPS = env.bool(
+ 'MINIO_EXTERNAL_ENDPOINT_USE_HTTPS',
+ default=False,
+)
MINIO_ACCESS_KEY = env.str('MINIO_ACCESS_KEY')
MINIO_SECRET_KEY = env.str('MINIO_SECRET_KEY')
+
+MINIO_STATIC_FILES_BUCKET = 'air-static'
+MINIO_DEFAULT_BUCKET = 'air-media'
+
MINIO_PRIVATE_BUCKETS = [
'air-messages',
'air-achievements',
@@ -235,12 +247,38 @@ MINIO_PRIVATE_BUCKETS = [
'air-models',
'air-media-presets',
'air-voices',
+ MINIO_STATIC_FILES_BUCKET,
+ MINIO_DEFAULT_BUCKET,
]
-MINIO_STATIC_FILES_BUCKET = 'air-static'
-MINIO_PRIVATE_BUCKETS.append(MINIO_STATIC_FILES_BUCKET)
-MINIO_MEDIA_FILES_BUCKET = 'air-media'
-MINIO_PRIVATE_BUCKETS.append(MINIO_MEDIA_FILES_BUCKET)
+MINIO_PUBLIC_BUCKETS = []
+
+STORAGES = {
+ 'default': {
+ 'BACKEND': 'django_minio_backend.models.MinioBackend',
+ 'OPTIONS': {
+ 'MINIO_ENDPOINT': MINIO_ENDPOINT,
+ 'MINIO_EXTERNAL_ENDPOINT': MINIO_EXTERNAL_ENDPOINT,
+ 'MINIO_EXTERNAL_ENDPOINT_USE_HTTPS': (
+ MINIO_EXTERNAL_ENDPOINT_USE_HTTPS
+ ),
+ 'MINIO_ACCESS_KEY': MINIO_ACCESS_KEY,
+ 'MINIO_SECRET_KEY': MINIO_SECRET_KEY,
+ 'MINIO_USE_HTTPS': MINIO_USE_HTTPS,
+ 'MINIO_PRIVATE_BUCKETS': MINIO_PRIVATE_BUCKETS,
+ 'MINIO_PUBLIC_BUCKETS': MINIO_PUBLIC_BUCKETS,
+ 'MINIO_DEFAULT_BUCKET': MINIO_DEFAULT_BUCKET,
+ 'MINIO_STATIC_FILES_BUCKET': MINIO_STATIC_FILES_BUCKET,
+ 'MINIO_BUCKET_CHECK_ON_SAVE': MINIO_BUCKET_CHECK_ON_SAVE,
+ 'MINIO_CONSISTENCY_CHECK_ON_START': False,
+ },
+ },
+ 'staticfiles': {
+ 'BACKEND': (
+ 'django.contrib.staticfiles.storage.StaticFilesStorage'
+ ),
+ },
+}
SPECTACULAR_SETTINGS = {
'TITLE': 'AIR',
@@ -463,22 +501,8 @@ if (SENTRY_URL := env.str('SENTRY_URL', '')) and RELEASE and ENVIRONMENT:
],
)
-CACHEOPS_REDIS = env.str('CACHEOPS_REDIS', CACHES['default']['LOCATION'])
-CACHEOPS_DEGRADE_ON_FAILURE = True
-
-if CACHEOPS_REDIS:
- CACHEOPS = {
- # 'authentication.*': {'ops': 'all', 'timeout': 60 * 60},
- 'authentication.companyipwhitelist': {'ops': 'all', 'timeout': 60 * 60},
- 'ml_model.*': {'ops': 'all', 'timeout': 60 * 60},
- 'tools.chats.*': {'ops': 'all', 'timeout': 60 * 60},
- 'tools.media.*': {'ops': 'all', 'timeout': 60 * 60},
- 'payments.paymentplan': {'ops': 'all', 'timeout': 60 * 60},
- 'payments.invoice': {'ops': 'all', 'timeout': 60 * 60 * 24 * 7},
- 'messages.*': {'ops': 'all', 'timeout': 60 * 60},
- 'reports.*': {'ops': 'all', 'timeout': 60 * 60},
- 'token_blacklist.outstandingtoken': {'ops': 'get', 'timeout': 60 * 60 * 24},
- }
+CACHALOT_CACHE = 'default'
+CACHALOT_TIMEOUT = 60 * 60
# Unleash settings
UNLEASH_API_URL = env.str('UNLEASH_API_URL', 'https://example.com')
@@ -0,0 +1,21 @@
+from collections.abc import Sequence
+
+import pytest
+
+
+class PytestTestRunner:
+ def __init__(self, verbosity: int = 1, **kwargs) -> None:
+ self.verbosity = verbosity
+
+ def run_tests(
+ self,
+ test_labels: Sequence[str] | None = None,
+ extra_args: Sequence[str] | None = None,
+ **kwargs,
+ ) -> int:
+ args = list(test_labels or ('tests',))
+ args.extend(extra_args or ())
+ if self.verbosity > 1:
+ args.insert(0, f'-{"v" * min(self.verbosity, 3)}')
+
+ return pytest.main(args)
@@ -1,7 +1,7 @@
from abc import abstractmethod
from typing import Any
-from cacheops import invalidate_all
+from cachalot.api import invalidate
from django.test import TestCase
from ninja.testing import TestClient
from rest_framework_simplejwt.tokens import RefreshToken
@@ -21,7 +21,7 @@ class BaseAPITest(TestCase):
def test_unauthorized_status_code(self) -> None: ...
def setUp(self):
- invalidate_all()
+ invalidate()
super().setUp()
@@ -104,4 +104,3 @@ class BaseAuthorizedAPITest(BaseAPITest):
@abstractmethod
def test_authorized_status_code(self) -> None: ...
-
@@ -0,0 +1,10 @@
+class MessageService:
+ @classmethod
+ def prepare_output_message(cls, reasoning_text: str, output_text: str) -> str:
+ reasoning = reasoning_text.strip()
+ output = output_text.strip()
+ if reasoning and output:
+ return f'{reasoning}\n{output}'
+ if reasoning:
+ return f'{reasoning}'
+ return output
@@ -7,6 +7,7 @@ from typing import Any, Generator, TypeAlias, TypedDict
import httpx
from backend import settings
+from messages.services.message_service import MessageService
from ml_model.exceptions import (
FileExtensionNotSupported,
GenerationException,
@@ -15,6 +16,7 @@ from ml_model.exceptions import (
RequestBlocked,
)
from poller.models import Proxy
+from tools.chats.domain import RawSSEChunk
logger = logging.getLogger(__name__)
@@ -70,7 +72,7 @@ class BytedanceVideoTaskResponse(TypedDict, total=False):
RunChatResult: TypeAlias = tuple[str, int, int]
-RunStreamChatResult: TypeAlias = Generator[str, None, BytedanceUsage]
+RunStreamChatResult: TypeAlias = Generator[RawSSEChunk, None, BytedanceUsage]
RunImageResult: TypeAlias = list[str]
RunVideoResult: TypeAlias = tuple[str, int]
BytedanceRunResult: TypeAlias = RunChatResult | RunImageResult | RunVideoResult
@@ -131,7 +133,6 @@ class BytedanceModelArkAdapter:
return str(error.get('code', '')) == 'InputTextSensitiveContentDetected'
-
@classmethod
def _extract_chat_answer(
cls,
@@ -144,11 +145,14 @@ class BytedanceModelArkAdapter:
for choice in choices
if choice.get('finish_reason') in (BytedanceFinishReason.STOP, BytedanceFinishReason.LENGTH)
and choice.get('message')
- and choice['message'].get('content') is not None
+ and (
+ choice['message'].get('content') is not None
+ or choice['message'].get('reasoning_content') is not None
+ )
]
if stop_choices:
- content = ','.join(str(choice['message']['content']) for choice in stop_choices)
- # TODO: включить после разделения reasoning и content в хранении
+ content = ','.join(choice['message'].get('content') or '' for choice in stop_choices)
+ reasoning = ''
if include_reasoning:
reasoning = ','.join(
str(reasoning)
@@ -156,11 +160,7 @@ class BytedanceModelArkAdapter:
if choice.get('message')
and (reasoning := choice['message'].get('reasoning_content')) is not None
)
- if reasoning and content:
- return f'**Рассуждение:**\n\n{reasoning}\n\n**Основная мысль:**\n\n{content}'
- if reasoning:
- return reasoning
- return content
+ return MessageService.prepare_output_message(reasoning, content)
cls._raise_by_error_payload(data, choices)
@@ -286,8 +286,6 @@ class BytedanceModelArkAdapter:
)
raise GenerationException
usage: BytedanceUsage = {}
- reasoning_started = False
- content_started = False
for line in resp.iter_lines():
if not line:
continue
@@ -305,16 +303,10 @@ class BytedanceModelArkAdapter:
delta = choices[0].get('delta', {})
reasoning_chunk = delta.get('reasoning_content') or ''
if include_reasoning and reasoning_chunk:
- if not reasoning_started:
- yield '**Рассуждение:**\n\n'
- reasoning_started = True
- yield reasoning_chunk
+ yield RawSSEChunk(event='think', data={'content': reasoning_chunk})
chunk = delta.get('content') or ''
if chunk:
- if include_reasoning and reasoning_started and not content_started:
- yield '\n\n**Основная мысль:**\n\n'
- content_started = True
- yield chunk
+ yield RawSSEChunk(event='token', data={'content': chunk})
if (
fr := choices[0].get('finish_reason')
) and fr not in (
@@ -7,7 +7,9 @@ import httpx
import tiktoken
from backend import settings
+from messages.services.message_service import MessageService
from poller.models import Proxy
+from tools.chats.domain import RawSSEChunk
from .models import ModelResponse
@@ -35,7 +37,7 @@ class OpenrouterAdapter:
@classmethod
def run_streaming_api(
cls, version: str, messages: list, callback_data: dict, model_name: str
- ) -> Iterator[str]:
+ ) -> Iterator[RawSSEChunk]:
for proxy in Proxy.objects.all():
with httpx.Client(
base_url=cls.BASE_URL,
@@ -55,8 +57,6 @@ class OpenrouterAdapter:
},
) as resp:
content = ''
- # reasoning используем только для фоллбэк-подсчёта токенизатора
- # в ответ не кладём, заполняет буфер истории сообщений
reasoning = ''
input_tokens = output_tokens = cost = 0
for line in resp.iter_lines():
@@ -70,11 +70,14 @@ class OpenrouterAdapter:
try:
data_obj = json.loads(data)
- chunk = data_obj['choices'][0]['delta'].get('content') or ''
- reasoning += data_obj['choices'][0]['delta'].get('reasoning') or ''
- if chunk:
- content += chunk
- yield chunk
+ content_chunk = data_obj['choices'][0]['delta'].get('content') or ''
+ reasoning_chunk = data_obj['choices'][0]['delta'].get('reasoning') or ''
+ if content_chunk:
+ content += content_chunk
+ yield RawSSEChunk(event='token', data={'content': content_chunk})
+ if reasoning_chunk:
+ reasoning += reasoning_chunk
+ yield RawSSEChunk(event='think', data={'content': reasoning_chunk})
if data_obj.get('usage'):
input_tokens = data_obj['usage']['prompt_tokens']
output_tokens = data_obj['usage']['completion_tokens']
@@ -100,15 +103,24 @@ class OpenrouterAdapter:
cls, version: str, messages: list, callback_data: dict, model_name: str
) -> ModelResponse:
stream = cls.run_streaming_api(version, messages, callback_data, model_name)
- content_parts: list[str] = []
+ reasoning = ''
+ content = ''
try:
while True:
chunk = next(stream)
if chunk:
- content_parts.append(chunk)
+ if chunk.event == 'think':
+ reasoning += chunk.data['content']
+ else:
+ content += chunk.data['content']
except StopIteration as exc:
input_tokens, output_tokens, cost = exc.value
- return ModelResponse(''.join(content_parts), input_tokens, output_tokens, cost)
+ return ModelResponse(
+ MessageService.prepare_output_message(reasoning, content),
+ input_tokens,
+ output_tokens,
+ cost,
+ )
@classmethod
def _fallback_tokenize(cls, model_name: str, messages: list, content: str) -> tuple[int, int]:
@@ -0,0 +1,48 @@
+from django.core.management.base import BaseCommand, CommandError
+
+from core.test_runner import PytestTestRunner
+
+
+class Command(BaseCommand):
+ help = 'Run the pytest suite through the Django management interface.'
+
+ def add_arguments(self, parser) -> None:
+ parser.add_argument('test_labels', nargs='*')
+ parser.add_argument(
+ '--model-slug',
+ action='append',
+ default=[],
+ help='Run ML model API contracts only for this slug. May be repeated.',
+ )
+ parser.add_argument(
+ '--profile-resources',
+ action='store_true',
+ help='Report wall time, CPU time, and peak RSS for each test.',
+ )
+ parser.add_argument(
+ '--provider-smoke',
+ action='store_true',
+ help='Run one explicitly selected model against its real provider.',
+ )
+ parser.add_argument(
+ '--pytest-arg',
+ action='append',
+ default=[],
+ help='Forward an additional argument to pytest. May be repeated.',
+ )
+
+ def handle(self, *args, **options) -> None:
+ pytest_args = list(options['pytest_arg'])
+ test_labels = options['test_labels']
+ for slug in options['model_slug']:
+ pytest_args.extend(('--model-slug', slug))
+ if options['profile_resources']:
+ pytest_args.append('--profile-resources')
+ if options['provider_smoke']:
+ pytest_args.extend(('--provider-smoke', '-s'))
+ test_labels = ('tests/ml_models/test_provider_smoke.py',)
+
+ runner = PytestTestRunner(verbosity=options['verbosity'])
+ exit_code = runner.run_tests(test_labels, extra_args=pytest_args)
+ if exit_code:
+ raise CommandError(f'pytest exited with status {exit_code}')
@@ -1,6 +1,6 @@
from abc import ABC, abstractmethod
from decimal import Decimal
-from typing import Any, Generator, Never
+from typing import Any, Iterator, Never
from asgiref.sync import async_to_sync
from googletrans import Translator
@@ -86,4 +86,4 @@ class SimpleService(ABC):
class StreamSimpleService(SimpleService):
@abstractmethod
- def make_stream(self, input_message: Message, save: bool = True) -> Generator: ...
+ def make_stream(self, input_message: Message, save: bool = True) -> Iterator: ...
@@ -37,6 +37,7 @@ from payments.exceptions.insufficient_balance import InsufficientBalance
from payments.selectors.payment_plan_selector import PaymentPlanSelector
from poller.models import Proxy
+from tools.chats.domain import RawSSEChunk
class Chatgpt(Chatgpt_4, StreamSimpleService, OpenAIStreamMixin):
@@ -265,7 +266,7 @@ class Chatgpt(Chatgpt_4, StreamSimpleService, OpenAIStreamMixin):
)
return self.save_results([response], process_time, generated_image, save)
- def make_stream(self, input_message: Message, save: bool = True) -> Iterator[str]:
+ def make_stream(self, input_message: Message, save: bool = True) -> Iterator[RawSSEChunk]:
ctx: dict[str, Any] = {}
content_parts: list[str] = []
input_tokens = output_tokens = 0
@@ -284,7 +285,7 @@ class Chatgpt(Chatgpt_4, StreamSimpleService, OpenAIStreamMixin):
chunk = next(stream)
if chunk:
content_parts.append(chunk)
- yield chunk
+ yield RawSSEChunk(event='token', data={'content': chunk})
except StopIteration as exc:
input_tokens, output_tokens = exc.value or (0, 0)
break
@@ -11,6 +11,7 @@ from messages.models import Message
from ml_model.exceptions import ModelVersionNotAvailable, PaidPlanRequiredError
from ml_model.services.chatgpt import Chatgpt
from poller.models import Proxy
+from tools.chats.domain import RawSSEChunk
class Chatgpt_5_4(Chatgpt):
@@ -94,7 +95,7 @@ class Chatgpt_5_4(Chatgpt):
price += self.TOKENS_COST[model]['generated_image']
return price.quantize(Decimal('0.1'), rounding='ROUND_UP')
- def make_stream(self, input_message: Message, save: bool = True) -> Iterator[str]:
+ def make_stream(self, input_message: Message, save: bool = True) -> Iterator[RawSSEChunk]:
return (yield from super().make_stream(input_message, save))
def _build_payload(
@@ -11,6 +11,7 @@ from PIL import Image
from django.utils.translation import gettext
from messages.models import Message
+from messages.services.message_service import MessageService
from ml_model.adapters.openrouter import OpenrouterAdapter
from ml_model.exceptions import (
CorruptedFileError,
@@ -25,6 +26,7 @@ from ml_model.services.serper_mixin import SerperMixin
from payments.exceptions.insufficient_balance import InsufficientBalance
from payments.selectors.payment_plan_selector import PaymentPlanSelector
from poller.models import Proxy
+from tools.chats.domain import RawSSEChunk
from tools.chats.models import Chat
from tools.copywrite.models import Copywrite
from tools.public_api.models import APIStore
@@ -132,7 +134,7 @@ class Claude(SerperMixin, StreamSimpleService):
)
return self.save_results(result.content, process_time, save)
- def make_stream(self, input_message: Message, save: bool = True) -> Iterator[str]:
+ def make_stream(self, input_message: Message, save: bool = True) -> Iterator[RawSSEChunk]:
start_time = time.time()
version_slug = input_message.info.get('version')
if version_slug is None or version_slug not in self.TOKENS_COST:
@@ -140,9 +142,10 @@ class Claude(SerperMixin, StreamSimpleService):
model_slug = f'anthropic/{version_slug}'
callback_data = self._build_callback_data(input_message)
messages, embedding_tokens = self._prepare_messages(input_message, version_slug, callback_data)
- content_parts: list[str] = []
input_tokens = output_tokens = 0
cost = 0
+ reasoning = ''
+ content = ''
result = ''
try:
@@ -151,13 +154,16 @@ class Claude(SerperMixin, StreamSimpleService):
while True:
chunk = next(stream)
if chunk:
- content_parts.append(chunk)
+ if chunk.event == 'think':
+ reasoning += chunk.data['content']
+ else:
+ content += chunk.data['content']
yield chunk
except StopIteration as exc:
input_tokens, output_tokens, cost = exc.value
finally:
- if content_parts:
- result = ''.join(content_parts)
+ if reasoning or content:
+ result = MessageService.prepare_output_message(reasoning, content)
process_time = timedelta(seconds=(time.time() - start_time))
self.handle_invoice(
input_message.content_object.model,
@@ -12,6 +12,7 @@ import filetype
from PIL import Image
from messages.models import Message
+from messages.services.message_service import MessageService
from ml_model.adapters.bytedance_model_ark import BytedanceContentType, BytedanceModelArkAdapter
from ml_model.services.FileService import FileProcessingService
from ml_model.exceptions import GenerationException, ModelVersionNotAvailable
@@ -19,6 +20,7 @@ from ml_model.services.base import SimpleService
from payments.exceptions.insufficient_balance import InsufficientBalance
from payments.selectors.payment_plan_selector import PaymentPlanSelector
from ml_model.tasks import bytedance_model_ark_run, stream_bytedance_model_ark_run
+from tools.chats.domain import RawSSEChunk
from tools.chats.models import Chat
from tools.copywrite.models import Copywrite
from tools.public_api.models import APIStore
@@ -284,12 +286,13 @@ class Dola_Seed(SimpleService):
msgs = self.save_results(result[0], process_time, save)
return msgs
- def make_stream(self, input_message: Message, save: bool = True) -> Iterator[str]:
+ def make_stream(self, input_message: Message, save: bool = True) -> Iterator[RawSSEChunk]:
version, callback_data, messages = self._prepare_data(input_message)
model = self.VERSION_MAPPING[version]
start_time = time.time()
- content_parts: list[str] = []
input_tokens = output_tokens = 0
+ reasoning = ''
+ content = ''
result = ''
try:
@@ -303,15 +306,18 @@ class Dola_Seed(SimpleService):
while True:
chunk = next(stream)
if chunk:
- content_parts.append(chunk)
+ if chunk.event == 'think':
+ reasoning += chunk.data['content']
+ else:
+ content += chunk.data['content']
yield chunk
except StopIteration as exc:
usage = exc.value or {}
input_tokens = int(usage.get('prompt_tokens') or 0)
output_tokens = int(usage.get('completion_tokens') or 0)
finally:
- if content_parts:
- result = ''.join(content_parts)
+ if reasoning or content:
+ result = MessageService.prepare_output_message(reasoning, content)
if not (input_tokens + output_tokens):
input_text_parts = []
for message in messages:
@@ -10,6 +10,7 @@ import filetype
from PIL import Image
from messages.models import Message
+from messages.services.message_service import MessageService
from ml_model.adapters.openrouter import OpenrouterAdapter
from ml_model.exceptions import CorruptedFileError, FileExtensionNotSupported, ModelVersionNotAvailable
from ml_model.services.EmbeddingService import EmbeddingService
@@ -17,6 +18,7 @@ from ml_model.services.FileService import FileProcessingService
from ml_model.services.base import StreamSimpleService
from ml_model.tasks import openrouter_run
from poller.models import Proxy
+from tools.chats.domain import RawSSEChunk
from tools.chats.models import Chat
from tools.copywrite.models import Copywrite
from tools.public_api.models import APIStore
@@ -94,7 +96,7 @@ class Gemini_3_1(StreamSimpleService):
)
return self.save_results(result[0], process_time, save)
- def make_stream(self, input_message: Message, save: bool = True) -> Iterator[str]:
+ def make_stream(self, input_message: Message, save: bool = True) -> Iterator[RawSSEChunk]:
start_time = time.time()
version_slug = input_message.info.get('version')
if version_slug is None or version_slug not in self.TOKENS_COST:
@@ -105,8 +107,9 @@ class Gemini_3_1(StreamSimpleService):
**input_message.info,
}
messages, embedding_tokens = self._prepare_messages(input_message)
- content_parts: list[str] = []
input_tokens = output_tokens = 0
+ reasoning = ''
+ content = ''
result = ''
try:
@@ -115,13 +118,16 @@ class Gemini_3_1(StreamSimpleService):
while True:
chunk = next(stream)
if chunk:
- content_parts.append(chunk)
+ if chunk.event == 'think':
+ reasoning += chunk.data['content']
+ else:
+ content += chunk.data['content']
yield chunk
except StopIteration as exc:
input_tokens, output_tokens, _ = exc.value
finally:
- if content_parts:
- result = ''.join(content_parts)
+ if reasoning or content:
+ result = MessageService.prepare_output_message(reasoning, content)
process_time = timedelta(seconds=(time.time() - start_time))
self.handle_invoice(
input_message.content_object.model,
@@ -5,12 +5,14 @@ from decimal import Decimal
from typing import Any, Iterator
from messages.models import Message
+from messages.services.message_service import MessageService
from ml_model.adapters.bytedance_model_ark import BytedanceContentType, BytedanceModelArkAdapter
from ml_model.exceptions import GenerationException
from ml_model.services.base import SimpleService
from ml_model.tasks import bytedance_model_ark_run, stream_bytedance_model_ark_run
from payments.exceptions.insufficient_balance import InsufficientBalance
from payments.selectors.payment_plan_selector import PaymentPlanSelector
+from tools.chats.domain import RawSSEChunk
from tools.chats.models import Chat
from tools.copywrite.models import Copywrite
from tools.public_api.models import APIStore
@@ -107,11 +109,12 @@ class Glm_4_7(SimpleService):
msgs = self.save_results(result[0], process_time, save)
return msgs
- def make_stream(self, input_message: Message, save: bool = True) -> Iterator[str]:
+ def make_stream(self, input_message: Message, save: bool = True) -> Iterator[RawSSEChunk]:
version, callback_data, messages = self._prepare_data(input_message)
start_time = time.time()
- content_parts: list[str] = []
input_tokens = output_tokens = 0
+ reasoning = ''
+ content = ''
result = ''
try:
@@ -125,15 +128,18 @@ class Glm_4_7(SimpleService):
while True:
chunk = next(stream)
if chunk:
- content_parts.append(chunk)
+ if chunk.event == 'think':
+ reasoning += chunk.data['content']
+ else:
+ content += chunk.data['content']
yield chunk
except StopIteration as exc:
usage = exc.value or {}
input_tokens = int(usage.get('prompt_tokens') or 0)
output_tokens = int(usage.get('completion_tokens') or 0)
finally:
- if content_parts:
- result = ''.join(content_parts)
+ if reasoning or content:
+ result = MessageService.prepare_output_message(reasoning, content)
if not (input_tokens + output_tokens):
input_tokens = BytedanceModelArkAdapter.tokenize(
version, ''.join([m['content'] for m in messages])
@@ -11,6 +11,7 @@ from PIL import Image
from django.utils.translation import gettext
from messages.models import Message
+from messages.services.message_service import MessageService
from ml_model.adapters.openrouter import OpenrouterAdapter
from ml_model.exceptions import CorruptedFileError, FileExtensionNotSupported, PaidPlanRequiredError
from ml_model.services.base import StreamSimpleService
@@ -20,6 +21,7 @@ from ml_model.services.serper_mixin import SerperMixin
from payments.exceptions.insufficient_balance import InsufficientBalance
from payments.selectors.payment_plan_selector import PaymentPlanSelector
from poller.models import Proxy
+from tools.chats.domain import RawSSEChunk
from tools.chats.models import Chat
from tools.copywrite.models import Copywrite
from tools.public_api.models import APIStore
@@ -81,14 +83,15 @@ class Grok(SerperMixin, StreamSimpleService):
)
return self.save_results(result.content, process_time, save)
- def make_stream(self, input_message: Message, save: bool = True) -> Iterator[str]:
+ def make_stream(self, input_message: Message, save: bool = True) -> Iterator[RawSSEChunk]:
start_time = time.time()
version = input_message.info.get('version') or 'grok-4.5'
callback_data = {**input_message.info, 'tools': []}
messages, embedding_tokens = self._prepare_messages(input_message, version, callback_data)
- content_parts: list[str] = []
input_tokens = output_tokens = 0
cost = 0
+ reasoning = ''
+ content = ''
result = ''
try:
@@ -97,13 +100,16 @@ class Grok(SerperMixin, StreamSimpleService):
while True:
chunk = next(stream)
if chunk:
- content_parts.append(chunk)
+ if chunk.event == 'think':
+ reasoning += chunk.data['content']
+ else:
+ content += chunk.data['content']
yield chunk
except StopIteration as exc:
input_tokens, output_tokens, cost = exc.value
finally:
- if content_parts:
- result = ''.join(content_parts)
+ if reasoning or content:
+ result = MessageService.prepare_output_message(reasoning, content)
process_time = timedelta(seconds=(time.time() - start_time))
self.handle_invoice(
input_message.content_object.model,
@@ -1,4 +1,3 @@
-import json
import time
from datetime import timedelta
from decimal import Decimal
@@ -6,10 +5,9 @@ from pathlib import Path
from typing import Iterator
import filetype
-import httpx
-from backend import settings
from messages.models import Message
+from messages.services.message_service import MessageService
from ml_model.adapters.openrouter import OpenrouterAdapter
from ml_model.exceptions import CorruptedFileError, FileExtensionNotSupported
from ml_model.services.EmbeddingService import EmbeddingService
@@ -18,6 +16,7 @@ from ml_model.exceptions import ModelVersionNotAvailable
from ml_model.services.base import StreamSimpleService
from ml_model.tasks import openrouter_run
from poller.models import Proxy
+from tools.chats.domain import RawSSEChunk
from tools.chats.models import Chat
from tools.copywrite.models import Copywrite
from tools.public_api.models import APIStore
@@ -92,32 +91,34 @@ class Qwen_3_7(StreamSimpleService):
msgs = self.save_results(result[0], process_time)
return msgs
- def make_stream(self, input_message: Message, save: bool = True) -> Iterator[str]:
+ def make_stream(self, input_message: Message, save: bool = True) -> Iterator[RawSSEChunk]:
start_time = time.time()
version_slug, model_slug, callback_data, messages, embedding_tokens = self._prepare_data(
input_message
)
- content_parts: list[str] = []
+ input_tokens = output_tokens = 0
cost = 0
+ reasoning = ''
+ content = ''
result = ''
try:
- stream = self._run_streaming_api(model_slug, messages, callback_data)
+ stream = OpenrouterAdapter.run_streaming_api(model_slug, messages, callback_data, 'Qwen')
try:
while True:
chunk = next(stream)
if chunk:
- content_parts.append(chunk)
+ if chunk.event == 'think':
+ reasoning += chunk.data['content']
+ else:
+ content += chunk.data['content']
yield chunk
except StopIteration as exc:
- cost = exc.value
+ input_tokens, output_tokens, cost = exc.value
finally:
- if content_parts:
- result = ''.join(content_parts)
+ if reasoning or content:
+ result = MessageService.prepare_output_message(reasoning, content)
if not cost:
- input_tokens, output_tokens = OpenrouterAdapter.count_tokens_fallback(
- 'Qwen', messages, result
- )
cost = self._estimate_cost(version_slug, input_tokens, output_tokens)
process_time = timedelta(seconds=(time.time() - start_time))
self.handle_invoice(
@@ -240,46 +241,3 @@ class Qwen_3_7(StreamSimpleService):
+ output_tokens * price_map['output'] / 1_000_000
)
return float(price / self.COEFFICIENT)
-
- def _run_streaming_api(
- self, version: str, messages: list, callback_data: dict
- ) -> Iterator[str]:
- for proxy in Proxy.objects.all():
- with httpx.Client(
- base_url='https://openrouter.ai/api/v1',
- headers={'Authorization': f'Bearer {settings.OPENROUTER_API_KEY}'},
- proxy=f'{proxy.protocol}://{proxy.address}',
- timeout=600,
- ) as client:
- with client.stream(
- 'POST',
- 'chat/completions',
- json={
- 'model': version,
- 'stream': True,
- 'messages': messages,
- 'transforms': ['middle-out'],
- **callback_data,
- },
- ) as resp:
- cost = 0
- for line in resp.iter_lines():
- line = line.strip()
- if not line or not line.startswith('data: '):
- continue
-
- data = line[6:]
- if data == '[DONE]':
- break
-
- try:
- data_obj = json.loads(data)
- chunk = data_obj['choices'][0]['delta'].get('content') or ''
- if chunk:
- yield chunk
- if data_obj.get('usage'):
- cost = data_obj['usage'].get('cost') or 0
- except json.JSONDecodeError:
- continue
-
- return cost
@@ -1,25 +1,32 @@
-import base64
import time
from datetime import timedelta
from decimal import Decimal
from io import BytesIO
from typing import Any
-import filetype
import requests
from django.core.files import File
from messages.models import Message
from ml_model.adapters.bytedance_model_ark import BytedanceContentType
+from ml_model.exceptions import InvalidParameterError
from ml_model.services.base import SimpleService
-from ml_model.tasks import bytedance_model_ark_run, replicate_run
+from ml_model.tasks import bytedance_model_ark_run
+
+# import base64
+# import filetype
+# from ml_model.tasks import replicate_run
class Reve(SimpleService):
# Reve временно не работает на репликейте. Временно используем сидрим
TEMPORARY_PROVIDER_MODEL = 'seedream-5-0-260128'
- PRICE = Decimal('25')
+ PRICE = {
+ '2K': Decimal('25'),
+ '3K': Decimal('50'),
+ '4K': Decimal('100'),
+ }
# PRICE = {
# 'create': Decimal('12.5'),
# 'edit-fast': Decimal('5'),
@@ -27,10 +34,12 @@ class Reve(SimpleService):
@classmethod
def predict_price(cls, content: str, file_exists: bool, info: dict[str, Any]) -> Decimal | None:
- return cls.PRICE.quantize(Decimal('0.1'), rounding='ROUND_UP')
+ price = cls.PRICE.get(info.get('size', '2K'))
+
+ return price.quantize(Decimal('0.1'), rounding='ROUND_UP')
- def calculate_price(self) -> Decimal:
- return self.PRICE.quantize(Decimal('0.1'), rounding='ROUND_UP')
+ def calculate_price(self, size: str) -> Decimal:
+ return self.PRICE[size].quantize(Decimal('0.1'), rounding='ROUND_UP')
# @classmethod
# def predict_price(cls, content: str, file_exists: bool, info: dict[str, Any]) -> Decimal | None:
@@ -62,10 +71,14 @@ class Reve(SimpleService):
def make(self, input_message: Message, save: bool = True) -> list[Message]:
start_time = time.time()
+ size = input_message.info.get('size', '2K')
+ if size not in self.PRICE:
+ raise InvalidParameterError(f'Unsupported size: {size}')
+
callback_data = {
'prompt': self.translate_prompt(input_message.content),
**input_message.info,
- 'size': '2K',
+ 'size': size,
'watermark': False,
}
if image := input_message.file:
@@ -76,7 +89,7 @@ class Reve(SimpleService):
content_type=BytedanceContentType.IMAGE,
)[0]
process_time = timedelta(seconds=(time.time() - start_time))
- self.handle_invoice(input_message.content_object.model)
+ self.handle_invoice(input_message.content_object.model, size=size)
msgs = self.save_results(input_message.content, image, process_time, save)
return msgs
@@ -9,10 +9,14 @@ from decimal import Decimal
from django.core.cache import cache
from messages.models import Message
+from messages.services.message_service import MessageService
+from ml_model.exceptions import GenerationException
from ml_model.services.base import StreamSimpleService
from typing import Iterator
+from tools.chats.domain import RawSSEChunk
+
type TestTextTokens = list[str]
@@ -20,6 +24,7 @@ type TestTextTokens = list[str]
class PrepareTestData:
input_tokens: TestTextTokens
output_tokens: TestTextTokens
+ reasoning_tokens: TestTextTokens
ttft: float
tbt: float
start_time: float
@@ -39,6 +44,13 @@ class Text_Test_Model(StreamSimpleService):
fermentum lorem sit amet tortor ultricies, id pulvinar nibh pulvinar.
"""
+ BASE_REASONING_MESSAGE = """
+ In a dapibus nulla. Aenean erat orci, egestas non orci at, varius tempus risus. Ut suscipit lorem magna,
+ quis auctor leo molestie ac. Integer ut efficitur neque. Curabitur sollicitudin ipsum dolor, et tempus massa
+ lacinia a. Donec efficitur egestas facilisis. Aliquam feugiat convallis arcu quis sollicitudin.
+ Nullam eleifend iaculis sapien id scelerisque.
+ """
+
def calculate_price(self, input_tokens: int, output_tokens: int) -> Decimal:
price = (
input_tokens * self.TOKENS_COST['input'] / 1_000_000
@@ -58,23 +70,49 @@ class Text_Test_Model(StreamSimpleService):
def make(self, input_message: Message, save: bool = True) -> list[Message]:
prepare = self._prepare(input_message)
- result = ''.join(self._stream(prepare))
+ reasoning = ''
+ output = ''
+ for token in self._stream(prepare):
+ if token.event == 'think':
+ reasoning += token.data['content']
+ else:
+ output += token.data['content']
+ result = MessageService.prepare_output_message(reasoning, output)
+ if not result:
+ raise GenerationException
return self._finalize(
- input_message, prepare, result, save, output_token_count=len(prepare.output_tokens)
+ input_message,
+ prepare,
+ result,
+ save,
+ output_token_count=len(prepare.output_tokens) + len(prepare.reasoning_tokens),
)
- def make_stream(self, input_message: Message, save: bool = True) -> Iterator[str]:
+ def make_stream(self, input_message: Message, save: bool = True) -> Iterator[RawSSEChunk]:
prepare = self._prepare(input_message)
- result = ''
+ output = ''
+ reasoning = ''
output_token_count = 0
try:
for token in self._stream(prepare):
- result += token
+ if token.event == 'think':
+ reasoning += token.data['content']
+ else:
+ output += token.data['content']
output_token_count += 1
yield token
finally:
+ result = MessageService.prepare_output_message(reasoning, output)
if result:
- self._finalize(input_message, prepare, result, save, output_token_count=output_token_count)
+ self._finalize(
+ input_message,
+ prepare,
+ result,
+ save,
+ output_token_count=output_token_count,
+ )
+ if not result:
+ raise GenerationException
return result
def _prepare(self, input_message: Message) -> PrepareTestData:
@@ -82,16 +120,21 @@ class Text_Test_Model(StreamSimpleService):
info = input_message.info.copy()
input_tokens = self._get_cached_tokens(input_message.content)
output_tokens = self._get_cached_tokens(info.get('cm') or self.BASE_OUTPUT_MESSAGE)
+ reasoning_tokens = []
+ if info.get('reasoning'):
+ reasoning_tokens = self._get_cached_tokens(info.get('rm') or self.BASE_REASONING_MESSAGE)
ttft = info.get('ttft', 0.5)
tbt = info.get('tbt', 0.35)
- return PrepareTestData(input_tokens, output_tokens, ttft, tbt, start_time)
+ return PrepareTestData(input_tokens, output_tokens, reasoning_tokens, ttft, tbt, start_time)
- def _stream(self, prepare: PrepareTestData) -> Iterator[str]:
+ def _stream(self, prepare: PrepareTestData) -> Iterator[RawSSEChunk]:
time.sleep(prepare.ttft)
- for i, token in enumerate(prepare.output_tokens, start=1):
- yield token
- if i < len(prepare.output_tokens):
- time.sleep(prepare.tbt)
+ streaming_data = {'think': prepare.reasoning_tokens, 'token': prepare.output_tokens}
+ for k, v in streaming_data.items():
+ for i, token in enumerate(v, start=1):
+ yield RawSSEChunk(event=k, data={'content': token})
+ if i < len(v):
+ time.sleep(prepare.tbt)
def _finalize(
self,
@@ -20,6 +20,7 @@ from replicate.exceptions import ModelError
from requests import Response
from backend import settings
+from messages.services.message_service import MessageService
from ml_model.adapters.bytedance_model_ark import BytedanceContentType, BytedanceModelArkAdapter
from ml_model.adapters.openrouter import OpenrouterAdapter
from ml_model.exceptions import (
@@ -167,7 +168,12 @@ def openrouter_run(version: str, messages: list, callback_data: dict, model_name
for c in data.get('choices', [])
if c.get('message') and c['message'].get('content') is not None
]
- if not raw_content:
+ raw_reasoning = ','.join(
+ reasoning
+ for choice in data.get('choices', [])
+ if (reasoning := choice['message'].get('reasoning')) is not None
+ )
+ if not raw_content and not raw_reasoning:
if any(
c.get('native_finish_reason') == 'SAFETY_CHECK_TYPE_CSAM'
for c in (data.get('choices') or [])
@@ -182,22 +188,8 @@ def openrouter_run(version: str, messages: list, callback_data: dict, model_name
)
raise GenerationException
content = ','.join(raw_content)
- reasoning = ','.join(
- reasoning
- for choice in data.get('choices', [])
- if (reasoning := choice['message'].get('reasoning')) is not None
- )
- reasoning = re.sub(r'Вывод:|Основная мысль:|Рассуждение:|\*\*', '', reasoning)
- answer = reasoning
- if any(m in data['model'] for m in ('google/gemini', 'x-ai/grok-4.3')) or re.match(
- r'^qwen/qwen3\.(?:5|6|7)-.*$', data['model']
- ):
- answer = content
- elif reasoning and content:
- # TODO: переделать рендеринг сообщения на Jinja 2
- answer = f'**Рассуждение:**\n\n{reasoning}\n\n**Основная мысль:**\n\n{content}'
- elif content:
- answer = content
+ reasoning = re.sub(r'Вывод:|Основная мысль:|Рассуждение:|\*\*', '', raw_reasoning)
+ answer = MessageService.prepare_output_message(reasoning, content)
if int(data.get('choices')[0].get('error', {}).get('code', 0)) == 502:
error_type = re.sub(r'["\']', '', str(data['choices'][0]['error']['message']))
if error_type == 'Overloaded':
@@ -1,7 +1,6 @@
from datetime import datetime
from django.contrib.auth import get_user_model
-from django.contrib.postgres.fields import ArrayField
from django.db import models
from django.utils.translation import gettext_lazy as _
@@ -76,7 +75,12 @@ class PaymentPlanUserInfo(BaseModel):
update_fields=None,
):
self.last_payment_at = datetime.now().date()
- return super().save(force_insert, force_update, using, update_fields)
+ return super().save(
+ force_insert=force_insert,
+ force_update=force_update,
+ using=using,
+ update_fields=update_fields,
+ )
@property
def is_recurring(self) -> bool:
@@ -100,6 +100,8 @@ class PaymentPlanUserInfoAdmin(admin.ModelAdmin):
class PaymentPlanFeatureAdmin(OrderedModelAdmin):
list_display = ('plan', 'model', 'move_up_down_links')
list_filter = ('plan', 'model__category')
+ search_fields = ('model__title', 'model__slug')
+ search_help_text = _('Search by model title or slug')
@admin.register(PaymentMethod)
@@ -250,4 +252,3 @@ class ReferralAccountAdmin(admin.ModelAdmin):
@admin.display(description='Получено бонусов')
def _accrued_bonuses(self, obj: ReferralAccount):
return f'{obj.accrued_bonuses.aggregate(total=Coalesce(Sum("amount"), Decimal(0), output_field=models.DecimalField()))["total"]} токенов'
-
@@ -0,0 +1,235 @@
+import base64
+from dataclasses import dataclass, field
+from decimal import Decimal
+from enum import StrEnum
+
+from django.core.files.uploadedfile import SimpleUploadedFile
+
+
+class ResultKind(StrEnum):
+ TEXT = 'text'
+ FILE = 'file'
+
+
+@dataclass(frozen=True)
+class InputFileFixture:
+ name: str
+ content: bytes
+ content_type: str
+
+
+@dataclass(frozen=True)
+class ModelCase:
+ slug: str
+ result_kind: ResultKind
+ file_suffix: str = ''
+ input_file: InputFileFixture | None = None
+ info: dict[str, object] = field(default_factory=dict)
+
+ def make_input_file(self) -> SimpleUploadedFile | None:
+ if not self.input_file:
+ return None
+
+ return SimpleUploadedFile(
+ self.input_file.name,
+ self.input_file.content,
+ content_type=self.input_file.content_type,
+ )
+
+
+TEXT_MODEL_SLUGS = (
+ 'chatgpt_4',
+ 'chatgpt',
+ 'chatgpt_5',
+ 'chatgpt_5_4',
+ 'claude',
+ 'codellama',
+ 'deepl',
+ 'deepseek',
+ 'dola_seed',
+ 'gemini',
+ 'gemini_3_1',
+ 'gemma',
+ 'glm_4_7',
+ 'granite',
+ 'grok',
+ 'grok_4_1_fast',
+ 'llama',
+ 'mistral',
+ 'perplexity',
+ 'qwen',
+ 'qwen_235B',
+ 'qwen_3_6',
+ 'qwen_3_7',
+ 'qwen_3_max_thinking',
+ 'raifgpt',
+ 'vicuna',
+ 'whisper',
+)
+
+FILE_MODEL_SUFFIXES = {
+ 'dalle': '.png',
+ 'djourney': '.png',
+ 'elevenlabs': '.mp3',
+ 'elevenlabs_music': '.mp3',
+ 'epicphotogasm': '.png',
+ 'flux': '.png',
+ 'flux_2': '.png',
+ 'fluxkrea': '.png',
+ 'fluxlorafast': '.png',
+ 'fluxproultra': '.png',
+ 'fluxpulid': '.png',
+ 'geminiimage': '.png',
+ 'gptimage': '.png',
+ 'grok_image': '.png',
+ 'grok_imagine_video': '.mp4',
+ 'hailuo': '.mp4',
+ 'hunyuan': '.mp4',
+ 'iconic': '.png',
+ 'ideogram': '.png',
+ 'imagen': '.png',
+ 'kandinsky': '.png',
+ 'kling': '.mp4',
+ 'leonardo': '.png',
+ 'lightning': '.png',
+ 'logoai': '.png',
+ 'ltx': '.mp4',
+ 'lyria': '.mp3',
+ 'midjourney': '.png',
+ 'minimaxmusic': '.mp3',
+ 'minimaxmusic_lite': '.mp3',
+ 'minimaxvideo': '.mp4',
+ 'musicgen': '.mp3',
+ 'nanobanana': '.png',
+ 'nanobanana_2': '.png',
+ 'photon': '.png',
+ 'pixverse': '.mp4',
+ 'pruna_v': '.mp4',
+ 'prunaai': '.mp4',
+ 'pulid': '.png',
+ 'ray': '.mp4',
+ 'recraft': '.png',
+ 'reve': '.png',
+ 'runway': '.mp4',
+ 'sdxlemoji': '.png',
+ 'seedance': '.mp4',
+ 'seedance_2_dreamina': '.mp4',
+ 'seedream': '.png',
+ 'sora': '.mp4',
+ 'stablediffusion': '.png',
+ 'stablemusic': '.mp3',
+ 'suno': '.mp3',
+ 'upscaleai': '.png',
+ 'veo': '.mp4',
+ 'wan': '.mp4',
+ 'wan_lite': '.mp4',
+}
+
+EXCLUDED_FAKE_MODEL_SLUGS = {
+ 'audio_test_model',
+ 'image_test_model',
+ 'text_test_model',
+ 'video_test_model',
+}
+
+PNG_INPUT_FILE = InputFileFixture(
+ name='input.png',
+ content=base64.b64decode(
+ 'iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNk+A8AAQUBAScY42YAAAAASUVORK5CYII='
+ ),
+ content_type='image/png',
+)
+MP3_INPUT_FILE = InputFileFixture(
+ name='input.mp3',
+ content=b'mocked audio input',
+ content_type='audio/mpeg',
+)
+REQUIRED_INPUT_FILES = {
+ 'fluxpulid': PNG_INPUT_FILE,
+ 'kling': PNG_INPUT_FILE,
+ 'runway': PNG_INPUT_FILE,
+ 'upscaleai': PNG_INPUT_FILE,
+ 'wan': PNG_INPUT_FILE,
+ 'whisper': MP3_INPUT_FILE,
+}
+
+CASE_INFO = {
+ 'chatgpt_4': {'version': 'o3-mini'},
+ 'chatgpt': {'version': 'gpt-5.5'},
+ 'chatgpt_5': {'version': 'gpt-5'},
+ 'chatgpt_5_4': {'version': 'gpt-5.4'},
+ 'claude': {'version': 'claude-sonnet-4.6'},
+ 'deepl': {'source_lang': 'ru', 'target_lang': 'en'},
+ 'deepseek': {'version': 'deepseek/deepseek-v4-pro'},
+ 'dola_seed': {'version': 'seed-2-0-pro'},
+ 'gemini': {'version': 'gemini-2.0-flash-001'},
+ 'gemini_3_1': {'version': 'gemini-3.1-pro-preview'},
+ 'grok': {'version': 'grok-4.3'},
+ 'llama': {'version': 'llama-3.3-70b-instruct'},
+ 'perplexity': {'version': 'sonar'},
+ 'qwen': {'version': 'qwq-32b'},
+ 'qwen_235B': {'version': 'qwen3-235b-a22b-thinking-2507'},
+ 'qwen_3_6': {'version': 'qwen3.6-flash'},
+ 'qwen_3_7': {'version': 'qwen3.7-max'},
+ 'elevenlabs_music': {'duration': 5},
+ 'flux_2': {'version': 'flux-2-pro'},
+ 'fluxproultra': {'version': 'flux-dev'},
+ 'hailuo': {'version': 'hailuo-2.3', 'resolution': '768p'},
+ 'hunyuan': {'version': 'hunyuan-video'},
+ 'kling': {'mode': 'standard', 'duration': 5},
+ 'leonardo': {'version': 'lucid-origin'},
+ 'ltx': {'resolution': '1080p', 'duration': 6},
+ 'lyria': {'version': 'lyria-3'},
+ 'minimaxmusic': {'style': 'ambient'},
+ 'nanobanana': {'version': 'nano-banana'},
+ 'nanobanana_2': {'resolution': '2K'},
+ 'pixverse': {'quality': '1080p', 'duration': 5},
+ 'pruna_v': {'resolution': '720p', 'generation_mode': 'standard', 'duration': 5},
+ 'prunaai': {'version': 'p-image'},
+ 'ray': {'version': 'ray-2-720p', 'duration': 5},
+ 'recraft': {'version': 'recraft-v3', 'style': 'любой'},
+ 'runway': {'duration': 5},
+ 'seedance': {'version': 'seedance-2.0-fast', 'resolution': '720p', 'duration': 5},
+ 'seedance_2_dreamina': {
+ 'version': 'dreamina-seedance-2-0',
+ 'resolution': '720p',
+ 'duration': 5,
+ },
+ 'seedream': {'version': 'seedream-boosted', 'size': '2K'},
+ 'sora': {'version': 'sora-2', 'seconds': 4},
+ 'stablediffusion': {'version': 'sd3'},
+ 'suno': {'style': 'ambient'},
+ 'veo': {'version': 'veo-3'},
+ 'wan': {'resolution': '720p', 'duration': 5},
+ 'wan_lite': {'resolution': '720p'},
+}
+
+def get_model_case(slug: str) -> ModelCase:
+ if slug in TEXT_MODEL_SLUGS:
+ result_kind = ResultKind.TEXT
+ file_suffix = ''
+ elif slug in FILE_MODEL_SUFFIXES:
+ result_kind = ResultKind.FILE
+ file_suffix = FILE_MODEL_SUFFIXES[slug]
+ else:
+ available_slugs = ', '.join((*TEXT_MODEL_SLUGS, *FILE_MODEL_SUFFIXES))
+
+ raise ValueError(f'Unknown ML model slug: {slug}. Available slugs: {available_slugs}')
+
+ return ModelCase(
+ slug=slug,
+ result_kind=result_kind,
+ file_suffix=file_suffix,
+ input_file=REQUIRED_INPUT_FILES.get(slug),
+ info=CASE_INFO.get(slug, {}),
+ )
+
+
+def get_model_cases() -> tuple[ModelCase, ...]:
+ slugs = (*TEXT_MODEL_SLUGS, *FILE_MODEL_SUFFIXES)
+
+ return tuple(get_model_case(slug) for slug in slugs)
+
+MOCK_OUTPUT = 'Mocked provider response'
+MOCK_MEDIA_BYTES = b'mocked media bytes'
+STARTING_BALANCE = Decimal('1000000000')
@@ -0,0 +1,322 @@
+import base64
+import importlib
+import inspect
+from dataclasses import dataclass
+from decimal import Decimal
+from io import BytesIO
+from types import ModuleType
+from types import SimpleNamespace
+from unittest.mock import Mock
+
+from langchain_core.messages import AIMessage
+
+from ml_model.adapters.models import ModelResponse
+from ml_model.services.base import SimpleService
+from tests.ml_models.cases import MOCK_MEDIA_BYTES, MOCK_OUTPUT, ModelCase, ResultKind
+
+
+class SmartMediaURL(str):
+ def __new__(cls):
+ return super().__new__(cls, 'https://provider.test/generated-file')
+
+ def __iter__(self):
+ yield self
+
+ @property
+ def url(self):
+ return self
+
+ def __getitem__(self, item):
+ if item == 0:
+ return self
+
+ return super().__getitem__(item)
+
+
+class SmartStatus:
+ SUCCESS_VALUES = {'COMPLETED', 'completed', 'succeeded'}
+
+ def __eq__(self, other) -> bool:
+ return other in self.SUCCESS_VALUES
+
+ def __hash__(self) -> int:
+ return hash('completed')
+
+
+class SmartResponse:
+ status_code = 200
+ content = MOCK_MEDIA_BYTES
+ text = MOCK_OUTPUT
+
+ def json(self) -> dict:
+ encoded_media = base64.b64encode(MOCK_MEDIA_BYTES).decode()
+
+ return {
+ 'id': 'mock-generation',
+ 'status': SmartStatus(),
+ 'logs': '',
+ 'input_tokens': 100,
+ 'output': [SmartMediaURL()],
+ 'response_url': SmartMediaURL(),
+ 'status_url': SmartMediaURL(),
+ 'urls': {'get': SmartMediaURL()},
+ 'images': [{'url': SmartMediaURL()}],
+ 'data': [{'b64_json': encoded_media}],
+ 'artifacts': [{'base64': encoded_media}],
+ 'choices': [{'message': {'content': MOCK_OUTPUT, 'annotations': []}}],
+ 'usage': {
+ 'prompt_tokens': 100,
+ 'completion_tokens': 40,
+ 'input_tokens': 100,
+ 'output_tokens': 40,
+ 'total_tokens': 140,
+ 'input_tokens_details': {'image_tokens': 0, 'text_tokens': 100},
+ },
+ }
+
+ def raise_for_status(self) -> None:
+ return None
+
+
+class SmartClient:
+ def __init__(self, *args, **kwargs) -> None:
+ self.response = SmartResponse()
+
+ def __enter__(self):
+ return self
+
+ def __exit__(self, exc_type, exc_value, traceback) -> None:
+ return None
+
+ def post(self, *args, **kwargs) -> SmartResponse:
+ return self.response
+
+ def get(self, *args, **kwargs) -> SmartResponse:
+ return self.response
+
+ def close(self) -> None:
+ return None
+
+
+class FakeLLM:
+ model_name = 'gpt-4o'
+
+ def __init__(self, *args, **kwargs) -> None:
+ self.model_name = kwargs.get('model', self.model_name)
+
+
+class FakeConversation:
+ def __init__(self, *args, **kwargs) -> None:
+ pass
+
+ def invoke(self, *args, **kwargs) -> AIMessage:
+ return AIMessage(content=MOCK_OUTPUT)
+
+
+class FakeAudioSegment:
+ @classmethod
+ def empty(cls):
+ return cls()
+
+ @classmethod
+ def from_file(cls, *args, **kwargs):
+ return cls()
+
+ def __iadd__(self, other):
+ return self
+
+ def export(self, output: BytesIO, format: str) -> None:
+ output.write(MOCK_MEDIA_BYTES)
+
+
+@dataclass
+class ProviderMockTracker:
+ call_count: int = 0
+
+ def record(self) -> None:
+ self.call_count += 1
+
+
+class SuccessfulTask:
+ def __init__(self, result) -> None:
+ self.result = result
+
+ def get(self):
+ return self.result
+
+ def successful(self) -> bool:
+ return True
+
+
+def install_provider_fakes(monkeypatch, case: ModelCase, model) -> ProviderMockTracker:
+ tracker = ProviderMockTracker()
+ media_url = SmartMediaURL()
+
+ monkeypatch.setattr(SimpleService, 'neuron_model', property(lambda service: model))
+ monkeypatch.setattr(SimpleService, 'translate_prompt', lambda service, prompt, to='en': prompt)
+ service_class = model.service
+ if not hasattr(service_class, 'price'):
+ monkeypatch.setattr(service_class, 'price', Decimal('1'), raising=False)
+
+ def provider_result(*args, **kwargs):
+ tracker.record()
+ if case.result_kind == ResultKind.TEXT:
+ return MOCK_OUTPUT
+
+ return media_url
+
+ def token_provider_result(*args, **kwargs):
+ tracker.record()
+
+ return MOCK_OUTPUT, 100, 40
+
+ def bytedance_result(*args, **kwargs):
+ tracker.record()
+ content_type = str(kwargs.get('content_type', '')).lower()
+ if 'chat' in content_type:
+ return MOCK_OUTPUT, 100, 40
+ if 'video' in content_type:
+ return media_url, 40
+
+ return media_url
+
+ def streaming_result(*args, **kwargs) -> ModelResponse:
+ tracker.record()
+
+ return ModelResponse(MOCK_OUTPUT, 100, 40, 0.001)
+
+ service_modules = _service_modules(service_class)
+
+ for module in service_modules:
+ for name, replacement in (
+ ('replicate_run', provider_result),
+ ('upscale_run', lambda *args, **kwargs: [media_url]),
+ ('openrouter_run', token_provider_result),
+ ('bytedance_model_ark_run', bytedance_result),
+ ):
+ if hasattr(module, name):
+ monkeypatch.setattr(module, name, replacement)
+
+ if hasattr(module, 'requests'):
+ monkeypatch.setattr(module.requests, 'get', Mock(return_value=SmartResponse()))
+ monkeypatch.setattr(module.requests, 'post', Mock(return_value=SmartResponse()))
+ if hasattr(module, 'httpx'):
+ monkeypatch.setattr(module.httpx, 'Client', SmartClient)
+ monkeypatch.setattr(module.httpx, 'get', Mock(return_value=SmartResponse()))
+ monkeypatch.setattr(module.httpx, 'post', Mock(return_value=SmartResponse()))
+ if hasattr(module, 'client'):
+ monkeypatch.setattr(module, 'client', SmartClient())
+
+ _install_adapter_fakes(monkeypatch, service_class, service_modules, streaming_result)
+ _install_task_fakes(monkeypatch, case.slug)
+ _install_model_specific_fakes(monkeypatch, case.slug, tracker, media_url)
+
+ return tracker
+
+
+def _service_modules(service_class) -> tuple[ModuleType, ...]:
+ modules = []
+ for base_class in service_class.__mro__:
+ module = inspect.getmodule(base_class)
+ if module and module.__name__.startswith('ml_model.services.') and module not in modules:
+ modules.append(module)
+
+ return tuple(modules)
+
+
+def _install_adapter_fakes(monkeypatch, service_class, service_modules, streaming_result) -> None:
+ if any(hasattr(module, 'OpenrouterAdapter') for module in service_modules):
+ from ml_model.adapters.openrouter import OpenrouterAdapter
+
+ monkeypatch.setattr(OpenrouterAdapter, 'collect_streaming_api', streaming_result)
+
+ if any(hasattr(module, 'BytedanceModelArkAdapter') for module in service_modules):
+ from ml_model.adapters.bytedance_model_ark import BytedanceModelArkAdapter
+
+ monkeypatch.setattr(BytedanceModelArkAdapter, 'batch_tokenize', lambda model, texts: 100)
+
+ from ml_model.services.chatgpt import Chatgpt
+ from ml_model.services.chatgpt_4 import Chatgpt_4
+
+ if issubclass(service_class, Chatgpt_4):
+ monkeypatch.setattr(
+ Chatgpt_4,
+ 'call_openai_api',
+ lambda service, proxy, endpoint, json_data: (100, 40, AIMessage(content=MOCK_OUTPUT)),
+ )
+ monkeypatch.setattr(Chatgpt_4, 'count_text_tokens', lambda service, messages: 100)
+
+ if issubclass(service_class, Chatgpt):
+ monkeypatch.setattr(Chatgpt, '_count_responses_input_tokens', lambda service, proxy, payload: 100)
+ monkeypatch.setattr(
+ Chatgpt,
+ '_stream_openai_responses',
+ lambda service, proxy, json_data, model_name: (100, 40, AIMessage(content=MOCK_OUTPUT)),
+ )
+
+
+def _install_task_fakes(monkeypatch, slug: str) -> None:
+ if slug == 'deepl':
+ module = importlib.import_module('ml_model.services.deepl')
+ monkeypatch.setattr(
+ module.translate,
+ 'delay',
+ lambda callback_data: SuccessfulTask(MOCK_OUTPUT),
+ )
+
+ if slug == 'whisper':
+ module = importlib.import_module('ml_model.services.whisper')
+ monkeypatch.setattr(
+ module.transcript_audio,
+ 'delay',
+ lambda file: SuccessfulTask({'text': MOCK_OUTPUT}),
+ )
+ audio_info = SimpleNamespace(info=SimpleNamespace(length=1))
+ monkeypatch.setattr(module, 'MP3', lambda audio: audio_info)
+ monkeypatch.setattr(module, 'WAVE', lambda audio: audio_info)
+
+
+def _install_model_specific_fakes(
+ monkeypatch,
+ slug: str,
+ tracker: ProviderMockTracker,
+ media_url: SmartMediaURL,
+) -> None:
+ if slug == 'minimaxmusic_lite':
+ module = importlib.import_module('ml_model.services.minimaxmusic_lite')
+ monkeypatch.setattr(
+ module.Preset.objects,
+ 'get',
+ lambda **kwargs: SimpleNamespace(file=SimpleNamespace(url=media_url)),
+ )
+
+ if slug == 'granite':
+ module = importlib.import_module('ml_model.services.granite')
+ monkeypatch.setattr(
+ module.Granite,
+ '_call_api',
+ lambda service, payload: (tracker.record() or {'output': [MOCK_OUTPUT]}),
+ )
+
+ if slug in {'chatgpt_4', 'raifgpt'}:
+ module = importlib.import_module(f'ml_model.services.{slug}')
+ monkeypatch.setattr(module, 'ChatOpenAI', FakeLLM)
+ if hasattr(module, 'RunnableWithMessageHistory'):
+ monkeypatch.setattr(module, 'RunnableWithMessageHistory', FakeConversation)
+
+ if slug == 'elevenlabs':
+ module = importlib.import_module('ml_model.services.elevenlabs')
+ monkeypatch.setattr(module, 'AudioSegment', FakeAudioSegment)
+ monkeypatch.setattr(
+ module.FileProcessingService,
+ 'get_voice_file',
+ classmethod(lambda cls, voice_id, preset_id, user: SimpleNamespace(url=media_url)),
+ )
+
+ if slug == 'gptimage':
+ module = importlib.import_module('ml_model.services.gptimage')
+ monkeypatch.setattr(
+ module.Gptimage,
+ 'count_predict_tokens',
+ classmethod(lambda cls, text, width, height, size, quality: (100, 0, 196)),
+ )
@@ -0,0 +1,88 @@
+import ast
+import json
+from pathlib import Path
+
+import pytest
+
+from messages.models import Message
+from payments.models import Invoice
+from tests.factories import ChatFactory, NeuronModelFactory
+from tests.ml_models.cases import (
+ EXCLUDED_FAKE_MODEL_SLUGS,
+ ResultKind,
+ get_model_cases,
+)
+from tests.ml_models.provider_fakes import install_provider_fakes
+
+
+MODEL_CASES = get_model_cases()
+MODEL_PARAMS = tuple(
+ pytest.param(case, id=case.slug, marks=pytest.mark.ml_model(case.slug))
+ for case in MODEL_CASES
+)
+
+
+@pytest.mark.django_db
+@pytest.mark.usefixtures('fake_provider_proxy')
+@pytest.mark.parametrize('case', MODEL_PARAMS)
+def test_model_api_contract(
+ case,
+ authenticated_client,
+ user,
+ local_message_storage,
+ monkeypatch,
+) -> None:
+ model = NeuronModelFactory(slug=case.slug)
+ chat = ChatFactory(user=user, model=model)
+ install_provider_fakes(monkeypatch, case, model)
+ request_data = {
+ 'content': 'Generic ML model API contract',
+ 'info': dict(case.info),
+ }
+ if input_file := case.make_input_file():
+ request_data['file'] = input_file
+ request_data['info'] = json.dumps(request_data['info'])
+
+ balance_before = user.payment_plan.current_token_balance
+ response = authenticated_client.post(
+ f'/api/v1/chats/{chat.uid}/messages/',
+ request_data,
+ format='multipart' if input_file else 'json',
+ )
+
+ assert response.status_code == 201, response.json()
+ payload = response.json()
+ assert len(payload) >= 2
+ assert payload[0]['from_model'] is False
+ assert all(message['from_model'] is True for message in payload[1:])
+
+ output_messages = Message.objects.filter(object_id=chat.uid, from_model=True)
+ assert output_messages.exists()
+ if case.result_kind == ResultKind.TEXT:
+ assert any(message.content for message in output_messages)
+ else:
+ assert all(message.file for message in output_messages)
+ for output_message in output_messages:
+ assert local_message_storage.exists(output_message.file.name)
+ with local_message_storage.open(output_message.file.name, 'rb') as saved_file:
+ assert saved_file.read()
+
+ invoice = Invoice.objects.get(user=user, model=model)
+ user.payment_plan.refresh_from_db()
+ charged_tokens = balance_before - user.payment_plan.current_token_balance
+ assert charged_tokens == invoice.cost
+
+
+def test_every_exported_model_has_api_case() -> None:
+ services_init = Path('ml_model/services/__init__.py')
+ module = ast.parse(services_init.read_text())
+ exported_slugs = {
+ node.module.rsplit('.', 1)[1]
+ for node in module.body
+ if isinstance(node, ast.ImportFrom)
+ and node.module
+ and node.module.startswith('ml_model.services.')
+ }
+ case_slugs = {case.slug for case in MODEL_CASES}
+
+ assert case_slugs == exported_slugs - EXCLUDED_FAKE_MODEL_SLUGS
@@ -0,0 +1,85 @@
+import json
+import os
+
+import pytest
+
+from messages.models import Message
+from payments.models import Invoice
+from poller.models import Proxy
+from tests.factories import ChatFactory, NeuronModelFactory
+from tests.ml_models.cases import ModelCase, ResultKind, get_model_case
+
+
+def _configure_provider_proxy() -> None:
+ proxy_address = os.getenv('PROVIDER_SMOKE_PROXY_ADDRESS', '').strip()
+ if not proxy_address:
+ return
+
+ proxy_protocol = os.getenv('PROVIDER_SMOKE_PROXY_PROTOCOL', 'http').strip().lower()
+ allowed_protocols = {choice[0] for choice in Proxy.ProtocolChoices.choices}
+ if proxy_protocol not in allowed_protocols:
+ raise pytest.UsageError(
+ 'PROVIDER_SMOKE_PROXY_PROTOCOL must be http, https, or socks.'
+ )
+ if '://' in proxy_address:
+ raise pytest.UsageError('PROVIDER_SMOKE_PROXY_ADDRESS must not contain a protocol.')
+
+ Proxy.objects.create(address=proxy_address, protocol=proxy_protocol)
+
+
+def pytest_generate_tests(metafunc: pytest.Metafunc) -> None:
+ selected_slugs = metafunc.config.getoption('--model-slug')
+ if len(selected_slugs) != 1:
+ raise pytest.UsageError('--provider-smoke requires exactly one --model-slug.')
+
+ try:
+ case = get_model_case(selected_slugs[0])
+ except ValueError as error:
+ raise pytest.UsageError(str(error)) from error
+
+ parameter = pytest.param(case, id=case.slug, marks=pytest.mark.ml_model(case.slug))
+ metafunc.parametrize('case', (parameter,))
+
+
+@pytest.mark.provider_smoke
+@pytest.mark.django_db
+def test_real_provider_api_flow(
+ case: ModelCase,
+ authenticated_client,
+ user,
+) -> None:
+ _configure_provider_proxy()
+ model = NeuronModelFactory(slug=case.slug)
+ chat = ChatFactory(user=user, model=model)
+ request_data = {
+ 'content': 'Short provider smoke test. Reply or generate a minimal result.',
+ 'info': dict(case.info),
+ }
+ if input_file := case.make_input_file():
+ request_data['file'] = input_file
+ request_data['info'] = json.dumps(request_data['info'])
+
+ balance_before = user.payment_plan.current_token_balance
+ response = authenticated_client.post(
+ f'/api/v1/chats/{chat.uid}/messages/',
+ request_data,
+ format='multipart' if input_file else 'json',
+ )
+
+ response_payload = response.json()
+
+ assert response.status_code == 201, response_payload
+ output_messages = Message.objects.filter(object_id=chat.uid, from_model=True)
+ assert output_messages.exists()
+ if case.result_kind == ResultKind.TEXT:
+ assert any(message.content for message in output_messages)
+ else:
+ assert all(message.file for message in output_messages)
+ assert all(
+ message.file.storage.exists(message.file.name) for message in output_messages
+ )
+
+ invoice = Invoice.objects.get(user=user, model=model)
+ user.payment_plan.refresh_from_db()
+ charged_tokens = balance_before - user.payment_plan.current_token_balance
+ assert charged_tokens == invoice.cost
@@ -0,0 +1,34 @@
+# Тесты ML-моделей
+
+Основные тесты проходят полный API-флоу чата: проверяют сообщения, файлы, счета и
+списание баланса. Внешние провайдеры в них замоканы. Для каждой экспортированной
+модели должен быть кейс в `tests/ml_models/cases.py`.
+
+```bash
+# Все тесты
+python manage.py runtests
+
+# Одна или несколько моделей
+python manage.py runtests --model-slug gemma --model-slug qwen_3_6
+
+# С замерами времени, CPU и памяти
+python manage.py runtests --profile-resources
+```
+
+GitLab CI сохраняет JUnit-отчёт в `test-results/junit.xml` и показывает статистику
+тестов в pipeline и merge request.
+
+## Проверка реального провайдера
+
+Ручной smoke-тест запускает ровно одну модель без моков и печатает полученный ответ:
+
+```bash
+python manage.py runtests --provider-smoke --model-slug grok
+```
+
+Он использует ключи провайдера из окружения. Для моделей, которым нужен прокси,
+задайте `PROVIDER_SMOKE_PROXY_ADDRESS` в формате поля `Proxy.address` (без протокола)
+и `PROVIDER_SMOKE_PROXY_PROTOCOL`: `http`, `https` или `socks`.
+
+Этот тест может расходовать реальные токены. Обычный `runtests` и GitLab CI его не
+собирают.
@@ -0,0 +1,166 @@
+import resource
+import time
+from collections.abc import Iterator
+
+import pytest
+from django.core.files.storage import FileSystemStorage
+from rest_framework.test import APIClient
+from rest_framework_simplejwt.tokens import RefreshToken
+
+from messages.models import Message
+from poller.models import Proxy
+from tests.factories import UserFactory
+from tests.ml_models.cases import STARTING_BALANCE
+
+RESOURCE_PROFILES: list[dict[str, str | float | int]] = []
+PROVIDER_SMOKE_TEST_FILE = 'test_provider_smoke.py'
+
+
+def pytest_ignore_collect(collection_path, config) -> bool:
+ is_provider_smoke_test = collection_path.name == PROVIDER_SMOKE_TEST_FILE
+ is_manual_provider_smoke_run = config.getoption('--provider-smoke')
+
+ return is_provider_smoke_test and not is_manual_provider_smoke_run
+
+
+@pytest.fixture
+def api_client() -> APIClient:
+ return APIClient()
+
+
+@pytest.fixture
+def user(db):
+ user = UserFactory()
+ user.payment_plan.current_token_balance = STARTING_BALANCE
+ user.payment_plan.save()
+
+ return user
+
+
+@pytest.fixture
+def authenticated_client(api_client: APIClient, user) -> APIClient:
+ access_token = RefreshToken.for_user(user).access_token
+ api_client.credentials(HTTP_AUTHORIZATION=f'Bearer {access_token}')
+
+ return api_client
+
+
+@pytest.fixture
+def local_message_storage(tmp_path, monkeypatch) -> FileSystemStorage:
+ storage = FileSystemStorage(location=tmp_path)
+ file_field = Message._meta.get_field('file')
+ monkeypatch.setattr(file_field, 'storage', storage)
+
+ return storage
+
+
+@pytest.fixture
+def fake_provider_proxy(db) -> Proxy:
+ return Proxy.objects.create(address='proxy.test:8080', protocol=Proxy.ProtocolChoices.HTTP)
+
+
+def pytest_addoption(parser) -> None:
+ group = parser.getgroup('AIR ML model contracts')
+ group.addoption(
+ '--model-slug',
+ action='append',
+ default=[],
+ help='Run ML model API contracts only for the selected slug.',
+ )
+ group.addoption(
+ '--profile-resources',
+ action='store_true',
+ help='Report wall time, CPU time, and peak RSS for each test.',
+ )
+ group.addoption(
+ '--provider-smoke',
+ action='store_true',
+ help='Run one explicitly selected ML model against its real provider.',
+ )
+
+
+def pytest_collection_modifyitems(config, items) -> None:
+ selected_slugs = set(config.getoption('--model-slug'))
+ provider_smoke = config.getoption('--provider-smoke')
+
+ available_slugs = {
+ marker.args[0]
+ for item in items
+ if (marker := item.get_closest_marker('ml_model')) and marker.args
+ }
+ if unknown_slugs := selected_slugs - available_slugs:
+ available = ', '.join(sorted(available_slugs))
+ unknown = ', '.join(sorted(unknown_slugs))
+
+ raise pytest.UsageError(f'Unknown ML model slug: {unknown}. Available slugs: {available}')
+
+ if provider_smoke and len(selected_slugs) != 1:
+ raise pytest.UsageError('--provider-smoke requires exactly one --model-slug.')
+
+ deselected_items = []
+ selected_items = []
+ for item in items:
+ is_provider_smoke = item.get_closest_marker('provider_smoke') is not None
+ marker = item.get_closest_marker('ml_model')
+ slug = marker.args[0] if marker and marker.args else None
+
+ if provider_smoke:
+ is_selected = is_provider_smoke and slug in selected_slugs
+ else:
+ is_selected = not is_provider_smoke
+
+ if is_selected:
+ selected_items.append(item)
+ else:
+ deselected_items.append(item)
+
+ if deselected_items:
+ config.hook.pytest_deselected(items=deselected_items)
+ items[:] = selected_items
+
+ if provider_smoke:
+ return
+
+ for item in items:
+ marker = item.get_closest_marker('ml_model')
+ slug = marker.args[0] if marker and marker.args else None
+ if selected_slugs and slug not in selected_slugs:
+ item.add_marker(pytest.mark.skip(reason='ML model slug was not selected'))
+
+
+@pytest.fixture(autouse=True)
+def resource_profile(request) -> Iterator[None]:
+ if not request.config.getoption('--profile-resources'):
+ yield
+
+ return
+
+ usage_before = resource.getrusage(resource.RUSAGE_SELF)
+ wall_started = time.perf_counter()
+ cpu_started = time.process_time()
+
+ yield
+
+ usage_after = resource.getrusage(resource.RUSAGE_SELF)
+ peak_rss_mb = usage_after.ru_maxrss / 1024
+ RESOURCE_PROFILES.append(
+ {
+ 'test': request.node.nodeid,
+ 'wall_seconds': round(time.perf_counter() - wall_started, 6),
+ 'cpu_seconds': round(time.process_time() - cpu_started, 6),
+ 'peak_rss_mb': round(peak_rss_mb, 3),
+ 'minor_page_faults': usage_after.ru_minflt - usage_before.ru_minflt,
+ }
+ )
+
+
+def pytest_terminal_summary(terminalreporter, config) -> None:
+ if not config.getoption('--profile-resources') or not RESOURCE_PROFILES:
+ return
+
+ terminalreporter.section('resource profile')
+ for profile in RESOURCE_PROFILES:
+ terminalreporter.write_line(
+ '{test}: wall={wall_seconds:.3f}s cpu={cpu_seconds:.3f}s '
+ 'peak_rss={peak_rss_mb:.1f}MB minor_faults={minor_page_faults}'.format(**profile)
+ )
@@ -0,0 +1,72 @@
+from decimal import Decimal
+
+import factory
+from factory.django import DjangoModelFactory
+
+from authentication.models import CustomUserModel
+from ml_model.models import ModelCategory, ModelInput, ModelSettings, NeuronModel
+from payments.models import PaymentPlan
+from tools.chats.models import Chat
+
+
+class PaymentPlanFactory(DjangoModelFactory):
+ class Meta:
+ model = PaymentPlan
+ django_get_or_create = ('price',)
+
+ price = Decimal('1')
+ tokens_per_plan = Decimal('10000')
+
+
+class UserFactory(DjangoModelFactory):
+ class Meta:
+ model = CustomUserModel
+ skip_postgeneration_save = True
+
+ email = factory.Sequence(lambda number: f'ml-model-{number}@test.local')
+ password = factory.PostGenerationMethodCall('set_password', 'test-password')
+
+ @classmethod
+ def _create(cls, model_class, *args, **kwargs):
+ paid_plan = PaymentPlanFactory()
+ user = model_class.objects.create_user(*args, **kwargs)
+ user.payment_plan.plan = paid_plan
+ user.payment_plan.current_token_balance = paid_plan.tokens_per_plan
+ user.payment_plan.save(update_fields=('plan', 'current_token_balance'))
+
+ return user
+
+
+class ModelCategoryFactory(DjangoModelFactory):
+ class Meta:
+ model = ModelCategory
+ django_get_or_create = ('slug',)
+
+ title = 'Chat-bots'
+ slug = 'chat-bots'
+
+
+class NeuronModelFactory(DjangoModelFactory):
+ class Meta:
+ model = NeuronModel
+ skip_postgeneration_save = True
+
+ title = factory.LazyAttribute(lambda model: model.slug.replace('_', ' ').title())
+ slug = factory.Sequence(lambda number: f'test_model_{number}')
+ category = factory.SubFactory(ModelCategoryFactory)
+
+ @factory.post_generation
+ def configure(model, create, extracted, **kwargs):
+ if not create:
+ return
+ ModelSettings.objects.create(model=model, is_active=True)
+ ModelInput.objects.create(model=model, type=ModelInput.TypeChoices.TEXT, required=True)
+
+
+class ChatFactory(DjangoModelFactory):
+ class Meta:
+ model = Chat
+
+ title = 'ML model API contract'
+ user = factory.SubFactory(UserFactory)
+ model = factory.SubFactory(NeuronModelFactory)
@@ -15,6 +15,10 @@ class SSEChunkService:
def token(cls, event_id: int, content: str) -> SSEChunk:
return cls._chunk(event_id, 'token', {'content': content})
+ @classmethod
+ def think(cls, event_id: int, content: str) -> SSEChunk:
+ return cls._chunk(event_id, 'think', {'content': content})
+
@classmethod
def error(cls, event_id: int, content: str) -> SSEChunk:
return cls._chunk(event_id, 'error', {'detail': content})
@@ -8,7 +8,7 @@ from dataclasses import dataclass
from unittest.mock import patch
import orjson
-from cacheops import invalidate_all
+from cachalot.api import invalidate
from django.db import close_old_connections, connections
from django.test import Client, TransactionTestCase
from rest_framework_simplejwt.tokens import RefreshToken
@@ -75,7 +75,7 @@ class SSEStreamLoadTest(TransactionTestCase):
scenario: LoadScenario
def setUp(self) -> None:
- invalidate_all()
+ invalidate()
self.user = CustomUserModel.objects.create_user(email='sse-load@test.test', password='test')
self.access_token = str(RefreshToken.for_user(self.user).access_token)
category = ModelCategory.objects.create(title='Chat-bots', slug='chat-bots')
@@ -3,7 +3,7 @@ import time
from unittest.mock import patch
import orjson
-from cacheops import invalidate_all
+from cachalot.api import invalidate
from django.test import Client, TestCase
from rest_framework_simplejwt.tokens import RefreshToken
@@ -37,7 +37,7 @@ class SSEStreamAPITest(TestCase):
)
def setUp(self) -> None:
- invalidate_all()
+ invalidate()
self.client = Client()
self.user = CustomUserModel.objects.create_user(email='sse-test@test.test', password='test')
self.access_token = str(RefreshToken.for_user(self.user).access_token)
@@ -5,12 +5,16 @@ from dataclasses import asdict, dataclass
from tools.chats.typing import SSEData, SSEEvent
-@dataclass
-class SSEChunk:
- event_id: int
+@dataclass(frozen=True, slots=True)
+class RawSSEChunk:
event: SSEEvent
data: SSEData
+
+@dataclass(frozen=True, slots=True)
+class SSEChunk(RawSSEChunk):
+ event_id: int
+
def encode(self):
payload = orjson.dumps(self.data, default=str).decode()
return f'id: {self.event_id}\nevent: {self.event}\ndata: {payload}\n\n'
@@ -8,6 +8,7 @@ from django.db.models.functions import Greatest
from messages.models import Message
from ml_model.models import NeuronModel
from ml_model.services.base import StreamSimpleService
+from tools.chats.domain import RawSSEChunk
from tools.chats.models import Chat
from tools.chats.services.sse_chunk_service import SSEChunkService
from tools.chats.services.sse_store import PublicSSEStoreService, SSEStoreService
@@ -35,10 +36,17 @@ def _run_stream(
stream = service.make_stream(message)
while True:
token = next(stream)
- if not token:
+ if not isinstance(token, RawSSEChunk):
+ continue
+ token_data = token.data.get('content')
+ if not token_data or not isinstance(token_data, str):
continue
event_id += 1
- store.push(SSEChunkService.token(event_id, token))
+ store.push(
+ SSEChunkService.think(event_id, token_data)
+ if token.event == 'think'
+ else SSEChunkService.token(event_id, token_data)
+ )
except StopIteration as exc:
event_id += 1
store.push(SSEChunkService.done(event_id, exc.value or ''), ttl=settings.SSE_DONE_STREAM_TTL)
@@ -1,4 +1,4 @@
from typing import Any, Literal
-type SSEEvent = Literal['pending', 'start', 'token', 'error', 'done']
+type SSEEvent = Literal['pending', 'start', 'token', 'think', 'error', 'done']
type SSEData = dict[str, Any]
@@ -2,7 +2,6 @@ from datetime import date
from decimal import Decimal
from django.contrib.admin.models import ADDITION, DELETION, LogEntry
-from django.contrib.contenttypes.models import ContentType
from django.db import IntegrityError
from django.utils.translation import gettext as _
@@ -15,7 +14,7 @@ from tools.public_api.serializers import APIKeyResultSerializer
class APIKeyService(BaseService):
def create(self, payload: dict, serialize: bool = False) -> APIKey | APIKeyResultSerializer:
- if not (self.user.account_type in ('business_host', 'regular', 'business_admin')):
+ if self.user.account_type not in ('business_host', 'regular', 'business_admin'):
raise Exception('Can not create API key from business sub-account.')
try:
api_key = APIKey.objects.create(user=self.user, **payload)
@@ -23,13 +22,11 @@ class APIKeyService(BaseService):
if 'unique' in str(exc):
raise DuplicateError(model=APIKey, attrs=(_('Name'), _('Owner')))
raise UnknownError from exc
- LogEntry.objects.log_action(
- self.user.pk,
- ContentType.objects.get_for_model(api_key).pk,
- api_key.pk,
- str(api_key),
- ADDITION,
- [
+ LogEntry.objects.log_actions(
+ user_id=self.user.pk,
+ queryset=[api_key],
+ action_flag=ADDITION,
+ change_message=[
{
'added': {
'name': 'API-ключ',
@@ -37,6 +34,7 @@ class APIKeyService(BaseService):
}
}
],
+ single_object=True,
)
if serialize:
return APIKeyResultSerializer(api_key)
@@ -65,13 +63,11 @@ class APIKeyService(BaseService):
def delete(self, key_name):
api_key = APIKeySelector(self.user).get_by_name(name=key_name)
- LogEntry.objects.log_action(
- self.user.pk,
- ContentType.objects.get_for_model(api_key).pk,
- api_key.pk,
- str(api_key),
- DELETION,
- [
+ LogEntry.objects.log_actions(
+ user_id=self.user.pk,
+ queryset=[api_key],
+ action_flag=DELETION,
+ change_message=[
{
'deleted': {
'name': 'API-ключ',
@@ -79,6 +75,7 @@ class APIKeyService(BaseService):
}
}
],
+ single_object=True,
)
api_key.is_deleted = True
api_key.save()
@@ -0,0 +1,24 @@
+SECRET_KEY=ci-only-secret-key-that-is-long-enough
+TELEGRAM_BOT_TOKEN=ci-only-token
+DOMAIN=localhost
+DJANGO_SUPERUSER_USERNAME=ci-admin
+DJANGO_SUPERUSER_EMAIL=ci-admin@example.test
+DJANGO_SUPERUSER_PASSWORD=ci-only-password
+
+POSTGRES_DB=backend_ci
+POSTGRES_USER=backend_ci
+POSTGRES_PASSWORD=backend_ci
+POSTGRES_HOST=db
+POSTGRES_PORT=5432
+
+MINIO_ENDPOINT=s3:9000
+MINIO_ACCESS_KEY=ci-only-access-key
+MINIO_SECRET_KEY=ci-only-secret-key
+
+REDIS_HOST=cache-mdb
+CACHE_BROKER_URL=redis://cache-mdb:6379/0
+CELERY_BROKER_URL=redis://cache-mdb:6379/1
+CELERY_RESULT_BACKEND=redis://cache-mdb:6379/2
+
+RELEASE=ci
+ENVIRONMENT=test
@@ -116,4 +116,8 @@ PYTHONWARNINGS=ignore::UserWarning:polymorphic # temporarily
# SSE STREAMING
FF__STREAMING_ENABLED=True
-DATA_UPLOAD_MAX_MEMORY_SIZE=5 # MB
\ No newline at end of file
+DATA_UPLOAD_MAX_MEMORY_SIZE=5 # MB
+
+# MANUAL PROVIDER SMOKE TEST
+PROVIDER_SMOKE_PROXY_ADDRESS=
+PROVIDER_SMOKE_PROXY_PROTOCOL=http
@@ -10,6 +10,7 @@ venv/
.venv/
virtualenv/
air_reports/
+test-results/
.python-version
**/locales/**/*.mo
@@ -1,5 +1,6 @@
stages:
- Build
+ - Test
- Deploy
default:
@@ -19,6 +20,33 @@ build:
- staging
when: on_success
+test:
+ stage: Test
+ variables:
+ IMAGE_TAG: $CI_REGISTRY_IMAGE:$CI_COMMIT_SHA
+ COMPOSE_FILE: docker-compose.yml:docker-compose.local.yml
+ COMPOSE_PROJECT_NAME: $CI_PROJECT_NAME-test-$CI_PIPELINE_ID
+ script:
+ - cp .env.ci .env
+ - docker pull $IMAGE_TAG
+ - docker compose up -d --wait db cache-mdb
+ - until docker compose exec -T db pg_isready --username=backend_ci --dbname=backend_ci; do sleep 1; done
+ - mkdir -p test-results
+ - docker compose run --rm app python manage.py runtests --profile-resources --pytest-arg=--junitxml=/app/test-results/junit.xml
+ after_script:
+ - docker compose down --volumes --remove-orphans
+ artifacts:
+ when: always
+ expire_in: 1 week
+ reports:
+ junit: test-results/junit.xml
+ paths:
+ - test-results/junit.xml
+ only:
+ - main
+ - staging
+ when: on_success
+
.deploy_template: &default_deploy_job
stage: Deploy
services:
@@ -66,4 +94,4 @@ deploy_production:
deployment_tier: production
url: $DOMAIN
only:
- - main
\ No newline at end of file
+ - main
@@ -8,14 +8,14 @@ dependencies = [
"channels[daphne]==4.2.0",
"deepl==1.21.1",
"dj-rest-auth==4.0.1",
- "django==5.0.*",
- "django-cacheops==7.0.2",
+ "django==6.0.*",
+ "django-cachalot==2.9.0",
"django-celery-beat>=2.9.0",
"django-cors-headers==4.2.0",
"django-filter==23.2",
"django-import-export==4.0.9",
- "django-minio-backend",
- "django-ninja==1.3.0",
+ "django-minio-backend==4.5.0",
+ "django-ninja==1.6.2",
"django-oauth-toolkit==2.3.0",
"django-ordered-model==3.7.4",
"django-polymorphic==3.1.0",
@@ -38,6 +38,7 @@ dependencies = [
"langchain-openai==0.3.6",
"langchainhub==0.1.15",
"langserve[client]==0.0.46",
+ "markupsafe>=3.0.2",
"minio>=7.0,<=8.0",
"mutagen==1.47.0",
"openpyxl==3.1.2",
@@ -54,6 +55,7 @@ dependencies = [
"sentry-sdk[django]==2.39.0",
"setuptools<81",
"social-auth-app-django==5.3.0",
+ "social-auth-core==4.9.1",
"tiktoken==0.9.0",
"unleashclient==6.4.0",
"yookassa==3.10.1",
@@ -64,6 +66,9 @@ debug = [
"debugpy>=1.8.20",
]
dev = [
+ "factory-boy>=3.3.3",
+ "pytest>=9.0.2",
+ "pytest-django>=4.11.1",
"ruff>=0.15.15",
]
@@ -124,8 +129,15 @@ line-ending = "lf"
docstring-code-format = false
docstring-code-line-length = "dynamic"
-[tool.uv.sources]
-django-minio-backend = { git = "https://github.com/theriverman/django-minio-backend", tag = "3.7.0" }
+[tool.pytest.ini_options]
+DJANGO_SETTINGS_MODULE = "backend.settings"
+python_files = ["test_*.py"]
+testpaths = ["tests"]
+addopts = "-ra"
+markers = [
+ "ml_model(slug): API contract for a concrete ML model service",
+ "provider_smoke: opt-in API test that calls a real external provider",
+]
[tool.ruff.lint.per-file-ignores]
"__init__.py" = ["E402", "F401"]