Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -279,6 +279,7 @@ fn finish_anthropic_stream(state: &mut StreamTranslationState) -> Vec<Value> {
}
}

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) {
Expand All @@ -287,6 +288,7 @@ fn finish_anthropic_stream(state: &mut StreamTranslationState) -> Vec<Value> {
}
tool.started = false;
}
state.active_anthropic_tool = None;
}

if !state.emitted_content_block {
Expand Down Expand Up @@ -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));
Expand Down
46 changes: 35 additions & 11 deletions crates/switchyard-translation/src/codecs/openai_chat/stream.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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()),
});
Expand Down Expand Up @@ -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 => {
Expand Down
4 changes: 4 additions & 0 deletions crates/switchyard-translation/src/codecs/stream.rs
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,10 @@ pub struct StreamTranslationState {
pub(crate) emitted_content_block: bool,
pub(crate) tool_states: BTreeMap<usize, StreamToolState>,
#[serde(default)]
pub(crate) pending_chat_tool_names: BTreeMap<usize, String>,
#[serde(default)]
pub(crate) active_anthropic_tool: Option<usize>,
#[serde(default)]
pub(crate) deferred_anthropic_tools: Vec<usize>,
/// Reasoning text observed while DECODING, per output index, so a completed item
/// that repeats already-streamed text is not decoded twice.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
86 changes: 76 additions & 10 deletions crates/switchyard-translation/tests/stream_translation.rs
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,68 @@ use common::{REASONING_MODEL, text_and_encrypted_reasoning_details};

type TestResult = std::result::Result<(), Box<dyn std::error::Error + Send + Sync>>;

#[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 (`<block>_start`, `<delta>`,
// `<block>_stop`) so ordering assertions stay readable without restating each payload.
fn event_labels(events: &[Value]) -> Vec<String> {
Expand Down Expand Up @@ -2071,24 +2133,28 @@ fn translated_responses_text_and_tool_events_keep_item_identity_until_done() ->
.iter()
.filter(|event| event["output_index"] == index)
.collect::<Vec<_>>();
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::<Vec<_>>(),
[
"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");
Expand Down
Loading