Trace scheduled, webhook, search, MCP and graph-extraction runs

run_agent_headless records each unattended run under its endpoint; the
scheduler passes its run id and the webhook worker its task id so Logs rows
can find their trace, while the LLM's own request id stays untouched for
quota counts. /api/search and MCP search_docs record their retrieval, and a
graph build records every extraction call under one step.
This commit is contained in:
arc53-machine committed 2026-09-23 17:37:24 +01:00
1 parent bae842d151
commit 5d0992eef8
9 files changed
+344 -15

No files matched your search

+46
View File
@@ -5,6 +5,7 @@ from __future__ import annotations
import logging
from typing import Any, Dict, Iterable, List, Optional
from docsgpt import tracing
from docsgpt.agents.agent_creator import AgentCreator
from docsgpt.agents.tool_executor import ToolExecutor
from docsgpt.api.answer.services.prompt_renderer import (
@@ -69,12 +70,57 @@ def run_agent_headless(
endpoint: str = "headless",
chat_history: Optional[List[Dict[str, Any]]] = None,
conversation_id: Optional[str] = None,
request_id: Optional[str] = None,
) -> Dict[str, Any]:
"""Run an agent with no live client; returns a structured outcome dict.
The run is recorded as one execution trace under ``endpoint`` as its
source. ``request_id`` links that trace to the caller's own record (the
scheduler passes its run id, the webhook worker its task id); it is kept
off the LLM's token-usage rows, whose request ids drive request counts.
Raises:
QuotaExceededError: If the agent owner's usage quota is exhausted.
"""
trace = tracing.start_trace(
source=endpoint,
request_id=request_id,
user_id=_resolve_owner(agent_config),
agent_id=_resolve_agent_id(agent_config),
conversation_id=conversation_id,
)
status = None
with tracing.activate(trace):
try:
outcome = _run_agent_headless(
agent_config,
query,
tool_allowlist=tool_allowlist,
model_id_override=model_id_override,
endpoint=endpoint,
chat_history=chat_history,
conversation_id=conversation_id,
)
if outcome.get("error"):
status = tracing.STATUS_ERROR
return outcome
except BaseException:
status = tracing.STATUS_ERROR
raise
finally:
tracing.flush(trace, status)
def _run_agent_headless(
agent_config: Dict[str, Any],
query: str,
*,
tool_allowlist: Optional[Iterable[str]] = None,
model_id_override: Optional[str] = None,
endpoint: str = "headless",
chat_history: Optional[List[Dict[str, Any]]] = None,
conversation_id: Optional[str] = None,
) -> Dict[str, Any]:
from docsgpt.core.model_utils import (
get_api_key_for_provider,
get_default_model_id,
+2
View File
@@ -277,6 +277,8 @@ def execute_scheduled_run_body(run_id: str, celery_task_id: Optional[str]) -> Di
endpoint="schedule",
conversation_id=schedule.get("origin_conversation_id"),
chat_history=chat_history,
# Links the run's execution trace to its Logs row.
request_id=str(run_id),
)
except SoftTimeLimitExceeded:
timed_out = True
+4 -1
View File
@@ -23,6 +23,7 @@ import logging
import re
from typing import Any, Callable, Dict, List, Optional
from docsgpt import tracing
from docsgpt.core.model_utils import (
get_api_key_for_provider,
get_provider_from_model_id,
@@ -353,7 +354,9 @@ def extract_graph_for_source(
if pool is not None:
# ``map`` yields in submission order, so chunks are still applied in
# the order they were given and a run stays reproducible.
prepared = pool.map(_prepare, items)
# Pool threads don't inherit context; carry the trace in so each
# chunk's extraction LLM call is recorded.
prepared = pool.map(tracing.wrap(_prepare), items)
else:
prepared = (_prepare(item) for item in items)
missed = []
+1 -1
View File
@@ -53,7 +53,7 @@ async def search_docs(query: str, chunks: int = 5) -> list[dict]:
if not api_key:
raise PermissionError("Missing Bearer token")
try:
return await asyncio.to_thread(search, api_key, query, chunks)
return await asyncio.to_thread(search, api_key, query, chunks, source="mcp")
except InvalidAPIKey as exc:
raise PermissionError("Invalid API key") from exc
except SearchFailed:
+27 -2
View File
@@ -10,10 +10,12 @@ from __future__ import annotations
import logging
from typing import Any, Dict, List, Optional
from docsgpt import tracing
from docsgpt.core.settings import settings
from docsgpt.retriever.fanout import fetch_per_source
from docsgpt.storage.db.repositories.agents import AgentsRepository
from docsgpt.storage.db.session import db_readonly
from docsgpt.tracing.retrieval import describe_documents, start_retrieval_span
from docsgpt.vectorstore.vector_creator import VectorCreator
logger = logging.getLogger(__name__)
@@ -219,14 +221,21 @@ def _search_sources(
return results[:chunks]
def search(api_key: str, query: str, chunks: int = 5) -> List[Dict[str, Any]]:
def search(
api_key: str, query: str, chunks: int = 5, *, source: str = "search"
) -> List[Dict[str, Any]]:
"""Resolve an agent by API key and search its sources.
Every search that reaches the sources is recorded as an execution trace
owned by the agent's owner, under ``source``.
Args:
api_key: Agent API key (the opaque string stored on
``agents.key`` in Postgres).
query: Free-text search query.
chunks: Max number of hits to return.
source: Trace source name: ``search`` for ``/api/search``, ``mcp``
for the MCP ``search_docs`` tool.
Returns:
List of hit dicts with ``text``, ``title``, ``source`` keys.
@@ -256,4 +265,20 @@ def search(api_key: str, query: str, chunks: int = 5) -> List[Dict[str, Any]]:
if not source_ids:
return []
return _search_sources(query, source_ids, chunks)
trace = tracing.start_trace(
source=source,
user_id=agent.get("user_id"),
agent_id=str(agent.get("id")) if agent.get("id") else None,
)
with tracing.activate(trace):
try:
with start_retrieval_span(
f"retrieval {source}",
sources=source_ids,
**{"docsgpt.top_k": chunks},
) as span:
results = _search_sources(query, source_ids, chunks)
describe_documents(span, results, query=query)
return results
finally:
tracing.flush(trace)
+31 -8
View File
@@ -16,6 +16,7 @@ from urllib.parse import urljoin, urlsplit
import requests
from docsgpt import tracing
from docsgpt.core.settings import settings
from docsgpt.events.publisher import publish_user_event
from docsgpt.parser.chunking_creator import ChunkerCreator
@@ -2129,6 +2130,7 @@ def agent_webhook_worker(self, agent_id, payload):
input_data,
tool_allowlist=_webhook_tool_allowlist(agent_config),
endpoint="webhook",
request_id=getattr(getattr(self, "request", None), "id", None),
)
result = {
"answer": outcome.get("answer", ""),
@@ -2942,16 +2944,36 @@ def extract_graph_worker(self, source_id, user):
},
)
trace = tracing.start_trace(
source="graph_extraction",
name=f"graph_extraction {source.get('name') or source_id}",
request_id=getattr(self.request, "id", None),
user_id=user,
)
try:
summary = extract_graph_for_source(
source_id,
user,
chunks,
config=cfg,
request_id=getattr(self.request, "id", None),
progress_cb=_progress,
)
with tracing.activate(trace), tracing.span(
tracing.KIND_STEP,
"graph_extraction",
attributes={"docsgpt.source_id": source_id, "docsgpt.chunk_count": total},
) as span:
summary = extract_graph_for_source(
source_id,
user,
chunks,
config=cfg,
request_id=getattr(self.request, "id", None),
progress_cb=_progress,
)
if isinstance(summary, dict):
span.set(
**{
"docsgpt.graph.nodes": summary.get("nodes"),
"docsgpt.graph.edges": summary.get("edges"),
"docsgpt.graph.chunks_processed": summary.get("chunks_processed"),
}
)
except Exception as e:
tracing.flush(trace, tracing.STATUS_ERROR)
_publish_graph_event(
user,
source_id,
@@ -2960,6 +2982,7 @@ def extract_graph_worker(self, source_id, user):
)
raise
tracing.flush(trace)
_publish_graph_event(
user,
source_id,
+3 -3
View File
@@ -99,7 +99,7 @@ class TestSearchDocsTool:
):
out = await search_docs(query="q", chunks=7)
assert out == hits
mock_search.assert_called_once_with("the-key", "q", 7)
mock_search.assert_called_once_with("the-key", "q", 7, source="mcp")
@pytest.mark.asyncio
async def test_default_chunks_is_5(self):
@@ -115,7 +115,7 @@ class TestSearchDocsTool:
) as mock_search,
):
await search_docs(query="q")
mock_search.assert_called_once_with("k", "q", 5)
mock_search.assert_called_once_with("k", "q", 5, source="mcp")
@pytest.mark.asyncio
async def test_bearer_scheme_case_insensitive(self):
@@ -131,4 +131,4 @@ class TestSearchDocsTool:
) as mock_search,
):
await search_docs(query="q")
mock_search.assert_called_once_with("lowercase-scheme", "q", 5)
mock_search.assert_called_once_with("lowercase-scheme", "q", 5, source="mcp")
+180
View File
@@ -0,0 +1,180 @@
"""Non-chat entry points each record and write one execution trace."""
from __future__ import annotations
from contextlib import contextmanager
from unittest.mock import MagicMock, patch
import pytest
from docsgpt import tracing
from docsgpt.core.settings import settings
@pytest.fixture(autouse=True)
def _tracing_on(monkeypatch):
monkeypatch.setattr(settings, "TRACES_ENABLED", True)
monkeypatch.setattr(settings, "TRACES_OTEL_EXPORT", False)
@pytest.fixture()
def flushed():
"""Capture flushed traces instead of writing them."""
captured = []
def _fake_flush(trace, status=None):
if trace is None or trace.flushed:
return
trace.flushed = True
trace.finish(status)
captured.append(trace)
with patch("docsgpt.tracing.flush", side_effect=_fake_flush):
yield captured
def _headless(events, monkeypatch, **kwargs):
from docsgpt.agents import headless_runner as hr
agent = MagicMock(name="agent")
def _gen(query):
with tracing.span(tracing.KIND_AGENT, "invoke_agent Fake"):
yield from events
agent.gen.side_effect = _gen
agent.llm.token_usage = {"prompt_tokens": 1, "generated_tokens": 1}
retriever = MagicMock(name="retriever")
def _search(query):
with tracing.span(tracing.KIND_RETRIEVAL, "retrieval"):
return []
retriever.search.side_effect = _search
tool_executor = MagicMock(name="tool_executor")
tool_executor.headless_denials = []
monkeypatch.setattr(hr, "get_prompt", lambda _pid: "system prompt")
monkeypatch.setattr(
hr.RetrieverCreator, "create_retriever", classmethod(lambda cls, *a, **kw: retriever)
)
monkeypatch.setattr(hr, "ToolExecutor", lambda *a, **kw: tool_executor)
monkeypatch.setattr(
hr.AgentCreator, "create_agent", classmethod(lambda cls, *a, **kw: agent)
)
with patch("docsgpt.core.model_utils.validate_model_id", return_value=True), \
patch("docsgpt.core.model_utils.get_default_model_id", return_value="m"), \
patch("docsgpt.core.model_utils.get_provider_from_model_id", return_value="openai"), \
patch("docsgpt.core.model_utils.get_api_key_for_provider", return_value="k"), \
patch("docsgpt.utils.calculate_doc_token_budget", return_value=1000):
return hr.run_agent_headless(
{"user_id": "owner-1", "id": "11111111-1111-1111-1111-111111111111"},
"do the thing",
**kwargs,
)
@pytest.mark.unit
class TestHeadless:
def test_scheduled_run_is_traced_with_its_run_id(self, monkeypatch, flushed):
_headless([{"answer": "done"}], monkeypatch, endpoint="schedule", request_id="run-1")
(trace,) = flushed
assert trace.source == "schedule"
assert trace.request_id == "run-1"
assert trace.user_id == "owner-1"
assert trace.status == "ok"
assert [s.kind for s in trace.spans] == ["retrieval", "agent"]
def test_stream_error_marks_trace_error(self, monkeypatch, flushed):
outcome = _headless([{"type": "error", "error": "boom"}], monkeypatch, endpoint="webhook")
assert outcome["error_type"] == "stream_error"
assert flushed[0].status == "error"
def test_raised_error_still_flushes(self, monkeypatch, flushed):
from docsgpt.agents import headless_runner as hr
monkeypatch.setattr(hr, "_run_agent_headless", MagicMock(side_effect=RuntimeError("x")))
with pytest.raises(RuntimeError):
hr.run_agent_headless({"user_id": "u"}, "q")
assert flushed[0].status == "error"
def test_trace_request_id_stays_off_llm_usage_rows(self, monkeypatch, flushed):
"""The headless LLM's own request id is untouched (quota counts depend on it)."""
from docsgpt.agents import headless_runner as hr
created = {}
def _create(cls, *a, **kw):
agent = MagicMock()
agent.gen.return_value = iter([{"answer": "x"}])
agent.llm.token_usage = {}
agent.llm._request_id = None
created["agent"] = agent
return agent
monkeypatch.setattr(hr.AgentCreator, "create_agent", classmethod(_create))
monkeypatch.setattr(hr, "get_prompt", lambda _pid: "p")
monkeypatch.setattr(
hr.RetrieverCreator, "create_retriever",
classmethod(lambda cls, *a, **kw: MagicMock(search=MagicMock(return_value=[]))),
)
monkeypatch.setattr(hr, "ToolExecutor", lambda *a, **kw: MagicMock(headless_denials=[]))
with patch("docsgpt.core.model_utils.validate_model_id", return_value=True), \
patch("docsgpt.core.model_utils.get_default_model_id", return_value="m"), \
patch("docsgpt.core.model_utils.get_provider_from_model_id", return_value="openai"), \
patch("docsgpt.core.model_utils.get_api_key_for_provider", return_value="k"), \
patch("docsgpt.utils.calculate_doc_token_budget", return_value=1000):
hr.run_agent_headless({"user_id": "u"}, "q", endpoint="schedule", request_id="run-9")
assert created["agent"].llm._request_id is None
@contextmanager
def _search_env(agent):
repo = MagicMock()
repo.find_by_key.return_value = agent
@contextmanager
def _conn():
yield MagicMock()
store = MagicMock()
store.search.return_value = [{"text": "hit", "metadata": {"title": "T", "source": "s"}}]
with patch("docsgpt.api.user.team_sharing.can_access", return_value=True), \
patch("docsgpt.services.search_service.db_readonly", _conn), \
patch("docsgpt.services.search_service.AgentsRepository", return_value=repo), \
patch(
"docsgpt.services.search_service.VectorCreator.create_vectorstore",
return_value=store,
):
yield
@pytest.mark.unit
class TestSearch:
def test_search_is_traced(self, flushed):
from docsgpt.services.search_service import search
agent = {"id": "a-1", "source_id": "src-1", "extra_source_ids": [], "user_id": "owner"}
with _search_env(agent):
results = search("k", "what is x", 3)
assert results
(trace,) = flushed
assert trace.source == "search"
assert trace.user_id == "owner"
retrieval = trace.spans[0]
assert retrieval.kind == tracing.KIND_RETRIEVAL
assert retrieval.attributes["docsgpt.chunk_count"] == 1
def test_mcp_source_name(self, flushed):
from docsgpt.services.search_service import search
agent = {"id": "a-1", "source_id": "src-1", "extra_source_ids": [], "user_id": "owner"}
with _search_env(agent):
search("k", "q", 3, source="mcp")
assert flushed[0].source == "mcp"
def test_no_sources_records_nothing(self, flushed):
from docsgpt.services.search_service import search
with _search_env({"id": "a-1", "source_id": None, "user_id": "owner"}):
assert search("k", "q", 3) == []
assert flushed == []
+50
View File
@@ -170,3 +170,53 @@ class TestExtractGraphWorker:
worker.extract_graph_worker(task_self, source_id, "alice")
assert "graph.extract.failed" in [e[0] for e in events]
@pytest.mark.unit
class TestExtractGraphTrace:
"""A graph build is one execution trace holding every extraction call."""
def _run(self, pg_conn, monkeypatch, task_self, extract):
from docsgpt import tracing, worker
from docsgpt.core.settings import settings
monkeypatch.setattr(settings, "TRACES_ENABLED", True)
monkeypatch.setattr(settings, "TRACES_OTEL_EXPORT", False)
source_id = _seed_source(pg_conn)
_patch_store(monkeypatch, [{"doc_id": "c1", "text": "alpha"}])
monkeypatch.setattr("docsgpt.graphrag.graphrag_available", lambda: True)
monkeypatch.setattr("docsgpt.graphrag.extraction.extract_graph_for_source", extract)
flushed = []
def _flush(trace, status=None):
trace.flushed = True
trace.finish(status)
flushed.append(trace)
monkeypatch.setattr(tracing, "flush", _flush)
return worker, source_id, flushed
def test_successful_build_is_traced(self, pg_conn, patch_worker_db, task_self, monkeypatch):
from docsgpt import tracing
def _extract(*_a, **_kw):
with tracing.span(tracing.KIND_LLM, "chat m"):
pass
return {"nodes": 3, "edges": 2, "chunks_processed": 1}
worker, source_id, flushed = self._run(pg_conn, monkeypatch, task_self, _extract)
worker.extract_graph_worker(task_self, source_id, "alice")
(trace,) = flushed
assert trace.source == "graph_extraction"
assert trace.user_id == "alice"
step, llm = trace.spans
assert step.attributes["docsgpt.graph.nodes"] == 3
assert llm.parent_id == step.id
def test_failed_build_is_an_error_trace(self, pg_conn, patch_worker_db, task_self, monkeypatch):
worker, source_id, flushed = self._run(
pg_conn, monkeypatch, task_self, MagicMock(side_effect=RuntimeError("llm down"))
)
with pytest.raises(RuntimeError):
worker.extract_graph_worker(task_self, source_id, "alice")
assert flushed[0].status == "error"