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"