From 5d0992eef85fdf13d1d9ccdddf72b48d70d79bc4 Mon Sep 17 00:00:00 2001 From: arc53-machine <232052973+arc53-machine@users.noreply.github.com> Date: Wed, 23 Sep 2026 17:37:24 +0100 Subject: [PATCH] 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. --- docsgpt/agents/headless_runner.py | 46 +++++++ docsgpt/api/user/scheduler_worker.py | 2 + docsgpt/graphrag/extraction.py | 5 +- docsgpt/mcp_server.py | 2 +- docsgpt/services/search_service.py | 29 ++++- docsgpt/worker.py | 39 ++++-- tests/services/test_mcp_server.py | 6 +- tests/tracing/test_entry_points.py | 180 +++++++++++++++++++++++++++ tests/worker/test_extract_graph.py | 50 ++++++++ 9 files changed, 344 insertions(+), 15 deletions(-) create mode 100644 tests/tracing/test_entry_points.py diff --git a/docsgpt/agents/headless_runner.py b/docsgpt/agents/headless_runner.py index c2a4d713..e50e2478 100644 --- a/docsgpt/agents/headless_runner.py +++ b/docsgpt/agents/headless_runner.py @@ -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, diff --git a/docsgpt/api/user/scheduler_worker.py b/docsgpt/api/user/scheduler_worker.py index 3e085f0e..15c73f76 100644 --- a/docsgpt/api/user/scheduler_worker.py +++ b/docsgpt/api/user/scheduler_worker.py @@ -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 diff --git a/docsgpt/graphrag/extraction.py b/docsgpt/graphrag/extraction.py index daa8f910..80efd1fa 100644 --- a/docsgpt/graphrag/extraction.py +++ b/docsgpt/graphrag/extraction.py @@ -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 = [] diff --git a/docsgpt/mcp_server.py b/docsgpt/mcp_server.py index f67120b2..9e8c7fbf 100644 --- a/docsgpt/mcp_server.py +++ b/docsgpt/mcp_server.py @@ -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: diff --git a/docsgpt/services/search_service.py b/docsgpt/services/search_service.py index b67a1aca..68756c91 100644 --- a/docsgpt/services/search_service.py +++ b/docsgpt/services/search_service.py @@ -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) diff --git a/docsgpt/worker.py b/docsgpt/worker.py index 9fbf57f3..f3311859 100755 --- a/docsgpt/worker.py +++ b/docsgpt/worker.py @@ -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, diff --git a/tests/services/test_mcp_server.py b/tests/services/test_mcp_server.py index d8441db5..b048eb77 100644 --- a/tests/services/test_mcp_server.py +++ b/tests/services/test_mcp_server.py @@ -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") diff --git a/tests/tracing/test_entry_points.py b/tests/tracing/test_entry_points.py new file mode 100644 index 00000000..221b0c29 --- /dev/null +++ b/tests/tracing/test_entry_points.py @@ -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 == [] diff --git a/tests/worker/test_extract_graph.py b/tests/worker/test_extract_graph.py index 7e040fbe..d6dd5701 100644 --- a/tests/worker/test_extract_graph.py +++ b/tests/worker/test_extract_graph.py @@ -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"