refactor: Enforce strict typed JsonSchema for all MCP tools
This commit is contained in:
1 parent
686fea683d
commit
2d3aaed289
31 files changed
+2892
-308
No files matched your search
+81
-59
@@ -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!(
|
||||
|
||||
Reference in new issue
Block a user