fix: wire TraceCollector into server chat endpoints (#513)

This commit is contained in:
Robby Manihani
2026-06-08 18:04:14 -07:00
committed by GitHub
parent b30bf43557
commit 726445433b
7 changed files with 305 additions and 13 deletions
+48
View File
@@ -554,6 +554,30 @@ async def prometheus_metrics(request: Request):
websocket_router = APIRouter(tags=["websocket"])
def _record_ws_trace(
trace_store,
*,
query: str,
result: str,
model: str,
started_at: float,
ended_at: float,
) -> None:
"""Record a trace for a completed WebSocket chat (best-effort)."""
if trace_store is None or not result:
return
from openjarvis.traces.collector import record_response_trace
record_response_trace(
trace_store,
query=query,
result=result,
model=model,
started_at=started_at,
ended_at=ended_at,
)
@websocket_router.websocket("/v1/chat/stream")
async def websocket_chat_stream(websocket: WebSocket):
"""Stream chat responses over a WebSocket connection.
@@ -608,6 +632,14 @@ async def websocket_chat_stream(websocket: WebSocket):
messages = [{"role": "user", "content": message}]
# This WS path streams straight from the engine (no agent /
# TraceCollector), so record the interaction directly once it
# finishes — otherwise WebSocket chats never reach traces.db.
import time as _time
trace_store = getattr(websocket.app.state, "trace_store", None)
_ws_started_at = _time.time()
try:
# Prefer streaming if the engine supports it
stream_fn = getattr(engine, "stream", None)
@@ -651,6 +683,14 @@ async def websocket_chat_stream(websocket: WebSocket):
await websocket.send_json(
{"type": "done", "content": full_content},
)
_record_ws_trace(
trace_store,
query=message,
result=full_content,
model=model,
started_at=_ws_started_at,
ended_at=_time.time(),
)
else:
# No stream method — single-shot generate
result = engine.generate(messages, model=model)
@@ -668,6 +708,14 @@ async def websocket_chat_stream(websocket: WebSocket):
await websocket.send_json(
{"type": "done", "content": content},
)
_record_ws_trace(
trace_store,
query=message,
result=content,
model=model,
started_at=_ws_started_at,
ended_at=_time.time(),
)
except WebSocketDisconnect:
raise
except Exception as exc:
+11 -6
View File
@@ -229,7 +229,16 @@ def create_app(
# AuthMiddleware never sees WS upgrade requests). Empty = auth disabled.
app.state.api_key = api_key
# Wire up trace store if traces are enabled
# Wire up trace store if traces are enabled.
#
# We deliberately do NOT subscribe the trace store to the bus. The chat
# endpoints persist through a TraceCollector that calls store.save()
# directly (mirroring system/orchestrator.py), and the collector ALSO
# publishes TRACE_COMPLETE. A store subscribed to that same bus would
# therefore save every agent trace twice — the second INSERT hitting the
# UNIQUE constraint on trace_id (a 500 on every completion). Keeping the
# collector the single writer is what makes the dual code path safe; only
# the telemetry store is bus-subscribed (see system/builder.py).
app.state.trace_store = None
try:
from openjarvis.core.config import load_config
@@ -237,11 +246,7 @@ def create_app(
cfg = config if config is not None else load_config()
if cfg.traces.enabled:
_trace_store = TraceStore(db_path=cfg.traces.db_path)
app.state.trace_store = _trace_store
_bus = getattr(app.state, "bus", None)
if _bus is not None:
_trace_store.subscribe_to_bus(_bus)
app.state.trace_store = TraceStore(db_path=cfg.traces.db_path)
except Exception:
pass # traces are optional; don't block server startup
+72 -5
View File
@@ -151,7 +151,13 @@ async def chat_completions(request_body: ChatCompletionRequest, request: Request
return await _handle_stream_tools(
engine, model, request_body, complexity_info
)
return await _handle_stream(engine, model, request_body, complexity_info)
return await _handle_stream(
engine,
model,
request_body,
complexity_info,
trace_store=getattr(request.app.state, "trace_store", None),
)
# Non-streaming: use agent if available, otherwise direct engine call.
#
@@ -170,7 +176,14 @@ async def chat_completions(request_body: ChatCompletionRequest, request: Request
# the agent to execute them), add an explicit opt-in header rather
# than removing this guard — silent re-routing is what produced #414.
if agent is not None and not request_body.tools:
return _handle_agent(agent, model, request_body, complexity_info)
return _handle_agent(
agent,
model,
request_body,
complexity_info,
trace_store=getattr(request.app.state, "trace_store", None),
bus=getattr(request.app.state, "bus", None),
)
bus = getattr(request.app.state, "bus", None)
return _handle_direct(
@@ -288,8 +301,19 @@ def _handle_agent(
model: str,
req: ChatCompletionRequest,
complexity_info=None,
*,
trace_store=None,
bus=None,
) -> ChatCompletionResponse:
"""Run through agent."""
"""Run through agent.
When *trace_store* is set, the agent run is wrapped in a
``TraceCollector`` (mirroring ``system/orchestrator.py``) so every
completion records a ``Trace`` to ``traces.db``. Previously this endpoint
called ``agent.run()`` raw, so the server never produced traces:
``traces.db`` stayed empty and spec_search's cold-start gate
(``check_readiness``, min 20 traces) could never open.
"""
from openjarvis.agents._stubs import AgentContext
# Build context from prior messages
@@ -307,7 +331,13 @@ def _handle_agent(
if model:
agent._model = model
try:
result = agent.run(input_text, context=ctx)
if trace_store is not None:
from openjarvis.traces.collector import TraceCollector
collector = TraceCollector(agent, store=trace_store, bus=bus)
result = collector.run(input_text, context=ctx)
else:
result = agent.run(input_text, context=ctx)
finally:
agent._model = original_model
@@ -459,8 +489,19 @@ async def _handle_stream(
model: str,
req: ChatCompletionRequest,
complexity_info=None,
*,
trace_store=None,
):
"""Stream response using SSE format."""
"""Stream response using SSE format.
This path streams straight from the engine, bypassing the agent /
``TraceCollector``. When *trace_store* is set we accumulate the streamed
tokens and record a minimal ``Trace`` once the stream completes
successfully — otherwise streamed chats (the desktop GUI's main path)
would never populate ``traces.db``.
"""
import time
from openjarvis.server.cloud_router import (
is_cloud_model,
stream_cloud,
@@ -470,11 +511,20 @@ async def _handle_stream(
messages = _to_messages(req.messages)
chunk_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
# Last user message — recorded as the trace query.
query_text = ""
for _m in reversed(req.messages):
if _m.role == "user" and _m.content:
query_text = _m.content
break
# Route directly to the right backend — bypasses engine routing entirely
# so broken MultiEngine state can never misdirect requests.
use_cloud = is_cloud_model(model)
async def generate():
started_at = time.time()
full_content = ""
# Send role chunk first
first_chunk = ChatCompletionChunk(
id=chunk_id,
@@ -527,6 +577,7 @@ async def _handle_stream(
max_tokens=req.max_tokens,
)
async for token in token_iter:
full_content += token
chunk = ChatCompletionChunk(
id=chunk_id,
model=model,
@@ -563,6 +614,22 @@ async def _handle_stream(
yield "data: [DONE]\n\n"
return
# Record a trace for the completed stream (best-effort; never breaks
# the response). Mirrors the agent path so streamed chats also
# populate traces.db.
if trace_store is not None and full_content:
from openjarvis.traces.collector import record_response_trace
record_response_trace(
trace_store,
query=query_text,
result=full_content,
model=model,
engine="cloud" if use_cloud else "ollama",
started_at=started_at,
ended_at=time.time(),
)
# Send finish chunk with usage data if available
import json as _json
+56 -1
View File
@@ -221,4 +221,59 @@ class TraceCollector:
)
__all__ = ["TraceCollector"]
def record_response_trace(
store: Optional[TraceStore],
*,
query: str,
result: str,
model: str = "",
engine: str = "",
agent: str = "server",
started_at: float,
ended_at: float,
) -> Optional[Trace]:
"""Persist a minimal single-step ``Trace`` for a non-agent response.
The streaming SSE and WebSocket chat paths stream straight from the
engine, bypassing the agent (and therefore ``TraceCollector``). They call
this so those interactions still land in ``traces.db`` — otherwise streamed
chats, which are the desktop GUI's main path, would never produce traces.
Best-effort: returns the saved ``Trace`` or ``None`` (when *store* is
``None`` or persistence raised), and never propagates an exception into the
caller's response path.
"""
if store is None:
return None
try:
duration = max(0.0, ended_at - started_at)
trace = Trace(
query=query,
agent=agent,
model=model,
engine=engine,
result=result,
started_at=started_at,
ended_at=ended_at,
steps=[
TraceStep(
step_type=StepType.RESPOND,
timestamp=ended_at,
duration_seconds=duration,
output={"content": result},
)
],
)
trace.total_latency_seconds = duration
store.save(trace)
return trace
except Exception:
import logging
logging.getLogger("openjarvis.traces").debug(
"record_response_trace failed", exc_info=True
)
return None
__all__ = ["TraceCollector", "record_response_trace"]
+9 -1
View File
@@ -106,7 +106,15 @@ class TraceStore:
self._conn.commit()
def save(self, trace: Trace) -> None:
"""Persist a complete trace with all its steps."""
"""Persist a complete trace with all its steps.
``trace_id`` is a primary key: saving a second, different trace under
an existing id raises ``sqlite3.IntegrityError`` (the external-corpus
adapter relies on this to surface duplicate record ids). The server
avoids re-saving the same trace by keeping the ``TraceCollector`` the
single writer — see ``server/app.py`` — rather than swallowing
collisions here.
"""
self._conn.execute(
_INSERT_TRACE,
(
+34
View File
@@ -0,0 +1,34 @@
"""Shared fixtures for server route tests.
Server tests build apps via ``create_app``, which (with traces enabled by
default) wires a ``TraceStore`` at the real ``~/.openjarvis/traces.db``. Now
that the chat endpoints actually *write* traces, an unguarded run would
pollute the developer's real trace DB and make tests non-hermetic. This
autouse fixture redirects the traces DB to a per-test temp path.
"""
from __future__ import annotations
import pytest
@pytest.fixture(autouse=True)
def _isolate_traces_db(tmp_path, monkeypatch):
"""Point ``config.traces.db_path`` at a temp file for every server test.
``load_config`` returns a fresh ``JarvisConfig`` per call (no caching), so
wrapping it to rewrite ``traces.db_path`` only affects calls made during
the test — there is no global leak.
"""
from openjarvis.core import config as _config
real_load_config = _config.load_config
db_path = str(tmp_path / "traces.db")
def _patched_load_config(*args, **kwargs):
cfg = real_load_config(*args, **kwargs)
cfg.traces.db_path = db_path
return cfg
monkeypatch.setattr(_config, "load_config", _patched_load_config)
return db_path
+75
View File
@@ -559,3 +559,78 @@ class TestCreateApp:
engine = _make_engine()
app = create_app(engine, "test-model")
assert app.state.agent is None
# ---------------------------------------------------------------------------
# Trace recording — regression coverage for the empty-traces.db bug
# (TraceCollector was never wired into the server chat endpoints).
# ---------------------------------------------------------------------------
class TestTraceRecording:
def test_agent_completion_creates_trace(self):
"""A non-streaming agent completion records exactly one trace.
The collector is the single writer: it saves directly and also
publishes TRACE_COMPLETE, but the store is NOT subscribed to the bus
(see server/app.py), so the trace is persisted exactly once. If the
store were re-subscribed, the collector's second save would raise
IntegrityError on the trace_id primary key and the request would 500 —
so asserting 200 + count == 1 guards that double-save regression.
"""
from openjarvis.core.events import EventBus
engine = _make_engine()
agent = _make_agent(content="traced reply")
app = create_app(
engine,
"test-model",
agent=agent,
bus=EventBus(record_history=False),
)
store = app.state.trace_store
assert store is not None, "traces enabled by default → store should exist"
assert store.count() == 0
client = TestClient(app)
resp = client.post(
"/v1/chat/completions",
json={
"model": "test-model",
"messages": [{"role": "user", "content": "What is 2+2?"}],
},
)
assert resp.status_code == 200
assert resp.json()["choices"][0]["message"]["content"] == "traced reply"
assert store.count() == 1 # not 2 — double-save must be idempotent
trace = store.list_traces(limit=1)[0]
assert trace.query == "What is 2+2?"
assert trace.result == "traced reply"
def test_streaming_completion_creates_trace(self):
"""A streamed completion (no agent) records the assembled response."""
engine = _make_engine()
app = create_app(engine, "test-model")
store = app.state.trace_store
assert store is not None
assert store.count() == 0
client = TestClient(app)
resp = client.post(
"/v1/chat/completions",
json={
"model": "test-model",
"messages": [{"role": "user", "content": "stream please"}],
"stream": True,
},
)
assert resp.status_code == 200
# Drain the SSE body so the streaming generator runs to completion.
assert "data:" in resp.text
assert store.count() == 1
trace = store.list_traces(limit=1)[0]
assert trace.query == "stream please"
# _make_engine streams "Hello", " ", "world".
assert trace.result == "Hello world"