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..33c4e30 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`]. @@ -15,21 +15,18 @@ pub fn normalize_sse_line(line: &str) -> Option { return None; } - let json: Value = deserialize_from_str_opt(data_str)?; + let wire: WireEvent = deserialize_from_str_opt(data_str)?; + let json = wire.to_value(); - 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); + let event_type = classify_event_type(&wire.event_type); let payload = extract_payload(event_type, &json); Some(EventFrame { event_type, payload, - sequence_number, + sequence_number: wire.sequence_number, + wire, }) } diff --git a/crates/agentic-server-core/src/events/types.rs b/crates/agentic-server-core/src/events/types.rs index 91a1e47..6b483d2 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,77 @@ pub enum SSEEventType { Other, } +impl SSEEventType { + #[must_use] + pub fn as_str(self) -> &'static str { + match self { + Self::ResponseCreated => "response.created", + Self::ResponseInProgress => "response.in_progress", + Self::ResponseCompleted => "response.completed", + Self::ResponseFailed => "response.failed", + Self::ResponseIncomplete => "response.incomplete", + Self::OutputItemAdded => "response.output_item.added", + Self::OutputItemDone => "response.output_item.done", + Self::OutputTextDelta => "response.output_text.delta", + Self::OutputTextDone => "response.output_text.done", + Self::ContentPartAdded => "response.content_part.added", + Self::ContentPartDone => "response.content_part.done", + Self::FunctionCallArgumentsDelta => "response.function_call_arguments.delta", + Self::FunctionCallArgumentsDone => "response.function_call_arguments.done", + Self::CustomToolCallInputDelta => "response.custom_tool_call_input.delta", + Self::CustomToolCallInputDone => "response.custom_tool_call_input.done", + Self::ReasoningTextDelta => "response.reasoning_text.delta", + Self::ReasoningTextDone => "response.reasoning_text.done", + Self::ReasoningPartAdded => "response.reasoning_part.added", + Self::ReasoningPartDone => "response.reasoning_part.done", + Self::ReasoningSummaryTextDelta => "response.reasoning_summary_text.delta", + Self::ReasoningSummaryTextDone => "response.reasoning_summary_text.done", + Self::FileSearchCallSearching => "response.file_search_call.searching", + Self::FileSearchCallCompleted => "response.file_search_call.completed", + Self::WebSearchCallInProgress => "response.web_search_call.in_progress", + Self::WebSearchCallSearching => "response.web_search_call.searching", + Self::WebSearchCallCompleted => "response.web_search_call.completed", + Self::McpToolCallInProgress => "response.mcp_tool_call.in_progress", + Self::McpToolCallCompleted => "response.mcp_tool_call.completed", + Self::Other => "unknown", + } + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct WireEvent { + #[serde(rename = "type")] + pub event_type: String, + #[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: event_type.into(), + sequence_number: None, + output_index: None, + rest: Map::new(), + } + } + + #[must_use] + pub fn into_value(self) -> Value { + serde_json::to_value(self).unwrap_or_else(|_| Value::Object(Map::new())) + } + + #[must_use] + pub fn to_value(&self) -> Value { + serde_json::to_value(self).unwrap_or_else(|_| Value::Object(Map::new())) + } +} + /// Typed payload extracted from an SSE event's JSON data. #[derive(Debug, Clone)] #[non_exhaustive] @@ -206,4 +278,19 @@ 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, mut wire: WireEvent) -> Self { + event_type.as_str().clone_into(&mut wire.event_type); + let payload = EventPayload::Raw(wire.to_value()); + Self { + event_type, + payload, + sequence_number: wire.sequence_number, + wire, + } + } } diff --git a/crates/agentic-server-core/src/executor/accumulator.rs b/crates/agentic-server-core/src/executor/accumulator.rs index 2e4466c..a409f8a 100644 --- a/crates/agentic-server-core/src/executor/accumulator.rs +++ b/crates/agentic-server-core/src/executor/accumulator.rs @@ -224,30 +224,29 @@ 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); } } - 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) { @@ -263,6 +262,7 @@ impl ResponseAccumulator { /// frame (e.g. [`StreamTee`](future)) can call this directly without /// re-parsing from a raw line. pub(crate) fn process_event(&mut self, frame: &EventFrame) { + self.capture_terminal_details_if_needed(frame); match (&frame.event_type, &frame.payload) { (SSEEventType::ResponseCreated, EventPayload::Response { id, .. }) if !id.is_empty() => { self.response_id.clone_from(id); @@ -431,6 +431,7 @@ impl ResponseAccumulator { #[cfg(test)] mod tests { use super::*; + use crate::events::WireEvent; #[test] fn test_accumulator_new() { @@ -526,6 +527,7 @@ mod tests { usage: None, }, sequence_number: Some(0), + wire: WireEvent::new("test"), }; acc.process_event(&frame); assert_eq!(acc.response_id, "resp_new"); @@ -542,6 +544,7 @@ mod tests { usage: None, }, sequence_number: Some(0), + wire: WireEvent::new("test"), }; acc.process_event(&frame); assert_eq!(acc.response_id, "resp_keep"); @@ -562,6 +565,7 @@ mod tests { call_id: None, }, sequence_number: Some(1), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -573,6 +577,7 @@ mod tests { content_index: 0, }, sequence_number: Some(2), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { event_type: SSEEventType::OutputTextDelta, @@ -583,6 +588,7 @@ mod tests { content_index: 0, }, sequence_number: Some(3), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -593,6 +599,7 @@ mod tests { usage: None, }, sequence_number: Some(4), + wire: WireEvent::new("test"), }); assert_eq!(acc.status, ResponseStatus::Completed); @@ -651,6 +658,7 @@ mod tests { }), }, sequence_number: Some(9), + wire: WireEvent::new("test"), }; acc.process_event(&frame); assert_eq!(acc.status, ResponseStatus::Completed); @@ -661,6 +669,14 @@ mod tests { #[test] fn test_process_event_failed_sets_error_status() { let mut acc = ResponseAccumulator::new("resp_1".into(), None); + let mut wire = WireEvent::new("response.failed"); + wire.rest.insert( + "response".to_owned(), + serde_json::json!({ + "error": {"code": "tool_catalog_too_large"}, + "incomplete_details": {"reason": "upstream_error"} + }), + ); acc.process_event(&EventFrame { event_type: SSEEventType::ResponseFailed, payload: EventPayload::Response { @@ -669,8 +685,20 @@ mod tests { usage: None, }, sequence_number: Some(4), + wire, }); assert_eq!(acc.status, ResponseStatus::Error); + assert_eq!( + acc.error + .as_ref() + .and_then(|error| error.get("code")) + .and_then(serde_json::Value::as_str), + Some("tool_catalog_too_large") + ); + assert_eq!( + acc.incomplete_details.and_then(|details| details.reason), + Some("upstream_error".to_owned()) + ); } #[test] @@ -684,6 +712,7 @@ mod tests { usage: None, }, sequence_number: Some(4), + wire: WireEvent::new("test"), }); assert_eq!(acc.status, ResponseStatus::Incomplete); } @@ -695,6 +724,7 @@ mod tests { 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"); @@ -815,6 +845,7 @@ mod tests { call_id: Some("call_abc".into()), }, sequence_number: Some(1), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -826,6 +857,7 @@ mod tests { output_index: 0, }, sequence_number: Some(2), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -837,6 +869,7 @@ mod tests { output_index: 0, }, sequence_number: Some(3), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -849,6 +882,7 @@ mod tests { output_index: 0, }, sequence_number: Some(4), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -859,6 +893,7 @@ mod tests { usage: None, }, sequence_number: Some(5), + wire: WireEvent::new("test"), }); assert_eq!(acc.status, ResponseStatus::Completed); @@ -890,6 +925,7 @@ mod tests { call_id: Some("call_1".into()), }, sequence_number: Some(1), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -901,6 +937,7 @@ mod tests { output_index: 0, }, sequence_number: Some(2), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -913,6 +950,7 @@ mod tests { output_index: 0, }, sequence_number: Some(3), + wire: WireEvent::new("test"), }); acc.finalize_all(); @@ -939,6 +977,7 @@ mod tests { call_id: Some("call_1".into()), }, sequence_number: Some(1), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { event_type: SSEEventType::FunctionCallArgumentsDone, @@ -950,6 +989,7 @@ mod tests { output_index: 0, }, sequence_number: Some(2), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -963,6 +1003,7 @@ mod tests { call_id: Some("call_2".into()), }, sequence_number: Some(3), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { event_type: SSEEventType::FunctionCallArgumentsDone, @@ -974,6 +1015,7 @@ mod tests { output_index: 1, }, sequence_number: Some(4), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -984,6 +1026,7 @@ mod tests { usage: None, }, sequence_number: Some(5), + wire: WireEvent::new("test"), }); assert_eq!(acc.output.len(), 2); @@ -1006,6 +1049,7 @@ mod tests { call_id: None, }, sequence_number: Some(1), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { event_type: SSEEventType::OutputTextDelta, @@ -1016,6 +1060,7 @@ mod tests { content_index: 0, }, sequence_number: Some(2), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -1029,6 +1074,7 @@ mod tests { call_id: Some("call_x".into()), }, sequence_number: Some(3), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { event_type: SSEEventType::FunctionCallArgumentsDone, @@ -1040,6 +1086,7 @@ mod tests { output_index: 1, }, sequence_number: Some(4), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -1050,6 +1097,7 @@ mod tests { usage: None, }, sequence_number: Some(5), + wire: WireEvent::new("test"), }); assert_eq!(acc.output.len(), 2); @@ -1072,6 +1120,7 @@ mod tests { call_id: Some("old_call".into()), }, sequence_number: Some(1), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -1084,6 +1133,7 @@ mod tests { output_index: 0, }, sequence_number: Some(2), + wire: WireEvent::new("test"), }); acc.finalize_all(); @@ -1110,6 +1160,7 @@ mod tests { call_id: Some("c1".into()), }, sequence_number: Some(1), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -1122,6 +1173,7 @@ mod tests { output_index: 0, }, sequence_number: Some(2), + wire: WireEvent::new("test"), }); acc.finalize_all(); @@ -1146,6 +1198,7 @@ mod tests { output_index: 0, }, sequence_number: Some(1), + wire: WireEvent::new("test"), }); assert!(acc.output.is_empty()); @@ -1167,6 +1220,7 @@ mod tests { call_id: Some("c1".into()), }, sequence_number: Some(1), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { event_type: SSEEventType::FunctionCallArgumentsDelta, @@ -1177,6 +1231,7 @@ mod tests { output_index: 0, }, sequence_number: Some(2), + wire: WireEvent::new("test"), }); acc.process_event(&EventFrame { @@ -1187,6 +1242,7 @@ mod tests { 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..d207734 100644 --- a/crates/agentic-server-core/src/executor/engine.rs +++ b/crates/agentic-server-core/src/executor/engine.rs @@ -9,13 +9,14 @@ use std::sync::Arc; use async_stream::stream; use either::Either; -use tokio::sync::mpsc; +use tokio::sync::{Mutex, 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, }; +use super::gateway_accumulator::{GatewayStreamAccumulator, GatewayStreamContext}; use crate::executor::error::ExecutorResult; use crate::executor::inference::DONE_MARKER; use crate::executor::persist::persist_if_needed; @@ -25,7 +26,6 @@ use crate::executor::upstream::{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 +57,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 +94,7 @@ async fn run_until_gateway_tools_complete( exec_ctx: &ExecutionContext, auth: Option<&str>, stream_upstream: bool, - stream_events: Option<&mpsc::UnboundedSender>, + mut stream_context: Option>, ) -> 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,8 +104,9 @@ async fn run_until_gateway_tools_complete( let mut combined_usage: Option = None; for round in 0..MAX_GATEWAY_TOOL_ROUNDS { + let output_offset = combined_output.len(); let mut payload: ResponsePayload = if stream_upstream { - fetch_stream_payload(&ctx, exec_ctx, auth, ®istry, stream_events).await? + fetch_stream_payload(&ctx, exec_ctx, auth, ®istry, stream_context.as_mut(), output_offset).await? } else { fetch_blocking_payload(&ctx, exec_ctx, auth).await? }; @@ -136,7 +126,7 @@ 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?; + execute_and_emit_output_calls(¤t_output, ®istry, output_offset, stream_context.as_mut()).await?; let public_output = public_output_items(¤t_output, ®istry, &gateway_results); combined_output.extend(public_output); @@ -224,13 +214,16 @@ 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 stream_accumulator = Arc::new(Mutex::new(GatewayStreamAccumulator::new())); + let stream_accumulator_for_run = Arc::clone(&stream_accumulator); let mut run_handle = AbortOnDrop::new(tokio::spawn(async move { + let mut stream_accumulator = stream_accumulator_for_run.lock().await; run_until_gateway_tools_complete( ctx, exec_ctx_for_run.as_ref(), auth.as_deref(), true, - Some(&event_tx), + Some(GatewayStreamContext::new(&event_tx, &mut stream_accumulator)), ) .await })); @@ -246,11 +239,13 @@ fn run_stream(ctx: RequestContext, exec_ctx: Arc, auth: Option } match result { Err(e) => { - yield error_sse_chunk(&format!("stream task failed: {e}")); + let mut stream_accumulator = stream_accumulator.lock().await; + yield stream_accumulator.error_chunk(&format!("stream task failed: {e}")); yield DONE_MARKER.to_string(); } Ok(Err(e)) => { - yield error_sse_chunk(&e.to_string()); + let mut stream_accumulator = stream_accumulator.lock().await; + yield stream_accumulator.error_chunk(&e.to_string()); yield DONE_MARKER.to_string(); } Ok(Ok((payload, ctx))) => { @@ -258,14 +253,18 @@ fn run_stream(ctx: RequestContext, exec_ctx: Arc, auth: Option // `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 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 { + if let Err(e) = persist_if_needed(payload.clone(), ctx, ch, rh).await { warn!("persist failed: {e}"); } - yield terminal_event; + let mut stream_accumulator = stream_accumulator.lock().await; + match stream_accumulator.terminal_response_chunk(&payload) { + Ok(chunk) => yield chunk, + Err(e) => yield stream_accumulator.error_chunk(&e.to_string()), + } + drop(stream_accumulator); yield DONE_MARKER.to_string(); } } diff --git a/crates/agentic-server-core/src/executor/gateway.rs b/crates/agentic-server-core/src/executor/gateway.rs index 74a5b35..fd9e126 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::{GatewayStreamContext, synthetic_event}; use crate::executor::request::RequestContext; use crate::tool::{GatewayDispatchResult, ToolError, ToolOutput, ToolRegistry, ToolType}; use crate::types::io::output::{FunctionToolCall, GatewayCallStatus}; @@ -264,57 +265,55 @@ 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( plans: &[GatewayCallEventPlan], - stream_events: Option<&mpsc::UnboundedSender>, + stream_context: &mut GatewayStreamContext<'_>, ) -> 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), + ], + ); + stream_context.process_event(&mut added_event, 0)?; 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)), + ], + ); + stream_context.process_event(&mut in_progress_event, 0)?; + 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)), + ], + ); + stream_context.process_event(&mut searching_event, 0)?; } 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)), + ], + ); + stream_context.process_event(&mut in_progress_event, 0)?; } OutputItem::Message(_) | OutputItem::FunctionCall(_) @@ -329,11 +328,8 @@ fn emit_gateway_start_events( fn emit_gateway_completed_events( results: &[GatewayCallResult], plans: &[GatewayCallEventPlan], - stream_events: Option<&mpsc::UnboundedSender>, + stream_context: &mut GatewayStreamContext<'_>, ) -> ExecutorResult<()> { - let Some(sender) = stream_events else { - return Ok(()); - }; for result in results { let Some(public_output) = &result.public_output else { continue; @@ -344,9 +340,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 +350,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()), + ], + ); + stream_context.process_event(&mut completed_event, 0)?; + let mut done_event = synthetic_event( + SSEEventType::OutputItemDone, + [ + ("output_index".to_owned(), serde_json::json!(output_index)), + ("item".to_owned(), item), + ], + ); + stream_context.process_event(&mut done_event, 0)?; } Ok(()) } @@ -375,12 +375,16 @@ pub(super) async fn execute_and_emit_output_calls( output_items: &[OutputItem], registry: &ToolRegistry, output_offset: usize, - stream_events: Option<&mpsc::UnboundedSender>, + mut stream_context: Option<&mut GatewayStreamContext<'_>>, ) -> ExecutorResult> { let event_plans = gateway_event_plans(output_items, registry, output_offset); - emit_gateway_start_events(&event_plans, stream_events)?; + if let Some(stream_context) = &mut stream_context { + emit_gateway_start_events(&event_plans, stream_context)?; + } let gateway_results = execute_output_calls(output_items, registry).await?; - emit_gateway_completed_events(&gateway_results, &event_plans, stream_events)?; + if let Some(stream_context) = &mut stream_context { + emit_gateway_completed_events(&gateway_results, &event_plans, stream_context)?; + } Ok(gateway_results) } 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..9c190dc --- /dev/null +++ b/crates/agentic-server-core/src/executor/gateway_accumulator.rs @@ -0,0 +1,180 @@ +use serde_json::Value; +use tokio::sync::mpsc; + +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; + +pub struct GatewayStreamAccumulator { + next_sequence_number: u64, + emitted_created: bool, + emitted_in_progress: bool, +} + +pub(crate) struct GatewayStreamContext<'a> { + sender: &'a mpsc::UnboundedSender, + accumulator: &'a mut GatewayStreamAccumulator, +} + +impl<'a> GatewayStreamContext<'a> { + pub(crate) fn new( + sender: &'a mpsc::UnboundedSender, + accumulator: &'a mut GatewayStreamAccumulator, + ) -> Self { + Self { sender, accumulator } + } + + pub(crate) fn process_event(&mut self, frame: &mut EventFrame, output_offset: usize) -> ExecutorResult<()> { + if self.accumulator.process_event(frame, output_offset) { + emit_sse_frame(self.sender, frame)?; + } + Ok(()) + } +} + +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) + } + + pub fn process_event(&mut self, frame: &mut EventFrame, output_offset: usize) -> bool { + if !self.should_emit_lifecycle(frame.event_type) { + return false; + } + let sequence_number = Some(self.take_sequence_number()); + frame.sequence_number = sequence_number; + frame.wire.sequence_number = sequence_number; + rebase_output_index(&mut frame.wire, output_offset); + true + } + + pub(crate) fn terminal_response_chunk(&mut self, payload: &ResponsePayload) -> ExecutorResult { + let mut frame = terminal_response_frame(payload)?; + self.process_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.process_event(&mut frame, 0); + serialize_sse_frame(&frame).unwrap_or_else(|_| "data: {\"type\":\"error\"}\n\n".to_owned()) + } + + 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 wire = WireEvent::new(event_type.as_str()); + wire.rest.insert( + "response".to_owned(), + serde_json::to_value(payload).map_err(ExecutorError::JsonError)?, + ); + Ok(EventFrame::synthetic(event_type, wire)) +} + +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::Raw(wire.to_value()), + sequence_number: wire.sequence_number, + wire, + } +} + +pub(super) fn synthetic_event(event_type: SSEEventType, rest: impl IntoIterator) -> EventFrame { + let mut wire = WireEvent::new(event_type.as_str()); + wire.rest.extend(rest); + EventFrame::synthetic(event_type, wire) +} + +fn emit_sse_frame(sender: &mpsc::UnboundedSender, frame: &EventFrame) -> ExecutorResult<()> { + sender + .send(serialize_sse_frame(frame)?) + .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"); + } +} diff --git a/crates/agentic-server-core/src/executor/mod.rs b/crates/agentic-server-core/src/executor/mod.rs index 0c2b27a..a967160 100644 --- a/crates/agentic-server-core/src/executor/mod.rs +++ b/crates/agentic-server-core/src/executor/mod.rs @@ -10,6 +10,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..05aa711 100644 --- a/crates/agentic-server-core/src/executor/upstream.rs +++ b/crates/agentic-server-core/src/executor/upstream.rs @@ -3,16 +3,23 @@ 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, normalize_sse_line}; use crate::executor::accumulator::ResponseAccumulator; use crate::executor::error::{ExecutorError, ExecutorResult}; +use crate::executor::gateway_accumulator::GatewayStreamContext; 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, 'stream> { + request: &'a RequestContext, + registry: &'a ToolRegistry, + stream: &'a mut GatewayStreamContext<'stream>, + output_offset: usize, +} pub(super) async fn fetch_blocking_payload( ctx: &RequestContext, @@ -42,7 +49,8 @@ pub(super) async fn fetch_stream_payload( exec_ctx: &ExecutionContext, auth: Option<&str>, registry: &ToolRegistry, - stream_events: Option<&mpsc::UnboundedSender>, + stream_context: Option<&mut GatewayStreamContext<'_>>, + output_offset: usize, ) -> ExecutorResult { let url = exec_ctx.responses_url(); let upstream_request = ctx.enriched_request.to_upstream_request(true)?; @@ -56,21 +64,28 @@ 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 stream_context = stream_context; 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(mut frame) = normalize_sse_line(&line) { + log_upstream_failure(&frame, &ctx.response_id); + if let Some(stream) = stream_context.as_deref_mut() { + let mut emit_ctx = StreamEmitContext { + request: ctx, + registry, + stream, + output_offset, + }; + emit_upstream_stream_event( + &mut frame, + &mut emit_ctx, + &mut hidden_gateway_item_ids, + &mut pending_unnamed_function_events, + )?; + } + acc.process_event(&frame); } - acc.process_sse_line(&line); } acc.finish_stream(); let mut payload = acc.finalize( @@ -82,21 +97,12 @@ pub(super) async fn fetch_stream_payload( Ok(payload) } -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,70 +126,46 @@ 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: &mut EventFrame, + emit_ctx: &mut StreamEmitContext<'_, '_>, hidden_gateway_item_ids: &mut HashSet, - pending_unnamed_function_events: &mut HashMap>, + pending_unnamed_function_events: &mut HashMap>, ) -> 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) + 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, + frame, + emit_ctx, hidden_gateway_item_ids, pending_unnamed_function_events, )? { return Ok(()); } - emit_stream_line(data, ctx, registry, sender) + emit_stream_frame(frame, emit_ctx) } -fn emit_stream_line( - data: &str, - ctx: &RequestContext, - registry: &ToolRegistry, - sender: &mpsc::UnboundedSender, -) -> 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())) +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); + emit_ctx.stream.process_event(frame, emit_ctx.output_offset) } fn defer_or_flush_function_event( - line: &str, - payload: &EventPayload, - ctx: &RequestContext, - registry: &ToolRegistry, - sender: &mpsc::UnboundedSender, + frame: &mut EventFrame, + emit_ctx: &mut StreamEmitContext<'_, '_>, hidden_gateway_item_ids: &mut HashSet, - pending_unnamed_function_events: &mut HashMap>, + pending_unnamed_function_events: &mut HashMap>, ) -> ExecutorResult { - match payload { + match &frame.payload { EventPayload::OutputItemAdded { item_id, item_type, @@ -193,7 +175,7 @@ fn defer_or_flush_function_event( pending_unnamed_function_events .entry(item_id.clone()) .or_default() - .push(line.to_owned()); + .push(frame.clone()); Ok(true) } EventPayload::FunctionCallArgsDelta { item_id, .. } @@ -202,16 +184,16 @@ fn defer_or_flush_function_event( pending_unnamed_function_events .entry(item_id.clone()) .or_default() - .push(line.to_owned()); + .push(frame.clone()); Ok(true) } 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()); pending_unnamed_function_events.remove(item_id); return Ok(true); } - flush_pending_function_events(item_id, ctx, registry, sender, pending_unnamed_function_events)?; + flush_pending_function_events(item_id, emit_ctx, pending_unnamed_function_events)?; Ok(false) } EventPayload::OutputItemDone { @@ -223,13 +205,13 @@ 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()); pending_unnamed_function_events.remove(item_id); return Ok(true); } - flush_pending_function_events(item_id, ctx, registry, sender, pending_unnamed_function_events)?; + flush_pending_function_events(item_id, emit_ctx, pending_unnamed_function_events)?; Ok(false) } _ => Ok(false), @@ -238,26 +220,21 @@ fn defer_or_flush_function_event( 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>, ) -> 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 mut frame in frames { + emit_stream_frame(&mut frame, emit_ctx)?; } 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 +294,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 ea04bd2..4870ed7 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, HashSet}; -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}; @@ -281,6 +282,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 { @@ -498,6 +507,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 0e11b29..45d2aa5 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}; @@ -118,9 +119,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, } @@ -201,6 +201,10 @@ impl ToolRegistry { 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. #[must_use] pub fn gateway_owned<'a>(&self, calls: &'a [FunctionToolCall]) -> Vec<&'a FunctionToolCall> { diff --git a/crates/agentic-server-core/src/types/request_response.rs b/crates/agentic-server-core/src/types/request_response.rs index ce18768..bfc91f9 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..8cfac7d 100644 --- a/crates/agentic-server-core/tests/event_normalizer_test.rs +++ b/crates/agentic-server-core/tests/event_normalizer_test.rs @@ -135,6 +135,19 @@ fn test_unknown_event_type() { assert!(matches!(frame.payload, EventPayload::Raw(_))); } +#[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()); 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..4624e55 100644 --- a/crates/agentic-server-core/tests/web_search_tool_test.rs +++ b/crates/agentic-server-core/tests/web_search_tool_test.rs @@ -436,6 +436,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!({ @@ -938,6 +1067,129 @@ 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, + }; + + 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), + ("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; @@ -1488,13 +1740,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"];