Skip to content

Commit ef539fb

Browse files
committed
feat: implement token coalescing and buffer flushing in NatsAnswerStreamer
1 parent ebe3249 commit ef539fb

2 files changed

Lines changed: 117 additions & 1 deletion

File tree

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)