From 3a3c29a7205ce6113d0f4ad30f8e83c2375545ba Mon Sep 17 00:00:00 2001 From: Banibrata De <3373242+banide@users.noreply.github.com> Date: Thu, 17 Sep 2026 14:44:37 -0700 Subject: [PATCH] Preserve budget failures in hosted responses Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 99113cf0-df56-4e0f-979a-a7d39a173e4d --- .../_responses.py | 67 ++++++++++++++++++- .../foundry_hosting/tests/test_responses.py | 47 +++++++++++++ .../agent_framework_openai/_chat_client.py | 47 ++++++++++--- .../agent_framework_openai/_exceptions.py | 17 ++++- .../tests/openai/test_openai_chat_client.py | 32 +++++++++ 5 files changed, 198 insertions(+), 12 deletions(-) diff --git a/python/packages/foundry_hosting/agent_framework_foundry_hosting/_responses.py b/python/packages/foundry_hosting/agent_framework_foundry_hosting/_responses.py index 77d56c5e631..3f239b479ce 100644 --- a/python/packages/foundry_hosting/agent_framework_foundry_hosting/_responses.py +++ b/python/packages/foundry_hosting/agent_framework_foundry_hosting/_responses.py @@ -88,6 +88,59 @@ _MODEL_OUTPUT_KIND_KEY = "model_output_kind" _MODEL_OUTPUT_REFUSAL = "refusal" _HOSTED_RESPONSES_HISTORY_SOURCE_ID = "_foundry_responses_history" +_COST_CONTROL_HEADER_NAMES = ( + "x-ms-budget-cause", + "x-ms-budget-state", + "x-ms-remaining-budget", + "x-ms-consumed-budget", + "x-ms-budget-counter-key", + "x-ms-budget-window-reset", + "retry-after", + "x-ms-budget-consumed", + "x-ms-budget-remaining", +) +_COST_CONTROL_METADATA_KEY = "costControl" + + +def _cost_control_failure(ex: BaseException) -> tuple[str, dict[str, str] | None]: + """Extract allowlisted CostControl headers from an exception chain.""" + pending: list[BaseException] = [ex] + seen: set[int] = set() + while pending: + current = pending.pop() + if id(current) in seen: + continue + seen.add(id(current)) + + response_headers = getattr(current, "response_headers", None) + if isinstance(response_headers, Mapping): + available = {str(key).lower(): str(value) for key, value in response_headers.items()} + normalized: dict[str, str] = {} + for header_name in _COST_CONTROL_HEADER_NAMES: + value = available.get(header_name) + if value is None: + continue + candidate = {**normalized, header_name: value} + encoded = json.dumps(candidate, ensure_ascii=True, separators=(",", ":"), sort_keys=True) + if len(encoded) <= 512: + normalized = candidate + if normalized.get("x-ms-budget-cause", "").lower() == "budget_exceeded": + return "budget_exceeded", { + _COST_CONTROL_METADATA_KEY: json.dumps( + normalized, + ensure_ascii=True, + separators=(",", ":"), + sort_keys=True, + ) + } + + if isinstance(current.__cause__, BaseException): + pending.append(current.__cause__) + if isinstance(current.__context__, BaseException): + pending.append(current.__context__) + pending.extend(arg for arg in current.args if isinstance(arg, BaseException)) + + return "server_error", None def _is_refusal_text_content(content: Content) -> bool: @@ -1140,7 +1193,19 @@ def _emit_failure( except Exception: logger.exception("Error while closing streaming tracker after failure") message = str(ex) or type(ex).__name__ - yield response_event_stream.emit_failed(message=message, usage=tracker.usage if tracker is not None else None) + code, metadata = _cost_control_failure(ex) + if metadata is not None: + response_metadata = response_event_stream.response.get("metadata") + if not isinstance(response_metadata, dict): + response_metadata = {} + response_event_stream.response["metadata"] = response_metadata + if _COST_CONTROL_METADATA_KEY in response_metadata or len(response_metadata) < 16: + response_metadata.update(metadata) + yield response_event_stream.emit_failed( + code=code, + message=message, + usage=tracker.usage if tracker is not None else None, + ) # endregion ResponsesHostServer diff --git a/python/packages/foundry_hosting/tests/test_responses.py b/python/packages/foundry_hosting/tests/test_responses.py index dc712b1d223..f49f630e1bc 100644 --- a/python/packages/foundry_hosting/tests/test_responses.py +++ b/python/packages/foundry_hosting/tests/test_responses.py @@ -5210,6 +5210,53 @@ def run_failure(*args: Any, **kwargs: Any) -> ResponseStream[AgentResponseUpdate error: dict[str, Any] = body.get("error") or {} assert error.get("message") == "non-stream kaboom" + @pytest.mark.parametrize("stream", [False, True]) + async def test_budget_failure_emits_typed_error_and_cost_control_metadata(self, stream: bool) -> None: + class _BudgetExceededError(RuntimeError): + response_headers = { + "retry-after": "120", + "x-ms-budget-cause": "budget_exceeded", + "x-ms-budget-state": "degraded", + "x-ms-consumed-budget": "0.26", + "x-ms-remaining-budget": "0.00", + "x-request-id": "not-forwarded", + } + + async def _raise_budget_failure() -> AsyncIterator[AgentResponseUpdate]: + raise _BudgetExceededError("budget exceeded") + yield # pragma: no cover + + agent = _make_agent( + response=AgentResponse(messages=[Message(role="assistant", contents=[Content.from_text("hi")])]) + ) + agent.run = MagicMock( + return_value=ResponseStream( + _raise_budget_failure(), + finalizer=AgentResponse.from_updates, + ) + ) + server = _make_server(agent) + + resp = await _post(server, input_text="hello", stream=stream) + + assert resp.status_code == 200 + if stream: + events = _parse_sse_events(resp.text) + failed = next(event["data"]["response"] for event in events if event["event"] == "response.failed") + else: + failed = resp.json() + + assert failed["status"] == "failed" + assert failed["error"]["code"] == "budget_exceeded" + cost_control = json.loads(failed["metadata"]["costControl"]) + assert cost_control == { + "retry-after": "120", + "x-ms-budget-cause": "budget_exceeded", + "x-ms-budget-state": "degraded", + "x-ms-consumed-budget": "0.26", + "x-ms-remaining-budget": "0.00", + } + async def test_streaming_run_failure_emits_response_failed(self) -> None: async def _raise_stream() -> AsyncIterator[AgentResponseUpdate]: yield AgentResponseUpdate(contents=[Content.from_text("partial ")], role="assistant") diff --git a/python/packages/openai/agent_framework_openai/_chat_client.py b/python/packages/openai/agent_framework_openai/_chat_client.py index e1fe3cc5cf9..54c9e3b8262 100644 --- a/python/packages/openai/agent_framework_openai/_chat_client.py +++ b/python/packages/openai/agent_framework_openai/_chat_client.py @@ -73,7 +73,6 @@ validate_tool_mode, ) from agent_framework.exceptions import ( - ChatClientException, ChatClientInvalidRequestException, ) from agent_framework.observability import ChatTelemetryLayer @@ -110,7 +109,7 @@ from openai.types.responses.web_search_tool_param import WebSearchToolParam from pydantic import BaseModel, TypeAdapter, ValidationError -from ._exceptions import OpenAIContentFilterException +from ._exceptions import OpenAIContentFilterException, _OpenAIChatClientException from ._feature_usage import FeatureIndex from ._shared import ( AzureTokenProvider, @@ -795,19 +794,39 @@ async def _prepare_request( run_options = await self._prepare_options(messages, validated_options) return client, run_options, validated_options - def _handle_request_error(self, ex: Exception) -> NoReturn: + @staticmethod + def _response_headers_from_error(ex: Exception) -> dict[str, str]: + response = getattr(ex, "response", None) + headers = getattr(response, "headers", None) + if headers is None: + return {} + try: + return {str(key): str(value) for key, value in headers.items()} + except (AttributeError, TypeError, ValueError): + return {} + + def _handle_request_error( + self, + ex: Exception, + *, + response_headers: Mapping[str, str] | None = None, + ) -> NoReturn: """Convert exceptions to appropriate service exceptions. Always raises.""" if isinstance(ex, BadRequestError) and ex.code == "content_filter": raise OpenAIContentFilterException( f"{type(self)} service encountered a content error: {ex}", inner_exception=ex, ) from ex - raise ChatClientException( + captured_headers = self._response_headers_from_error(ex) + if response_headers is not None: + captured_headers.update((str(key), str(value)) for key, value in response_headers.items()) + raise _OpenAIChatClientException( maybe_append_azure_endpoint_guidance( f"{type(self)} service failed to complete the prompt: {ex}", azure_endpoint=self.azure_endpoint, ), inner_exception=ex, + response_headers=captured_headers, ) from ex @override @@ -847,6 +866,7 @@ async def _stream() -> AsyncIterable[ChatResponseUpdate]: validated_options = await self._validate_options(options) validated_options[_SEEN_FUNCTION_CALL_OUTPUT_IDS_OPTION] = seen_function_call_output_ids response_format = validated_options.get("response_format") + retrieve_response_headers: Mapping[str, str] | None = None try: raw_stream_response = await client.responses.with_raw_response.retrieve( continuation_token["response_id"], @@ -857,7 +877,8 @@ async def _stream() -> AsyncIterable[ChatResponseUpdate]: # experimental tracing) wrap the streaming response in objects that do not # proxy ``.headers``. Degrade gracefully so the served-model surfacing is # best-effort instead of crashing the whole call. - served_model = self._extract_served_model(getattr(raw_stream_response, "headers", None)) + retrieve_response_headers = getattr(raw_stream_response, "headers", None) + served_model = self._extract_served_model(retrieve_response_headers) async with _open_event_stream(raw_stream_response) as stream_response: async for chunk in stream_response: update = self._parse_chunk_from_openai( @@ -883,7 +904,7 @@ async def _stream() -> AsyncIterable[ChatResponseUpdate]: options.pop("continuation_token", None) yield update except Exception as ex: - self._handle_request_error(ex) + self._handle_request_error(ex, response_headers=retrieve_response_headers) else: ( client, @@ -894,6 +915,7 @@ async def _stream() -> AsyncIterable[ChatResponseUpdate]: if extra_headers is not None: run_options["extra_headers"] = dict(extra_headers) response_format = validated_options.get("response_format") + create_response_headers: Mapping[str, str] | None = None try: if "text_format" in run_options: # The SDK's ``responses.stream(text_format=...)`` helper preserves @@ -915,7 +937,8 @@ async def _stream() -> AsyncIterable[ChatResponseUpdate]: stream=True, **run_options ) # See note above on ``raw_stream_response.headers``. - served_model = self._extract_served_model(getattr(raw_create_response, "headers", None)) + create_response_headers = getattr(raw_create_response, "headers", None) + served_model = self._extract_served_model(create_response_headers) async with _open_event_stream(raw_create_response) as stream_response: async for chunk in stream_response: update = self._parse_chunk_from_openai( @@ -928,7 +951,7 @@ async def _stream() -> AsyncIterable[ChatResponseUpdate]: update.model = served_model yield update except Exception as ex: - self._handle_request_error(ex) + self._handle_request_error(ex, response_headers=create_response_headers) return ResponseStream(_stream(), finalizer=_finalize_with_captured_format) @@ -940,14 +963,16 @@ async def _get_response() -> ChatResponse: if self._FEATURE_USAGE_INDEX is not None: mark_feature_used(self._FEATURE_USAGE_INDEX) validated_options = await self._validate_options(options) + retrieve_response_headers: Mapping[str, str] | None = None try: raw_response = await client.responses.with_raw_response.retrieve( continuation_token["response_id"], extra_headers=extra_headers, ) + retrieve_response_headers = getattr(raw_response, "headers", None) response = raw_response.parse() except Exception as ex: - self._handle_request_error(ex) + self._handle_request_error(ex, response_headers=retrieve_response_headers) chat_response = self._parse_response_from_openai(response, options=validated_options) # See note above on ``raw_stream_response.headers``. served_model = self._extract_served_model(getattr(raw_response, "headers", None)) @@ -965,14 +990,16 @@ async def _get_response() -> ChatResponse: client, run_options, validated_options = await self._prepare_request(messages, options) if extra_headers is not None: run_options["extra_headers"] = dict(extra_headers) + create_response_headers: Mapping[str, str] | None = None try: if "text_format" in run_options: raw_response = await client.responses.with_raw_response.parse(stream=False, **run_options) else: raw_response = await client.responses.with_raw_response.create(stream=False, **run_options) + create_response_headers = getattr(raw_response, "headers", None) response = raw_response.parse() except Exception as ex: - self._handle_request_error(ex) + self._handle_request_error(ex, response_headers=create_response_headers) chat_response = self._parse_response_from_openai(response, options=validated_options) # See note above on ``raw_stream_response.headers``. served_model = self._extract_served_model(getattr(raw_response, "headers", None)) diff --git a/python/packages/openai/agent_framework_openai/_exceptions.py b/python/packages/openai/agent_framework_openai/_exceptions.py index 106430ab5f2..8b16213a062 100644 --- a/python/packages/openai/agent_framework_openai/_exceptions.py +++ b/python/packages/openai/agent_framework_openai/_exceptions.py @@ -2,11 +2,12 @@ from __future__ import annotations +from collections.abc import Mapping from dataclasses import dataclass from enum import Enum from typing import Any -from agent_framework.exceptions import ChatClientContentFilterException +from agent_framework.exceptions import ChatClientContentFilterException, ChatClientException from openai import BadRequestError @@ -54,6 +55,20 @@ class ContentFilterCodes(Enum): UNKNOWN = "Unknown" +class _OpenAIChatClientException(ChatClientException): + """Chat-client failure carrying response headers captured before streaming.""" + + def __init__( + self, + message: str, + *, + inner_exception: Exception, + response_headers: Mapping[str, str] | None = None, + ) -> None: + super().__init__(message, inner_exception=inner_exception) + self.response_headers = dict(response_headers or {}) + + @dataclass class OpenAIContentFilterException(ChatClientContentFilterException): """AI exception for an error from Azure OpenAI's content filter.""" diff --git a/python/packages/openai/tests/openai/test_openai_chat_client.py b/python/packages/openai/tests/openai/test_openai_chat_client.py index d30448c1e04..35021ab0f60 100644 --- a/python/packages/openai/tests/openai/test_openai_chat_client.py +++ b/python/packages/openai/tests/openai/test_openai_chat_client.py @@ -920,6 +920,38 @@ async def test_streaming_survives_telemetry_wrapped_raw_response() -> None: assert all(update.model == "test-model" for update in updates) +async def test_streaming_failure_preserves_raw_response_headers() -> None: + """Headers received before an SSE error remain available on the wrapped exception.""" + client = OpenAIChatClient(model="test-model", api_key="test-key") + + class _FailingStream(_FakeAsyncEventStream): + async def __anext__(self) -> object: + raise RuntimeError("budget exceeded") + + raw_stream = _FailingStream( + [], + headers={ + "x-ms-budget-cause": "budget_exceeded", + "x-ms-remaining-budget": "0.00", + }, + ) + + with ( + patch.object(client, "_prepare_request", new=AsyncMock(return_value=(client.client, {}, {}))), + patch.object(client.client.responses.with_raw_response, "create", new=AsyncMock(return_value=raw_stream)), + pytest.raises(ChatClientException) as exc_info, + ): + stream = _as_chat_response_stream( + client._inner_get_response(messages=[Message(role="user", contents=["Hi"])], options={}, stream=True) + ) + _ = [update async for update in stream] + + assert getattr(exc_info.value, "response_headers", {}) == { + "x-ms-budget-cause": "budget_exceeded", + "x-ms-remaining-budget": "0.00", + } + + async def test_streaming_accepts_raw_response_that_is_already_an_event_stream() -> None: """An object with no ``parse`` and no wrapped raw response is iterated directly.""" client = OpenAIChatClient(model="test-model", api_key="test-key")