From bd1bccc69263942ef9502f5e9e0b3056c9e63be6 Mon Sep 17 00:00:00 2001 From: Algis Dumbris Date: Thu, 26 Mar 2026 07:17:23 +0200 Subject: [PATCH] feat(015): SQL query interface for agents + split read/write pools Split Connection Pools: - writeDB: MaxOpenConns=1, serializes all writes (no SQLITE_BUSY) - readDB: MaxOpenConns=8, query_only=ON, for all SELECTs - QueryDB() helper returns read pool when available SQL Query Interface: - New 'query' action via execute MCP tool - Read-only enforcement (PRAGMA query_only=ON + SQL validation) - Curated views: my_messages, my_channels, channel_messages - Per-agent access control via CTE injection - Auto LIMIT 100, 5s timeout, SELECT-only validation - Blocks: INSERT, UPDATE, DELETE, DROP, PRAGMA, etc. - 12 new tests (access control, validation, limits, CTEs) Migration 016: agent query views (v_agent_messages, etc.) Action registry: 30 actions (was 29, added 'query') All 29 test packages pass. Co-Authored-By: Claude Opus 4.6 (1M context) --- cmd/synapbus/main.go | 8 + internal/actions/registry.go | 28 ++ internal/actions/registry_test.go | 8 +- internal/agentquery/executor.go | 236 ++++++++++++ internal/agentquery/executor_test.go | 341 ++++++++++++++++++ internal/mcp/bridge.go | 29 ++ internal/mcp/server.go | 34 +- internal/mcp/tools_hybrid.go | 10 + .../storage/schema/016_agent_query_views.sql | 58 +++ internal/storage/sqlite.go | 123 +++++-- internal/storage/sqlite_test.go | 76 +++- 11 files changed, 906 insertions(+), 45 deletions(-) create mode 100644 internal/agentquery/executor.go create mode 100644 internal/agentquery/executor_test.go create mode 100644 internal/storage/schema/016_agent_query_views.sql diff --git a/cmd/synapbus/main.go b/cmd/synapbus/main.go index 6c8b1c9..b0a4219 100644 --- a/cmd/synapbus/main.go +++ b/cmd/synapbus/main.go @@ -39,6 +39,7 @@ import ( "github.com/synapbus/synapbus/internal/jsruntime" k8spkg "github.com/synapbus/synapbus/internal/k8s" mcpserver "github.com/synapbus/synapbus/internal/mcp" + "github.com/synapbus/synapbus/internal/agentquery" reactorpkg "github.com/synapbus/synapbus/internal/reactor" "github.com/synapbus/synapbus/internal/messaging" prommetrics "github.com/synapbus/synapbus/internal/metrics" @@ -492,6 +493,13 @@ func runServe(cmd *cobra.Command, args []string) error { // Create MCP server (4 hybrid tools: my_status, send_message, search, execute) mcpSrv := mcpserver.NewMCPServer(msgService, agentService, channelService, swarmService, attachmentService, searchService, reactionService, trustService, con, jsPool, actionRegistry, actionIndex, db.DB) + + // Set up SQL query executor for agents (uses read pool if available) + queryDB := db.QueryDB() + queryExec := agentquery.New(queryDB, slog.Default()) + mcpSrv.SetQueryExecutor(queryExec) + slog.Info("agent SQL query executor initialized", "read_pool", db.ReadDB != nil) + startTime := time.Now() // Start task expiry worker diff --git a/internal/actions/registry.go b/internal/actions/registry.go index f17c8ef..9ecda51 100644 --- a/internal/actions/registry.go +++ b/internal/actions/registry.go @@ -571,5 +571,33 @@ func allActions() []Action { }, }, }, + // ── SQL Query (1 action) ──────────────────────────────────── + { + Name: "query", + Category: "data", + Description: "Execute a read-only SQL query against your accessible messages, channels, and reactions. Use tables: my_messages (your DMs + joined channels), my_channels (channels you are in), channel_messages (messages in your channels). Results are limited to 100 rows. Only SELECT statements are allowed.", + Params: []Param{ + {Name: "sql", Type: "string", Description: "SQL SELECT query. Available tables: my_messages (id, body, from_agent, to_agent, priority, status, metadata, created_at, channel_name), my_channels (id, name, description, type), channel_messages (id, body, from_agent, priority, channel_name, created_at). CTEs (WITH) are supported.", Required: true}, + }, + Returns: "JSON with columns (array of column names), rows (array of row arrays), row_count, and truncated (boolean if > 100 rows)", + Examples: []Example{ + { + Description: "Find high-priority messages in a channel", + Code: `call("query", {"sql": "SELECT id, body, from_agent, priority FROM channel_messages WHERE channel_name = 'news-mcpproxy' AND priority >= 7 ORDER BY created_at DESC LIMIT 10"})`, + }, + { + Description: "List your channels", + Code: `call("query", {"sql": "SELECT name, description FROM my_channels ORDER BY name"})`, + }, + { + Description: "Count messages per channel", + Code: `call("query", {"sql": "SELECT channel_name, COUNT(*) as msg_count FROM channel_messages GROUP BY channel_name ORDER BY msg_count DESC"})`, + }, + { + Description: "Search messages with keyword", + Code: `call("query", {"sql": "SELECT id, body, from_agent, created_at FROM my_messages WHERE body LIKE '%MCP%' ORDER BY created_at DESC LIMIT 20"})`, + }, + }, + }, } } diff --git a/internal/actions/registry_test.go b/internal/actions/registry_test.go index 57aff94..1925e62 100644 --- a/internal/actions/registry_test.go +++ b/internal/actions/registry_test.go @@ -4,11 +4,11 @@ import ( "testing" ) -func TestRegistryHas29Actions(t *testing.T) { +func TestRegistryHas30Actions(t *testing.T) { r := NewRegistry() got := len(r.List()) - if got != 29 { - t.Errorf("expected 29 actions, got %d", got) + if got != 30 { + t.Errorf("expected 30 actions, got %d", got) } } @@ -58,6 +58,8 @@ func TestRegistryGetByName(t *testing.T) { "get_replies", // trust "get_trust", + // data + "query", } for _, name := range allNames { diff --git a/internal/agentquery/executor.go b/internal/agentquery/executor.go new file mode 100644 index 0000000..696224a --- /dev/null +++ b/internal/agentquery/executor.go @@ -0,0 +1,236 @@ +// Package agentquery provides a sandboxed SQL query executor for agents. +// Agents can run read-only SELECT queries against curated views with +// per-agent access control, automatic LIMIT enforcement, and timeouts. +package agentquery + +import ( + "context" + "database/sql" + "fmt" + "log/slog" + "strings" + "time" +) + +const ( + // MaxRows is the maximum number of rows returned by a query. + MaxRows = 100 + // QueryTimeout is the maximum duration for a query. + QueryTimeout = 5 * time.Second +) + +// Allowed view names that agents can query. +var allowedTables = map[string]bool{ + "my_messages": true, + "my_channels": true, + "channel_messages": true, +} + +// Executor runs sandboxed SQL queries on behalf of agents. +type Executor struct { + db *sql.DB // read-only pool (query_only=ON) + logger *slog.Logger +} + +// New creates a new query executor using the provided read-only database connection. +func New(readDB *sql.DB, logger *slog.Logger) *Executor { + return &Executor{ + db: readDB, + logger: logger.With("component", "agentquery"), + } +} + +// QueryResult holds the results of a SQL query. +type QueryResult struct { + Columns []string `json:"columns"` + Rows [][]interface{} `json:"rows"` + RowCount int `json:"row_count"` + Truncated bool `json:"truncated"` +} + +// Execute runs a SQL query on behalf of an agent with access control. +func (e *Executor) Execute(ctx context.Context, agentName, sqlQuery string) (*QueryResult, error) { + // 1. Validate the SQL statement + if err := validateSQL(sqlQuery); err != nil { + return nil, fmt.Errorf("query validation failed: %w", err) + } + + // 2. Rewrite the query to inject access control and enforce LIMIT + rewritten := rewriteQuery(agentName, sqlQuery) + + // 3. Execute with timeout + queryCtx, cancel := context.WithTimeout(ctx, QueryTimeout) + defer cancel() + + rows, err := e.db.QueryContext(queryCtx, rewritten) + if err != nil { + if queryCtx.Err() == context.DeadlineExceeded { + return nil, fmt.Errorf("query timed out after %s", QueryTimeout) + } + return nil, fmt.Errorf("query execution failed: %w", err) + } + defer rows.Close() + + // 4. Collect results + columns, err := rows.Columns() + if err != nil { + return nil, fmt.Errorf("get columns: %w", err) + } + + var resultRows [][]interface{} + truncated := false + + for rows.Next() { + if len(resultRows) >= MaxRows { + truncated = true + break + } + + values := make([]interface{}, len(columns)) + scanArgs := make([]interface{}, len(columns)) + for i := range values { + scanArgs[i] = &values[i] + } + + if err := rows.Scan(scanArgs...); err != nil { + return nil, fmt.Errorf("scan row: %w", err) + } + + // Convert []byte to string for JSON serialization + row := make([]interface{}, len(columns)) + for i, v := range values { + if b, ok := v.([]byte); ok { + row[i] = string(b) + } else { + row[i] = v + } + } + resultRows = append(resultRows, row) + } + + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterate rows: %w", err) + } + + if resultRows == nil { + resultRows = [][]interface{}{} + } + + e.logger.Info("agent query executed", + "agent", agentName, + "rows", len(resultRows), + "truncated", truncated, + ) + + return &QueryResult{ + Columns: columns, + Rows: resultRows, + RowCount: len(resultRows), + Truncated: truncated, + }, nil +} + +// validateSQL checks that the query is a read-only SELECT statement. +func validateSQL(query string) error { + trimmed := strings.TrimSpace(query) + if trimmed == "" { + return fmt.Errorf("empty query") + } + + // Remove comments + upper := strings.ToUpper(trimmed) + + // Must start with SELECT or WITH (CTEs) + if !strings.HasPrefix(upper, "SELECT") && !strings.HasPrefix(upper, "WITH") { + return fmt.Errorf("only SELECT statements are allowed (got %q)", firstWord(upper)) + } + + // Block dangerous keywords (check as whole words or with common delimiters) + blocked := []string{ + "INSERT ", "UPDATE ", "DELETE ", "DROP ", "ALTER ", "CREATE ", + "ATTACH ", "DETACH ", "PRAGMA", "REINDEX ", "VACUUM ", + "REPLACE ", "GRANT ", "REVOKE ", + } + for _, kw := range blocked { + if strings.Contains(upper, kw) { + return fmt.Errorf("statement contains blocked keyword: %s", strings.TrimSpace(kw)) + } + } + + // Block multiple statements (semicolon followed by non-whitespace) + parts := strings.Split(trimmed, ";") + nonEmpty := 0 + for _, p := range parts { + if strings.TrimSpace(p) != "" { + nonEmpty++ + } + } + if nonEmpty > 1 { + return fmt.Errorf("multiple statements not allowed") + } + + return nil +} + +// rewriteQuery wraps the agent's query with access control CTEs. +// It replaces references to my_messages, my_channels, channel_messages +// with CTEs that filter by the agent's access. +func rewriteQuery(agentName, query string) string { + // Build access-control CTEs that the agent's query can reference + cte := fmt.Sprintf(` +WITH my_messages AS ( + SELECT v.* FROM v_agent_messages v + LEFT JOIN channel_members cm ON cm.channel_id = v.channel_id AND cm.agent_name = %[1]s + WHERE v.to_agent = %[1]s + OR v.from_agent = %[1]s + OR (v.channel_id IS NOT NULL AND cm.agent_name IS NOT NULL) +), +my_channels AS ( + SELECT c.id, c.name, c.description, c.type, c.topic, c.is_private, c.created_at, + cm.joined_at AS member_since + FROM channels c + JOIN channel_members cm ON cm.channel_id = c.id AND cm.agent_name = %[1]s +), +channel_messages AS ( + SELECT v.* FROM v_channel_messages v + WHERE v.channel_id IN ( + SELECT channel_id FROM channel_members WHERE agent_name = %[1]s + ) +) +`, quoteSQLString(agentName)) + + trimmed := strings.TrimSpace(query) + upper := strings.ToUpper(trimmed) + + // Remove trailing semicolon if present + trimmed = strings.TrimRight(trimmed, "; \t\n") + + if strings.HasPrefix(upper, "WITH") { + // User has their own CTEs. We need to merge them. + // Strategy: our CTEs come first, then append user's CTEs after a comma. + // Remove the user's "WITH " prefix since our CTE block already has WITH. + userCTEs := strings.TrimSpace(trimmed[4:]) // skip "WITH" + return cte + ", " + userCTEs + " LIMIT " + fmt.Sprintf("%d", MaxRows+1) + } + + // Simple SELECT — prepend our CTEs + return cte + trimmed + " LIMIT " + fmt.Sprintf("%d", MaxRows+1) +} + +// quoteSQLString safely quotes a string for use in SQL. +func quoteSQLString(s string) string { + escaped := strings.ReplaceAll(s, "'", "''") + return "'" + escaped + "'" +} + +func firstWord(s string) string { + for i, c := range s { + if c == ' ' || c == '\t' || c == '\n' || c == '\r' || c == '(' { + return s[:i] + } + } + if len(s) > 20 { + return s[:20] + } + return s +} diff --git a/internal/agentquery/executor_test.go b/internal/agentquery/executor_test.go new file mode 100644 index 0000000..0a8c863 --- /dev/null +++ b/internal/agentquery/executor_test.go @@ -0,0 +1,341 @@ +package agentquery + +import ( + "context" + "database/sql" + "log/slog" + "testing" + + _ "modernc.org/sqlite" +) + +func setupTestDB(t *testing.T) *sql.DB { + t.Helper() + db, err := sql.Open("sqlite", ":memory:") + if err != nil { + t.Fatalf("open db: %v", err) + } + + // Create the schema needed for views + schema := ` + CREATE TABLE channels ( + id INTEGER PRIMARY KEY, + name TEXT NOT NULL UNIQUE, + description TEXT DEFAULT '', + type TEXT DEFAULT 'standard', + topic TEXT DEFAULT '', + is_private INTEGER DEFAULT 0, + created_at DATETIME DEFAULT CURRENT_TIMESTAMP + ); + CREATE TABLE channel_members ( + channel_id INTEGER, + agent_name TEXT, + joined_at DATETIME DEFAULT CURRENT_TIMESTAMP, + PRIMARY KEY (channel_id, agent_name) + ); + CREATE TABLE messages ( + id INTEGER PRIMARY KEY, + conversation_id INTEGER DEFAULT 0, + from_agent TEXT, + to_agent TEXT, + channel_id INTEGER, + reply_to INTEGER, + body TEXT, + priority INTEGER DEFAULT 5, + status TEXT DEFAULT 'pending', + metadata TEXT DEFAULT '{}', + created_at DATETIME DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME DEFAULT CURRENT_TIMESTAMP + ); + + -- Views matching the migration + CREATE VIEW v_agent_messages AS + SELECT m.id, m.body, m.from_agent, m.to_agent, m.priority, m.status, m.metadata, + m.created_at, m.updated_at, c.name AS channel_name, m.channel_id, m.reply_to, m.conversation_id + FROM messages m LEFT JOIN channels c ON c.id = m.channel_id; + + CREATE VIEW v_agent_channels AS + SELECT c.id, c.name, c.description, c.type, c.topic, c.is_private, c.created_at, + cm.joined_at AS member_since + FROM channels c JOIN channel_members cm ON cm.channel_id = c.id; + + CREATE VIEW v_channel_messages AS + SELECT m.id, m.body, m.from_agent, m.priority, m.status, m.metadata, m.created_at, + c.name AS channel_name, m.channel_id, m.reply_to + FROM messages m JOIN channels c ON c.id = m.channel_id; + ` + if _, err := db.Exec(schema); err != nil { + t.Fatalf("create schema: %v", err) + } + + // Seed test data + seed := ` + INSERT INTO channels (id, name) VALUES (1, 'general'), (2, 'news-mcpproxy'), (3, 'private-channel'); + INSERT INTO channel_members (channel_id, agent_name) VALUES + (1, 'agent-a'), (1, 'agent-b'), + (2, 'agent-a'), + (3, 'agent-b'); + + -- DMs + INSERT INTO messages (id, from_agent, to_agent, body, priority) VALUES + (1, 'algis', 'agent-a', 'Hello agent A', 7), + (2, 'agent-a', 'algis', 'Hi there', 5), + (3, 'algis', 'agent-b', 'Hello agent B', 5); + + -- Channel messages + INSERT INTO messages (id, from_agent, channel_id, body, priority) VALUES + (4, 'agent-a', 1, 'General post from A', 5), + (5, 'agent-b', 1, 'General post from B', 5), + (6, 'agent-a', 2, 'News post high prio', 8), + (7, 'agent-b', 3, 'Private channel msg', 5); + ` + if _, err := db.Exec(seed); err != nil { + t.Fatalf("seed data: %v", err) + } + + return db +} + +func TestExecuteBasicQuery(t *testing.T) { + db := setupTestDB(t) + defer db.Close() + + exec := New(db, slog.Default()) + result, err := exec.Execute(context.Background(), "agent-a", + "SELECT id, body, priority FROM my_messages ORDER BY id") + if err != nil { + t.Fatalf("query failed: %v", err) + } + + if len(result.Columns) != 3 { + t.Errorf("expected 3 columns, got %d", len(result.Columns)) + } + if result.Columns[0] != "id" || result.Columns[1] != "body" || result.Columns[2] != "priority" { + t.Errorf("unexpected columns: %v", result.Columns) + } + + // agent-a should see: DM to it (1), DM from it (2), general posts (4,5), news post (6) + // Should NOT see: DM to agent-b (3), private channel msg (7) + if result.RowCount < 4 { + t.Errorf("expected at least 4 rows for agent-a, got %d", result.RowCount) + } + + // Verify agent-b's DM and private channel msg are NOT visible + for _, row := range result.Rows { + id := row[0] + if id == int64(3) { + t.Error("agent-a should NOT see message 3 (DM to agent-b)") + } + if id == int64(7) { + t.Error("agent-a should NOT see message 7 (private channel, not joined)") + } + } +} + +func TestAccessControlAgentB(t *testing.T) { + db := setupTestDB(t) + defer db.Close() + + exec := New(db, slog.Default()) + result, err := exec.Execute(context.Background(), "agent-b", + "SELECT id, body FROM my_messages ORDER BY id") + if err != nil { + t.Fatalf("query failed: %v", err) + } + + // agent-b should see: DM to it (3), general posts (4,5), private channel (7) + // Should NOT see: DM to agent-a (1), DM from agent-a (2), news post (6) + hasMsg3 := false + hasMsg7 := false + for _, row := range result.Rows { + id := row[0] + if id == int64(3) { + hasMsg3 = true + } + if id == int64(7) { + hasMsg7 = true + } + if id == int64(1) { + t.Error("agent-b should NOT see message 1 (DM to agent-a)") + } + if id == int64(6) { + t.Error("agent-b should NOT see message 6 (news channel, not joined)") + } + } + if !hasMsg3 { + t.Error("agent-b should see message 3 (DM to it)") + } + if !hasMsg7 { + t.Error("agent-b should see message 7 (private channel, joined)") + } +} + +func TestQueryChannelMessages(t *testing.T) { + db := setupTestDB(t) + defer db.Close() + + exec := New(db, slog.Default()) + result, err := exec.Execute(context.Background(), "agent-a", + "SELECT id, body, channel_name FROM channel_messages WHERE channel_name = 'news-mcpproxy'") + if err != nil { + t.Fatalf("query failed: %v", err) + } + + if result.RowCount != 1 { + t.Errorf("expected 1 news message, got %d", result.RowCount) + } +} + +func TestQueryMyChannels(t *testing.T) { + db := setupTestDB(t) + defer db.Close() + + exec := New(db, slog.Default()) + result, err := exec.Execute(context.Background(), "agent-a", + "SELECT name FROM my_channels ORDER BY name") + if err != nil { + t.Fatalf("query failed: %v", err) + } + + // agent-a is in: general, news-mcpproxy (not private-channel) + if result.RowCount != 2 { + t.Errorf("expected 2 channels for agent-a, got %d", result.RowCount) + } +} + +func TestValidationRejectsInsert(t *testing.T) { + db := setupTestDB(t) + defer db.Close() + + exec := New(db, slog.Default()) + _, err := exec.Execute(context.Background(), "agent-a", + "INSERT INTO messages (body) VALUES ('evil')") + if err == nil { + t.Fatal("expected INSERT to be rejected") + } + if !contains(err.Error(), "only SELECT") { + t.Errorf("expected 'only SELECT' error, got: %v", err) + } +} + +func TestValidationRejectsDrop(t *testing.T) { + db := setupTestDB(t) + defer db.Close() + + exec := New(db, slog.Default()) + _, err := exec.Execute(context.Background(), "agent-a", + "SELECT 1; DROP TABLE messages") + if err == nil { + t.Fatal("expected multi-statement to be rejected") + } +} + +func TestValidationRejectsUpdate(t *testing.T) { + db := setupTestDB(t) + defer db.Close() + + exec := New(db, slog.Default()) + _, err := exec.Execute(context.Background(), "agent-a", + "UPDATE messages SET body = 'hacked'") + if err == nil { + t.Fatal("expected UPDATE to be rejected") + } +} + +func TestValidationRejectsPragma(t *testing.T) { + db := setupTestDB(t) + defer db.Close() + + exec := New(db, slog.Default()) + _, err := exec.Execute(context.Background(), "agent-a", + "SELECT * FROM pragma_table_info('messages')") + if err == nil { + t.Fatal("expected PRAGMA in SELECT to be rejected") + } +} + +func TestEmptyQuery(t *testing.T) { + db := setupTestDB(t) + defer db.Close() + + exec := New(db, slog.Default()) + _, err := exec.Execute(context.Background(), "agent-a", "") + if err == nil { + t.Fatal("expected empty query to be rejected") + } +} + +func TestCTEQuery(t *testing.T) { + db := setupTestDB(t) + defer db.Close() + + exec := New(db, slog.Default()) + result, err := exec.Execute(context.Background(), "agent-a", + "WITH high_prio AS (SELECT * FROM my_messages WHERE priority >= 7) SELECT id, priority FROM high_prio") + if err != nil { + t.Fatalf("CTE query failed: %v", err) + } + + // agent-a should see high-priority messages it has access to + if result.RowCount == 0 { + t.Error("expected at least 1 high-priority message") + } +} + +func TestEmptyResultSet(t *testing.T) { + db := setupTestDB(t) + defer db.Close() + + exec := New(db, slog.Default()) + result, err := exec.Execute(context.Background(), "agent-a", + "SELECT * FROM my_messages WHERE body = 'nonexistent'") + if err != nil { + t.Fatalf("query failed: %v", err) + } + if result.RowCount != 0 { + t.Errorf("expected 0 rows, got %d", result.RowCount) + } + if result.Rows == nil { + t.Error("rows should be empty array, not nil") + } + if result.Truncated { + t.Error("should not be truncated") + } +} + +func TestLimitEnforcement(t *testing.T) { + db := setupTestDB(t) + defer db.Close() + + // Insert 150 messages to test limit + for i := 100; i < 250; i++ { + _, _ = db.Exec("INSERT INTO messages (id, from_agent, to_agent, body) VALUES (?, 'algis', 'agent-a', 'msg')", i) + } + + exec := New(db, slog.Default()) + result, err := exec.Execute(context.Background(), "agent-a", + "SELECT id FROM my_messages") + if err != nil { + t.Fatalf("query failed: %v", err) + } + + if result.RowCount > MaxRows { + t.Errorf("expected max %d rows, got %d", MaxRows, result.RowCount) + } + if !result.Truncated { + t.Error("expected truncated=true for large result set") + } +} + +func contains(s, substr string) bool { + return len(s) >= len(substr) && (s == substr || len(s) > 0 && containsStr(s, substr)) +} + +func containsStr(s, sub string) bool { + for i := 0; i <= len(s)-len(sub); i++ { + if s[i:i+len(sub)] == sub { + return true + } + } + return false +} diff --git a/internal/mcp/bridge.go b/internal/mcp/bridge.go index 7ccd095..2402575 100644 --- a/internal/mcp/bridge.go +++ b/internal/mcp/bridge.go @@ -12,6 +12,7 @@ import ( "time" "github.com/synapbus/synapbus/internal/agents" + "github.com/synapbus/synapbus/internal/agentquery" "github.com/synapbus/synapbus/internal/attachments" "github.com/synapbus/synapbus/internal/channels" "github.com/synapbus/synapbus/internal/messaging" @@ -31,6 +32,7 @@ type ServiceBridge struct { searchService *search.Service reactionService *reactions.Service trustService *trust.Service + queryExecutor *agentquery.Executor agentName string } @@ -130,6 +132,10 @@ func (b *ServiceBridge) Call(ctx context.Context, actionName string, args map[st case "get_trust": return b.callGetTrust(ctx, args) + // --- SQL Query --- + case "query": + return b.callQuery(ctx, args) + // --- DM send (also accessible via bridge for execute tool) --- case "send_message": return b.callSendMessage(ctx, args) @@ -1215,6 +1221,29 @@ func (b *ServiceBridge) callGetTrust(ctx context.Context, args map[string]any) ( }, nil } +// SetQueryExecutor sets the SQL query executor for the bridge. +func (b *ServiceBridge) SetQueryExecutor(exec *agentquery.Executor) { + b.queryExecutor = exec +} + +func (b *ServiceBridge) callQuery(ctx context.Context, args map[string]any) (any, error) { + if b.queryExecutor == nil { + return nil, fmt.Errorf("SQL query not available") + } + + sqlStr := getString(args, "sql", "") + if sqlStr == "" { + return nil, fmt.Errorf("sql parameter is required") + } + + result, err := b.queryExecutor.Execute(ctx, b.agentName, sqlStr) + if err != nil { + return nil, err + } + + return result, nil +} + // --- Helpers --- // resolveChannelID resolves a channel ID from either channel_id or channel_name in args. diff --git a/internal/mcp/server.go b/internal/mcp/server.go index 25dbc3e..38cee2d 100644 --- a/internal/mcp/server.go +++ b/internal/mcp/server.go @@ -12,6 +12,7 @@ import ( "github.com/mark3labs/mcp-go/server" "github.com/synapbus/synapbus/internal/actions" + "github.com/synapbus/synapbus/internal/agentquery" "github.com/synapbus/synapbus/internal/agents" "github.com/synapbus/synapbus/internal/attachments" "github.com/synapbus/synapbus/internal/channels" @@ -26,12 +27,13 @@ import ( // MCPServer wraps the mcp-go server with SynapBus services. type MCPServer struct { - mcpServer *server.MCPServer - httpServer *server.StreamableHTTPServer - connMgr *ConnectionManager - agentService *agents.AgentService - logger *slog.Logger - console *console.Printer + mcpServer *server.MCPServer + httpServer *server.StreamableHTTPServer + connMgr *ConnectionManager + agentService *agents.AgentService + hybridRegistrar *HybridToolRegistrar + logger *slog.Logger + console *console.Printer } // NewMCPServer creates and configures a new MCP server with 4 hybrid tools registered. @@ -187,18 +189,26 @@ func NewMCPServer( ) s := &MCPServer{ - mcpServer: mcpSrv, - httpServer: httpServer, - connMgr: connMgr, - agentService: agentService, - logger: logger, - console: consolePrinter, + mcpServer: mcpSrv, + httpServer: httpServer, + connMgr: connMgr, + agentService: agentService, + hybridRegistrar: hybridRegistrar, + logger: logger, + console: consolePrinter, } logger.Info("MCP server initialized (4 hybrid tools, 4 prompts, streamable HTTP transport)") return s } +// SetQueryExecutor sets the SQL query executor for agent queries via the execute tool. +func (s *MCPServer) SetQueryExecutor(exec *agentquery.Executor) { + if s.hybridRegistrar != nil { + s.hybridRegistrar.SetQueryExecutor(exec) + } +} + // Handler returns the HTTP handler for mounting on a router. func (s *MCPServer) Handler() http.Handler { return s.httpServer diff --git a/internal/mcp/tools_hybrid.go b/internal/mcp/tools_hybrid.go index 3dac2f0..a26f5c6 100644 --- a/internal/mcp/tools_hybrid.go +++ b/internal/mcp/tools_hybrid.go @@ -17,6 +17,7 @@ import ( "github.com/synapbus/synapbus/internal/attachments" "github.com/synapbus/synapbus/internal/channels" "github.com/synapbus/synapbus/internal/jsruntime" + "github.com/synapbus/synapbus/internal/agentquery" "github.com/synapbus/synapbus/internal/messaging" "github.com/synapbus/synapbus/internal/reactions" "github.com/synapbus/synapbus/internal/search" @@ -37,9 +38,15 @@ type HybridToolRegistrar struct { actionRegistry *actions.Registry actionIndex *actions.Index db *sql.DB + queryExecutor *agentquery.Executor logger *slog.Logger } +// SetQueryExecutor sets the SQL query executor for all agent bridges. +func (h *HybridToolRegistrar) SetQueryExecutor(exec *agentquery.Executor) { + h.queryExecutor = exec +} + // NewHybridToolRegistrar creates a new hybrid tool registrar. func NewHybridToolRegistrar( msgService *messaging.MessagingService, @@ -495,6 +502,9 @@ func (h *HybridToolRegistrar) handleExecute(ctx context.Context, req mcplib.Call h.trustService, agentName, ) + if h.queryExecutor != nil { + bridge.SetQueryExecutor(h.queryExecutor) + } result, err := h.jsPool.Execute(ctx, code, bridge, jsruntime.ExecuteOptions{ Timeout: timeout, diff --git a/internal/storage/schema/016_agent_query_views.sql b/internal/storage/schema/016_agent_query_views.sql new file mode 100644 index 0000000..5b0b565 --- /dev/null +++ b/internal/storage/schema/016_agent_query_views.sql @@ -0,0 +1,58 @@ +-- 016: Agent SQL query views +-- These views are used by the 'query' action to give agents read access +-- to messages they can see. The views expose a stable schema that agents +-- can query via SQL. Access control is enforced at the Go layer by +-- rewriting queries to filter by agent name. + +-- Note: SQLite views cannot be parameterized. The Go query executor +-- wraps agent queries in a CTE that filters by the authenticated agent's +-- access (own DMs + joined channels). These views provide the base schema. + +-- my_messages: All messages accessible to the calling agent +CREATE VIEW IF NOT EXISTS v_agent_messages AS +SELECT + m.id, + m.body, + m.from_agent, + m.to_agent, + m.priority, + m.status, + m.metadata, + m.created_at, + m.updated_at, + c.name AS channel_name, + m.channel_id, + m.reply_to, + m.conversation_id +FROM messages m +LEFT JOIN channels c ON c.id = m.channel_id; + +-- my_channels: Channels the calling agent has joined +CREATE VIEW IF NOT EXISTS v_agent_channels AS +SELECT + c.id, + c.name, + c.description, + c.type, + c.topic, + c.is_private, + c.created_at, + cm.joined_at AS member_since +FROM channels c +JOIN channel_members cm ON cm.channel_id = c.id; + +-- channel_messages: Messages in channels (filtered by membership at Go layer) +CREATE VIEW IF NOT EXISTS v_channel_messages AS +SELECT + m.id, + m.body, + m.from_agent, + m.priority, + m.status, + m.metadata, + m.created_at, + c.name AS channel_name, + m.channel_id, + m.reply_to +FROM messages m +JOIN channels c ON c.id = m.channel_id; diff --git a/internal/storage/sqlite.go b/internal/storage/sqlite.go index adcfb9a..97863a0 100644 --- a/internal/storage/sqlite.go +++ b/internal/storage/sqlite.go @@ -12,17 +12,22 @@ import ( _ "modernc.org/sqlite" ) -// DB wraps a *sql.DB with SynapBus-specific configuration. +// DB wraps a write-only *sql.DB and an optional read-only *sql.DB +// for split connection pool architecture. The write pool has MaxOpenConns=1 +// to serialize writes and eliminate SQLITE_BUSY errors. The read pool has +// MaxOpenConns=8 and query_only=ON for safe concurrent reads. type DB struct { - *sql.DB + *sql.DB // Write pool (MaxOpenConns=1) + ReadDB *sql.DB // Read pool (MaxOpenConns=8, query_only=ON) — nil for :memory: DBs } -// New opens a SQLite database with WAL mode, busy_timeout, and foreign keys enabled. -// If dataDir is empty or ":memory:", an in-memory database is used. +// New opens a SQLite database with WAL mode, split read/write pools, and foreign keys. +// If dataDir is empty or ":memory:", an in-memory database is used (single pool, no split). func New(ctx context.Context, dataDir string) (*DB, error) { var dsn string + isMemory := dataDir == "" || dataDir == ":memory:" - if dataDir == "" || dataDir == ":memory:" { + if isMemory { dsn = ":memory:" } else { if err := os.MkdirAll(dataDir, 0o755); err != nil { @@ -31,12 +36,67 @@ func New(ctx context.Context, dataDir string) (*DB, error) { dsn = filepath.Join(dataDir, "synapbus.db") } - db, err := sql.Open("sqlite", dsn) + // Open WRITE pool (single connection, serializes all writes) + writeDB, err := openPool(ctx, dsn, poolConfig{ + maxOpen: 1, + maxIdle: 1, + queryOnly: false, + label: "write", + }) if err != nil { - return nil, fmt.Errorf("open database: %w", err) + return nil, fmt.Errorf("open write pool: %w", err) } - // Configure SQLite pragmas + result := &DB{DB: writeDB} + + // For file-based databases, open a separate READ pool + if !isMemory { + readDB, err := openPool(ctx, dsn, poolConfig{ + maxOpen: 8, + maxIdle: 4, + queryOnly: true, + label: "read", + }) + if err != nil { + writeDB.Close() + return nil, fmt.Errorf("open read pool: %w", err) + } + result.ReadDB = readDB + } + + // Verify settings on write pool + var journalMode string + if err := writeDB.QueryRowContext(ctx, "PRAGMA journal_mode").Scan(&journalMode); err != nil { + result.Close() + return nil, fmt.Errorf("verify journal_mode: %w", err) + } + + slog.Info("database opened", + "dsn", dsn, + "journal_mode", journalMode, + "write_pool", "MaxOpenConns=1", + "read_pool_enabled", result.ReadDB != nil, + ) + + return result, nil +} + +type poolConfig struct { + maxOpen int + maxIdle int + queryOnly bool + label string +} + +func openPool(ctx context.Context, dsn string, cfg poolConfig) (*sql.DB, error) { + db, err := sql.Open("sqlite", dsn) + if err != nil { + return nil, fmt.Errorf("open %s pool: %w", cfg.label, err) + } + + db.SetMaxOpenConns(cfg.maxOpen) + db.SetMaxIdleConns(cfg.maxIdle) + pragmas := []string{ "PRAGMA journal_mode=WAL", "PRAGMA busy_timeout=15000", @@ -44,13 +104,9 @@ func New(ctx context.Context, dataDir string) (*DB, error) { "PRAGMA synchronous=NORMAL", "PRAGMA wal_autocheckpoint=1000", } - - // Limit connection pool to reduce write contention. - // SQLite allows one writer at a time; multiple connections competing - // for the write lock cause SQLITE_BUSY errors. Keeping MaxOpenConns - // low reduces lock contention while still allowing concurrent reads. - db.SetMaxOpenConns(4) - db.SetMaxIdleConns(2) + if cfg.queryOnly { + pragmas = append(pragmas, "PRAGMA query_only=ON") + } for _, pragma := range pragmas { if _, err := db.ExecContext(ctx, pragma); err != nil { @@ -59,22 +115,31 @@ func New(ctx context.Context, dataDir string) (*DB, error) { } } - // Verify settings - var journalMode string - if err := db.QueryRowContext(ctx, "PRAGMA journal_mode").Scan(&journalMode); err != nil { - db.Close() - return nil, fmt.Errorf("verify journal_mode: %w", err) + return db, nil +} + +// QueryDB returns the read pool if available, otherwise falls back to the write pool. +// Use this for all SELECT queries to avoid blocking writers. +func (db *DB) QueryDB() *sql.DB { + if db.ReadDB != nil { + return db.ReadDB } - - slog.Info("database opened", - "dsn", dsn, - "journal_mode", journalMode, - ) - - return &DB{DB: db}, nil + return db.DB } -// Close closes the database connection. +// Close closes both the write and read database connections. func (db *DB) Close() error { - return db.DB.Close() + var errs []error + if db.ReadDB != nil { + if err := db.ReadDB.Close(); err != nil { + errs = append(errs, fmt.Errorf("close read pool: %w", err)) + } + } + if err := db.DB.Close(); err != nil { + errs = append(errs, fmt.Errorf("close write pool: %w", err)) + } + if len(errs) > 0 { + return errs[0] + } + return nil } diff --git a/internal/storage/sqlite_test.go b/internal/storage/sqlite_test.go index d2f2e0c..79fbd88 100644 --- a/internal/storage/sqlite_test.go +++ b/internal/storage/sqlite_test.go @@ -72,7 +72,7 @@ func TestNew(t *testing.T) { t.Errorf("busy_timeout = %d, want 15000", timeout) } - // Verify database is usable + // Verify database is usable via write pool _, err = db.Exec("CREATE TABLE test (id INTEGER PRIMARY KEY)") if err != nil { t.Fatalf("failed to create test table: %v", err) @@ -80,3 +80,77 @@ func TestNew(t *testing.T) { }) } } + +func TestSplitPools(t *testing.T) { + ctx := context.Background() + dir := t.TempDir() + + db, err := New(ctx, dir) + if err != nil { + t.Fatalf("New() error: %v", err) + } + defer db.Close() + + // Run migrations to create tables + if err := RunMigrations(ctx, db.DB); err != nil { + t.Fatalf("migrations: %v", err) + } + + // Verify read pool exists for file-based DB + if db.ReadDB == nil { + t.Fatal("expected ReadDB to be non-nil for file-based database") + } + + // Verify QueryDB returns read pool + if db.QueryDB() != db.ReadDB { + t.Error("QueryDB() should return ReadDB when available") + } + + // Create user first (FK requirement) + _, err = db.Exec("INSERT INTO users (id, username, password_hash, display_name) VALUES (1, 'testuser', 'hash', 'Test')") + if err != nil { + t.Fatalf("create user: %v", err) + } + + // Verify write pool can write + _, err = db.Exec("INSERT INTO agents (name, display_name, type, capabilities, owner_id, api_key_hash, status) VALUES ('test-agent', 'Test', 'ai', '{}', 1, 'hash', 'active')") + if err != nil { + t.Fatalf("write pool should allow writes: %v", err) + } + + // Verify read pool can read + var name string + err = db.ReadDB.QueryRow("SELECT name FROM agents WHERE name = 'test-agent'").Scan(&name) + if err != nil { + t.Fatalf("read pool should allow reads: %v", err) + } + if name != "test-agent" { + t.Errorf("expected 'test-agent', got %q", name) + } + + // Verify read pool rejects writes + _, err = db.ReadDB.Exec("INSERT INTO agents (name, display_name, type, capabilities, owner_id, api_key_hash, status) VALUES ('bad', 'Bad', 'ai', '{}', 1, 'hash', 'active')") + if err == nil { + t.Fatal("read pool should reject writes (query_only=ON)") + } +} + +func TestInMemoryNoSplitPool(t *testing.T) { + ctx := context.Background() + + db, err := New(ctx, ":memory:") + if err != nil { + t.Fatalf("New() error: %v", err) + } + defer db.Close() + + // In-memory DB should NOT have a separate read pool + if db.ReadDB != nil { + t.Error("in-memory DB should not have a separate ReadDB") + } + + // QueryDB should fall back to write pool + if db.QueryDB() != db.DB { + t.Error("QueryDB() should return write pool for in-memory DB") + } +}