diff --git a/crates/gateway/src/proxy/responses_ws.rs b/crates/gateway/src/proxy/responses_ws.rs index 8180af0a..41a1e09a 100644 --- a/crates/gateway/src/proxy/responses_ws.rs +++ b/crates/gateway/src/proxy/responses_ws.rs @@ -22,10 +22,21 @@ //! logged exactly like the HTTP request it stands for, and the API key was //! checked on the upgrade by the same middleware. //! -//! What this does not carry over: a connection-local store of the previous -//! response. Each turn goes upstream as its own request, so a -//! `previous_response_id` works only against an upstream that stored that -//! response. +//! # The previous response is kept on the connection +//! +//! OpenAI's socket mode keeps the connection's most recent response in +//! memory, so a turn can name it in `previous_response_id` even with +//! `store: false` — which is how Codex runs. The desktop gateway gets that +//! for free by piping the whole socket to one upstream socket. Here each +//! turn is a separate upstream request, possibly to an upstream in another +//! format with no such store, so the connection keeps it instead: the +//! conversation so far (the turn's full input and the output items of its +//! `response.completed`). A turn whose `previous_response_id` names it goes +//! upstream with that history written into `input` and no +//! `previous_response_id` — valid against any upstream, and billed as +//! what it is. Like OpenAI, only the most recent response is kept; a +//! `previous_response_id` naming anything else goes upstream as sent, for +//! an upstream that stored it. //! //! A refusal — before the stream opens or during it — is a //! `response.failed` frame, the event Responses clients dispatch on; the @@ -74,6 +85,8 @@ async fn serve( let (mut tx, mut rx) = socket.split(); // Turns the client sent while one was still running. let mut queued: VecDeque = VecDeque::new(); + // The connection's most recent response. + let mut last: Option = None; loop { let message = match queued.pop_front() { Some(m) => m, @@ -90,20 +103,28 @@ async fn serve( Message::Ping(_) | Message::Pong(_) => continue, }; let turn = match body { - Ok(body) => { + Ok(mut body) => { + let history = continue_from(&mut body, last.as_ref()); let answer = generate( state.clone(), headers.clone(), identity.clone(), None, - Bytes::from(body), + Bytes::from(Value::Object(body).to_string()), RESPONSES, "/v1/responses", None, ) .await; match answer { - Ok(resp) => relay(resp, &mut tx, &mut rx, &mut queued).await, + Ok(resp) => { + let (turn, completed) = relay(resp, &mut tx, &mut rx, &mut queued).await; + // A failed turn leaves the chain where it was. + if let (Some(history), Some(done)) = (history, completed) { + last = Last::of(history, &done).or(last); + } + turn + } Err(e) => { let e = e.error(); refuse(&mut tx, e.status_code(), &e.to_string()).await @@ -120,7 +141,7 @@ async fn serve( } /// The request body a `response.create` frame stands for, as a stream. -fn request_of(text: &str) -> Result, String> { +fn request_of(text: &str) -> Result, String> { let not_create = || "Only response.create messages are accepted on this connection.".to_string(); let Ok(Value::Object(mut frame)) = serde_json::from_str::(text) else { @@ -135,10 +156,75 @@ fn request_of(text: &str) -> Result, String> { _ => frame, }; body.insert("stream".into(), Value::Bool(true)); - Ok(Value::Object(body).to_string().into_bytes()) + Ok(body) +} + +/// A response this connection produced, as the conversation up to and +/// including it: every input item the turn went upstream with, then the +/// response's output items. +struct Last { + id: String, + items: Vec, } -/// Send the turn's SSE to the client, a frame per event. +impl Last { + /// From a turn's full input and its `response.completed` response. + fn of(mut items: Vec, response: &Value) -> Option { + let id = response.get("id")?.as_str()?.to_string(); + for item in response.get("output")?.as_array()? { + let mut item = item.clone(); + // Output item ids refer to the upstream's store, which a + // `store: false` turn never wrote to; sent back as input they + // would be looked up and not found. `call_id` is what ties a + // tool result to its call, and stays. + if let Some(o) = item.as_object_mut() { + o.remove("id"); + o.remove("status"); + } + items.push(item); + } + Some(Last { id, items }) + } +} + +/// The turn's `input` as a list of items: a bare string is one user +/// message. +fn input_items(body: &serde_json::Map) -> Vec { + match body.get("input") { + Some(Value::Array(items)) => items.clone(), + Some(Value::String(text)) => { + vec![serde_json::json!({"type": "message", "role": "user", "content": text})] + } + _ => Vec::new(), + } +} + +/// Continue from the connection's last response if the turn names it: +/// its history goes into `input` and `previous_response_id` goes away. +/// +/// Returns the turn's whole conversation, to keep once it completes — +/// `None` when the turn still points at an earlier response this +/// connection does not have, so its history is not all here. +fn continue_from( + body: &mut serde_json::Map, + last: Option<&Last>, +) -> Option> { + let previous = body.get("previous_response_id").and_then(Value::as_str); + let mut items = match (previous, last) { + (None, _) => Vec::new(), + (Some(p), Some(l)) if p == l.id => l.items.clone(), + (Some(_), _) => return None, + }; + items.extend(input_items(body)); + if previous.is_some() { + body.remove("previous_response_id"); + body.insert("input".into(), Value::Array(items.clone())); + } + Some(items) +} + +/// Send the turn's SSE to the client, a frame per event, and hand back the +/// `response` of its `response.completed`, if it got that far. /// /// Dropping the response body cancels the turn: the pipeline's tail then /// records it as cancelled by the client, as for an HTTP stream. @@ -147,9 +233,10 @@ async fn relay( tx: &mut futures::stream::SplitSink, rx: &mut futures::stream::SplitStream, queued: &mut VecDeque, -) -> Turn { +) -> (Turn, Option) { let mut body = resp.into_body().into_data_stream(); let mut decoder = Decoder::default(); + let mut completed = None; loop { tokio::select! { chunk = body.next() => { @@ -160,19 +247,22 @@ async fn relay( for f in frames { // Every Responses event is a JSON object; nothing else // is a frame. - if serde_json::from_str::(&f.data).is_err() { + let Ok(mut event) = serde_json::from_str::(&f.data) else { continue; + }; + if event.get("type").and_then(Value::as_str) == Some("response.completed") { + completed = event.get_mut("response").map(Value::take); } if tx.send(Message::Text(f.data.into())).await.is_err() { - return Turn::Gone; + return (Turn::Gone, None); } } if end { - return Turn::Done; + return (Turn::Done, completed); } } incoming = rx.next() => match incoming { - Some(Ok(Message::Close(_))) | Some(Err(_)) | None => return Turn::Gone, + Some(Ok(Message::Close(_))) | Some(Err(_)) | None => return (Turn::Gone, None), Some(Ok(m @ (Message::Text(_) | Message::Binary(_)))) => queued.push_back(m), Some(Ok(_)) => {} }, @@ -204,7 +294,52 @@ mod tests { use super::*; fn body(text: &str) -> Value { - serde_json::from_slice(&request_of(text).unwrap()).unwrap() + Value::Object(request_of(text).unwrap()) + } + + fn create(v: Value) -> serde_json::Map { + request_of(&v.to_string()).unwrap() + } + + #[test] + fn a_turn_naming_the_last_response_carries_its_history() { + let mut first = + create(serde_json::json!({"type": "response.create", "model": "m", "input": "one"})); + let history = continue_from(&mut first, None).unwrap(); + let done = serde_json::json!({"id": "resp_1", "output": [ + {"id": "msg_1", "type": "message", "role": "assistant", "status": "completed", + "content": [{"type": "output_text", "text": "hi"}]}, + {"id": "fc_1", "type": "function_call", "call_id": "call_1", "name": "f", "arguments": "{}"}, + ]}); + let last = Last::of(history, &done).unwrap(); + + let mut second = create(serde_json::json!({"type": "response.create", "model": "m", + "previous_response_id": "resp_1", + "input": [{"type": "function_call_output", "call_id": "call_1", "output": "ok"}]})); + let kept = continue_from(&mut second, Some(&last)).unwrap(); + assert!(second.get("previous_response_id").is_none()); + let input = second["input"].as_array().unwrap(); + assert_eq!(input.len(), 4); + assert_eq!(input[0]["content"], "one"); + assert_eq!(input[1]["role"], "assistant"); + assert!(input[1].get("id").is_none() && input[1].get("status").is_none()); + assert_eq!(input[2]["call_id"], "call_1"); + assert!(input[2].get("id").is_none()); + assert_eq!(input[3]["type"], "function_call_output"); + assert_eq!(kept.len(), 4); + } + + #[test] + fn a_response_this_connection_does_not_have_goes_upstream_as_sent() { + let last = Last { + id: "resp_1".into(), + items: vec![serde_json::json!({"x": 1})], + }; + let mut turn = create(serde_json::json!({"type": "response.create", "model": "m", + "previous_response_id": "resp_0", "input": "two"})); + assert!(continue_from(&mut turn, Some(&last)).is_none()); + assert_eq!(turn["previous_response_id"], "resp_0"); + assert_eq!(turn["input"], "two"); } #[test] diff --git a/crates/test-support/tests/gateway_responses_ws.rs b/crates/test-support/tests/gateway_responses_ws.rs index 9f6165fe..72fcf75f 100644 --- a/crates/test-support/tests/gateway_responses_ws.rs +++ b/crates/test-support/tests/gateway_responses_ws.rs @@ -184,3 +184,65 @@ async fn limits_apply_per_turn_and_a_refusal_keeps_the_connection() { assert_eq!(refused.last().unwrap()["type"], "response.failed"); socket.send(Message::Ping(vec![1])).await.unwrap(); } + +/// OpenAI's socket mode keeps the connection's last response, so a turn +/// can continue from it with `store: false` (how Codex runs). The +/// connection keeps it here: the next turn goes upstream — to a Chat +/// upstream, which has no such store — with the whole conversation. +#[ignore = "integration test — run via `make test-it`"] +#[tokio::test] +async fn a_turn_continues_from_the_connections_last_response() { + let app = TestApp::spawn().await; + let upstream = MockProvider::openai_chat_stream_ok("ws-chain").await; + let (_, key) = seed(&app, &upstream.uri(), "ws-chain").await; + + let mut socket = connect(&app, &key).await; + socket + .send(Message::Text( + json!({"type": "response.create", "model": "ws-chain", "store": false, "input": "one"}) + .to_string(), + )) + .await + .unwrap(); + let first = turn(&mut socket).await; + let id = first.last().unwrap()["response"]["id"] + .as_str() + .expect("response id") + .to_string(); + + socket + .send(Message::Text( + json!({"type": "response.create", "model": "ws-chain", "store": false, + "previous_response_id": id, "input": "two"}) + .to_string(), + )) + .await + .unwrap(); + let second = turn(&mut socket).await; + assert_eq!( + second.last().unwrap()["type"], + "response.completed", + "{second:?}" + ); + + let sent = upstream.received_requests().await; + assert_eq!(sent.len(), 2); + let body: Value = sent[1].body_json().unwrap(); + let said: Vec<(&str, String)> = body["messages"] + .as_array() + .unwrap() + .iter() + .map(|m| { + let text = match &m["content"] { + Value::String(s) => s.clone(), + other => other.to_string(), + }; + (m["role"].as_str().unwrap(), text) + }) + .collect(); + assert_eq!(said.len(), 3, "{body}"); + assert_eq!(said[0], ("user", "one".to_string()), "{body}"); + assert_eq!(said[1].0, "assistant", "{body}"); + assert!(said[1].1.contains("hi there"), "{body}"); + assert_eq!(said[2], ("user", "two".to_string()), "{body}"); +}