Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion examples/fastapi_server/_runner_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

# ------------------------------------------------------------------
Expand Down
16 changes: 16 additions & 0 deletions examples/llmagent_with_model_retry/agent/agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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."""
Expand All @@ -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,
)


Expand Down
4 changes: 4 additions & 0 deletions examples/llmagent_with_model_retry/run_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
"""Run the model retry weather agent example."""

import asyncio
import json
import uuid

from dotenv import load_dotenv
Expand Down Expand Up @@ -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()
Expand Down
93 changes: 93 additions & 0 deletions tests/models/test_openai_model_ext.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down Expand Up @@ -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):
Comment on lines +1528 to 1590

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

问题: 新测试 test_streaming_extracts_provider_metadata_from_usage_chunk 覆盖了 final usage-chunk 场景,但未覆盖提取器抛异常、返回非 dict、返回 None 及元数据在首个/中间 chunk(非末帧)的路径,未覆盖 _generate_responses_stream(Responses API)与 _generate_single 异常分支。这些分支的行为各不相同。

触发条件: 回归测试执行时,凡依赖上述未覆盖分支的用户场景(厂商字段早到、提取器异常、Responses API 流)都没有防回归保障。

实际影响: 异常/非 dict 返回被静默吞掉、末帧覆盖中间帧等缺陷无法被 CI 发现,且将来改动易引入回归。

修正方向: 补齐 except Exception 返回空 dict、非 dict 类型、首/中帧携带元数据、Responses 流提取等用例的断言。

"""Null response from API raises ValueError wrapped in error response."""
Expand Down
20 changes: 20 additions & 0 deletions tests/telemetry/test_trace.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
2 changes: 2 additions & 0 deletions trpc_agent_sdk/models/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -81,6 +82,7 @@
"TOOL_STREAMING_ARGS",
"THINKING_ENABLED",
"THINKING_TOKENS",
"PROVIDER_RESPONSE_METADATA",
"AnthropicModel",
"LiteLLMModel",
"LLMModel",
Expand Down
3 changes: 3 additions & 0 deletions trpc_agent_sdk/models/_constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down
70 changes: 63 additions & 7 deletions trpc_agent_sdk/models/_openai_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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.)

Expand Down Expand Up @@ -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)
Expand All @@ -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))
Expand Down Expand Up @@ -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
Comment on lines +489 to +495

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

问题: _attach_provider_response_metadata 直接原地 response.custom_metadata[const.PROVIDER_RESPONSE_METADATA] = metadata,没有走 model_copy。终态 LlmResponse 是非流式 _generate_singleyield 前的最后一步,流式路径也只作用于新建的 final response,属安全;但 _responses_error 与非流式 _generate_singlecustom_metadata 存在其他分支赋值(如 _retry.py 的 error response),若后续在此基类上复用该方法,原地修改可能污染同一实例。

触发条件: 后续扩展在 _attach_provider_response_metadata 被调用时引用同一 LlmResponsecustom_metadata 的其他字段。

实际影响: 当前无实害,custom_metadata 的赋值本就要求整体兼容;仅提示未来改动与 partial 响应路径共用一个实例时可能产生隐式共享。

修正方向: 改为 response.model_copy(update={'custom_metadata': {**(response.custom_metadata or {}), const.PROVIDER_RESPONSE_METADATA: metadata}}) 或先 dict() 再赋值,保持不变性。


def is_retriable_exception(self, ex: Exception) -> bool:
if isinstance(ex, httpx.TimeoutException):
return True
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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)
Comment on lines +1836 to +1837

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

问题: _generate_single(非流式)中提取器收到的是 response.model_dump() 的顶层 dict。多数 OpenAI 兼容后端把厂商字段放在响应体顶层(如 providerRequestIdsome_field),测试也如此构造;但部分后端(如 OpenAI Responses 结构或 gateway 实现)把元数据放在 choices[].messageresponse 嵌套对象里,提取器只拿到顶层 dict,看不到嵌套字段。已有测试 test_non_streaming_extracts_provider_metadata 只覆盖顶层场景。

触发条件: 后端把目标字段放在 choices[0].messageresponse 或其他嵌套层而非顶层时。

实际影响: 嵌套位置携带的厂商字段无法上报,链路观测数据缺失;且提取器失败仅产生一条 warning,用户难以定位。

修正方向: 向提取器传入完整响应 dict 的同时文档明确约定字段位置,或提取时对 choices[0].message 等常见嵌套位置做兜底查找;并补充非顶层字段的测试。

finally:
await self._http_client_provider.close_http_client(client)

Expand Down Expand Up @@ -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))

Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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"):
Expand Down Expand Up @@ -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)
Comment on lines +2693 to +2694

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

问题: 流式路径(Chat Completions _generate_stream 与 Responses API _generate_responses_stream)仅在 async for 结束后用最后一块 chunk 调用一次提取器。OpenAI 兼容后端在 SSE 流终止前往往还会发送额外的元数据事件(如无 choicesusage 收尾块、response.completed 等),这些事件在真实流中是最后一条;而测试 test_streaming_extracts_provider_metadata_from_usage_chunk 中最后一块恰好是携带 venusMarker 的 usage-only chunk——该字段在该块中才出现。即实测中只有当后端调度恰好把元数据放在最后一帧时才成功;若元数据早到(如首个非空 chunk 或中间 chunk 携带),或最后还有一个不含该字段的收尾 chunk,提取器见不到它,元数据被丢弃。Responses 流中 last_event_dict 也常被 terminal 的 response.completed 事件(response_data 中通常不含顶层元数据)覆盖。

触发条件: 响应头/首个 chunk 携带元数据、或 SSE 流有多个收尾事件(真实后端常见),末帧不含目标字段。

实际影响: 依赖该功能的可观测性场景(如 Venus 追踪链路)中元数据静默丢失,最终 LlmResponse.custom_metadata 及 trace span 中缺少本应上报的 provider_response_metadata

修正方向: 累积候选事件(如保留首个携带且提取器返回非空的 last_event_dict,即第一次提取成功后不再覆盖),或在每个 chunk 到达时调用提取器并合并非空结果,而不只在循环结束后用最后一帧。

finally:
await _aclose_openai_stream(response)
await self._http_client_provider.close_http_client(client)
Loading
Loading