Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 4 additions & 4 deletions .github/workflows/tests.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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 .
Expand All @@ -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 .
Expand All @@ -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 .
Expand All @@ -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
3 changes: 3 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
@@ -1,3 +1,6 @@
[tool.poetry]
packages = [{ include = "app", from = "src" }]

[project]
name = "assistant-chat-backend"
version = "0.1.0"
Expand Down
43 changes: 42 additions & 1 deletion src/app/infrastructure/streaming/answer_streamer_nats.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import asyncio
import json
import logging
import re
from collections.abc import AsyncGenerator

from nats.aio.client import Client
Expand All @@ -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,
Expand Down Expand Up @@ -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)
75 changes: 75 additions & 0 deletions src/tests/unit/infrastructure/test_answer_streamer_coalescing.py
Original file line number Diff line number Diff line change
@@ -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
Loading