use std::collections::HashMap;
use rmcp::{
model::{CallToolRequestParams, CallToolResult, Tool},
service::ServerSink,
};
use serde::{Deserialize, Serialize};
use serde_json::Value;
use tracing::{debug, warn};
use crate::{
client::error::codes,
model::{Function, Tools},
};
#[inline]
pub fn mcp_tool_to_function(t: &Tool) -> Tools {
let desc = t.description.as_deref().unwrap_or("Remote MCP tool");
let schema = t.schema_as_json_value();
Tools::Function {
function: Function::new(t.name.to_string(), desc.to_string(), schema),
}
}
#[inline]
pub fn mcp_tools_to_functions(tools: &[Tool]) -> Vec<Tools> {
tools.iter().map(mcp_tool_to_function).collect()
}
#[inline]
pub fn call_tool_result_to_json(res: &CallToolResult) -> Value {
if res.is_error == Some(true) {
return serde_json::to_value(res).unwrap_or_else(|_| serialization_error_value());
}
if let Some(structured) = &res.structured_content {
return structured.clone();
}
serde_json::to_value(res).unwrap_or_else(|_| serialization_error_value())
}
fn serialization_error_value() -> Value {
serde_json::json!({
"error": {"type": "serialization_error", "message": "failed to serialize tool result"}
})
}
fn text_tool_message(
content: String,
id: Option<&str>,
) -> crate::model::chat_message_types::TextMessage {
match id {
Some(id) => crate::model::chat_message_types::TextMessage::tool_with_id(content, id),
None => crate::model::chat_message_types::TextMessage::tool(content),
}
}
#[derive(Clone, Serialize, Deserialize)]
pub struct McpCallSpec {
pub name: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub arguments: Option<Value>,
}
impl std::fmt::Debug for McpCallSpec {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("McpCallSpec")
.field("name", &self.name)
.field("arguments", &self.arguments.as_ref().map(|_| "<redacted>"))
.finish()
}
}
impl McpCallSpec {
pub fn new(name: impl Into<String>, arguments: Option<Value>) -> Self {
Self {
name: name.into(),
arguments,
}
}
pub fn validate(&self) -> crate::ZaiResult<()> {
crate::toolkits::core::validate_tool_name(&self.name)
.map_err(|error| validation_error(&error.to_string()))?;
if self
.arguments
.as_ref()
.is_some_and(|arguments| !arguments.is_object())
{
return Err(validation_error("arguments must be a JSON object"));
}
Ok(())
}
}
fn validation_error(message: &str) -> crate::client::error::ZaiError {
crate::client::error::ZaiError::ApiError {
code: codes::SDK_VALIDATION,
message: message.to_string(),
}
}
pub async fn call_mcp_tool(
server: &ServerSink,
name: impl Into<String>,
args: Option<Value>,
) -> crate::ZaiResult<(String, Value)> {
let spec = McpCallSpec::new(name, args);
spec.validate()?;
let McpCallSpec { name, arguments } = spec;
let arguments = match arguments {
Some(Value::Object(arguments)) => Some(arguments),
None => None,
Some(_) => return Err(validation_error("arguments must be a JSON object")),
};
let mut request = CallToolRequestParams::new(name.clone());
if let Some(arguments) = arguments {
request = request.with_arguments(arguments);
}
let res = server.call_tool(request).await.map_err(|_| {
warn!(code = codes::SDK_EXTERNAL_TOOL, "RMCP call_tool failed");
crate::client::error::ZaiError::ApiError {
code: codes::SDK_EXTERNAL_TOOL,
message: "RMCP service error".to_string(),
}
})?;
Ok((name, call_tool_result_to_json(&res)))
}
pub async fn call_mcp_tools_collect<I>(
server: &ServerSink,
calls: I,
) -> crate::ZaiResult<HashMap<String, Value>>
where
I: IntoIterator<Item = (String, Option<Value>)>,
{
use futures_util::{StreamExt, TryStreamExt};
futures_util::stream::iter(calls)
.map(|(name, arguments)| call_mcp_tool(server, name, arguments))
.buffered(8)
.try_collect::<HashMap<_, _>>()
.await
}
#[derive(Clone)]
pub struct McpToolCaller {
server: ServerSink,
}
impl McpToolCaller {
pub fn new(server: ServerSink) -> Self {
Self { server }
}
pub async fn call(
&self,
name: impl Into<String>,
args: Option<Value>,
) -> crate::ZaiResult<(String, Value)> {
call_mcp_tool(&self.server, name, args).await
}
pub async fn call_collect<I>(&self, calls: I) -> crate::ZaiResult<HashMap<String, Value>>
where
I: IntoIterator<Item = (String, Option<Value>)>,
{
call_mcp_tools_collect(&self.server, calls).await
}
}
pub async fn execute_tool_calls_as_messages(
caller: &McpToolCaller,
resp: &crate::model::chat_base_response::ChatCompletionResponse,
) -> crate::ZaiResult<Vec<crate::model::chat_message_types::TextMessage>> {
use crate::model::{chat_base_response::ToolCallMessage, chat_message_types::TextMessage};
let calls: Option<&[ToolCallMessage]> = resp
.choices()
.and_then(|v| v.first())
.and_then(|choice| choice.message())
.and_then(|message| message.tool_calls());
let Some(calls) = calls else {
return Ok(Vec::new());
};
debug!(tool_calls = calls.len(), "Dispatching tool calls");
async fn execute_one(
caller: &McpToolCaller,
call: &ToolCallMessage,
) -> crate::ZaiResult<TextMessage> {
let id = call.id();
let message = |error_type: &str, message: String| {
let payload = serde_json::json!({
"error": {"type": error_type, "message": message}
})
.to_string();
text_tool_message(payload, id)
};
let Some(function) = call.function() else {
return Ok(message(
"missing_function",
"tool_call.function is missing".to_string(),
));
};
let name = function.name();
if name.trim().is_empty() {
return Ok(message(
"missing_function_name",
"tool_call.function.name is blank".to_string(),
));
}
let arguments = match serde_json::from_str(function.arguments()) {
Ok(Value::Object(arguments)) => Some(Value::Object(arguments)),
Ok(_) => {
return Ok(message(
"invalid_arguments",
"tool arguments must decode to a JSON object".to_string(),
));
},
Err(error) => {
return Ok(message(
"invalid_arguments",
format!("tool arguments are not valid JSON: {error}"),
));
},
};
let (_, payload) = caller.call(name, arguments).await?;
Ok(text_tool_message(payload.to_string(), id))
}
use futures_util::{StreamExt, TryStreamExt};
futures_util::stream::iter(calls)
.map(|call| execute_one(caller, call))
.buffered(8)
.try_collect()
.await
}
fn assistant_request_message(
response: &crate::model::chat_base_response::ChatCompletionResponse,
) -> crate::ZaiResult<Option<crate::model::TextMessage>> {
use crate::model::{FunctionParams, TextMessage, ToolCall};
let Some(message) = response
.choices()
.and_then(|choices| choices.first())
.and_then(|choice| choice.message())
else {
return Ok(None);
};
let Some(calls) = message.tool_calls().filter(|calls| !calls.is_empty()) else {
return Ok(None);
};
let request_calls = calls
.iter()
.map(|call| {
let id = call
.id()
.filter(|id| !id.trim().is_empty())
.ok_or_else(|| validation_error("tool call id must not be blank"))?;
let function = call
.function()
.ok_or_else(|| validation_error("tool call function is required"))?;
crate::toolkits::core::validate_tool_name(function.name())
.map_err(|_| validation_error("tool call function name is invalid"))?;
if !matches!(
serde_json::from_str(function.arguments()),
Ok(Value::Object(_))
) {
return Err(validation_error(
"tool call arguments must decode to a JSON object",
));
}
Ok(ToolCall::new_function(
id,
FunctionParams::new(function.name(), function.arguments()),
))
})
.collect::<crate::ZaiResult<Vec<_>>>()?;
Ok(Some(TextMessage::assistant_with_tools(
message.content_str().map(str::to_owned),
request_calls,
)))
}
pub async fn run_mcp_tool_roundtrip<N>(
caller: &McpToolCaller,
client: &crate::client::ZaiClient,
mut chat: crate::model::chat::ChatCompletion<
N,
crate::model::chat_message_types::TextMessage,
crate::model::traits::StreamOff,
>,
system_hint_after_tools: Option<&str>,
) -> crate::ZaiResult<crate::model::chat_base_response::ChatCompletionResponse>
where
N: crate::model::traits::Chat
+ crate::model::traits::ChatToolSupport<Tool = crate::model::tools::Tools>
+ serde::Serialize,
(N, crate::model::chat_message_types::TextMessage): crate::model::traits::Bounded,
{
use crate::model::TextMessage;
let first_resp = chat.send_via(client).await?;
let Some(assistant_message) = assistant_request_message(&first_resp)? else {
return Ok(first_resp);
};
let tool_msgs: Vec<crate::model::chat_message_types::TextMessage> =
execute_tool_calls_as_messages(caller, &first_resp).await?;
chat = chat.add_message(assistant_message);
for m in tool_msgs {
chat = chat.add_message(m);
}
chat = chat.clear_tools();
if let Some(hint) = system_hint_after_tools {
chat = chat.add_message(TextMessage::system(hint));
}
let final_resp = chat.send_via(client).await?;
Ok(final_resp)
}
pub fn extract_final_text(
resp: &crate::model::chat_base_response::ChatCompletionResponse,
) -> Option<String> {
let msg = resp.choices()?.first()?.message()?;
match msg.content() {
Some(crate::model::chat_base_response::MessageContent::Text(text)) => Some(text.clone()),
Some(crate::model::chat_base_response::MessageContent::Parts(parts)) => parts
.iter()
.find(|part| {
matches!(
part.type_,
Some(crate::model::chat_base_response::MessageContentPartType::Text)
)
})
.and_then(|part| part.text.clone()),
_ => None,
}
}
#[cfg(test)]
mod tests {
use super::*;
fn response_with_arguments(
arguments: &str,
) -> crate::model::chat_base_response::ChatCompletionResponse {
serde_json::from_value(serde_json::json!({
"id": "response-id",
"choices": [{
"index": 0,
"message": {
"role": "assistant",
"tool_calls": [{
"id": "call-id",
"type": "function",
"function": {"name": "lookup", "arguments": arguments}
}]
},
"finish_reason": "tool_calls"
}]
}))
.unwrap()
}
#[test]
fn call_spec_debug_redacts_arguments() {
let spec = McpCallSpec::new("lookup", Some(serde_json::json!({"token": "secret-value"})));
let debug = format!("{spec:?}");
assert!(debug.contains("<redacted>"));
assert!(!debug.contains("secret-value"));
}
#[test]
fn roundtrip_replays_the_assistant_tool_request() {
let response = response_with_arguments(r#"{"query":"weather"}"#);
let message = assistant_request_message(&response).unwrap().unwrap();
let value = serde_json::to_value(message).unwrap();
assert_eq!(value["role"], "assistant");
assert_eq!(value["tool_calls"][0]["id"], "call-id");
assert_eq!(value["tool_calls"][0]["function"]["name"], "lookup");
}
#[test]
fn roundtrip_rejects_malformed_calls_before_dispatch() {
let response = response_with_arguments("[]");
assert!(assistant_request_message(&response).is_err());
}
}