diff --git a/docs/mkdocs/en/model.md b/docs/mkdocs/en/model.md index a67812c8b..c483771b4 100644 --- a/docs/mkdocs/en/model.md +++ b/docs/mkdocs/en/model.md @@ -11,6 +11,7 @@ Models in tRPC-Agent have the following core features: - **Multimodal capabilities**: Supports multimodal content processing including text, images, etc. (e.g., Hunyuan multimodal models) - **Prompt Cache support**: Provides unified prompt cache configuration across OpenAI, Anthropic, and LiteLLM routes to reduce repeated input cost for long prompts and multi-turn conversations - **Model retry support**: Supports configuring retry at the model layer. The SDK automatically retries when exceptions such as rate limits occur, and backs off using an exponential backoff strategy +- **Provider metadata extraction**: Allows OpenAI-compatible provider-specific response fields to be allowlisted and attached to model responses and tracing spans - **Extensible configuration**: Supports custom configuration options such as GenerateContentConfig, HttpOptions, client_args to meet various scenario requirements ## Quick Start @@ -167,6 +168,69 @@ see the OpenAI reasoning guide for model-specific support. The SDK deliberately `thinking_budget` to `effort`; it only requests `reasoning.summary` when thinking is enabled so the reasoning output stays readable. +#### Extracting Provider-Specific Response Fields + +Some OpenAI-compatible providers add proprietary fields to the response, such as a request or tracing +identifier. Configure `response_metadata_extractor` to explicitly allowlist and normalize the fields +that should be exposed to application code and tracing: + +```python +from typing import Any + +from trpc_agent_sdk.models import OpenAIModel + + +def extract_provider_metadata( + response_data: dict[str, Any], +) -> dict[str, Any] | None: + some_marker = response_data.get("some_marker") + if not isinstance(some_marker, dict): + return None + some_field = some_marker.get("some_field") + if not isinstance(some_field, str) or not some_field: + return None + return { + "some_marker": { + "some_field": some_field, + }, + } + + +model = OpenAIModel( + model_name="your-model", + api_key="your-api-key", + base_url="https://your-openai-compatible-endpoint/v1", + response_metadata_extractor=extract_provider_metadata, +) +``` + +The callback receives the complete `model_dump()` dictionary for the provider response. It may read +top-level or nested fields, but should return only the small, JSON-serializable values required by the +application. Returning `None` means that no metadata was extracted. Invalid return values and callback +exceptions are logged and ignored without failing the model call. + +For streaming requests, the extractor is invoked exactly once with the first response event to avoid +per-token overhead. The compatible provider must therefore include its metadata in the first event. + +The extracted value is stored under the stable `provider_response_metadata` namespace: + +```python +from trpc_agent_sdk.models import PROVIDER_RESPONSE_METADATA + +async for event in runner.run_async(...): + provider_metadata = (event.custom_metadata or {}).get( + PROVIDER_RESPONSE_METADATA + ) + if provider_metadata: + print(provider_metadata) +``` + +The same metadata is also attached to the model invocation tracing span. Do not return the full raw +provider response from the extractor, because it may contain sensitive or unnecessarily large data. + +For a complete runnable example, see +[examples/llmagent_with_model_extra_fields](../../../examples/llmagent_with_model_extra_fields/README.md). + #### Advanced Usage Since version `1.1.10`, `OpenAIModel` supports passing a shared HTTP client provider to enable connection reuse. By default, `OpenAIModel` creates a temporary HTTP client for each model-service request. If you want to reuse connections, use the following configuration: diff --git a/docs/mkdocs/zh/model.md b/docs/mkdocs/zh/model.md index f74ed7717..6d2b5bcd8 100644 --- a/docs/mkdocs/zh/model.md +++ b/docs/mkdocs/zh/model.md @@ -11,6 +11,7 @@ tRPC-Agent 内的模型具有以下核心特性: - **多模态能力**:支持文本、图像等多模态内容处理(如 hunyuan 多模态模型) - **Prompt Cache 支持**:支持跨 OpenAI、Anthropic 与 LiteLLM 路由的统一 prompt cache 配置,降低长提示词和多轮会话的重复输入成本 - **模型重试支持**:支持在模型层配置重试,SDK 将在限流等异常发生时自动重试,并按指数退避策略进行退避 +- **Provider 元数据提取**:支持将 OpenAI 兼容服务返回的厂商扩展字段加入白名单,并附加到模型响应和 tracing span - **可扩展配置**:支持 GenerateContentConfig、HttpOptions、client_args 等自定义配置项,满足不同场景需求 ## 快速上手 @@ -164,6 +165,67 @@ model = OpenAIModel( 具体以 OpenAI reasoning 指南及模型文档为准。SDK 刻意不做 `thinking_budget` → `effort` 的映射; 仅在开启 thinking 时请求 `reasoning.summary`,保证推理输出可读。 +#### 提取厂商响应扩展字段 + +部分 OpenAI 兼容服务会在响应中增加请求标识、链路标识等厂商扩展字段。可以配置 +`response_metadata_extractor`,明确选择并规范化允许暴露给业务代码和 tracing 的字段: + +```python +from typing import Any + +from trpc_agent_sdk.models import OpenAIModel + + +def extract_provider_metadata( + response_data: dict[str, Any], +) -> dict[str, Any] | None: + some_marker = response_data.get("some_marker") + if not isinstance(some_marker, dict): + return None + some_field = some_marker.get("some_field") + if not isinstance(some_field, str) or not some_field: + return None + return { + "some_marker": { + "some_field": some_field, + }, + } + + +model = OpenAIModel( + model_name="your-model", + api_key="your-api-key", + base_url="https://your-openai-compatible-endpoint/v1", + response_metadata_extractor=extract_provider_metadata, +) +``` + +回调接收厂商响应完整的 `model_dump()` 字典,可以读取顶层或嵌套字段,但应只返回业务所需的、 +体积较小且可 JSON 序列化的数据。返回 `None` 表示没有提取到元数据。返回值非法或回调抛出异常时, +SDK 会记录日志并忽略该元数据,不会导致模型调用失败。 + +对于流式请求,为避免产生逐 token 开销,提取器只会使用第一个响应事件调用一次。因此兼容服务 +必须在首个事件中携带需要提取的厂商字段。 + +提取结果存放在稳定的 `provider_response_metadata` 命名空间中: + +```python +from trpc_agent_sdk.models import PROVIDER_RESPONSE_METADATA + +async for event in runner.run_async(...): + provider_metadata = (event.custom_metadata or {}).get( + PROVIDER_RESPONSE_METADATA + ) + if provider_metadata: + print(provider_metadata) +``` + +同一份元数据也会附加到模型调用的 tracing span。不要从提取器返回完整原始响应,以免泄漏敏感信息 +或引入不必要的大体积数据。 + +完整可运行示例参见 +[examples/llmagent_with_model_extra_fields](../../../examples/llmagent_with_model_extra_fields/README.md)。 + #### 高级用法 从版本 `1.1.10`之后 OpenAIModel 支持传入共享的 http client 来解决连接复用的场景,当前的 OpenAIModel 默认每次都会创建临时的 http client 去访问模型服务;如果期望连接复用可以使用如下的方式 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_extra_fields/.env b/examples/llmagent_with_model_extra_fields/.env new file mode 100644 index 000000000..133025196 --- /dev/null +++ b/examples/llmagent_with_model_extra_fields/.env @@ -0,0 +1,8 @@ +# Copy this file or edit it in place before running the example. +# The example uses an OpenAI-compatible endpoint. + +TRPC_AGENT_API_KEY= +TRPC_AGENT_BASE_URL= +TRPC_AGENT_MODEL_NAME= + + diff --git a/examples/llmagent_with_model_extra_fields/README.md b/examples/llmagent_with_model_extra_fields/README.md new file mode 100644 index 000000000..89d6c8687 --- /dev/null +++ b/examples/llmagent_with_model_extra_fields/README.md @@ -0,0 +1,145 @@ +# LLM Agent 模型额外字段上报示例 + +本示例演示如何通过 `OpenAIModel.response_metadata_extractor`,从 OpenAI +兼容服务的响应中提取厂商扩展字段,并通过 `LlmResponse.custom_metadata` +传递给业务代码和 tracing 链路。 + +## 关键特性 + +- **显式白名单提取**:业务通过回调决定允许上报哪些厂商字段,避免透传完整原始响应。 +- **兼容流式与非流式响应**:回调接收响应的完整 `model_dump()` 字典;流式模式只读取第一个事件。 +- **稳定的元数据命名空间**:结果保存在 + `custom_metadata["provider_response_metadata"]` 中。 +- **失败不影响模型调用**:回调返回 `None`、非字典、不可 JSON 序列化数据或抛出异常时,SDK 会忽略该元数据。 +- **自动进入 tracing**:提取结果会作为 provider response metadata 上报到模型调用 span。 + +## Agent 层级结构说明 + +本例是单 Agent 示例,额外字段提取器绑定在模型上: + +```text +weather_agent (LlmAgent) +├── model: OpenAIModel(..., response_metadata_extractor=_extract_some_field) +├── tool: get_weather_report(city) +└── runner: 从 Event.custom_metadata 读取并打印厂商元数据 +``` + +关键文件: + +- [agent/agent.py](./agent/agent.py):定义白名单提取器并注入 `OpenAIModel` +- [agent/config.py](./agent/config.py):读取模型连接环境变量 +- [agent/tools.py](./agent/tools.py):天气工具实现 +- [agent/prompts.py](./agent/prompts.py):Agent 提示词 +- [run_agent.py](./run_agent.py):运行入口,读取并打印 provider metadata + +## 关键代码解释 + +这里的 `some_marker` `some_field` 表示厂商返回的额外字段,具体视模型本身返回的真实数据为例,这里只是一个描述 + +### 1) 定义厂商字段提取器 + +```python +def _extract_some_field( + response_data: dict[str, Any], +) -> dict[str, Any] | None: + some_marker = response_data.get("some_marker") + if not isinstance(some_marker, dict): + return None + some_field = some_marker.get("some_field") + if not isinstance(some_field, str) or not some_field: + return None + return { + "some_marker": { + "some_field": some_field, + }, + } +``` + +`response_data` 是 OpenAI SDK 响应对象的完整 `model_dump()` 结果。回调只返回 +允许进入业务事件和 tracing 的字段,并可在这里统一转换命名格式。 + +> 流式模式只会使用第一个响应事件执行一次提取,因此兼容服务需要在首个事件中携带厂商字段。 + +### 2) 将提取器注入模型 + +`agent/agent.py` 将回调传给 `OpenAIModel`: + +```python +OpenAIModel( + model_name=model_name, + api_key=api_key, + base_url=base_url, + response_metadata_extractor=_extract_some_field, +) +``` + +### 3) 从事件中读取元数据 + +```python +from trpc_agent_sdk.models import PROVIDER_RESPONSE_METADATA + +provider_metadata = (event.custom_metadata or {}).get( + PROVIDER_RESPONSE_METADATA +) +``` + +本例的一次天气问答包含两次模型调用:第一次模型选择天气工具,第二次模型根据 +工具结果生成最终回答。因此运行输出中会看到两个不同的 `some_field`。 + +## 环境要求 + +- Python3.10+,推荐 Python3.12 + +## 构建步骤 + +```bash +git clone https://github.com/trpc-group/trpc-agent-python.git +cd trpc-agent-python +./build.sh +source .venv/bin/activate +``` + +## 运行步骤 + +### 配置环境变量 + +在 [`.env`](./.env) 中配置(或通过 `export` 设置): + +```bash +TRPC_AGENT_API_KEY= +TRPC_AGENT_BASE_URL= +TRPC_AGENT_MODEL_NAME= +``` + +模型服务需要返回提取器识别的扩展字段。本例期望响应中包含: + +```json +{ + "some_marker": { + "some_field": "" + } +} +``` + +### 运行命令 + +```bash +cd examples/llmagent_with_model_extra_fields +python3 run_agent.py +``` + +## 运行结果示例 + +```text +Session ID: 926d3fa3... +User: What's the current weather in Beijing? +Assistant: +Provider metadata: {"some_marker": {"some_field": "d3f87b419c8e871c"}} + +Invoke Tool: get_weather_report({'city': 'Beijing'}) +Tool Result: {'temperature': '25°C', 'condition': 'Sunny', 'humidity': '60%'} +Assistant: It's currently **25°C and sunny** in Beijing. +Provider metadata: {"some_marker": {"some_field": "5ca6786f4a74f5b1"}} +``` + +`some_field` 由模型服务生成,每次运行都会不同。 diff --git a/examples/llmagent_with_model_extra_fields/agent/__init__.py b/examples/llmagent_with_model_extra_fields/agent/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/examples/llmagent_with_model_extra_fields/agent/agent.py b/examples/llmagent_with_model_extra_fields/agent/agent.py new file mode 100644 index 000000000..c5b46f2ef --- /dev/null +++ b/examples/llmagent_with_model_extra_fields/agent/agent.py @@ -0,0 +1,55 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Agent module for the model extra fields 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 +from trpc_agent_sdk.tools import FunctionTool + +from .config import get_model_config +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("venusMarker") + if not isinstance(marker, dict): + return None + span_id = marker.get("spanId") + if not isinstance(span_id, str) or not span_id: + return None + return { + "venusMarker": { + "span_id": span_id, + }, + } + +def _create_model() -> LLMModel: + """Create an OpenAI-compatible model with SDK-managed extra fields enabled.""" + api_key, base_url, model_name = get_model_config() + return OpenAIModel( + model_name=model_name, + api_key=api_key, + base_url=base_url, + response_metadata_extractor=_extract_some_field, + ) + + +def create_agent() -> LlmAgent: + """Create a weather agent that uses model-level extra fields.""" + return LlmAgent( + name="weather_agent", + description="A weather assistant with SDK-managed extra fields enabled.", + model=_create_model(), + instruction=INSTRUCTION, + tools=[FunctionTool(get_weather_report)], + ) + + +root_agent = create_agent() diff --git a/examples/llmagent_with_model_extra_fields/agent/config.py b/examples/llmagent_with_model_extra_fields/agent/config.py new file mode 100644 index 000000000..8a47cf133 --- /dev/null +++ b/examples/llmagent_with_model_extra_fields/agent/config.py @@ -0,0 +1,21 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Configuration helpers for the model extra fields example.""" + +import os + +def get_model_config() -> tuple[str, str, str]: + """Get model connection settings from environment variables.""" + api_key = os.getenv("TRPC_AGENT_API_KEY", "") + base_url = os.getenv("TRPC_AGENT_BASE_URL", "") + model_name = os.getenv("TRPC_AGENT_MODEL_NAME", "") + if not api_key or not base_url or not model_name: + raise ValueError("TRPC_AGENT_API_KEY, TRPC_AGENT_BASE_URL, and " + "TRPC_AGENT_MODEL_NAME must be set in environment variables") + return api_key, base_url, model_name + + + diff --git a/examples/llmagent_with_model_extra_fields/agent/prompts.py b/examples/llmagent_with_model_extra_fields/agent/prompts.py new file mode 100644 index 000000000..707ea6841 --- /dev/null +++ b/examples/llmagent_with_model_extra_fields/agent/prompts.py @@ -0,0 +1,13 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Prompts for the model extra fields example agent.""" + +INSTRUCTION = """You are a practical weather assistant. + +When the user asks for weather, identify the city and call get_weather_report. +If the city is missing, ask one short clarification question. +After receiving tool results, summarize the weather clearly and mention the extra fields only if the user asks about them. +""" diff --git a/examples/llmagent_with_model_extra_fields/agent/tools.py b/examples/llmagent_with_model_extra_fields/agent/tools.py new file mode 100644 index 000000000..4f7d45b3f --- /dev/null +++ b/examples/llmagent_with_model_extra_fields/agent/tools.py @@ -0,0 +1,40 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Tools for the model extra fields example agent.""" + + +def get_weather_report(city: str) -> dict: + """Get weather information for the specified city.""" + weather_data = { + "Beijing": { + "temperature": "25°C", + "condition": "Sunny", + "humidity": "60%", + }, + "Shanghai": { + "temperature": "28°C", + "condition": "Cloudy", + "humidity": "70%", + }, + "Guangzhou": { + "temperature": "32°C", + "condition": "Thunderstorm", + "humidity": "85%", + }, + "Shenzhen": { + "temperature": "30°C", + "condition": "Light rain", + "humidity": "78%", + }, + } + return weather_data.get( + city, + { + "temperature": "Unknown", + "condition": "Data not available", + "humidity": "Unknown", + }, + ) diff --git a/examples/llmagent_with_model_extra_fields/run_agent.py b/examples/llmagent_with_model_extra_fields/run_agent.py new file mode 100644 index 000000000..90abe641d --- /dev/null +++ b/examples/llmagent_with_model_extra_fields/run_agent.py @@ -0,0 +1,87 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Run the model extra fields weather agent example.""" + +import asyncio +import json +import uuid + +from dotenv import load_dotenv +from trpc_agent_sdk.runners import Runner +from trpc_agent_sdk.sessions import InMemorySessionService +from trpc_agent_sdk.types import Content +from trpc_agent_sdk.types import Part +from trpc_agent_sdk.models import PROVIDER_RESPONSE_METADATA + +load_dotenv() + + +async def run_weather_agent() -> None: + """Run the weather query agent with model-level extra fields enabled.""" + app_name = "model_extra_fields_weather_demo" + + from agent.agent import root_agent + + session_service = InMemorySessionService() + runner = Runner(app_name=app_name, agent=root_agent, session_service=session_service) + + user_id = "demo_user" + session_id = str(uuid.uuid4()) + await session_service.create_session( + app_name=app_name, + user_id=user_id, + session_id=session_id, + ) + + query = "What's the current weather in Beijing?" + print(f"Session ID: {session_id[:8]}...") + print(f"User: {query}") + print("Assistant: ", end="", flush=True) + + user_content = Content(parts=[Part.from_text(text=query)]) + 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() + assistant_started = False + print(f"Error: {event.error_code}: {event.error_message}") + continue + + if not event.content or not event.content.parts: + continue + + if event.partial: + for part in event.content.parts: + if part.text and not part.thought: + if not assistant_started: + print("Assistant: ", end="", flush=True) + assistant_started = True + print(part.text, end="", flush=True) + continue + + for part in event.content.parts: + if part.thought: + continue + if part.function_call: + print(f"\nInvoke Tool: {part.function_call.name}({part.function_call.args})") + assistant_started = False + elif part.function_response: + print(f"Tool Result: {part.function_response.response}") + elif part.text and not assistant_started: + print("Assistant: ", end="", flush=True) + print(part.text, end="", flush=True) + assistant_started = True + + print("\n") + + +if __name__ == "__main__": + asyncio.run(run_weather_agent()) diff --git a/tests/models/test_openai_model_ext.py b/tests/models/test_openai_model_ext.py index 2fa1aa6b4..c3fbf2b11 100644 --- a/tests/models/test_openai_model_ext.py +++ b/tests/models/test_openai_model_ext.py @@ -1441,6 +1441,99 @@ 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_non_streaming_extracts_nested_message_metadata(self): + """Extractor can read vendor fields nested under choices[0].message.""" + + def extract_metadata(response_data): + choices = response_data.get("choices") or [] + message = choices[0].get("message", {}) if choices else {} + marker = message.get("venusMarker") + if not isinstance(marker, dict): + 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")]) + mock_response = Mock() + mock_response.model_dump.return_value = { + "choices": [{ + "message": { + "content": "ok", + "role": "assistant", + "venusMarker": { + "spanId": "nested-span" + }, + }, + "finish_reason": "stop", + }], + "usage": None, + } + + 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": { + "venus_marker": { + "span_id": "nested-span" + } + } + } + + def test_extractor_exception_and_invalid_results_are_ignored(self): + """Extractor failures must not break model calls.""" + + def boom(_data): + raise RuntimeError("bad extractor") + + model = _model(response_metadata_extractor=boom) + assert model._extract_provider_response_metadata({"providerRequestId": "x"}) == {} + + model = _model(response_metadata_extractor=lambda _: "not-a-dict") + assert model._extract_provider_response_metadata({"providerRequestId": "x"}) == {} + + model = _model(response_metadata_extractor=lambda _: None) + assert model._extract_provider_response_metadata({"providerRequestId": "x"}) == {} + @pytest.mark.asyncio async def test_streaming_with_thinking_content(self): """Streaming mode correctly tags reasoning_content as thought.""" @@ -1493,6 +1586,255 @@ 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_first_chunk(self): + """Provider metadata is extracted once from the first stream 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, + "venusMarker": { + "spanId": "9d3e43a402a76a5b" + }, + } + 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_keeps_metadata_from_earlier_chunk(self): + """A later usage-only chunk without metadata must not drop earlier fields.""" + + 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, + "venusMarker": { + "spanId": "early-span" + }, + } + usage_chunk = Mock() + usage_chunk.model_dump.return_value = { + "id": "resp_1", + "choices": [], + "usage": { + "prompt_tokens": 1, + "completion_tokens": 1, + "total_tokens": 2, + }, + } + + 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": "early-span" + } + }, + } + + @pytest.mark.asyncio + async def test_responses_streaming_keeps_metadata_from_earlier_event(self): + """Responses stream merges metadata even if response.completed lacks it.""" + + def extract_metadata(response_data): + marker = response_data.get("venusMarker") + if not marker: + return None + return {"venus_marker": {"span_id": marker["spanId"]}} + + model = _model(use_responses_api=True, response_metadata_extractor=extract_metadata) + request = _request([Content(parts=[Part.from_text(text="hi")], role="user")]) + + async def stream_events(): + yield { + "type": "response.created", + "response": { + "id": "resp_meta", + }, + "venusMarker": { + "spanId": "responses-span" + }, + } + yield {"type": "response.output_text.delta", "delta": "hello"} + yield { + "type": "response.completed", + "response": { + "id": "resp_meta", + "status": "completed", + "output": [{ + "type": "message", + "role": "assistant", + "content": [{ + "type": "output_text", + "text": "hello" + }], + }], + }, + } + + with patch.object(model, "_create_async_client") as mock_factory: + mock_client = AsyncMock() + mock_client.responses.create = AsyncMock(return_value=stream_events()) + 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["provider_response_metadata"] == { + "venus_marker": { + "span_id": "responses-span" + } + } + + @pytest.mark.asyncio + async def test_streaming_skips_extractor_after_first_nonempty(self): + """Later chunks must not rerun the extractor once metadata is found.""" + calls = {"count": 0} + + def extract_metadata(response_data): + calls["count"] += 1 + 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")]) + + first_chunk = Mock() + first_chunk.model_dump.return_value = { + "id": "resp_1", + "choices": [{ + "delta": { + "content": "hello" + }, + "finish_reason": "stop", + }], + "usage": None, + "venusMarker": { + "spanId": "first-span" + }, + } + later_chunk = Mock() + later_chunk.model_dump.return_value = { + "id": "resp_1", + "choices": [], + "usage": { + "prompt_tokens": 1, + "completion_tokens": 1, + "total_tokens": 2, + }, + "venusMarker": { + "spanId": "later-span" + }, + } + + async def mock_stream(): + yield first_chunk + yield later_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["provider_response_metadata"] == { + "venus_marker": { + "span_id": "first-span" + } + } + assert calls["count"] == 1 + @pytest.mark.asyncio async def test_streaming_null_response_raises(self): """Null response from API raises ValueError wrapped in error response.""" 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..d95d77c11 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,15 @@ 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 a + provider response. The callback receives the full + ``model_dump()`` dictionary and is responsible for + reading any nested fields. For streams, it is + called exactly once with the first event; providers + must include their metadata in that 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 +408,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 +419,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 +465,41 @@ 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.""" + extractor = self._response_metadata_extractor + if extractor is None or not response_data: + return {} + try: + 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 +1801,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 +1836,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 +2313,13 @@ 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) - + provider_response_metadata: dict[str, Any] = {} + metadata_extraction_attempted = False async for event in response: event_dict = self._model_dump(event) + if not metadata_extraction_attempted: + provider_response_metadata = self._extract_provider_response_metadata(event_dict) + metadata_extraction_attempted = True event_type = event_dict.get("type", "") logger.debug("OpenAI Responses event: %s", json.dumps(event_dict, ensure_ascii=False)) @@ -2374,7 +2434,8 @@ 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 + + yield self._attach_provider_response_metadata(final_response, provider_response_metadata) finally: await _aclose_openai_stream(response) try: @@ -2424,12 +2485,17 @@ async def _generate_stream( raise ValueError("Empty response from API") _patch_stream_response_to_drain_http_body(response) + provider_response_metadata: dict[str, Any] = {} + metadata_extraction_attempted = False 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)) + if not metadata_extraction_attempted: + provider_response_metadata = self._extract_provider_response_metadata(chunk_dict) + metadata_extraction_attempted = True # 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 +2694,14 @@ 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}, ) + 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)