Fix code review issues in WebSocket Responses service

- Use finally block in _disconnect to ensure state is always cleaned
  up, even if websocket.close() throws — prevents stale cancellation
  state (e.g. _cancel_pending_response) from polluting a new connection
- Catch ConnectionClosed in _drain_cancelled_response alongside
  TimeoutError — prevents _needs_drain from staying True and bricking
  the service on every subsequent inference attempt
- Fall back to OPENAI_API_KEY env var when api_key is not passed,
  since the WebSocket connection uses raw websockets (not the
  AsyncOpenAI client which handles this automatically)
- Use _clear_cancellation_state() instead of piecemeal resets where
  appropriate
This commit is contained in:
Paul Kompfner
2026-03-30 10:54:47 -04:00
parent b5683556d4
commit 0efef19d60

View File

@@ -9,6 +9,7 @@
import asyncio import asyncio
import hashlib import hashlib
import json import json
import os
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Any, Dict, List, Mapping, Optional from typing import Any, Dict, List, Mapping, Optional
@@ -162,7 +163,10 @@ class _BaseOpenAIResponsesLLMService(LLMService):
**kwargs, **kwargs,
) )
self._api_key = api_key # Resolve the API key from the environment if not provided. The
# AsyncOpenAI HTTP client does this automatically, but the WebSocket
# variant connects via raw websockets and needs the key explicitly.
self._api_key = api_key or os.environ.get("OPENAI_API_KEY")
self._service_tier = service_tier self._service_tier = service_tier
self._client = self._create_client( self._client = self._create_client(
api_key=api_key, api_key=api_key,
@@ -390,7 +394,7 @@ class OpenAIResponsesLLMService(_BaseOpenAIResponsesLLMService):
self._previous_response_output: Optional[list] = None self._previous_response_output: Optional[list] = None
# Response cancellation state # Response cancellation state
self._current_response_id: Optional[str] = None self._current_response_id: Optional[str] = None # ID of current non-cancelled response
self._cancel_pending_response: bool = False self._cancel_pending_response: bool = False
self._needs_drain: bool = False self._needs_drain: bool = False
@@ -448,12 +452,13 @@ class OpenAIResponsesLLMService(_BaseOpenAIResponsesLLMService):
await self.stop_all_metrics() await self.stop_all_metrics()
if self._websocket: if self._websocket:
await self._websocket.close() await self._websocket.close()
self._websocket = None except Exception as e:
await self.push_error(error_msg=f"Error disconnecting from WebSocket: {e}", exception=e)
finally:
self._websocket = None
self._clear_previous_response_state() self._clear_previous_response_state()
self._clear_cancellation_state() self._clear_cancellation_state()
self._disconnecting = False self._disconnecting = False
except Exception as e:
await self.push_error(error_msg=f"Error disconnecting from WebSocket: {e}", exception=e)
async def _reconnect(self): async def _reconnect(self):
"""Reconnect to the WebSocket, clearing previous_response_id state.""" """Reconnect to the WebSocket, clearing previous_response_id state."""
@@ -704,11 +709,10 @@ class OpenAIResponsesLLMService(_BaseOpenAIResponsesLLMService):
f"{self}: Cancelled response terminated with {event_type}" f"{self}: Cancelled response terminated with {event_type}"
f"connection is clean" f"connection is clean"
) )
self._needs_drain = False self._clear_cancellation_state()
return return
except asyncio.TimeoutError: except (asyncio.TimeoutError, ConnectionClosed) as e:
logger.warning(f"{self}: Timed out draining cancelled response — reconnecting") logger.warning(f"{self}: Error draining cancelled response: {e} — reconnecting")
self._needs_drain = False
await self._reconnect() await self._reconnect()
# -- frame processing ----------------------------------------------------- # -- frame processing -----------------------------------------------------