@@ -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"]