diff --git a/Cargo.lock b/Cargo.lock index a485c5cc4..c1561e49b 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2766,6 +2766,7 @@ dependencies = [ "rand 0.9.4", "serde", "serde_json", + "sha2 0.10.9", "strum 0.28.0", "thiserror 2.0.18", "tokio", diff --git a/lib/components/fabro-llm/Cargo.toml b/lib/components/fabro-llm/Cargo.toml index b10321253..d7bcb8df7 100644 --- a/lib/components/fabro-llm/Cargo.toml +++ b/lib/components/fabro-llm/Cargo.toml @@ -21,6 +21,7 @@ anyhow.workspace = true thiserror.workspace = true serde.workspace = true serde_json.workspace = true +sha2.workspace = true strum.workspace = true tokio.workspace = true uuid.workspace = true diff --git a/lib/components/fabro-llm/src/codec/bedrock_converse/decode.rs b/lib/components/fabro-llm/src/codec/bedrock_converse/decode.rs index 36ec5878d..f9cf0f87e 100644 --- a/lib/components/fabro-llm/src/codec/bedrock_converse/decode.rs +++ b/lib/components/fabro-llm/src/codec/bedrock_converse/decode.rs @@ -279,6 +279,21 @@ mod tests { assert!(decode_content_block(&serde_json::json!({"text": ""})).is_none()); } + #[test] + fn tool_use_names_are_preserved_verbatim() { + let block = serde_json::json!({ + "toolUse": { + "toolUseId": "tool-1", + "name": "search???", + "input": {} + } + }); + let Some(ContentPart::ToolCall(tool_call)) = decode_content_block(&block) else { + panic!("expected tool call"); + }; + assert_eq!(tool_call.name, "search???"); + } + #[test] fn reasoning_text_block_round_trips_signature() { let block = serde_json::json!({ diff --git a/lib/components/fabro-llm/src/codec/bedrock_converse/encode.rs b/lib/components/fabro-llm/src/codec/bedrock_converse/encode.rs index 39c9c5c10..131a97cd3 100644 --- a/lib/components/fabro-llm/src/codec/bedrock_converse/encode.rs +++ b/lib/components/fabro-llm/src/codec/bedrock_converse/encode.rs @@ -4,6 +4,7 @@ use base64::Engine; use base64::engine::general_purpose::STANDARD as BASE64; use serde_json::{Map, Value, json}; +use super::sanitize; use crate::codec::{CodecCtx, EncodedRequest, extract_system_prompt, merge_named_provider_options}; use crate::error::Error; use crate::types::{ContentPart, Message, Request, Role, ToolChoice}; @@ -129,7 +130,7 @@ fn encode_message(message: &Message) -> Option { let text = message.text(); blocks.push(json!({ "toolResult": { - "toolUseId": tool_call_id, + "toolUseId": sanitize::tool_use_id(tool_call_id), "content": [{ "text": text }], } })); @@ -185,8 +186,8 @@ fn encode_content_part(part: &ContentPart) -> Option { }; Some(json!({ "toolUse": { - "toolUseId": tool_call.id, - "name": tool_call.name, + "toolUseId": sanitize::tool_use_id(&tool_call.id), + "name": sanitize::tool_name(&tool_call.name), "input": input, } })) @@ -197,7 +198,10 @@ fn encode_content_part(part: &ContentPart) -> Option { other => json!([{ "json": other }]), }; let mut block = Map::new(); - block.insert("toolUseId".to_string(), json!(result.tool_call_id)); + block.insert( + "toolUseId".to_string(), + json!(sanitize::tool_use_id(&result.tool_call_id)), + ); block.insert("content".to_string(), content); if result.is_error { block.insert("status".to_string(), json!("error")); @@ -522,6 +526,119 @@ mod tests { assert_eq!(tool_use["input"], json!({})); } + #[test] + fn historical_tool_names_are_sanitized_without_mutating_the_request() { + let mut request = base_request("claude"); + request.messages = vec![Message { + role: Role::Assistant, + content: vec![ContentPart::ToolCall(ToolCall::new( + "tool-1", + "search???", + json!({}), + ))], + name: None, + tool_call_id: None, + }]; + + let encoded = encode_with(&request); + let tool_use = &encoded.body["messages"][0]["content"][0]["toolUse"]; + assert_eq!(tool_use["name"], "search___"); + + let ContentPart::ToolCall(original) = &request.messages[0].content[0] else { + panic!("expected original tool call"); + }; + assert_eq!(original.name, "search???"); + } + + #[test] + fn sanitized_tool_use_ids_remain_paired() { + for id in ["bad id!".to_string(), "x".repeat(100)] { + let mut request = base_request("claude"); + request.messages = vec![ + Message { + role: Role::Assistant, + content: vec![ContentPart::ToolCall(ToolCall::new( + &id, + "search", + json!({}), + ))], + name: None, + tool_call_id: None, + }, + Message { + role: Role::Tool, + content: vec![ContentPart::ToolResult(ToolResult::success( + &id, + json!("done"), + ))], + name: None, + tool_call_id: Some(id.clone()), + }, + ]; + + let encoded = encode_with(&request); + let tool_use_id = &encoded.body["messages"][0]["content"][0]["toolUse"]["toolUseId"]; + let tool_result_id = + &encoded.body["messages"][1]["content"][0]["toolResult"]["toolUseId"]; + assert_eq!(tool_use_id, tool_result_id); + assert!(tool_use_id.as_str().is_some_and(|value| value.len() <= 64)); + } + } + + #[test] + fn tool_role_fallback_sanitizes_the_tool_use_id() { + let mut request = base_request("claude"); + request.messages = vec![Message { + role: Role::Tool, + content: vec![], + name: None, + tool_call_id: Some("bad id!".to_string()), + }]; + + let encoded = encode_with(&request); + assert_eq!( + encoded.body["messages"][0]["content"][0]["toolResult"]["toolUseId"], + "bad_id_" + ); + } + + #[test] + fn overlength_tool_names_encode_within_the_bedrock_limit() { + let mut request = base_request("claude"); + request.messages = vec![Message { + role: Role::Assistant, + content: vec![ContentPart::ToolCall(ToolCall::new( + "tool-1", + "x".repeat(100), + json!({}), + ))], + name: None, + tool_call_id: None, + }]; + + let encoded = encode_with(&request); + let name = encoded.body["messages"][0]["content"][0]["toolUse"]["name"] + .as_str() + .unwrap(); + assert_eq!(name.len(), 64); + } + + #[test] + fn tool_definition_names_remain_unsanitized() { + let mut request = base_request("claude"); + request.tools = Some(vec![ToolDefinition::function( + "weird.name", + "Deliberately invalid for Bedrock", + json!({"type": "object"}), + )]); + + let encoded = encode_with(&request); + assert_eq!( + encoded.body["toolConfig"]["tools"][0]["toolSpec"]["name"], + "weird.name" + ); + } + #[test] fn thinking_parts_restructure_into_reasoning_text_blocks() { let mut request = base_request("claude"); diff --git a/lib/components/fabro-llm/src/codec/bedrock_converse/mod.rs b/lib/components/fabro-llm/src/codec/bedrock_converse/mod.rs index 12b315094..c55c0a3b3 100644 --- a/lib/components/fabro-llm/src/codec/bedrock_converse/mod.rs +++ b/lib/components/fabro-llm/src/codec/bedrock_converse/mod.rs @@ -12,6 +12,7 @@ mod decode; mod encode; +mod sanitize; mod stream; use crate::codec::{Codec, CodecCtx, EncodedRequest, StreamDecoder}; diff --git a/lib/components/fabro-llm/src/codec/bedrock_converse/sanitize.rs b/lib/components/fabro-llm/src/codec/bedrock_converse/sanitize.rs new file mode 100644 index 000000000..a28e3c911 --- /dev/null +++ b/lib/components/fabro-llm/src/codec/bedrock_converse/sanitize.rs @@ -0,0 +1,135 @@ +//! Bedrock Converse tool identifier sanitization. +//! +//! Tool names must match `[a-zA-Z0-9_-]+`; tool-use IDs additionally allow +//! `.` and `:`. Both are limited to 64 characters. These helpers rewrite only +//! the Bedrock wire view: the canonical transcript retains provider output +//! verbatim. Every encoded tool-use ID site must use the same helper so +//! `toolUse` and `toolResult` blocks remain paired. + +use std::borrow::Cow; + +use sha2::{Digest, Sha256}; + +const MAX_LENGTH: usize = 64; +const HASH_HEX_LENGTH: usize = 16; +const PREFIX_LENGTH: usize = MAX_LENGTH - 1 - HASH_HEX_LENGTH; + +pub(super) fn tool_name(name: &str) -> Cow<'_, str> { + sanitize(name, "unknown_tool", is_tool_name_char) +} + +pub(super) fn tool_use_id(id: &str) -> Cow<'_, str> { + sanitize(id, "unknown_tool_use_id", is_tool_use_id_char) +} + +fn sanitize<'a>( + value: &'a str, + empty_fallback: &'static str, + is_allowed: fn(char) -> bool, +) -> Cow<'a, str> { + if value.is_empty() { + return Cow::Borrowed(empty_fallback); + } + if value.len() <= MAX_LENGTH && value.chars().all(is_allowed) { + return Cow::Borrowed(value); + } + + let mut sanitized = String::with_capacity(value.len()); + for character in value.chars() { + sanitized.push(if is_allowed(character) { + character + } else { + '_' + }); + } + + if sanitized.len() <= MAX_LENGTH { + return Cow::Owned(sanitized); + } + + Cow::Owned(truncate_with_hash(&sanitized, value)) +} + +fn is_tool_name_char(character: char) -> bool { + character.is_ascii_alphanumeric() || matches!(character, '_' | '-') +} + +fn is_tool_use_id_char(character: char) -> bool { + is_tool_name_char(character) || matches!(character, '.' | ':') +} + +fn truncate_with_hash(sanitized: &str, original: &str) -> String { + debug_assert!(sanitized.is_ascii()); + let digest = Sha256::digest(original.as_bytes()); + let digest_hex = format!("{digest:x}"); + format!( + "{}-{}", + &sanitized[..PREFIX_LENGTH], + &digest_hex[..HASH_HEX_LENGTH] + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn valid_values_are_returned_borrowed() { + for name in ["search", "TaskList", "a-b_c9"] { + assert!(matches!(tool_name(name), Cow::Borrowed(value) if value == name)); + } + + let max_length = "a".repeat(64); + assert!(matches!( + tool_name(&max_length), + Cow::Borrowed(value) if value == max_length + )); + + let id = "functions.read_file:4"; + assert!(matches!( + tool_use_id(id), + Cow::Borrowed(value) if value == id + )); + assert_eq!(tool_name(id), "functions_read_file_4"); + } + + #[test] + fn invalid_characters_are_replaced() { + assert_eq!(tool_name("search???"), "search___"); + assert_eq!(tool_name("bad name"), "bad_name"); + assert_eq!(tool_use_id("bad id!"), "bad_id_"); + } + + #[test] + fn non_ascii_characters_become_single_underscores() { + let sanitized = tool_name("before🙂after"); + assert_eq!(sanitized, "before_after"); + assert!(sanitized.is_ascii()); + } + + #[test] + fn empty_values_use_nonempty_fallbacks() { + assert_eq!(tool_name(""), "unknown_tool"); + assert_eq!(tool_use_id(""), "unknown_tool_use_id"); + } + + #[test] + fn overlength_values_use_deterministic_hash_suffixes() { + let boundary = "a".repeat(65); + let first = tool_name(&boundary); + let second = tool_name(&boundary); + assert!(matches!(&first, Cow::Owned(_))); + assert_eq!(first, second); + assert_eq!(first.len(), 64); + assert!( + first + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-')) + ); + + let shared_prefix = "x".repeat(99); + let left = tool_name(&format!("{shared_prefix}a")).into_owned(); + let right = tool_name(&format!("{shared_prefix}b")).into_owned(); + assert_ne!(left, right); + } +} diff --git a/lib/components/fabro-llm/src/codec/bedrock_converse/stream.rs b/lib/components/fabro-llm/src/codec/bedrock_converse/stream.rs index 0facf2583..0d413cabf 100644 --- a/lib/components/fabro-llm/src/codec/bedrock_converse/stream.rs +++ b/lib/components/fabro-llm/src/codec/bedrock_converse/stream.rs @@ -406,6 +406,22 @@ mod tests { assert!(!tool_call.arguments.is_null()); } + #[test] + fn streamed_tool_use_names_are_preserved_verbatim() { + let mut d = decoder(); + feed(&mut d, "messageStart", r#"{"role":"assistant"}"#); + feed( + &mut d, + "contentBlockStart", + r#"{"start":{"toolUse":{"toolUseId":"tool-1","name":"search???"}},"contentBlockIndex":0}"#, + ); + let stop = feed(&mut d, "contentBlockStop", r#"{"contentBlockIndex":0}"#); + let StreamEvent::ToolCallEnd { tool_call } = &stop[0] else { + panic!("expected ToolCallEnd"); + }; + assert_eq!(tool_call.name, "search???"); + } + #[test] fn tool_use_accumulates_string_input_fragments() { let mut d = decoder();