added validation and debug for invalid call tool result (#6368)
This commit is contained in:
@@ -33,6 +33,7 @@ use crate::conversation::message::{
|
|||||||
ActionRequiredData, Message, MessageContent, ProviderMetadata, SystemNotificationType,
|
ActionRequiredData, Message, MessageContent, ProviderMetadata, SystemNotificationType,
|
||||||
ToolRequest,
|
ToolRequest,
|
||||||
};
|
};
|
||||||
|
use crate::conversation::tool_result_serde::call_tool_result;
|
||||||
use crate::conversation::{debug_conversation_fix, fix_conversation, Conversation};
|
use crate::conversation::{debug_conversation_fix, fix_conversation, Conversation};
|
||||||
use crate::mcp_utils::ToolResult;
|
use crate::mcp_utils::ToolResult;
|
||||||
use crate::permission::permission_inspector::PermissionInspector;
|
use crate::permission::permission_inspector::PermissionInspector;
|
||||||
@@ -1123,6 +1124,8 @@ impl Agent {
|
|||||||
|
|
||||||
match item {
|
match item {
|
||||||
ToolStreamItem::Result(output) => {
|
ToolStreamItem::Result(output) => {
|
||||||
|
let output = call_tool_result::validate(output);
|
||||||
|
|
||||||
if enable_extension_request_ids.contains(&request_id)
|
if enable_extension_request_ids.contains(&request_id)
|
||||||
&& output.is_err()
|
&& output.is_err()
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ use thiserror::Error;
|
|||||||
use utoipa::ToSchema;
|
use utoipa::ToSchema;
|
||||||
|
|
||||||
pub mod message;
|
pub mod message;
|
||||||
mod tool_result_serde;
|
pub mod tool_result_serde;
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize, ToSchema, PartialEq)]
|
#[derive(Debug, Clone, Serialize, Deserialize, ToSchema, PartialEq)]
|
||||||
pub struct Conversation(Vec<Message>);
|
pub struct Conversation(Vec<Message>);
|
||||||
|
|||||||
@@ -145,7 +145,17 @@ pub mod call_tool_result {
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
let format = ResultFormat::deserialize(deserializer)?;
|
let original_value = serde_json::Value::deserialize(deserializer)?;
|
||||||
|
|
||||||
|
let format = ResultFormat::deserialize(&original_value).map_err(|e| {
|
||||||
|
tracing::debug!(
|
||||||
|
"Failed to deserialize call_tool_result: {}. Original data: {}",
|
||||||
|
e,
|
||||||
|
serde_json::to_string(&original_value)
|
||||||
|
.unwrap_or_else(|_| "<invalid json>".to_string())
|
||||||
|
);
|
||||||
|
serde::de::Error::custom(e)
|
||||||
|
})?;
|
||||||
|
|
||||||
match format {
|
match format {
|
||||||
ResultFormat::SuccessWithCallToolResult { status, value } => {
|
ResultFormat::SuccessWithCallToolResult { status, value } => {
|
||||||
@@ -184,4 +194,87 @@ pub mod call_tool_result {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub fn validate(result: ToolResult<CallToolResult>) -> ToolResult<CallToolResult> {
|
||||||
|
match &result {
|
||||||
|
Ok(call_tool_result) => match serde_json::to_string(call_tool_result) {
|
||||||
|
Ok(json_str) => match serde_json::from_str::<CallToolResult>(&json_str) {
|
||||||
|
Ok(_) => result,
|
||||||
|
Err(e) => {
|
||||||
|
tracing::error!("CallToolResult failed validation by deserialization: {}. Original data: {}", e, json_str);
|
||||||
|
Err(ErrorData {
|
||||||
|
code: ErrorCode::INTERNAL_ERROR,
|
||||||
|
message: Cow::from(format!("Tool result validation failed: {}", e)),
|
||||||
|
data: None,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
},
|
||||||
|
Err(e) => {
|
||||||
|
tracing::error!("CallToolResult failed serialization: {}", e);
|
||||||
|
Err(ErrorData {
|
||||||
|
code: ErrorCode::INTERNAL_ERROR,
|
||||||
|
message: Cow::from(format!("Tool result serialization failed: {}", e)),
|
||||||
|
data: None,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
},
|
||||||
|
Err(_) => result,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
use rmcp::model::{CallToolResult, Content, ErrorCode, ErrorData};
|
||||||
|
use std::borrow::Cow;
|
||||||
|
#[test]
|
||||||
|
fn test_validate_accepts_valid_call_tool_result() {
|
||||||
|
let valid_result = CallToolResult {
|
||||||
|
content: vec![Content::text("test")],
|
||||||
|
is_error: Some(false),
|
||||||
|
structured_content: None,
|
||||||
|
meta: None,
|
||||||
|
};
|
||||||
|
|
||||||
|
let tool_result: ToolResult<CallToolResult> = Ok(valid_result);
|
||||||
|
let validated = call_tool_result::validate(tool_result);
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
validated.is_ok(),
|
||||||
|
"Expected validation to pass for valid CallToolResult"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
#[test]
|
||||||
|
fn test_validate_returns_error_for_invalid_calltoolresult() {
|
||||||
|
let valid_result = CallToolResult {
|
||||||
|
content: vec![],
|
||||||
|
is_error: Some(false),
|
||||||
|
structured_content: None,
|
||||||
|
meta: None,
|
||||||
|
};
|
||||||
|
|
||||||
|
let tool_result: ToolResult<CallToolResult> = Ok(valid_result);
|
||||||
|
let validated = call_tool_result::validate(tool_result);
|
||||||
|
|
||||||
|
assert!(validated.is_err());
|
||||||
|
assert!(validated
|
||||||
|
.unwrap_err()
|
||||||
|
.message
|
||||||
|
.contains("Tool result validation failed"))
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_validate_passes_through_errors() {
|
||||||
|
let error_result: ToolResult<CallToolResult> = Err(ErrorData {
|
||||||
|
code: ErrorCode::INTERNAL_ERROR,
|
||||||
|
message: Cow::from("test error"),
|
||||||
|
data: None,
|
||||||
|
});
|
||||||
|
|
||||||
|
let validated = call_tool_result::validate(error_result.clone());
|
||||||
|
|
||||||
|
assert!(validated.is_err());
|
||||||
|
assert_eq!(validated.unwrap_err().message, "test error");
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user