@@ -1,16 +1,21 @@ +import json import time from datetime import timedelta from decimal import Decimal from pathlib import Path +from typing import Iterator import filetype +import httpx +from backend import settings from messages.models import Message +from ml_model.adapters.openrouter import OpenrouterAdapter from ml_model.exceptions import CorruptedFileError, FileExtensionNotSupported from ml_model.services.EmbeddingService import EmbeddingService from ml_model.services.FileService import FileProcessingService from ml_model.exceptions import ModelVersionNotAvailable -from ml_model.services.base import SimpleService +from ml_model.services.base import StreamSimpleService from ml_model.tasks import openrouter_run from poller.models import Proxy from tools.chats.models import Chat @@ -18,7 +23,7 @@ from tools.copywrite.models import Copywrite from tools.public_api.models import APIStore -class Qwen_3_7(SimpleService): +class Qwen_3_7(StreamSimpleService): COEFFICIENT = Decimal('300.0') TOKENS_COST = { @@ -73,6 +78,61 @@ class Qwen_3_7(SimpleService): def make(self, input_message: Message, save: bool = True) -> list[Message]: start_time = time.time() + version_slug, model_slug, callback_data, messages, embedding_tokens = self._prepare_data( + input_message + ) + result = openrouter_run(model_slug, messages, callback_data, 'Qwen 3.7') + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice( + input_message.content_object.model, + version=version_slug, + cost=result[1], + embedding_tokens=embedding_tokens, + ) + msgs = self.save_results(result[0], process_time) + return msgs + + def make_stream(self, input_message: Message, save: bool = True) -> Iterator[str]: + start_time = time.time() + version_slug, model_slug, callback_data, messages, embedding_tokens = self._prepare_data( + input_message + ) + content_parts: list[str] = [] + cost = 0 + result = '' + + try: + stream = self._run_streaming_api(model_slug, messages, callback_data) + try: + while True: + chunk = next(stream) + if chunk: + content_parts.append(chunk) + yield chunk + except StopIteration as exc: + cost = exc.value + finally: + if content_parts: + result = ''.join(content_parts) + 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( + input_message.content_object.model, + version=version_slug, + cost=cost, + embedding_tokens=embedding_tokens, + ) + self.save_results(result, process_time, save) + if result: + return result + + def _prepare_data( + self, input_message: Message + ) -> tuple[str, str, dict, list[dict[str, str | list]], int]: version_slug = input_message.info.pop('version', None) if version_slug is None or version_slug not in self.TOKENS_COST: raise ModelVersionNotAvailable(version_slug, self.TOKENS_COST) @@ -133,16 +193,7 @@ class Qwen_3_7(SimpleService): } ) model_slug = f'qwen/{version_slug}' - result = openrouter_run(model_slug, messages, callback_data, 'Qwen 3.7') - process_time = timedelta(seconds=(time.time() - start_time)) - self.handle_invoice( - input_message.content_object.model, - version=version_slug, - cost=result[1], - embedding_tokens=embedding_tokens, - ) - msgs = self.save_results(result[0], process_time) - return msgs + return version_slug, model_slug, callback_data, messages, embedding_tokens def get_chat_history( self, message_limit: int = 10, max_character_limit: int = 1500 @@ -181,3 +232,54 @@ class Qwen_3_7(SimpleService): character_length -= len(memory.pop(0)['content']) return memory + + def _estimate_cost(self, version_slug: str, input_tokens: int, output_tokens: int) -> float: + price_map = self.TOKENS_COST[version_slug] + price = ( + input_tokens * price_map['input'] / 1_000_000 + + 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