diff --git a/.github/workflows/tests.yml b/.github/workflows/tests.yml index b0f491f..95f4688 100644 --- a/.github/workflows/tests.yml +++ b/.github/workflows/tests.yml @@ -26,7 +26,7 @@ jobs: version: latest - name: Install dependencies - run: poetry install --no-interaction --no-ansi + run: poetry install --no-interaction --no-ansi --extras dev - name: Check code with ruff run: poetry run ruff check . @@ -50,7 +50,7 @@ jobs: version: latest - name: Install dependencies - run: poetry install --no-interaction --no-ansi + run: poetry install --no-interaction --no-ansi --extras dev - name: Check code formatting with black run: poetry run black --check . @@ -74,7 +74,7 @@ jobs: version: latest - name: Install dependencies - run: poetry install --no-interaction --no-ansi + run: poetry install --no-interaction --no-ansi --extras dev - name: Run mypy run: poetry run mypy . @@ -98,7 +98,7 @@ jobs: version: latest - name: Install dependencies - run: poetry install --no-interaction --no-ansi + run: poetry install --no-interaction --no-ansi --extras dev - name: Run tests run: poetry run pytest src/tests/ -v diff --git a/pyproject.toml b/pyproject.toml index dca75b7..bb36d55 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,3 +1,6 @@ +[tool.poetry] +packages = [{ include = "app", from = "src" }] + [project] name = "assistant-chat-backend" version = "0.1.0" diff --git a/src/app/infrastructure/streaming/answer_streamer_nats.py b/src/app/infrastructure/streaming/answer_streamer_nats.py index 83660eb..fdebd0d 100644 --- a/src/app/infrastructure/streaming/answer_streamer_nats.py +++ b/src/app/infrastructure/streaming/answer_streamer_nats.py @@ -1,6 +1,7 @@ import asyncio import json import logging +import re from collections.abc import AsyncGenerator from nats.aio.client import Client @@ -13,12 +14,32 @@ logger = logging.getLogger(__name__) _SSE_TIMEOUT = 60.0 +_MAX_COALESCE_BUFFER = 4096 def _sse_event(event: str, data: str) -> str: return f"event: {event}\ndata: {data}\n\n" +def _take_whitespace_delimited_prefixes(buffer: str) -> tuple[str, list[str]]: + """Split leading segments that end at the first whitespace run (inclusive).""" + emitted: list[str] = [] + rest = buffer + while True: + m = re.search(r"\s+", rest) + if not m: + break + emitted.append(rest[: m.end()]) + rest = rest[m.end() :] + return rest, emitted + + +def _flush_oversized_tail(buffer: str, max_len: int) -> tuple[str, list[str]]: + if len(buffer) <= max_len: + return buffer, [] + return "", [buffer] + + class NatsAnswerStreamer(AnswerStreamer): def __init__( self, @@ -51,24 +72,44 @@ async def _on_message(msg: Msg) -> None: sub = await self._nats.subscribe(subject, cb=_on_message) logger.info("SSE subscribed to %s", subject) + token_buffer = "" try: while True: try: item = await asyncio.wait_for(queue.get(), timeout=_SSE_TIMEOUT) except TimeoutError: logger.warning("SSE timeout for message_id=%s", message_id) + if token_buffer: + yield _sse_event("token", token_buffer) yield _sse_event("done", "end") break if item is None: + if token_buffer: + yield _sse_event("token", token_buffer) yield _sse_event("done", "end") break parsed = json.loads(item) if parsed["event"] == "done": + if token_buffer: + yield _sse_event("token", token_buffer) yield _sse_event("done", "end") break - yield _sse_event("token", parsed["data"]) + + chunk = parsed["data"] + if not isinstance(chunk, str): + chunk = str(chunk) + token_buffer += chunk + token_buffer, emitted = _take_whitespace_delimited_prefixes(token_buffer) + for part in emitted: + yield _sse_event("token", part) + token_buffer, forced = _flush_oversized_tail( + token_buffer, + _MAX_COALESCE_BUFFER, + ) + for part in forced: + yield _sse_event("token", part) finally: await sub.unsubscribe() logger.info("SSE unsubscribed from %s", subject) diff --git a/src/tests/unit/infrastructure/test_answer_streamer_coalescing.py b/src/tests/unit/infrastructure/test_answer_streamer_coalescing.py new file mode 100644 index 0000000..421f50b --- /dev/null +++ b/src/tests/unit/infrastructure/test_answer_streamer_coalescing.py @@ -0,0 +1,75 @@ +from app.infrastructure.streaming.answer_streamer_nats import ( + _MAX_COALESCE_BUFFER, + _flush_oversized_tail, + _take_whitespace_delimited_prefixes, +) + + +def _feed_tokens(parts: list[str]) -> tuple[str, list[str]]: + """Mimic NatsAnswerStreamer token accumulation + coalescing.""" + buf = "" + out: list[str] = [] + for p in parts: + buf += p + buf, emitted = _take_whitespace_delimited_prefixes(buf) + out.extend(emitted) + buf, forced = _flush_oversized_tail(buf, _MAX_COALESCE_BUFFER) + out.extend(forced) + return buf, out + + +def test_take_prefixes_single_word_with_trailing_space() -> None: + rest, emitted = _take_whitespace_delimited_prefixes("Куликов ") + assert rest == "" + assert emitted == ["Куликов "] + + +def test_take_prefixes_incremental_kulikov() -> None: + tail, emitted = _feed_tokens(["К", "ули", "ков", " "]) + assert tail == "" + assert emitted == ["Куликов "] + + +def test_take_prefixes_multiple_words_one_buffer() -> None: + rest, emitted = _take_whitespace_delimited_prefixes("a b c ") + assert rest == "" + assert emitted == ["a ", "b ", "c "] + + +def test_take_prefixes_tail_without_whitespace_stays() -> None: + rest, emitted = _take_whitespace_delimited_prefixes("Куликов") + assert rest == "Куликов" + assert emitted == [] + + +def test_take_prefixes_leading_spaces_then_word() -> None: + rest, emitted = _take_whitespace_delimited_prefixes(" Дмитрий ") + assert rest == "" + assert emitted == [" ", "Дмитрий "] + + +def test_incremental_flushes_tail_on_done_semantics() -> None: + tail, emitted = _feed_tokens(["часть"]) + assert tail == "часть" + assert emitted == [] + + +def test_flush_oversized_tail_no_whitespace() -> None: + long = "x" * (_MAX_COALESCE_BUFFER + 1) + rest, forced = _flush_oversized_tail(long, _MAX_COALESCE_BUFFER) + assert rest == "" + assert forced == [long] + + +def test_flush_oversized_tail_under_limit_unchanged() -> None: + rest, forced = _flush_oversized_tail("word", _MAX_COALESCE_BUFFER) + assert rest == "word" + assert forced == [] + + +def test_incremental_chunks_crossing_max_without_space() -> None: + piece = "n" * 2000 + tail, emitted = _feed_tokens([piece, piece, piece]) + assert tail == "" + assert len(emitted) == 1 + assert len(emitted[0]) == 6000