Do not resend output items in incremental websockets connections (#11383)

In the incremental websocket output items are already part of the
context, no need to send them again and duplicate.
This commit is contained in:
pakrym-oai
2026-02-10 19:38:08 -08:00
committed by GitHub
parent cc8c293378
commit 4473147985
3 changed files with 365 additions and 161 deletions
+72 -55
View File
@@ -167,8 +167,7 @@ pub struct ModelClientSession {
client: ModelClient,
connection: Option<ApiWebSocketConnection>,
websocket_last_request: Option<ResponsesApiRequest>,
websocket_last_response_id: Option<String>,
websocket_last_response_id_rx: Option<oneshot::Receiver<String>>,
websocket_last_response_rx: Option<oneshot::Receiver<LastResponse>>,
/// Turn state for sticky routing.
///
/// This is an `OnceLock` that stores the turn state value received from the server
@@ -182,6 +181,12 @@ pub struct ModelClientSession {
turn_state: Arc<OnceLock<String>>,
}
#[derive(Debug, Clone)]
struct LastResponse {
response_id: String,
items_added: Vec<ResponseItem>,
}
enum WebsocketStreamOutcome {
Stream(ResponseStream),
FallbackToHttp,
@@ -231,8 +236,7 @@ impl ModelClient {
client: self.clone(),
connection: None,
websocket_last_request: None,
websocket_last_response_id: None,
websocket_last_response_id_rx: None,
websocket_last_response_rx: None,
turn_state: Arc::new(OnceLock::new()),
}
}
@@ -531,10 +535,15 @@ impl ModelClientSession {
}
}
fn get_incremental_items(&self, request: &ResponsesApiRequest) -> Option<Vec<ResponseItem>> {
fn get_incremental_items(
&self,
request: &ResponsesApiRequest,
last_response: Option<&LastResponse>,
) -> Option<Vec<ResponseItem>> {
// Checks whether the current request is an incremental append to the previous request.
// We only append when non-input request fields are unchanged and `input` is a strict
// extension of the previous input.
// extension of the previous known input. Server-returned output items are treated as part
// of the baseline so we do not resend them.
let previous_request = self.websocket_last_request.as_ref()?;
let mut previous_without_input = previous_request.clone();
previous_without_input.input.clear();
@@ -544,38 +553,29 @@ impl ModelClientSession {
return None;
}
let previous_len = previous_request.input.len();
if previous_len > 0
&& request.input.starts_with(&previous_request.input)
&& previous_len < request.input.len()
let mut baseline = previous_request.input.clone();
if let Some(last_response) = last_response {
baseline.extend(last_response.items_added.clone());
}
let baseline_len = baseline.len();
if baseline_len > 0
&& request.input.starts_with(&baseline)
&& baseline_len < request.input.len()
{
Some(request.input[previous_len..].to_vec())
Some(request.input[baseline_len..].to_vec())
} else {
None
}
}
fn refresh_websocket_last_response_id(&mut self) {
if let Some(mut receiver) = self.websocket_last_response_id_rx.take() {
match receiver.try_recv() {
Ok(response_id) if !response_id.is_empty() => {
self.websocket_last_response_id = Some(response_id);
}
Ok(_) | Err(TryRecvError::Closed) => {
self.websocket_last_response_id = None;
}
Err(TryRecvError::Empty) => {
self.websocket_last_response_id_rx = Some(receiver);
}
}
}
}
fn websocket_previous_response_id(&mut self) -> Option<String> {
self.refresh_websocket_last_response_id();
self.websocket_last_response_id
.clone()
.filter(|id| !id.is_empty())
fn get_last_response(&mut self) -> Option<LastResponse> {
self.websocket_last_response_rx
.take()
.and_then(|mut receiver| match receiver.try_recv() {
Ok(last_response) => Some(last_response),
Err(TryRecvError::Closed) | Err(TryRecvError::Empty) => None,
})
}
fn prepare_websocket_request(
@@ -583,11 +583,15 @@ impl ModelClientSession {
payload: ResponseCreateWsRequest,
request: &ResponsesApiRequest,
) -> ResponsesWsRequest {
let last_response = self.get_last_response();
let responses_websockets_v2_enabled = self.client.responses_websockets_v2_enabled();
let incremental_items = self.get_incremental_items(request);
let incremental_items = self.get_incremental_items(request, last_response.as_ref());
if let Some(append_items) = incremental_items {
if responses_websockets_v2_enabled
&& let Some(previous_response_id) = self.websocket_previous_response_id()
&& let Some(previous_response_id) = last_response
.as_ref()
.map(|last_response| last_response.response_id.clone())
.filter(|id| !id.is_empty())
{
let payload = ResponseCreateWsRequest {
previous_response_id: Some(previous_response_id),
@@ -660,8 +664,7 @@ impl ModelClientSession {
if needs_new {
self.websocket_last_request = None;
self.websocket_last_response_id = None;
self.websocket_last_response_id_rx = None;
self.websocket_last_response_rx = None;
let turn_state = options
.turn_state
.clone()
@@ -716,7 +719,8 @@ impl ModelClientSession {
self.client.state.provider.stream_idle_timeout(),
)
.map_err(map_api_error)?;
return Ok(map_response_stream(stream, otel_manager.clone()));
let (stream, _last_request_rx) = map_response_stream(stream, otel_manager.clone());
return Ok(stream);
}
let auth_manager = self.client.state.auth_manager.clone();
@@ -747,7 +751,8 @@ impl ModelClientSession {
match stream_result {
Ok(stream) => {
return Ok(map_response_stream(stream, otel_manager.clone()));
let (stream, _) = map_response_stream(stream, otel_manager.clone());
return Ok(stream);
}
Err(ApiError::Transport(
unauthorized_transport @ TransportError::Http { status, .. },
@@ -829,22 +834,11 @@ impl ModelClientSession {
.await
.map_err(map_api_error)?;
self.websocket_last_request = Some(request);
let (last_response_id_sender, last_response_id_receiver) = oneshot::channel();
self.websocket_last_response_id_rx = Some(last_response_id_receiver);
let mut last_response_id_sender = Some(last_response_id_sender);
let stream_result = stream_result.inspect(move |event| {
if let Ok(ResponseEvent::Completed { response_id, .. }) = event
&& !response_id.is_empty()
&& let Some(sender) = last_response_id_sender.take()
{
let _ = sender.send(response_id.clone());
}
});
let (stream, last_request_rx) =
map_response_stream(stream_result, otel_manager.clone());
self.websocket_last_response_rx = Some(last_request_rx);
return Ok(WebsocketStreamOutcome::Stream(map_response_stream(
stream_result,
otel_manager.clone(),
)));
return Ok(WebsocketStreamOutcome::Stream(stream));
}
}
@@ -942,6 +936,7 @@ impl ModelClientSession {
self.connection = None;
self.websocket_last_request = None;
self.websocket_last_response_rx = None;
}
activated
}
@@ -986,7 +981,10 @@ fn build_responses_headers(
headers
}
fn map_response_stream<S>(api_stream: S, otel_manager: OtelManager) -> ResponseStream
fn map_response_stream<S>(
api_stream: S,
otel_manager: OtelManager,
) -> (ResponseStream, oneshot::Receiver<LastResponse>)
where
S: futures::Stream<Item = std::result::Result<ResponseEvent, ApiError>>
+ Unpin
@@ -994,12 +992,25 @@ where
+ 'static,
{
let (tx_event, rx_event) = mpsc::channel::<Result<ResponseEvent>>(1600);
let (tx_last_response, rx_last_response) = oneshot::channel::<LastResponse>();
tokio::spawn(async move {
let mut logged_error = false;
let mut tx_last_response = Some(tx_last_response);
let mut items_added: Vec<ResponseItem> = Vec::new();
let mut api_stream = api_stream;
while let Some(event) = api_stream.next().await {
match event {
Ok(ResponseEvent::OutputItemDone(item)) => {
items_added.push(item.clone());
if tx_event
.send(Ok(ResponseEvent::OutputItemDone(item)))
.await
.is_err()
{
return;
}
}
Ok(ResponseEvent::Completed {
response_id,
token_usage,
@@ -1013,6 +1024,12 @@ where
usage.total_tokens,
);
}
if let Some(sender) = tx_last_response.take() {
let _ = sender.send(LastResponse {
response_id: response_id.clone(),
items_added: std::mem::take(&mut items_added),
});
}
if tx_event
.send(Ok(ResponseEvent::Completed {
response_id,
@@ -1043,7 +1060,7 @@ where
}
});
ResponseStream { rx_event }
(ResponseStream { rx_event }, rx_last_response)
}
/// Handles a 401 response by optionally refreshing ChatGPT tokens once.