fix(conversation): sanitize nested tool responses (#10609)

This commit is contained in:
Jasper
2026-08-10 09:54:34 -06:00
committed by GitHub
parent 433f621e14
commit 69ddf55cd1
4 changed files with 411 additions and 29 deletions
@@ -2,10 +2,11 @@ use crate::conversation::token_usage::{CostSource, ProviderUsage};
use crate::conversation::tool_result_serde;
use crate::mcp_utils::extract_text_from_resource;
use crate::utils::sanitize_unicode_tags;
use base64::Engine;
use chrono::Utc;
use rmcp::model::{
CallToolRequestParams, CallToolResult, ContentBlock, ElicitationAction, ImageContent,
JsonObject, PromptMessage, Role, TextContent,
JsonObject, PromptMessage, ResourceContents, Role, TextContent,
};
use serde::{Deserialize, Deserializer, Serialize};
use std::collections::HashSet;
@@ -53,18 +54,24 @@ where
.map_err(|e| Error::custom(format!("Failed to deserialize MessageContent: {}", e)))?;
for message_content in &mut content {
if let MessageContentBlock::Text(text_content) = message_content {
let original = &text_content.text;
let sanitized = sanitize_unicode_tags(original);
if *original != sanitized {
tracing::info!(
original = %original,
sanitized = %sanitized,
removed_count = original.len() - sanitized.len(),
"Unicode Tags sanitized during Message deserialization"
);
text_content.text = sanitized;
match message_content {
MessageContentBlock::Text(text_content) => {
let original = &text_content.text;
let sanitized = sanitize_unicode_tags(original);
if *original != sanitized {
tracing::info!(
original = %original,
sanitized = %sanitized,
removed_count = original.len() - sanitized.len(),
"Unicode Tags sanitized during Message deserialization"
);
text_content.text = sanitized;
}
}
MessageContentBlock::ToolResponse(response) => {
sanitize_tool_result_in_place(&mut response.tool_result);
}
_ => {}
}
}
@@ -76,6 +83,51 @@ where
pub type ProviderMetadata = serde_json::Map<String, serde_json::Value>;
pub type ToolResult<T> = Result<T, rmcp::model::ErrorData>;
pub(crate) fn sanitize_tool_result_in_place(tool_result: &mut ToolResult<CallToolResult>) {
match tool_result {
Ok(result) => {
for content in &mut result.content {
match content {
ContentBlock::Text(text) => {
text.text = sanitize_unicode_tags(&text.text);
}
ContentBlock::Resource(resource) => match &mut resource.resource {
ResourceContents::TextResourceContents { text, .. } => {
*text = sanitize_unicode_tags(text);
}
ResourceContents::BlobResourceContents { blob, .. } => {
let Ok(bytes) =
base64::engine::general_purpose::STANDARD.decode(blob.as_bytes())
else {
*blob = sanitize_unicode_tags(blob);
continue;
};
let Ok(text) = String::from_utf8(bytes) else {
continue;
};
let sanitized = sanitize_unicode_tags(&text);
if text != sanitized {
*blob = base64::engine::general_purpose::STANDARD
.encode(sanitized.as_bytes());
}
}
_ => {}
},
_ => {}
}
}
}
Err(error) => {
error.message = sanitize_unicode_tags(error.message.as_ref()).into();
}
}
}
fn sanitize_tool_result(mut tool_result: ToolResult<CallToolResult>) -> ToolResult<CallToolResult> {
sanitize_tool_result_in_place(&mut tool_result);
tool_result
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ToolRequest {
@@ -422,7 +474,7 @@ impl MessageContentBlock {
pub fn tool_response<S: Into<String>>(id: S, tool_result: ToolResult<CallToolResult>) -> Self {
MessageContentBlock::ToolResponse(ToolResponse {
id: id.into(),
tool_result,
tool_result: sanitize_tool_result(tool_result),
metadata: None,
})
}
@@ -434,7 +486,7 @@ impl MessageContentBlock {
) -> Self {
MessageContentBlock::ToolResponse(ToolResponse {
id: id.into(),
tool_result,
tool_result: sanitize_tool_result(tool_result),
metadata: metadata.cloned(),
})
}
@@ -1268,13 +1320,15 @@ pub struct TokenState {
#[cfg(test)]
mod tests {
use crate::conversation::message::{
ActionRequiredData, Message, MessageContentBlock, MessageMetadata,
ActionRequiredData, Message, MessageContentBlock, MessageMetadata, ProviderMetadata,
ToolResponse,
};
use base64::Engine;
use rmcp::model::{
Annotations, CallToolResult, ElicitationAction, ErrorCode, ErrorData, ImageContent,
TextContent,
ResourceContents, TextContent,
};
use rmcp::model::{CallToolRequestParams, ContentBlock, PromptMessage, ResourceContents, Role};
use rmcp::model::{CallToolRequestParams, ContentBlock, EmbeddedResource, PromptMessage, Role};
use rmcp::object;
use serde_json::Value;
@@ -1292,6 +1346,251 @@ mod tests {
assert_eq!(message.as_concat_text(), clean_text);
}
#[test]
fn test_tool_response_sanitizes_unicode_tags() {
let content = MessageContentBlock::tool_response(
"tool-1",
Ok(CallToolResult::success(vec![ContentBlock::text(
"visible\u{E0041}\u{E0042}text",
)])),
);
let MessageContentBlock::ToolResponse(response) = content else {
panic!("expected tool response");
};
let result = response.tool_result.unwrap();
let ContentBlock::Text(text) = &result.content[0] else {
panic!("expected text content");
};
assert_eq!(text.text, "visibletext");
}
#[test]
fn test_tool_response_with_metadata_sanitizes_unicode_tags() {
let mut metadata = ProviderMetadata::new();
metadata.insert("provider".to_string(), serde_json::json!("test"));
let tagged = ContentBlock::Text(
TextContent::new("result\u{E0041}")
.with_annotations(Annotations::default().with_audience(vec![Role::Assistant])),
);
let mut message = Message::user();
message.add_tool_response_with_metadata(
"tool-1",
Ok(CallToolResult::success(vec![tagged])),
Some(&metadata),
);
let MessageContentBlock::ToolResponse(response) = &message.content[0] else {
panic!("expected tool response");
};
assert_eq!(response.metadata.as_ref(), Some(&metadata));
let result = response.tool_result.as_ref().unwrap();
let text = &result.content[0];
let ContentBlock::Text(text) = text else {
panic!("expected text content");
};
assert_eq!(
text.annotations
.as_ref()
.and_then(|value| value.audience.as_ref()),
Some(&vec![Role::Assistant])
);
assert_eq!(text.text, "result");
}
#[test]
fn test_tool_response_sanitizes_error_message() {
let data = serde_json::json!({"retry": false});
let content = MessageContentBlock::tool_response(
"tool-1",
Err(ErrorData::new(
ErrorCode::INTERNAL_ERROR,
"error\u{E0041}text",
Some(data.clone()),
)),
);
let MessageContentBlock::ToolResponse(response) = content else {
panic!("expected tool response");
};
let error = response.tool_result.unwrap_err();
assert_eq!(error.message, "errortext");
assert_eq!(error.code, ErrorCode::INTERNAL_ERROR);
assert_eq!(error.data, Some(data));
}
#[test]
fn test_tool_response_sanitizes_text_resource() {
let resource = ResourceContents::TextResourceContents {
uri: "file:///result.txt".to_string(),
mime_type: Some("text/plain".to_string()),
text: "resource\u{E0041}text".to_string(),
meta: None,
};
let content = MessageContentBlock::tool_response(
"tool-1",
Ok(CallToolResult::success(vec![ContentBlock::Resource(
EmbeddedResource::new(resource),
)])),
);
let MessageContentBlock::ToolResponse(response) = content else {
panic!("expected tool response");
};
let result = response.tool_result.unwrap();
let ContentBlock::Resource(resource) = &result.content[0] else {
panic!("expected resource content");
};
let ResourceContents::TextResourceContents {
uri,
mime_type,
text,
meta,
} = &resource.resource
else {
panic!("expected text resource");
};
assert_eq!(uri, "file:///result.txt");
assert_eq!(mime_type.as_deref(), Some("text/plain"));
assert_eq!(text, "resourcetext");
assert!(meta.is_none());
}
#[test]
fn test_tool_response_sanitizes_utf8_blob_resource() {
let blob =
base64::engine::general_purpose::STANDARD.encode("resource\u{E0041}text".as_bytes());
let resource = ResourceContents::BlobResourceContents {
uri: "file:///result.txt".to_string(),
mime_type: Some("text/plain".to_string()),
blob,
meta: None,
};
let content = MessageContentBlock::tool_response(
"tool-1",
Ok(CallToolResult::success(vec![ContentBlock::Resource(
EmbeddedResource::new(resource),
)])),
);
let MessageContentBlock::ToolResponse(response) = content else {
panic!("expected tool response");
};
let result = response.tool_result.unwrap();
let ContentBlock::Resource(resource) = &result.content[0] else {
panic!("expected resource content");
};
let ResourceContents::BlobResourceContents { blob, .. } = &resource.resource else {
panic!("expected blob resource");
};
assert_eq!(
base64::engine::general_purpose::STANDARD
.decode(blob)
.unwrap(),
b"resourcetext"
);
}
#[test]
fn test_tool_response_sanitizes_malformed_blob_resource() {
let resource = ResourceContents::BlobResourceContents {
uri: "file:///result.txt".to_string(),
mime_type: Some("text/plain".to_string()),
blob: "malformed\u{E0041}text".to_string(),
meta: None,
};
let content = MessageContentBlock::tool_response(
"tool-1",
Ok(CallToolResult::success(vec![ContentBlock::Resource(
EmbeddedResource::new(resource),
)])),
);
let MessageContentBlock::ToolResponse(response) = content else {
panic!("expected tool response");
};
let result = response.tool_result.unwrap();
let ContentBlock::Resource(resource) = &result.content[0] else {
panic!("expected resource content");
};
let ResourceContents::BlobResourceContents { blob, .. } = &resource.resource else {
panic!("expected blob resource");
};
assert_eq!(blob, "malformedtext");
}
#[test]
fn test_deserialization_sanitizes_persisted_tool_response() {
let message = Message::new(
Role::User,
1,
vec![MessageContentBlock::ToolResponse(ToolResponse {
id: "tool-1".to_string(),
tool_result: Ok(CallToolResult::success(vec![ContentBlock::text(
"persisted\u{E0041}text",
)])),
metadata: None,
})],
);
let json = serde_json::to_string(&message).unwrap();
let deserialized: Message = serde_json::from_str(&json).unwrap();
let MessageContentBlock::ToolResponse(response) = &deserialized.content[0] else {
panic!("expected tool response");
};
let result = response.tool_result.as_ref().unwrap();
let ContentBlock::Text(text) = &result.content[0] else {
panic!("expected text content");
};
assert_eq!(text.text, "persistedtext");
}
#[test]
fn test_content_deserialization_sanitizes_persisted_tool_response() {
let content = vec![MessageContentBlock::ToolResponse(ToolResponse {
id: "tool-1".to_string(),
tool_result: Ok(CallToolResult::success(vec![ContentBlock::text(
"persisted\u{E0041}text",
)])),
metadata: None,
})];
let json = serde_json::to_string(&content).unwrap();
let deserialized: Vec<MessageContentBlock> = serde_json::from_str(&json).unwrap();
let MessageContentBlock::ToolResponse(response) = &deserialized[0] else {
panic!("expected tool response");
};
let result = response.tool_result.as_ref().unwrap();
let ContentBlock::Text(text) = &result.content[0] else {
panic!("expected text content");
};
assert_eq!(text.text, "persistedtext");
}
#[test]
fn test_tool_response_sanitization_preserves_legitimate_content() {
let text = ContentBlock::Text(
TextContent::new("世界 🌍 café")
.with_annotations(Annotations::default().with_audience(vec![Role::Assistant])),
);
let image = ContentBlock::Image(
ImageContent::new("image-data", "image/png")
.with_annotations(Annotations::default().with_audience(vec![Role::User])),
);
let mut result = CallToolResult::success(vec![text, image]);
result.structured_content = Some(serde_json::json!({"safe": "世界"}));
result.meta = Some(rmcp::model::Meta(object!({"source": "test"})));
let expected = result.clone();
let content = MessageContentBlock::tool_response("tool-1", Ok(result));
let MessageContentBlock::ToolResponse(response) = content else {
panic!("expected tool response");
};
assert_eq!(response.tool_result.unwrap(), expected);
}
#[test]
fn test_message_serialization() {
let message = Message::assistant()
@@ -115,6 +115,7 @@ where
pub mod call_tool_result {
use super::*;
use crate::conversation::message::sanitize_tool_result_in_place;
use rmcp::model::{CallToolResult, ContentBlock};
pub fn serialize<S>(
@@ -150,41 +151,44 @@ pub mod call_tool_result {
let format = ResultFormat::deserialize(deserializer)?;
match format {
let mut result = match format {
ResultFormat::SuccessWithCallToolResult { status, value } => {
if status == "success" {
Ok(Ok(value))
Ok(value)
} else {
Err(serde::de::Error::custom(format!(
return Err(serde::de::Error::custom(format!(
"Expected status 'success', got '{}'",
status
)))
)));
}
}
ResultFormat::SuccessWithContentVec { status, value } => {
if status == "success" {
Ok(Ok(CallToolResult::success(value)))
Ok(CallToolResult::success(value))
} else {
Err(serde::de::Error::custom(format!(
return Err(serde::de::Error::custom(format!(
"Expected status 'success', got '{}'",
status
)))
)));
}
}
ResultFormat::Error { status, error } => {
if status == "error" {
Ok(Err(ErrorData {
Err(ErrorData {
code: ErrorCode::INTERNAL_ERROR,
message: Cow::from(error),
data: None,
}))
})
} else {
Err(serde::de::Error::custom(format!(
return Err(serde::de::Error::custom(format!(
"Expected status 'error', got '{}'",
status
)))
)));
}
}
}
};
sanitize_tool_result_in_place(&mut result);
Ok(result)
}
}
@@ -333,6 +333,57 @@ mod tests {
assert!(content.iter().any(|c| c["type"] == "toolResponse"));
}
#[test]
fn sanitizes_unicode_tags_in_tool_result() {
let jsonl = serde_json::json!({
"type": "user",
"sessionId": "s",
"uuid": "u1",
"timestamp": "2026-01-01T00:00:00Z",
"cwd": "/tmp",
"message": {
"role": "user",
"content": [{
"type": "tool_result",
"tool_use_id": "toolu_1",
"content": [{"type": "text", "text": "visible\u{E0041}世界"}]
}]
}
})
.to_string();
let json = convert(&jsonl).unwrap();
assert!(json.contains("visible世界"));
assert!(!json.contains('\u{E0041}'));
}
#[test]
fn sanitizes_unicode_tags_in_tool_result_error() {
let jsonl = serde_json::json!({
"type": "user",
"sessionId": "s",
"uuid": "u1",
"timestamp": "2026-01-01T00:00:00Z",
"cwd": "/tmp",
"message": {
"role": "user",
"content": [{
"type": "tool_result",
"tool_use_id": "toolu_1",
"is_error": true,
"content": "failed\u{E0041}café"
}]
}
})
.to_string();
let json = convert(&jsonl).unwrap();
assert!(json.contains("failedcafé"));
assert!(!json.contains('\u{E0041}'));
}
#[test]
fn emits_cache_token_breakdown() {
let jsonl = r#"{"type":"user","sessionId":"s","uuid":"u1","timestamp":"2026-01-01T00:00:01Z","cwd":"/tmp","message":{"role":"user","content":"hi"}}
@@ -341,6 +341,34 @@ mod tests {
.any(|c| c["type"] == "toolResponse"));
}
#[test]
fn sanitizes_unicode_tags_in_function_call_output() {
let jsonl = [
serde_json::json!({
"timestamp": "2026-05-22T13:37:22Z",
"type": "session_meta",
"payload": {"id": "s", "cwd": "/w"}
})
.to_string(),
serde_json::json!({
"timestamp": "2026-05-22T13:37:23Z",
"type": "response_item",
"payload": {
"type": "function_call_output",
"call_id": "call_1",
"output": "visible\u{E0041}世界"
}
})
.to_string(),
]
.join("\n");
let json = convert(&jsonl).unwrap();
assert!(json.contains("visible世界"));
assert!(!json.contains('\u{E0041}'));
}
#[test]
fn first_user_text_skips_context_blobs() {
let jsonl = r#"{"timestamp":"2026-05-22T13:37:22Z","type":"session_meta","payload":{"id":"s","cwd":"/w"}}