Skip to content

Commit a04ad2a

Browse files
authored
Merge pull request #1 from UniCortex/fix/backend-01-token-stream-format
feat: implement token coalescing and buffer flushing in NatsAnswerStreamer
2 parents ebe3249 + ec08e4d commit a04ad2a

4 files changed

Lines changed: 124 additions & 5 deletions

File tree

.github/workflows/tests.yml

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -26,7 +26,7 @@ jobs:
2626
version: latest
2727

2828
- name: Install dependencies
29-
run: poetry install --no-interaction --no-ansi
29+
run: poetry install --no-interaction --no-ansi --extras dev
3030

3131
- name: Check code with ruff
3232
run: poetry run ruff check .
@@ -50,7 +50,7 @@ jobs:
5050
version: latest
5151

5252
- name: Install dependencies
53-
run: poetry install --no-interaction --no-ansi
53+
run: poetry install --no-interaction --no-ansi --extras dev
5454

5555
- name: Check code formatting with black
5656
run: poetry run black --check .
@@ -74,7 +74,7 @@ jobs:
7474
version: latest
7575

7676
- name: Install dependencies
77-
run: poetry install --no-interaction --no-ansi
77+
run: poetry install --no-interaction --no-ansi --extras dev
7878

7979
- name: Run mypy
8080
run: poetry run mypy .
@@ -98,7 +98,7 @@ jobs:
9898
version: latest
9999

100100
- name: Install dependencies
101-
run: poetry install --no-interaction --no-ansi
101+
run: poetry install --no-interaction --no-ansi --extras dev
102102

103103
- name: Run tests
104104
run: poetry run pytest src/tests/ -v

pyproject.toml

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,6 @@
1+
[tool.poetry]
2+
packages = [{ include = "app", from = "src" }]
3+
14
[project]
25
name = "assistant-chat-backend"
36
version = "0.1.0"

src/app/infrastructure/streaming/answer_streamer_nats.py

Lines changed: 42 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
import asyncio
22
import json
33
import logging
4+
import re
45
from collections.abc import AsyncGenerator
56

67
from nats.aio.client import Client
@@ -13,12 +14,32 @@
1314
logger = logging.getLogger(__name__)
1415

1516
_SSE_TIMEOUT = 60.0
17+
_MAX_COALESCE_BUFFER = 4096
1618

1719

1820
def _sse_event(event: str, data: str) -> str:
1921
return f"event: {event}\ndata: {data}\n\n"
2022

2123

24+
def _take_whitespace_delimited_prefixes(buffer: str) -> tuple[str, list[str]]:
25+
"""Split leading segments that end at the first whitespace run (inclusive)."""
26+
emitted: list[str] = []
27+
rest = buffer
28+
while True:
29+
m = re.search(r"\s+", rest)
30+
if not m:
31+
break
32+
emitted.append(rest[: m.end()])
33+
rest = rest[m.end() :]
34+
return rest, emitted
35+
36+
37+
def _flush_oversized_tail(buffer: str, max_len: int) -> tuple[str, list[str]]:
38+
if len(buffer) <= max_len:
39+
return buffer, []
40+
return "", [buffer]
41+
42+
2243
class NatsAnswerStreamer(AnswerStreamer):
2344
def __init__(
2445
self,
@@ -51,24 +72,44 @@ async def _on_message(msg: Msg) -> None:
5172
sub = await self._nats.subscribe(subject, cb=_on_message)
5273
logger.info("SSE subscribed to %s", subject)
5374

75+
token_buffer = ""
5476
try:
5577
while True:
5678
try:
5779
item = await asyncio.wait_for(queue.get(), timeout=_SSE_TIMEOUT)
5880
except TimeoutError:
5981
logger.warning("SSE timeout for message_id=%s", message_id)
82+
if token_buffer:
83+
yield _sse_event("token", token_buffer)
6084
yield _sse_event("done", "end")
6185
break
6286

6387
if item is None:
88+
if token_buffer:
89+
yield _sse_event("token", token_buffer)
6490
yield _sse_event("done", "end")
6591
break
6692

6793
parsed = json.loads(item)
6894
if parsed["event"] == "done":
95+
if token_buffer:
96+
yield _sse_event("token", token_buffer)
6997
yield _sse_event("done", "end")
7098
break
71-
yield _sse_event("token", parsed["data"])
99+
100+
chunk = parsed["data"]
101+
if not isinstance(chunk, str):
102+
chunk = str(chunk)
103+
token_buffer += chunk
104+
token_buffer, emitted = _take_whitespace_delimited_prefixes(token_buffer)
105+
for part in emitted:
106+
yield _sse_event("token", part)
107+
token_buffer, forced = _flush_oversized_tail(
108+
token_buffer,
109+
_MAX_COALESCE_BUFFER,
110+
)
111+
for part in forced:
112+
yield _sse_event("token", part)
72113
finally:
73114
await sub.unsubscribe()
74115
logger.info("SSE unsubscribed from %s", subject)
Lines changed: 75 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,75 @@
1+
from app.infrastructure.streaming.answer_streamer_nats import (
2+
_MAX_COALESCE_BUFFER,
3+
_flush_oversized_tail,
4+
_take_whitespace_delimited_prefixes,
5+
)
6+
7+
8+
def _feed_tokens(parts: list[str]) -> tuple[str, list[str]]:
9+
"""Mimic NatsAnswerStreamer token accumulation + coalescing."""
10+
buf = ""
11+
out: list[str] = []
12+
for p in parts:
13+
buf += p
14+
buf, emitted = _take_whitespace_delimited_prefixes(buf)
15+
out.extend(emitted)
16+
buf, forced = _flush_oversized_tail(buf, _MAX_COALESCE_BUFFER)
17+
out.extend(forced)
18+
return buf, out
19+
20+
21+
def test_take_prefixes_single_word_with_trailing_space() -> None:
22+
rest, emitted = _take_whitespace_delimited_prefixes("Куликов ")
23+
assert rest == ""
24+
assert emitted == ["Куликов "]
25+
26+
27+
def test_take_prefixes_incremental_kulikov() -> None:
28+
tail, emitted = _feed_tokens(["К", "ули", "ков", " "])
29+
assert tail == ""
30+
assert emitted == ["Куликов "]
31+
32+
33+
def test_take_prefixes_multiple_words_one_buffer() -> None:
34+
rest, emitted = _take_whitespace_delimited_prefixes("a b c ")
35+
assert rest == ""
36+
assert emitted == ["a ", "b ", "c "]
37+
38+
39+
def test_take_prefixes_tail_without_whitespace_stays() -> None:
40+
rest, emitted = _take_whitespace_delimited_prefixes("Куликов")
41+
assert rest == "Куликов"
42+
assert emitted == []
43+
44+
45+
def test_take_prefixes_leading_spaces_then_word() -> None:
46+
rest, emitted = _take_whitespace_delimited_prefixes(" Дмитрий ")
47+
assert rest == ""
48+
assert emitted == [" ", "Дмитрий "]
49+
50+
51+
def test_incremental_flushes_tail_on_done_semantics() -> None:
52+
tail, emitted = _feed_tokens(["часть"])
53+
assert tail == "часть"
54+
assert emitted == []
55+
56+
57+
def test_flush_oversized_tail_no_whitespace() -> None:
58+
long = "x" * (_MAX_COALESCE_BUFFER + 1)
59+
rest, forced = _flush_oversized_tail(long, _MAX_COALESCE_BUFFER)
60+
assert rest == ""
61+
assert forced == [long]
62+
63+
64+
def test_flush_oversized_tail_under_limit_unchanged() -> None:
65+
rest, forced = _flush_oversized_tail("word", _MAX_COALESCE_BUFFER)
66+
assert rest == "word"
67+
assert forced == []
68+
69+
70+
def test_incremental_chunks_crossing_max_without_space() -> None:
71+
piece = "n" * 2000
72+
tail, emitted = _feed_tokens([piece, piece, piece])
73+
assert tail == ""
74+
assert len(emitted) == 1
75+
assert len(emitted[0]) == 6000

0 commit comments

Comments
 (0)