Skip to content

Commit 201dd28

Browse files
fix: stream provider deltas through SSE
Connect provider streaming to the SSE endpoint while preserving tool receipts, usage, memory, completion, and legacy payload fields. Add a timing regression and make the precision test use a fixed peak timestamp.\n\nFixes #86
1 parent ad7f62c commit 201dd28

5 files changed

Lines changed: 233 additions & 38 deletions

File tree

‎docs/self-evolution.md‎

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@ Historical verification snapshot: `51be33361422e55e1f2f00c33a0e0f8c56132a91`
99
(the post-#54 `main` revision, captured before this #55 documentation-only
1010
update). Snapshot date: 2026-09-04.
1111

12-
Current repository test count at this snapshot: **160 unittest cases**.
12+
Current repository test count at this snapshot: **161 unittest cases**.
1313

1414
## Verified surface
1515

@@ -106,8 +106,11 @@ curl -sS http://127.0.0.1:8000/api/chat \
106106
-d '{"message":"Write a sourced migration note", "profile":"researcher", "session_id":"demo", "speaker":"authenticated", "audience":"team", "channel":"chat"}'
107107
```
108108

109-
The streaming endpoint `/api/chat/stream` accepts the same body and emits the
110-
receipt as an SSE record when one is available.
109+
The streaming endpoint `/api/chat/stream` accepts the same body and forwards
110+
provider content deltas as they arrive. Control responses are buffered only
111+
long enough to parse actions; durable tool receipts, task progress, usage,
112+
memory, and completion are emitted as typed SSE records. The existing `chunk`,
113+
`cost`, `memory_receipt`, `error`, and `[DONE]` payloads remain compatible.
111114

112115
### Inspect learning evidence
113116

‎main.py‎

Lines changed: 47 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -508,8 +508,20 @@ def _retrieve_failure(query: str, n: int = 3) -> list[str]:
508508
_last_prompt_tokens: int = 0
509509
_last_completion_tokens: int = 0
510510
_active_usage_run_id: ContextVar[str | None] = ContextVar("active_usage_run_id", default=None)
511+
_stream_event_callback: ContextVar[Any] = ContextVar("stream_event_callback", default=None)
511512
_turn_cost_log: list[dict] = [] # {"tokens":int, "time":float, "tool_calls":int}
512513

514+
515+
def _emit_stream_event(event: dict[str, Any]) -> None:
516+
"""Forward typed stream progress without letting a client affect the turn."""
517+
callback = _stream_event_callback.get()
518+
if not callable(callback):
519+
return
520+
try:
521+
callback(event)
522+
except Exception:
523+
pass
524+
513525
def _track_tool_performance(action: str, result: str, elapsed: float) -> None:
514526
stats = _tool_stats.setdefault(action, {"calls":0,"successes":0,"total_time":0.0})
515527
stats["calls"] += 1
@@ -997,6 +1009,13 @@ def _spinner_worker(stop_event: threading.Event) -> None:
9971009

9981010
def _call_llm_with_spinner(messages: list[dict], model: str | None = None) -> str:
9991011
global _SPINNER_STOP, _SPINNER_THREAD
1012+
streaming = callable(_stream_event_callback.get())
1013+
if streaming:
1014+
return _get_llm_response(
1015+
messages, model=model, stream=True,
1016+
on_chunk=lambda chunk: _emit_stream_event({"event": "content", "chunk": str(chunk)}),
1017+
on_stream_end=lambda: _emit_stream_event({"event": "model_complete"}),
1018+
)
10001019
_SPINNER_STOP.clear()
10011020
_SPINNER_THREAD = threading.Thread(target=_spinner_worker, args=(_SPINNER_STOP,), daemon=True)
10021021
_SPINNER_THREAD.start()
@@ -3303,12 +3322,18 @@ def _record_turn_receipt(receipt: ExecutionReceipt) -> dict[str, Any]:
33033322
"execution.receipt", receipt.as_dict(), user_id=tasks.user_id,
33043323
workspace_id=tasks.workspace_id, session_id=tasks.session_id, task_id=task_id,
33053324
)
3306-
return {
3325+
result = {
33073326
"receipt_id": receipt.receipt_id, "operation_id": receipt.operation_id,
33083327
"action": receipt.action, "args": receipt.args, "result": receipt.result,
33093328
"success": receipt.success, "authorized": receipt.authorized,
33103329
"acceptance": receipt.acceptance, "failure": receipt.failure,
33113330
}
3331+
_emit_stream_event({"event": "tool_receipt", "tool_receipt": result})
3332+
_emit_stream_event({"event": "tasks", "tasks": [
3333+
{"id": item["id"], "description": item["description"], "status": item["status"]}
3334+
for item in tasks.tasks
3335+
]})
3336+
return result
33123337

33133338

33143339
def _execute_durable_task(task: dict[str, Any]) -> dict[str, Any]:
@@ -3449,7 +3474,8 @@ def run() -> None:
34493474
return value
34503475

34513476

3452-
def _get_llm_response(messages: list[dict[str, str]], model: str | None = None, stream: bool = False) -> str:
3477+
def _get_llm_response(messages: list[dict[str, str]], model: str | None = None, stream: bool = False,
3478+
on_chunk: Any = None, on_stream_end: Any = None) -> str:
34533479
global _last_prompt_tokens, _last_completion_tokens, _total_prompt_tokens, _total_completion_tokens
34543480
if llm_provider is None:
34553481
return "[Error] LLM provider not initialised"
@@ -3459,15 +3485,29 @@ def _get_llm_response(messages: list[dict[str, str]], model: str | None = None,
34593485
workspace_id=memory_bank.workspace_id, session_id=memory_bank.session_id,
34603486
run_id=_active_usage_run_id.get(), surface=_EXECUTION_SURFACE):
34613487
if stream and hasattr(llm_provider, 'chat_stream'):
3462-
# Streaming usage frames are persisted by the streaming implementation.
3488+
before = memory_bank.store.usage_totals(
3489+
user_id=memory_bank.user_id, workspace_id=memory_bank.workspace_id,
3490+
session_id=memory_bank.session_id, run_id=_active_usage_run_id.get(),
3491+
)
34633492
collected: list[str] = []
34643493
for chunk in llm_provider.chat_stream(messages, model or DEEPSEEK_MODEL):
34653494
collected.append(chunk)
3466-
sys.stdout.write(chunk)
3467-
sys.stdout.flush()
3495+
if on_chunk:
3496+
on_chunk(chunk)
3497+
else:
3498+
sys.stdout.write(chunk)
3499+
sys.stdout.flush()
34683500
text = "".join(collected).strip()
3469-
_last_prompt_tokens = 0
3470-
_last_completion_tokens = len(text) // 4 # presentation-only until a usage frame arrives
3501+
after = memory_bank.store.usage_totals(
3502+
user_id=memory_bank.user_id, workspace_id=memory_bank.workspace_id,
3503+
session_id=memory_bank.session_id, run_id=_active_usage_run_id.get(),
3504+
)
3505+
_last_prompt_tokens = max(0, int(after["prompt_tokens"]) - int(before["prompt_tokens"]))
3506+
_last_completion_tokens = max(0, int(after["completion_tokens"]) - int(before["completion_tokens"]))
3507+
_total_prompt_tokens += _last_prompt_tokens
3508+
_total_completion_tokens += _last_completion_tokens
3509+
if on_stream_end:
3510+
on_stream_end()
34713511
else:
34723512
text, usage_dict = _bounded_provider_call(
34733513
lambda: llm_provider.chat(messages, model or DEEPSEEK_MODEL)

‎server.py‎

Lines changed: 83 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616
import copy
1717
import re
1818
import asyncio
19+
import queue
1920
from pathlib import Path
2021
from typing import Any
2122
from urllib.parse import urlparse
@@ -778,26 +779,89 @@ async def api_chat_stream(request: Request):
778779
_set_memory_context(session, body)
779780
_audit("CHAT_STREAM", f"user={session['user_id']} msg={msg[:80]}", session["user_id"])
780781

781-
async def generate():
782+
class StreamProjection:
783+
"""Pass plain deltas through while holding model control prefixes."""
784+
785+
prefixes = ("Thought:", "Plan:", "TaskList:", "TaskDone:", "Action:", "DefineTool:")
786+
787+
def __init__(self, sink):
788+
self.sink = sink
789+
self.buffer = ""
790+
791+
def __call__(self, event: dict[str, Any]) -> None:
792+
kind = event.get("event")
793+
if kind == "content":
794+
self.buffer += str(event.get("chunk", ""))
795+
candidate = self.buffer.lstrip()
796+
lowered = candidate.lower()
797+
if candidate and (
798+
any(prefix.lower().startswith(lowered) for prefix in self.prefixes)
799+
or any(lowered.startswith(prefix.lower()) for prefix in self.prefixes)
800+
):
801+
return
802+
if self.buffer:
803+
self.sink({"event": "content", "chunk": self.buffer})
804+
self.buffer = ""
805+
return
806+
if kind == "model_complete":
807+
self._flush()
808+
return
809+
self.sink(event)
810+
811+
def _flush(self) -> None:
812+
if not self.buffer:
813+
return
814+
text = self.buffer
815+
self.buffer = ""
816+
if _agent._collect_tool_calls(text) or any(
817+
text.lstrip().lower().startswith(prefix.lower()) for prefix in self.prefixes
818+
):
819+
text = _agent._clean_final_response(text)
820+
if text:
821+
self.sink({"event": "content", "chunk": text})
822+
823+
events: queue.Queue[dict[str, Any]] = queue.Queue()
824+
projection = StreamProjection(events.put)
825+
826+
def run_streaming_turn() -> None:
827+
callback_token = _agent._stream_event_callback.set(projection)
782828
try:
783-
# The legacy agent is synchronous and may perform network and disk
784-
# I/O. Keep it off the ASGI event loop so one chat cannot stall
785-
# health checks or unrelated requests.
786-
reply = await asyncio.to_thread(_run_session_chat, session, msg)
787-
# Send chunks (simulated streaming for non-streaming providers)
788-
chunk_size = 20
789-
for i in range(0, len(reply), chunk_size):
790-
chunk = reply[i:i+chunk_size]
791-
yield f"data: {json.dumps({'chunk': chunk})}\n\n"
792-
await asyncio_sleep(0.01)
793-
yield f"data: {json.dumps({'cost': _cost_summary()})}\n\n"
794-
if session.get("last_memory_receipt"):
795-
yield f"data: {json.dumps({'memory_receipt': session['last_memory_receipt']})}\n\n"
796-
yield "data: [DONE]\n\n"
797-
_emit_chat_completed(session, reply, streamed=True)
798-
_audit("REPLY_STREAM", f"len={len(reply)}", session["user_id"])
799-
except Exception as e:
800-
yield f"data: {json.dumps({'error': str(e)})}\n\n"
829+
reply = _run_session_chat(session, msg)
830+
if str(reply).startswith("[LLM Error]"):
831+
events.put({"event": "error", "error": str(reply)})
832+
else:
833+
events.put({"event": "complete", "reply": reply})
834+
except Exception as exc:
835+
events.put({"event": "error", "error": str(exc)})
836+
finally:
837+
_agent._stream_event_callback.reset(callback_token)
838+
839+
async def generate():
840+
# The synchronous agent runs in a dedicated worker and safely drains if
841+
# a client disconnects; it never blocks the ASGI event loop.
842+
threading.Thread(target=run_streaming_turn, daemon=True).start()
843+
while True:
844+
event = await asyncio.to_thread(events.get)
845+
kind = event.get("event")
846+
if kind == "content":
847+
yield f"data: {json.dumps({'event': 'content', 'chunk': event.get('chunk', '')}, ensure_ascii=False)}\n\n"
848+
elif kind == "tool_receipt":
849+
yield f"data: {json.dumps({'event': 'tool_receipt', 'tool_receipt': event.get('tool_receipt')}, ensure_ascii=False)}\n\n"
850+
elif kind == "tasks":
851+
yield f"data: {json.dumps({'event': 'tasks', 'tasks': event.get('tasks', [])}, ensure_ascii=False)}\n\n"
852+
elif kind == "error":
853+
yield f"data: {json.dumps({'event': 'error', 'error': event.get('error', 'stream failed')}, ensure_ascii=False)}\n\n"
854+
break
855+
elif kind == "complete":
856+
reply = str(event.get("reply", ""))
857+
yield f"data: {json.dumps({'event': 'usage', 'cost': _cost_summary()})}\n\n"
858+
if session.get("last_memory_receipt"):
859+
yield f"data: {json.dumps({'event': 'memory_receipt', 'memory_receipt': session['last_memory_receipt']}, ensure_ascii=False)}\n\n"
860+
yield f"data: {json.dumps({'event': 'completion', 'status': 'completed'})}\n\n"
861+
yield "data: [DONE]\n\n"
862+
_emit_chat_completed(session, reply, streamed=True)
863+
_audit("REPLY_STREAM", f"len={len(reply)}", session["user_id"])
864+
break
801865

802866
return StreamingResponse(generate(), media_type="text/event-stream")
803867

@@ -1660,11 +1724,6 @@ async def pwa_manifest():
16601724
}
16611725

16621726

1663-
# Async sleep helper
1664-
async def asyncio_sleep(seconds: float):
1665-
await asyncio.sleep(seconds)
1666-
1667-
16681727
# ---------------------------------------------------------------------------
16691728
# Plugin system
16701729
# ---------------------------------------------------------------------------

‎tests/test_streaming.py‎

Lines changed: 92 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,92 @@
1+
import asyncio
2+
import json
3+
import tempfile
4+
import threading
5+
import time
6+
import unittest
7+
from pathlib import Path
8+
from unittest.mock import patch
9+
10+
import main
11+
import server
12+
from memory import MemoryBank
13+
from providers import ProviderConfig
14+
15+
16+
class _Request:
17+
def __init__(self, body):
18+
self.body = body
19+
20+
async def json(self):
21+
return self.body
22+
23+
24+
class _DelayedProvider:
25+
def __init__(self):
26+
self.config = ProviderConfig(provider="deepseek", model_simple="deepseek-v4-flash")
27+
self.release = threading.Event()
28+
self.completed = threading.Event()
29+
30+
def chat_stream(self, _messages, _model=None):
31+
yield "FIRST"
32+
self.release.wait(timeout=3)
33+
yield " SECOND"
34+
self.completed.set()
35+
36+
37+
class StreamingEndpointTests(unittest.TestCase):
38+
def test_first_sse_content_arrives_before_provider_stream_completes(self):
39+
session_id = "stream-timing-regression"
40+
provider = _DelayedProvider()
41+
42+
async def exercise():
43+
response = await server.api_chat_stream(_Request({
44+
"message": "hello", "session_id": session_id,
45+
}))
46+
iterator = response.body_iterator.__aiter__()
47+
started = time.monotonic()
48+
first = await asyncio.wait_for(anext(iterator), timeout=1)
49+
first_elapsed = time.monotonic() - started
50+
provider_complete_before_release = provider.completed.is_set()
51+
provider.release.set()
52+
remaining = []
53+
try:
54+
while True:
55+
remaining.append(await asyncio.wait_for(anext(iterator), timeout=2))
56+
except StopAsyncIteration:
57+
pass
58+
return first, first_elapsed, provider_complete_before_release, remaining
59+
60+
with tempfile.TemporaryDirectory(prefix="openkyrozen-stream-") as directory:
61+
memory = MemoryBank(Path(directory) / "state.sqlite3", workspace_id="stream-test")
62+
original_sessions = server._sessions
63+
server._sessions = {}
64+
try:
65+
with (patch.object(server._agent, "memory_bank", memory),
66+
patch.object(server._agent, "llm_provider", provider),
67+
patch.object(server._agent, "DEEPSEEK_MODEL", "deepseek-v4-flash"),
68+
patch.object(server._agent, "_chat_turn", side_effect=lambda message, **_: (
69+
main._call_llm_with_spinner([{"role": "user", "content": message}])
70+
))):
71+
first, elapsed, provider_complete_before_release, remaining = asyncio.run(exercise())
72+
finally:
73+
server._sessions = original_sessions
74+
75+
def text(item):
76+
return item.decode() if isinstance(item, bytes) else item
77+
78+
first_payload = json.loads(text(first).split("data: ", 1)[1].splitlines()[0])
79+
payloads = [json.loads(text(item).split("data: ", 1)[1].splitlines()[0])
80+
for item in remaining
81+
if text(item).startswith("data: ") and text(item) != "data: [DONE]\n\n"]
82+
self.assertLess(elapsed, 1)
83+
self.assertFalse(provider_complete_before_release)
84+
self.assertEqual(first_payload, {"event": "content", "chunk": "FIRST"})
85+
self.assertEqual([item.get("event") for item in payloads[:3]], ["content", "usage", "completion"])
86+
self.assertEqual(payloads[0]["chunk"], " SECOND")
87+
self.assertEqual(sum(text(item) == "data: [DONE]\n\n" for item in remaining), 1)
88+
self.assertTrue(provider.completed.is_set())
89+
90+
91+
if __name__ == "__main__":
92+
unittest.main()

‎tests/test_usage_ledger.py‎

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -121,20 +121,21 @@ def test_reset_requires_explicit_confirmation_and_preserves_installation_ledger(
121121
def test_short_and_mixed_calls_aggregate_before_display_rounding(self):
122122
with tempfile.TemporaryDirectory(prefix="openkyrozen-usage-precision-") as directory:
123123
store = EventStore(Path(directory) / "state.sqlite3")
124+
peak_time = datetime(2026, 9, 7, 2, tzinfo=timezone.utc)
124125
with usage_scope(store=store, workspace_id="project", session_id="short"):
125126
for _ in range(1000):
126127
_track_cost("deepseek", {"prompt_tokens": 0, "completion_tokens": 100},
127-
model="deepseek-chat")
128+
model="deepseek-chat", occurred_at=peak_time)
128129
with usage_scope(store=store, workspace_id="project", session_id="mixed"):
129130
for _ in range(1000):
130131
_track_cost("deepseek", {"prompt_tokens": 100, "completion_tokens": 100},
131-
model="deepseek-chat")
132+
model="deepseek-chat", occurred_at=peak_time)
132133
with usage_scope(store=store, workspace_id="project", session_id="combined-short"):
133134
_track_cost("deepseek", {"prompt_tokens": 0, "completion_tokens": 100_000},
134-
model="deepseek-chat")
135+
model="deepseek-chat", occurred_at=peak_time)
135136
with usage_scope(store=store, workspace_id="project", session_id="combined-mixed"):
136137
_track_cost("deepseek", {"prompt_tokens": 100_000, "completion_tokens": 100_000},
137-
model="deepseek-chat")
138+
model="deepseek-chat", occurred_at=peak_time)
138139

139140
short = store.usage_totals(workspace_id="project", session_id="short", user_id="local")
140141
mixed = store.usage_totals(workspace_id="project", session_id="mixed", user_id="local")

0 commit comments

Comments
 (0)