Skip to main content

agentdb/
mcp.rs

1//! # MCP (Model Context Protocol) Server Interface
2//!
3//! Implements the MCP JSON-RPC transport for AgentDB, exposing the database
4//! as an MCP-compatible tool server. Supports:
5//!
6//! - `initialize` / `initialized` handshake
7//! - `tools/list` — enumerate all AgentDB capabilities as MCP tools
8//! - `tools/call` — invoke any AgentDB operation
9//! - `resources/list` / `resources/read` — expose database stats and collections
10//!
11//! ## Usage
12//!
13//! ```rust,no_run
14//! use agentdb::{AgentDB, mcp::McpServer};
15//!
16//! let db = AgentDB::open("agent.db").unwrap();
17//! let server = McpServer::new(db);
18//!
19//! // Process a JSON-RPC request (from stdin, HTTP, WebSocket, etc.)
20//! let request = r#"{"jsonrpc":"2.0","id":1,"method":"tools/list"}"#;
21//! let response = server.handle_message(request);
22//! println!("{}", response);
23//! ```
24
25use crate::db::AgentDB;
26use serde_json::{json, Value};
27use std::collections::HashMap;
28
29/// MCP server wrapping an AgentDB instance.
30pub struct McpServer {
31    db: AgentDB,
32}
33
34impl McpServer {
35    /// Create a new MCP server backed by the given database.
36    pub fn new(db: AgentDB) -> Self {
37        Self { db }
38    }
39
40    /// Handle a single JSON-RPC message string and return the response.
41    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(&params),
60            "initialized" => return String::new(),
61            "tools/list" => self.handle_tools_list(),
62            "tools/call" => self.handle_tools_call(&params),
63            "resources/list" => self.handle_resources_list(),
64            "resources/read" => self.handle_resources_read(&params),
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}