@@ -13,7 +13,12 @@ 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.exceptions import ( + CorruptedFileError, + FileExtensionNotSupported, + PaidPlanRequiredError, + ModelVersionNotAvailable, +) from ml_model.services.base import StreamSimpleService from ml_model.services.EmbeddingService import EmbeddingService from ml_model.services.FileService import FileProcessingService @@ -67,16 +72,21 @@ class Grok(SerperMixin, StreamSimpleService): def make(self, input_message: Message, save: bool = True) -> list[Message]: start_time = time.time() - version = input_message.info.get('version') or 'grok-4.5' + version_slug = input_message.info.get('version') + if version_slug is None or version_slug not in self.TOKENS_COST: + raise ModelVersionNotAvailable(version_slug, self.TOKENS_COST) + callback_data = {**input_message.info, 'tools': []} - messages, embedding_tokens = self._prepare_messages(input_message, version, callback_data) + messages, embedding_tokens = self._prepare_messages(input_message, version_slug, callback_data) - result = OpenrouterAdapter.collect_streaming_api(f'x-ai/{version}', messages, callback_data, 'Grok') + result = OpenrouterAdapter.collect_streaming_api( + f'x-ai/{version_slug}', messages, callback_data, 'Grok' + ) process_time = timedelta(seconds=(time.time() - start_time)) self.handle_invoice( input_message.content_object.model, - version=version, + version=version_slug, cost=result.cost, input_tokens=result.input_tokens, output_tokens=result.output_tokens, @@ -86,9 +96,11 @@ class Grok(SerperMixin, StreamSimpleService): 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' + version_slug = input_message.info.get('version') + if version_slug is None or version_slug not in self.TOKENS_COST: + raise ModelVersionNotAvailable(version_slug, self.TOKENS_COST) callback_data = {**input_message.info, 'tools': []} - messages, embedding_tokens = self._prepare_messages(input_message, version, callback_data) + messages, embedding_tokens = self._prepare_messages(input_message, version_slug, callback_data) input_tokens = output_tokens = 0 cost = 0 reasoning = '' @@ -96,7 +108,9 @@ class Grok(SerperMixin, StreamSimpleService): result = '' try: - stream = OpenrouterAdapter.run_streaming_api(f'x-ai/{version}', messages, callback_data, 'Grok') + stream = OpenrouterAdapter.run_streaming_api( + f'x-ai/{version_slug}', messages, callback_data, 'Grok' + ) try: while True: chunk = next(stream) @@ -114,7 +128,7 @@ class Grok(SerperMixin, StreamSimpleService): process_time = timedelta(seconds=(time.time() - start_time)) self.handle_invoice( input_message.content_object.model, - version=version, + version=version_slug, cost=cost, input_tokens=input_tokens, output_tokens=output_tokens, @@ -194,10 +208,7 @@ class Grok(SerperMixin, StreamSimpleService): if (chunks_length := sum(len(chunk) for chunk in chunks)) > 20_000: predict_embedding_tokens = len(chunks) * 2020 predicted_input_price += ( - ( - Decimal('210') - + Decimal(chunks_length) / Decimal(len(chunks)) * Decimal('10') - ) + (Decimal('210') + Decimal(chunks_length) / Decimal(len(chunks)) * Decimal('10')) / Decimal('2.0') * self.TOKENS_COST[version]['input'] / Decimal('1_000_000') @@ -251,21 +262,14 @@ class Grok(SerperMixin, StreamSimpleService): if image: predicted_image_tokens = min((image_width * image_height + 999) // 1000, 2500) predicted_input_price += ( - Decimal(predicted_image_tokens) - * self.TOKENS_COST[version]['input'] - / Decimal('1_000_000') + Decimal(predicted_image_tokens) * self.TOKENS_COST[version]['input'] / Decimal('1_000_000') ) - estimated_input_tokens = ( - Decimal( - sum( - len(message['content']) if isinstance(message['content'], str) else 0 - for message in messages - ) - + (len(input_message.content) if image else 0) + estimated_input_tokens = Decimal( + sum( + len(message['content']) if isinstance(message['content'], str) else 0 for message in messages ) - / Decimal('2.0') - + (150 if is_free_plan else 250) - ) + + (len(input_message.content) if image else 0) + ) / Decimal('2.0') + (150 if is_free_plan else 250) predicted_input_price += ( estimated_input_tokens * self.TOKENS_COST[version]['input'] / Decimal('1_000_000') + predict_embedding_tokens * self.TOOLS_TOKEN_COSTS['text-embedding-3-small']['output'] @@ -14,7 +14,24 @@ from rest_framework.views import APIView from messages.models import Message from messages.serializers import MessageSerializer from ml_model.choices import ContentTypes -from ml_model.exceptions import FileNotProvided, InvalidParameterError +from ml_model.exceptions import ( + CorruptedFileError, + ExceededContextLengthError, + FileExtensionNotSupported, + FileNotProvided, + FileTooLargeError, + FileUploadUnsupported, + ImageAnalysisError, + ImageTooLargeError, + InputImageSensitiveContentError, + InvalidParameterError, + ModelVersionNotAvailable, + OutputSensitiveImageContentError, + PaidPlanRequiredError, + PromptLengthExceeded, + RequestBlocked, + UnrecognizedFileError, +) from ml_model.models import NeuronModel from ml_model.selectors.ml_models_selector import NeuronModelSelector from ml_model.serializers import PublicNeuronModelSerializer @@ -94,13 +111,31 @@ class BaseGenerationView(APIView): # WARNING: output должен быть списком! try: output_message = service(store).make(input_message) + except PaidPlanRequiredError as exc: + return Response({'detail': str(exc)}, status=HTTP_402_PAYMENT_REQUIRED) + except ( + FileExtensionNotSupported, + ExceededContextLengthError, + RequestBlocked, + PromptLengthExceeded, + CorruptedFileError, + FileTooLargeError, + ImageAnalysisError, + FileUploadUnsupported, + UnrecognizedFileError, + InvalidParameterError, + ModelVersionNotAvailable, + InputImageSensitiveContentError, + OutputSensitiveImageContentError, + ImageTooLargeError, + ValidationError, + ) as exc: + return Response({'detail': str(exc)}, status=HTTP_400_BAD_REQUEST) except Exception as exc: input_message.is_sent = False input_message.save() if isinstance(exc, InsufficientBalance): return Response({'detail': str(exc)}, status=HTTP_402_PAYMENT_REQUIRED) - elif isinstance(exc, (InvalidParameterError, ValidationError)): - return Response({'detail': str(exc)}, status=HTTP_400_BAD_REQUEST) logger.exception(exc) return Response( {