diff --git a/crates/agentic-server-core/src/events/mod.rs b/crates/agentic-server-core/src/events/mod.rs index 2b7b673..19dd971 100644 --- a/crates/agentic-server-core/src/events/mod.rs +++ b/crates/agentic-server-core/src/events/mod.rs @@ -2,4 +2,4 @@ pub mod normalize; pub mod types; pub use normalize::normalize_sse_line; -pub use types::{EventFrame, EventPayload, SSEEventType, SSEItemType}; +pub use types::{EventFrame, EventPayload, SSEEventType, SSEItemType, WireEvent}; diff --git a/crates/agentic-server-core/src/events/normalize.rs b/crates/agentic-server-core/src/events/normalize.rs index f81b914..afcb0c0 100644 --- a/crates/agentic-server-core/src/events/normalize.rs +++ b/crates/agentic-server-core/src/events/normalize.rs @@ -1,6 +1,6 @@ use serde_json::Value; -use super::types::{EventFrame, EventPayload, SSEEventType, SSEItemType}; +use super::types::{EventFrame, EventPayload, SSEEventType, SSEItemType, WireEvent}; use crate::utils::common::{deserialize_from_str_opt, deserialize_from_value_opt}; /// Normalize a raw SSE data line into a typed [`EventFrame`]. @@ -16,58 +16,21 @@ pub fn normalize_sse_line(line: &str) -> Option { } let json: Value = deserialize_from_str_opt(data_str)?; - let event_type = json .get("type") .and_then(Value::as_str) - .map_or(SSEEventType::Other, classify_event_type); - - let sequence_number = json.get("sequence_number").and_then(Value::as_u64); + .map_or(SSEEventType::Other, SSEEventType::from); let payload = extract_payload(event_type, &json); + let wire: WireEvent = deserialize_from_value_opt(json)?; Some(EventFrame { event_type, payload, - sequence_number, + wire, }) } -/// Map a wire-format event type string to our enum. -fn classify_event_type(type_str: &str) -> SSEEventType { - match type_str { - "response.created" => SSEEventType::ResponseCreated, - "response.in_progress" => SSEEventType::ResponseInProgress, - "response.completed" | "response.done" => SSEEventType::ResponseCompleted, - "response.failed" => SSEEventType::ResponseFailed, - "response.incomplete" => SSEEventType::ResponseIncomplete, - "response.output_item.added" => SSEEventType::OutputItemAdded, - "response.output_item.done" => SSEEventType::OutputItemDone, - "response.output_text.delta" => SSEEventType::OutputTextDelta, - "response.output_text.done" => SSEEventType::OutputTextDone, - "response.content_part.added" => SSEEventType::ContentPartAdded, - "response.content_part.done" => SSEEventType::ContentPartDone, - "response.function_call_arguments.delta" => SSEEventType::FunctionCallArgumentsDelta, - "response.function_call_arguments.done" => SSEEventType::FunctionCallArgumentsDone, - "response.custom_tool_call_input.delta" => SSEEventType::CustomToolCallInputDelta, - "response.custom_tool_call_input.done" => SSEEventType::CustomToolCallInputDone, - "response.reasoning_text.delta" => SSEEventType::ReasoningTextDelta, - "response.reasoning_text.done" => SSEEventType::ReasoningTextDone, - "response.reasoning_part.added" => SSEEventType::ReasoningPartAdded, - "response.reasoning_part.done" => SSEEventType::ReasoningPartDone, - "response.reasoning_summary_text.delta" => SSEEventType::ReasoningSummaryTextDelta, - "response.reasoning_summary_text.done" => SSEEventType::ReasoningSummaryTextDone, - "response.file_search_call.searching" => SSEEventType::FileSearchCallSearching, - "response.file_search_call.completed" => SSEEventType::FileSearchCallCompleted, - "response.web_search_call.in_progress" => SSEEventType::WebSearchCallInProgress, - "response.web_search_call.searching" => SSEEventType::WebSearchCallSearching, - "response.web_search_call.completed" => SSEEventType::WebSearchCallCompleted, - "response.mcp_tool_call.in_progress" => SSEEventType::McpToolCallInProgress, - "response.mcp_tool_call.completed" => SSEEventType::McpToolCallCompleted, - _ => SSEEventType::Other, - } -} - /// Extract a typed payload from the JSON body based on the classified event type. fn extract_payload(event_type: SSEEventType, json: &Value) -> EventPayload { match event_type { diff --git a/crates/agentic-server-core/src/events/types.rs b/crates/agentic-server-core/src/events/types.rs index 91a1e47..bf1a348 100644 --- a/crates/agentic-server-core/src/events/types.rs +++ b/crates/agentic-server-core/src/events/types.rs @@ -1,4 +1,5 @@ -use serde_json::Value; +use serde::{Deserialize, Serialize}; +use serde_json::{Map, Value}; use crate::types::io::ResponseUsage; @@ -109,6 +110,104 @@ pub enum SSEEventType { Other, } +impl From<&str> for SSEEventType { + fn from(value: &str) -> Self { + match value { + "response.created" => Self::ResponseCreated, + "response.in_progress" => Self::ResponseInProgress, + "response.completed" | "response.done" => Self::ResponseCompleted, + "response.failed" => Self::ResponseFailed, + "response.incomplete" => Self::ResponseIncomplete, + "response.output_item.added" => Self::OutputItemAdded, + "response.output_item.done" => Self::OutputItemDone, + "response.output_text.delta" => Self::OutputTextDelta, + "response.output_text.done" => Self::OutputTextDone, + "response.content_part.added" => Self::ContentPartAdded, + "response.content_part.done" => Self::ContentPartDone, + "response.function_call_arguments.delta" => Self::FunctionCallArgumentsDelta, + "response.function_call_arguments.done" => Self::FunctionCallArgumentsDone, + "response.custom_tool_call_input.delta" => Self::CustomToolCallInputDelta, + "response.custom_tool_call_input.done" => Self::CustomToolCallInputDone, + "response.reasoning_text.delta" => Self::ReasoningTextDelta, + "response.reasoning_text.done" => Self::ReasoningTextDone, + "response.reasoning_part.added" => Self::ReasoningPartAdded, + "response.reasoning_part.done" => Self::ReasoningPartDone, + "response.reasoning_summary_text.delta" => Self::ReasoningSummaryTextDelta, + "response.reasoning_summary_text.done" => Self::ReasoningSummaryTextDone, + "response.file_search_call.searching" => Self::FileSearchCallSearching, + "response.file_search_call.completed" => Self::FileSearchCallCompleted, + "response.web_search_call.in_progress" => Self::WebSearchCallInProgress, + "response.web_search_call.searching" => Self::WebSearchCallSearching, + "response.web_search_call.completed" => Self::WebSearchCallCompleted, + "response.mcp_tool_call.in_progress" => Self::McpToolCallInProgress, + "response.mcp_tool_call.completed" => Self::McpToolCallCompleted, + _ => Self::Other, + } + } +} + +impl TryFrom for &'static str { + type Error = (); + + fn try_from(value: SSEEventType) -> Result { + match value { + SSEEventType::ResponseCreated => Ok("response.created"), + SSEEventType::ResponseInProgress => Ok("response.in_progress"), + SSEEventType::ResponseCompleted => Ok("response.completed"), + SSEEventType::ResponseFailed => Ok("response.failed"), + SSEEventType::ResponseIncomplete => Ok("response.incomplete"), + SSEEventType::OutputItemAdded => Ok("response.output_item.added"), + SSEEventType::OutputItemDone => Ok("response.output_item.done"), + SSEEventType::OutputTextDelta => Ok("response.output_text.delta"), + SSEEventType::OutputTextDone => Ok("response.output_text.done"), + SSEEventType::ContentPartAdded => Ok("response.content_part.added"), + SSEEventType::ContentPartDone => Ok("response.content_part.done"), + SSEEventType::FunctionCallArgumentsDelta => Ok("response.function_call_arguments.delta"), + SSEEventType::FunctionCallArgumentsDone => Ok("response.function_call_arguments.done"), + SSEEventType::CustomToolCallInputDelta => Ok("response.custom_tool_call_input.delta"), + SSEEventType::CustomToolCallInputDone => Ok("response.custom_tool_call_input.done"), + SSEEventType::ReasoningTextDelta => Ok("response.reasoning_text.delta"), + SSEEventType::ReasoningTextDone => Ok("response.reasoning_text.done"), + SSEEventType::ReasoningPartAdded => Ok("response.reasoning_part.added"), + SSEEventType::ReasoningPartDone => Ok("response.reasoning_part.done"), + SSEEventType::ReasoningSummaryTextDelta => Ok("response.reasoning_summary_text.delta"), + SSEEventType::ReasoningSummaryTextDone => Ok("response.reasoning_summary_text.done"), + SSEEventType::FileSearchCallSearching => Ok("response.file_search_call.searching"), + SSEEventType::FileSearchCallCompleted => Ok("response.file_search_call.completed"), + SSEEventType::WebSearchCallInProgress => Ok("response.web_search_call.in_progress"), + SSEEventType::WebSearchCallSearching => Ok("response.web_search_call.searching"), + SSEEventType::WebSearchCallCompleted => Ok("response.web_search_call.completed"), + SSEEventType::McpToolCallInProgress => Ok("response.mcp_tool_call.in_progress"), + SSEEventType::McpToolCallCompleted => Ok("response.mcp_tool_call.completed"), + SSEEventType::Other => Err(()), + } + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct WireEvent { + #[serde(rename = "type", skip_serializing_if = "Option::is_none")] + pub event_type: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub sequence_number: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub output_index: Option, + #[serde(flatten)] + pub rest: Map, +} + +impl WireEvent { + #[must_use] + pub fn new(event_type: impl Into) -> Self { + Self { + event_type: Some(event_type.into()), + sequence_number: None, + output_index: None, + rest: Map::new(), + } + } +} + /// Typed payload extracted from an SSE event's JSON data. #[derive(Debug, Clone)] #[non_exhaustive] @@ -205,5 +304,70 @@ pub enum EventPayload { pub struct EventFrame { pub event_type: SSEEventType, pub payload: EventPayload, - pub sequence_number: Option, + pub wire: WireEvent, +} + +impl EventFrame { + #[must_use] + pub fn synthetic(event_type: SSEEventType, rest: Map) -> Option { + let event_type_name = <&str>::try_from(event_type).ok()?; + Some(Self { + event_type, + payload: EventPayload::None, + wire: WireEvent { + event_type: Some(event_type_name.to_owned()), + sequence_number: None, + output_index: None, + rest, + }, + }) + } + + #[must_use] + pub fn sequence_number(&self) -> Option { + self.wire.sequence_number + } +} + +#[cfg(test)] +mod tests { + use super::SSEEventType; + + #[test] + fn sse_event_type_wire_names_round_trip() { + for event_type in [ + SSEEventType::ResponseCreated, + SSEEventType::ResponseInProgress, + SSEEventType::ResponseCompleted, + SSEEventType::ResponseFailed, + SSEEventType::ResponseIncomplete, + SSEEventType::OutputItemAdded, + SSEEventType::OutputItemDone, + SSEEventType::OutputTextDelta, + SSEEventType::OutputTextDone, + SSEEventType::ContentPartAdded, + SSEEventType::ContentPartDone, + SSEEventType::FunctionCallArgumentsDelta, + SSEEventType::FunctionCallArgumentsDone, + SSEEventType::CustomToolCallInputDelta, + SSEEventType::CustomToolCallInputDone, + SSEEventType::ReasoningTextDelta, + SSEEventType::ReasoningTextDone, + SSEEventType::ReasoningPartAdded, + SSEEventType::ReasoningPartDone, + SSEEventType::ReasoningSummaryTextDelta, + SSEEventType::ReasoningSummaryTextDone, + SSEEventType::FileSearchCallSearching, + SSEEventType::FileSearchCallCompleted, + SSEEventType::WebSearchCallInProgress, + SSEEventType::WebSearchCallSearching, + SSEEventType::WebSearchCallCompleted, + SSEEventType::McpToolCallInProgress, + SSEEventType::McpToolCallCompleted, + ] { + let wire_name = <&str>::try_from(event_type).expect("known event type has a wire name"); + assert_eq!(SSEEventType::from(wire_name), event_type); + } + assert!(<&str>::try_from(SSEEventType::Other).is_err()); + } } diff --git a/crates/agentic-server-core/src/executor/accumulator.rs b/crates/agentic-server-core/src/executor/accumulator.rs index 2e4466c..19b40b5 100644 --- a/crates/agentic-server-core/src/executor/accumulator.rs +++ b/crates/agentic-server-core/src/executor/accumulator.rs @@ -194,7 +194,7 @@ impl ResponseAccumulator { fn process_stream_chunks(rx: mpsc::Receiver, conversation_id: Option) -> Self { let mut acc = Self::new(uuid7_str("resp_"), conversation_id); for line in rx { - acc.process_sse_line(&line); + let _ = acc.process_sse_line(&line); } acc.finish_stream(); acc @@ -209,7 +209,7 @@ impl ResponseAccumulator { pub fn from_sse_lines(lines: impl IntoIterator, conversation_id: Option<&str>) -> Self { let mut acc = Self::new(uuid7_str("resp_"), conversation_id.map(str::to_string)); for line in lines { - acc.process_sse_line(&line); + let _ = acc.process_sse_line(&line); } acc.finalize_all(); acc @@ -222,32 +222,32 @@ impl ResponseAccumulator { } } - pub(crate) fn process_sse_line(&mut self, line: &str) { - if let Some(frame) = normalize_sse_line(line) { - if matches!( - frame.event_type, - SSEEventType::ResponseFailed | SSEEventType::ResponseIncomplete - ) { - self.capture_terminal_details(line); - } - self.process_event(&frame); - } + pub(crate) fn process_sse_line(&mut self, line: &str) -> Option { + let frame = normalize_sse_line(line)?; + self.capture_terminal_details_if_needed(&frame); + self.process_event(&frame); + Some(frame) } - fn capture_terminal_details(&mut self, line: &str) { - let Some(data) = line.strip_prefix("data: ") else { - return; - }; - let Ok(mut event) = deserialize_from_str::(data) else { - return; - }; - let Some(response) = event.get_mut("response") else { + fn capture_terminal_details(&mut self, frame: &EventFrame) { + let Some(response) = frame.wire.rest.get("response") else { return; }; - self.incomplete_details = - deserialize_from_value_opt::(response["incomplete_details"].take()); - self.error = (!response["error"].is_null()).then(|| response["error"].take()); + self.incomplete_details = response + .get("incomplete_details") + .cloned() + .and_then(deserialize_from_value_opt::); + self.error = response.get("error").filter(|error| !error.is_null()).cloned(); + } + + fn capture_terminal_details_if_needed(&mut self, frame: &EventFrame) { + if matches!( + frame.event_type, + SSEEventType::ResponseFailed | SSEEventType::ResponseIncomplete + ) { + self.capture_terminal_details(frame); + } } pub(crate) fn finish_stream(&mut self) { @@ -431,6 +431,7 @@ impl ResponseAccumulator { #[cfg(test)] mod tests { use super::*; + use crate::events::WireEvent; #[test] fn test_accumulator_new() { @@ -525,7 +526,7 @@ mod tests { status: "in_progress".into(), usage: None, }, - sequence_number: Some(0), + wire: WireEvent::new("test"), }; acc.process_event(&frame); assert_eq!(acc.response_id, "resp_new"); @@ -541,7 +542,7 @@ mod tests { status: "in_progress".into(), usage: None, }, - sequence_number: Some(0), + wire: WireEvent::new("test"), }; acc.process_event(&frame); assert_eq!(acc.response_id, "resp_keep"); @@ -561,7 +562,7 @@ mod tests { namespace: None, call_id: None, }, - sequence_number: Some(1), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -572,7 +573,7 @@ mod tests { output_index: 0, content_index: 0, }, - sequence_number: Some(2), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { event_type: SSEEventType::OutputTextDelta, @@ -582,7 +583,7 @@ mod tests { output_index: 0, content_index: 0, }, - sequence_number: Some(3), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -592,7 +593,7 @@ mod tests { status: "completed".into(), usage: None, }, - sequence_number: Some(4), + wire: WireEvent::new("test"), }); assert_eq!(acc.status, ResponseStatus::Completed); @@ -650,7 +651,7 @@ mod tests { ..Default::default() }), }, - sequence_number: Some(9), + wire: WireEvent::new("test"), }; acc.process_event(&frame); assert_eq!(acc.status, ResponseStatus::Completed); @@ -668,7 +669,7 @@ mod tests { status: "failed".into(), usage: None, }, - sequence_number: Some(4), + wire: WireEvent::new("response.failed"), }); assert_eq!(acc.status, ResponseStatus::Error); } @@ -683,7 +684,7 @@ mod tests { status: "incomplete".into(), usage: None, }, - sequence_number: Some(4), + wire: WireEvent::new("test"), }); assert_eq!(acc.status, ResponseStatus::Incomplete); } @@ -694,7 +695,7 @@ mod tests { let frame = EventFrame { event_type: SSEEventType::ContentPartAdded, payload: EventPayload::Raw(serde_json::json!({"type": "response.content_part.added"})), - sequence_number: Some(3), + wire: WireEvent::new("test"), }; acc.process_event(&frame); assert_eq!(acc.response_id, "resp_1"); @@ -814,7 +815,7 @@ mod tests { namespace: Some("mcp__weather".into()), call_id: Some("call_abc".into()), }, - sequence_number: Some(1), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -825,7 +826,7 @@ mod tests { item_id: "fc_1".into(), output_index: 0, }, - sequence_number: Some(2), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -836,7 +837,7 @@ mod tests { item_id: "fc_1".into(), output_index: 0, }, - sequence_number: Some(3), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -848,7 +849,7 @@ mod tests { name: "get_weather".into(), output_index: 0, }, - sequence_number: Some(4), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -858,7 +859,7 @@ mod tests { status: "completed".into(), usage: None, }, - sequence_number: Some(5), + wire: WireEvent::new("test"), }); assert_eq!(acc.status, ResponseStatus::Completed); @@ -889,7 +890,7 @@ mod tests { namespace: None, call_id: Some("call_1".into()), }, - sequence_number: Some(1), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -900,7 +901,7 @@ mod tests { item_id: "fc_1".into(), output_index: 0, }, - sequence_number: Some(2), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -912,7 +913,7 @@ mod tests { name: "search".into(), output_index: 0, }, - sequence_number: Some(3), + wire: WireEvent::new("test"), }); acc.finalize_all(); @@ -938,7 +939,7 @@ mod tests { namespace: None, call_id: Some("call_1".into()), }, - sequence_number: Some(1), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { event_type: SSEEventType::FunctionCallArgumentsDone, @@ -949,7 +950,7 @@ mod tests { name: "get_weather".into(), output_index: 0, }, - sequence_number: Some(2), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -962,7 +963,7 @@ mod tests { namespace: None, call_id: Some("call_2".into()), }, - sequence_number: Some(3), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { event_type: SSEEventType::FunctionCallArgumentsDone, @@ -973,7 +974,7 @@ mod tests { name: "get_time".into(), output_index: 1, }, - sequence_number: Some(4), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -983,7 +984,7 @@ mod tests { status: "completed".into(), usage: None, }, - sequence_number: Some(5), + wire: WireEvent::new("test"), }); assert_eq!(acc.output.len(), 2); @@ -1005,7 +1006,7 @@ mod tests { namespace: None, call_id: None, }, - sequence_number: Some(1), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { event_type: SSEEventType::OutputTextDelta, @@ -1015,7 +1016,7 @@ mod tests { output_index: 0, content_index: 0, }, - sequence_number: Some(2), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -1028,7 +1029,7 @@ mod tests { namespace: None, call_id: Some("call_x".into()), }, - sequence_number: Some(3), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { event_type: SSEEventType::FunctionCallArgumentsDone, @@ -1039,7 +1040,7 @@ mod tests { name: "lookup".into(), output_index: 1, }, - sequence_number: Some(4), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -1049,7 +1050,7 @@ mod tests { status: "completed".into(), usage: None, }, - sequence_number: Some(5), + wire: WireEvent::new("test"), }); assert_eq!(acc.output.len(), 2); @@ -1071,7 +1072,7 @@ mod tests { namespace: None, call_id: Some("old_call".into()), }, - sequence_number: Some(1), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -1083,7 +1084,7 @@ mod tests { name: "new_name".into(), output_index: 0, }, - sequence_number: Some(2), + wire: WireEvent::new("test"), }); acc.finalize_all(); @@ -1109,7 +1110,7 @@ mod tests { namespace: None, call_id: Some("c1".into()), }, - sequence_number: Some(1), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -1121,7 +1122,7 @@ mod tests { name: "tool".into(), output_index: 0, }, - sequence_number: Some(2), + wire: WireEvent::new("test"), }); acc.finalize_all(); @@ -1145,7 +1146,7 @@ mod tests { item_id: String::new(), output_index: 0, }, - sequence_number: Some(1), + wire: WireEvent::new("test"), }); assert!(acc.output.is_empty()); @@ -1166,7 +1167,7 @@ mod tests { namespace: None, call_id: Some("c1".into()), }, - sequence_number: Some(1), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { event_type: SSEEventType::FunctionCallArgumentsDelta, @@ -1176,7 +1177,7 @@ mod tests { item_id: "fc_1".into(), output_index: 0, }, - sequence_number: Some(2), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -1186,7 +1187,7 @@ mod tests { status: "completed".into(), usage: None, }, - sequence_number: Some(3), + wire: WireEvent::new("test"), }); assert_eq!(acc.output.len(), 1); diff --git a/crates/agentic-server-core/src/executor/engine.rs b/crates/agentic-server-core/src/executor/engine.rs index 3cde48b..fe4fa12 100644 --- a/crates/agentic-server-core/src/executor/engine.rs +++ b/crates/agentic-server-core/src/executor/engine.rs @@ -13,19 +13,22 @@ use tokio::sync::mpsc; use tracing::{debug, warn}; use super::gateway::{ - LoopDecision, append_gateway_calls_to_new_input, append_output_items_to_input, append_tool_outputs, classify_round, - execute_and_emit_output_calls, has_client_owned_calls, public_output_items, + GatewayCallResult, LoopDecision, append_gateway_calls_to_new_input, append_output_items_to_input, + append_tool_outputs, classify_round, emit_gateway_completed_events, emit_gateway_start_events, + execute_and_emit_output_calls, execute_output_calls, gateway_event_plans, has_client_owned_calls, + is_gateway_owned_call, public_output_items, }; +use super::gateway_accumulator::{GatewayStreamAccumulator, StreamEvent, error_sse_chunk}; +use crate::events::EventFrame; use crate::executor::error::ExecutorResult; use crate::executor::inference::DONE_MARKER; use crate::executor::persist::persist_if_needed; use crate::executor::rehydrate::rehydrate_conversation; use crate::executor::request::{ExecutionContext, RequestContext}; -use crate::executor::upstream::{fetch_blocking_payload, fetch_stream_payload}; +use crate::executor::upstream::{emit_deferred_stream_events, fetch_blocking_payload, fetch_stream_payload}; use crate::tool::ToolRegistry; use crate::types::io::{OutputItem, ResponseUsage, ToolChoice}; use crate::types::request_response::{IncompleteDetails, RequestPayload, ResponsePayload}; -use crate::utils::common::serialize_to_string; pub use crate::executor::inference::BoxStream; @@ -57,17 +60,6 @@ fn accumulate_usage(total: &mut Option, usage: Option String { - let event = serde_json::json!({ - "type": "error", - "error": { - "message": message, - }, - }); - let event_json = serialize_to_string(&event).unwrap_or_else(|_| "{\"error\":\"stream error\"}".to_owned()); - format!("data: {event_json}\n\n") -} - struct AbortOnDrop { handle: tokio::task::JoinHandle, } @@ -105,7 +97,7 @@ async fn run_until_gateway_tools_complete( exec_ctx: &ExecutionContext, auth: Option<&str>, stream_upstream: bool, - stream_events: Option<&mpsc::UnboundedSender>, + mut stream: Option<(&mut GatewayStreamAccumulator, &mpsc::UnboundedSender)>, ) -> ExecutorResult<(ResponsePayload, RequestContext)> { let registry: ToolRegistry = match ctx.enriched_request.tools.as_ref() { Some(tools) => ToolRegistry::build_with_handlers(tools, &exec_ctx.gateway_executors).await?, @@ -115,10 +107,22 @@ async fn run_until_gateway_tools_complete( let mut combined_usage: Option = None; for round in 0..MAX_GATEWAY_TOOL_ROUNDS { - let mut payload: ResponsePayload = if stream_upstream { - fetch_stream_payload(&ctx, exec_ctx, auth, ®istry, stream_events).await? + let output_offset = combined_output.len(); + let (mut payload, deferred_stream_events): (ResponsePayload, Vec<_>) = if stream_upstream { + let stream_payload = fetch_stream_payload( + &ctx, + exec_ctx, + auth, + ®istry, + stream + .as_mut() + .map(|(accumulator, sender)| (&mut **accumulator, *sender)), + output_offset, + ) + .await?; + (stream_payload.payload, stream_payload.deferred_events) } else { - fetch_blocking_payload(&ctx, exec_ctx, auth).await? + (fetch_blocking_payload(&ctx, exec_ctx, auth).await?, Vec::new()) }; registry.restore_final_payload_output(&mut payload.output); accumulate_usage(&mut combined_usage, payload.usage.take()); @@ -135,8 +139,17 @@ async fn run_until_gateway_tools_complete( } } let has_client_owned = has_client_owned_calls(¤t_output, ®istry); - let gateway_results = - execute_and_emit_output_calls(¤t_output, ®istry, combined_output.len(), stream_events).await?; + let gateway_results = execute_and_emit_round_output_calls( + ¤t_output, + ®istry, + output_offset, + deferred_stream_events, + &ctx, + stream + .as_mut() + .map(|(accumulator, sender)| (&mut **accumulator, *sender)), + ) + .await?; let public_output = public_output_items(¤t_output, ®istry, &gateway_results); combined_output.extend(public_output); @@ -190,6 +203,112 @@ async fn run_until_gateway_tools_complete( unreachable!("the final round returns Done, RequiresClientAction, or Incomplete"); } +async fn execute_and_emit_round_output_calls( + output_items: &[OutputItem], + registry: &ToolRegistry, + output_offset: usize, + deferred_events: Vec, + ctx: &RequestContext, + stream: Option<(&mut GatewayStreamAccumulator, &mpsc::UnboundedSender)>, +) -> ExecutorResult> { + match (deferred_events.is_empty(), stream) { + (true, stream) => execute_and_emit_output_calls(output_items, registry, output_offset, stream).await, + (false, Some((stream_accumulator, stream_sender))) => { + execute_and_emit_ordered_output_calls( + output_items, + registry, + output_offset, + deferred_events, + ctx, + stream_accumulator, + stream_sender, + ) + .await + } + (false, None) => execute_and_emit_output_calls(output_items, registry, output_offset, None).await, + } +} + +async fn execute_and_emit_ordered_output_calls( + output_items: &[OutputItem], + registry: &ToolRegistry, + output_offset: usize, + deferred_events: Vec, + ctx: &RequestContext, + stream_accumulator: &mut GatewayStreamAccumulator, + stream_sender: &mpsc::UnboundedSender, +) -> ExecutorResult> { + let mut events_by_output = Vec::with_capacity(output_items.len()); + events_by_output.resize_with(output_items.len(), Vec::new); + let mut remaining_events = Vec::new(); + for frame in deferred_events { + let Some(output_index) = frame + .wire + .output_index + .and_then(|index| usize::try_from(index).ok()) + .filter(|index| *index < events_by_output.len()) + else { + remaining_events.push(frame); + continue; + }; + events_by_output[output_index].push(frame); + } + + let event_plans = gateway_event_plans(output_items, registry, output_offset); + let first_gateway_index = output_items + .iter() + .position(|item| matches!(item, OutputItem::FunctionCall(call) if is_gateway_owned_call(call, registry))); + let first_gateway_run_end = first_gateway_index.map_or(0, |start| { + output_items[start..] + .iter() + .take_while(|item| matches!(item, OutputItem::FunctionCall(call) if is_gateway_owned_call(call, registry))) + .count() + .saturating_add(start) + }); + let first_gateway_run_len = first_gateway_run_end.saturating_sub(first_gateway_index.unwrap_or(0)); + emit_gateway_start_events(&event_plans[..first_gateway_run_len], stream_accumulator, stream_sender)?; + + let gateway_results = execute_output_calls(output_items, registry).await?; + let mut gateway_index = 0; + for (index, item) in output_items.iter().enumerate() { + if matches!(item, OutputItem::FunctionCall(call) if is_gateway_owned_call(call, registry)) { + let plan = &event_plans[gateway_index..=gateway_index]; + let result = &gateway_results[gateway_index..=gateway_index]; + if index >= first_gateway_run_end { + emit_gateway_start_events(plan, stream_accumulator, stream_sender)?; + } + emit_gateway_completed_events(result, plan, stream_accumulator, stream_sender)?; + emit_deferred_stream_events( + std::mem::take(&mut events_by_output[index]), + ctx, + registry, + stream_accumulator, + stream_sender, + output_offset, + )?; + gateway_index += 1; + } else { + emit_deferred_stream_events( + std::mem::take(&mut events_by_output[index]), + ctx, + registry, + stream_accumulator, + stream_sender, + output_offset, + )?; + } + } + emit_deferred_stream_events( + remaining_events, + ctx, + registry, + stream_accumulator, + stream_sender, + output_offset, + )?; + Ok(gateway_results) +} + /// Move accumulated output/usage onto the terminating round's payload and /// inject the response/conversation IDs. The payload's `model`/`created_at`/ /// `status` from the latest inference turn are preserved. @@ -224,48 +343,60 @@ fn run_stream(ctx: RequestContext, exec_ctx: Arc, auth: Option Box::pin(stream! { let (event_tx, mut event_rx) = mpsc::unbounded_channel(); let exec_ctx_for_run = Arc::clone(&exec_ctx); + let event_tx_for_run = event_tx.clone(); + let stream_accumulator = GatewayStreamAccumulator::new(); let mut run_handle = AbortOnDrop::new(tokio::spawn(async move { - run_until_gateway_tools_complete( + let mut stream_accumulator = stream_accumulator; + let result = run_until_gateway_tools_complete( ctx, exec_ctx_for_run.as_ref(), auth.as_deref(), true, - Some(&event_tx), + Some((&mut stream_accumulator, &event_tx_for_run)), ) - .await + .await; + (result, stream_accumulator) })); + let mut next_sequence_number = 0; loop { tokio::select! { Some(event) = event_rx.recv() => { - yield event; + yield consume_stream_event(event, &mut next_sequence_number); } result = &mut run_handle.handle => { - while let Ok(event) = event_rx.try_recv() { - yield event; - } match result { Err(e) => { - yield error_sse_chunk(&format!("stream task failed: {e}")); - yield DONE_MARKER.to_string(); + for chunk in panicked_stream_chunks(&e, &mut event_rx, &mut next_sequence_number) { + yield chunk; + } } - Ok(Err(e)) => { - yield error_sse_chunk(&e.to_string()); + Ok((Err(e), mut stream_accumulator)) => { + while let Ok(event) = event_rx.try_recv() { + yield consume_stream_event(event, &mut next_sequence_number); + } + yield stream_accumulator.error_chunk(&e.to_string()); yield DONE_MARKER.to_string(); } - Ok(Ok((payload, ctx))) => { + Ok((Ok((payload, ctx)), mut stream_accumulator)) => { + while let Ok(event) = event_rx.try_recv() { + yield consume_stream_event(event, &mut next_sequence_number); + } // Codex may close its WebSocket as soon as it receives // `response.completed`. Persist before exposing that // event so a custom call/output continuation cannot be // cancelled by the client disconnect. - let terminal_event = payload.as_terminal_response_chunk(); + let terminal_chunk = stream_accumulator.terminal_response_chunk(&payload); let ch = exec_ctx.conv_handler.clone(); let rh = exec_ctx.resp_handler.clone(); if let Err(e) = persist_if_needed(payload, ctx, ch, rh).await { warn!("persist failed: {e}"); } - yield terminal_event; + match terminal_chunk { + Ok(chunk) => yield chunk, + Err(e) => yield stream_accumulator.error_chunk(&e.to_string()), + } yield DONE_MARKER.to_string(); } } @@ -276,6 +407,29 @@ fn run_stream(ctx: RequestContext, exec_ctx: Arc, auth: Option }) } +fn consume_stream_event(event: StreamEvent, next_sequence_number: &mut u64) -> String { + *next_sequence_number = event.sequence_number.saturating_add(1); + event.content +} + +fn stream_task_failure_chunk(error: &tokio::task::JoinError, sequence_number: u64) -> String { + error_sse_chunk(&format!("stream task failed: {error}"), sequence_number) +} + +fn panicked_stream_chunks( + error: &tokio::task::JoinError, + event_rx: &mut mpsc::UnboundedReceiver, + next_sequence_number: &mut u64, +) -> Vec { + let mut chunks = Vec::new(); + while let Ok(event) = event_rx.try_recv() { + chunks.push(consume_stream_event(event, next_sequence_number)); + } + chunks.push(stream_task_failure_chunk(error, *next_sequence_number)); + chunks.push(DONE_MARKER.to_owned()); + chunks +} + /// Create a new conversation and return its data. /// /// Exposes the conversation-creation step as a standalone function so callers @@ -357,3 +511,42 @@ pub async fn execute( ) -> ExecutorResult> { ExecuteRequest::new(request, exec_ctx).run().await } + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn stream_task_panic_after_event_uses_next_sequence_number_for_error() { + let accumulator = GatewayStreamAccumulator::new(); + let (event_tx, mut event_rx) = mpsc::unbounded_channel(); + let task = tokio::spawn(async move { + let mut accumulator = accumulator; + let event = accumulator + .process_sse_line(r#"data: {"type":"response.created"}"#, 0) + .expect("event should be emitted"); + event_tx + .send(StreamEvent { + content: "event".to_owned(), + sequence_number: event.sequence_number().expect("event should be numbered"), + }) + .expect("test receiver should remain open"); + panic!("test task panic"); + }); + + let error = task.await.expect_err("task should panic"); + let mut next_sequence_number = 0; + let chunks = panicked_stream_chunks(&error, &mut event_rx, &mut next_sequence_number); + let error_event: serde_json::Value = serde_json::from_str( + chunks[1] + .trim_end_matches('\n') + .strip_prefix("data: ") + .expect("SSE data prefix"), + ) + .expect("error chunk should be valid JSON"); + + assert_eq!(chunks[0], "event"); + assert_eq!(error_event["sequence_number"], 1); + assert_eq!(chunks[2], DONE_MARKER); + } +} diff --git a/crates/agentic-server-core/src/executor/gateway.rs b/crates/agentic-server-core/src/executor/gateway.rs index 74a5b35..5880445 100644 --- a/crates/agentic-server-core/src/executor/gateway.rs +++ b/crates/agentic-server-core/src/executor/gateway.rs @@ -2,9 +2,10 @@ use std::time::Duration; use futures::StreamExt; use futures::stream as futures_stream; -use tokio::sync::mpsc; +use crate::events::SSEEventType; use crate::executor::error::{ExecutorError, ExecutorResult}; +use crate::executor::gateway_accumulator::{GatewayStreamAccumulator, StreamEvent, emit_sse_frame, synthetic_event}; use crate::executor::request::RequestContext; use crate::tool::{GatewayDispatchResult, ToolError, ToolOutput, ToolRegistry, ToolType}; use crate::types::io::output::{FunctionToolCall, GatewayCallStatus}; @@ -84,7 +85,7 @@ pub(super) struct GatewayCallResult { pub(super) public_output: Option, } -struct GatewayCallEventPlan { +pub(super) struct GatewayCallEventPlan { call_id: String, output_index: u32, started_output: Option, @@ -100,7 +101,7 @@ fn function_calls(output_items: &[OutputItem]) -> Vec { .collect() } -fn is_gateway_owned_call(call: &FunctionToolCall, registry: &ToolRegistry) -> bool { +pub(super) fn is_gateway_owned_call(call: &FunctionToolCall, registry: &ToolRegistry) -> bool { registry .lookup(&call.name) .is_some_and(|entry| entry.tool_type.is_gateway_owned()) @@ -234,7 +235,7 @@ pub(super) fn public_output_items( .collect() } -fn gateway_event_plans( +pub(super) fn gateway_event_plans( output_items: &[OutputItem], registry: &ToolRegistry, output_offset: usize, @@ -264,57 +265,56 @@ fn gateway_event_plans( plans } -fn emit_sse_json(sender: &mpsc::UnboundedSender, event: &serde_json::Value) -> ExecutorResult<()> { - let event_json = serialize_to_string(&event).map_err(ExecutorError::JsonError)?; - sender - .send(format!("data: {event_json}\n\n")) - .map_err(|_| ExecutorError::StreamError("stream receiver closed while emitting gateway event".to_owned())) -} - fn output_item_value(item: &OutputItem) -> ExecutorResult { serde_json::to_value(item).map_err(ExecutorError::JsonError) } -fn emit_gateway_start_events( +pub(super) fn emit_gateway_start_events( plans: &[GatewayCallEventPlan], - stream_events: Option<&mpsc::UnboundedSender>, + stream_accumulator: &mut GatewayStreamAccumulator, + stream_sender: &tokio::sync::mpsc::UnboundedSender, ) -> ExecutorResult<()> { - let Some(sender) = stream_events else { - return Ok(()); - }; for plan in plans { let Some(output_item) = &plan.started_output else { continue; }; let item = output_item_value(output_item)?; - let added_event = serde_json::json!({ - "type": "response.output_item.added", - "output_index": plan.output_index, - "item": item - }); - emit_sse_json(sender, &added_event)?; + let mut added_event = synthetic_event( + SSEEventType::OutputItemAdded, + [ + ("output_index".to_owned(), serde_json::json!(plan.output_index)), + ("item".to_owned(), item), + ], + )?; + emit_gateway_event(&mut added_event, stream_accumulator, stream_sender)?; match output_item { OutputItem::WebSearchCall(web_search_call) => { - let in_progress_event = serde_json::json!({ - "type": "response.web_search_call.in_progress", - "item_id": web_search_call.id, - "output_index": plan.output_index - }); - emit_sse_json(sender, &in_progress_event)?; - let searching_event = serde_json::json!({ - "type": "response.web_search_call.searching", - "item_id": web_search_call.id, - "output_index": plan.output_index - }); - emit_sse_json(sender, &searching_event)?; + let mut in_progress_event = synthetic_event( + SSEEventType::WebSearchCallInProgress, + [ + ("item_id".to_owned(), serde_json::json!(web_search_call.id)), + ("output_index".to_owned(), serde_json::json!(plan.output_index)), + ], + )?; + emit_gateway_event(&mut in_progress_event, stream_accumulator, stream_sender)?; + let mut searching_event = synthetic_event( + SSEEventType::WebSearchCallSearching, + [ + ("item_id".to_owned(), serde_json::json!(web_search_call.id)), + ("output_index".to_owned(), serde_json::json!(plan.output_index)), + ], + )?; + emit_gateway_event(&mut searching_event, stream_accumulator, stream_sender)?; } OutputItem::McpToolCall(mcp_tool_call) => { - let in_progress_event = serde_json::json!({ - "type": "response.mcp_tool_call.in_progress", - "item_id": mcp_tool_call.id, - "output_index": plan.output_index - }); - emit_sse_json(sender, &in_progress_event)?; + let mut in_progress_event = synthetic_event( + SSEEventType::McpToolCallInProgress, + [ + ("item_id".to_owned(), serde_json::json!(mcp_tool_call.id)), + ("output_index".to_owned(), serde_json::json!(plan.output_index)), + ], + )?; + emit_gateway_event(&mut in_progress_event, stream_accumulator, stream_sender)?; } OutputItem::Message(_) | OutputItem::FunctionCall(_) @@ -326,14 +326,12 @@ fn emit_gateway_start_events( Ok(()) } -fn emit_gateway_completed_events( +pub(super) fn emit_gateway_completed_events( results: &[GatewayCallResult], plans: &[GatewayCallEventPlan], - stream_events: Option<&mpsc::UnboundedSender>, + stream_accumulator: &mut GatewayStreamAccumulator, + stream_sender: &tokio::sync::mpsc::UnboundedSender, ) -> ExecutorResult<()> { - let Some(sender) = stream_events else { - return Ok(()); - }; for result in results { let Some(public_output) = &result.public_output else { continue; @@ -344,9 +342,9 @@ fn emit_gateway_completed_events( .map_or(0, |plan| plan.output_index); let (event_type, item_id) = match public_output { OutputItem::WebSearchCall(web_search_call) => { - ("response.web_search_call.completed", web_search_call.id.as_str()) + (SSEEventType::WebSearchCallCompleted, web_search_call.id.as_str()) } - OutputItem::McpToolCall(mcp_tool_call) => ("response.mcp_tool_call.completed", mcp_tool_call.id.as_str()), + OutputItem::McpToolCall(mcp_tool_call) => (SSEEventType::McpToolCallCompleted, mcp_tool_call.id.as_str()), OutputItem::Message(_) | OutputItem::FunctionCall(_) | OutputItem::CustomToolCall(_) @@ -354,19 +352,23 @@ fn emit_gateway_completed_events( | OutputItem::Unknown => continue, }; let item = output_item_value(public_output)?; - let completed_event = serde_json::json!({ - "type": event_type, - "item_id": item_id, - "output_index": output_index, - "item": item.clone() - }); - emit_sse_json(sender, &completed_event)?; - let done_event = serde_json::json!({ - "type": "response.output_item.done", - "output_index": output_index, - "item": item - }); - emit_sse_json(sender, &done_event)?; + let mut completed_event = synthetic_event( + event_type, + [ + ("item_id".to_owned(), serde_json::json!(item_id)), + ("output_index".to_owned(), serde_json::json!(output_index)), + ("item".to_owned(), item.clone()), + ], + )?; + emit_gateway_event(&mut completed_event, stream_accumulator, stream_sender)?; + let mut done_event = synthetic_event( + SSEEventType::OutputItemDone, + [ + ("output_index".to_owned(), serde_json::json!(output_index)), + ("item".to_owned(), item), + ], + )?; + emit_gateway_event(&mut done_event, stream_accumulator, stream_sender)?; } Ok(()) } @@ -375,15 +377,33 @@ pub(super) async fn execute_and_emit_output_calls( output_items: &[OutputItem], registry: &ToolRegistry, output_offset: usize, - stream_events: Option<&mpsc::UnboundedSender>, + mut stream: Option<( + &mut GatewayStreamAccumulator, + &tokio::sync::mpsc::UnboundedSender, + )>, ) -> ExecutorResult> { let event_plans = gateway_event_plans(output_items, registry, output_offset); - emit_gateway_start_events(&event_plans, stream_events)?; + if let Some((stream_accumulator, stream_sender)) = stream.as_mut() { + emit_gateway_start_events(&event_plans, stream_accumulator, stream_sender)?; + } let gateway_results = execute_output_calls(output_items, registry).await?; - emit_gateway_completed_events(&gateway_results, &event_plans, stream_events)?; + if let Some((stream_accumulator, stream_sender)) = stream.as_mut() { + emit_gateway_completed_events(&gateway_results, &event_plans, stream_accumulator, stream_sender)?; + } Ok(gateway_results) } +fn emit_gateway_event( + frame: &mut crate::events::EventFrame, + stream_accumulator: &mut GatewayStreamAccumulator, + stream_sender: &tokio::sync::mpsc::UnboundedSender, +) -> ExecutorResult<()> { + if stream_accumulator.process_event(frame, 0) { + emit_sse_frame(stream_sender, frame)?; + } + Ok(()) +} + pub(super) fn append_input_item(input: &mut ResponsesInput, item: InputItem) { match input { ResponsesInput::Items(items) => items.push(item), diff --git a/crates/agentic-server-core/src/executor/gateway_accumulator.rs b/crates/agentic-server-core/src/executor/gateway_accumulator.rs new file mode 100644 index 0000000..e545286 --- /dev/null +++ b/crates/agentic-server-core/src/executor/gateway_accumulator.rs @@ -0,0 +1,227 @@ +use crate::events::{EventFrame, EventPayload, SSEEventType, WireEvent, normalize_sse_line}; +use crate::executor::error::{ExecutorError, ExecutorResult}; +use crate::types::request_response::ResponsePayload; +use crate::utils::common::{serialize_to_string, serialize_to_value}; +use serde_json::Value; + +pub struct GatewayStreamAccumulator { + next_sequence_number: u64, + emitted_created: bool, + emitted_in_progress: bool, +} + +pub(super) struct StreamEvent { + pub(super) content: String, + pub(super) sequence_number: u64, +} + +impl GatewayStreamAccumulator { + #[must_use] + pub fn new() -> Self { + Self { + next_sequence_number: 0, + emitted_created: false, + emitted_in_progress: false, + } + } + + pub fn process_sse_line(&mut self, line: &str, output_offset: usize) -> Option { + let mut frame = normalize_sse_line(line)?; + self.process_event(&mut frame, output_offset).then_some(frame) + } + + #[must_use] + pub fn process_event(&mut self, frame: &mut EventFrame, output_offset: usize) -> bool { + if !self.should_emit_lifecycle(frame.event_type) { + return false; + } + self.stamp_event(frame, output_offset); + true + } + + fn stamp_event(&mut self, frame: &mut EventFrame, output_offset: usize) { + frame.wire.sequence_number = Some(self.take_sequence_number()); + rebase_output_index(&mut frame.wire, output_offset); + } + + pub(crate) fn terminal_response_chunk(&mut self, payload: &ResponsePayload) -> ExecutorResult { + let mut frame = terminal_response_frame(payload)?; + self.stamp_event(&mut frame, 0); + serialize_sse_frame(&frame) + } + + pub(crate) fn error_chunk(&mut self, message: &str) -> String { + let mut frame = error_frame(message); + self.stamp_event(&mut frame, 0); + serialize_sse_frame(&frame).unwrap_or_else(|_| error_sse_chunk(message, frame.sequence_number().unwrap_or(0))) + } + + fn should_emit_lifecycle(&mut self, event_type: SSEEventType) -> bool { + match event_type { + SSEEventType::ResponseCreated => take_once(&mut self.emitted_created), + SSEEventType::ResponseInProgress => take_once(&mut self.emitted_in_progress), + _ => true, + } + } + + fn take_sequence_number(&mut self) -> u64 { + let sequence_number = self.next_sequence_number; + self.next_sequence_number = self.next_sequence_number.saturating_add(1); + sequence_number + } +} + +impl Default for GatewayStreamAccumulator { + fn default() -> Self { + Self::new() + } +} + +fn take_once(already_taken: &mut bool) -> bool { + if *already_taken { + false + } else { + *already_taken = true; + true + } +} + +fn rebase_output_index(wire: &mut WireEvent, output_offset: usize) { + let Some(offset) = u64::try_from(output_offset).ok().filter(|offset| *offset > 0) else { + return; + }; + if let Some(index) = wire.output_index { + wire.output_index = Some(index.saturating_add(offset)); + } +} + +fn terminal_response_frame(payload: &ResponsePayload) -> ExecutorResult { + let event_type = match payload.terminal_event_type() { + "response.incomplete" => SSEEventType::ResponseIncomplete, + "response.failed" => SSEEventType::ResponseFailed, + "response.in_progress" => SSEEventType::ResponseInProgress, + _ => SSEEventType::ResponseCompleted, + }; + let mut rest = serde_json::Map::new(); + rest.insert( + "response".to_owned(), + serialize_to_value(payload).map_err(ExecutorError::JsonError)?, + ); + EventFrame::synthetic(event_type, rest) + .ok_or_else(|| ExecutorError::StreamError("terminal response event has no wire representation".to_owned())) +} + +fn error_frame(message: &str) -> EventFrame { + let mut wire = WireEvent::new("error"); + wire.rest.insert( + "error".to_owned(), + serde_json::json!({ + "message": message, + }), + ); + EventFrame { + event_type: SSEEventType::Other, + payload: EventPayload::None, + wire, + } +} + +pub(super) fn error_sse_chunk(message: &str, sequence_number: u64) -> String { + let mut frame = error_frame(message); + frame.wire.sequence_number = Some(sequence_number); + serialize_sse_frame(&frame) + .unwrap_or_else(|_| format!("data: {{\"type\":\"error\",\"sequence_number\":{sequence_number}}}\n\n")) +} + +pub(super) fn synthetic_event( + event_type: SSEEventType, + rest: impl IntoIterator, +) -> ExecutorResult { + EventFrame::synthetic(event_type, rest.into_iter().collect()) + .ok_or_else(|| ExecutorError::StreamError("synthetic event has no wire representation".to_owned())) +} + +pub(super) fn emit_sse_frame( + sender: &tokio::sync::mpsc::UnboundedSender, + frame: &EventFrame, +) -> ExecutorResult<()> { + let sequence_number = frame + .sequence_number() + .ok_or_else(|| ExecutorError::StreamError("stream event has no sequence number".to_owned()))?; + sender + .send(StreamEvent { + content: serialize_sse_frame(frame)?, + sequence_number, + }) + .map_err(|_| ExecutorError::StreamError("stream receiver closed while emitting gateway event".to_owned())) +} + +fn serialize_sse_frame(frame: &EventFrame) -> ExecutorResult { + let event_json = serialize_to_string(&frame.wire).map_err(ExecutorError::JsonError)?; + Ok(format!("data: {event_json}\n\n")) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn process_sse_line_numbers_and_rebases_output_index() { + let mut accumulator = GatewayStreamAccumulator::new(); + let frame = accumulator + .process_sse_line( + r#"data: {"type":"response.output_text.delta","output_index":2,"delta":"hi"}"#, + 3, + ) + .expect("line should normalize"); + + assert_eq!(frame.sequence_number(), Some(0)); + assert_eq!(frame.wire.sequence_number, Some(0)); + assert_eq!(frame.wire.output_index, Some(5)); + assert_eq!(frame.wire.rest["delta"], "hi"); + } + + #[test] + fn error_sse_chunk_escapes_error_messages() { + let chunk = error_sse_chunk("task failed: \"unexpected\"\nretry", 7); + let data = chunk + .trim_end_matches('\n') + .strip_prefix("data: ") + .expect("SSE data prefix"); + let event: serde_json::Value = serde_json::from_str(data).expect("valid error event JSON"); + + assert_eq!(event["type"], "error"); + assert_eq!(event["sequence_number"], 7); + assert_eq!(event["error"]["message"], "task failed: \"unexpected\"\nretry"); + } + + #[test] + fn emits_in_progress_terminal_event_after_lifecycle_event() { + let mut accumulator = GatewayStreamAccumulator::new(); + accumulator + .process_sse_line(r#"data: {"type":"response.in_progress"}"#, 0) + .expect("first lifecycle event should be emitted"); + + let payload: ResponsePayload = serde_json::from_value(serde_json::json!({ + "id": "resp_1", + "object": "response", + "created_at": 0, + "model": "test", + "status": "in_progress", + "output": [], + "usage": null, + "incomplete_details": null, + "error": null, + "previous_response_id": null, + "conversation_id": null, + "instructions": null + })) + .expect("valid response payload"); + + let chunk = accumulator + .terminal_response_chunk(&payload) + .expect("terminal event serializes"); + assert!(chunk.contains("\"type\":\"response.in_progress\"")); + assert!(chunk.contains("\"sequence_number\":1")); + } +} diff --git a/crates/agentic-server-core/src/executor/mod.rs b/crates/agentic-server-core/src/executor/mod.rs index daf7713..134fd12 100644 --- a/crates/agentic-server-core/src/executor/mod.rs +++ b/crates/agentic-server-core/src/executor/mod.rs @@ -12,6 +12,7 @@ pub mod rehydrate; pub mod request; mod gateway; +pub mod gateway_accumulator; mod upstream; pub use engine::{BoxStream, ExecuteRequest, create_conversation, execute}; diff --git a/crates/agentic-server-core/src/executor/upstream.rs b/crates/agentic-server-core/src/executor/upstream.rs index 895b142..14eea04 100644 --- a/crates/agentic-server-core/src/executor/upstream.rs +++ b/crates/agentic-server-core/src/executor/upstream.rs @@ -3,16 +3,29 @@ use std::sync::Arc; use futures::StreamExt; use serde_json::Value; -use tokio::sync::mpsc; -use crate::events::{EventPayload, SSEEventType, SSEItemType, normalize_sse_line}; +use crate::events::{EventFrame, EventPayload, SSEEventType, SSEItemType, WireEvent}; use crate::executor::accumulator::ResponseAccumulator; use crate::executor::error::{ExecutorError, ExecutorResult}; +use crate::executor::gateway_accumulator::{GatewayStreamAccumulator, StreamEvent, emit_sse_frame}; use crate::executor::inference::{call_inference, fetch_response_json}; use crate::executor::request::{ExecutionContext, RequestContext}; use crate::tool::ToolRegistry; use crate::types::request_response::ResponsePayload; -use crate::utils::common::{deserialize_from_str, serialize_to_string}; +use crate::utils::common::serialize_to_string; + +struct StreamEmitContext<'a> { + request: &'a RequestContext, + registry: &'a ToolRegistry, + sender: &'a tokio::sync::mpsc::UnboundedSender, + accumulator: &'a mut GatewayStreamAccumulator, + output_offset: usize, +} + +pub(super) struct StreamPayload { + pub(super) payload: ResponsePayload, + pub(super) deferred_events: Vec, +} pub(super) async fn fetch_blocking_payload( ctx: &RequestContext, @@ -42,8 +55,12 @@ pub(super) async fn fetch_stream_payload( exec_ctx: &ExecutionContext, auth: Option<&str>, registry: &ToolRegistry, - stream_events: Option<&mpsc::UnboundedSender>, -) -> ExecutorResult { + mut stream: Option<( + &mut GatewayStreamAccumulator, + &tokio::sync::mpsc::UnboundedSender, + )>, + output_offset: usize, +) -> ExecutorResult { let url = exec_ctx.responses_url(); let upstream_request = ctx.enriched_request.to_upstream_request(true)?; let upstream_json = serialize_to_string(&upstream_request).map_err(ExecutorError::JsonError)?; @@ -56,21 +73,31 @@ pub(super) async fn fetch_stream_payload( )); let mut acc = ResponseAccumulator::new(ctx.response_id.clone(), ctx.conversation_id.clone()); let mut hidden_gateway_item_ids = HashSet::new(); - let mut pending_unnamed_function_events = HashMap::>::new(); + let mut pending_unnamed_function_events = HashMap::>::new(); + let mut defer_from_output_index = None; + let mut deferred_events = Vec::new(); while let Some(line_result) = line_stream.next().await { let line = line_result?; - log_upstream_failure(&line, &ctx.response_id); - if let Some(sender) = stream_events { - emit_upstream_stream_event( - &line, - ctx, - registry, - sender, - &mut hidden_gateway_item_ids, - &mut pending_unnamed_function_events, - )?; + if let Some(frame) = acc.process_sse_line(&line) { + log_upstream_failure(&frame, &ctx.response_id); + if let Some((accumulator, sender)) = stream.as_mut() { + let mut emit_ctx = StreamEmitContext { + request: ctx, + registry, + sender, + accumulator, + output_offset, + }; + emit_upstream_stream_event( + frame, + &mut emit_ctx, + &mut hidden_gateway_item_ids, + &mut pending_unnamed_function_events, + &mut defer_from_output_index, + &mut deferred_events, + )?; + } } - acc.process_sse_line(&line); } acc.finish_stream(); let mut payload = acc.finalize( @@ -79,24 +106,18 @@ pub(super) async fn fetch_stream_payload( ctx.original_request.instructions.as_deref(), ); ctx.inject_ids(&mut payload); - Ok(payload) + Ok(StreamPayload { + payload, + deferred_events, + }) } -fn log_upstream_failure(line: &str, gateway_response_id: &str) { - let Some(frame) = normalize_sse_line(line) else { - return; - }; +fn log_upstream_failure(frame: &EventFrame, gateway_response_id: &str) { if frame.event_type != SSEEventType::ResponseFailed { return; } - let Some(data) = line.strip_prefix("data: ") else { - return; - }; - let Ok(event) = deserialize_from_str::(data) else { - return; - }; - let response = &event["response"]; + let response = frame.wire.rest.get("response").unwrap_or(&Value::Null); let error = &response["error"]; let error_code = error.get("code").and_then(Value::as_str).unwrap_or_default(); let error_message = error @@ -120,99 +141,155 @@ fn log_upstream_failure(line: &str, gateway_response_id: &str) { } fn emit_upstream_stream_event( - line: &str, - ctx: &RequestContext, - registry: &ToolRegistry, - sender: &mpsc::UnboundedSender, + frame: EventFrame, + emit_ctx: &mut StreamEmitContext<'_>, hidden_gateway_item_ids: &mut HashSet, - pending_unnamed_function_events: &mut HashMap>, + pending_unnamed_function_events: &mut HashMap>, + defer_from_output_index: &mut Option, + deferred_events: &mut Vec, ) -> ExecutorResult<()> { - let Some(data) = line.strip_prefix("data: ") else { - return Ok(()); - }; - let data = data.trim(); - if data == "[DONE]" { - return Ok(()); - } - - let Some(frame) = normalize_sse_line(line) else { - return Ok(()); - }; - if should_hide_upstream_event(frame.event_type, &frame.payload, registry, hidden_gateway_item_ids) - || is_terminal_response_event(frame.event_type) + defer_after_gateway_call(&frame, emit_ctx.registry, defer_from_output_index); + if should_hide_upstream_event( + frame.event_type, + &frame.payload, + emit_ctx.registry, + hidden_gateway_item_ids, + ) || is_terminal_response_event(frame.event_type) { drop_pending_function_events(&frame.payload, pending_unnamed_function_events); return Ok(()); } - if defer_or_flush_function_event( - line, - &frame.payload, - ctx, - registry, - sender, + let Some(frame) = defer_or_flush_function_event( + frame, + emit_ctx, hidden_gateway_item_ids, pending_unnamed_function_events, - )? { + defer_from_output_index, + deferred_events, + )? + else { return Ok(()); - } + }; - emit_stream_line(data, ctx, registry, sender) + emit_or_defer_stream_frame(frame, emit_ctx, *defer_from_output_index, deferred_events) } -fn emit_stream_line( - data: &str, - ctx: &RequestContext, +pub(super) fn emit_deferred_stream_events( + deferred_events: Vec, + request: &RequestContext, registry: &ToolRegistry, - sender: &mpsc::UnboundedSender, + accumulator: &mut GatewayStreamAccumulator, + sender: &tokio::sync::mpsc::UnboundedSender, + output_offset: usize, ) -> ExecutorResult<()> { - let mut value = serde_json::from_str::(data).map_err(ExecutorError::JsonError)?; - apply_context_response_ids(&mut value, ctx); - registry.restore_stream_event_value(&mut value); - let event_json = serialize_to_string(&value).map_err(ExecutorError::JsonError)?; - sender - .send(format!("data: {event_json}\n\n")) - .map_err(|_| ExecutorError::StreamError("stream receiver closed while emitting upstream event".to_owned())) + let mut emit_ctx = StreamEmitContext { + request, + registry, + sender, + accumulator, + output_offset, + }; + for mut frame in deferred_events { + emit_stream_frame(&mut frame, &mut emit_ctx)?; + } + Ok(()) +} + +fn defer_after_gateway_call(frame: &EventFrame, registry: &ToolRegistry, defer_from_output_index: &mut Option) { + let EventPayload::OutputItemAdded { + item_type: SSEItemType::FunctionCall, + name: Some(name), + .. + } = &frame.payload + else { + return; + }; + if registry.is_gateway_owned_name(name) { + record_first_hidden_gateway_output_index(frame, defer_from_output_index); + } +} + +fn record_first_hidden_gateway_output_index(frame: &EventFrame, defer_from_output_index: &mut Option) { + let Some(output_index) = frame.wire.output_index else { + return; + }; + if defer_from_output_index.is_none_or(|first_hidden_index| output_index < first_hidden_index) { + *defer_from_output_index = Some(output_index); + } +} + +fn should_defer_stream_event(frame: &EventFrame, defer_from_output_index: Option) -> bool { + defer_from_output_index.is_some_and(|first_hidden_index| { + frame + .wire + .output_index + .is_some_and(|output_index| output_index >= first_hidden_index) + }) +} + +fn emit_stream_frame(frame: &mut EventFrame, emit_ctx: &mut StreamEmitContext<'_>) -> ExecutorResult<()> { + apply_context_response_ids(&mut frame.wire, emit_ctx.request); + emit_ctx.registry.restore_stream_event_wire(&mut frame.wire); + if emit_ctx.accumulator.process_event(frame, emit_ctx.output_offset) { + emit_sse_frame(emit_ctx.sender, frame)?; + } + Ok(()) +} + +fn emit_or_defer_stream_frame( + mut frame: EventFrame, + emit_ctx: &mut StreamEmitContext<'_>, + defer_from_output_index: Option, + deferred_events: &mut Vec, +) -> ExecutorResult<()> { + if should_defer_stream_event(&frame, defer_from_output_index) { + deferred_events.push(frame); + return Ok(()); + } + emit_stream_frame(&mut frame, emit_ctx) } fn defer_or_flush_function_event( - line: &str, - payload: &EventPayload, - ctx: &RequestContext, - registry: &ToolRegistry, - sender: &mpsc::UnboundedSender, + frame: EventFrame, + emit_ctx: &mut StreamEmitContext<'_>, hidden_gateway_item_ids: &mut HashSet, - pending_unnamed_function_events: &mut HashMap>, -) -> ExecutorResult { - match payload { + pending_unnamed_function_events: &mut HashMap>, + defer_from_output_index: &mut Option, + deferred_events: &mut Vec, +) -> ExecutorResult> { + match &frame.payload { EventPayload::OutputItemAdded { item_id, item_type, name: None, .. } if *item_type == SSEItemType::FunctionCall => { - pending_unnamed_function_events - .entry(item_id.clone()) - .or_default() - .push(line.to_owned()); - Ok(true) + let item_id = item_id.clone(); + pending_unnamed_function_events.entry(item_id).or_default().push(frame); + Ok(None) } EventPayload::FunctionCallArgsDelta { item_id, .. } if pending_unnamed_function_events.contains_key(item_id) => { - pending_unnamed_function_events - .entry(item_id.clone()) - .or_default() - .push(line.to_owned()); - Ok(true) + let item_id = item_id.clone(); + pending_unnamed_function_events.entry(item_id).or_default().push(frame); + Ok(None) } EventPayload::FunctionCallArgsDone { item_id, name, .. } => { - if registry.is_gateway_owned_name(name) { + if emit_ctx.registry.is_gateway_owned_name(name) { hidden_gateway_item_ids.insert(item_id.clone()); + record_first_hidden_gateway_output_index(&frame, defer_from_output_index); pending_unnamed_function_events.remove(item_id); - return Ok(true); + return Ok(None); } - flush_pending_function_events(item_id, ctx, registry, sender, pending_unnamed_function_events)?; - Ok(false) + flush_pending_function_events( + item_id, + emit_ctx, + pending_unnamed_function_events, + *defer_from_output_index, + deferred_events, + )?; + Ok(Some(frame)) } EventPayload::OutputItemDone { item_id, @@ -223,41 +300,45 @@ fn defer_or_flush_function_event( if item .get("name") .and_then(Value::as_str) - .is_some_and(|name| registry.is_gateway_owned_name(name)) + .is_some_and(|name| emit_ctx.registry.is_gateway_owned_name(name)) { hidden_gateway_item_ids.insert(item_id.clone()); + record_first_hidden_gateway_output_index(&frame, defer_from_output_index); pending_unnamed_function_events.remove(item_id); - return Ok(true); + return Ok(None); } - flush_pending_function_events(item_id, ctx, registry, sender, pending_unnamed_function_events)?; - Ok(false) + flush_pending_function_events( + item_id, + emit_ctx, + pending_unnamed_function_events, + *defer_from_output_index, + deferred_events, + )?; + Ok(Some(frame)) } - _ => Ok(false), + _ => Ok(Some(frame)), } } fn flush_pending_function_events( item_id: &str, - ctx: &RequestContext, - registry: &ToolRegistry, - sender: &mpsc::UnboundedSender, - pending_unnamed_function_events: &mut HashMap>, + emit_ctx: &mut StreamEmitContext<'_>, + pending_unnamed_function_events: &mut HashMap>, + defer_from_output_index: Option, + deferred_events: &mut Vec, ) -> ExecutorResult<()> { - let Some(lines) = pending_unnamed_function_events.remove(item_id) else { + let Some(frames) = pending_unnamed_function_events.remove(item_id) else { return Ok(()); }; - for line in lines { - let Some(data) = line.strip_prefix("data: ") else { - continue; - }; - emit_stream_line(data.trim(), ctx, registry, sender)?; + for frame in frames { + emit_or_defer_stream_frame(frame, emit_ctx, defer_from_output_index, deferred_events)?; } Ok(()) } fn drop_pending_function_events( payload: &EventPayload, - pending_unnamed_function_events: &mut HashMap>, + pending_unnamed_function_events: &mut HashMap>, ) { match payload { EventPayload::OutputItemDone { item_id, .. } @@ -317,8 +398,8 @@ fn is_terminal_response_event(event_type: SSEEventType) -> bool { ) } -fn apply_context_response_ids(value: &mut Value, ctx: &RequestContext) { - let Some(response) = value.get_mut("response").and_then(Value::as_object_mut) else { +fn apply_context_response_ids(wire: &mut WireEvent, ctx: &RequestContext) { + let Some(response) = wire.rest.get_mut("response").and_then(Value::as_object_mut) else { return; }; response.insert("id".to_owned(), Value::String(ctx.response_id.clone())); diff --git a/crates/agentic-server-core/src/tool/codex.rs b/crates/agentic-server-core/src/tool/codex.rs index 099ee6b..beb6258 100644 --- a/crates/agentic-server-core/src/tool/codex.rs +++ b/crates/agentic-server-core/src/tool/codex.rs @@ -1,7 +1,8 @@ use std::collections::HashMap; -use serde_json::Value; +use serde_json::{Map, Value}; +use crate::events::WireEvent; use crate::types::io::{FunctionTool, FunctionToolCall, OutputItem, ToolChoice}; use crate::types::tools::{CodexNamespaceMember, CodexNamespaceToolParam, NonEmptyToolName, ResponsesTool}; @@ -302,6 +303,14 @@ impl CodexNamespaceHandler { }; restore_response_value_with_map(value, map) } + + #[must_use] + pub fn restore_response_wire(&self, wire: &mut WireEvent, map: Option<&NamespaceMap>) -> bool { + let Some(map) = map else { + return false; + }; + restore_response_map_with_map(&mut wire.rest, map) + } } impl ToolHandler for CodexNamespaceHandler { @@ -521,6 +530,24 @@ fn restore_call_value_with_map(value: &mut Value, map: &NamespaceMap) -> bool { true } +fn restore_response_map_with_map(object: &mut Map, map: &NamespaceMap) -> bool { + let mut changed = false; + if let Some(item) = object.get_mut("item") { + changed |= restore_call_value_with_map(item, map); + } + for key in ["response", "payload"] { + if let Some(nested) = object.get_mut(key) { + changed |= restore_response_value_with_map(nested, map); + } + } + if let Some(Value::Array(items)) = object.get_mut("output") { + for item in items { + changed |= restore_call_value_with_map(item, map); + } + } + changed +} + #[cfg(test)] mod tests { use super::*; diff --git a/crates/agentic-server-core/src/tool/registry.rs b/crates/agentic-server-core/src/tool/registry.rs index 17c0e8e..1f506f4 100644 --- a/crates/agentic-server-core/src/tool/registry.rs +++ b/crates/agentic-server-core/src/tool/registry.rs @@ -10,6 +10,7 @@ use super::function::insert_function_entry; use super::mcp::{insert_mcp_entry, maybe_mcp_function}; use super::web_search::insert_web_search_entry; use super::{CodexNamespaceHandler, GatewayExecutor, NamespaceMap, ToolError, ToolOutput}; +use crate::events::WireEvent; use crate::types::io::OutputItem; use crate::types::io::output::FunctionToolCall; use crate::types::tools::{CodeInterpreterToolParam, FileSearchToolParam, ResponsesTool}; @@ -130,9 +131,8 @@ fn insert_code_interpreter_entry( #[derive(Debug, Default)] pub struct ToolRegistry { entries: HashMap, - /// Built once from the declared tools, so `restore_final_payload_output` - /// and `restore_stream_event_value` — the latter called once per SSE line - /// during streaming — don't rebuild it on every call. + /// Built once from the declared tools, so final payload and streaming event + /// restoration don't rebuild it on every call. namespace_map: Option, } @@ -209,8 +209,8 @@ impl ToolRegistry { CodexNamespaceHandler.restore_output_items(output, self.namespace_map.as_ref()); } - pub fn restore_stream_event_value(&self, value: &mut Value) -> bool { - CodexNamespaceHandler.restore_response_value(value, self.namespace_map.as_ref()) + pub fn restore_stream_event_wire(&self, wire: &mut WireEvent) -> bool { + CodexNamespaceHandler.restore_response_wire(wire, self.namespace_map.as_ref()) } /// Returns the subset of `calls` whose names map to gateway-owned tools. diff --git a/crates/agentic-server-core/src/types/request_response.rs b/crates/agentic-server-core/src/types/request_response.rs index 3648bd2..56e43fe 100644 --- a/crates/agentic-server-core/src/types/request_response.rs +++ b/crates/agentic-server-core/src/types/request_response.rs @@ -246,7 +246,7 @@ impl ResponsePayload { format!("data: {json_str}\n\n") } - fn terminal_event_type(&self) -> &'static str { + pub(crate) fn terminal_event_type(&self) -> &'static str { match self.status.as_str() { "incomplete" => "response.incomplete", "failed" | "error" => "response.failed", diff --git a/crates/agentic-server-core/tests/event_normalizer_test.rs b/crates/agentic-server-core/tests/event_normalizer_test.rs index 4822a8e..1304e69 100644 --- a/crates/agentic-server-core/tests/event_normalizer_test.rs +++ b/crates/agentic-server-core/tests/event_normalizer_test.rs @@ -8,7 +8,7 @@ fn test_text_delta() { let line = r#"data: {"type":"response.output_text.delta","delta":"hello","item_id":"msg_1","output_index":0,"content_index":0,"sequence_number":4}"#; let frame = normalize_sse_line(line).unwrap(); assert_eq!(frame.event_type, SSEEventType::OutputTextDelta); - assert_eq!(frame.sequence_number, Some(4)); + assert_eq!(frame.sequence_number(), Some(4)); if let EventPayload::TextDelta { delta, item_id, @@ -30,7 +30,7 @@ fn test_function_call_args_delta() { let line = r#"data: {"type":"response.function_call_arguments.delta","delta":"{\"city\":","call_id":"call_abc","item_id":"fc_1","output_index":0,"sequence_number":7}"#; let frame = normalize_sse_line(line).unwrap(); assert_eq!(frame.event_type, SSEEventType::FunctionCallArgumentsDelta); - assert_eq!(frame.sequence_number, Some(7)); + assert_eq!(frame.sequence_number(), Some(7)); if let EventPayload::FunctionCallArgsDelta { delta, call_id, @@ -135,6 +135,27 @@ fn test_unknown_event_type() { assert!(matches!(frame.payload, EventPayload::Raw(_))); } +#[test] +fn test_typeless_json_event_is_preserved() { + let frame = normalize_sse_line(r#"data: {"foo":1}"#).unwrap(); + + assert_eq!(frame.event_type, SSEEventType::Other); + assert_eq!(serde_json::to_value(frame.wire).unwrap(), serde_json::json!({"foo": 1})); +} + +#[test] +fn test_wire_event_preserves_unknown_fields() { + let line = r#"data: {"type":"response.output_text.delta","sequence_number":4,"output_index":2,"item_id":"msg_1","content_index":0,"delta":"hello","provider_extra":{"nested":true},"future_array":[1,2]}"#; + let frame = normalize_sse_line(line).unwrap(); + let wire = serde_json::to_value(&frame.wire).unwrap(); + + assert_eq!(wire["type"], "response.output_text.delta"); + assert_eq!(wire["sequence_number"], 4); + assert_eq!(wire["output_index"], 2); + assert_eq!(wire["provider_extra"]["nested"], true); + assert_eq!(wire["future_array"], serde_json::json!([1, 2])); +} + #[test] fn test_malformed_json_returns_none() { assert!(normalize_sse_line("data: {not valid json}").is_none()); @@ -146,7 +167,7 @@ fn test_response_created() { let line = r#"data: {"type":"response.created","response":{"id":"resp_abc","status":"in_progress","usage":null},"sequence_number":0}"#; let frame = normalize_sse_line(line).unwrap(); assert_eq!(frame.event_type, SSEEventType::ResponseCreated); - assert_eq!(frame.sequence_number, Some(0)); + assert_eq!(frame.sequence_number(), Some(0)); if let EventPayload::Response { id, status, .. } = &frame.payload { assert_eq!(id, "resp_abc"); assert_eq!(status, "in_progress"); @@ -213,7 +234,7 @@ fn test_no_sequence_number() { let line = r#"data: {"type":"response.output_text.delta","delta":"x","item_id":"m","output_index":0,"content_index":0}"#; let frame = normalize_sse_line(line).unwrap(); - assert_eq!(frame.sequence_number, None); + assert_eq!(frame.sequence_number(), None); } #[test] @@ -442,7 +463,7 @@ fn test_sequence_numbers_increasing() { let mut last_seq: Option = None; for line in SIMULATED_SSE { if let Some(frame) = normalize_sse_line(line) { - if let Some(seq) = frame.sequence_number { + if let Some(seq) = frame.sequence_number() { if let Some(prev) = last_seq { assert!(seq > prev, "sequence {seq} should be > {prev}"); } diff --git a/crates/agentic-server-core/tests/tool_normalization_test.rs b/crates/agentic-server-core/tests/tool_normalization_test.rs index 213811f..341aae8 100644 --- a/crates/agentic-server-core/tests/tool_normalization_test.rs +++ b/crates/agentic-server-core/tests/tool_normalization_test.rs @@ -6,6 +6,7 @@ use serde::Deserialize; use serde_json::Value; +use agentic_core::events::WireEvent; use agentic_core::executor::RequestContext; use agentic_core::tool::{ CodexNamespaceHandler, GatewayExecutors, ToolRegistry, ToolType, model_visible_namespace_member_name, @@ -416,6 +417,41 @@ fn codex_namespace_cassettes_flatten_to_safe_upstream_function_name() { } } +#[tokio::test] +async fn tool_registry_restores_wire_event_namespace_losslessly() { + let tools: Vec = serde_json::from_value(serde_json::json!([ + { + "type": "namespace", + "name": "mcp__agentic_fixture", + "tools": [{"type": "function", "name": "add_numbers"}] + } + ])) + .unwrap(); + let registry = ToolRegistry::build_with_handlers(&tools, &GatewayExecutors::default()) + .await + .expect("valid registry"); + let mut wire = WireEvent::new("response.output_item.done"); + wire.output_index = Some(0); + wire.rest.insert( + "item".to_owned(), + serde_json::json!({ + "type": "function_call", + "name": "agentic_ns__mcp__agentic_fixture__add_numbers", + "call_id": "call_1", + "arguments": "{\"numbers\":[8,0]}", + "provider_extra": {"kept": true} + }), + ); + + assert!(registry.restore_stream_event_wire(&mut wire)); + + let item = &wire.rest["item"]; + assert_eq!(item["namespace"], "mcp__agentic_fixture"); + assert_eq!(item["name"], "add_numbers"); + assert_eq!(item["arguments"], "{\"numbers\":[8,0]}"); + assert_eq!(item["provider_extra"]["kept"], true); +} + #[test] fn codex_direct_vllm_flat_namespace_cassette_is_plain_function_tool() { let filename = "codex-direct-vllm-http-flat-namespace-tool-Qwen-Qwen3.6-35B-A3B-streaming.yaml"; diff --git a/crates/agentic-server-core/tests/web_search_tool_test.rs b/crates/agentic-server-core/tests/web_search_tool_test.rs index e57b8c3..bd5bcb7 100644 --- a/crates/agentic-server-core/tests/web_search_tool_test.rs +++ b/crates/agentic-server-core/tests/web_search_tool_test.rs @@ -359,6 +359,11 @@ fn web_search_function_call_sse_response() -> support::MockResponse { "name": "web_search", "arguments": "{\"query\":\"rust async\",\"count\":2}" }), + serde_json::json!({ + "type": "provider.gateway_metadata", + "output_index": 0, + "metadata": {"trace_id": "trace_search"} + }), serde_json::json!({ "type": "response.completed", "response": {"id": "resp_tool_call", "status": "completed", "usage": null} @@ -436,6 +441,135 @@ fn text_sse_response(text: &str) -> support::MockResponse { ]) } +fn text_sse_response_with_output_index(text: &str, output_index: u32) -> support::MockResponse { + sse_response([ + serde_json::json!({ + "type": "response.created", + "response": {"id": "resp_final", "status": "in_progress", "usage": null} + }), + serde_json::json!({ + "type": "response.in_progress", + "response": {"id": "resp_final", "status": "in_progress", "usage": null} + }), + serde_json::json!({ + "type": "response.output_item.added", + "output_index": output_index, + "item": { + "id": "msg_final", + "type": "message", + "role": "assistant", + "status": "in_progress", + "content": [] + } + }), + serde_json::json!({ + "type": "response.output_text.delta", + "item_id": "msg_final", + "output_index": output_index, + "content_index": 0, + "delta": text + }), + serde_json::json!({ + "type": "response.output_item.done", + "output_index": output_index, + "item": { + "id": "msg_final", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": text, "annotations": []}] + } + }), + serde_json::json!({ + "type": "response.completed", + "response": {"id": "resp_final", "status": "completed", "usage": null} + }), + ]) +} + +fn two_messages_then_web_search_sse_response() -> support::MockResponse { + sse_response([ + serde_json::json!({ + "type": "response.created", + "response": {"id": "resp_mid", "status": "in_progress", "usage": null} + }), + serde_json::json!({ + "type": "response.in_progress", + "response": {"id": "resp_mid", "status": "in_progress", "usage": null} + }), + serde_json::json!({ + "type": "response.output_item.added", + "output_index": 0, + "item": {"id": "msg_mid_0", "type": "message", "role": "assistant", "status": "in_progress", "content": []} + }), + serde_json::json!({ + "type": "response.output_text.delta", + "item_id": "msg_mid_0", + "output_index": 0, + "content_index": 0, + "delta": "First result." + }), + serde_json::json!({ + "type": "response.output_item.done", + "output_index": 0, + "item": { + "id": "msg_mid_0", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "First result.", "annotations": []}] + } + }), + serde_json::json!({ + "type": "response.output_item.added", + "output_index": 1, + "item": {"id": "msg_mid_1", "type": "message", "role": "assistant", "status": "in_progress", "content": []} + }), + serde_json::json!({ + "type": "response.output_text.delta", + "item_id": "msg_mid_1", + "output_index": 1, + "content_index": 0, + "delta": "Second result." + }), + serde_json::json!({ + "type": "response.output_item.done", + "output_index": 1, + "item": { + "id": "msg_mid_1", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "Second result.", "annotations": []}] + } + }), + serde_json::json!({ + "type": "response.output_item.added", + "output_index": 2, + "item": { + "id": "fc_search_mid", + "type": "function_call", + "call_id": "call_search_mid", + "name": "web_search", + "arguments": "", + "status": "in_progress" + } + }), + serde_json::json!({ + "type": "response.function_call_arguments.done", + "item_id": "fc_search_mid", + "output_index": 2, + "call_id": "call_search_mid", + "name": "web_search", + "arguments": "{\"query\":\"tokio streams\",\"count\":2}" + }), + serde_json::json!({ + "type": "response.completed", + "response": {"id": "resp_mid", "status": "completed", "usage": null} + }), + ]) +} + fn mixed_web_search_and_client_function_response() -> support::MockResponse { support::MockResponse::Json( serde_json::json!({ @@ -473,6 +607,90 @@ fn mixed_web_search_and_client_function_response() -> support::MockResponse { ) } +fn mixed_web_search_and_client_function_sse_response() -> support::MockResponse { + sse_response([ + serde_json::json!({ + "type": "response.created", + "response": {"id": "resp_mixed_tool_call", "status": "in_progress", "usage": null} + }), + serde_json::json!({ + "type": "response.output_item.added", + "output_index": 0, + "item": { + "id": "fc_search", + "type": "function_call", + "call_id": "call_search", + "name": "web_search", + "arguments": "", + "status": "in_progress" + } + }), + serde_json::json!({ + "type": "response.function_call_arguments.done", + "item_id": "fc_search", + "output_index": 0, + "call_id": "call_search", + "name": "web_search", + "arguments": "{\"query\":\"rust async\",\"count\":2}" + }), + serde_json::json!({ + "type": "response.output_item.added", + "output_index": 1, + "item": { + "id": "fc_weather", + "type": "function_call", + "call_id": "call_weather", + "arguments": "", + "status": "in_progress" + } + }), + serde_json::json!({ + "type": "response.function_call_arguments.done", + "item_id": "fc_weather", + "output_index": 1, + "call_id": "call_weather", + "name": "get_weather", + "arguments": "{\"city\":\"San Francisco\"}" + }), + serde_json::json!({ + "type": "response.output_item.done", + "output_index": 1, + "item": { + "id": "fc_weather", + "type": "function_call", + "call_id": "call_weather", + "name": "get_weather", + "arguments": "{\"city\":\"San Francisco\"}", + "status": "completed" + } + }), + serde_json::json!({ + "type": "response.output_item.added", + "output_index": 2, + "item": { + "id": "fc_search_second", + "type": "function_call", + "call_id": "call_search_second", + "name": "web_search", + "arguments": "", + "status": "in_progress" + } + }), + serde_json::json!({ + "type": "response.function_call_arguments.done", + "item_id": "fc_search_second", + "output_index": 2, + "call_id": "call_search_second", + "name": "web_search", + "arguments": "{\"query\":\"tokio streams\",\"count\":2}" + }), + serde_json::json!({ + "type": "response.completed", + "response": {"id": "resp_mixed_tool_call", "status": "completed", "usage": null} + }), + ]) +} + fn text_response_with_usage(text: &str, input_tokens: i64, output_tokens: i64) -> support::MockResponse { let id_suffix = text.replace(' ', "_"); support::MockResponse::Json( @@ -938,6 +1156,131 @@ async fn stream_emits_web_search_lifecycle_events_before_final_payload() { assert!(output.iter().any(|item| item["type"] == "message")); } +fn assert_single_logical_lifecycle(json_events: &[serde_json::Value]) { + let event_types: Vec<&str> = json_events.iter().filter_map(|event| event["type"].as_str()).collect(); + assert_eq!( + event_types + .iter() + .filter(|event_type| **event_type == "response.created") + .count(), + 1, + "multi-round stream should expose one logical response.created: {event_types:?}" + ); + assert_eq!( + event_types + .iter() + .filter(|event_type| **event_type == "response.in_progress") + .count(), + 1, + "multi-round stream should expose one logical response.in_progress: {event_types:?}" + ); +} + +fn assert_contiguous_sequence_numbers(json_events: &[serde_json::Value], message: &str) { + let sequence_numbers: Vec = json_events + .iter() + .map(|event| { + event["sequence_number"] + .as_u64() + .unwrap_or_else(|| panic!("event missing sequence_number: {event}")) + }) + .collect(); + assert_eq!( + sequence_numbers, + (0..u64::try_from(sequence_numbers.len()).unwrap()).collect::>(), + "{message}" + ); +} + +fn assert_output_event_indices_in_order(json_events: &[serde_json::Value], expected_events: &[(&str, u64)]) { + let output_events: Vec<(&str, u64)> = json_events + .iter() + .filter_map(|event| Some((event["type"].as_str()?, event["output_index"].as_u64()?))) + .collect(); + assert_eq!(output_events, expected_events); +} + +#[tokio::test] +async fn multi_round_stream_has_single_lifecycle_and_monotonic_public_sequence() { + let (you_url, mut captured_you, _you_handle) = spawn_mock_you().await; + let llm = support::MockServer::start_deque(vec![ + web_search_function_call_sse_response(), + two_messages_then_web_search_sse_response(), + text_sse_response_with_output_index("Use async carefully.", 0), + ]) + .await; + let exec_ctx = build_exec_ctx(llm.url(), you_url).await; + let web_search: ResponsesTool = serde_json::from_value(serde_json::json!({"type": "web_search_preview"})).unwrap(); + let payload = RequestPayload { + model: "test-model".to_owned(), + input: ResponsesInput::Text("look up rust async".to_owned()), + instructions: None, + previous_response_id: None, + conversation_id: None, + tools: Some(vec![web_search]), + tool_choice: None, + stream: true, + store: true, + include: None, + temperature: None, + top_p: None, + max_output_tokens: Some(1024), + truncation: None, + cache_salt: None, + metadata: None, + parallel_tool_calls: None, + }; + + let result = ExecuteRequest::new(payload, Arc::clone(&exec_ctx)).run().await.unwrap(); + let Either::Right(stream) = result else { + panic!("expected streaming response"); + }; + let chunks: Vec = stream.collect().await; + captured_you.recv().await.expect("mock You.com should receive request"); + captured_you + .recv() + .await + .expect("mock You.com should receive second request"); + + let json_events: Vec = chunks + .iter() + .filter_map(|chunk| { + let data = chunk.trim_end_matches('\n').strip_prefix("data: ")?; + (data != "[DONE]").then(|| serde_json::from_str(data).ok())? + }) + .collect(); + assert_single_logical_lifecycle(&json_events); + assert_contiguous_sequence_numbers( + &json_events, + "public sequence_number must be contiguous across upstream, synthetic, and terminal frames", + ); + assert_output_event_indices_in_order( + &json_events, + &[ + ("response.output_item.added", 0), + ("response.web_search_call.in_progress", 0), + ("response.web_search_call.searching", 0), + ("response.web_search_call.completed", 0), + ("response.output_item.done", 0), + ("provider.gateway_metadata", 0), + ("response.output_item.added", 1), + ("response.output_text.delta", 1), + ("response.output_item.done", 1), + ("response.output_item.added", 2), + ("response.output_text.delta", 2), + ("response.output_item.done", 2), + ("response.output_item.added", 3), + ("response.web_search_call.in_progress", 3), + ("response.web_search_call.searching", 3), + ("response.web_search_call.completed", 3), + ("response.output_item.done", 3), + ("response.output_item.added", 4), + ("response.output_text.delta", 4), + ("response.output_item.done", 4), + ], + ); +} + #[tokio::test] async fn stream_hides_web_search_function_events_when_name_arrives_on_done() { let (you_url, mut captured_you, _you_handle) = spawn_mock_you().await; @@ -1005,6 +1348,99 @@ async fn stream_hides_web_search_function_events_when_name_arrives_on_done() { assert!(output.iter().any(|item| item["type"] == "message")); } +#[tokio::test] +async fn stream_orders_gateway_lifecycle_before_later_client_function_events() { + let (you_url, mut captured_you, _you_handle) = spawn_mock_you_waiting_for_two_searches().await; + let llm = support::MockServer::start_deque(vec![mixed_web_search_and_client_function_sse_response()]).await; + let exec_ctx = build_exec_ctx(llm.url(), you_url).await; + let web_search: ResponsesTool = serde_json::from_value(serde_json::json!({"type": "web_search_preview"})).unwrap(); + let client_function: ResponsesTool = serde_json::from_value(serde_json::json!({ + "type": "function", + "name": "get_weather", + "parameters": { + "type": "object", + "properties": {"city": {"type": "string"}} + } + })) + .unwrap(); + let payload = RequestPayload { + model: "test-model".to_owned(), + input: ResponsesInput::Text("look up rust async and weather".to_owned()), + instructions: None, + previous_response_id: None, + conversation_id: None, + tools: Some(vec![web_search, client_function]), + tool_choice: None, + stream: true, + store: true, + include: None, + temperature: None, + top_p: None, + max_output_tokens: Some(1024), + truncation: None, + metadata: None, + parallel_tool_calls: None, + cache_salt: None, + }; + + let result = ExecuteRequest::new(payload, exec_ctx).run().await.unwrap(); + let Either::Right(stream) = result else { + panic!("expected streaming response"); + }; + let chunks: Vec = tokio::time::timeout(Duration::from_secs(2), stream.collect()) + .await + .expect("interleaved gateway calls should execute concurrently"); + captured_you + .recv() + .await + .expect("mock You.com should receive first request"); + captured_you + .recv() + .await + .expect("mock You.com should receive second request"); + + let json_events: Vec = chunks + .iter() + .filter_map(|chunk| { + let data = chunk.trim_end_matches('\n').strip_prefix("data: ")?; + (data != "[DONE]").then(|| serde_json::from_str(data).ok())? + }) + .collect(); + assert_contiguous_sequence_numbers( + &json_events, + "mixed-call stream must retain contiguous sequence numbers", + ); + assert!( + json_events + .iter() + .any(|event| { event["item"]["type"] == "function_call" && event["item"]["name"] == "get_weather" }) + ); + assert!( + !json_events + .iter() + .any(|event| { event["item"]["type"] == "function_call" && event["item"]["name"] == "web_search" }) + ); + + assert_output_event_indices_in_order( + &json_events, + &[ + ("response.output_item.added", 0), + ("response.web_search_call.in_progress", 0), + ("response.web_search_call.searching", 0), + ("response.web_search_call.completed", 0), + ("response.output_item.done", 0), + ("response.output_item.added", 1), + ("response.function_call_arguments.done", 1), + ("response.output_item.done", 1), + ("response.output_item.added", 2), + ("response.web_search_call.in_progress", 2), + ("response.web_search_call.searching", 2), + ("response.web_search_call.completed", 2), + ("response.output_item.done", 2), + ], + ); +} + #[tokio::test] async fn execute_runs_multiple_web_search_calls_concurrently() { let (you_url, mut captured_you, _you_handle) = spawn_mock_you_waiting_for_two_searches().await; @@ -1488,13 +1924,29 @@ async fn stream_returns_incomplete_after_max_gateway_tool_rounds() { }; let chunks: Vec = stream.collect().await; - let final_event = chunks + let json_events: Vec = chunks .iter() .filter_map(|chunk| { let data = chunk.trim_end_matches('\n').strip_prefix("data: ")?; (data != "[DONE]").then(|| serde_json::from_str::(data).ok())? }) - .next_back() + .collect(); + let sequence_numbers: Vec = json_events + .iter() + .map(|event| { + event["sequence_number"] + .as_u64() + .unwrap_or_else(|| panic!("event missing sequence_number: {event}")) + }) + .collect(); + assert_eq!( + sequence_numbers, + (0..u64::try_from(sequence_numbers.len()).unwrap()).collect::>(), + "incomplete terminal stream should keep sequence_number contiguous" + ); + + let final_event = json_events + .last() .expect("stream should carry a final response payload"); // The terminal SSE event wraps the payload: {"type":"response.incomplete","response":{...}} let response = &final_event["response"];