mirror of
https://github.com/pchuan98/codex.git
synced 2026-07-01 00:31:56 +08:00
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:
+72
-55
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user