refactor: Enforce strict typed JsonSchema for all MCP tools

This commit is contained in:
Riz Ashraf committed 2026-09-12 22:20:01 +01:00
1 parent 686fea683d
commit 2d3aaed289
31 files changed
+2892 -308

No files matched your search

+81 -59
View File
@@ -39,7 +39,8 @@ impl MemoryHandler {
}
"tools/list" => {
let tools = vec![
crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entities in the knowledge graph."),
crate::mcp::tool_def::<crate::tools::QueryGraphPathTool>("query_graph_path", "Traverse the knowledge graph to find a path between two entities."),
crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entities in the knowledge graph."),
crate::mcp::tool_def::<CreateRelationsTool>("create_relations", "Create new relations between entities in the knowledge graph."),
crate::mcp::tool_def::<AddObservationsTool>("add_observations", "Add new observations to existing entities in the knowledge graph."),
crate::mcp::tool_def::<DeleteEntitiesTool>("delete_entities", "Delete entities from the knowledge graph."),
@@ -120,7 +121,68 @@ impl MemoryHandler {
.unwrap_or(serde_json::Value::Object(Default::default()));
let result: Result<String, String> = match name {
"create_entities" => {
"query_graph_path" => {
let req: crate::tools::QueryGraphPathTool = match parse_args(args.clone()) {
Ok(r) => r,
Err(e) => return Some(crate::mcp::success(id.clone(), serde_json::json!({"isError": true, "content": [{"type": "text", "text": format!("Invalid args: {}", e)}] })))
};
let graph = self.state.get_full_graph();
let max_depth = req.max_depth.unwrap_or(5);
let mut queue = std::collections::VecDeque::new();
let mut visited = std::collections::HashSet::new();
let mut parents: std::collections::HashMap<String, (String, String)> = std::collections::HashMap::new();
queue.push_back(req.start_node.clone());
visited.insert(req.start_node.clone());
let mut found = false;
let mut current_depth = 0;
let mut nodes_at_current_depth = 1;
let mut nodes_at_next_depth = 0;
while let Some(current) = queue.pop_front() {
if current == req.end_node {
found = true;
break;
}
nodes_at_current_depth -= 1;
if current_depth < max_depth {
for rel in &graph.relations {
if rel.from == current && !visited.contains(&rel.to) {
visited.insert(rel.to.clone());
parents.insert(rel.to.clone(), (current.clone(), rel.relation_type.clone()));
queue.push_back(rel.to.clone());
nodes_at_next_depth += 1;
} else if rel.to == current && !visited.contains(&rel.from) {
visited.insert(rel.from.clone());
parents.insert(rel.from.clone(), (current.clone(), format!("inverse({})", rel.relation_type)));
queue.push_back(rel.from.clone());
nodes_at_next_depth += 1;
}
}
}
if nodes_at_current_depth == 0 {
current_depth += 1;
nodes_at_current_depth = nodes_at_next_depth;
nodes_at_next_depth = 0;
}
}
if found {
let mut path = Vec::new();
let mut curr = req.end_node.clone();
while curr != req.start_node {
let (parent, rel) = parents.get(&curr).unwrap();
path.push(format!("({}) --[{}]--> ({})", parent, rel, curr));
curr = parent.clone();
}
path.reverse();
Ok(format!("Path found:\n{}", path.join("\n")))
} else {
Ok(format!("No path found between {} and {} within depth {}", req.start_node, req.end_node, max_depth))
}
}
"create_entities" => {
let req: CreateEntitiesTool = match parse_args(args.clone()) {
Ok(r) => r,
Err(e) => {
@@ -131,19 +193,12 @@ impl MemoryHandler {
}
};
self.state.write_to_local_delta(|g| {
for e_val in req.entities {
if let Some(e) = if e_val.is_string() {
serde_json::from_str(e_val.as_str().unwrap()).ok()
} else {
serde_json::from_value(e_val).ok()
} {
let entity: Entity = e;
if !entity.name.is_empty() {
if let Ok(idx) = self.state.search_index.read() {
let _ = idx.index_entity(&entity);
}
g.entities.insert(entity.name.clone(), entity);
for entity in req.entities {
if !entity.name.is_empty() {
if let Ok(idx) = self.state.search_index.read() {
let _ = idx.index_entity(&entity);
}
g.entities.insert(entity.name.clone(), entity);
}
}
}).await;
@@ -160,16 +215,9 @@ impl MemoryHandler {
}
};
self.state.write_to_local_delta(|g| {
for r_val in req.relations {
if let Some(r) = if r_val.is_string() {
serde_json::from_str(r_val.as_str().unwrap()).ok()
} else {
serde_json::from_value(r_val).ok()
} {
let relation: Relation = r;
if !relation.from.is_empty() && !relation.to.is_empty() {
g.relations.push(relation);
}
for relation in req.relations {
if !relation.from.is_empty() && !relation.to.is_empty() {
g.relations.push(relation);
}
}
}).await;
@@ -185,20 +233,10 @@ impl MemoryHandler {
));
}
};
#[derive(Deserialize)]
struct ObsInput {
#[serde(rename = "entityName")]
entity_name: String,
contents: Vec<String>,
}
let full = self.state.get_full_graph();
self.state.write_to_local_delta(|g| {
for o_val in req.observations {
if let Some(o) = if o_val.is_string() {
serde_json::from_str::<ObsInput>(o_val.as_str().unwrap()).ok()
} else {
serde_json::from_value(o_val).ok()
} && let Some(full_e) = full.entities.get(&o.entity_name)
for o in req.observations {
if let Some(full_e) = full.entities.get(&o.entity_name)
{
let mut e =
g.entities.get(&o.entity_name).cloned().unwrap_or_else(
@@ -248,19 +286,9 @@ impl MemoryHandler {
));
}
};
#[derive(Deserialize)]
struct ObsDel {
#[serde(rename = "entityName")]
entity_name: String,
observations: Vec<String>,
}
self.state.apply_sync_write(|master| {
for d_val in req.deletions {
if let Some(d) = if d_val.is_string() {
serde_json::from_str::<ObsDel>(d_val.as_str().unwrap()).ok()
} else {
serde_json::from_value(d_val).ok()
} && let Some(e) = master.entities.get_mut(&d.entity_name)
for d in req.deletions {
if let Some(e) = master.entities.get_mut(&d.entity_name)
{
let to_rem: HashSet<_> = d.observations.into_iter().collect();
e.observations.retain(|o| !to_rem.contains(o));
@@ -281,17 +309,11 @@ impl MemoryHandler {
};
self.state.apply_sync_write(|master| {
let mut to_rem = HashSet::new();
for r_val in req.relations {
if let Some(r) = if r_val.is_string() {
serde_json::from_str::<Relation>(r_val.as_str().unwrap()).ok()
} else {
serde_json::from_value(r_val).ok()
} {
to_rem.insert(format!(
"{}|{}|{}|{}",
r.from, r.to, r.relation_type, r.namespace
));
}
for r in req.relations {
to_rem.insert(format!(
"{}|{}|{}|{}",
r.from, r.to, r.relation_type, r.namespace
));
}
master.relations.retain(|r| {
!to_rem.contains(&format!(