diff --git a/Cargo.toml b/Cargo.toml index 9273ca9034..b1b2ef2c10 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -303,4 +303,3 @@ bare_urls = "deny" invalid_html_tags = "deny" missing_crate_level_docs = "deny" private_doc_tests = "allow" - diff --git a/apis/src/openai/responses/README.md b/apis/src/openai/responses/README.md index aebf137417..25a0164612 100644 --- a/apis/src/openai/responses/README.md +++ b/apis/src/openai/responses/README.md @@ -29,20 +29,20 @@ Body-phase columns show `Access / Mode` when the hook is implemented. | Filter | `on_request` | `on_request_body` | `on_response` | `on_response_body` | |--------|:------------:|:-----------------:|:--------------:|:------------------:| -| `openai_agentic_loop` | — | ReadOnly / StreamBuffer | — | ReadWrite / StreamBuffer | +| `openai_agentic_loop` | — | ReadOnly / StreamBuffer | — | ReadWrite / Stream | | `openai_doc_extract` | — | ReadWrite / StreamBuffer | — | — | | `openai_file_resolve` | — | ReadWrite / StreamBuffer | — | — | | `openai_file_search_callout` | ✓ | ReadOnly / StreamBuffer | — | ReadWrite / StreamBuffer | -| `openai_mcp_dispatch` | — | ReadOnly / StreamBuffer | — | ReadOnly / StreamBuffer | +| `openai_mcp_dispatch` | — | ReadOnly / StreamBuffer | — | ReadOnly / Stream | | `openai_mcp_tool_resolve` | — | ReadWrite / StreamBuffer | — | — | -| `openai_response_store` | ✓ | ReadOnly / Stream | ✓ | ReadOnly / StreamBuffer | +| `openai_response_store` | ✓ | ReadOnly / Stream | ✓ | ReadOnly / Stream | | `openai_responses_compact` | — | ReadOnly / StreamBuffer | — | — | | `openai_responses_format` | — | ReadOnly / StreamBuffer | — | — | | `openai_responses_model_rewrite` | ✓ | ReadWrite / StreamBuffer | — | — | | `openai_responses_proxy` | — | ReadWrite / StreamBuffer | — | — | | `openai_responses_rehydrate` | — | ReadOnly / StreamBuffer | — | — | | `openai_responses_validate` | — | ReadOnly / StreamBuffer | — | — | -| `openai_stream_events` | ✓ | — | ✓ | ReadOnly / Stream | +| `openai_stream_events` | ✓ | — | ✓ | ReadWrite / Stream | | `openai_tool_parse` | ✓ | ReadOnly / StreamBuffer | — | — | -| `openai_web_search` | — | ReadOnly / StreamBuffer | — | ReadOnly / StreamBuffer | +| `openai_web_search` | — | ReadOnly / StreamBuffer | — | ReadOnly / Stream | | `responses_to_chat_completions` | — | ReadWrite / StreamBuffer | ✓ | ReadWrite / Stream | diff --git a/apis/src/openai/responses/agentic_loop/mod.rs b/apis/src/openai/responses/agentic_loop/mod.rs index e3346b658c..d601fafcf9 100644 --- a/apis/src/openai/responses/agentic_loop/mod.rs +++ b/apis/src/openai/responses/agentic_loop/mod.rs @@ -19,16 +19,16 @@ //! - `openai_agentic_loop.action = "loop"` — tool calls present, loop back //! - `openai_agentic_loop.action = "done"` — exit to client //! -//! # Non-streaming tool call extraction +//! # Tool call extraction //! -//! For non-streaming responses (the only mode supported by IRR), +//! For non-streaming responses, //! this filter parses the response body JSON and extracts //! `function_call` items from the `output` array into //! `state.tool_calls` and `web_search_call` items into //! `state.web_search_calls`. It also appends these items to //! `state.messages` so the model sees its own calls on re-entry. //! -//! For streaming responses (future), `stream_events` populates +//! For streaming responses, `stream_events` populates //! `state.tool_calls` via SSE event parsing. When the body is //! `None` at end-of-stream (consumed by streaming filters), this //! filter skips body parsing and checks `state.tool_calls` as-is. @@ -81,12 +81,6 @@ //! done: true //! ``` //! -//! # Streaming limitation -//! -//! Streaming requests (`stream: true`) are rejected with a 400 -//! error. `iterative_request_router` fully buffers all responses -//! within the loop and cannot forward incremental SSE events. -//! //! # State dependency //! //! Requires [`ResponsesState`] in request extensions. Without it @@ -216,9 +210,7 @@ impl HttpFilter for AgenticLoopFilter { } fn response_body_mode(&self) -> BodyMode { - BodyMode::StreamBuffer { - max_bytes: Some(self.config.max_body_bytes), - } + BodyMode::Stream } async fn on_request(&self, _ctx: &mut HttpFilterContext<'_>) -> Result { @@ -239,16 +231,6 @@ impl HttpFilter for AgenticLoopFilter { return Ok(FilterAction::Continue); }; - if state.request_body.get("stream") == Some(&Value::Bool(true)) { - ctx.extensions.insert(state); - return Ok(FilterAction::Reject(responses_error_rejection( - 400, - "invalid_request_error", - "streaming is not supported with openai_agentic_loop", - false, - ))); - } - prepare_iteration(ctx, &mut state); trace!(iteration = state.iteration, "openai_agentic_loop on_request_body"); ctx.extensions.insert(state); @@ -269,16 +251,13 @@ impl HttpFilter for AgenticLoopFilter { return Ok(FilterAction::Continue); }; - if let Some(bytes) = body.as_ref() - && let Err(msg) = extract_tool_calls_from_body(bytes, &mut state) - { + if let Some(bytes) = body.as_ref() { + if let Err(msg) = extract_tool_calls_from_body(bytes, &mut state) { + return Ok(reject_invalid_function_cardinality(ctx, state, msg)); + } + } else if !prepare_streamed_round(ctx, &mut state)? { ctx.extensions.insert(state); - return Ok(FilterAction::Reject(responses_error_rejection( - 400, - "invalid_request_error", - msg, - false, - ))); + return Ok(FilterAction::Continue); } let result = evaluate_loop_decision(ctx, &mut state, body, &self.config)?; @@ -287,6 +266,60 @@ impl HttpFilter for AgenticLoopFilter { } } +/// Preserve state while rejecting a buffered round with invalid cardinality. +fn reject_invalid_function_cardinality( + ctx: &mut HttpFilterContext<'_>, + state: ResponsesState, + message: &'static str, +) -> FilterAction { + ctx.extensions.insert(state); + FilterAction::Reject(responses_error_rejection(400, "invalid_request_error", message, false)) +} + +/// Only a successfully terminated stream may authorize external side effects. +fn streamed_round_is_dispatchable(ctx: &HttpFilterContext<'_>, state: &ResponsesState) -> bool { + state.request_body.get("stream").and_then(Value::as_bool) != Some(true) + || (ctx.get_metadata("responses.stream_completion") == Some("terminal") + && state.response_object.get("status").and_then(Value::as_str) == Some("completed") + && ctx.get_metadata("responses.stream_parse_error") != Some("true")) +} + +/// Collect an authoritative successful stream or terminate without dispatch. +fn prepare_streamed_round(ctx: &mut HttpFilterContext<'_>, state: &mut ResponsesState) -> Result { + if !streamed_round_is_dispatchable(ctx, state) { + collect_streaming_output_items(state); + state.tool_calls.clear(); + state.web_search_calls.clear(); + set_action(ctx, ACTION_DONE)?; + return Ok(false); + } + collect_streaming_output_items(state); + if state.tool_calls.len() > 1 { + end_stream_with_error( + ctx, + state, + "invalid_request_error", + "openai_agentic_loop supports exactly one function call per round", + )?; + } + Ok(true) +} + +/// End an already-committed stream with a local SSE error and no side effects. +fn end_stream_with_error( + ctx: &mut HttpFilterContext<'_>, + state: &mut ResponsesState, + code: &'static str, + message: &'static str, +) -> Result<(), FilterError> { + state.tool_calls.clear(); + state.web_search_calls.clear(); + ctx.set_metadata("responses.stream_error_code", code); + ctx.set_metadata("responses.stream_error_message", message); + ctx.set_metadata("responses.skip_persist", "true"); + set_action(ctx, ACTION_DONE) +} + // ----------------------------------------------------------------------------- // Request-Side Bookkeeping // ----------------------------------------------------------------------------- @@ -344,12 +377,7 @@ fn evaluate_loop_decision( set_action(ctx, ACTION_DONE)?; Ok(FilterAction::Continue) }, - Some(ExitReason::IterationLimit) => Ok(FilterAction::Reject(responses_error_rejection( - 508, - "server_error", - "agentic loop iteration limit exceeded", - false, - ))), + Some(ExitReason::IterationLimit) => end_at_iteration_limit(ctx, state), None => { state.iteration += 1; let (tc, wsc) = (state.tool_calls.len(), state.web_search_calls.len()); @@ -361,6 +389,23 @@ fn evaluate_loop_decision( } } +/// Terminate at the iteration cap using the transport that is still writable. +fn end_at_iteration_limit( + ctx: &mut HttpFilterContext<'_>, + state: &mut ResponsesState, +) -> Result { + if state.request_body.get("stream").and_then(Value::as_bool) == Some(true) { + end_stream_with_error(ctx, state, "server_error", "agentic loop iteration limit exceeded")?; + return Ok(FilterAction::Continue); + } + Ok(FilterAction::Reject(responses_error_rejection( + 508, + "server_error", + "agentic loop iteration limit exceeded", + false, + ))) +} + // ----------------------------------------------------------------------------- // Body Parsing // ----------------------------------------------------------------------------- @@ -418,6 +463,27 @@ fn collect_output_items(response: &Value, state: &mut ResponsesState) { } } +/// Retain the current streamed round after `openai_stream_events` has built +/// its authoritative response object and tool-call list incrementally. +fn collect_streaming_output_items(state: &mut ResponsesState) { + let output = state.output_items().to_vec(); + for item in output { + state.accumulated_output.push(item.clone()); + match item.get("type").and_then(Value::as_str) { + Some("function_call" | "reasoning") => { + state.messages.push(item.clone()); + state.persisted_messages.push(item); + }, + Some("web_search_call") => { + state.web_search_calls.push(item.clone()); + state.messages.push(item.clone()); + state.persisted_messages.push(item); + }, + _ => {}, + } + } +} + /// Check whether a parsed response is a valid Responses API output. /// /// Returns `false` for error bodies (`"object": "error"`) and diff --git a/apis/src/openai/responses/agentic_loop/tests.rs b/apis/src/openai/responses/agentic_loop/tests.rs index 3ee0084d8b..84c6ac14fa 100644 --- a/apis/src/openai/responses/agentic_loop/tests.rs +++ b/apis/src/openai/responses/agentic_loop/tests.rs @@ -260,11 +260,11 @@ async fn preserves_unmodified_parallel_tool_calls_false() { } // ----------------------------------------------------------------------------- -// on_request_body: Reject Streaming +// on_request_body: Streaming // ----------------------------------------------------------------------------- #[tokio::test] -async fn rejects_streaming_request() { +async fn accepts_streaming_request() { let filter = make_filter(); let req = make_request(Method::POST, "/v1/responses"); let mut ctx = make_filter_context(&req); @@ -275,13 +275,13 @@ async fn rejects_streaming_request() { let action = filter.on_request_body(&mut ctx, &mut None, true).await.unwrap(); assert!( - matches!(&action, FilterAction::Reject(r) if r.status == 400), - "stream:true should produce a 400 rejection" + matches!(action, FilterAction::Continue), + "stream:true should continue into IRR streaming" ); } #[tokio::test] -async fn streaming_rejection_preserves_state() { +async fn streaming_request_preserves_state() { let filter = make_filter(); let req = make_request(Method::POST, "/v1/responses"); let mut ctx = make_filter_context(&req); @@ -294,7 +294,220 @@ async fn streaming_rejection_preserves_state() { assert!( ctx.extensions.get::().is_some(), - "ResponsesState must remain in extensions after streaming rejection" + "ResponsesState must remain in extensions for streaming response accumulation" + ); +} + +#[test] +fn incomplete_stream_does_not_dispatch_accumulated_tool_call() { + let filter = make_filter(); + let req = make_request(Method::POST, "/v1/responses"); + let mut ctx = make_filter_context(&req); + + let mut state = ResponsesState::from_request_body(json!({ + "model": "gpt-4o", + "input": "test", + "stream": true + })); + state.tool_calls.push(json!({ + "type": "function_call", + "call_id": "call_partial", + "name": "must_not_run", + "arguments": "{}", + "status": "completed" + })); + state.response_object = json!({ + "id": "resp_incomplete", + "object": "response", + "status": "incomplete", + "output": [{ + "type": "message", + "id": "msg_partial", + "status": "incomplete", + "role": "assistant", + "content": [{"type": "output_text", "text": "partial"}] + }] + }); + let partial_output = state + .response_object + .get("output") + .and_then(Value::as_array) + .expect("incomplete response should have output") + .clone(); + state.output_items_mut().clone_from(&partial_output); + ctx.set_metadata("responses.stream_completion", "terminal"); + ctx.extensions.insert(state); + + let action = filter.on_response_body(&mut ctx, &mut None, true).unwrap(); + assert!(matches!(action, FilterAction::Continue)); + assert_action(&ctx, "done"); + assert!( + ctx.extensions.get::().unwrap().tool_calls.is_empty(), + "a tool from a truncated stream must not remain dispatchable" + ); + assert_eq!( + ctx.extensions.get::().unwrap().accumulated_output[0]["id"], + "msg_partial", + "partial output from an incomplete response must remain client-visible" + ); +} + +#[test] +fn failed_stream_does_not_dispatch_accumulated_tool_call() { + let filter = make_filter(); + let req = make_request(Method::POST, "/v1/responses"); + let mut ctx = make_filter_context(&req); + + let mut state = ResponsesState::from_request_body(json!({ + "model": "gpt-4o", + "input": "test", + "stream": true + })); + state.tool_calls.push(json!({ + "type": "function_call", + "call_id": "call_failed", + "name": "must_not_run", + "arguments": "{}", + "status": "completed" + })); + state.response_object = json!({ + "id": "resp_failed", + "object": "response", + "status": "failed", + "output": [{ + "type": "message", + "id": "msg_before_failure", + "status": "incomplete", + "role": "assistant", + "content": [{"type": "output_text", "text": "before failure"}] + }] + }); + let partial_output = state + .response_object + .get("output") + .and_then(Value::as_array) + .expect("failed response should retain partial output") + .clone(); + state.output_items_mut().clone_from(&partial_output); + ctx.set_metadata("responses.stream_completion", "terminal"); + ctx.extensions.insert(state); + + let action = filter.on_response_body(&mut ctx, &mut None, true).unwrap(); + assert!(matches!(action, FilterAction::Continue)); + assert_action(&ctx, "done"); + assert!( + ctx.extensions.get::().unwrap().tool_calls.is_empty(), + "a provider-failed response must not authorize tool execution" + ); + assert_eq!( + ctx.extensions.get::().unwrap().accumulated_output[0]["id"], + "msg_before_failure", + "partial output preceding a provider failure must remain client-visible" + ); +} + +#[test] +fn streamed_web_search_call_is_available_to_dispatch_filter() { + let filter = make_filter(); + let req = make_request(Method::POST, "/v1/responses"); + let mut ctx = make_filter_context(&req); + + let mut state = ResponsesState::from_request_body(json!({ + "model": "gpt-4o", + "input": "test", + "stream": true + })); + state.response_object = json!({ + "id": "resp_search", + "object": "response", + "status": "completed", + "output": [{ + "type": "web_search_call", + "id": "ws_1", + "status": "completed", + "action": {"type": "search", "query": "Praxis"} + }] + }); + ctx.set_metadata("responses.stream_completion", "terminal"); + ctx.extensions.insert(state); + + let action = filter.on_response_body(&mut ctx, &mut None, true).unwrap(); + assert!(matches!(action, FilterAction::Continue)); + assert_action(&ctx, "loop"); + assert_eq!( + ctx.extensions.get::().unwrap().web_search_calls.len(), + 1, + "streamed web-search calls must be visible to openai_web_search" + ); +} + +#[test] +fn multiple_streamed_function_calls_end_with_sse_error() { + let filter = make_filter(); + let req = make_request(Method::POST, "/v1/responses"); + let mut ctx = make_filter_context(&req); + + let mut state = ResponsesState::from_request_body(json!({ + "model": "gpt-4o", + "input": "test", + "stream": true + })); + state.tool_calls = vec![ + json!({"type": "function_call", "call_id": "call_1", "name": "first", "status": "completed"}), + json!({"type": "function_call", "call_id": "call_2", "name": "second", "status": "completed"}), + ]; + state.response_object = json!({"id": "resp_multiple", "object": "response", "status": "completed", "output": []}); + ctx.set_metadata("responses.stream_completion", "terminal"); + ctx.extensions.insert(state); + + let action = filter.on_response_body(&mut ctx, &mut None, true).unwrap(); + assert!(matches!(action, FilterAction::Continue)); + assert_action(&ctx, "done"); + assert_eq!( + ctx.get_metadata("responses.stream_error_code"), + Some("invalid_request_error"), + "post-commit validation must select a terminal SSE error" + ); + assert!( + ctx.extensions.get::().unwrap().tool_calls.is_empty(), + "invalid streamed calls must not remain dispatchable" + ); +} + +#[test] +fn streaming_iteration_limit_ends_with_sse_error() { + let yaml: serde_yaml::Value = serde_yaml::from_str("max_infer_iters: 1").unwrap(); + let filter = super::AgenticLoopFilter::from_config(&yaml).unwrap(); + let req = make_request(Method::POST, "/v1/responses"); + let mut ctx = make_filter_context(&req); + + let mut state = ResponsesState::from_request_body(json!({ + "model": "gpt-4o", + "input": "test", + "stream": true + })); + state.iteration = 1; + state.tool_calls = vec![json!({ + "type": "function_call", + "call_id": "call_limit", + "name": "must_not_run", + "status": "completed" + })]; + state.response_object = json!({"id": "resp_limit", "object": "response", "status": "completed", "output": []}); + ctx.set_metadata("responses.stream_completion", "terminal"); + ctx.extensions.insert(state); + + let action = filter.on_response_body(&mut ctx, &mut None, true).unwrap(); + assert!(matches!(action, FilterAction::Continue)); + assert_action(&ctx, "done"); + assert_eq!( + ctx.get_metadata("responses.stream_error_code"), + Some("server_error"), + "an exhausted committed stream must end with an SSE error" + ); + assert!( + ctx.extensions.get::().unwrap().tool_calls.is_empty(), + "iteration-limit errors must not leave calls dispatchable" ); } diff --git a/apis/src/openai/responses/error.rs b/apis/src/openai/responses/error.rs index 83bdbddf9f..97e6281432 100644 --- a/apis/src/openai/responses/error.rs +++ b/apis/src/openai/responses/error.rs @@ -31,7 +31,13 @@ pub(crate) fn responses_error_body(code: &str, message: &str) -> Bytes { /// /// Produces `event: error\ndata: \n\n`. pub(crate) fn responses_error_sse_body(code: &str, message: &str) -> Bytes { - let json = serde_json::json!({ + let json = responses_error_sse_payload(code, message); + Bytes::from(format!("event: error\ndata: {json}\n\n")) +} + +/// Build the JSON payload for a Responses API SSE error event. +pub(crate) fn responses_error_sse_payload(code: &str, message: &str) -> serde_json::Value { + serde_json::json!({ "type": "error", "sequence_number": 0, "error": { @@ -40,8 +46,7 @@ pub(crate) fn responses_error_sse_body(code: &str, message: &str) -> Bytes { "message": message, "param": null, }, - }); - Bytes::from(format!("event: error\ndata: {json}\n\n")) + }) } /// Build a [`Rejection`] with the appropriate OpenAI error format. diff --git a/apis/src/openai/responses/mcp_dispatch/mod.rs b/apis/src/openai/responses/mcp_dispatch/mod.rs index 6597cae883..3fe7a4e2a6 100644 --- a/apis/src/openai/responses/mcp_dispatch/mod.rs +++ b/apis/src/openai/responses/mcp_dispatch/mod.rs @@ -167,9 +167,7 @@ impl HttpFilter for McpDispatchFilter { } fn response_body_mode(&self) -> BodyMode { - BodyMode::StreamBuffer { - max_bytes: Some(self.max_body_bytes), - } + BodyMode::Stream } async fn on_request(&self, _ctx: &mut HttpFilterContext<'_>) -> Result { @@ -224,7 +222,7 @@ impl HttpFilter for McpDispatchFilter { end_of_stream: bool, ) -> Result { if !end_of_stream { - return Ok(FilterAction::Release); + return Ok(FilterAction::Continue); } let Some(state) = ctx.extensions.get::() else { diff --git a/apis/src/openai/responses/mcp_dispatch/tests.rs b/apis/src/openai/responses/mcp_dispatch/tests.rs index b26efee8f4..91460ed238 100644 --- a/apis/src/openai/responses/mcp_dispatch/tests.rs +++ b/apis/src/openai/responses/mcp_dispatch/tests.rs @@ -195,11 +195,8 @@ fn filter_response_body_mode() { let config = serde_yaml::from_str::("max_body_bytes: 1024").unwrap(); let filter = McpDispatchFilter::from_config(&config).unwrap(); assert!( - matches!( - filter.response_body_mode(), - praxis_filter::BodyMode::StreamBuffer { max_bytes: Some(1024) } - ), - "should return StreamBuffer with configured max_bytes" + matches!(filter.response_body_mode(), praxis_filter::BodyMode::Stream), + "agentic responses must remain stream-compatible" ); } @@ -928,13 +925,16 @@ fn assert_dispatch_action(ctx: &praxis_filter::HttpFilterContext<'_>, expected: } #[test] -fn on_response_body_not_end_of_stream_returns_release() { +fn on_response_body_not_end_of_stream_continues_to_stream_parser() { let filter = make_dispatch_filter(); let req = make_request(http::Method::POST, "/v1/responses"); let mut ctx = make_filter_context(&req); let mut body = Some(Bytes::from("data")); let result = filter.on_response_body(&mut ctx, &mut body, false).unwrap(); - assert!(matches!(result, FilterAction::Release)); + assert!( + matches!(result, FilterAction::Continue), + "stream chunks must reach the downstream openai_stream_events filter" + ); } #[test] diff --git a/apis/src/openai/responses/openai_responses_proxy/config.rs b/apis/src/openai/responses/openai_responses_proxy/config.rs index 3be246b7d9..e4dc94f836 100644 --- a/apis/src/openai/responses/openai_responses_proxy/config.rs +++ b/apis/src/openai/responses/openai_responses_proxy/config.rs @@ -17,6 +17,7 @@ use serde::Deserialize; /// /// ```yaml /// filter: openai_responses_proxy +/// terminal_streaming: false /// ``` #[derive(Debug, Deserialize)] #[serde(deny_unknown_fields)] @@ -24,12 +25,19 @@ pub(super) struct ResponsesProxyConfig { /// Maximum body size in bytes for `StreamBuffer` mode. #[serde(default = "default_max_body_bytes")] pub max_body_bytes: usize, + + /// Select Praxis streaming transport for effective `stream: true` + /// requests. IRR may resume the same downstream stream after a step + /// transition when its response filters use streaming body mode. + #[serde(default)] + pub terminal_streaming: bool, } impl Default for ResponsesProxyConfig { fn default() -> Self { Self { max_body_bytes: MAX_JSON_BODY_BYTES, + terminal_streaming: false, } } } diff --git a/apis/src/openai/responses/openai_responses_proxy/mod.rs b/apis/src/openai/responses/openai_responses_proxy/mod.rs index e6b33171c6..537201ee6c 100644 --- a/apis/src/openai/responses/openai_responses_proxy/mod.rs +++ b/apis/src/openai/responses/openai_responses_proxy/mod.rs @@ -37,9 +37,10 @@ use async_trait::async_trait; use base64::Engine as _; use bytes::Bytes; use praxis_filter::{ - BodyAccess, BodyMode, FilterAction, FilterError, HttpFilter, HttpFilterContext, parse_filter_config, + BodyAccess, BodyMode, FilterAction, FilterError, HttpFilter, HttpFilterContext, SubRequestResponseMode, + parse_filter_config, }; -use serde::ser::SerializeMap as _; +use serde::{Deserialize, ser::SerializeMap as _}; use tracing::{debug, trace}; use self::config::{ResponsesProxyConfig, build_config}; @@ -61,6 +62,13 @@ use crate::json_body::{SerializedJson, serialize_json_body}; /// When no `ResponsesState` exists, preserves the request body apart /// from removing the Praxis-owned `conversation` field. /// +/// Set `terminal_streaming: true` inside an iterative request router step to +/// select Praxis's streaming transport when the effective outbound body +/// contains `"stream": true`. Classifier metadata remains descriptive client +/// intent; this final serializer owns the transport decision. IRR can resume +/// one downstream stream across response-dependent transitions, but every +/// response-body filter in a streaming-capable step must use `BodyMode::Stream`. +/// /// # YAML /// /// ```yaml @@ -72,6 +80,7 @@ use crate::json_body::{SerializedJson, serialize_json_body}; /// ```yaml /// filter: openai_responses_proxy /// max_body_bytes: 67108864 +/// terminal_streaming: false /// ``` /// /// # Example @@ -154,6 +163,10 @@ impl HttpFilter for ResponsesProxyFilter { } } + fn may_select_streaming_subrequest_response(&self) -> bool { + self.config.terminal_streaming + } + async fn on_request(&self, _ctx: &mut HttpFilterContext<'_>) -> Result { Ok(FilterAction::Continue) } @@ -171,12 +184,14 @@ impl HttpFilter for ResponsesProxyFilter { let Some(state) = ctx.extensions.get::() else { strip_conversation_field(body, self.name()); + select_terminal_response_mode(&self.config, ctx, body); debug!("no ResponsesState in extensions, passthrough"); return Ok(FilterAction::Continue); }; if !request_needs_rebuild(state) { strip_conversation_field(body, self.name()); + select_terminal_response_mode(&self.config, ctx, body); debug!("ResponsesState does not require an outbound rewrite, passthrough"); return Ok(FilterAction::Continue); } @@ -191,15 +206,49 @@ impl HttpFilter for ResponsesProxyFilter { }; SerializedJson::from_bytes(serialized).commit(body, self.name(), "body"); + select_terminal_response_mode(&self.config, ctx, body); Ok(FilterAction::Continue) } } +/// Narrow deserialization target for the provider-visible stream bit. +/// +/// Only `stream` participates in transport selection; all other request fields +/// are intentionally ignored. +#[derive(Deserialize)] +struct EffectiveResponseMode { + /// Whether the effective outbound Responses request asks for SSE. + #[serde(default)] + stream: bool, +} + // ----------------------------------------------------------------------------- // Helpers // ----------------------------------------------------------------------------- +/// Align the typed Praxis response mode with the effective serialized request. +/// +/// Classifier metadata describes client intent, but request transformations can +/// change the provider-visible body. The final serializer therefore owns this +/// transport decision and reads the bytes it actually leaves for the upstream. +fn select_terminal_response_mode(config: &ResponsesProxyConfig, ctx: &mut HttpFilterContext<'_>, body: &Option) { + if !config.terminal_streaming { + return; + } + + let mode = if body + .as_deref() + .and_then(|bytes| serde_json::from_slice::(bytes).ok()) + .is_some_and(|selection| selection.stream) + { + SubRequestResponseMode::Streaming + } else { + SubRequestResponseMode::Buffered + }; + ctx.set_subrequest_response_mode(mode); +} + /// Defensively strip `conversation` from a passthrough body so it never /// leaks to the backend even when no [`ResponsesState`] was produced. fn strip_conversation_field(body: &mut Option, filter_name: &'static str) { diff --git a/apis/src/openai/responses/openai_responses_proxy/tests.rs b/apis/src/openai/responses/openai_responses_proxy/tests.rs index 254688de8e..508b2035da 100644 --- a/apis/src/openai/responses/openai_responses_proxy/tests.rs +++ b/apis/src/openai/responses/openai_responses_proxy/tests.rs @@ -6,7 +6,7 @@ use base64::Engine as _; use bytes::Bytes; use http::Method; -use praxis_filter::{BodyAccess, BodyMode, FilterAction, HttpFilter}; +use praxis_filter::{BodyAccess, BodyMode, FilterAction, HttpFilter, SubRequestResponseMode}; use serde_json::json; use super::super::state::ResponsesState; @@ -64,6 +64,128 @@ fn body_mode_is_stream_buffer() { ); } +#[test] +fn terminal_streaming_capability_is_opt_in() { + assert!( + !make_filter().may_select_streaming_subrequest_response(), + "default configuration must preserve existing buffered IRR pipelines" + ); + assert!( + make_terminal_streaming_filter().may_select_streaming_subrequest_response(), + "terminal_streaming must declare the Praxis streaming capability" + ); +} + +#[tokio::test] +async fn terminal_streaming_selects_streaming_from_effective_passthrough_body() { + let filter = make_terminal_streaming_filter(); + let req = make_request(Method::POST, "/v1/responses"); + let mut ctx = make_filter_context(&req); + ctx.set_metadata("openai_responses_format.stream", "false"); + let mut body = Some(Bytes::from_static( + br#"{"model":"gpt-4.1","input":"hello","stream":true}"#, + )); + + let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); + + assert!( + matches!(action, FilterAction::Continue), + "terminal streaming passthrough should continue" + ); + assert_eq!( + ctx.subrequest_response_mode(), + SubRequestResponseMode::Streaming, + "effective provider body, not descriptive classifier metadata, must select transport" + ); +} + +#[tokio::test] +async fn terminal_streaming_preserves_buffered_mode_when_stream_is_false_or_absent() { + let filter = make_terminal_streaming_filter(); + for original in [ + br#"{"model":"gpt-4.1","input":"hello","stream":false}"#.as_slice(), + br#"{"model":"gpt-4.1","input":"hello"}"#.as_slice(), + ] { + let req = make_request(Method::POST, "/v1/responses"); + let mut ctx = make_filter_context(&req); + ctx.set_subrequest_response_mode(SubRequestResponseMode::Streaming); + let mut body = Some(Bytes::copy_from_slice(original)); + + let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); + + assert!( + matches!(action, FilterAction::Continue), + "non-streaming terminal request should continue" + ); + assert_eq!( + ctx.subrequest_response_mode(), + SubRequestResponseMode::Buffered, + "non-streaming effective body must keep the buffered transport" + ); + } +} + +#[tokio::test] +async fn terminal_streaming_uses_rebuilt_state_body_not_client_intent_metadata() { + let filter = make_terminal_streaming_filter(); + let req = make_request(Method::POST, "/v1/responses"); + let mut ctx = make_filter_context(&req); + ctx.set_metadata("openai_responses_format.stream", "true"); + let mut state = ResponsesState::from_request_body(json!({ + "model": "gpt-4.1", + "input": "hello", + "stream": false + })); + state + .messages + .push(json!({"type":"function_call_output","call_id":"call_1","output":"done"})); + ctx.extensions.insert(state); + ctx.set_subrequest_response_mode(SubRequestResponseMode::Streaming); + let mut body = Some(Bytes::from_static( + br#"{"model":"gpt-4.1","input":"hello","stream":true}"#, + )); + + let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); + + assert!( + matches!(action, FilterAction::Continue), + "rebuilt buffered request should continue" + ); + let outbound: serde_json::Value = serde_json::from_slice(body.as_ref().unwrap()).unwrap(); + assert_eq!(outbound["stream"], false); + assert_eq!(ctx.subrequest_response_mode(), SubRequestResponseMode::Buffered); +} + +#[tokio::test] +async fn terminal_streaming_selects_streaming_for_rebuilt_effective_body() { + let filter = make_terminal_streaming_filter(); + let req = make_request(Method::POST, "/v1/responses"); + let mut ctx = make_filter_context(&req); + ctx.set_metadata("openai_responses_format.stream", "false"); + let mut state = ResponsesState::from_request_body(json!({ + "model": "gpt-4.1", + "input": "hello", + "stream": true + })); + state + .messages + .push(json!({"type":"function_call_output","call_id":"call_1","output":"done"})); + ctx.extensions.insert(state); + let mut body = Some(Bytes::from_static( + br#"{"model":"gpt-4.1","input":"hello","stream":false}"#, + )); + + let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); + + assert!( + matches!(action, FilterAction::Continue), + "rebuilt streaming request should continue" + ); + let outbound: serde_json::Value = serde_json::from_slice(body.as_ref().unwrap()).unwrap(); + assert_eq!(outbound["stream"], true); + assert_eq!(ctx.subrequest_response_mode(), SubRequestResponseMode::Streaming); +} + #[tokio::test] async fn on_request_returns_continue() { let filter = make_filter(); @@ -111,7 +233,10 @@ async fn initialized_state_preserves_scalar_input_on_first_pass() { let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); - assert!(matches!(action, FilterAction::Continue)); + assert!( + matches!(action, FilterAction::Continue), + "initialized scalar request should continue" + ); assert_eq!(body.as_deref(), Some(original.as_slice())); assert!( ctx.request_headers_to_set.is_empty(), @@ -135,9 +260,15 @@ async fn provider_previous_response_id_is_byte_exact_without_rehydrate() { let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); - assert!(matches!(action, FilterAction::Continue)); + assert!( + matches!(action, FilterAction::Continue), + "provider previous_response_id passthrough should continue" + ); assert_eq!(body.as_deref(), Some(original.as_slice())); - assert!(ctx.request_headers_to_set.is_empty()); + assert!( + ctx.request_headers_to_set.is_empty(), + "byte-exact previous_response_id passthrough must not synthesize headers" + ); } #[tokio::test] @@ -163,7 +294,10 @@ async fn rebuild_preserves_provider_previous_response_id_without_rehydrate() { let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); - assert!(matches!(action, FilterAction::Continue)); + assert!( + matches!(action, FilterAction::Continue), + "rebuilt previous_response_id request should continue" + ); let outbound: serde_json::Value = serde_json::from_slice(body.as_ref().unwrap()).unwrap(); assert_eq!(outbound["previous_response_id"], "resp_provider"); assert_eq!(outbound["input"].as_array().unwrap().len(), 2); @@ -184,7 +318,10 @@ async fn rebuild_serializes_from_state_request_body() { let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); - assert!(matches!(action, FilterAction::Continue)); + assert!( + matches!(action, FilterAction::Continue), + "state-backed rebuild should continue" + ); let outbound: serde_json::Value = serde_json::from_slice(body.as_ref().unwrap()).unwrap(); assert_eq!(outbound["model"], "client-model", "serializes from state.request_body"); assert_eq!(outbound["input"].as_array().unwrap().len(), 2); @@ -361,7 +498,10 @@ async fn strips_conversation_from_outbound_body() { r#"{"model":"gpt-4o","input":"hello","conversation":{"id":"conv_abc123"}}"#, )); let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); - assert!(matches!(action, FilterAction::Continue)); + assert!( + matches!(action, FilterAction::Continue), + "conversation stripping should continue" + ); let rebuilt: serde_json::Value = serde_json::from_slice(body.as_ref().unwrap()).unwrap(); assert!( @@ -394,7 +534,10 @@ async fn strips_both_previous_response_id_and_conversation() { r#"{"model":"gpt-4o","input":"hello","previous_response_id":"resp_abc123","conversation":"conv_xyz789"}"#, )); let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); - assert!(matches!(action, FilterAction::Continue)); + assert!( + matches!(action, FilterAction::Continue), + "identifier stripping should continue" + ); let rebuilt: serde_json::Value = serde_json::from_slice(body.as_ref().unwrap()).unwrap(); assert!( @@ -412,7 +555,10 @@ async fn passthrough_strips_conversation_from_body() { let mut body = Some(Bytes::from(r#"{"model":"gpt-4.1","input":"hello","conversation":42}"#)); let action = filter.on_request_body(&mut ctx, &mut body, true).await.unwrap(); - assert!(matches!(action, FilterAction::Continue)); + assert!( + matches!(action, FilterAction::Continue), + "passthrough conversation stripping should continue" + ); let parsed: serde_json::Value = serde_json::from_slice(body.as_ref().unwrap()).unwrap(); assert!( @@ -543,10 +689,16 @@ fn messages_for_backend_translates_compaction_item() { let encoded = base64::engine::general_purpose::STANDARD.encode("summary text"); let msgs = vec![json!({"type": "compaction", "id": "c_1", "encrypted_content": encoded})]; let result = super::messages_for_backend(&msgs); - assert!(matches!(result, std::borrow::Cow::Owned(_))); + assert!( + matches!(result, std::borrow::Cow::Owned(_)), + "summary insertion must return an owned input array" + ); assert_eq!(result.len(), 1); assert_eq!(result[0]["role"], "assistant"); - assert!(result[0]["content"].as_str().unwrap().contains("summary text")); + assert!( + result[0]["content"].as_str().unwrap().contains("summary text"), + "inserted summary message must contain the supplied summary" + ); } #[test] @@ -568,7 +720,10 @@ fn compaction_to_assistant_message_decodes_encrypted_content() { let item = json!({"type": "compaction", "id": "c_1", "encrypted_content": encoded}); let msg = super::compaction_to_assistant_message(&item); assert_eq!(msg["role"], "assistant"); - assert!(msg["content"].as_str().unwrap().contains("decoded summary")); + assert!( + msg["content"].as_str().unwrap().contains("decoded summary"), + "inserted summary message must contain the decoded summary" + ); } #[test] @@ -598,3 +753,8 @@ fn compaction_to_assistant_message_handles_invalid_base64() { fn make_filter() -> Box { super::ResponsesProxyFilter::from_config(&serde_yaml::Value::Null).unwrap() } + +fn make_terminal_streaming_filter() -> Box { + let yaml = serde_yaml::from_str("terminal_streaming: true").unwrap(); + super::ResponsesProxyFilter::from_config(&yaml).unwrap() +} diff --git a/apis/src/openai/responses/state.rs b/apis/src/openai/responses/state.rs index 81e22d68a6..847bca73bf 100644 --- a/apis/src/openai/responses/state.rs +++ b/apis/src/openai/responses/state.rs @@ -55,6 +55,13 @@ pub(crate) struct ResponsesState { /// optional sections to populate. pub include: Vec, + /// Stable response ID used while several streamed inference rounds are + /// exposed as one logical Responses stream. + pub logical_stream_response_id: Option, + + /// Next downstream sequence number for a logical Responses stream. + pub logical_stream_sequence: u64, + /// Whether stored history was successfully resolved into this state. /// /// The proxy uses this to distinguish locally consumed history identifiers @@ -186,6 +193,8 @@ impl Default for ResponsesState { conversation: None, file_search_output_items: Vec::new(), include: Vec::new(), + logical_stream_response_id: None, + logical_stream_sequence: 0, history_rehydrated: false, input: Vec::new(), iteration: 0, diff --git a/apis/src/openai/responses/store/filter.rs b/apis/src/openai/responses/store/filter.rs index a2a6968432..cd5e040367 100644 --- a/apis/src/openai/responses/store/filter.rs +++ b/apis/src/openai/responses/store/filter.rs @@ -689,8 +689,8 @@ impl HttpFilter for ResponseStoreFilter { BodyAccess::ReadOnly } - /// `StreamBuffer` so the protocol layer assembles the complete - /// response body before delivering it at end-of-stream. + /// Streaming by default. Non-streaming Responses requests select a + /// bounded `StreamBuffer` dynamically in [`Self::on_request`]. /// /// Non-streaming Responses API payloads are bounded by output /// token limits (typically under 2 MiB). The 64 MiB ceiling is @@ -699,12 +699,16 @@ impl HttpFilter for ResponseStoreFilter { /// for the full model inference, so the hold-back latency from /// `StreamBuffer` is negligible. fn response_body_mode(&self) -> BodyMode { - BodyMode::StreamBuffer { - max_bytes: Some(MAX_JSON_BODY_BYTES), - } + BodyMode::Stream } async fn on_request(&self, ctx: &mut HttpFilterContext<'_>) -> Result { + if is_responses_format(ctx) && !is_streaming_request(ctx) { + ctx.set_response_body_mode(BodyMode::StreamBuffer { + max_bytes: Some(MAX_JSON_BODY_BYTES), + }); + } + if ctx.request.method == http::Method::GET { if let Some(action) = self.try_get_retrieval(ctx).await? { return Ok(action); diff --git a/apis/src/openai/responses/store/tests.rs b/apis/src/openai/responses/store/tests.rs index 38ae87ac34..1a7bd4f8c2 100644 --- a/apis/src/openai/responses/store/tests.rs +++ b/apis/src/openai/responses/store/tests.rs @@ -248,14 +248,32 @@ fn request_body_access_is_read_only() { } #[test] -fn response_body_mode_is_bounded_stream_buffer() { +fn response_body_mode_defaults_to_stream() { let filter = make_filter(); assert_eq!( filter.response_body_mode(), + BodyMode::Stream, + "streaming requests must not inherit a pipeline-level StreamBuffer" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn on_request_selects_bounded_stream_buffer_for_non_streaming_responses() { + let filter = make_filter(); + let req = crate::test_utils::make_request(http::Method::POST, "/v1/responses"); + let mut ctx = crate::test_utils::make_filter_context(&req); + ctx.set_metadata("openai_responses_format.format", "openai_responses"); + ctx.set_metadata("openai_responses_format.stream", "false"); + + let action = filter.on_request(&mut ctx).await.unwrap(); + + assert!(matches!(action, FilterAction::Continue), "request should continue"); + assert_eq!( + ctx.response_body_mode, BodyMode::StreamBuffer { - max_bytes: Some(67_108_864) // 64 MiB + max_bytes: Some(67_108_864) }, - "response body mode should be StreamBuffer capped at 64 MiB" + "non-streaming Responses requests should remain bounded" ); } diff --git a/apis/src/openai/responses/stream_events/accumulator.rs b/apis/src/openai/responses/stream_events/accumulator.rs index 3f971560fe..f7f6c1bfe4 100644 --- a/apis/src/openai/responses/stream_events/accumulator.rs +++ b/apis/src/openai/responses/stream_events/accumulator.rs @@ -100,6 +100,7 @@ pub(super) fn accumulate_response_object( } if let Some(Value::Array(output)) = response.get("output") { state.output_items_mut().clone_from(output); + replace_completed_tool_calls(state, output); } state.response_object = response; had_prior_usage @@ -110,6 +111,20 @@ pub(super) fn accumulate_response_object( had_prior_usage } +/// Replace incremental function calls from the authoritative terminal output. +fn replace_completed_tool_calls(state: &mut ResponsesState, output: &[Value]) { + state.tool_calls.clear(); + state.tool_calls.extend( + output + .iter() + .filter(|item| { + item.get("type").and_then(Value::as_str) == Some("function_call") + && item.get("status").and_then(Value::as_str) == Some("completed") + }) + .cloned(), + ); +} + /// Push a new output item to the incremental accumulator. fn handle_output_item_added(ctx: &mut HttpFilterContext<'_>, payload: &Value) { let state = ctx.extensions.get_or_insert_with(ResponsesState::default); diff --git a/apis/src/openai/responses/stream_events/config.rs b/apis/src/openai/responses/stream_events/config.rs index 0dc5657278..d7e185c4cb 100644 --- a/apis/src/openai/responses/stream_events/config.rs +++ b/apis/src/openai/responses/stream_events/config.rs @@ -17,6 +17,12 @@ use crate::openai::sse::SseParserConfig; #[derive(Deserialize)] #[serde(deny_unknown_fields)] pub(crate) struct StreamEventsConfig { + /// Treat successive IRR inference streams as one logical Responses + /// stream. Per-iteration lifecycle events are normalized and only the + /// final terminal event is exposed downstream. + #[serde(default)] + pub logical_stream: bool, + /// Maximum bytes buffered for incomplete SSE lines/data across /// chunk boundaries. Default: 10 MiB. #[serde(default)] diff --git a/apis/src/openai/responses/stream_events/mod.rs b/apis/src/openai/responses/stream_events/mod.rs index f3f30caf57..5590737fce 100644 --- a/apis/src/openai/responses/stream_events/mod.rs +++ b/apis/src/openai/responses/stream_events/mod.rs @@ -5,7 +5,8 @@ //! //! Parses backend SSE chunks using [`SseFrameParser`], dispatches //! typed events to update [`ResponsesState`] in request extensions. -//! The response body passes through unchanged. +//! With `logical_stream: true`, successive IRR inference streams are +//! normalized into one downstream Responses lifecycle. //! //! [`SseFrameParser`]: crate::openai::sse::SseFrameParser //! [`ResponsesState`]: super::state::ResponsesState @@ -18,8 +19,10 @@ use std::time::{Duration, Instant}; use async_trait::async_trait; use bytes::Bytes; use praxis_filter::{ - BodyAccess, BodyMode, FilterAction, FilterError, HttpFilter, HttpFilterContext, parse_filter_config, + BodyAccess, BodyMode, FilterAction, FilterError, HttpFilter, HttpFilterContext, SubRequestResponseMode, + parse_filter_config, }; +use serde_json::Value; use tracing::{debug, trace, warn}; #[cfg(test)] @@ -28,9 +31,20 @@ use self::{accumulator::accumulate_event, config::StreamEventsConfig}; use crate::{ classifier::is_responses_create, is_event_stream_content_type, - openai::sse::{SseFrameParser, SseParseError, SseParserConfig, responses::ResponsesEvent}, + openai::{ + responses::{error::responses_error_sse_payload, state::ResponsesState}, + sse::{SseFrameParser, SseParseError, SseParserConfig, responses::ResponsesEvent}, + }, }; +/// A per-turn terminal event held until the agentic transition is known. +struct DeferredTerminalEvent { + /// Canonical event type. + event_type: String, + /// Parsed event payload. + payload: Value, +} + /// Completion state observed while parsing a Responses SSE stream. #[derive(Clone, Copy, Debug, Eq, PartialEq)] pub(super) enum CompletionState { @@ -62,6 +76,16 @@ pub(super) struct StreamEventsState { tool_call_args: std::collections::HashMap, /// Cap on accumulated bytes per tool-call argument string. max_tool_call_argument_bytes: usize, + /// Whether this parser normalizes an IRR multi-round logical stream. + logical_stream: bool, + /// Inference iteration number for lifecycle suppression and index offsets. + iteration: u32, + /// Output index offset contributed by preceding inference/tool rounds. + output_index_offset: u64, + /// Terminal event withheld until completion filters publish a transition. + deferred_terminal: Option, + /// Whether a provider `[DONE]` sentinel should follow the logical terminal. + deferred_done: bool, } /// Accumulates state from native Responses API SSE event streams. @@ -71,6 +95,7 @@ pub(super) struct StreamEventsState { /// ```yaml /// filter: openai_stream_events /// # All fields optional: +/// # logical_stream: false /// # max_buffer_bytes: 10485760 /// # max_events: 100000 /// # timeout_secs: 300 @@ -81,6 +106,8 @@ pub struct OpenaiStreamEventsFilter { parser_config: SseParserConfig, /// Cap on accumulated bytes per tool-call argument string. max_tool_call_argument_bytes: usize, + /// Normalize successive IRR turns into one logical Responses stream. + logical_stream: bool, } impl OpenaiStreamEventsFilter { @@ -95,6 +122,7 @@ impl OpenaiStreamEventsFilter { Ok(Box::new(Self { parser_config: cfg.to_parser_config(), max_tool_call_argument_bytes: cfg.max_tool_call_argument_bytes(), + logical_stream: cfg.logical_stream, })) } @@ -102,6 +130,33 @@ impl OpenaiStreamEventsFilter { fn is_armed(ctx: &HttpFilterContext<'_>) -> bool { ctx.get_filter_state::().is_some() } + + /// Install fresh parser state for one inference stream. + fn arm(&self, ctx: &mut HttpFilterContext<'_>) { + let (iteration, output_index_offset) = ctx.extensions.get::().map_or((0, 0), |state| { + ( + state.iteration, + u64::try_from(state.accumulated_output.len()).unwrap_or(u64::MAX), + ) + }); + ctx.insert_filter_state(StreamEventsState { + frame_parser: SseFrameParser::new(self.parser_config.max_buffer_bytes), + event_count: 0, + max_events: self.parser_config.max_events, + timeout: self.parser_config.timeout, + started_at: None, + completed_at: None, + completion_state: CompletionState::Open, + tool_call_args: std::collections::HashMap::new(), + max_tool_call_argument_bytes: self.max_tool_call_argument_bytes, + logical_stream: self.logical_stream, + iteration, + output_index_offset, + deferred_terminal: None, + deferred_done: false, + }); + ctx.set_metadata("responses.stream_completion", "open"); + } } #[async_trait] @@ -119,7 +174,11 @@ impl HttpFilter for OpenaiStreamEventsFilter { } fn response_body_access(&self) -> BodyAccess { - BodyAccess::ReadOnly + if self.logical_stream { + BodyAccess::ReadWrite + } else { + BodyAccess::ReadOnly + } } fn response_body_mode(&self) -> BodyMode { @@ -127,23 +186,14 @@ impl HttpFilter for OpenaiStreamEventsFilter { } async fn on_request(&self, ctx: &mut HttpFilterContext<'_>) -> Result { + let typed_streaming = ctx.subrequest_response_mode() == SubRequestResponseMode::Streaming; let is_responses = is_responses_create(&ctx.request.method, ctx.request.uri.path()) - && ctx.get_metadata("openai_responses_format.format") == Some("openai_responses"); - let is_streaming = ctx.get_metadata("openai_responses_format.stream") == Some("true"); + && (typed_streaming || ctx.get_metadata("openai_responses_format.format") == Some("openai_responses")); + let is_streaming = typed_streaming || ctx.get_metadata("openai_responses_format.stream") == Some("true"); if is_responses && is_streaming { trace!("arming stream_events for streaming Responses API request"); - ctx.insert_filter_state(StreamEventsState { - frame_parser: SseFrameParser::new(self.parser_config.max_buffer_bytes), - event_count: 0, - max_events: self.parser_config.max_events, - timeout: self.parser_config.timeout, - started_at: None, - completed_at: None, - completion_state: CompletionState::Open, - tool_call_args: std::collections::HashMap::new(), - max_tool_call_argument_bytes: self.max_tool_call_argument_bytes, - }); + self.arm(ctx); } Ok(FilterAction::Continue) @@ -178,14 +228,15 @@ impl HttpFilter for OpenaiStreamEventsFilter { if end_of_stream { validate_stream_end(ctx); + finalize_logical_stream(ctx, body); } Ok(FilterAction::Continue) } } -/// Parse SSE frames and accumulate state without modifying the body. -fn process_chunk(ctx: &mut HttpFilterContext<'_>, body: &Option) { +/// Parse SSE frames, accumulating state and optionally normalizing output. +fn process_chunk(ctx: &mut HttpFilterContext<'_>, body: &mut Option) { let Some(bytes) = body.as_ref() else { return; }; @@ -197,21 +248,64 @@ fn process_chunk(ctx: &mut HttpFilterContext<'_>, body: &Option) { let now = Instant::now(); state.started_at.get_or_insert(now); - if let Err(e) = parse_and_accumulate(&mut state, ctx, bytes, now) { - warn!(error = %e, "SSE parse error in stream_events"); - ctx.set_metadata("responses.stream_parse_error", "true".to_owned()); - } + let parsed = parse_and_accumulate(&mut state, ctx, bytes, now); + handle_parse_result(ctx, body, &state, parsed); ctx.insert_filter_state(state); } +/// Publish parser state and rewrite logical-stream output when needed. +fn handle_parse_result( + ctx: &mut HttpFilterContext<'_>, + body: &mut Option, + state: &StreamEventsState, + parsed: Result, SseParseError>, +) { + let parsed = match parsed { + Ok(parsed) => parsed, + Err(error) => { + handle_parse_error(ctx, body, state, &error); + return; + }, + }; + let completion = match state.completion_state { + CompletionState::Open => "open", + CompletionState::TerminalLifecycle => "terminal", + CompletionState::Error => "error", + }; + ctx.set_metadata("responses.stream_completion", completion); + if state.logical_stream { + *body = parsed; + } +} + +/// Record a parse failure and suppress unnormalized logical-stream bytes. +fn handle_parse_error( + ctx: &mut HttpFilterContext<'_>, + body: &mut Option, + state: &StreamEventsState, + error: &SseParseError, +) { + warn!(%error, "SSE parse error in stream_events"); + ctx.set_metadata("responses.stream_parse_error", "true".to_owned()); + if state.logical_stream { + ctx.set_metadata("responses.stream_error_code", "server_error"); + ctx.set_metadata( + "responses.stream_error_message", + "upstream Responses stream could not be parsed", + ); + ctx.set_metadata("responses.skip_persist", "true"); + *body = None; + } +} + /// Parse frames from raw bytes and accumulate events. fn parse_and_accumulate( state: &mut StreamEventsState, ctx: &mut HttpFilterContext<'_>, bytes: &Bytes, now: Instant, -) -> Result<(), SseParseError> { +) -> Result, SseParseError> { check_timeout(state, now)?; let frames = state.frame_parser.parse_chunk_with_counted_event_limit( @@ -221,8 +315,12 @@ fn parse_and_accumulate( |frame| frame.data != b"[DONE]", )?; + let mut logical_output = Vec::new(); for frame in &frames { if frame.data == b"[DONE]" { + if state.logical_stream { + state.deferred_done = true; + } continue; } @@ -230,9 +328,160 @@ fn parse_and_accumulate( let event = ResponsesEvent::from_frame(frame)?; record_completion(state, &event, now)?; accumulate_event(ctx, state, &event); + if state.logical_stream { + append_logical_event(state, ctx, &event, &mut logical_output); + } } - Ok(()) + Ok(state + .logical_stream + .then(|| Bytes::from(logical_output)) + .filter(|bytes| !bytes.is_empty())) +} + +/// Append one provider event to the logical stream or defer/suppress it. +fn append_logical_event( + state: &mut StreamEventsState, + ctx: &mut HttpFilterContext<'_>, + event: &ResponsesEvent, + output: &mut Vec, +) { + if event.is_terminal() { + state.deferred_terminal = Some(DeferredTerminalEvent { + event_type: event.event_type().to_owned(), + payload: event.payload().clone(), + }); + return; + } + if state.iteration > 0 + && matches!( + event, + ResponsesEvent::ResponseCreated(_) + | ResponsesEvent::ResponseQueued(_) + | ResponsesEvent::ResponseInProgress(_) + ) + { + return; + } + + let mut payload = event.payload().clone(); + normalize_logical_payload(ctx, &mut payload, state.output_index_offset); + encode_sse_event(event.event_type(), &payload, output); +} + +/// Normalize response identity, sequence numbers, and output indices. +#[expect( + clippy::too_many_lines, + reason = "single-pass normalization of three related SSE fields" +)] +fn normalize_logical_payload(ctx: &mut HttpFilterContext<'_>, payload: &mut Value, output_index_offset: u64) { + let state = ctx.extensions.get_or_insert_with(ResponsesState::default); + if state.logical_stream_response_id.is_none() { + state.logical_stream_response_id = payload + .get("response") + .and_then(|response| response.get("id")) + .or_else(|| payload.get("response_id")) + .and_then(Value::as_str) + .map(ToOwned::to_owned); + } + let response_id = state.logical_stream_response_id.as_deref(); + if let Some(object) = payload.as_object_mut() { + if let Some(index) = object.get("output_index").and_then(Value::as_u64) { + object.insert( + "output_index".to_owned(), + Value::Number(serde_json::Number::from(index.saturating_add(output_index_offset))), + ); + } + if let Some(response_id) = response_id { + if object.contains_key("response_id") { + object.insert("response_id".to_owned(), Value::String(response_id.to_owned())); + } + if let Some(response) = object.get_mut("response").and_then(Value::as_object_mut) { + response.insert("id".to_owned(), Value::String(response_id.to_owned())); + } + } + if object.contains_key("sequence_number") { + object.insert( + "sequence_number".to_owned(), + Value::Number(serde_json::Number::from(state.logical_stream_sequence)), + ); + } + } + state.logical_stream_sequence = state.logical_stream_sequence.saturating_add(1); +} + +/// Encode one canonical single-line SSE event. +fn encode_sse_event(event_type: &str, payload: &Value, output: &mut Vec) { + output.extend_from_slice(b"event: "); + output.extend_from_slice(event_type.as_bytes()); + output.extend_from_slice(b"\ndata: "); + output.extend_from_slice(payload.to_string().as_bytes()); + output.extend_from_slice(b"\n\n"); +} + +/// Emit the held terminal event only when the current IRR step is terminal. +fn finalize_logical_stream(ctx: &mut HttpFilterContext<'_>, body: &mut Option) { + let Some(mut parser_state) = ctx.remove_filter_state::() else { + return; + }; + if !parser_state.logical_stream { + ctx.insert_filter_state(parser_state); + return; + } + + let continues = logical_stream_continues(ctx); + let mut output = Vec::new(); + if !continues && let Some(mut error) = logical_stream_error(ctx) { + normalize_logical_payload(ctx, &mut error, parser_state.output_index_offset); + encode_sse_event("error", &error, &mut output); + } else if !continues && let Some(mut terminal) = parser_state.deferred_terminal.take() { + let state = ctx.extensions.get_or_insert_with(ResponsesState::default); + let (accumulated_output, usage) = canonicalize_logical_response(state); + if let Some(response) = terminal.payload.get_mut("response").and_then(Value::as_object_mut) { + response.insert("output".to_owned(), Value::Array(accumulated_output)); + if !usage.is_null() { + response.insert("usage".to_owned(), usage); + } + } + normalize_logical_payload(ctx, &mut terminal.payload, parser_state.output_index_offset); + encode_sse_event(&terminal.event_type, &terminal.payload, &mut output); + if parser_state.deferred_done { + output.extend_from_slice(b"data: [DONE]\n\n"); + } + } + *body = (!output.is_empty()).then(|| Bytes::from(output)); + ctx.insert_filter_state(parser_state); +} + +/// Whether a dispatch filter requested another inference step. +fn logical_stream_continues(ctx: &HttpFilterContext<'_>) -> bool { + ["openai_mcp_dispatch", "openai_web_search"] + .iter() + .any(|filter| ctx.filter_results.get(filter).and_then(|results| results.get("action")) == Some("loop")) +} + +/// Return a locally generated terminal error for an already-committed stream. +fn logical_stream_error(ctx: &HttpFilterContext<'_>) -> Option { + let code = ctx.get_metadata("responses.stream_error_code")?; + let message = ctx.get_metadata("responses.stream_error_message")?; + Some(responses_error_sse_payload(code, message)) +} + +/// Make the response-store source agree with the logical SSE terminal. +fn canonicalize_logical_response(state: &mut ResponsesState) -> (Vec, Value) { + let logical_id = state.logical_stream_response_id.clone(); + let accumulated_output = state.accumulated_output.clone(); + let usage = state.usage.clone(); + if let Some(response) = state.response_object.as_object_mut() { + if let Some(logical_id) = logical_id { + response.insert("id".to_owned(), Value::String(logical_id)); + } + response.insert("output".to_owned(), Value::Array(accumulated_output.clone())); + if !usage.is_null() { + response.insert("usage".to_owned(), usage.clone()); + } + } + (accumulated_output, usage) } /// Check whether the stream has exceeded its wall-clock timeout. @@ -284,14 +533,27 @@ fn mark_complete(state: &mut StreamEventsState, new_state: CompletionState, now: /// Check that the SSE stream terminated with a terminal event. fn validate_stream_end(ctx: &mut HttpFilterContext<'_>) { - if let Some(state) = ctx.get_filter_state::() { + let incomplete_logical_stream = ctx.get_filter_state::().and_then(|state| { let checked_at = state.completed_at.unwrap_or_else(Instant::now); if let Err(e) = check_timeout(state, checked_at) { warn!(error = %e, "stream did not terminate cleanly"); - ctx.set_metadata("responses.stream_incomplete", "true".to_owned()); + Some(state.logical_stream) } else if state.completion_state == CompletionState::Open { warn!("stream did not terminate cleanly: missing terminal event"); - ctx.set_metadata("responses.stream_incomplete", "true".to_owned()); + Some(state.logical_stream) + } else { + None + } + }); + if let Some(logical_stream) = incomplete_logical_stream { + ctx.set_metadata("responses.stream_incomplete", "true".to_owned()); + if logical_stream && ctx.get_metadata("responses.stream_error_code").is_none() { + ctx.set_metadata("responses.stream_error_code", "server_error"); + ctx.set_metadata( + "responses.stream_error_message", + "upstream Responses stream did not terminate cleanly", + ); + ctx.set_metadata("responses.skip_persist", "true"); } } debug!("stream_events processing complete"); diff --git a/apis/src/openai/responses/stream_events/tests.rs b/apis/src/openai/responses/stream_events/tests.rs index 3709f988d6..420477ab6e 100644 --- a/apis/src/openai/responses/stream_events/tests.rs +++ b/apis/src/openai/responses/stream_events/tests.rs @@ -10,7 +10,7 @@ )] use bytes::Bytes; -use praxis_filter::{FilterAction, HttpFilter}; +use praxis_filter::{FilterAction, HttpFilter, SubRequestResponseMode}; use serde_json::json; use super::{CompletionState, OpenaiStreamEventsFilter, StreamEventsState, accumulate_response_object}; @@ -34,6 +34,11 @@ fn make_armed_context() -> (Box, praxis_filter::HttpFilterContex (filter, ctx) } +fn make_logical_filter() -> Box { + let yaml: serde_yaml::Value = serde_yaml::from_str("logical_stream: true").unwrap(); + OpenaiStreamEventsFilter::from_config(&yaml).unwrap() +} + #[test] fn default_config_parses() { let yaml: serde_yaml::Value = serde_yaml::from_str("{}").unwrap(); @@ -49,6 +54,16 @@ fn custom_config_overrides_apply() { assert!(filter.is_ok(), "custom config should parse"); } +#[test] +fn logical_stream_requires_response_write_access() { + let filter = make_logical_filter(); + assert_eq!( + filter.response_body_access(), + praxis_filter::BodyAccess::ReadWrite, + "logical lifecycle normalization rewrites emitted SSE frames" + ); +} + #[test] fn unknown_config_field_rejected() { let yaml: serde_yaml::Value = serde_yaml::from_str("bogus_field: true").unwrap(); @@ -105,13 +120,36 @@ fn oversized_max_tool_call_argument_bytes_rejected() { async fn arms_for_streaming_responses_request() { let (filter, mut ctx) = make_armed_context(); let action = filter.on_request(&mut ctx).await.unwrap(); - assert!(matches!(action, FilterAction::Continue)); + assert!( + matches!(action, FilterAction::Continue), + "metadata-selected streaming request should continue" + ); assert!( ctx.get_filter_state::().is_some(), "filter should be armed" ); } +#[tokio::test] +async fn arms_for_typed_streaming_selection_without_classifier_metadata() { + let filter = make_filter(); + let req = make_request(http::Method::POST, "/v1/responses"); + let mut ctx = make_filter_context(Box::leak(Box::new(req))); + ctx.set_subrequest_response_mode(SubRequestResponseMode::Streaming); + ctx.current_filter_id = Some(0); + + let action = filter.on_request(&mut ctx).await.unwrap(); + + assert!( + matches!(action, FilterAction::Continue), + "typed terminal streaming selection should continue" + ); + assert!( + ctx.get_filter_state::().is_some(), + "typed terminal streaming selection should arm the SSE parser" + ); +} + #[tokio::test] async fn does_not_arm_for_non_streaming() { let filter = make_filter(); @@ -159,7 +197,10 @@ async fn does_not_arm_for_other_responses_routes() { let action = filter.on_request(&mut ctx).await.unwrap(); - assert!(matches!(action, FilterAction::Continue)); + assert!( + matches!(action, FilterAction::Continue), + "non-create Responses route should continue" + ); assert!( ctx.get_filter_state::().is_none(), "filter should not arm for {path}" @@ -176,7 +217,10 @@ fn unarmed_filter_passes_through_body() { let mut body = Some(Bytes::from("data: {}\n\n")); let action = filter.on_response_body(&mut ctx, &mut body, false).unwrap(); - assert!(matches!(action, FilterAction::Continue)); + assert!( + matches!(action, FilterAction::Continue), + "unarmed response body should continue" + ); assert!(body.is_some(), "body should not be consumed"); } @@ -190,6 +234,210 @@ fn make_sse_chunk(event_type: &str, data: &serde_json::Value) -> Bytes { Bytes::from(format!("event: {event_type}\ndata: {data_str}\n\n")) } +#[tokio::test] +async fn logical_stream_suppresses_intermediate_terminal_and_normalizes_resumed_turn() { + let filter = make_logical_filter(); + let req = make_request(http::Method::POST, "/v1/responses"); + let mut ctx = make_filter_context(Box::leak(Box::new(req))); + ctx.set_subrequest_response_mode(SubRequestResponseMode::Streaming); + ctx.current_filter_id = Some(0); + ctx.extensions.insert(ResponsesState::from_request_body(json!({ + "model": "test-model", + "input": "hello", + "stream": true + }))); + + filter.on_request(&mut ctx).await.unwrap(); + + let mut created = Some(make_sse_chunk( + "response.created", + &json!({ + "response": {"id": "resp_first", "status": "in_progress", "output": []}, + "sequence_number": 0 + }), + )); + filter.on_response_body(&mut ctx, &mut created, false).unwrap(); + assert!( + String::from_utf8_lossy(created.as_ref().unwrap()).contains("resp_first"), + "the first response lifecycle should be emitted" + ); + + let function_call = json!({ + "type": "function_call", + "id": "fc_1", + "call_id": "call_1", + "name": "weather__get", + "arguments": "{}", + "status": "completed" + }); + let mut terminal = Some(make_sse_chunk( + "response.completed", + &json!({ + "response": {"id": "resp_first", "status": "completed", "output": [function_call.clone()]}, + "sequence_number": 1 + }), + )); + filter.on_response_body(&mut ctx, &mut terminal, false).unwrap(); + assert!(terminal.is_none(), "the per-turn terminal must be withheld"); + ctx.filter_results + .entry("openai_mcp_dispatch") + .or_default() + .set("action", "loop") + .unwrap(); + let mut first_eos = None; + filter.on_response_body(&mut ctx, &mut first_eos, true).unwrap(); + assert!( + first_eos.is_none(), + "an agentic transition must suppress the intermediate terminal" + ); + + let state = ctx.extensions.get_mut::().unwrap(); + state.iteration = 1; + state.accumulated_output = vec![function_call, json!({"type": "mcp_call", "id": "mcp_1"})]; + ctx.filter_results.remove("openai_mcp_dispatch"); + filter.on_request(&mut ctx).await.unwrap(); + + let mut resumed_created = Some(make_sse_chunk( + "response.created", + &json!({ + "response": {"id": "resp_second", "status": "in_progress", "output": []}, + "sequence_number": 0 + }), + )); + filter.on_response_body(&mut ctx, &mut resumed_created, false).unwrap(); + assert!( + resumed_created.is_none(), + "resumed lifecycle creation must be suppressed" + ); + + let mut delta = Some(make_sse_chunk( + "response.output_text.delta", + &json!({ + "response_id": "resp_second", + "output_index": 0, + "content_index": 0, + "delta": "done", + "sequence_number": 1 + }), + )); + filter.on_response_body(&mut ctx, &mut delta, false).unwrap(); + let delta = String::from_utf8(delta.unwrap().to_vec()).unwrap(); + assert!( + delta.contains(r#""response_id":"resp_first""#), + "logical response ID should remain stable: {delta}" + ); + assert!( + delta.contains(r#""output_index":2"#), + "resumed output index should include prior tool items: {delta}" + ); + + let final_message = json!({ + "type": "message", + "id": "msg_1", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "done"}] + }); + let mut final_terminal = Some(make_sse_chunk( + "response.completed", + &json!({ + "response": { + "id": "resp_second", + "status": "completed", + "output": [final_message.clone()], + "usage": {"input_tokens": 2, "output_tokens": 1, "total_tokens": 3} + }, + "sequence_number": 2 + }), + )); + filter.on_response_body(&mut ctx, &mut final_terminal, false).unwrap(); + assert!(final_terminal.is_none(), "final terminal should be held until EOS"); + ctx.extensions + .get_mut::() + .unwrap() + .accumulated_output + .push(final_message); + let mut final_eos = None; + filter.on_response_body(&mut ctx, &mut final_eos, true).unwrap(); + let final_eos = String::from_utf8(final_eos.unwrap().to_vec()).unwrap(); + assert!( + final_eos.contains("event: response.completed"), + "final terminal should be emitted: {final_eos}" + ); + assert!( + final_eos.contains(r#""id":"resp_first""#), + "terminal response ID should remain stable: {final_eos}" + ); + assert!( + final_eos.contains("mcp_1"), + "terminal output should contain the accumulated tool result: {final_eos}" + ); + assert!( + final_eos.contains("msg_1"), + "terminal output should contain the final message: {final_eos}" + ); + let state = ctx.extensions.get::().unwrap(); + assert_eq!( + state.response_object["id"], "resp_first", + "the persisted response object must use the client-visible logical ID" + ); + assert_eq!( + state.response_object["output"].as_array().map(Vec::len), + Some(3), + "the persisted response object must contain every logical-stream output item" + ); +} + +#[tokio::test] +async fn logical_stream_suppresses_malformed_chunk_and_emits_terminal_error() { + let filter = make_logical_filter(); + let req = make_request(http::Method::POST, "/v1/responses"); + let mut ctx = make_filter_context(Box::leak(Box::new(req))); + ctx.set_subrequest_response_mode(SubRequestResponseMode::Streaming); + ctx.current_filter_id = Some(0); + ctx.extensions.insert(ResponsesState::from_request_body(json!({ + "model": "test-model", + "input": "hello", + "stream": true + }))); + filter.on_request(&mut ctx).await.unwrap(); + + let mut created = Some(make_sse_chunk( + "response.created", + &json!({ + "response": {"id": "resp_first", "status": "in_progress", "output": []}, + "sequence_number": 0 + }), + )); + filter.on_response_body(&mut ctx, &mut created, false).unwrap(); + + let mut malformed = Some(Bytes::from( + "event: response.output_text.delta\ndata: {\"response_id\":\"resp_second\",bad}\n\n", + )); + filter.on_response_body(&mut ctx, &mut malformed, false).unwrap(); + assert!( + malformed.is_none(), + "a malformed logical-stream chunk must never bypass normalization" + ); + + let mut eos = None; + filter.on_response_body(&mut ctx, &mut eos, true).unwrap(); + let eos = String::from_utf8(eos.unwrap().to_vec()).unwrap(); + assert!( + eos.contains("event: error"), + "the logical stream should terminate with an SSE error: {eos}" + ); + assert!( + !eos.contains("resp_second"), + "the malformed resumed response identity must not leak downstream: {eos}" + ); + assert_eq!( + ctx.get_metadata("responses.skip_persist"), + Some("true"), + "parse-error streams must not be persisted" + ); +} + fn make_done_chunk() -> Bytes { Bytes::from("data: [DONE]\n\n") } @@ -221,6 +469,51 @@ async fn terminal_event_writes_response_object() { assert_eq!(ctx.get_metadata("responses.status"), Some("completed"),); } +#[tokio::test] +async fn terminal_event_authoritatively_populates_completed_function_calls() { + let (filter, mut ctx) = make_armed_context(); + filter.on_request(&mut ctx).await.unwrap(); + ctx.extensions.insert(ResponsesState::default()); + ctx.extensions + .get_mut::() + .unwrap() + .tool_calls + .push(json!({ + "type": "function_call", + "id": "fc_stale", + "call_id": "call_stale", + "name": "stale", + "arguments": "{}", + "status": "completed" + })); + + let completed = json!({ + "id": "resp_123", + "status": "completed", + "output": [{ + "type": "function_call", + "id": "fc_final", + "call_id": "call_final", + "name": "lookup", + "arguments": r#"{"query":"Praxis"}"#, + "status": "completed" + }] + }); + let mut body = Some(make_sse_chunk("response.completed", &completed)); + filter.on_response_body(&mut ctx, &mut body, false).unwrap(); + + let state = ctx.extensions.get::().unwrap(); + assert_eq!( + state.tool_calls.len(), + 1, + "the authoritative terminal response must replace incremental tool calls" + ); + assert_eq!( + state.tool_calls[0]["call_id"], "call_final", + "a terminal-only completed function call must be dispatchable" + ); +} + #[test] fn response_accumulation_sums_usage_across_iterations() { let req = make_request(http::Method::POST, "/v1/responses"); @@ -247,8 +540,14 @@ fn response_accumulation_sums_usage_across_iterations() { } }); - assert!(!accumulate_response_object(&mut ctx, first, None)); - assert!(accumulate_response_object(&mut ctx, second, None)); + assert!( + !accumulate_response_object(&mut ctx, first, None), + "in-progress response must not report terminal completion" + ); + assert!( + accumulate_response_object(&mut ctx, second, None), + "completed response must report terminal completion" + ); let state = ctx.extensions.get::().unwrap(); assert_eq!(state.usage["input_tokens"], 17); assert_eq!(state.usage["output_tokens"], 6); @@ -257,7 +556,10 @@ fn response_accumulation_sums_usage_across_iterations() { assert_eq!(state.response_object["usage"], state.usage); let final_without_usage = json!({"status":"completed","output":[]}); - assert!(accumulate_response_object(&mut ctx, final_without_usage, None)); + assert!( + accumulate_response_object(&mut ctx, final_without_usage, None), + "completed response without usage must remain terminal" + ); let state = ctx.extensions.get::().unwrap(); assert_eq!(state.response_object["usage"], state.usage); assert_eq!(state.usage["total_tokens"], 23); @@ -422,6 +724,46 @@ async fn eos_without_terminal_sets_incomplete() { ); } +#[tokio::test] +async fn logical_eos_without_terminal_emits_error() { + let filter = make_logical_filter(); + let req = make_request(http::Method::POST, "/v1/responses"); + let mut ctx = make_filter_context(Box::leak(Box::new(req))); + ctx.set_subrequest_response_mode(SubRequestResponseMode::Streaming); + ctx.current_filter_id = Some(0); + ctx.extensions.insert(ResponsesState::from_request_body(json!({ + "model": "test-model", + "input": "hello", + "stream": true + }))); + filter.on_request(&mut ctx).await.unwrap(); + + let mut delta = Some(make_sse_chunk( + "response.output_text.delta", + &json!({ + "response_id": "resp_partial", + "output_index": 0, + "content_index": 0, + "delta": "partial", + "sequence_number": 0 + }), + )); + filter.on_response_body(&mut ctx, &mut delta, false).unwrap(); + + let mut eos = None; + filter.on_response_body(&mut ctx, &mut eos, true).unwrap(); + let eos = String::from_utf8(eos.unwrap().to_vec()).unwrap(); + assert!( + eos.contains("event: error"), + "a logical stream must explicitly terminate when upstream omits its terminal event: {eos}" + ); + assert_eq!( + ctx.get_metadata("responses.skip_persist"), + Some("true"), + "a stream missing its terminal event must not be persisted" + ); +} + #[test] fn body_passes_through_unchanged() { let (filter, mut ctx) = make_armed_context(); @@ -435,6 +777,11 @@ fn body_passes_through_unchanged() { completion_state: CompletionState::Open, tool_call_args: std::collections::HashMap::new(), max_tool_call_argument_bytes: 1024 * 1024, + logical_stream: false, + iteration: 0, + output_index_offset: 0, + deferred_terminal: None, + deferred_done: false, }); let original = Bytes::from("event: response.created\ndata: {\"type\":\"response.created\",\"id\":\"r1\"}\n\n"); @@ -461,6 +808,11 @@ fn parse_error_sets_metadata() { completion_state: CompletionState::Open, tool_call_args: std::collections::HashMap::new(), max_tool_call_argument_bytes: 1024 * 1024, + logical_stream: false, + iteration: 0, + output_index_offset: 0, + deferred_terminal: None, + deferred_done: false, }); let large_chunk = @@ -833,7 +1185,10 @@ async fn tool_call_argument_bytes_within_limit() { async fn on_response_disarms_for_non_2xx_status() { let (filter, mut ctx) = make_armed_context(); filter.on_request(&mut ctx).await.unwrap(); - assert!(ctx.get_filter_state::().is_some()); + assert!( + ctx.get_filter_state::().is_some(), + "test setup should arm the SSE parser" + ); let resp = Box::leak(Box::new(crate::test_utils::make_response())); resp.status = http::StatusCode::BAD_REQUEST; diff --git a/apis/src/openai/responses/web_search/mod.rs b/apis/src/openai/responses/web_search/mod.rs index 9d34face9b..92f88e72ef 100644 --- a/apis/src/openai/responses/web_search/mod.rs +++ b/apis/src/openai/responses/web_search/mod.rs @@ -206,9 +206,7 @@ impl HttpFilter for WebSearchFilter { } fn response_body_mode(&self) -> BodyMode { - BodyMode::StreamBuffer { - max_bytes: Some(self.max_body_bytes), - } + BodyMode::Stream } async fn on_request(&self, _ctx: &mut HttpFilterContext<'_>) -> Result { diff --git a/apis/src/openai/sse/responses/event.rs b/apis/src/openai/sse/responses/event.rs index 1f83d953fe..93fc769866 100644 --- a/apis/src/openai/sse/responses/event.rs +++ b/apis/src/openai/sse/responses/event.rs @@ -8,7 +8,6 @@ use serde_json::Value; use super::super::{SseFrame, SseParseError}; /// A parsed Responses API streaming event. -#[expect(dead_code, reason = "variants consumed by stream_events filter (#433)")] #[derive(Debug)] pub(crate) enum ResponsesEvent { // Response lifecycle @@ -157,6 +156,37 @@ impl ResponsesEvent { ) } + /// Borrow the JSON payload carried by this event. + pub fn payload(&self) -> &Value { + match self { + Self::ResponseCreated(payload) + | Self::ResponseQueued(payload) + | Self::ResponseInProgress(payload) + | Self::ResponseCompleted(payload) + | Self::ResponseIncomplete(payload) + | Self::ResponseFailed(payload) + | Self::OutputItemAdded(payload) + | Self::OutputItemDone(payload) + | Self::ContentPartAdded(payload) + | Self::ContentPartDone(payload) + | Self::OutputTextDelta(payload) + | Self::OutputTextDone(payload) + | Self::OutputTextAnnotationAdded(payload) + | Self::FunctionCallArgumentsDelta(payload) + | Self::FunctionCallArgumentsDone(payload) + | Self::RefusalDelta(payload) + | Self::RefusalDone(payload) + | Self::ReasoningDelta(payload) + | Self::ReasoningDone(payload) + | Self::ReasoningSummaryTextDelta(payload) + | Self::ReasoningSummaryTextDone(payload) + | Self::ReasoningSummaryPartAdded(payload) + | Self::ReasoningSummaryPartDone(payload) + | Self::Error(payload) + | Self::Unknown { data: payload, .. } => payload, + } + } + /// Return the event type string. pub fn event_type(&self) -> &str { match self { diff --git a/apis/src/subrequest.rs b/apis/src/subrequest.rs index 902d06e5bd..9ee620d0e7 100644 --- a/apis/src/subrequest.rs +++ b/apis/src/subrequest.rs @@ -127,7 +127,15 @@ async fn execute_url_inner( ) -> Result { let parsed = parse_url_components(url)?; let addrs = resolve_addrs(&parsed.host, parsed.port).await?; - execute_resolved_url(client, parsed, request, &addrs, max_response_bytes, timeout).await + Box::pin(execute_resolved_url( + client, + parsed, + request, + &addrs, + max_response_bytes, + timeout, + )) + .await } /// Try each resolved address until one connects successfully. @@ -232,7 +240,7 @@ mod tests { #[test] fn parse_url_ipv6_loopback() { let parsed = parse_url_components("http://[::1]:9090/metrics").unwrap(); - assert!(!parsed.tls); + assert!(!parsed.tls, "HTTP URL should not enable TLS"); assert_eq!(parsed.host, "::1"); assert_eq!(parsed.port, 9090); assert_eq!(parsed.authority, "[::1]:9090"); @@ -241,12 +249,18 @@ mod tests { #[test] fn parse_url_missing_host_returns_error() { - assert!(parse_url_components("/relative/path").is_err()); + assert!( + parse_url_components("/relative/path").is_err(), + "relative URL should be rejected" + ); } #[test] fn parse_url_invalid_returns_error() { - assert!(parse_url_components("://bad").is_err()); + assert!( + parse_url_components("://bad").is_err(), + "malformed URL should be rejected" + ); } #[test] @@ -282,7 +296,10 @@ mod tests { std::future::pending::>(), ) .await; - assert!(matches!(result, Err(SubRequestError::DeadlineExceeded))); + assert!( + matches!(result, Err(SubRequestError::DeadlineExceeded)), + "deadline should bound the pending operation" + ); } #[tokio::test] diff --git a/docs/filters/openai_responses_proxy.md b/docs/filters/openai_responses_proxy.md index 8f8600b4cf..bb5b867503 100644 --- a/docs/filters/openai_responses_proxy.md +++ b/docs/filters/openai_responses_proxy.md @@ -11,11 +11,14 @@ Reads the assembled conversation history from `ResponsesState::messages` and rep When no `ResponsesState` exists, preserves the request body apart from removing the Praxis-owned `conversation` field. +Set `terminal_streaming: true` inside an iterative request router step to select Praxis's streaming transport when the effective outbound body contains `"stream": true`. Classifier metadata remains descriptive client intent; this final serializer owns the transport decision. IRR can resume one downstream stream across response-dependent transitions, but every response-body filter in a streaming-capable step must use `BodyMode::Stream`. + ## Configuration | Field | Type | Required | Description | |-------|------|---------|-------------| | `max_body_bytes` | integer | no | Maximum body size in bytes for `StreamBuffer` mode. | +| `terminal_streaming` | bool | no | Select Praxis streaming transport for effective `stream: true` requests. IRR may resume the same downstream stream after a step transition when its response filters use streaming body mode. | ## Examples @@ -30,4 +33,5 @@ filter: openai_responses_proxy ```yaml filter: openai_responses_proxy max_body_bytes: 67108864 +terminal_streaming: false ``` diff --git a/docs/filters/openai_stream_events.md b/docs/filters/openai_stream_events.md index 0138d14d00..480eff742e 100644 --- a/docs/filters/openai_stream_events.md +++ b/docs/filters/openai_stream_events.md @@ -13,6 +13,7 @@ All fields are optional; omitted values fall back to [`SseParserConfig`] default | Field | Type | Required | Description | |-------|------|---------|-------------| +| `logical_stream` | bool | no | Treat successive IRR inference streams as one logical Responses stream. Per-iteration lifecycle events are normalized and only the final terminal event is exposed downstream. | | `max_buffer_bytes` | integer | no | Maximum bytes buffered for incomplete SSE lines/data across chunk boundaries. Default: 10 MiB. | | `max_events` | integer | no | Maximum number of SSE events before the parser errors. Default: 100,000. | | `timeout_secs` | integer | no | Maximum seconds from first chunk to stream completion. Default: 300 (5 minutes). | @@ -23,6 +24,7 @@ All fields are optional; omitted values fall back to [`SseParserConfig`] default ```yaml filter: openai_stream_events # All fields optional: +# logical_stream: false # max_buffer_bytes: 10485760 # max_events: 100000 # timeout_secs: 300 diff --git a/examples/README.md b/examples/README.md index 6849626ef0..18e4526a56 100644 --- a/examples/README.md +++ b/examples/README.md @@ -78,6 +78,7 @@ before sending requests. | [format-routing.yaml](configs/openai/responses/format-routing.yaml) | Routes AI API traffic by detected body format | | [full-flow-agentic.yaml](configs/openai/responses/full-flow-agentic.yaml) | Extends the full-flow pipeline with an iterative_request_router (IRR) around the inference step, enabling server-side file search execution | | [full-flow.yaml](configs/openai/responses/full-flow.yaml) | Combines conversations, format classification, request validation, file resolution, and backend routing into a single pipeline | +| [irr-terminal-streaming.yaml](configs/openai/responses/irr-terminal-streaming.yaml) | Demonstrates a single-step iterative_request_router pipeline that exposes a native OpenAI Responses SSE body incrementally. `terminal_streaming: true` lets the final serializer select Praxis's typed streaming transport for an effective `"stream": true` request | | [mcp-dispatch.yaml](configs/openai/responses/mcp-dispatch.yaml) | Demonstrates the `openai_mcp_dispatch` filter configuration | | [mcp-tool-resolve.yaml](configs/openai/responses/mcp-tool-resolve.yaml) | Demonstrates the `openai_mcp_tool_resolve` filter, which resolves MCP tool entries in the Responses API `tools` array into concrete tool definitions by calling `tools/list` on each upstream MCP server | | [model-rewrite.yaml](configs/openai/responses/model-rewrite.yaml) | Rewrites or injects the top-level `model` field in Responses API request bodies before forwarding to the inference backend | diff --git a/examples/configs/openai/responses/agentic-loop.yaml b/examples/configs/openai/responses/agentic-loop.yaml index 21e348fff0..d30c508a98 100644 --- a/examples/configs/openai/responses/agentic-loop.yaml +++ b/examples/configs/openai/responses/agentic-loop.yaml @@ -47,9 +47,11 @@ # at least max_infer_iters + 1 (initial inference plus # the allowed MCP-backed inference continuations). # -# Streaming limitation: -# IRR does not support incremental streaming. All responses within -# the loop are fully buffered. +# Streaming: +# With stream=true, each inference round is streamed through the same +# downstream Responses SSE lifecycle. Intermediate response lifecycle +# events are normalized by openai_stream_events while MCP and web-search +# transitions resume inference after each upstream stream completes. # # Example requests: # @@ -103,9 +105,16 @@ filter_chains: - filter: iterative_request_router initial_step: inference max_iterations: 11 + max_stream_response_bytes: 67108864 steps: - name: inference filters: + # Parses every inference SSE stream, preserves one logical + # response identity across rounds, and withholds per-round + # terminal events until the IRR transition is known. + - filter: openai_stream_events + logical_stream: true + # Runs first on request re-entry to execute pending web # searches, and last on the model response to detect # web_search_call items. @@ -119,6 +128,7 @@ filter_chains: - filter: openai_agentic_loop max_infer_iters: 10 - filter: openai_responses_proxy + terminal_streaming: true - filter: router routes: - path_prefix: "/" diff --git a/examples/configs/openai/responses/irr-terminal-streaming.yaml b/examples/configs/openai/responses/irr-terminal-streaming.yaml new file mode 100644 index 0000000000..ffa8b0e6b8 --- /dev/null +++ b/examples/configs/openai/responses/irr-terminal-streaming.yaml @@ -0,0 +1,72 @@ +# Terminal Responses Streaming through IRR +# +# Demonstrates a single-step iterative_request_router pipeline that exposes a +# native OpenAI Responses SSE body incrementally. `terminal_streaming: true` +# lets the final serializer select Praxis's typed streaming transport for an +# effective `"stream": true` request. It is not a Praxis IRR YAML transport +# mode. +# +# `openai_responses_format` publishes descriptive client intent. The final +# outbound serializer, `openai_responses_proxy`, independently keeps the +# provider-visible `"stream": true` field aligned with Praxis's typed +# `SubRequestResponseMode::Streaming` selection. +# +# Every response filter in a streaming-capable step must use `BodyMode::Stream`. +# IRR may also resume the same downstream stream after response-dependent +# transitions; the agentic-loop example demonstrates that multi-round form. +# +# The typed response mode selected by `openai_responses_proxy` arms the +# step-local `openai_stream_events` filter. It receives each SSE chunk before it +# is sent to the client and preserves parser state through stream completion. +# +# Example: +# +# curl -N http://localhost:8080/v1/responses \ +# -H "Content-Type: application/json" \ +# -d '{"model":"gpt-4.1","input":"Say hello","stream":true}' + +listeners: + - name: ai-gateway + address: "127.0.0.1:8080" + filter_chains: [responses-terminal-stream] + +filter_chains: + - name: responses-terminal-stream + filters: + - filter: openai_responses_format + on_invalid: reject + + - filter: openai_responses_validate + + - filter: iterative_request_router + initial_step: inference + max_iterations: 1 + steps: + - name: inference + filters: + - filter: openai_responses_proxy + terminal_streaming: true + + - filter: openai_stream_events + + - filter: headers + request_set: + - name: Content-Type + value: application/json + + - filter: router + routes: + - path: "/v1/responses" + cluster: "inference-backend" + + - filter: load_balancer + clusters: + - name: "inference-backend" + endpoints: + - "127.0.0.1:3001" + on_result: + - default: true + done: true + +insecure_options: + allow_private_endpoints: true # example proxies to local backends diff --git a/filters/src/guardrails/tests.rs b/filters/src/guardrails/tests.rs index 14a7a83855..f5b197eaa1 100644 --- a/filters/src/guardrails/tests.rs +++ b/filters/src/guardrails/tests.rs @@ -270,7 +270,10 @@ provider: let mut ctx = crate::test_utils::make_filter_context(&req); let action = filter.on_request(&mut ctx).await.unwrap(); - assert!(matches!(action, praxis_filter::FilterAction::Continue)); + assert!( + matches!(action, praxis_filter::FilterAction::Continue), + "guardrails request-header phase should continue" + ); } #[tokio::test] diff --git a/server/src/server.rs b/server/src/server.rs index a98994183e..a79bee780c 100644 --- a/server/src/server.rs +++ b/server/src/server.rs @@ -307,13 +307,13 @@ fn insecure_warn(active: bool, msg: &str) { /// /// ``` /// let msg = praxis_ai::check_root_privilege(false, 0); -/// assert!(msg.is_some()); +/// assert!(msg.is_some(), "root without override should be rejected"); /// /// let msg = praxis_ai::check_root_privilege(true, 0); -/// assert!(msg.is_none()); +/// assert!(msg.is_none(), "root override should be accepted"); /// /// let msg = praxis_ai::check_root_privilege(false, 1000); -/// assert!(msg.is_none()); +/// assert!(msg.is_none(), "non-root user should be accepted"); /// ``` pub fn check_root_privilege(allow_root: bool, euid: u32) -> Option { if euid != 0 { diff --git a/tests/integration/fixtures/inference/README.md b/tests/integration/fixtures/inference/README.md index 13831e1152..c603d3ae42 100644 --- a/tests/integration/fixtures/inference/README.md +++ b/tests/integration/fixtures/inference/README.md @@ -19,7 +19,7 @@ than editing the table. -The manifest declares **13 features** across **5 scopes**, linked to **11 scenarios**. +The manifest declares **14 features** across **5 scopes**, linked to **12 scenarios**. | Scope | Feature | Status | Scenarios | Provider coverage | | --- | --- | --- | --- | --- | @@ -35,6 +35,7 @@ The manifest declares **13 features** across **5 scopes**, linked to **11 scenar | `responses_to_chat_completions` | `responses.chat.request` | `synthetic_only` | `responses/chat-basic-nonstream` | `synthetic`: `synthetic_only` | | `responses_to_chat_completions` | `responses.chat.response.text` | `synthetic_only` | `responses/chat-basic-nonstream` | `synthetic`: `synthetic_only` | | `responses_agentic_loop` | `responses.agentic.parallel_tool_calls` | `synthetic_only` | `responses/agentic-parallel-tool-calls` | `synthetic`: `synthetic_only` | +| `responses_agentic_loop` | `responses.agentic.irr_terminal_streaming` | `synthetic_only` | `responses/irr-terminal-streaming` | `synthetic`: `synthetic_only` | | `responses_to_chat_completions` | `responses.chat.continuation` | `synthetic_only` | `responses/chat-basic-nonstream` | `synthetic`: `synthetic_only` | diff --git a/tests/integration/fixtures/inference/coverage.yaml b/tests/integration/fixtures/inference/coverage.yaml index 99b53c819d..761392201d 100644 --- a/tests/integration/fixtures/inference/coverage.yaml +++ b/tests/integration/fixtures/inference/coverage.yaml @@ -132,6 +132,15 @@ features: providers: synthetic: status: synthetic_only + - id: responses.agentic.irr_terminal_streaming + scopes: + - responses_agentic_loop + status: synthetic_only + scenarios: + - responses/irr-terminal-streaming + providers: + synthetic: + status: synthetic_only - id: responses.chat.continuation scopes: - responses_to_chat_completions diff --git a/tests/integration/fixtures/inference/recordings/synthetic/responses/irr-terminal-streaming.json b/tests/integration/fixtures/inference/recordings/synthetic/responses/irr-terminal-streaming.json new file mode 100644 index 0000000000..7f24cf780e --- /dev/null +++ b/tests/integration/fixtures/inference/recordings/synthetic/responses/irr-terminal-streaming.json @@ -0,0 +1,127 @@ +{ + "version": 1, + "scenario_id": "responses/irr-terminal-streaming", + "protocol": "openai_responses", + "provenance": { + "kind": "synthetic", + "provider": "synthetic", + "model": "synthetic-responses-model", + "source_id": "controlled-irr-terminal-streaming" + }, + "normalization": { + "version": 1, + "linked_ids": { + "msg_synthetic_1": "msg_recorded_0002", + "resp_synthetic_1": "resp_recorded_0001" + } + }, + "turns": [ + { + "name": "initial", + "client": { + "request": { + "method": "POST", + "path": "/v1/responses", + "headers": { + "content-type": [ + "application/json" + ] + }, + "body": { + "kind": "json", + "value": { + "input": "Reply with exactly `STREAM-OK` and nothing else.", + "model": "synthetic-responses-model", + "store": false, + "stream": true + } + } + }, + "response": { + "status": 200, + "headers": { + "content-type": [ + "text/event-stream" + ] + }, + "body": { + "kind": "sse", + "frames": [ + { + "event": null, + "data": "{\"response\":{\"id\":\"resp_recorded_0001\",\"model\":\"synthetic-responses-model\",\"object\":\"response\",\"output\":[],\"status\":\"in_progress\"},\"sequence_number\":0,\"type\":\"response.created\"}", + "id": null, + "retry": null + }, + { + "event": null, + "data": "{\"content_index\":0,\"delta\":\"STREAM-OK\",\"output_index\":0,\"response_id\":\"resp_recorded_0001\",\"sequence_number\":1,\"type\":\"response.output_text.delta\"}", + "id": null, + "retry": null + }, + { + "event": null, + "data": "{\"response\":{\"id\":\"resp_recorded_0001\",\"model\":\"synthetic-responses-model\",\"object\":\"response\",\"output\":[{\"content\":[{\"annotations\":[],\"text\":\"STREAM-OK\",\"type\":\"output_text\"}],\"id\":\"msg_recorded_0002\",\"role\":\"assistant\",\"status\":\"completed\",\"type\":\"message\"}],\"status\":\"completed\",\"usage\":{\"input_tokens\":12,\"output_tokens\":3,\"total_tokens\":15}},\"sequence_number\":2,\"type\":\"response.completed\"}", + "id": null, + "retry": null + } + ], + "done": true + } + } + }, + "upstream": { + "request": { + "method": "POST", + "path": "/v1/responses", + "headers": { + "content-type": [ + "application/json" + ] + }, + "body": { + "kind": "json", + "value": { + "input": "Reply with exactly `STREAM-OK` and nothing else.", + "model": "synthetic-responses-model", + "store": false, + "stream": true + } + } + }, + "response": { + "status": 200, + "headers": { + "content-type": [ + "text/event-stream" + ] + }, + "body": { + "kind": "sse", + "frames": [ + { + "event": null, + "data": "{\"response\":{\"id\":\"resp_recorded_0001\",\"model\":\"synthetic-responses-model\",\"object\":\"response\",\"output\":[],\"status\":\"in_progress\"},\"sequence_number\":0,\"type\":\"response.created\"}", + "id": null, + "retry": null + }, + { + "event": null, + "data": "{\"content_index\":0,\"delta\":\"STREAM-OK\",\"output_index\":0,\"response_id\":\"resp_recorded_0001\",\"sequence_number\":1,\"type\":\"response.output_text.delta\"}", + "id": null, + "retry": null + }, + { + "event": null, + "data": "{\"response\":{\"id\":\"resp_recorded_0001\",\"model\":\"synthetic-responses-model\",\"object\":\"response\",\"output\":[{\"content\":[{\"annotations\":[],\"text\":\"STREAM-OK\",\"type\":\"output_text\"}],\"id\":\"msg_recorded_0002\",\"role\":\"assistant\",\"status\":\"completed\",\"type\":\"message\"}],\"status\":\"completed\",\"usage\":{\"input_tokens\":12,\"output_tokens\":3,\"total_tokens\":15}},\"sequence_number\":2,\"type\":\"response.completed\"}", + "id": null, + "retry": null + } + ], + "done": true + } + } + } + } + ] +} diff --git a/tests/integration/fixtures/inference/scenarios/responses/irr-terminal-streaming.yaml b/tests/integration/fixtures/inference/scenarios/responses/irr-terminal-streaming.yaml new file mode 100644 index 0000000000..8a6f27a721 --- /dev/null +++ b/tests/integration/fixtures/inference/scenarios/responses/irr-terminal-streaming.yaml @@ -0,0 +1,32 @@ +version: 1 +id: responses/irr-terminal-streaming +description: A native Responses SSE exchange remains streaming through a terminal iterative request router step. +protocol: openai_responses +example_config: openai/responses/irr-terminal-streaming.yaml +upstream_authority: 127.0.0.1:3001 +features: + - responses.agentic.irr_terminal_streaming +turns: + - name: initial + request: + method: POST + path: /v1/responses + headers: + content-type: + - application/json + body: + kind: json + value: + model: ${MODEL} + input: Reply with exactly `STREAM-OK` and nothing else. + store: false + stream: true + expect: + client_status: 200 + client_body_kind: sse + upstream_path: /v1/responses + upstream_body_kind: json + # Controlled imports use data-only SSE frames; replay still compares the + # complete normalized upstream and client exchanges byte-for-byte. + client_sse_events: [] + upstream_sse_events: [] diff --git a/tests/integration/sdk/openai/test_openai_responses_vllm.py b/tests/integration/sdk/openai/test_openai_responses_vllm.py index 21b1d7c452..7660d76ade 100644 --- a/tests/integration/sdk/openai/test_openai_responses_vllm.py +++ b/tests/integration/sdk/openai/test_openai_responses_vllm.py @@ -45,6 +45,9 @@ DATABASE_URL = os.environ.get("DATABASE_URL", "") CONFIG_PATH = "examples/configs/openai/responses/full-flow.yaml" AGENTIC_CONFIG_PATH = "examples/configs/openai/responses/agentic-loop.yaml" +IRR_STREAMING_CONFIG_PATH = ( + "examples/configs/openai/responses/irr-terminal-streaming.yaml" +) # --------------------------------------------------------------------------- # Helpers @@ -118,6 +121,19 @@ def _write_config(praxis_port: int, db_path: str) -> str: return path +def _write_irr_streaming_config(praxis_port: int) -> str: + with open(IRR_STREAMING_CONFIG_PATH) as f: + config = f.read() + + config = config.replace("127.0.0.1:8080", f"127.0.0.1:{praxis_port}") + config = config.replace("127.0.0.1:3001", _vllm_endpoint()) + + fd, path = tempfile.mkstemp(suffix=".yaml") + with os.fdopen(fd, "w") as f: + f.write(config) + return path + + def _wait_for_proxy(port: int, timeout: float = 30.0) -> None: deadline = time.monotonic() + timeout while time.monotonic() < deadline: @@ -323,6 +339,44 @@ def praxis_proxy(tmp_path_factory, request): os.unlink(config_path) +@pytest.fixture(scope="session") +def irr_streaming_proxy(tmp_path_factory, request): + """Start a Praxis proxy with terminal Responses streaming through IRR.""" + port = _free_port() + config_path = _write_irr_streaming_config(port) + binary = _find_binary() + + log_dir = tmp_path_factory.mktemp("irr-terminal-streaming") + log_path = str(log_dir / "praxis.log") + log_file = open(log_path, "w") + started = False + + proc = subprocess.Popen( + [binary, "-c", config_path], + stdout=log_file, + stderr=subprocess.STDOUT, + ) + try: + _wait_for_proxy(port) + started = True + yield port + finally: + proc.send_signal(signal.SIGINT) + try: + proc.wait(timeout=5) + except subprocess.TimeoutExpired: + proc.kill() + proc.wait() + log_file.close() + if not started or request.session.testsfailed > 0: + with open(log_path) as f: + print( + f"\n=== IRR streaming Praxis logs ===\n{f.read()}", + file=sys.stderr, + ) + os.unlink(config_path) + + @pytest.fixture(scope="session") def openai_client(praxis_proxy): """Return an OpenAI client pointed at the local Praxis proxy.""" @@ -334,6 +388,17 @@ def openai_client(praxis_proxy): ) +@pytest.fixture(scope="session") +def irr_streaming_client(irr_streaming_proxy): + """Return an OpenAI client using the terminal-streaming IRR proxy.""" + return OpenAI( + base_url=f"http://127.0.0.1:{irr_streaming_proxy}/v1", + api_key="test", + max_retries=0, + timeout=300, + ) + + # --------------------------------------------------------------------------- # Tests # --------------------------------------------------------------------------- @@ -546,8 +611,9 @@ def test_client_function_call_returns_to_client(self, openai_client): args = json.loads(fc.arguments) assert "city" in args, f"function arguments should contain city: {args}" - def test_streaming(self, openai_client): - stream = openai_client.responses.create( + def test_streaming_through_irr(self, irr_streaming_client): + """Stream a Responses request through a terminal IRR step.""" + stream = irr_streaming_client.responses.create( model=VLLM_MODEL, input="Say exactly: STREAM-OK /no_think", store=False, @@ -566,10 +632,10 @@ def test_streaming(self, openai_client): if event.type == "response.completed": final_status = event.response.status - assert event_types[0] == "response.created" - assert event_types[-1] == "response.completed" - assert final_status == "completed" - assert "STREAM-OK" in "".join(text_parts) + assert event_types[0] == "response.created", event_types + assert event_types[-1] == "response.completed", event_types + assert final_status == "completed", final_status + assert "STREAM-OK" in "".join(text_parts), text_parts # --------------------------------------------------------------------------- diff --git a/tests/integration/tests/suite/examples/irr_terminal_streaming.rs b/tests/integration/tests/suite/examples/irr_terminal_streaming.rs new file mode 100644 index 0000000000..149c639bf7 --- /dev/null +++ b/tests/integration/tests/suite/examples/irr_terminal_streaming.rs @@ -0,0 +1,355 @@ +// SPDX-License-Identifier: MIT +// Copyright (c) 2026 Praxis Contributors + +//! Functional tests for filter-selected terminal Responses streaming in IRR. + +use std::{ + collections::HashMap, + io::{Read as _, Write as _}, + net::{TcpListener, TcpStream}, + sync::mpsc, + thread, + time::Duration, +}; + +use praxis_test_utils::{free_port, json_post, load_example_config, parse_body, parse_status, start_proxy}; + +const EXAMPLE: &str = "openai/responses/irr-terminal-streaming.yaml"; +const FIRST_EVENT: &str = concat!( + "event: response.output_text.delta\n", + "data: {\"type\":\"response.output_text.delta\",\"delta\":\"hel\"}\n\n", +); +const FINAL_EVENT: &str = concat!( + "event: response.completed\n", + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_1\",", + "\"object\":\"response\",\"status\":\"completed\",\"output\":[]}}\n\n", +); + +#[test] +fn sse_event_reaches_client_before_upstream_completes() { + let (backend_port, first_sent, release, backend_thread) = start_gated_backend( + vec![ + "event: response.output_text.delta\nda".to_owned(), + "ta: {\"type\":\"response.output_text.delta\",\"delta\":\"hel\"}\n\n".to_owned(), + ], + vec![FINAL_EVENT.to_owned()], + "text/event-stream", + ); + let proxy = start_example_proxy(backend_port); + let (observed_tx, observed_rx) = mpsc::channel(); + let (complete_tx, complete_rx) = mpsc::channel(); + let proxy_addr = proxy.addr().to_owned(); + + let client = thread::spawn(move || { + let raw = read_response_incrementally( + &proxy_addr, + r#"{"model":"gpt-4.1","input":"hello","stream":true}"#, + "\"delta\":\"hel\"}\n\n", + &observed_tx, + ); + complete_tx.send(raw).expect("test receiver should remain available"); + }); + + first_sent + .recv_timeout(Duration::from_secs(2)) + .expect("backend should send the fragmented first event"); + observed_rx + .recv_timeout(Duration::from_secs(2)) + .expect("client should observe one complete SSE event while upstream is still gated"); + release + .send(()) + .expect("backend release receiver should remain available"); + + let raw = complete_rx + .recv_timeout(Duration::from_secs(3)) + .expect("client should receive the completed stream"); + assert_eq!(parse_status(&raw), 200, "terminal stream should return 200: {raw}"); + let body = parse_body(&raw); + assert!( + body.contains(FIRST_EVENT), + "fragmented event should pass through intact: {body}" + ); + assert!( + body.contains("response.completed"), + "terminal event should reach the client: {body}" + ); + + client.join().expect("client thread should not panic"); + backend_thread.join().expect("backend thread should not panic"); +} + +#[test] +fn stream_false_preserves_buffered_irr_response() { + let (backend_port, first_sent, release, backend_thread) = start_gated_backend( + vec![r#"{"id":"resp_buffered","object":"res"#.to_owned()], + vec![r#"ponse","status":"completed"}"#.to_owned()], + "application/json", + ); + let proxy = start_example_proxy(backend_port); + let (first_byte_tx, first_byte_rx) = mpsc::channel(); + let (complete_tx, complete_rx) = mpsc::channel(); + let proxy_addr = proxy.addr().to_owned(); + + let client = thread::spawn(move || { + let mut stream = connect_and_send(&proxy_addr, r#"{"model":"gpt-4.1","input":"hello","stream":false}"#); + let mut first = [0_u8; 1024]; + let count = stream.read(&mut first).expect("proxy response read should succeed"); + first_byte_tx + .send(()) + .expect("first-byte receiver should remain available"); + let mut raw = first[..count].to_vec(); + stream.read_to_end(&mut raw).expect("proxy response should complete"); + complete_tx + .send(String::from_utf8_lossy(&raw).into_owned()) + .expect("test receiver should remain available"); + }); + + first_sent + .recv_timeout(Duration::from_secs(2)) + .expect("backend should send the first JSON chunk"); + let arrived_while_incomplete = first_byte_rx.recv_timeout(Duration::from_millis(250)).is_ok(); + release + .send(()) + .expect("backend release receiver should remain available"); + + let raw = complete_rx + .recv_timeout(Duration::from_secs(3)) + .expect("buffered response should complete after upstream EOF"); + assert!( + !arrived_while_incomplete, + "stream=false must not expose headers or body before the upstream response is complete" + ); + assert_eq!(parse_status(&raw), 200, "buffered response should return 200: {raw}"); + assert_eq!( + parse_body(&raw), + r#"{"id":"resp_buffered","object":"response","status":"completed"}"#, + "buffered body should be preserved" + ); + + client.join().expect("client thread should not panic"); + backend_thread.join().expect("backend thread should not panic"); +} + +#[test] +fn downstream_cancellation_closes_terminal_upstream_stream() { + let listener = TcpListener::bind("127.0.0.1:0").expect("backend should bind"); + let backend_port = listener.local_addr().expect("backend should have an address").port(); + let (first_sent_tx, first_sent_rx) = mpsc::channel(); + let (client_dropped_tx, client_dropped_rx) = mpsc::channel(); + let (cancelled_tx, cancelled_rx) = mpsc::channel(); + let backend_thread = thread::spawn(move || { + let (mut stream, _) = listener.accept().expect("backend should accept request"); + read_request(&mut stream); + write_stream_headers(&mut stream, "text/event-stream"); + write_chunk(&mut stream, FIRST_EVENT); + stream.flush().expect("first event should flush"); + first_sent_tx.send(()).expect("test receiver should remain available"); + + client_dropped_rx + .recv_timeout(Duration::from_secs(2)) + .expect("client should drop the downstream stream"); + thread::sleep(Duration::from_millis(50)); + let payload = "x".repeat(64 * 1024); + let mut upstream_closed = false; + for _ in 0..256 { + if write!(stream, "{:x}\r\n{payload}\r\n", payload.len()) + .and_then(|()| stream.flush()) + .is_err() + { + upstream_closed = true; + break; + } + } + if !upstream_closed { + stream + .set_read_timeout(Some(Duration::from_secs(3))) + .expect("backend read timeout should be set"); + let mut byte = [0_u8; 1]; + upstream_closed = matches!(stream.read(&mut byte), Ok(0)); + } + cancelled_tx + .send(upstream_closed) + .expect("test receiver should remain available"); + }); + let proxy = start_example_proxy(backend_port); + let proxy_addr = proxy.addr().to_owned(); + + let client = thread::spawn(move || { + let mut stream = connect_and_send(&proxy_addr, r#"{"model":"gpt-4.1","input":"hello","stream":true}"#); + let mut received = Vec::new(); + let mut buffer = [0_u8; 1024]; + while !String::from_utf8_lossy(&received).contains("response.output_text.delta") { + let count = stream + .read(&mut buffer) + .expect("streaming response read should succeed"); + assert!(count > 0, "stream should not end before the first event"); + received.extend_from_slice(&buffer[..count]); + } + drop(stream); + client_dropped_tx + .send(()) + .expect("backend cancellation receiver should remain available"); + }); + + first_sent_rx + .recv_timeout(Duration::from_secs(2)) + .expect("backend should send the first event"); + client.join().expect("client thread should not panic"); + assert!( + cancelled_rx + .recv_timeout(Duration::from_secs(5)) + .expect("backend should observe cancellation"), + "dropping the downstream response must close the upstream streaming exchange" + ); + backend_thread.join().expect("backend thread should not panic"); +} + +#[test] +fn late_upstream_failure_does_not_replace_committed_sse() { + let listener = TcpListener::bind("127.0.0.1:0").expect("backend should bind"); + let backend_port = listener.local_addr().expect("backend should have an address").port(); + let backend_thread = thread::spawn(move || { + let (mut stream, _) = listener.accept().expect("backend should accept request"); + read_request(&mut stream); + write_stream_headers(&mut stream, "text/event-stream"); + write_chunk(&mut stream, FIRST_EVENT); + stream.flush().expect("first event should flush"); + stream + .write_all(b"20\r\nincomplete") + .expect("malformed late chunk should be written"); + }); + let proxy = start_example_proxy(backend_port); + let raw = read_response_to_end(proxy.addr(), r#"{"model":"gpt-4.1","input":"hello","stream":true}"#); + + assert_eq!( + parse_status(&raw), + 200, + "headers are committed before the late failure: {raw}" + ); + assert!( + raw.contains("response.output_text.delta"), + "the event delivered before the transport failure must be preserved: {raw}" + ); + assert!( + !raw.to_ascii_lowercase().contains("bad gateway") && !raw.contains("\"error\""), + "a late failure must not replace the committed SSE response: {raw}" + ); + backend_thread.join().expect("backend thread should not panic"); +} + +fn start_example_proxy(backend_port: u16) -> praxis_test_utils::ProxyGuard { + let proxy_port = free_port(); + let config = load_example_config(EXAMPLE, proxy_port, HashMap::from([("127.0.0.1:3001", backend_port)])); + start_proxy(&config) +} + +fn start_gated_backend( + first_chunks: Vec, + final_chunks: Vec, + content_type: &'static str, +) -> (u16, mpsc::Receiver<()>, mpsc::Sender<()>, thread::JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").expect("backend should bind"); + let port = listener.local_addr().expect("backend should have an address").port(); + let (first_sent_tx, first_sent_rx) = mpsc::channel(); + let (release_tx, release_rx) = mpsc::channel(); + let handle = thread::spawn(move || { + let (mut stream, _) = listener.accept().expect("backend should accept request"); + read_request(&mut stream); + write_stream_headers(&mut stream, content_type); + for chunk in first_chunks { + write_chunk(&mut stream, &chunk); + } + stream.flush().expect("initial chunks should flush"); + first_sent_tx.send(()).expect("test receiver should remain available"); + release_rx + .recv_timeout(Duration::from_secs(3)) + .expect("test should release the backend"); + for chunk in final_chunks { + write_chunk(&mut stream, &chunk); + } + stream.write_all(b"0\r\n\r\n").expect("chunked response should finish"); + stream.flush().expect("terminal chunks should flush"); + }); + (port, first_sent_rx, release_tx, handle) +} + +fn read_request(stream: &mut TcpStream) { + stream + .set_read_timeout(Some(Duration::from_secs(3))) + .expect("backend read timeout should be set"); + let mut request = Vec::new(); + let mut buffer = [0_u8; 4096]; + loop { + let count = stream.read(&mut buffer).expect("backend request read should succeed"); + assert!(count > 0, "request must complete before connection closes"); + request.extend_from_slice(&buffer[..count]); + let Some(header_end) = request.windows(4).position(|window| window == b"\r\n\r\n") else { + continue; + }; + let headers = String::from_utf8_lossy(&request[..header_end]); + let content_length = headers + .lines() + .find_map(|line| { + line.split_once(':').and_then(|(name, value)| { + name.eq_ignore_ascii_case("content-length") + .then(|| value.trim().parse::().ok()) + .flatten() + }) + }) + .unwrap_or(0); + if request.len() >= header_end + 4 + content_length { + return; + } + } +} + +fn write_stream_headers(stream: &mut TcpStream, content_type: &str) { + write!( + stream, + "HTTP/1.1 200 OK\r\nContent-Type: {content_type}\r\nTransfer-Encoding: chunked\r\nConnection: close\r\n\r\n" + ) + .expect("response headers should be written"); +} + +fn write_chunk(stream: &mut TcpStream, chunk: &str) { + write!(stream, "{:x}\r\n{chunk}\r\n", chunk.len()).expect("response chunk should be written"); +} + +fn connect_and_send(proxy_addr: &str, body: &str) -> TcpStream { + let mut stream = TcpStream::connect(proxy_addr).expect("client should connect to proxy"); + stream + .set_read_timeout(Some(Duration::from_secs(4))) + .expect("client read timeout should be set"); + stream + .write_all(json_post("/v1/responses", body).as_bytes()) + .expect("client request should be written"); + stream +} + +fn read_response_incrementally(proxy_addr: &str, body: &str, needle: &str, observed: &mpsc::Sender<()>) -> String { + let mut stream = connect_and_send(proxy_addr, body); + let mut raw = Vec::new(); + let mut buffer = [0_u8; 1024]; + let mut notified = false; + loop { + match stream.read(&mut buffer) { + Ok(0) => break, + Ok(count) => { + raw.extend_from_slice(&buffer[..count]); + if !notified && String::from_utf8_lossy(&raw).contains(needle) { + observed.send(()).expect("test receiver should remain available"); + notified = true; + } + }, + Err(error) => panic!("streaming response read failed: {error}"), + } + } + String::from_utf8_lossy(&raw).into_owned() +} + +fn read_response_to_end(proxy_addr: &str, body: &str) -> String { + let mut stream = connect_and_send(proxy_addr, body); + let mut raw = Vec::new(); + stream.read_to_end(&mut raw).expect("proxy response should end"); + String::from_utf8_lossy(&raw).into_owned() +} diff --git a/tests/integration/tests/suite/examples/mod.rs b/tests/integration/tests/suite/examples/mod.rs index 1c3427b716..d90c0047ee 100644 --- a/tests/integration/tests/suite/examples/mod.rs +++ b/tests/integration/tests/suite/examples/mod.rs @@ -22,6 +22,7 @@ mod full_flow_agentic; mod gcp_adc; mod guardrails; mod inference_fallback; +mod irr_terminal_streaming; #[cfg(feature = "http-callout-filter")] mod lakera_guard; #[cfg(feature = "llmd-ext-proc")] diff --git a/tests/integration/tests/suite/examples/openai_agentic_loop.rs b/tests/integration/tests/suite/examples/openai_agentic_loop.rs index 52061ddd84..b96d0976eb 100644 --- a/tests/integration/tests/suite/examples/openai_agentic_loop.rs +++ b/tests/integration/tests/suite/examples/openai_agentic_loop.rs @@ -7,7 +7,14 @@ //! These tests verify that IRR, request-supplied MCP resolution, //! MCP dispatch, and the agentic inference loop function together. -use std::collections::HashMap; +use std::{ + collections::HashMap, + io::{Read as _, Write as _}, + net::{TcpListener, TcpStream}, + sync::{Arc, Mutex}, + thread, + time::Duration, +}; use praxis_test_utils::{ McpMockConfig, McpToolFixture, StatefulCapturingBackend, build_pipeline, example_config_path, free_port, http_send, @@ -294,6 +301,222 @@ fn round_trip_captures_tool_and_model_requests() { ); } +#[test] +fn streaming_mcp_round_trip_uses_one_logical_sse_response() { + let first_response = vec![ + sse_event( + "response.created", + serde_json::json!({ + "response": {"id": "resp_stream_1", "object": "response", "status": "in_progress", "output": []}, + "sequence_number": 0 + }), + ), + sse_event( + "response.output_item.added", + serde_json::json!({ + "response_id": "resp_stream_1", + "output_index": 0, + "item": { + "type": "function_call", + "id": "fc_stream_1", + "call_id": "call_stream_1", + "name": "weather__get_weather", + "arguments": "", + "status": "in_progress" + }, + "sequence_number": 1 + }), + ), + sse_event( + "response.function_call_arguments.delta", + serde_json::json!({ + "response_id": "resp_stream_1", + "item_id": "fc_stream_1", + "output_index": 0, + "delta": r#"{"location":"SF"}"#, + "sequence_number": 2 + }), + ), + sse_event( + "response.function_call_arguments.done", + serde_json::json!({ + "response_id": "resp_stream_1", + "item_id": "fc_stream_1", + "output_index": 0, + "arguments": r#"{"location":"SF"}"#, + "sequence_number": 3 + }), + ), + sse_event( + "response.completed", + serde_json::json!({ + "response": { + "id": "resp_stream_1", + "object": "response", + "status": "completed", + "output": [{ + "type": "function_call", + "id": "fc_stream_1", + "call_id": "call_stream_1", + "name": "weather__get_weather", + "arguments": r#"{"location":"SF"}"#, + "status": "completed" + }], + "usage": {"input_tokens": 10, "output_tokens": 4, "total_tokens": 14} + }, + "sequence_number": 4 + }), + ), + ]; + let second_response = vec![ + sse_event( + "response.created", + serde_json::json!({ + "response": {"id": "resp_stream_2", "object": "response", "status": "in_progress", "output": []}, + "sequence_number": 0 + }), + ), + sse_event( + "response.output_item.added", + serde_json::json!({ + "response_id": "resp_stream_2", + "output_index": 0, + "item": {"type": "message", "id": "msg_stream_2", "role": "assistant", "status": "in_progress", "content": []}, + "sequence_number": 1 + }), + ), + sse_event( + "response.output_text.delta", + serde_json::json!({ + "response_id": "resp_stream_2", + "item_id": "msg_stream_2", + "output_index": 0, + "content_index": 0, + "delta": "The weather in SF is sunny.", + "sequence_number": 2 + }), + ), + sse_event( + "response.completed", + serde_json::json!({ + "response": { + "id": "resp_stream_2", + "object": "response", + "status": "completed", + "output": [{ + "type": "message", + "id": "msg_stream_2", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "The weather in SF is sunny."}] + }], + "usage": {"input_tokens": 20, "output_tokens": 7, "total_tokens": 27} + }, + "sequence_number": 3 + }), + ), + ]; + let (model_port, model_requests, model_thread) = start_streaming_model(vec![first_response, second_response]); + let mcp = start_mcp_mock_server_with_config(McpMockConfig { + tools: vec![ + McpToolFixture::new("get_weather") + .with_description("Get the weather for a location") + .with_input_schema(serde_json::json!({ + "type": "object", + "properties": {"location": {"type": "string"}}, + "required": ["location"], + "additionalProperties": false + })), + ], + ..McpMockConfig::default() + }); + let proxy_port = free_port(); + let config = load_loopback_mcp_config(proxy_port, model_port); + let proxy = start_proxy(&config); + let request = serde_json::json!({ + "model": "gpt-4.1", + "input": "What is the weather in SF?", + "stream": true, + "store": false, + "tools": [{ + "type": "mcp", + "server_label": "weather", + "server_url": format!("http://127.0.0.1:{}/mcp", mcp.port()), + "allowed_tools": ["get_weather"], + "require_approval": "never" + }] + }); + + let raw = http_send( + proxy.addr(), + &json_post("/v1/responses", &serde_json::to_string(&request).unwrap()), + ); + let body = parse_body(&raw); + + assert_eq!( + parse_status(&raw), + 200, + "streamed agentic request should return 200 (model requests: {}, MCP list: {}, MCP calls: {}): {raw}", + model_requests + .lock() + .expect("model request lock should not be poisoned") + .len(), + mcp.method_count("tools/list"), + mcp.method_count("tools/call"), + ); + assert_eq!( + body.matches("event: response.created").count(), + 1, + "one logical stream must expose one response.created event: {body}" + ); + assert_eq!( + body.matches("event: response.completed").count(), + 1, + "intermediate completion must be suppressed: {body}" + ); + assert!( + body.contains("response.function_call_arguments.delta"), + "tool-call argument deltas should reach the client: {body}" + ); + assert!( + body.contains("The weather in SF is sunny."), + "the resumed inference text should reach the same stream: {body}" + ); + assert!( + !body.contains("resp_stream_2"), + "resumed turns must retain the first logical response ID: {body}" + ); + assert!( + body.contains(r#""output_index":2"#), + "resumed model output should follow function and MCP output items: {body}" + ); + assert_eq!( + mcp.method_count("tools/call"), + 1, + "MCP tool should execute exactly once" + ); + + model_thread.join().expect("streaming model thread should finish"); + let second_request: serde_json::Value = { + let requests = model_requests + .lock() + .expect("model request lock should not be poisoned"); + assert_eq!(requests.len(), 2, "IRR should make two streamed model requests"); + serde_json::from_str(&requests[1]).expect("second request should be JSON") + }; + let input = second_request["input"] + .as_array() + .expect("second request input should be an array"); + assert!( + input.iter().any(|item| item["type"] == "function_call"), + "second inference should receive the streamed function call" + ); + assert!( + input.iter().any(|item| item["type"] == "function_call_output"), + "second inference should receive the MCP result" + ); +} + // ----------------------------------------------------------------------------- // Round-Trip: Web Search via IRR // ----------------------------------------------------------------------------- @@ -328,7 +551,7 @@ fn web_search_round_trip_executes_and_re_enters_inference() { ]) .start_with_shutdown(); - let search_listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap(); + let search_listener = TcpListener::bind("127.0.0.1:0").unwrap(); let search_port = search_listener.local_addr().unwrap().port(); spawn_search_mock(search_listener); @@ -373,7 +596,143 @@ fn web_search_round_trip_executes_and_re_enters_inference() { ); } -fn spawn_search_mock(listener: std::net::TcpListener) { +#[test] +fn streaming_web_search_round_trip_resumes_one_logical_response() { + let search_call = serde_json::json!({ + "type": "web_search_call", + "id": "ws_stream_1", + "status": "completed", + "action": {"type": "search", "query": "Rust 2025 edition"} + }); + let first_response = vec![ + sse_event( + "response.created", + serde_json::json!({ + "response": {"id": "resp_ws_stream_1", "object": "response", "status": "in_progress", "output": []}, + "sequence_number": 0 + }), + ), + sse_event( + "response.output_item.added", + serde_json::json!({ + "response_id": "resp_ws_stream_1", + "output_index": 0, + "item": search_call, + "sequence_number": 1 + }), + ), + sse_event( + "response.completed", + serde_json::json!({ + "response": { + "id": "resp_ws_stream_1", + "object": "response", + "status": "completed", + "output": [search_call], + "usage": {"input_tokens": 8, "output_tokens": 2, "total_tokens": 10} + }, + "sequence_number": 2 + }), + ), + ]; + let final_message = serde_json::json!({ + "type": "message", + "id": "msg_ws_stream_2", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "Rust search completed."}] + }); + let second_response = vec![ + sse_event( + "response.created", + serde_json::json!({ + "response": {"id": "resp_ws_stream_2", "object": "response", "status": "in_progress", "output": []}, + "sequence_number": 0 + }), + ), + sse_event( + "response.output_text.delta", + serde_json::json!({ + "response_id": "resp_ws_stream_2", + "item_id": "msg_ws_stream_2", + "output_index": 0, + "content_index": 0, + "delta": "Rust search completed.", + "sequence_number": 1 + }), + ), + sse_event( + "response.completed", + serde_json::json!({ + "response": { + "id": "resp_ws_stream_2", + "object": "response", + "status": "completed", + "output": [final_message], + "usage": {"input_tokens": 15, "output_tokens": 4, "total_tokens": 19} + }, + "sequence_number": 2 + }), + ), + ]; + let (model_port, model_requests, model_thread) = start_streaming_model(vec![first_response, second_response]); + let search_listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let search_port = search_listener.local_addr().unwrap().port(); + spawn_search_mock(search_listener); + let proxy_port = free_port(); + let config = load_web_search_config(proxy_port, model_port, search_port); + let proxy = start_proxy(&config); + let request = serde_json::json!({ + "model": "gpt-4.1", + "input": "Search for Rust 2025 edition features", + "stream": true, + "store": false, + "tools": [{"type": "web_search_preview"}] + }); + + let raw = http_send( + proxy.addr(), + &json_post("/v1/responses", &serde_json::to_string(&request).unwrap()), + ); + let body = parse_body(&raw); + + assert_eq!(parse_status(&raw), 200, "streamed web search should return 200: {raw}"); + assert_eq!( + body.matches("event: response.created").count(), + 1, + "web search should preserve one logical response lifecycle: {body}" + ); + assert_eq!( + body.matches("event: response.completed").count(), + 1, + "the web-search inference terminal should be suppressed: {body}" + ); + assert!( + body.contains("Rust search completed."), + "the post-search inference should resume in the same stream: {body}" + ); + assert!( + !body.contains("resp_ws_stream_2"), + "the resumed inference must retain the first response ID: {body}" + ); + + model_thread.join().expect("streaming model thread should finish"); + let requests = model_requests + .lock() + .expect("model request lock should not be poisoned"); + assert_eq!(requests.len(), 2, "web search should trigger a second model stream"); + let second_request: serde_json::Value = + serde_json::from_str(&requests[1]).expect("second model request should be JSON"); + drop(requests); + assert!( + second_request["input"] + .as_array() + .is_some_and(|input| input.iter().any(|item| item["type"] == "web_search_call")), + "the second inference should receive the completed web-search result" + ); +} + +fn spawn_search_mock(listener: TcpListener) { use std::io::{Read as _, Write as _}; let body = serde_json::json!({ "web": { @@ -385,7 +744,7 @@ fn spawn_search_mock(listener: std::net::TcpListener) { } }) .to_string(); - std::thread::spawn(move || { + thread::spawn(move || { let (mut stream, _) = listener.accept().unwrap(); let mut buf = [0_u8; 4096]; let _n = stream.read(&mut buf).unwrap(); @@ -397,6 +756,85 @@ fn spawn_search_mock(listener: std::net::TcpListener) { }); } +/// Encode a typed Responses event as one SSE frame. +fn sse_event(event_type: &str, mut payload: serde_json::Value) -> String { + payload + .as_object_mut() + .expect("SSE payload should be an object") + .insert("type".to_owned(), serde_json::Value::String(event_type.to_owned())); + format!("event: {event_type}\ndata: {payload}\n\n") +} + +/// Handle returned by the synthetic streaming model backend. +type StreamingModel = (u16, Arc>>, thread::JoinHandle<()>); + +/// Start a two-turn model backend that emits each SSE event as a chunk. +fn start_streaming_model(responses: Vec>) -> StreamingModel { + let listener = TcpListener::bind("127.0.0.1:0").expect("streaming model should bind"); + let port = listener + .local_addr() + .expect("streaming model should have an address") + .port(); + let requests = Arc::new(Mutex::new(Vec::new())); + let captured = Arc::clone(&requests); + let handle = thread::spawn(move || { + for response in responses { + let (mut stream, _) = listener.accept().expect("streaming model should accept request"); + stream + .set_read_timeout(Some(Duration::from_secs(5))) + .expect("streaming model should set read timeout"); + let request = read_json_request(&mut stream); + captured + .lock() + .expect("model request lock should not be poisoned") + .push(request); + stream + .write_all( + b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nTransfer-Encoding: chunked\r\nConnection: close\r\n\r\n", + ) + .expect("streaming model should write response headers"); + for event in response { + write!(stream, "{:x}\r\n{event}\r\n", event.len()).expect("streaming model should write event chunk"); + stream.flush().expect("streaming model should flush event chunk"); + } + stream + .write_all(b"0\r\n\r\n") + .expect("streaming model should finish chunked response"); + } + }); + (port, requests, handle) +} + +/// Read one content-length JSON request and return its body. +fn read_json_request(stream: &mut TcpStream) -> String { + let mut raw = Vec::new(); + let mut buffer = [0_u8; 8192]; + loop { + let read = stream.read(&mut buffer).expect("streaming model should read request"); + if read == 0 { + break; + } + raw.extend_from_slice(&buffer[..read]); + let text = String::from_utf8_lossy(&raw); + let Some((headers, body)) = text.split_once("\r\n\r\n") else { + continue; + }; + let content_length = headers + .lines() + .find_map(|line| { + let (name, value) = line.split_once(':')?; + name.eq_ignore_ascii_case("content-length") + .then(|| value.trim().parse::().ok()) + .flatten() + }) + .unwrap_or(0); + if body.len() >= content_length { + return body.get(..content_length).unwrap_or_default().to_owned(); + } + } + String::new() +} + fn load_web_search_config(proxy_port: u16, model_port: u16, search_port: u16) -> praxis_core::config::Config { let path = example_config_path("openai/responses/agentic-loop.yaml"); let yaml = std::fs::read_to_string(path).expect("read agentic-loop example"); diff --git a/tests/utils/src/inference_fixture/coverage.rs b/tests/utils/src/inference_fixture/coverage.rs index 72586b1d6f..cfbe86f810 100644 --- a/tests/utils/src/inference_fixture/coverage.rs +++ b/tests/utils/src/inference_fixture/coverage.rs @@ -1233,6 +1233,7 @@ mod tests { vec!["responses_to_chat_completions"], vec!["responses_to_chat_completions"], vec!["responses_agentic_loop"], + vec!["responses_agentic_loop"], vec!["responses_to_chat_completions"], ] ); @@ -1256,11 +1257,12 @@ mod tests { CoverageStatus::SyntheticOnly, CoverageStatus::SyntheticOnly, CoverageStatus::SyntheticOnly, + CoverageStatus::SyntheticOnly, ] ); - assert_eq!(report.features_total, 13); - assert_eq!(report.scenarios_total, 11); - assert_eq!(report.recordings_total, 16); + assert_eq!(report.features_total, 14); + assert_eq!(report.scenarios_total, 12); + assert_eq!(report.recordings_total, 17); assert_eq!( scenarios.keys().collect::>(), vec![ @@ -1272,12 +1274,13 @@ mod tests { "messages/upstream-error", "responses/agentic-parallel-tool-calls", "responses/chat-basic-nonstream", + "responses/irr-terminal-streaming", "responses/native-basic-nonstream", "responses/native-basic-stream", "responses/native-tool-call", ] ); - assert_eq!(manifest.features.len(), 13); + assert_eq!(manifest.features.len(), 14); assert_eq!(manifest.version, 1); assert_eq!( manifest @@ -1354,6 +1357,10 @@ mod tests { &"responses.agentic.parallel_tool_calls".to_owned(), &vec!["responses/agentic-parallel-tool-calls".to_owned()] ), + ( + &"responses.agentic.irr_terminal_streaming".to_owned(), + &vec!["responses/irr-terminal-streaming".to_owned()] + ), ( &"responses.chat.continuation".to_owned(), &vec!["responses/chat-basic-nonstream".to_owned()]