@@ -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 @@ -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: ... @@ -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, @@ -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}) @@ -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 @@ -25,10 +26,17 @@ def _run_stream(store: SSEStoreService, message_uuid: str, service: StreamSimple 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]