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) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.6
parent
b6fc298595
commit
bd1bccc692
@@ -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
|
||||
|
||||
@@ -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"})`,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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.
|
||||
|
||||
+22
-12
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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;
|
||||
+94
-29
@@ -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
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user