use crate::internal_proxy::InternalProxy;
use crate::jrpc::{Error, Notification, Request, Response};
use serde::de::{MapAccess, Visitor};
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use std::collections::HashMap;
use std::fmt;
use std::sync::{LazyLock, RwLock};
pub trait Tool: Send + Sync {
fn name(&self) -> &str;
fn description(&self) -> &str;
fn input_schema(&self) -> InputSchema;
fn call(
&self,
params: HashMap<String, serde_json::Value>,
) -> Result<ToolCallResponse, ToolCallError>;
}
pub(crate) static TOOLS: LazyLock<RwLock<Vec<Box<dyn Tool>>>> =
LazyLock::new(|| RwLock::new(vec![]));
pub(crate) static SHARED_TOOLS: LazyLock<Vec<Box<dyn Tool>>> = LazyLock::new(|| {
vec![
Box::new(crate::mcp::latest_tools::LatestTools),
Box::new(crate::mcp::latest_tools::RunLatestTool),
]
});
#[derive(Debug, Clone, PartialEq, Eq, Hash, serde::Serialize, serde::Deserialize)]
pub struct ToolList {
pub(crate) tools: Vec<ToolInfo>,
}
impl ToolList {
pub fn empty() -> Self {
ToolList { tools: Vec::new() }
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, serde::Serialize, serde::Deserialize)]
pub(crate) struct ToolInfo {
name: String,
description: String,
#[serde(rename = "inputSchema")]
input_schema: InputSchema,
}
impl ToolInfo {
pub(crate) fn from_tool(tool: &dyn Tool) -> Self {
ToolInfo {
name: tool.name().to_string(),
description: tool.description().to_string(),
input_schema: tool.input_schema(),
}
}
}
#[derive(Debug, serde::Serialize, serde::Deserialize, Clone, PartialEq, Eq)]
pub struct InputSchema {
r#type: String,
properties: HashMap<String, HashMap<String, serde_json::Value>>,
required: Vec<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct Argument {
name: String,
r#type: String,
description: String,
required: bool,
}
impl Argument {
pub fn new(name: String, r#type: String, description: String, required: bool) -> Self {
Self {
name,
r#type,
description,
required,
}
}
}
impl InputSchema {
pub fn new<A: IntoIterator<Item = Argument>>(arguments: A) -> Self {
let mut properties = HashMap::new();
let mut required = Vec::new();
for argument in arguments {
let mut inner_map: HashMap<String, serde_json::Value> = HashMap::new();
inner_map.insert("type".to_string(), argument.r#type.into());
inner_map.insert("description".to_string(), argument.description.into());
if argument.required {
required.push(argument.name.clone());
}
properties.insert(argument.name, inner_map);
}
InputSchema {
r#type: "object".to_string(),
properties,
required,
}
}
}
pub(crate) fn list_int() -> ToolList {
let tool_infos: Vec<ToolInfo> = TOOLS
.read()
.unwrap()
.iter()
.chain(SHARED_TOOLS.iter())
.map(|tool| ToolInfo::from_tool(tool.as_ref()))
.collect();
ToolList { tools: tool_infos }
}
pub(crate) fn list_process(request: Request) -> Response<ToolList> {
let tool_list = list_int();
Response::new(tool_list, request.id)
}
pub fn add_tool(tool: Box<dyn Tool>) {
TOOLS.write().unwrap().push(tool);
let n = Notification::new("notifications/tools/list_changed".to_string(), None);
let r = InternalProxy::current().send_notification(n);
match r {
Ok(_) => {}
Err(crate::internal_proxy::Error::NotConnected) => {
}
}
}
#[derive(Debug, serde::Deserialize, Clone)]
pub(crate) struct ToolCallParams {
pub(crate) name: String,
pub(crate) arguments: HashMap<String, serde_json::Value>,
}
impl ToolCallParams {
pub(crate) fn new(name: String, arguments: HashMap<String, serde_json::Value>) -> Self {
ToolCallParams { name, arguments }
}
}
#[derive(Debug, serde::Serialize, serde::Deserialize, Default)]
pub struct ToolCallResponse {
pub(crate) content: Vec<ToolContent>,
is_error: bool,
}
impl ToolCallResponse {
pub fn new(content: Vec<ToolContent>) -> Self {
ToolCallResponse {
content,
is_error: false,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, serde::Serialize, thiserror::Error)]
#[error("Tool call failed: {}", format_content(&self.content))]
pub struct ToolCallError {
content: Vec<ToolContent>,
is_error: bool,
}
impl ToolCallError {
pub fn new(content: Vec<ToolContent>) -> Self {
ToolCallError {
content,
is_error: true,
}
}
pub(crate) fn into_response(self) -> ToolCallResponse {
ToolCallResponse {
content: self.content,
is_error: true,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub enum ToolContent {
Text(String),
}
impl ToolContent {
#[cfg(feature = "transit")]
pub(crate) fn as_str(&self) -> Option<&str> {
match self {
ToolContent::Text(text) => Some(text),
}
}
}
impl Serialize for ToolContent {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
use serde::ser::SerializeStruct;
match self {
ToolContent::Text(text) => {
let mut s = serializer.serialize_struct("ToolContent", 2)?;
s.serialize_field("type", "text")?;
s.serialize_field("text", text)?;
s.end()
}
}
}
}
impl<'de> Deserialize<'de> for ToolContent {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
use serde::de;
struct ToolContentVisitor;
impl<'de> Visitor<'de> for ToolContentVisitor {
type Value = ToolContent;
fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
formatter.write_str("a tool content object with type and data")
}
fn visit_map<M>(self, mut map: M) -> Result<Self::Value, M::Error>
where
M: MapAccess<'de>,
{
let mut content_type: Option<String> = None;
let mut text: Option<String> = None;
while let Some(key) = map.next_key::<String>()? {
match key.as_str() {
"type" => {
if content_type.is_some() {
return Err(de::Error::duplicate_field("type"));
}
content_type = Some(map.next_value()?);
}
"text" => {
if text.is_some() {
return Err(de::Error::duplicate_field("text"));
}
text = Some(map.next_value()?);
}
_ => {
let _: de::IgnoredAny = map.next_value()?;
}
}
}
match content_type.as_deref() {
Some("text") => {
let text = text.ok_or_else(|| de::Error::missing_field("text"))?;
Ok(ToolContent::Text(text))
}
Some(other) => Err(de::Error::unknown_variant(other, &["text"])),
None => Err(de::Error::missing_field("type")),
}
}
}
deserializer.deserialize_map(ToolContentVisitor)
}
}
impl From<String> for ToolContent {
fn from(value: String) -> Self {
ToolContent::Text(value)
}
}
impl From<&str> for ToolContent {
fn from(value: &str) -> Self {
ToolContent::Text(value.to_string())
}
}
pub(crate) fn call_imp(params: ToolCallParams) -> Result<ToolCallResponse, crate::jrpc::Error> {
let tools = TOOLS.read().unwrap();
let tool = tools
.iter()
.chain(SHARED_TOOLS.iter())
.find(|t| t.name() == params.name)
.map(|t| t.as_ref());
match tool {
Some(tool) => {
let call = tool.call(params.arguments);
match call {
Ok(response) => Ok(response),
Err(err) => Ok(err.into_response()),
}
}
None => Err(Error::unknown_tool(params.name)),
}
}
pub(crate) fn call(request: Request) -> Response<ToolCallResponse> {
let params = match request.params {
Some(params) => match serde_json::from_value::<ToolCallParams>(params) {
Ok(params) => params,
Err(err) => return Response::err(Error::invalid_params(err.to_string()), request.id),
},
None => {
return Response::err(
Error::invalid_params("No parameters provided".to_string()),
request.id,
);
}
};
let r = call_imp(params);
match r {
Ok(r) => Response::new(r, request.id),
Err(e) => Response::err(e, request.id),
}
}
fn format_content(content: &[ToolContent]) -> String {
content
.iter()
.map(|c| match c {
ToolContent::Text(text) => text.clone(),
})
.collect::<Vec<_>>()
.join("; ")
}
impl Clone for ToolCallResponse {
fn clone(&self) -> Self {
Self {
content: self.content.clone(),
is_error: self.is_error,
}
}
}
impl PartialEq for ToolCallResponse {
fn eq(&self, other: &Self) -> bool {
self.content == other.content && self.is_error == other.is_error
}
}
impl Eq for ToolCallResponse {}
impl std::hash::Hash for ToolCallResponse {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
self.content.hash(state);
self.is_error.hash(state);
}
}
impl fmt::Display for ToolCallResponse {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
if self.is_error {
write!(
f,
"ToolCallResponse(Error): {}",
format_content(&self.content)
)
} else {
write!(
f,
"ToolCallResponse(Success): {}",
format_content(&self.content)
)
}
}
}
impl From<Vec<ToolContent>> for ToolCallResponse {
fn from(content: Vec<ToolContent>) -> Self {
Self::new(content)
}
}
impl From<String> for ToolCallResponse {
fn from(message: String) -> Self {
Self::new(vec![message.into()])
}
}
impl From<&str> for ToolCallResponse {
fn from(message: &str) -> Self {
Self::new(vec![message.into()])
}
}
impl From<String> for ToolCallError {
fn from(message: String) -> Self {
ToolCallError::new(vec![message.into()])
}
}
impl From<&str> for ToolCallError {
fn from(message: &str) -> Self {
ToolCallError::new(vec![message.into()])
}
}
impl From<Vec<String>> for ToolCallError {
fn from(messages: Vec<String>) -> Self {
ToolCallError::new(messages.into_iter().map(|m| m.into()).collect())
}
}
impl From<ToolContent> for ToolCallError {
fn from(content: ToolContent) -> Self {
ToolCallError::new(vec![content])
}
}
impl std::hash::Hash for InputSchema {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
self.r#type.hash(state);
let mut properties_vec: Vec<_> = self.properties.iter().collect();
properties_vec.sort_by_key(|(k, _)| *k);
for (k, v) in properties_vec {
k.hash(state);
format!("{:?}", v).hash(state);
}
self.required.hash(state);
}
}
impl Default for InputSchema {
fn default() -> Self {
Self {
r#type: "object".to_string(),
properties: HashMap::new(),
required: Vec::new(),
}
}
}
impl fmt::Display for InputSchema {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"InputSchema(type: {}, properties: {}, required: {:?})",
self.r#type,
self.properties.len(),
self.required
)
}
}
impl fmt::Display for Argument {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"{}: {} ({}) - {}",
self.name,
self.r#type,
if self.required {
"required"
} else {
"optional"
},
self.description
)
}
}
impl Default for ToolContent {
fn default() -> Self {
ToolContent::Text(String::new())
}
}
impl fmt::Display for ToolContent {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
ToolContent::Text(text) => write!(f, "{}", text),
}
}
}
impl AsRef<str> for ToolContent {
fn as_ref(&self) -> &str {
match self {
ToolContent::Text(text) => text,
}
}
}
impl std::ops::Deref for ToolContent {
type Target = str;
fn deref(&self) -> &Self::Target {
match self {
ToolContent::Text(text) => text,
}
}
}
impl Default for ToolList {
fn default() -> Self {
Self::empty()
}
}
impl fmt::Display for ToolList {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "ToolList({} tools)", self.tools.len())
}
}
impl From<Vec<ToolInfo>> for ToolList {
fn from(tools: Vec<ToolInfo>) -> Self {
ToolList { tools }
}
}
impl From<ToolList> for Vec<ToolInfo> {
fn from(tool_list: ToolList) -> Self {
tool_list.tools
}
}
impl AsRef<Vec<ToolInfo>> for ToolList {
fn as_ref(&self) -> &Vec<ToolInfo> {
&self.tools
}
}