Skip to content

Commit cd62899

Browse files
committed
fix: address review comments
add tests, make disconnect safe. Signed-off-by: Charlie Doern <cdoern@redhat.com>
1 parent 4a73cad commit cd62899

3 files changed

Lines changed: 151 additions & 11 deletions

File tree

src/ogx_api/responses/fastapi_routes.py

Lines changed: 21 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -238,16 +238,24 @@ async def _send_ws_error(
238238
message: str,
239239
param: str | None = None,
240240
) -> None:
241-
"""Send a WebSocket error envelope (matches the OpenResponses error event)."""
242-
await websocket.send_text(
243-
json.dumps(
244-
{
245-
"type": "error",
246-
"status": status,
247-
"error": {"code": code, "message": message, "param": param},
248-
}
241+
"""Send a WebSocket error envelope (matches the OpenResponses error event).
242+
243+
Sending is best-effort: if the client has already disconnected (a common
244+
cause of the failure we are reporting), the send raises and is suppressed so
245+
it does not mask the original error or produce a spurious traceback.
246+
"""
247+
try:
248+
await websocket.send_text(
249+
json.dumps(
250+
{
251+
"type": "error",
252+
"status": status,
253+
"error": {"code": code, "message": message, "param": param},
254+
}
255+
)
249256
)
250-
)
257+
except Exception:
258+
logger.debug("Failed to send WebSocket error envelope; client likely disconnected")
251259

252260

253261
async def _handle_ws_responses_turn(
@@ -317,6 +325,10 @@ async def _handle_ws_responses_turn(
317325
async for event in result:
318326
await websocket.send_text(event.model_dump_json())
319327
event_type = getattr(event, "type", None)
328+
# An incomplete response (e.g. truncated at max_output_tokens) is a
329+
# successful terminal state and remains continuable, matching the
330+
# HTTP previous_response_id path which can continue any stored
331+
# terminal response. Only response.failed is treated as a failure.
320332
if event_type in ("response.completed", "response.incomplete"):
321333
final_response = getattr(event, "response", None)
322334
elif event_type == "response.failed":

tests/unit/core/routers/test_responses_router.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44
# This source code is licensed under the terms described in the LICENSE file in
55
# the root directory of this source tree.
66

7+
import json
78
from unittest.mock import AsyncMock
89

910
import httpx
@@ -552,8 +553,6 @@ def test_create_response_form_urlencoded_with_json_encoded_complex_fields():
552553
assert router is not None
553554
app.include_router(router)
554555

555-
import json
556-
557556
tools_json = json.dumps([{"type": "web_search_preview"}])
558557
client = TestClient(app, raise_server_exceptions=False)
559558
resp = client.post(
Lines changed: 129 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,129 @@
1+
# Copyright (c) The OGX Contributors.
2+
# All rights reserved.
3+
#
4+
# This source code is licensed under the terms described in the LICENSE file in
5+
# the root directory of this source tree.
6+
7+
import json
8+
from unittest.mock import AsyncMock
9+
10+
from fastapi import FastAPI
11+
from starlette.testclient import TestClient
12+
13+
from ogx.core.server.fastapi_router_registry import build_fastapi_router
14+
from ogx_api import Api, Responses
15+
from ogx_api.openai_responses import (
16+
OpenAIResponseObject,
17+
OpenAIResponseObjectStreamResponseCompleted,
18+
OpenAIResponseObjectStreamResponseCreated,
19+
)
20+
21+
# WebSocket transport tests
22+
23+
24+
def _ws_app(impl: Responses) -> FastAPI:
25+
app = FastAPI()
26+
router = build_fastapi_router(Api.responses, impl)
27+
assert router is not None
28+
app.include_router(router)
29+
return app
30+
31+
32+
def _ws_response(response_id: str, status: str = "completed") -> OpenAIResponseObject:
33+
return OpenAIResponseObject(
34+
id=response_id,
35+
created_at=1234567890,
36+
model="test-model",
37+
object="response",
38+
output=[],
39+
status=status,
40+
store=False,
41+
)
42+
43+
44+
def test_websocket_invalid_json_returns_error_envelope():
45+
impl = AsyncMock(spec=Responses)
46+
client = TestClient(_ws_app(impl))
47+
48+
with client.websocket_connect("/v1/responses") as ws:
49+
ws.send_text("this is not json")
50+
event = ws.receive_json()
51+
52+
assert event["type"] == "error"
53+
assert event["error"]["code"] == "invalid_json"
54+
impl.create_openai_response.assert_not_called()
55+
56+
57+
def test_websocket_validation_error_returns_invalid_request():
58+
impl = AsyncMock(spec=Responses)
59+
client = TestClient(_ws_app(impl))
60+
61+
with client.websocket_connect("/v1/responses") as ws:
62+
# Missing required model/input fields.
63+
ws.send_text(json.dumps({"type": "response.create"}))
64+
event = ws.receive_json()
65+
66+
assert event["type"] == "error"
67+
assert event["error"]["code"] == "invalid_request"
68+
impl.create_openai_response.assert_not_called()
69+
70+
71+
def test_websocket_unknown_previous_response_not_found():
72+
impl = AsyncMock(spec=Responses)
73+
client = TestClient(_ws_app(impl))
74+
75+
with client.websocket_connect("/v1/responses") as ws:
76+
ws.send_text(
77+
json.dumps(
78+
{
79+
"type": "response.create",
80+
"model": "test",
81+
"store": False,
82+
"previous_response_id": "resp_does_not_exist",
83+
"input": "continue please",
84+
}
85+
)
86+
)
87+
event = ws.receive_json()
88+
89+
assert event["type"] == "error"
90+
assert event["status"] == 404
91+
assert event["error"]["code"] == "previous_response_not_found"
92+
# No inference is attempted for a connection-local cache miss.
93+
impl.create_openai_response.assert_not_called()
94+
95+
96+
def test_websocket_impl_exception_returns_server_error():
97+
impl = AsyncMock(spec=Responses)
98+
impl.create_openai_response.side_effect = RuntimeError("boom")
99+
client = TestClient(_ws_app(impl))
100+
101+
with client.websocket_connect("/v1/responses") as ws:
102+
ws.send_text(json.dumps({"type": "response.create", "model": "test", "input": "hi"}))
103+
event = ws.receive_json()
104+
105+
assert event["type"] == "error"
106+
assert event["error"]["code"] == "server_error"
107+
108+
109+
def test_websocket_streams_events_until_terminal():
110+
impl = AsyncMock(spec=Responses)
111+
112+
async def _stream():
113+
yield OpenAIResponseObjectStreamResponseCreated(response=_ws_response("resp_ws_1"), sequence_number=0)
114+
yield OpenAIResponseObjectStreamResponseCompleted(response=_ws_response("resp_ws_1"), sequence_number=1)
115+
116+
impl.create_openai_response.return_value = _stream()
117+
client = TestClient(_ws_app(impl))
118+
119+
with client.websocket_connect("/v1/responses") as ws:
120+
ws.send_text(json.dumps({"type": "response.create", "model": "test", "input": "hi"}))
121+
first = ws.receive_json()
122+
second = ws.receive_json()
123+
124+
assert first["type"] == "response.created"
125+
assert second["type"] == "response.completed"
126+
# The HTTP-only streaming discriminator must not leak into the request.
127+
sent_request = impl.create_openai_response.call_args.args[0]
128+
assert sent_request.stream is True
129+
assert not hasattr(sent_request, "type") or "type" not in sent_request.model_dump()

0 commit comments

Comments
 (0)