diff --git a/docsgpt/agents/base.py b/docsgpt/agents/base.py index 4f620236..22ba60d5 100644 --- a/docsgpt/agents/base.py +++ b/docsgpt/agents/base.py @@ -30,7 +30,13 @@ from docsgpt.guardrails.stream import StreamingOutputGuard from docsgpt.guardrails.types import Action, Stage, resolve_tool_result from docsgpt.llm.handlers.handler_creator import LLMHandlerCreator from docsgpt.llm.llm_creator import LLMCreator -from docsgpt.logging import build_stack_data, log_activity, LogContext, start_agent_span +from docsgpt.logging import ( + agent_log_context, + build_stack_data, + log_activity, + LogContext, + start_agent_span, +) logger = logging.getLogger(__name__) @@ -494,7 +500,8 @@ class BaseAgent(ABC): hands back to the LLM to continue the conversation. Unlike :meth:`gen` this is not wrapped by ``@log_activity``, so the - continuation's ``invoke_agent`` trace span is opened here. + continuation's ``invoke_agent`` trace span is opened here, and the + user/agent/endpoint log context is bound here too. Args: messages: The saved messages array from the pause point. @@ -502,7 +509,7 @@ class BaseAgent(ABC): pending_tool_calls: The pending tool call descriptors from the pause. tool_actions: Client-provided actions resolving the pending calls. """ - with start_agent_span(self, continuation=True): + with agent_log_context(self), start_agent_span(self, continuation=True): yield from self._gen_continuation_inner( messages, tools_dict, pending_tool_calls, tool_actions, reasoning_content ) diff --git a/docsgpt/logging.py b/docsgpt/logging.py index 5440631f..2ece4e1f 100644 --- a/docsgpt/logging.py +++ b/docsgpt/logging.py @@ -5,7 +5,8 @@ import time import logging import uuid -from typing import Any, Callable, Dict, Generator, List, Optional +from contextlib import contextmanager +from typing import Any, Callable, Dict, Generator, Iterator, List, Optional from docsgpt import tracing from docsgpt.core import log_context @@ -87,22 +88,52 @@ def build_stack_data( return data +def _agent_log_keys(agent: Any, data: Optional[Dict] = None) -> Dict[str, Any]: + """Return the log-context keys identifying an agent run.""" + if data is None: + data = build_stack_data(agent) + return { + "user_id": data.get("user", "local"), + "agent_id": getattr(agent, "agent_id", None), + "conversation_id": getattr(agent, "conversation_id", None), + "endpoint": data.get("endpoint", ""), + "model": getattr(agent, "gpt_model", None) or getattr(agent, "model", None), + } + + +@contextmanager +def agent_log_context(agent: Any) -> Iterator[None]: + """Bind an agent's identity to the log context without opening an activity. + + Any enclosing ``activity_id`` is kept. + + Args: + agent: The agent whose identity the log lines should carry. + + Yields: + None, with the context bound until the block exits. + """ + token = log_context.bind(**_agent_log_keys(agent)) + try: + yield + finally: + log_context.reset(token) + + def log_activity() -> Callable: def decorator(func: Callable) -> Callable: @functools.wraps(func) def wrapper(*args: Any, **kwargs: Any) -> Any: activity_id = str(uuid.uuid4()) data = build_stack_data(args[0]) - endpoint = data.get("endpoint", "") - user = data.get("user", "local") + keys = _agent_log_keys(args[0], data) + endpoint = keys["endpoint"] + user = keys["user_id"] api_key = data.get("user_api_key", "") query = kwargs.get("query", getattr(args[0], "query", "")) - agent_id = getattr(args[0], "agent_id", None) or kwargs.get("agent_id") - conversation_id = ( - kwargs.get("conversation_id") - or getattr(args[0], "conversation_id", None) - ) - model = getattr(args[0], "gpt_model", None) or getattr(args[0], "model", None) + agent_id = keys["agent_id"] or kwargs.get("agent_id") + conversation_id = kwargs.get("conversation_id") or keys["conversation_id"] + model = keys["model"] # Capture the surrounding activity_id before overlaying ours, # so nested activities record the parent → child link. diff --git a/tests/agents/test_continuation_log_context.py b/tests/agents/test_continuation_log_context.py new file mode 100644 index 00000000..28a3d21a --- /dev/null +++ b/tests/agents/test_continuation_log_context.py @@ -0,0 +1,97 @@ +"""A tool continuation stamps its log lines like the turn it resumes.""" + +from unittest.mock import Mock + +import pytest + +from docsgpt.core import log_context + + +def _agent(seen): + from docsgpt.agents.classic_agent import ClassicAgent + + def _model_call(*_args, **_kwargs): + seen.append(dict(log_context.snapshot())) + return iter(["Answer"]) + + llm = Mock() + llm._supports_tools = True + llm._supports_structured_output = Mock(return_value=False) + llm.__class__.__name__ = "MockLLM" + llm.gen_stream = Mock(side_effect=_model_call) + llm.gen = Mock(side_effect=_model_call) + + handler = Mock() + handler.process_message_flow = Mock(return_value=iter([])) + handler.create_tool_message = Mock(return_value={"role": "tool", "tool_call_id": "c1", "content": "r"}) + + executor = Mock() + executor.tool_calls = [] + executor.prepare_tools_for_llm = Mock(return_value=[]) + executor.get_truncated_tool_calls = Mock(return_value=[]) + + def _execute(_tools, _call, _llm_class): + yield {"type": "tool_call", "data": {"status": "pending"}} + return ("result", "c1") + + executor.execute = Mock(side_effect=_execute) + return ClassicAgent( + endpoint="stream", + llm_name="openai", + model_id="gpt-4", + api_key="test", + agent_id="aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa", + decoded_token={"sub": "user-cont"}, + llm=llm, + llm_handler=handler, + tool_executor=executor, + ) + + +def _resume(agent): + pending = [ + { + "call_id": "c1", + "name": "search_0", + "tool_name": "search", + "tool_id": "0", + "action_name": "search", + "arguments": {"q": "x"}, + "pause_type": "requires_client_execution", + "thought_signature": None, + } + ] + actions = [{"call_id": "c1", "decision": "approved"}] + return list(agent.gen_continuation([{"role": "system", "content": "s"}], {"0": {"name": "search"}}, pending, actions)) + + +@pytest.mark.unit +def test_model_calls_in_a_continuation_carry_the_turns_identity(): + seen = [] + _resume(_agent(seen)) + + assert seen, "the continuation should hand back to the model" + assert seen[0]["user_id"] == "user-cont" + assert seen[0]["agent_id"] == "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa" + assert seen[0]["endpoint"] == "stream" + assert "activity_id" not in seen[0], "a continuation is not an activity of its own" + + +@pytest.mark.unit +def test_the_binding_does_not_outlive_the_continuation(): + _resume(_agent([])) + + assert "user_id" not in log_context.snapshot() + + +@pytest.mark.unit +def test_an_enclosing_activity_keeps_its_id(): + seen = [] + token = log_context.bind(activity_id="act-1") + try: + _resume(_agent(seen)) + finally: + log_context.reset(token) + + assert seen[0]["activity_id"] == "act-1" + assert seen[0]["user_id"] == "user-cont"