diff --git a/crates/switchyard-translation/src/codecs/anthropic/stream.rs b/crates/switchyard-translation/src/codecs/anthropic/stream.rs index b3a0b5f7e..647d4a3bd 100644 --- a/crates/switchyard-translation/src/codecs/anthropic/stream.rs +++ b/crates/switchyard-translation/src/codecs/anthropic/stream.rs @@ -279,6 +279,7 @@ fn finish_anthropic_stream(state: &mut StreamTranslationState) -> Vec { } } + state.active_anthropic_tool = None; for index in std::mem::take(&mut state.deferred_anthropic_tools) { out.extend(encode_anthropic_tool_delta(state, index, None, None, None)); if let Some(tool) = state.tool_states.get_mut(&index) { @@ -287,6 +288,7 @@ fn finish_anthropic_stream(state: &mut StreamTranslationState) -> Vec { } tool.started = false; } + state.active_anthropic_tool = None; } if !state.emitted_content_block { @@ -514,10 +516,8 @@ fn encode_anthropic_tool_delta( state.text_block_started = false; } - let other_tool_started = state - .tool_states - .iter() - .any(|(&tool_index, tool)| tool_index != index && tool.started); + // Reserve the first call even while its name is still being assembled. + let other_tool_started = *state.active_anthropic_tool.get_or_insert(index) != index; let tool = state.tool_states.entry(index).or_default(); if let Some(id) = id.filter(|id| !id.is_empty()) { tool.id = Some(sanitize_anthropic_tool_use_id(&id)); diff --git a/crates/switchyard-translation/src/codecs/openai_chat/stream.rs b/crates/switchyard-translation/src/codecs/openai_chat/stream.rs index 5e7c41409..7a902627b 100644 --- a/crates/switchyard-translation/src/codecs/openai_chat/stream.rs +++ b/crates/switchyard-translation/src/codecs/openai_chat/stream.rs @@ -141,27 +141,46 @@ fn decode_openai_chat_stream( for tool_call in tool_calls { if let Some(tool_call) = tool_call.as_object() { let function = tool_call.get("function").and_then(Value::as_object); + let index = + tool_call.get("index").and_then(Value::as_u64).unwrap_or(0) as usize; + if let Some(name) = function + .and_then(|function| function.get("name")) + .and_then(Value::as_str) + .filter(|name| !name.is_empty()) + { + state + .pending_chat_tool_names + .entry(index) + .or_default() + .push_str(name); + } + let arguments_delta = function + .and_then(|function| function.get("arguments")) + .and_then(Value::as_str) + .map(ToOwned::to_owned); out.push(LlmResponseChunk::ToolCallDelta { - index: tool_call.get("index").and_then(Value::as_u64).unwrap_or(0) - as usize, + index, id: tool_call .get("id") .and_then(Value::as_str) .map(ToOwned::to_owned), - name: function - .and_then(|function| function.get("name")) - .and_then(Value::as_str) - .map(ToOwned::to_owned), - arguments_delta: function - .and_then(|function| function.get("arguments")) - .and_then(Value::as_str) - .map(ToOwned::to_owned), + // Names may continue after arguments begin; emit them at finish_reason. + name: None, + arguments_delta, }); } } } } if let Some(reason) = choice.get("finish_reason").and_then(Value::as_str) { + for (index, name) in std::mem::take(&mut state.pending_chat_tool_names) { + out.push(LlmResponseChunk::ToolCallDelta { + index, + id: None, + name: Some(name), + arguments_delta: None, + }); + } out.push(LlmResponseChunk::MessageStop { reason: Some(reason.to_string()), }); @@ -243,11 +262,16 @@ fn encode_openai_chat_stream( name, arguments_delta, } => { + let tool = state.tool_states.entry(index).or_default(); + // Repeated full names would be concatenated by Chat clients as new fragments. + let name = name.filter(|name| tool.name.as_ref() != Some(name)); + if name.is_some() { + tool.name = name.clone(); + } // The source index counts every content block (Anthropic) or output item // (Responses), so text ahead of the first tool call shifts it. Chat clients use // the index as a subscript into `tool_calls`, so number calls in that array // instead, in order of first appearance. - let tool = state.tool_states.entry(index).or_default(); let chat_index = match tool.chat_tool_index { Some(chat_index) => chat_index, None => { diff --git a/crates/switchyard-translation/src/codecs/stream.rs b/crates/switchyard-translation/src/codecs/stream.rs index 3732ff2b6..6c8072d50 100644 --- a/crates/switchyard-translation/src/codecs/stream.rs +++ b/crates/switchyard-translation/src/codecs/stream.rs @@ -53,6 +53,10 @@ pub struct StreamTranslationState { pub(crate) emitted_content_block: bool, pub(crate) tool_states: BTreeMap, #[serde(default)] + pub(crate) pending_chat_tool_names: BTreeMap, + #[serde(default)] + pub(crate) active_anthropic_tool: Option, + #[serde(default)] pub(crate) deferred_anthropic_tools: Vec, /// Reasoning text observed while DECODING, per output index, so a completed item /// that repeats already-streamed text is not decoded twice. diff --git a/crates/switchyard-translation/tests/anthropic_parallel_tools.rs b/crates/switchyard-translation/tests/anthropic_parallel_tools.rs index bfe1e6682..51251a10b 100644 --- a/crates/switchyard-translation/tests/anthropic_parallel_tools.rs +++ b/crates/switchyard-translation/tests/anthropic_parallel_tools.rs @@ -55,10 +55,8 @@ fn parallel_chat_tools_stream_as_ordered_nonoverlapping_anthropic_blocks() -> Te ]}), )?; assert!( - first_fragments - .iter() - .any(|event| { event["delta"]["partial_json"] == "{\"city\":\"Pa" }), - "the first tool should still stream its arguments before EOF" + first_fragments.is_empty(), + "tool names are not complete yet" ); events.extend(first_fragments); events.extend(translate( diff --git a/crates/switchyard-translation/tests/stream_translation.rs b/crates/switchyard-translation/tests/stream_translation.rs index 7893ff22b..256e1a94f 100644 --- a/crates/switchyard-translation/tests/stream_translation.rs +++ b/crates/switchyard-translation/tests/stream_translation.rs @@ -18,6 +18,68 @@ use common::{REASONING_MODEL, text_and_encrypted_reasoning_details}; type TestResult = std::result::Result<(), Box>; +#[test] +fn fragmented_chat_tool_names_are_complete_in_cross_format_streams() -> TestResult { + let engine = TranslationEngine::default(); + for target in [WireFormat::AnthropicMessages, WireFormat::OpenAiResponses] { + for (early_arguments, arguments) in [("{", "}"), ("", "{}"), ("", "")] { + let mut state = StreamTranslationState::new(WireFormat::OpenAiChat, target); + let mut events = Vec::new(); + for event in [ + json!({"choices": [{"delta": {"tool_calls": [{"index": 0, + "id": "call_weather", "function": {"name": "wea", "arguments": ""}}]}}]}), + json!({"choices": [{"delta": {"tool_calls": [{"index": 0, + "function": {"arguments": early_arguments}}]}}]}), + json!({"choices": [{"delta": {"tool_calls": [{"index": 0, + "function": {"name": "ther", "arguments": arguments}}]}}]}), + json!({"choices": [{"delta": {}, "finish_reason": "tool_calls"}]}), + ] { + let translated = + engine.translate_event(&mut state, WireFormat::OpenAiChat, target, &event)?; + if event["choices"][0]["finish_reason"].is_null() { + assert!( + translated.iter().all(|event| { + event["type"] != "content_block_start" + && event["type"] != "response.output_item.added" + }), + "tool name must not be announced before completion" + ); + } + events.extend(translated); + } + events.extend(engine.finish_stream(&mut state, target)?); + let names: Vec<_> = events + .iter() + .filter_map(|event| { + event + .get("content_block") + .or_else(|| event.get("item"))? + .get("name")? + .as_str() + }) + .collect(); + let expected = if target == WireFormat::AnthropicMessages { + vec!["weather"] + } else { + vec!["weather", "weather"] + }; + assert_eq!(names, expected, "{target:?}, arguments={arguments:?}"); + let emitted_arguments: String = events + .iter() + .filter_map(|event| { + if event["type"] == "response.function_call_arguments.delta" { + event["delta"].as_str() + } else { + event["delta"]["partial_json"].as_str() + } + }) + .collect(); + assert_eq!(emitted_arguments, format!("{early_arguments}{arguments}")); + } + } + Ok(()) +} + // Reduces Anthropic stream events to ordered labels (`_start`, ``, // `_stop`) so ordering assertions stay readable without restating each payload. fn event_labels(events: &[Value]) -> Vec { @@ -2071,24 +2133,28 @@ fn translated_responses_text_and_tool_events_keep_item_identity_until_done() -> .iter() .filter(|event| event["output_index"] == index) .collect::>(); + let mut expected_types = vec![ + "response.output_item.added", + "response.function_call_arguments.delta", + "response.function_call_arguments.done", + "response.output_item.done", + ]; + if source == WireFormat::AnthropicMessages { + expected_types.insert(2, "response.function_call_arguments.delta"); + } assert_eq!( tool_events .iter() .map(|event| event["type"].as_str().unwrap_or_default()) .collect::>(), - [ - "response.output_item.added", - "response.function_call_arguments.delta", - "response.function_call_arguments.delta", - "response.function_call_arguments.done", - "response.output_item.done" - ] + expected_types ); - let done = tool_events[3]; + let done = tool_events[tool_events.len() - 2]; assert_eq!(done["name"], name); assert_eq!(done["arguments"], arguments); - assert_eq!(tool_events[4]["item"]["id"], items[&index]["id"]); - assert_eq!(tool_events[4]["item"]["call_id"], call_id); + let item_done = tool_events[tool_events.len() - 1]; + assert_eq!(item_done["item"]["id"], items[&index]["id"]); + assert_eq!(item_done["item"]["call_id"], call_id); } let completed = events.last().ok_or("missing response completion")?; assert_eq!(completed["type"], "response.completed");