Skip to content
Merged
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
5 changes: 5 additions & 0 deletions src/agents/models/openai_responses.py
Original file line number Diff line number Diff line change
Expand Up @@ -262,6 +262,7 @@ class _ResponseStreamWithRequestId:
"response.completed",
"response.failed",
"response.incomplete",
"error",
"response.error",
}

Expand Down Expand Up @@ -778,6 +779,10 @@ async def stream_response(
trace_include_sensitive_data=tracing.include_data(),
)
yield chunk
if terminal_failure_error is not None:
# Close explicitly rather than advancing a failed stream: transport
# teardown during iteration must not replace the provider's error.
break
except asyncio.CancelledError:
close_stream_in_background = True
self._schedule_async_iterator_close(stream)
Expand Down
140 changes: 126 additions & 14 deletions tests/models/test_openai_responses.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
import copy
import json
import logging
from contextlib import nullcontext
from dataclasses import asdict, fields, replace
from types import SimpleNamespace
from typing import Any, cast
Expand Down Expand Up @@ -908,7 +909,10 @@ def __init__(self):

@pytest.mark.allow_call_model_methods
@pytest.mark.asyncio
async def test_stream_response_ignores_streaming_context_exit_failure_after_terminal_event():
@pytest.mark.parametrize("event_type", ["response.completed", "error"])
async def test_stream_response_ignores_streaming_context_exit_failure_after_terminal_event(
event_type,
):
class DummyHTTPStream:
def __init__(self):
self._yielded = False
Expand All @@ -920,6 +924,14 @@ async def __anext__(self):
if self._yielded:
raise StopAsyncIteration
self._yielded = True
if event_type == "error":
return ResponseErrorEvent(
type="error",
code="server_error",
message="synthetic provider failure",
param=None,
sequence_number=0,
)
return ResponseCompletedEvent(
type="response.completed",
response=get_response_obj([], response_id="resp-stream-request-id"),
Expand Down Expand Up @@ -958,23 +970,123 @@ def __init__(self):

model = OpenAIResponsesModel(model="gpt-4", openai_client=DummyResponsesClient()) # type: ignore[arg-type]

events: list[ResponseCompletedEvent] = []
async for event in model.stream_response(
system_instructions=None,
input="hi",
model_settings=ModelSettings(),
tools=[],
output_schema=None,
handoffs=[],
tracing=ModelTracing.DISABLED,
):
assert isinstance(event, ResponseCompletedEvent)
events.append(event)
events = []
expected_error = (
pytest.raises(ModelBehaviorError, match="synthetic provider failure")
if event_type == "error"
else nullcontext()
)
with expected_error:
async for event in model.stream_response(
system_instructions=None,
input="hi",
model_settings=ModelSettings(),
tools=[],
output_schema=None,
handoffs=[],
tracing=ModelTracing.DISABLED,
):
events.append(event)

assert len(events) == 1
assert [event.type for event in events] == [event_type]
assert aexit_calls == [(None, None, None)]


@pytest.mark.asyncio
async def test_response_stream_preserves_error_event_when_cleanup_fails() -> None:
event = ResponseErrorEvent(
type="error",
code="server_error",
message="synthetic provider failure",
param=None,
sequence_number=0,
)
cleanup_calls = 0

async def events():
yield event

async def cleanup():
nonlocal cleanup_calls
cleanup_calls += 1
raise RuntimeError("transport close failed")

stream = _ResponseStreamWithRequestId(events(), request_id=None, cleanup=cleanup)
assert [item async for item in stream] == [event]
assert cleanup_calls == 1


@pytest.mark.allow_call_model_methods
@pytest.mark.asyncio
@pytest.mark.parametrize("event_type", ["error", "response.failed", "response.incomplete", None])
async def test_stream_response_preserves_terminal_failure_when_http_close_fails(
event_type: str | None,
) -> None:
# Exercise real SSE parsing and transport teardown, which a model fake would bypass.
close_calls = 0
if event_type == "error":
payload = ResponseErrorEvent(
type="error",
code="server_error",
message="synthetic provider failure",
param=None,
sequence_number=0,
).model_dump()
expected_message = "code=server_error; message=synthetic provider failure"
elif event_type is not None:
status = event_type.removeprefix("response.")
payload = {
"type": event_type,
"sequence_number": 0,
"response": _response_with_terminal_status(status).model_dump(),
}
expected_message = f"status={status}"
else:
payload = None
expected_message = "transport close failed"

class FailingCloseStream(httpx2.AsyncByteStream):
async def __aiter__(self):
if payload is not None:
yield f"event: {event_type}\ndata: {json.dumps(payload)}\n\n".encode()
yield b"data: [DONE]\n\n"

async def aclose(self):
nonlocal close_calls
close_calls += 1
raise RuntimeError("transport close failed")

async def handler(request: httpx2.Request) -> httpx2.Response:
return httpx2.Response(
200,
headers={"content-type": "text/event-stream"},
stream=FailingCloseStream(),
request=request,
)

events = []
async with httpx2.AsyncClient(transport=httpx2.MockTransport(handler)) as http_client:
model = OpenAIResponsesModel(
model="gpt-4",
openai_client=AsyncOpenAI(api_key="test-key", http_client=http_client),
)
expected_exception = ModelBehaviorError if event_type is not None else RuntimeError
with pytest.raises(expected_exception, match=expected_message):
async for event in model.stream_response(
system_instructions=None,
input="hi",
model_settings=ModelSettings(),
tools=[],
output_schema=None,
handoffs=[],
tracing=ModelTracing.DISABLED,
):
events.append(event)

assert [event.type for event in events] == ([event_type] if event_type is not None else [])
assert close_calls == 1


@pytest.mark.allow_call_model_methods
@pytest.mark.asyncio
async def test_stream_response_close_closes_inner_http_stream_with_async_close(monkeypatch):
Expand Down
Loading