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