1use crate::db::AgentDB;
26use serde_json::{json, Value};
27use std::collections::HashMap;
28
29pub struct McpServer {
31 db: AgentDB,
32}
33
34impl McpServer {
35 pub fn new(db: AgentDB) -> Self {
37 Self { db }
38 }
39
40 pub fn handle_message(&self, input: &str) -> String {
42 let req: Value = match serde_json::from_str(input) {
43 Ok(v) => v,
44 Err(e) => {
45 return json!({
46 "jsonrpc": "2.0",
47 "id": null,
48 "error": { "code": -32700, "message": format!("Parse error: {e}") }
49 })
50 .to_string();
51 }
52 };
53
54 let id = req.get("id").cloned().unwrap_or(Value::Null);
55 let method = req.get("method").and_then(|m| m.as_str()).unwrap_or("");
56 let params = req.get("params").cloned().unwrap_or(Value::Object(Default::default()));
57
58 let result = match method {
59 "initialize" => self.handle_initialize(¶ms),
60 "initialized" => return String::new(),
61 "tools/list" => self.handle_tools_list(),
62 "tools/call" => self.handle_tools_call(¶ms),
63 "resources/list" => self.handle_resources_list(),
64 "resources/read" => self.handle_resources_read(¶ms),
65 _ => Err((-32601, format!("Method not found: {method}"))),
66 };
67
68 match result {
69 Ok(value) => json!({ "jsonrpc": "2.0", "id": id, "result": value }).to_string(),
70 Err((code, msg)) => {
71 json!({ "jsonrpc": "2.0", "id": id, "error": { "code": code, "message": msg } })
72 .to_string()
73 }
74 }
75 }
76
77 fn handle_initialize(&self, _params: &Value) -> std::result::Result<Value, (i32, String)> {
78 Ok(json!({
79 "protocolVersion": "2024-11-05",
80 "capabilities": {
81 "tools": { "listChanged": false },
82 "resources": { "subscribe": false, "listChanged": false }
83 },
84 "serverInfo": {
85 "name": "agentdb",
86 "version": env!("CARGO_PKG_VERSION")
87 }
88 }))
89 }
90
91 fn handle_tools_list(&self) -> std::result::Result<Value, (i32, String)> {
92 Ok(json!({ "tools": self.tool_definitions() }))
93 }
94
95 fn handle_tools_call(&self, params: &Value) -> std::result::Result<Value, (i32, String)> {
96 let name = params
97 .get("name")
98 .and_then(|n| n.as_str())
99 .ok_or((-32602, "Missing 'name' parameter".to_string()))?;
100 let arguments = params
101 .get("arguments")
102 .cloned()
103 .unwrap_or(Value::Object(Default::default()));
104
105 let result = self.dispatch_tool(name, &arguments)?;
106
107 Ok(json!({
108 "content": [{
109 "type": "text",
110 "text": result.to_string()
111 }]
112 }))
113 }
114
115 fn handle_resources_list(&self) -> std::result::Result<Value, (i32, String)> {
116 Ok(json!({
117 "resources": [
118 {
119 "uri": "agentdb://stats",
120 "name": "Database Statistics",
121 "description": "Current AgentDB database statistics",
122 "mimeType": "application/json"
123 }
124 ]
125 }))
126 }
127
128 fn handle_resources_read(&self, params: &Value) -> std::result::Result<Value, (i32, String)> {
129 let uri = params
130 .get("uri")
131 .and_then(|u| u.as_str())
132 .ok_or((-32602, "Missing 'uri' parameter".to_string()))?;
133
134 match uri {
135 "agentdb://stats" => {
136 let stats = self
137 .db
138 .stats()
139 .map_err(|e| (-32000, format!("Stats error: {e}")))?;
140 Ok(json!({
141 "contents": [{
142 "uri": "agentdb://stats",
143 "mimeType": "application/json",
144 "text": serde_json::to_string(&stats).unwrap_or_default()
145 }]
146 }))
147 }
148 _ => Err((-32002, format!("Resource not found: {uri}"))),
149 }
150 }
151
152 fn dispatch_tool(
153 &self,
154 name: &str,
155 args: &Value,
156 ) -> std::result::Result<Value, (i32, String)> {
157 let err = |e: crate::error::AgentDbError| (-32000, e.to_string());
158
159 match name {
160 "execute" => {
161 let sql = get_str(args, "sql")?;
162 let n = self.db.execute(sql).map_err(err)?;
163 Ok(json!({ "rows_affected": n }))
164 }
165 "query" => {
166 let sql = get_str(args, "sql")?;
167 let rows = self.db.query_json(sql).map_err(err)?;
168 Ok(Value::Array(rows))
169 }
170 "vector_upsert" => {
171 let collection = get_str(args, "collection")?;
172 let id = get_str(args, "id")?;
173 let vector: Vec<f32> = args
174 .get("vector")
175 .and_then(|v| serde_json::from_value(v.clone()).ok())
176 .ok_or((-32602, "Missing 'vector' array".to_string()))?;
177 let metadata = args.get("metadata").cloned();
178 let dim = vector.len();
179 let col = self.db.vectors().collection(collection, dim).map_err(err)?;
180 col.upsert(crate::vectors::VectorEntry {
181 id: id.to_string(),
182 vector,
183 metadata,
184 })
185 .map_err(err)?;
186 Ok(json!({ "ok": true }))
187 }
188 "vector_search" => {
189 let collection = get_str(args, "collection")?;
190 let query: Vec<f32> = args
191 .get("query")
192 .and_then(|v| serde_json::from_value(v.clone()).ok())
193 .ok_or((-32602, "Missing 'query' array".to_string()))?;
194 let top_k = args
195 .get("top_k")
196 .and_then(|v| v.as_u64())
197 .unwrap_or(10) as usize;
198 let filter = args.get("filter").cloned();
199 let dim = query.len();
200 let col = self.db.vectors().collection(collection, dim).map_err(err)?;
201 let results = col
202 .search(
203 &query,
204 crate::vectors::SearchOptions {
205 top_k,
206 metric: crate::vectors::DistanceMetric::Cosine,
207 filter,
208 },
209 )
210 .map_err(err)?;
211 let arr: Vec<Value> = results
212 .iter()
213 .map(|r| json!({"id": r.id, "score": r.score, "metadata": r.metadata}))
214 .collect();
215 Ok(Value::Array(arr))
216 }
217 "graph_add_node" => {
218 let id = get_str(args, "id")?;
219 let kind = get_str(args, "kind")?;
220 let data = args.get("data").cloned();
221 self.db.memory().add_node(id, kind, data).map_err(err)?;
222 Ok(json!({ "ok": true }))
223 }
224 "graph_add_edge" => {
225 let src = get_str(args, "src")?;
226 let dst = get_str(args, "dst")?;
227 let relation = get_str(args, "relation")?;
228 let weight = args.get("weight").and_then(|v| v.as_f64()).unwrap_or(1.0);
229 self.db
230 .memory()
231 .add_edge(src, dst, relation, weight)
232 .map_err(err)?;
233 Ok(json!({ "ok": true }))
234 }
235 "graph_neighbors" => {
236 let node_id = get_str(args, "node_id")?;
237 let max_depth = args.get("max_depth").and_then(|v| v.as_u64()).unwrap_or(2) as usize;
238 let min_weight = args.get("min_weight").and_then(|v| v.as_f64()).unwrap_or(0.0);
239 let relation = args.get("relation").and_then(|v| v.as_str());
240 let results = self
241 .db
242 .memory()
243 .neighbors(
244 node_id,
245 crate::memory::TraversalOptions {
246 max_depth,
247 min_weight: Some(min_weight),
248 relation: relation.map(|s| s.to_string()),
249 },
250 )
251 .map_err(err)?;
252 let arr: Vec<Value> = results
253 .iter()
254 .map(|r| {
255 json!({"id": r.node.id, "kind": r.node.kind, "depth": r.depth, "weight": r.weight, "data": r.node.data})
256 })
257 .collect();
258 Ok(Value::Array(arr))
259 }
260 "tool_register" => {
261 let tool_name = get_str(args, "name")?;
262 let description = args.get("description").and_then(|v| v.as_str());
263 let schema = args.get("parameters_schema").cloned();
264 let version = args.get("version").and_then(|v| v.as_str());
265 let id = self
266 .db
267 .tools()
268 .register_tool(tool_name, description, schema, version)
269 .map_err(err)?;
270 Ok(json!({ "id": id }))
271 }
272 "tool_list" => {
273 let tools = self.db.tools().list_tools().map_err(err)?;
274 let arr: Vec<Value> = tools
275 .iter()
276 .map(|t| {
277 json!({
278 "id": t.id, "name": t.name,
279 "description": t.description,
280 "parameters_schema": t.parameters_schema,
281 "version": t.version
282 })
283 })
284 .collect();
285 Ok(Value::Array(arr))
286 }
287 "tool_log_call" => {
288 let tool_name = get_str(args, "tool_name")?;
289 let session_id = args.get("session_id").and_then(|v| v.as_str());
290 let arguments = args.get("arguments").cloned();
291 let result = args.get("result").cloned();
292 let error = args.get("error").and_then(|v| v.as_str());
293 let latency_ms = args.get("latency_ms").and_then(|v| v.as_i64()).unwrap_or(0);
294 let id = self
295 .db
296 .tools()
297 .log_tool_call(session_id, tool_name, arguments, result, error, Some(latency_ms))
298 .map_err(err)?;
299 Ok(json!({ "id": id }))
300 }
301 "audit_log" => {
302 let action = get_str(args, "action")?;
303 let table_name = get_str(args, "table_name")?;
304 let record_id = get_str(args, "record_id")?;
305 let actor = args.get("actor").and_then(|v| v.as_str());
306 let old_value = args.get("old_value").cloned();
307 let new_value = args.get("new_value").cloned();
308 let reason = args.get("reason").and_then(|v| v.as_str());
309 let id = self
310 .db
311 .audit()
312 .log(actor, action, table_name, record_id, old_value, new_value, reason)
313 .map_err(err)?;
314 Ok(json!({ "id": id }))
315 }
316 "audit_query_recent" => {
317 let limit = args.get("limit").and_then(|v| v.as_u64()).unwrap_or(100) as usize;
318 let entries = self.db.audit().query_recent(Some(limit)).map_err(err)?;
319 let arr: Vec<Value> = entries
320 .iter()
321 .map(|e| {
322 json!({
323 "id": e.id, "timestamp": e.timestamp, "actor": e.actor,
324 "action": e.action, "table_name": e.table_name,
325 "record_id": e.record_id, "reason": e.reason
326 })
327 })
328 .collect();
329 Ok(Value::Array(arr))
330 }
331 "context_add" => {
332 let session_id = get_str(args, "session_id")?;
333 let source_type = get_str(args, "source_type")?;
334 let source_id = get_str(args, "source_id")?;
335 let content_preview = args.get("content_preview").and_then(|v| v.as_str());
336 let token_count = args
337 .get("token_count")
338 .and_then(|v| v.as_i64())
339 .ok_or((-32602, "Missing 'token_count'".to_string()))?;
340 let relevance_score = args
341 .get("relevance_score")
342 .and_then(|v| v.as_f64())
343 .unwrap_or(0.5);
344 let priority = args.get("priority").and_then(|v| v.as_i64()).unwrap_or(0);
345 let id = self
346 .db
347 .context()
348 .add_entry(
349 session_id,
350 source_type,
351 source_id,
352 content_preview,
353 token_count,
354 relevance_score,
355 priority,
356 )
357 .map_err(err)?;
358 Ok(json!({ "id": id }))
359 }
360 "context_build_window" => {
361 let session_id = get_str(args, "session_id")?;
362 let max_tokens = args
363 .get("max_tokens")
364 .and_then(|v| v.as_i64())
365 .ok_or((-32602, "Missing 'max_tokens'".to_string()))?;
366 let entries = self
367 .db
368 .context()
369 .build_window(session_id, max_tokens)
370 .map_err(err)?;
371 let arr: Vec<Value> = entries
372 .iter()
373 .map(|e| {
374 json!({
375 "id": e.id, "source_type": e.source_type,
376 "source_id": e.source_id, "content_preview": e.content_preview,
377 "token_count": e.token_count, "priority": e.priority
378 })
379 })
380 .collect();
381 Ok(Value::Array(arr))
382 }
383 "context_clear" => {
384 let session_id = get_str(args, "session_id")?;
385 self.db.context().clear_session(session_id).map_err(err)?;
386 Ok(json!({ "ok": true }))
387 }
388 "prompt_create" => {
389 let name = get_str(args, "name")?;
390 let template = get_str(args, "template")?;
391 let model_hint = args.get("model_hint").and_then(|v| v.as_str());
392 let max_tokens = args.get("max_tokens").and_then(|v| v.as_i64());
393 let metadata = args.get("metadata").cloned();
394 let id = self
395 .db
396 .prompts()
397 .create_template(name, template, model_hint, max_tokens, metadata)
398 .map_err(err)?;
399 Ok(json!({ "id": id }))
400 }
401 "prompt_render" => {
402 let name = get_str(args, "name")?;
403 let vars: HashMap<String, String> = args
404 .get("vars")
405 .and_then(|v| serde_json::from_value(v.clone()).ok())
406 .unwrap_or_default();
407 let rendered = self.db.prompts().render(name, &vars).map_err(err)?;
408 Ok(json!({ "text": rendered }))
409 }
410 "label_tag" => {
411 let table_name = get_str(args, "table_name")?;
412 let record_id = get_str(args, "record_id")?;
413 let label = get_str(args, "label")?;
414 let tagged_by = args.get("tagged_by").and_then(|v| v.as_str());
415 self.db
416 .labels()
417 .tag(table_name, record_id, label, tagged_by)
418 .map_err(err)?;
419 Ok(json!({ "ok": true }))
420 }
421 "label_untag" => {
422 let table_name = get_str(args, "table_name")?;
423 let record_id = get_str(args, "record_id")?;
424 let label = get_str(args, "label")?;
425 self.db
426 .labels()
427 .untag(table_name, record_id, label)
428 .map_err(err)?;
429 Ok(json!({ "ok": true }))
430 }
431 "label_get" => {
432 let table_name = get_str(args, "table_name")?;
433 let record_id = get_str(args, "record_id")?;
434 let labels = self
435 .db
436 .labels()
437 .get_labels(table_name, record_id)
438 .map_err(err)?;
439 let arr: Vec<Value> = labels
440 .iter()
441 .map(|l| {
442 json!({
443 "label": l.label, "tagged_by": l.tagged_by,
444 "tagged_at": l.tagged_at
445 })
446 })
447 .collect();
448 Ok(Value::Array(arr))
449 }
450 "label_has" => {
451 let table_name = get_str(args, "table_name")?;
452 let record_id = get_str(args, "record_id")?;
453 let label = get_str(args, "label")?;
454 let has = self
455 .db
456 .labels()
457 .has_label(table_name, record_id, label)
458 .map_err(err)?;
459 Ok(json!({ "has": has }))
460 }
461 "stats" => {
462 let stats = self.db.stats().map_err(err)?;
463 Ok(json!({
464 "collections": stats.collections,
465 "vectors": stats.vectors,
466 "nodes": stats.nodes,
467 "edges": stats.edges,
468 "conversations": stats.conversations,
469 "messages": stats.messages,
470 "workflows": stats.workflows,
471 "workflow_steps": stats.workflow_steps,
472 "traces": stats.traces,
473 "tools": stats.tools,
474 "tool_calls": stats.tool_calls,
475 "audit_entries": stats.audit_entries,
476 "prompt_templates": stats.prompt_templates
477 }))
478 }
479 _ => Err((-32601, format!("Unknown tool: {name}"))),
480 }
481 }
482
483 fn tool_definitions(&self) -> Value {
484 json!([
485 tool_def("execute", "Execute a raw SQL statement (DDL/DML)", json!({
486 "type": "object",
487 "properties": { "sql": { "type": "string", "description": "SQL statement" } },
488 "required": ["sql"]
489 })),
490 tool_def("query", "Execute a SELECT and return rows as JSON", json!({
491 "type": "object",
492 "properties": { "sql": { "type": "string", "description": "SELECT statement" } },
493 "required": ["sql"]
494 })),
495 tool_def("vector_upsert", "Insert or update a vector embedding", json!({
496 "type": "object",
497 "properties": {
498 "collection": { "type": "string" },
499 "id": { "type": "string" },
500 "vector": { "type": "array", "items": { "type": "number" } },
501 "metadata": { "type": "object" }
502 },
503 "required": ["collection", "id", "vector"]
504 })),
505 tool_def("vector_search", "Approximate nearest-neighbor search", json!({
506 "type": "object",
507 "properties": {
508 "collection": { "type": "string" },
509 "query": { "type": "array", "items": { "type": "number" } },
510 "top_k": { "type": "integer", "default": 10 },
511 "filter": { "type": "object" }
512 },
513 "required": ["collection", "query"]
514 })),
515 tool_def("graph_add_node", "Add or update a memory graph node", json!({
516 "type": "object",
517 "properties": {
518 "id": { "type": "string" },
519 "kind": { "type": "string" },
520 "data": { "type": "object" }
521 },
522 "required": ["id", "kind"]
523 })),
524 tool_def("graph_add_edge", "Add or update a directed graph edge", json!({
525 "type": "object",
526 "properties": {
527 "src": { "type": "string" },
528 "dst": { "type": "string" },
529 "relation": { "type": "string" },
530 "weight": { "type": "number", "default": 1.0 }
531 },
532 "required": ["src", "dst", "relation"]
533 })),
534 tool_def("graph_neighbors", "Traverse the memory graph from a node", json!({
535 "type": "object",
536 "properties": {
537 "node_id": { "type": "string" },
538 "max_depth": { "type": "integer", "default": 2 },
539 "min_weight": { "type": "number", "default": 0.0 },
540 "relation": { "type": "string" }
541 },
542 "required": ["node_id"]
543 })),
544 tool_def("tool_register", "Register or update a tool definition", json!({
545 "type": "object",
546 "properties": {
547 "name": { "type": "string" },
548 "description": { "type": "string" },
549 "parameters_schema": { "type": "object" },
550 "version": { "type": "string" }
551 },
552 "required": ["name"]
553 })),
554 tool_def("tool_list", "List all registered tools", json!({
555 "type": "object", "properties": {}
556 })),
557 tool_def("tool_log_call", "Log a tool invocation", json!({
558 "type": "object",
559 "properties": {
560 "tool_name": { "type": "string" },
561 "session_id": { "type": "string" },
562 "arguments": { "type": "object" },
563 "result": { "type": "object" },
564 "error": { "type": "string" },
565 "latency_ms": { "type": "integer" }
566 },
567 "required": ["tool_name"]
568 })),
569 tool_def("audit_log", "Append an entry to the audit log", json!({
570 "type": "object",
571 "properties": {
572 "action": { "type": "string" },
573 "table_name": { "type": "string" },
574 "record_id": { "type": "string" },
575 "actor": { "type": "string" },
576 "old_value": { "type": "object" },
577 "new_value": { "type": "object" },
578 "reason": { "type": "string" }
579 },
580 "required": ["action", "table_name", "record_id"]
581 })),
582 tool_def("audit_query_recent", "Query recent audit log entries", json!({
583 "type": "object",
584 "properties": { "limit": { "type": "integer", "default": 100 } }
585 })),
586 tool_def("context_add", "Add an entry to the context window", json!({
587 "type": "object",
588 "properties": {
589 "session_id": { "type": "string" },
590 "source_type": { "type": "string" },
591 "source_id": { "type": "string" },
592 "content_preview": { "type": "string" },
593 "token_count": { "type": "integer" },
594 "relevance_score": { "type": "number" },
595 "priority": { "type": "integer" }
596 },
597 "required": ["session_id", "source_type", "source_id", "token_count"]
598 })),
599 tool_def("context_build_window", "Build a token-budgeted context window", json!({
600 "type": "object",
601 "properties": {
602 "session_id": { "type": "string" },
603 "max_tokens": { "type": "integer" }
604 },
605 "required": ["session_id", "max_tokens"]
606 })),
607 tool_def("context_clear", "Clear all context entries for a session", json!({
608 "type": "object",
609 "properties": { "session_id": { "type": "string" } },
610 "required": ["session_id"]
611 })),
612 tool_def("prompt_create", "Create a new prompt template version", json!({
613 "type": "object",
614 "properties": {
615 "name": { "type": "string" },
616 "template": { "type": "string" },
617 "model_hint": { "type": "string" },
618 "max_tokens": { "type": "integer" },
619 "metadata": { "type": "object" }
620 },
621 "required": ["name", "template"]
622 })),
623 tool_def("prompt_render", "Render a prompt template with variables", json!({
624 "type": "object",
625 "properties": {
626 "name": { "type": "string" },
627 "vars": { "type": "object", "additionalProperties": { "type": "string" } }
628 },
629 "required": ["name"]
630 })),
631 tool_def("label_tag", "Tag a record with a classification label", json!({
632 "type": "object",
633 "properties": {
634 "table_name": { "type": "string" },
635 "record_id": { "type": "string" },
636 "label": { "type": "string" },
637 "tagged_by": { "type": "string" }
638 },
639 "required": ["table_name", "record_id", "label"]
640 })),
641 tool_def("label_untag", "Remove a label from a record", json!({
642 "type": "object",
643 "properties": {
644 "table_name": { "type": "string" },
645 "record_id": { "type": "string" },
646 "label": { "type": "string" }
647 },
648 "required": ["table_name", "record_id", "label"]
649 })),
650 tool_def("label_get", "Get all labels for a record", json!({
651 "type": "object",
652 "properties": {
653 "table_name": { "type": "string" },
654 "record_id": { "type": "string" }
655 },
656 "required": ["table_name", "record_id"]
657 })),
658 tool_def("label_has", "Check if a record has a specific label", json!({
659 "type": "object",
660 "properties": {
661 "table_name": { "type": "string" },
662 "record_id": { "type": "string" },
663 "label": { "type": "string" }
664 },
665 "required": ["table_name", "record_id", "label"]
666 })),
667 tool_def("stats", "Get database-wide statistics", json!({
668 "type": "object", "properties": {}
669 })),
670 ])
671 }
672}
673
674fn tool_def(name: &str, description: &str, input_schema: Value) -> Value {
675 json!({
676 "name": name,
677 "description": description,
678 "inputSchema": input_schema
679 })
680}
681
682fn get_str<'a>(args: &'a Value, key: &str) -> std::result::Result<&'a str, (i32, String)> {
683 args.get(key)
684 .and_then(|v| v.as_str())
685 .ok_or((-32602, format!("Missing required parameter: '{key}'")))
686}