Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
167 changes: 151 additions & 16 deletions crates/gateway/src/proxy/responses_ws.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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<Message> = VecDeque::new();
// The connection's most recent response.
let mut last: Option<Last> = None;
loop {
let message = match queued.pop_front() {
Some(m) => m,
Expand All @@ -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
Expand All @@ -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<Vec<u8>, String> {
fn request_of(text: &str) -> Result<serde_json::Map<String, Value>, 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::<Value>(text) else {
Expand All @@ -135,10 +156,75 @@ fn request_of(text: &str) -> Result<Vec<u8>, 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<Value>,
}

/// 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<Value>, response: &Value) -> Option<Last> {
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<String, Value>) -> Vec<Value> {
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<String, Value>,
last: Option<&Last>,
) -> Option<Vec<Value>> {
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.
Expand All @@ -147,9 +233,10 @@ async fn relay(
tx: &mut futures::stream::SplitSink<WebSocket, Message>,
rx: &mut futures::stream::SplitStream<WebSocket>,
queued: &mut VecDeque<Message>,
) -> Turn {
) -> (Turn, Option<Value>) {
let mut body = resp.into_body().into_data_stream();
let mut decoder = Decoder::default();
let mut completed = None;
loop {
tokio::select! {
chunk = body.next() => {
Expand All @@ -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::<Value>(&f.data).is_err() {
let Ok(mut event) = serde_json::from_str::<Value>(&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(_)) => {}
},
Expand Down Expand Up @@ -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<String, Value> {
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]
Expand Down
62 changes: 62 additions & 0 deletions crates/test-support/tests/gateway_responses_ws.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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}");
}
Loading