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
@@ -0,0 +1,92 @@
|
||||
import re
|
||||
import sys
|
||||
|
||||
handlers_path = r"C:\Users\reazul.ashraf\workspace\rust\mcp-memory\server\src\handlers.rs"
|
||||
|
||||
with open(handlers_path, "r", encoding="utf-8") as f:
|
||||
content = f.read()
|
||||
|
||||
# Add to list_tools
|
||||
list_tools_idx = content.find('crate::mcp::tool_def::<CreateEntitiesTool>')
|
||||
if list_tools_idx == -1:
|
||||
print("Could not find list_tools")
|
||||
sys.exit(1)
|
||||
|
||||
tool_def = ' crate::mcp::tool_def::<crate::tools::QueryGraphPathTool>("query_graph_path", "Traverse the knowledge graph to find a path between two entities."),'
|
||||
content = content[:list_tools_idx] + tool_def + "\n" + content[list_tools_idx:]
|
||||
|
||||
# Add to match name
|
||||
match_idx = content.find('"create_entities" => {')
|
||||
if match_idx == -1:
|
||||
print("Could not find match name")
|
||||
sys.exit(1)
|
||||
|
||||
handler_code = """ "query_graph_path" => {
|
||||
let req: crate::tools::QueryGraphPathTool = match parse_args(args.clone()) {
|
||||
Ok(r) => r,
|
||||
Err(e) => return Some(crate::mcp::success(req_id, true, &format!("Invalid args: {}", e))),
|
||||
};
|
||||
let graph = state.graph.read();
|
||||
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))
|
||||
}
|
||||
}
|
||||
"""
|
||||
|
||||
content = content[:match_idx] + handler_code + content[match_idx:]
|
||||
|
||||
with open(handlers_path, "w", encoding="utf-8") as f:
|
||||
f.write(content)
|
||||
|
||||
print("Injected query_graph_path tool!")
|
||||
Reference in new issue
Block a user