From e5ea8548aca46952f65e1055c336b14e663ba6ab Mon Sep 17 00:00:00 2001 From: raychen <815315825@qq.com> Date: Tue, 8 Sep 2026 17:38:12 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=94=AF=E6=8C=81=E4=B8=8A=E6=8A=A5?= =?UTF-8?q?=E6=A8=A1=E5=9E=8B=E8=BF=94=E5=9B=9E=E7=9A=84=E9=A2=9D=E5=A4=96?= =?UTF-8?q?=E5=AD=97=E6=AE=B5?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- examples/fastapi_server/_runner_manager.py | 2 +- .../llmagent_with_model_retry/agent/agent.py | 16 ++++ .../llmagent_with_model_retry/run_agent.py | 4 + tests/models/test_openai_model_ext.py | 93 +++++++++++++++++++ tests/telemetry/test_trace.py | 20 ++++ trpc_agent_sdk/models/__init__.py | 2 + trpc_agent_sdk/models/_constants.py | 3 + trpc_agent_sdk/models/_openai_model.py | 70 ++++++++++++-- trpc_agent_sdk/telemetry/_trace.py | 10 ++ 9 files changed, 212 insertions(+), 8 deletions(-) diff --git a/examples/fastapi_server/_runner_manager.py b/examples/fastapi_server/_runner_manager.py index ebbe1be49..62d36a741 100644 --- a/examples/fastapi_server/_runner_manager.py +++ b/examples/fastapi_server/_runner_manager.py @@ -88,7 +88,7 @@ def new_session_id() -> str: async def close(self) -> None: """Gracefully close the runner and release resources.""" - self._runner.close() + await self._runner.close() logger.info("RunnerManager closed: app=%s", self.app_name) # ------------------------------------------------------------------ diff --git a/examples/llmagent_with_model_retry/agent/agent.py b/examples/llmagent_with_model_retry/agent/agent.py index 25b3e7d43..f4a1a0d7b 100644 --- a/examples/llmagent_with_model_retry/agent/agent.py +++ b/examples/llmagent_with_model_retry/agent/agent.py @@ -5,6 +5,8 @@ # tRPC-Agent-Python is licensed under Apache-2.0. """Agent module for the model retry example.""" +from typing import Any + from trpc_agent_sdk.agents import LlmAgent from trpc_agent_sdk.models import LLMModel from trpc_agent_sdk.models import OpenAIModel @@ -15,6 +17,19 @@ from .prompts import INSTRUCTION from .tools import get_weather_report +def _extract_some_field(response_data: dict[str, Any]) -> dict[str, Any] | None: + """Allowlist Venus tracing metadata from an OpenAI-compatible response.""" + marker = response_data.get("some_field") + if not isinstance(marker, dict): + return None + some_value = marker.get("some_value") + if not isinstance(some_value, str) or not some_value: + return None + return { + "some_field": { + "some_value": some_value, + }, + } def _create_model() -> LLMModel: """Create an OpenAI-compatible model with SDK-managed retry enabled.""" @@ -26,6 +41,7 @@ def _create_model() -> LLMModel: api_key=api_key, base_url=base_url, model_retry_config=retry_config, + response_metadata_extractor=_extract_some_field, ) diff --git a/examples/llmagent_with_model_retry/run_agent.py b/examples/llmagent_with_model_retry/run_agent.py index ccd028781..29aff39a0 100644 --- a/examples/llmagent_with_model_retry/run_agent.py +++ b/examples/llmagent_with_model_retry/run_agent.py @@ -6,6 +6,7 @@ """Run the model retry weather agent example.""" import asyncio +import json import uuid from dotenv import load_dotenv @@ -47,6 +48,9 @@ async def run_weather_agent() -> None: assistant_started = True async for event in runner.run_async(user_id=user_id, session_id=session_id, new_message=user_content): + provider_metadata = (event.custom_metadata or {}).get("provider_response_metadata") + if provider_metadata: + print("\nProvider metadata: ", f"{json.dumps(provider_metadata, ensure_ascii=False)}") if event.is_error(): if assistant_started: print() diff --git a/tests/models/test_openai_model_ext.py b/tests/models/test_openai_model_ext.py index 2fa1aa6b4..1b45ef3f5 100644 --- a/tests/models/test_openai_model_ext.py +++ b/tests/models/test_openai_model_ext.py @@ -1441,6 +1441,37 @@ async def capture_create(**kwargs): assert captured[ApiParamsKey.SEED] == 42 assert captured[ApiParamsKey.N] == 2 + @pytest.mark.asyncio + async def test_non_streaming_extracts_provider_metadata(self): + """Provider metadata is attached to a non-streaming response.""" + model = _model(response_metadata_extractor=lambda data: {"provider_request_id": data["providerRequestId"]} + if data.get("providerRequestId") else None) + request = _request([Content(parts=[Part.from_text(text="hi")], role="user")]) + mock_response = Mock() + mock_response.model_dump.return_value = { + "choices": [{ + "message": { + "content": "ok", + "role": "assistant" + }, + "finish_reason": "stop", + }], + "usage": None, + "providerRequestId": "request-123", + } + + with patch.object(model, "_create_async_client") as mock_factory: + mock_client = AsyncMock() + mock_client.chat.completions.create = AsyncMock(return_value=mock_response) + mock_client.close = AsyncMock() + mock_factory.return_value = mock_client + + responses = [] + async for response in model.generate_async(request, stream=False): + responses.append(response) + + assert responses[0].custom_metadata == {"provider_response_metadata": {"provider_request_id": "request-123"}} + @pytest.mark.asyncio async def test_streaming_with_thinking_content(self): """Streaming mode correctly tags reasoning_content as thought.""" @@ -1493,6 +1524,68 @@ async def mock_stream(): thought_partials = [r for r in partial_responses if r.content and r.content.parts[0].thought] assert len(thought_partials) >= 1 + @pytest.mark.asyncio + async def test_streaming_extracts_provider_metadata_from_usage_chunk(self): + """Provider metadata survives a final usage-only chunk.""" + + def extract_metadata(response_data): + marker = response_data.get("venusMarker") + if not marker: + return None + return {"venus_marker": {"span_id": marker["spanId"]}} + + model = _model(response_metadata_extractor=extract_metadata) + request = _request([Content(parts=[Part.from_text(text="hi")], role="user")]) + + content_chunk = Mock() + content_chunk.model_dump.return_value = { + "id": "resp_1", + "choices": [{ + "delta": { + "content": "hello" + }, + "finish_reason": "stop", + }], + "usage": None, + } + usage_chunk = Mock() + usage_chunk.model_dump.return_value = { + "id": "resp_1", + "choices": [], + "usage": { + "prompt_tokens": 1, + "completion_tokens": 1, + "total_tokens": 2, + }, + "venusMarker": { + "spanId": "9d3e43a402a76a5b" + }, + } + + async def mock_stream(): + yield content_chunk + yield usage_chunk + + with patch.object(model, "_create_async_client") as mock_factory: + mock_client = AsyncMock() + mock_client.chat.completions.create = AsyncMock(return_value=mock_stream()) + mock_client.close = AsyncMock() + mock_factory.return_value = mock_client + + responses = [] + async for response in model.generate_async(request, stream=True): + responses.append(response) + + final_response = next(response for response in responses if not response.partial) + assert final_response.custom_metadata == { + "stream_complete": True, + "provider_response_metadata": { + "venus_marker": { + "span_id": "9d3e43a402a76a5b" + } + }, + } + @pytest.mark.asyncio async def test_streaming_null_response_raises(self): """Null response from API raises ValueError wrapped in error response.""" diff --git a/tests/telemetry/test_trace.py b/tests/telemetry/test_trace.py index 0a9fbf477..8c40df5dc 100644 --- a/tests/telemetry/test_trace.py +++ b/tests/telemetry/test_trace.py @@ -932,6 +932,26 @@ def test_basic_llm_trace(self, mock_get_span): span.set_attribute.assert_any_call("trpc.python.agent.session_id", "sess-1") span.set_attribute.assert_any_call("trpc.python.agent.event_id", "e-1") + @patch("trpc_agent_sdk.telemetry._trace.trace.get_current_span") + def test_reports_provider_response_metadata(self, mock_get_span): + span = _mock_span() + mock_get_span.return_value = span + ctx = _make_invocation_context() + req = self._make_llm_request() + resp = self._make_llm_response( + custom_metadata={"provider_response_metadata": { + "some_field": { + "some_value": "some_value" + } + }}) + + trace_call_llm(ctx, event_id="e-1", llm_request=req, llm_response=resp) + + span.set_attribute.assert_any_call( + "trpc.python.agent.provider_response_metadata", + '{"some_field": {"some_value": "some_value"}}', + ) + @patch("trpc_agent_sdk.telemetry._trace.trace.get_current_span") def test_explicit_error_sets_status_message_and_keeps_llm_response_output(self, mock_get_span): span = _mock_span() diff --git a/trpc_agent_sdk/models/__init__.py b/trpc_agent_sdk/models/__init__.py index fabaa68bf..60fbeeb24 100644 --- a/trpc_agent_sdk/models/__init__.py +++ b/trpc_agent_sdk/models/__init__.py @@ -37,6 +37,7 @@ from ._constants import TOOL_STREAMING_ARGS from ._constants import USAGE from ._constants import USER +from ._constants import PROVIDER_RESPONSE_METADATA from ._litellm_model import LiteLLMModel from ._llm_model import LLMModel from ._llm_request import LlmRequest @@ -81,6 +82,7 @@ "TOOL_STREAMING_ARGS", "THINKING_ENABLED", "THINKING_TOKENS", + "PROVIDER_RESPONSE_METADATA", "AnthropicModel", "LiteLLMModel", "LLMModel", diff --git a/trpc_agent_sdk/models/_constants.py b/trpc_agent_sdk/models/_constants.py index d274e874d..c8ffcd503 100644 --- a/trpc_agent_sdk/models/_constants.py +++ b/trpc_agent_sdk/models/_constants.py @@ -79,6 +79,9 @@ CHUNK: str = 'chunk' """Chunk field name in streaming responses.""" +PROVIDER_RESPONSE_METADATA: str = 'provider_response_metadata' +"""Allowlisted provider-specific metadata extracted from model responses.""" + TOOL_STREAMING: str = 'tool_streaming' """Tool streaming mode indicator name.""" diff --git a/trpc_agent_sdk/models/_openai_model.py b/trpc_agent_sdk/models/_openai_model.py index f3105bf55..ae40118c4 100644 --- a/trpc_agent_sdk/models/_openai_model.py +++ b/trpc_agent_sdk/models/_openai_model.py @@ -19,6 +19,7 @@ from enum import Enum from typing import Any from typing import AsyncGenerator +from typing import Callable from typing import Dict from typing import List from typing import Optional @@ -56,6 +57,7 @@ _HTTPCORE2_ATHROW_ERROR = "generator didn't stop after athrow" _HTTP_BODY_DRAIN_TIMEOUT_S = 2.0 +ResponseMetadataExtractor = Callable[[dict[str, Any]], Optional[dict[str, Any]]] def _is_httpx2_response(http_response: Any) -> bool: @@ -358,6 +360,10 @@ class OpenAIModel(LLMModel): the openai SDK's ``ResponseCreateParams`` and passed through verbatim to ``responses.create``. The model, input, and stream parameters remain managed by this class. + response_metadata_extractor: Optional callback that extracts a small, + JSON-serializable metadata dictionary from each + provider response or stream event. Extracted values + are attached to the final ``LlmResponse`` and trace. **kwargs: Additional arguments passed to parent LLMModel class (e.g., api_key, base_url, etc.) @@ -397,6 +403,7 @@ def __init__( http_client_provider_factory: HttpClientProviderFactory = temporary_http_client_provider_factory, use_responses_api: bool = False, responses_api_params: Optional[ResponseCreateParams] = None, + response_metadata_extractor: Optional[ResponseMetadataExtractor] = None, **kwargs, ): super().__init__(model_name, filters_name, **kwargs) @@ -407,6 +414,7 @@ def __init__( self.client_args = kwargs.get(const.CLIENT_ARGS, {}) self.use_responses_api = use_responses_api self.responses_api_params = dict(responses_api_params or {}) + self._response_metadata_extractor = response_metadata_extractor reserved_response_params = {"model", "input", "stream"}.intersection(self.responses_api_params) if reserved_response_params: names = ", ".join(sorted(reserved_response_params)) @@ -452,6 +460,40 @@ def _refresh_adapter(self) -> None: def is_retriable_status_code(self, status_code: int) -> Optional[bool]: return status_code in {408, 409, 429} or status_code >= 500 + def _extract_provider_response_metadata(self, response_data: dict[str, Any]) -> dict[str, Any]: + """Extract allowlisted provider metadata without affecting model calls.""" + if self._response_metadata_extractor is None: + return {} + try: + metadata = self._response_metadata_extractor(response_data) + if metadata is None: + return {} + if not isinstance(metadata, dict): + logger.warning( + "response_metadata_extractor returned %s instead of dict; ignoring it", + type(metadata).__name__, + ) + return {} + # LlmResponse.custom_metadata must remain JSON serializable. + json.dumps(metadata) + return metadata + except Exception: # pylint: disable=broad-except + logger.warning("Failed to extract provider response metadata", exc_info=True) + return {} + + @staticmethod + def _attach_provider_response_metadata( + response: LlmResponse, + metadata: dict[str, Any], + ) -> LlmResponse: + """Attach extracted metadata under a stable, provider-neutral namespace.""" + if not metadata: + return response + custom_metadata = dict(response.custom_metadata or {}) + custom_metadata[const.PROVIDER_RESPONSE_METADATA] = metadata + response.custom_metadata = custom_metadata + return response + def is_retriable_exception(self, ex: Exception) -> bool: if isinstance(ex, httpx.TimeoutException): return True @@ -1753,7 +1795,12 @@ async def _generate_responses_single( **self._prepare_responses_api_params(client, api_params), **(http_options or {}), ) - return self._create_responses_response(self._model_dump(response)) + response_dict = self._model_dump(response) + llm_response = self._create_responses_response(response_dict) + return self._attach_provider_response_metadata( + llm_response, + self._extract_provider_response_metadata(response_dict), + ) finally: await self._http_client_provider.close_http_client(client) @@ -1783,8 +1830,11 @@ async def _generate_single( # Create response with content if we have text or tool calls if has_text_content or has_tool_calls: - return self._create_response_with_content(response_dict) - return self._create_response_without_content(response_dict) + llm_response = self._create_response_with_content(response_dict) + else: + llm_response = self._create_response_without_content(response_dict) + provider_response_metadata = self._extract_provider_response_metadata(response_dict) + return self._attach_provider_response_metadata(llm_response, provider_response_metadata) finally: await self._http_client_provider.close_http_client(client) @@ -2257,9 +2307,10 @@ def upsert_function(item: dict) -> tuple[str, Dict[str, Any]]: if response is None: raise ValueError("Empty response from Responses API") _patch_stream_response_to_drain_http_body(response) - + last_event_dict: dict[str, Any] = {} async for event in response: event_dict = self._model_dump(event) + last_event_dict = event_dict event_type = event_dict.get("type", "") logger.debug("OpenAI Responses event: %s", json.dumps(event_dict, ensure_ascii=False)) @@ -2374,7 +2425,9 @@ def upsert_function(item: dict) -> tuple[str, Dict[str, Any]]: final_response = self._create_responses_response(completed_response) final_response.partial = False final_response.custom_metadata = {"stream_complete": True} - yield final_response + + provider_response_metadata = self._extract_provider_response_metadata(last_event_dict) + yield self._attach_provider_response_metadata(final_response, provider_response_metadata) finally: await _aclose_openai_stream(response) try: @@ -2424,12 +2477,14 @@ async def _generate_stream( raise ValueError("Empty response from API") _patch_stream_response_to_drain_http_body(response) + last_event_dict: dict[str, Any] = {} async for chunk in response: if chunk is None: continue chunk_dict: dict = chunk.model_dump() logger.debug("🔥 RAW LLM CHUNK: %s", json.dumps(chunk_dict, ensure_ascii=False)) + last_event_dict = chunk_dict # Capture response ID from chunk (only set once from first chunk that has it) if response_id is None and chunk_dict.get("id"): @@ -2628,14 +2683,15 @@ async def _generate_stream( if last_usage: # Create a compatible usage metadata object final_usage = last_usage # Use the existing usage object for now - - yield LlmResponse( + final_response = LlmResponse( content=final_content, usage_metadata=final_usage, partial=False, response_id=response_id, custom_metadata={"stream_complete": True}, ) + provider_response_metadata = self._extract_provider_response_metadata(last_event_dict) + yield self._attach_provider_response_metadata(final_response, provider_response_metadata) finally: await _aclose_openai_stream(response) await self._http_client_provider.close_http_client(client) diff --git a/trpc_agent_sdk/telemetry/_trace.py b/trpc_agent_sdk/telemetry/_trace.py index d0f036861..f599ecb4f 100644 --- a/trpc_agent_sdk/telemetry/_trace.py +++ b/trpc_agent_sdk/telemetry/_trace.py @@ -37,6 +37,7 @@ from trpc_agent_sdk.events import Event from trpc_agent_sdk.models import LlmRequest from trpc_agent_sdk.models import LlmResponse +from trpc_agent_sdk.models import PROVIDER_RESPONSE_METADATA from trpc_agent_sdk.tools import BaseTool from trpc_agent_sdk.types import Content from trpc_agent_sdk.types import InstructionMetadata @@ -503,6 +504,15 @@ def trace_call_llm( llm_response_json, ) + custom_metadata = llm_response.custom_metadata + if isinstance(custom_metadata, dict): + provider_metadata = custom_metadata.get(PROVIDER_RESPONSE_METADATA) + if isinstance(provider_metadata, dict) and provider_metadata: + span.set_attribute( + f"{_trpc_agent_span_name}.{PROVIDER_RESPONSE_METADATA}", + _safe_json_serialize(provider_metadata), + ) + # The caller-supplied error_type reflects an exception that propagated out # of the model call. But the SDK-managed retry layer (retry_model_call) # can also swallow a raised exception and yield a normal-looking