Skip to content
Open
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
5 changes: 4 additions & 1 deletion apis/src/openai/responses/file_search_callout/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -719,6 +719,7 @@ fn continuation_state_fits(
for value in [
state.context_management.as_ref(),
state.conversation.as_ref(),
state.original_tool_choice.as_ref(),
state.previous_usage.as_ref(),
]
.into_iter()
Expand All @@ -735,6 +736,7 @@ fn continuation_state_fits(
.map(|(key, value)| key.len().saturating_add(value.len()))
.chain(state.include.iter().map(String::len))
.chain(state.previous_response_id.iter().map(String::len))
.chain(state.response_id.iter().map(String::len))
.chain(
state
.mcp_tool_map
Expand Down Expand Up @@ -825,7 +827,8 @@ fn mixed_tool_response_rejection() -> FilterAction {

/// Allow the model to answer after satisfying the first forced search call.
fn reset_tool_choice(state: &mut ResponsesState) {
state.tool_choice = Value::String("auto".to_owned());
let original = std::mem::replace(&mut state.tool_choice, Value::String("auto".to_owned()));
state.original_tool_choice.get_or_insert(original);
if let Some(request) = state.request_body.as_object_mut() {
request.remove("tool_choice");
}
Expand Down
1 change: 1 addition & 0 deletions apis/src/openai/responses/file_search_callout/tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -765,6 +765,7 @@ async fn forced_tool_choice_resets_after_search_execution() {
));
let state = ctx.extensions.get::<ResponsesState>().unwrap();
assert_eq!(state.tool_choice, "auto");
assert_eq!(state.original_tool_choice, Some(json!({"type":"file_search"})));
assert!(state.request_body.get("tool_choice").is_none());
}

Expand Down
6 changes: 4 additions & 2 deletions apis/src/openai/responses/rehydrate/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -139,7 +139,8 @@ impl RehydrateFilter {
};
let previous_tools = collect_mcp_tool_listings(&record);
let previous_usage = record.response_object.get("usage").filter(|u| !u.is_null()).cloned();
let state = build_state(parsed_body, stored, previous_tools, previous_usage);
let mut state = build_state(parsed_body, stored, previous_tools, previous_usage);
state.response_id = ctx.get_metadata("responses.response_id").map(ToOwned::to_owned);
write_previous_usage_metadata(ctx, state.previous_usage.as_ref());
ctx.extensions.insert(state);
debug!(previous_response_id = %prev_id, "previous response validated, state populated");
Expand Down Expand Up @@ -176,7 +177,8 @@ impl RehydrateFilter {
Ok(s) => s,
Err(action) => return Ok(action),
};
let state = build_state(parsed_body, stored, vec![], None);
let mut state = build_state(parsed_body, stored, vec![], None);
state.response_id = ctx.get_metadata("responses.response_id").map(ToOwned::to_owned);
write_previous_usage_metadata(ctx, state.previous_usage.as_ref());
ctx.extensions.insert(state);
debug!(conversation_id = %conv_id, "conversation rehydrated, state populated");
Expand Down
4 changes: 4 additions & 0 deletions apis/src/openai/responses/rehydrate/tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -181,6 +181,7 @@ async fn validates_previous_response_and_sets_metadata() {
let mut ctx = crate::test_utils::make_filter_context(&req);
ctx.extensions.insert(registry.clone());
ctx.set_metadata("openai_responses_format.format", "openai_responses");
ctx.set_metadata("responses.response_id", "resp_current");
let original = r#"{"model":"gpt-4.1","input":"What next?","previous_response_id":"resp_prev"}"#;
let mut body = Some(Bytes::from(original));

Expand Down Expand Up @@ -216,6 +217,7 @@ async fn validates_previous_response_and_sets_metadata() {
state.messages[2]["content"], "What next?",
"current input should be last"
);
assert_eq!(state.response_id.as_deref(), Some("resp_current"));
}

#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
Expand Down Expand Up @@ -1351,6 +1353,7 @@ async fn rehydrates_from_conversation_string_id() {
let mut ctx = crate::test_utils::make_filter_context(&req);
ctx.extensions.insert(registry.clone());
ctx.set_metadata("openai_responses_format.format", "openai_responses");
ctx.set_metadata("responses.response_id", "resp_conversation");
let mut body = Some(Bytes::from(
r#"{"model":"gpt-4.1","input":"turn two","conversation":"conv_abc"}"#,
));
Expand Down Expand Up @@ -1379,6 +1382,7 @@ async fn rehydrates_from_conversation_string_id() {
3,
"persisted_messages should mirror messages for conversation rehydration"
);
assert_eq!(state.response_id.as_deref(), Some("resp_conversation"));
}

#[tokio::test]
Expand Down
42 changes: 34 additions & 8 deletions apis/src/openai/responses/responses_to_chat_completions/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -282,7 +282,12 @@ impl HttpFilter for ResponsesToChatCompletionsFilter {
ctx.request_headers_to_remove.push(http::header::ACCEPT_ENCODING);
*body = Some(Bytes::from(serialized));
ctx.set_metadata(ARMED_KEY, "true");
ctx.set_metadata(CREATED_AT_KEY, ctx.time_source.now().as_secs().to_string());
let now = ctx.time_source.now().as_secs();
let created_at = ctx
.extensions
.get_mut::<ResponsesState>()
.map_or(now, |state| *state.response_created_at.get_or_insert(now));
ctx.set_metadata(CREATED_AT_KEY, created_at.to_string());

Ok(FilterAction::Continue)
}
Expand All @@ -299,6 +304,14 @@ fn request_disposition(ctx: &HttpFilterContext<'_>) -> Option<FilterAction> {
trace!(format, "releasing request classified as a different API format");
Some(FilterAction::Release)
},
None if ctx
.extensions
.get::<ResponsesState>()
.is_some_and(|state| state.response_id.is_some()) =>
{
trace!("using canonical Responses state across an iterative router metadata boundary");
None
},
None => {
warn!(
prerequisite = "openai_responses_format",
Expand Down Expand Up @@ -346,8 +359,16 @@ fn ensure_previous_response_rehydrated(state: &ResponsesState, streaming: bool)

/// Return the client stream preference captured by the classifier.
fn request_is_streaming(ctx: &HttpFilterContext<'_>) -> bool {
ctx.get_metadata("openai_responses_format.stream")
.is_some_and(|value| value == "true")
ctx.get_metadata("openai_responses_format.stream").map_or_else(
|| {
ctx.extensions
.get::<ResponsesState>()
.and_then(|state| state.request_body.get("stream"))
.and_then(serde_json::Value::as_bool)
.unwrap_or(false)
},
|value| value == "true",
)
}

/// Detect an SSE media type while response headers are still available.
Expand Down Expand Up @@ -444,20 +465,25 @@ fn prepare_transformed_response_headers(ctx: &mut HttpFilterContext<'_>) {

/// Convert a finite successful Chat response into a Responses resource.
fn translate_success_response(ctx: &HttpFilterContext<'_>, body: &[u8]) -> Result<Bytes, FilterError> {
let state = ctx
.extensions
.get::<ResponsesState>()
.ok_or_else(|| -> FilterError { "responses_to_chat_completions: missing Responses state".into() })?;
let response_id = ctx
.get_metadata("responses.response_id")
.or(state.response_id.as_deref())
.ok_or_else(|| -> FilterError { "responses_to_chat_completions: missing response id".into() })?;
let created_at = ctx
.get_metadata(CREATED_AT_KEY)
.and_then(|value| value.parse::<u64>().ok())
.or(state.response_created_at)
.ok_or_else(|| -> FilterError { "responses_to_chat_completions: missing creation timestamp".into() })?;
let state = ctx
.extensions
.get::<ResponsesState>()
.ok_or_else(|| -> FilterError { "responses_to_chat_completions: missing Responses state".into() })?;
let response_context =
let mut response_context =
ResponseContext::from_responses_request(&state.request_body, response_id.to_owned(), created_at)
.with_completed_at(ctx.time_source.now().as_secs());
if let Some(tool_choice) = state.original_tool_choice.as_ref() {
response_context.tool_choice = Some(tool_choice);
}
let provider_response: serde_json::Value = serde_json::from_slice(body)
.map_err(|error| -> FilterError { format!("responses_to_chat_completions: {error}").into() })?;
let translated = chat_response_to_response_resource(&provider_response, &response_context)
Expand Down
100 changes: 100 additions & 0 deletions apis/src/openai/responses/responses_to_chat_completions/tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -157,6 +157,48 @@ async fn classified_responses_create_without_state_fails_closed() {
assert_server_error(action);
}

#[tokio::test]
async fn canonical_state_translates_across_iterative_metadata_boundary() {
let filter = ResponsesToChatCompletionsFilter::from_config(&serde_yaml::Value::Null).unwrap();
let request = crate::test_utils::make_request(http::Method::POST, "/v1/responses");
let mut context = crate::test_utils::make_filter_context(&request);
let mut state = ResponsesState::from_request_body(json!({
"model": "gpt-4.1-mini",
"input": "hello",
"stream": false
}));
state.response_id = Some("resp_iterative".to_owned());
context.extensions.insert(state);
let mut body = Some(Bytes::from_static(
br#"{"model":"gpt-4.1-mini","input":"hello","stream":false}"#,
));

let action = filter.on_request_body(&mut context, &mut body, true).await.unwrap();

assert!(matches!(action, FilterAction::Continue));
let translated: serde_json::Value = serde_json::from_slice(body.as_deref().unwrap()).unwrap();
assert_eq!(translated["messages"][0]["content"], "hello");
assert_eq!(translated["stream"], false);
assert_eq!(context.get_metadata(ARMED_KEY), Some("true"));
}

#[tokio::test]
async fn unvalidated_state_does_not_bypass_missing_classifier_metadata() {
let filter = ResponsesToChatCompletionsFilter::from_config(&serde_yaml::Value::Null).unwrap();
let request = crate::test_utils::make_request(http::Method::POST, "/v1/responses");
let mut context = crate::test_utils::make_filter_context(&request);
context.extensions.insert(ResponsesState::from_request_body(json!({
"model": "gpt-4.1-mini",
"input": "hello"
})));
let mut body = Some(Bytes::from_static(br#"{"model":"gpt-4.1-mini","input":"hello"}"#));

let action = filter.on_request_body(&mut context, &mut body, true).await.unwrap();

assert_server_error(action);
assert!(context.get_metadata(ARMED_KEY).is_none());
}

#[tokio::test]
async fn unresolved_previous_response_id_fails_closed() {
let filter = ResponsesToChatCompletionsFilter::from_config(&serde_yaml::Value::Null).unwrap();
Expand Down Expand Up @@ -324,6 +366,13 @@ async fn canonical_state_is_translated_and_arms_response() {
);
assert_eq!(context.get_metadata(ARMED_KEY), Some("true"));
assert_eq!(context.get_metadata(CREATED_AT_KEY), Some("1700000000"));
assert_eq!(
context
.extensions
.get::<ResponsesState>()
.and_then(|state| state.response_created_at),
Some(1_700_000_000)
);
}

#[tokio::test]
Expand Down Expand Up @@ -773,6 +822,57 @@ async fn non_streaming_chat_response_becomes_response_resource() {
assert_eq!(translated["usage"]["output_tokens"], 2);
}

#[tokio::test]
async fn chat_file_search_function_call_becomes_responses_function_call() {
let filter = ResponsesToChatCompletionsFilter::from_config(&serde_yaml::Value::Null).unwrap();
let request = crate::test_utils::make_request(http::Method::POST, "/v1/responses");
let fixed_time = FixedTimeSource::new(Duration::from_secs(1_700_000_000));
let mut context = crate::test_utils::make_filter_context(&request);
context.time_source = &fixed_time;
let request_value = json!({
"model": "chat-only-model",
"input": "find revenue",
"stream": false,
"store": false,
"tools": [{"type": "file_search", "vector_store_ids": ["vs_q4"]}],
"tool_choice": {"type": "file_search"}
});
let mut state = ResponsesState::from_request_body(request_value);
state.response_id = Some("resp_file_search".to_owned());
context.extensions.insert(state);
let mut request_body = Some(Bytes::from_static(
br#"{"model":"chat-only-model","input":"find revenue"}"#,
));
let request_action = filter
.on_request_body(&mut context, &mut request_body, true)
.await
.unwrap();
assert!(matches!(request_action, FilterAction::Continue));

let response = Box::leak(Box::new(crate::test_utils::make_response()));
response.headers.insert(
http::header::CONTENT_TYPE,
http::HeaderValue::from_static("application/json"),
);
context.response_header = Some(response);
let response_action = filter.on_response(&mut context).await.unwrap();
assert!(matches!(response_action, FilterAction::Continue));
context.response_header = None;
let mut response_body = Some(Bytes::from_static(
br#"{"id":"chatcmpl_search","object":"chat.completion","model":"chat-only-model","choices":[{"index":0,"message":{"role":"assistant","content":null,"tool_calls":[{"id":"call_search","type":"function","function":{"name":"file_search","arguments":"{\"query\":\"Q4 revenue\"}"}}]},"finish_reason":"tool_calls"}],"usage":{"prompt_tokens":12,"completion_tokens":5,"total_tokens":17}}"#,
));

let body_action = filter.on_response_body(&mut context, &mut response_body, true).unwrap();

assert!(matches!(body_action, FilterAction::Continue));
let translated: serde_json::Value = serde_json::from_slice(response_body.as_deref().unwrap()).unwrap();
assert_eq!(translated["id"], "resp_file_search");
assert_eq!(translated["tools"][0]["type"], "file_search");
assert_eq!(translated["output"][0]["type"], "function_call");
assert_eq!(translated["output"][0]["name"], "file_search");
assert_eq!(translated["output"][0]["arguments"], "{\"query\":\"Q4 revenue\"}");
}

#[tokio::test]
async fn malformed_success_aborts_after_headers_are_sent() {
let yaml = serde_yaml::from_str("{}").unwrap();
Expand Down
16 changes: 16 additions & 0 deletions apis/src/openai/responses/state.rs
Original file line number Diff line number Diff line change
Expand Up @@ -124,6 +124,19 @@ pub(crate) struct ResponsesState {
/// Parsed request body as received from the client.
pub request_body: serde_json::Value,

/// Original client tool choice retained when continuation widens the
/// provider-visible choice to `auto`.
pub original_tool_choice: Option<serde_json::Value>,

/// Stable creation timestamp for the public response across iterations.
pub response_created_at: Option<u64>,

/// Stable public response ID assigned by request validation.
///
/// Stored with canonical state because iterative router steps preserve
/// extensions while resetting per-step metadata.
pub response_id: Option<String>,

/// Whether provider-visible request fields require outbound serialization.
pub request_body_rebuild: RequestBodyRebuild,

Expand Down Expand Up @@ -200,6 +213,9 @@ impl Default for ResponsesState {
previous_tools: Vec::new(),
previous_usage: None,
request_body: serde_json::Value::Null,
original_tool_choice: None,
response_created_at: None,
response_id: None,
request_body_rebuild: RequestBodyRebuild::PreserveOriginal,
response_object: serde_json::Value::Null,
tool_calls: Vec::new(),
Expand Down
14 changes: 13 additions & 1 deletion apis/src/openai/responses/validate/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -121,7 +121,7 @@ impl HttpFilter for OpenaiResponsesValidateFilter {
let conversation_id = resolve_conversation_id(ctx, &parsed);

enrich_context(ctx, &response_id, &conversation_id);
ctx.extensions.insert(ResponsesState::from_request_body(parsed));
insert_responses_state(ctx, parsed, &response_id);

debug!(
response_id = %response_id,
Expand All @@ -137,6 +137,13 @@ impl HttpFilter for OpenaiResponsesValidateFilter {
// Helpers
// -----------------------------------------------------------------------------

/// Initialize canonical request state, including metadata that survives IRR steps.
fn insert_responses_state(ctx: &mut HttpFilterContext<'_>, parsed: serde_json::Value, response_id: &str) {
let mut state = ResponsesState::from_request_body(parsed);
state.response_id = Some(response_id.to_owned());
ctx.extensions.insert(state);
}

/// Parse the request body as JSON.
fn parse_request_body(ctx: &HttpFilterContext<'_>, body: &Option<Bytes>) -> Result<serde_json::Value, FilterAction> {
let streaming = ctx
Expand Down Expand Up @@ -317,6 +324,11 @@ mod tests {
assert_eq!(state.tools.len(), 1, "tools should be populated");
assert_eq!(state.iteration, 0, "iteration should start at 0");
assert!(state.tool_calls.is_empty(), "tool_calls should start empty");
assert_eq!(
state.response_id.as_deref(),
ctx.filter_metadata.get("responses.response_id").map(String::as_str),
"canonical state should retain the public ID across iterative metadata boundaries"
);
}

#[tokio::test]
Expand Down
Loading
Loading