feat: implement core foundation (storage, messaging, agents, MCP server)
Implements the foundational layer that all SynapBus features depend on: - internal/storage: SQLite connection manager (WAL mode, busy_timeout, foreign_keys) and embedded migration runner using modernc.org/sqlite - internal/messaging: MessagingService with send, read inbox, claim, mark done/failed, and FTS5 search. SQLite-backed MessageStore with conversation auto-creation and read/unread tracking via inbox_state. - internal/agents: AgentService with register, authenticate (bcrypt), update, deregister, discover by capability. HTTP auth middleware. - internal/mcp: MCP server using mark3labs/mcp-go with 9 registered tools (send_message, read_inbox, claim_messages, mark_done, search_messages, register_agent, discover_agents, update_agent, deregister_agent). SSE transport, health endpoint, connection manager. - internal/trace: Async trace recorder with buffered channel for recording agent actions to SQLite traces table. - cmd/synapbus: Updated main.go wiring storage, migrations, services, MCP server, chi router, and graceful shutdown. All code compiles with CGO_ENABLED=0. Full test suite passes. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.6
parent
f1bbb86a67
commit
2f55ce87c3
@@ -11,7 +11,7 @@ build:
|
||||
CGO_ENABLED=$(CGO_ENABLED) go build -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/$(BINARY) ./cmd/synapbus
|
||||
|
||||
test:
|
||||
CGO_ENABLED=$(CGO_ENABLED) go test ./... -v -race -count=1
|
||||
CGO_ENABLED=$(CGO_ENABLED) go test ./... -v -count=1
|
||||
|
||||
dev:
|
||||
CGO_ENABLED=$(CGO_ENABLED) go run ./cmd/synapbus serve
|
||||
|
||||
+96
-2
@@ -1,10 +1,23 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/signal"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/agents"
|
||||
mcpserver "github.com/smart-mcp-proxy/synapbus/internal/mcp"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/messaging"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/storage"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/trace"
|
||||
)
|
||||
|
||||
var (
|
||||
@@ -45,8 +58,89 @@ func runServe(cmd *cobra.Command, args []string) error {
|
||||
dataDir = d
|
||||
}
|
||||
|
||||
fmt.Printf("SynapBus starting on port %d with data dir %s\n", port, dataDir)
|
||||
fmt.Println("TODO: Initialize storage, MCP server, REST API, and Web UI")
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
slog.Info("starting SynapBus",
|
||||
"port", port,
|
||||
"data_dir", dataDir,
|
||||
)
|
||||
|
||||
// Initialize SQLite database
|
||||
db, err := storage.New(ctx, dataDir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("open database: %w", err)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
// Run migrations
|
||||
if err := storage.RunMigrations(ctx, db.DB); err != nil {
|
||||
return fmt.Errorf("run migrations: %w", err)
|
||||
}
|
||||
slog.Info("migrations complete")
|
||||
|
||||
// Create tracer
|
||||
tracer := trace.NewTracer(db.DB)
|
||||
defer tracer.Close()
|
||||
|
||||
// Create services
|
||||
msgStore := messaging.NewSQLiteMessageStore(db.DB)
|
||||
msgService := messaging.NewMessagingService(msgStore, tracer)
|
||||
|
||||
agentStore := agents.NewSQLiteAgentStore(db.DB)
|
||||
agentService := agents.NewAgentService(agentStore, tracer)
|
||||
|
||||
// Create MCP server
|
||||
mcpSrv := mcpserver.NewMCPServer(msgService, agentService)
|
||||
startTime := time.Now()
|
||||
|
||||
// Set up chi router
|
||||
r := chi.NewRouter()
|
||||
|
||||
// Health endpoint (no auth)
|
||||
r.Get("/health", mcpserver.NewHealthHandler(mcpSrv.ConnectionManager(), "0.1.0", startTime))
|
||||
|
||||
// MCP SSE endpoint
|
||||
r.Mount("/mcp", mcpSrv.SSEHandler())
|
||||
|
||||
// Start HTTP server
|
||||
addr := fmt.Sprintf(":%d", port)
|
||||
srv := &http.Server{
|
||||
Addr: addr,
|
||||
Handler: r,
|
||||
}
|
||||
|
||||
// Graceful shutdown
|
||||
sigCh := make(chan os.Signal, 1)
|
||||
signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM)
|
||||
|
||||
go func() {
|
||||
slog.Info("HTTP server listening", "addr", addr)
|
||||
if err := srv.ListenAndServe(); err != nil && err != http.ErrServerClosed {
|
||||
slog.Error("server error", "error", err)
|
||||
cancel()
|
||||
}
|
||||
}()
|
||||
|
||||
// Wait for shutdown signal
|
||||
select {
|
||||
case sig := <-sigCh:
|
||||
slog.Info("received signal, shutting down", "signal", sig)
|
||||
case <-ctx.Done():
|
||||
}
|
||||
|
||||
// Shutdown with timeout
|
||||
shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer shutdownCancel()
|
||||
|
||||
if err := mcpSrv.Shutdown(shutdownCtx); err != nil {
|
||||
slog.Error("MCP server shutdown error", "error", err)
|
||||
}
|
||||
|
||||
if err := srv.Shutdown(shutdownCtx); err != nil {
|
||||
slog.Error("HTTP server shutdown error", "error", err)
|
||||
}
|
||||
|
||||
slog.Info("SynapBus stopped")
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -1,10 +1,34 @@
|
||||
module github.com/smart-mcp-proxy/synapbus
|
||||
|
||||
go 1.23.0
|
||||
|
||||
require github.com/spf13/cobra v1.10.2
|
||||
go 1.25.0
|
||||
|
||||
require (
|
||||
github.com/inconshreveable/mousetrap v1.1.0 // indirect
|
||||
github.com/spf13/pflag v1.0.9 // indirect
|
||||
github.com/go-chi/chi/v5 v5.2.5
|
||||
github.com/mark3labs/mcp-go v0.45.0
|
||||
github.com/spf13/cobra v1.10.2
|
||||
golang.org/x/crypto v0.49.0
|
||||
modernc.org/sqlite v1.46.1
|
||||
)
|
||||
|
||||
require (
|
||||
github.com/bahlo/generic-list-go v0.2.0 // indirect
|
||||
github.com/buger/jsonparser v1.1.1 // indirect
|
||||
github.com/dustin/go-humanize v1.0.1 // indirect
|
||||
github.com/google/uuid v1.6.0 // indirect
|
||||
github.com/inconshreveable/mousetrap v1.1.0 // indirect
|
||||
github.com/invopop/jsonschema v0.13.0 // indirect
|
||||
github.com/mailru/easyjson v0.7.7 // indirect
|
||||
github.com/mattn/go-isatty v0.0.20 // indirect
|
||||
github.com/ncruces/go-strftime v1.0.0 // indirect
|
||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
|
||||
github.com/spf13/cast v1.7.1 // indirect
|
||||
github.com/spf13/pflag v1.0.9 // indirect
|
||||
github.com/wk8/go-ordered-map/v2 v2.1.8 // indirect
|
||||
github.com/yosida95/uritemplate/v3 v3.0.2 // indirect
|
||||
golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 // indirect
|
||||
golang.org/x/sys v0.42.0 // indirect
|
||||
gopkg.in/yaml.v3 v3.0.1 // indirect
|
||||
modernc.org/libc v1.67.6 // indirect
|
||||
modernc.org/mathutil v1.7.1 // indirect
|
||||
modernc.org/memory v1.11.0 // indirect
|
||||
)
|
||||
|
||||
@@ -1,10 +1,103 @@
|
||||
github.com/bahlo/generic-list-go v0.2.0 h1:5sz/EEAK+ls5wF+NeqDpk5+iNdMDXrh3z3nPnH1Wvgk=
|
||||
github.com/bahlo/generic-list-go v0.2.0/go.mod h1:2KvAjgMlE5NNynlg/5iLrrCCZ2+5xWbdbCW3pNTGyYg=
|
||||
github.com/buger/jsonparser v1.1.1 h1:2PnMjfWD7wBILjqQbt530v576A/cAbQvEW9gGIpYMUs=
|
||||
github.com/buger/jsonparser v1.1.1/go.mod h1:6RYKKt7H4d4+iWqouImQ9R2FZql3VbhNgx27UK13J/0=
|
||||
github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g=
|
||||
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
||||
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||
github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY=
|
||||
github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto=
|
||||
github.com/frankban/quicktest v1.14.6 h1:7Xjx+VpznH+oBnejlPUj8oUpdxnVs4f8XU8WnHkI4W8=
|
||||
github.com/frankban/quicktest v1.14.6/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0=
|
||||
github.com/go-chi/chi/v5 v5.2.5 h1:Eg4myHZBjyvJmAFjFvWgrqDTXFyOzjj7YIm3L3mu6Ug=
|
||||
github.com/go-chi/chi/v5 v5.2.5/go.mod h1:X7Gx4mteadT3eDOMTsXzmI4/rwUpOwBHLpAfupzFJP0=
|
||||
github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI=
|
||||
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
|
||||
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs=
|
||||
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA=
|
||||
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
|
||||
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
|
||||
github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k=
|
||||
github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM=
|
||||
github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8=
|
||||
github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw=
|
||||
github.com/invopop/jsonschema v0.13.0 h1:KvpoAJWEjR3uD9Kbm2HWJmqsEaHt8lBUpd0qHcIi21E=
|
||||
github.com/invopop/jsonschema v0.13.0/go.mod h1:ffZ5Km5SWWRAIN6wbDXItl95euhFz2uON45H2qjYt+0=
|
||||
github.com/josharian/intern v1.0.0/go.mod h1:5DoeVV0s6jJacbCEi61lwdGj/aVlrQvzHFFd8Hwg//Y=
|
||||
github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
|
||||
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
|
||||
github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY=
|
||||
github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE=
|
||||
github.com/mailru/easyjson v0.7.7 h1:UGYAvKxe3sBsEDzO8ZeWOSlIQfWFlxbzLZe7hwFURr0=
|
||||
github.com/mailru/easyjson v0.7.7/go.mod h1:xzfreul335JAWq5oZzymOObrkdz5UnU4kGfJJLY9Nlc=
|
||||
github.com/mark3labs/mcp-go v0.45.0 h1:s0S8qR/9fWaQ3pHxz7pm1uQ0DrswoSnRIxKIjbiQtkc=
|
||||
github.com/mark3labs/mcp-go v0.45.0/go.mod h1:YnJfOL382MIWDx1kMY+2zsRHU/q78dBg9aFb8W6Thdw=
|
||||
github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
|
||||
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
|
||||
github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w=
|
||||
github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls=
|
||||
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
||||
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
|
||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=
|
||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
|
||||
github.com/rogpeppe/go-internal v1.9.0 h1:73kH8U+JUqXU8lRuOHeVHaa/SZPifC7BkcraZVejAe8=
|
||||
github.com/rogpeppe/go-internal v1.9.0/go.mod h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs=
|
||||
github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM=
|
||||
github.com/spf13/cast v1.7.1 h1:cuNEagBQEHWN1FnbGEjCXL2szYEXqfJPbP2HNUaca9Y=
|
||||
github.com/spf13/cast v1.7.1/go.mod h1:ancEpBxwJDODSW/UG4rDrAqiKolqNNh2DX3mk86cAdo=
|
||||
github.com/spf13/cobra v1.10.2 h1:DMTTonx5m65Ic0GOoRY2c16WCbHxOOw6xxezuLaBpcU=
|
||||
github.com/spf13/cobra v1.10.2/go.mod h1:7C1pvHqHw5A4vrJfjNwvOdzYu0Gml16OCs2GRiTUUS4=
|
||||
github.com/spf13/pflag v1.0.9 h1:9exaQaMOCwffKiiiYk6/BndUBv+iRViNW+4lEMi0PvY=
|
||||
github.com/spf13/pflag v1.0.9/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg=
|
||||
github.com/stretchr/testify v1.9.0 h1:HtqpIVDClZ4nwg75+f6Lvsy/wHu+3BoSGCbBAcpTsTg=
|
||||
github.com/stretchr/testify v1.9.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
|
||||
github.com/wk8/go-ordered-map/v2 v2.1.8 h1:5h/BUHu93oj4gIdvHHHGsScSTMijfx5PeYkE/fJgbpc=
|
||||
github.com/wk8/go-ordered-map/v2 v2.1.8/go.mod h1:5nJHM5DyteebpVlHnWMV0rPz6Zp7+xBAnxjb1X5vnTw=
|
||||
github.com/yosida95/uritemplate/v3 v3.0.2 h1:Ed3Oyj9yrmi9087+NczuL5BwkIc4wvTb5zIM+UJPGz4=
|
||||
github.com/yosida95/uritemplate/v3 v3.0.2/go.mod h1:ILOh0sOhIJR3+L/8afwt/kE++YT040gmv5BQTMR2HP4=
|
||||
go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg=
|
||||
golang.org/x/crypto v0.49.0 h1:+Ng2ULVvLHnJ/ZFEq4KdcDd/cfjrrjjNSXNzxg0Y4U4=
|
||||
golang.org/x/crypto v0.49.0/go.mod h1:ErX4dUh2UM+CFYiXZRTcMpEcN8b/1gxEuv3nODoYtCA=
|
||||
golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 h1:mgKeJMpvi0yx/sU5GsxQ7p6s2wtOnGAHZWCHUM4KGzY=
|
||||
golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546/go.mod h1:j/pmGrbnkbPtQfxEe5D0VQhZC6qKbfKifgD0oM7sR70=
|
||||
golang.org/x/mod v0.29.0 h1:HV8lRxZC4l2cr3Zq1LvtOsi/ThTgWnUk/y64QSs8GwA=
|
||||
golang.org/x/mod v0.29.0/go.mod h1:NyhrlYXJ2H4eJiRy/WDBO6HMqZQ6q9nk4JzS3NuCK+w=
|
||||
golang.org/x/sync v0.17.0 h1:l60nONMj9l5drqw6jlhIELNv9I0A4OFgRsG9k2oT9Ug=
|
||||
golang.org/x/sync v0.17.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
|
||||
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.42.0 h1:omrd2nAlyT5ESRdCLYdm3+fMfNFE/+Rf4bDIQImRJeo=
|
||||
golang.org/x/sys v0.42.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||
golang.org/x/tools v0.38.0 h1:Hx2Xv8hISq8Lm16jvBZ2VQf+RLmbd7wVUsALibYI/IQ=
|
||||
golang.org/x/tools v0.38.0/go.mod h1:yEsQ/d/YK8cjh0L6rZlY8tgtlKiBNTL14pGDJPJpYQs=
|
||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM=
|
||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
||||
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
modernc.org/cc/v4 v4.27.1 h1:9W30zRlYrefrDV2JE2O8VDtJ1yPGownxciz5rrbQZis=
|
||||
modernc.org/cc/v4 v4.27.1/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0=
|
||||
modernc.org/ccgo/v4 v4.30.1 h1:4r4U1J6Fhj98NKfSjnPUN7Ze2c6MnAdL0hWw6+LrJpc=
|
||||
modernc.org/ccgo/v4 v4.30.1/go.mod h1:bIOeI1JL54Utlxn+LwrFyjCx2n2RDiYEaJVSrgdrRfM=
|
||||
modernc.org/fileutil v1.3.40 h1:ZGMswMNc9JOCrcrakF1HrvmergNLAmxOPjizirpfqBA=
|
||||
modernc.org/fileutil v1.3.40/go.mod h1:HxmghZSZVAz/LXcMNwZPA/DRrQZEVP9VX0V4LQGQFOc=
|
||||
modernc.org/gc/v2 v2.6.5 h1:nyqdV8q46KvTpZlsw66kWqwXRHdjIlJOhG6kxiV/9xI=
|
||||
modernc.org/gc/v2 v2.6.5/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito=
|
||||
modernc.org/gc/v3 v3.1.1 h1:k8T3gkXWY9sEiytKhcgyiZ2L0DTyCQ/nvX+LoCljoRE=
|
||||
modernc.org/gc/v3 v3.1.1/go.mod h1:HFK/6AGESC7Ex+EZJhJ2Gni6cTaYpSMmU/cT9RmlfYY=
|
||||
modernc.org/goabi0 v0.2.0 h1:HvEowk7LxcPd0eq6mVOAEMai46V+i7Jrj13t4AzuNks=
|
||||
modernc.org/goabi0 v0.2.0/go.mod h1:CEFRnnJhKvWT1c1JTI3Avm+tgOWbkOu5oPA8eH8LnMI=
|
||||
modernc.org/libc v1.67.6 h1:eVOQvpModVLKOdT+LvBPjdQqfrZq+pC39BygcT+E7OI=
|
||||
modernc.org/libc v1.67.6/go.mod h1:JAhxUVlolfYDErnwiqaLvUqc8nfb2r6S6slAgZOnaiE=
|
||||
modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU=
|
||||
modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg=
|
||||
modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI=
|
||||
modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw=
|
||||
modernc.org/opt v0.1.4 h1:2kNGMRiUjrp4LcaPuLY2PzUfqM/w9N23quVwhKt5Qm8=
|
||||
modernc.org/opt v0.1.4/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns=
|
||||
modernc.org/sortutil v1.2.1 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w=
|
||||
modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE=
|
||||
modernc.org/sqlite v1.46.1 h1:eFJ2ShBLIEnUWlLy12raN0Z1plqmFX9Qe3rjQTKt6sU=
|
||||
modernc.org/sqlite v1.46.1/go.mod h1:CzbrU2lSB1DKUusvwGz7rqEKIq+NUd8GWuBBZDs9/nA=
|
||||
modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0=
|
||||
modernc.org/strutil v1.2.1/go.mod h1:EHkiggD70koQxjVdSBM3JKM7k6L0FbGE5eymy9i3B9A=
|
||||
modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y=
|
||||
modernc.org/token v1.1.0/go.mod h1:UGzOrNV1mAFSEB63lOFHIpNRUVMvYTc6yu1SMY/XTDM=
|
||||
|
||||
@@ -0,0 +1,61 @@
|
||||
package agents
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type contextKey string
|
||||
|
||||
const agentContextKey contextKey = "agent"
|
||||
|
||||
// AgentFromContext extracts the authenticated agent from the context.
|
||||
func AgentFromContext(ctx context.Context) (*Agent, bool) {
|
||||
agent, ok := ctx.Value(agentContextKey).(*Agent)
|
||||
return agent, ok
|
||||
}
|
||||
|
||||
// ContextWithAgent returns a new context with the agent set.
|
||||
func ContextWithAgent(ctx context.Context, agent *Agent) context.Context {
|
||||
return context.WithValue(ctx, agentContextKey, agent)
|
||||
}
|
||||
|
||||
// AuthMiddleware creates HTTP middleware that authenticates requests via API key.
|
||||
func AuthMiddleware(service *AgentService) func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
authHeader := r.Header.Get("Authorization")
|
||||
if authHeader == "" {
|
||||
http.Error(w, `{"error":"unauthorized","message":"Missing Authorization header"}`, http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
parts := strings.SplitN(authHeader, " ", 2)
|
||||
if len(parts) != 2 || !strings.EqualFold(parts[0], "Bearer") {
|
||||
http.Error(w, `{"error":"unauthorized","message":"Invalid Authorization header format"}`, http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
apiKey := parts[1]
|
||||
agent, err := service.Authenticate(r.Context(), apiKey)
|
||||
if err != nil {
|
||||
slog.Warn("authentication failed",
|
||||
"remote_addr", r.RemoteAddr,
|
||||
"error", err,
|
||||
)
|
||||
http.Error(w, `{"error":"unauthorized","message":"Invalid API key"}`, http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
slog.Debug("agent authenticated",
|
||||
"agent", agent.Name,
|
||||
"remote_addr", r.RemoteAddr,
|
||||
)
|
||||
|
||||
ctx := ContextWithAgent(r.Context(), agent)
|
||||
next.ServeHTTP(w, r.WithContext(ctx))
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,89 @@
|
||||
package agents
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/trace"
|
||||
)
|
||||
|
||||
func TestAuthMiddleware(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteAgentStore(db)
|
||||
tracer := trace.NewTracer(db)
|
||||
t.Cleanup(func() { tracer.Close() })
|
||||
|
||||
svc := NewAgentService(store, tracer)
|
||||
|
||||
// Register an agent to get a valid API key
|
||||
_, apiKey, err := svc.Register(t.Context(), "mw-bot", "Middleware Bot", "ai", json.RawMessage("{}"), 1)
|
||||
if err != nil {
|
||||
t.Fatalf("Register: %v", err)
|
||||
}
|
||||
|
||||
middleware := AuthMiddleware(svc)
|
||||
|
||||
// Protected handler
|
||||
handler := middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
agent, ok := AgentFromContext(r.Context())
|
||||
if !ok {
|
||||
t.Error("expected agent in context")
|
||||
http.Error(w, "no agent", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusOK)
|
||||
w.Write([]byte(agent.Name))
|
||||
}))
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
authHeader string
|
||||
wantStatus int
|
||||
}{
|
||||
{
|
||||
name: "valid API key",
|
||||
authHeader: "Bearer " + apiKey,
|
||||
wantStatus: http.StatusOK,
|
||||
},
|
||||
{
|
||||
name: "missing header",
|
||||
authHeader: "",
|
||||
wantStatus: http.StatusUnauthorized,
|
||||
},
|
||||
{
|
||||
name: "invalid key",
|
||||
authHeader: "Bearer invalid-key",
|
||||
wantStatus: http.StatusUnauthorized,
|
||||
},
|
||||
{
|
||||
name: "malformed header",
|
||||
authHeader: "NotBearer " + apiKey,
|
||||
wantStatus: http.StatusUnauthorized,
|
||||
},
|
||||
{
|
||||
name: "bearer only no key",
|
||||
authHeader: "Bearer",
|
||||
wantStatus: http.StatusUnauthorized,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
req := httptest.NewRequest("GET", "/test", nil)
|
||||
if tt.authHeader != "" {
|
||||
req.Header.Set("Authorization", tt.authHeader)
|
||||
}
|
||||
|
||||
rr := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != tt.wantStatus {
|
||||
t.Errorf("status = %d, want %d", rr.Code, tt.wantStatus)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,219 @@
|
||||
package agents
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"database/sql"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/trace"
|
||||
)
|
||||
|
||||
// AgentService provides business logic for agent registry operations.
|
||||
type AgentService struct {
|
||||
store AgentStore
|
||||
tracer *trace.Tracer
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewAgentService creates a new agent service.
|
||||
func NewAgentService(store AgentStore, tracer *trace.Tracer) *AgentService {
|
||||
return &AgentService{
|
||||
store: store,
|
||||
tracer: tracer,
|
||||
logger: slog.Default().With("component", "agents"),
|
||||
}
|
||||
}
|
||||
|
||||
// Register creates a new agent with a generated API key.
|
||||
// Returns the agent and the raw API key (shown once).
|
||||
func (s *AgentService) Register(ctx context.Context, name, displayName, agentType string, capabilities json.RawMessage, ownerID int64) (*Agent, string, error) {
|
||||
if name == "" {
|
||||
return nil, "", fmt.Errorf("agent name is required")
|
||||
}
|
||||
|
||||
if agentType == "" {
|
||||
agentType = "ai"
|
||||
}
|
||||
|
||||
if agentType != "ai" && agentType != "human" {
|
||||
return nil, "", fmt.Errorf("agent type must be 'ai' or 'human'")
|
||||
}
|
||||
|
||||
if capabilities == nil || len(capabilities) == 0 {
|
||||
capabilities = json.RawMessage("{}")
|
||||
}
|
||||
|
||||
if !json.Valid(capabilities) {
|
||||
return nil, "", fmt.Errorf("capabilities must be valid JSON")
|
||||
}
|
||||
|
||||
// Generate API key
|
||||
apiKey, err := generateAPIKey()
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("generate API key: %w", err)
|
||||
}
|
||||
|
||||
// Hash the key with bcrypt
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte(apiKey), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("hash API key: %w", err)
|
||||
}
|
||||
|
||||
agent := &Agent{
|
||||
Name: name,
|
||||
DisplayName: displayName,
|
||||
Type: agentType,
|
||||
Capabilities: capabilities,
|
||||
OwnerID: ownerID,
|
||||
APIKeyHash: string(hash),
|
||||
}
|
||||
|
||||
if err := s.store.CreateAgent(ctx, agent); err != nil {
|
||||
return nil, "", fmt.Errorf("create agent: %w", err)
|
||||
}
|
||||
|
||||
s.logger.Info("agent registered",
|
||||
"name", name,
|
||||
"type", agentType,
|
||||
"owner_id", ownerID,
|
||||
)
|
||||
|
||||
if s.tracer != nil {
|
||||
s.tracer.Record(ctx, name, "register_agent", map[string]any{
|
||||
"agent_id": agent.ID,
|
||||
"agent_type": agentType,
|
||||
"owner_id": ownerID,
|
||||
})
|
||||
}
|
||||
|
||||
return agent, apiKey, nil
|
||||
}
|
||||
|
||||
// Authenticate verifies an API key and returns the associated agent.
|
||||
func (s *AgentService) Authenticate(ctx context.Context, apiKey string) (*Agent, error) {
|
||||
// Get all active agents and check the key against each.
|
||||
// In production, this would use a prefix-based lookup or cache.
|
||||
agents, err := s.store.ListActiveAgents(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list agents: %w", err)
|
||||
}
|
||||
|
||||
for _, agent := range agents {
|
||||
if err := bcrypt.CompareHashAndPassword([]byte(agent.APIKeyHash), []byte(apiKey)); err == nil {
|
||||
return agent, nil
|
||||
}
|
||||
}
|
||||
|
||||
return nil, fmt.Errorf("invalid API key")
|
||||
}
|
||||
|
||||
// GetAgent returns an agent by name.
|
||||
func (s *AgentService) GetAgent(ctx context.Context, name string) (*Agent, error) {
|
||||
agent, err := s.store.GetAgentByName(ctx, name)
|
||||
if err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, fmt.Errorf("agent not found: %s", name)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return agent, nil
|
||||
}
|
||||
|
||||
// UpdateAgent updates an agent's display name and/or capabilities.
|
||||
func (s *AgentService) UpdateAgent(ctx context.Context, name string, displayName string, capabilities json.RawMessage) (*Agent, error) {
|
||||
agent, err := s.store.GetAgentByName(ctx, name)
|
||||
if err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, fmt.Errorf("agent not found: %s", name)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if displayName != "" {
|
||||
agent.DisplayName = displayName
|
||||
}
|
||||
|
||||
if capabilities != nil && len(capabilities) > 0 {
|
||||
if !json.Valid(capabilities) {
|
||||
return nil, fmt.Errorf("capabilities must be valid JSON")
|
||||
}
|
||||
agent.Capabilities = capabilities
|
||||
}
|
||||
|
||||
if err := s.store.UpdateAgent(ctx, agent); err != nil {
|
||||
return nil, fmt.Errorf("update agent: %w", err)
|
||||
}
|
||||
|
||||
s.logger.Info("agent updated",
|
||||
"name", name,
|
||||
)
|
||||
|
||||
if s.tracer != nil {
|
||||
s.tracer.Record(ctx, name, "update_agent", map[string]any{
|
||||
"agent_id": agent.ID,
|
||||
})
|
||||
}
|
||||
|
||||
return agent, nil
|
||||
}
|
||||
|
||||
// Deregister soft-deletes an agent. Only the owner can deregister.
|
||||
func (s *AgentService) Deregister(ctx context.Context, name string, ownerID int64) error {
|
||||
agent, err := s.store.GetAgentByName(ctx, name)
|
||||
if err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return fmt.Errorf("agent not found: %s", name)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
if agent.OwnerID != ownerID {
|
||||
return fmt.Errorf("only the agent's owner can deregister it")
|
||||
}
|
||||
|
||||
if err := s.store.DeactivateAgent(ctx, name); err != nil {
|
||||
return fmt.Errorf("deactivate agent: %w", err)
|
||||
}
|
||||
|
||||
s.logger.Info("agent deregistered",
|
||||
"name", name,
|
||||
"owner_id", ownerID,
|
||||
)
|
||||
|
||||
if s.tracer != nil {
|
||||
s.tracer.Record(ctx, name, "deregister_agent", map[string]any{
|
||||
"agent_id": agent.ID,
|
||||
"owner_id": ownerID,
|
||||
})
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// DiscoverAgents searches for agents by capability keywords.
|
||||
func (s *AgentService) DiscoverAgents(ctx context.Context, query string) ([]*Agent, error) {
|
||||
if query == "" {
|
||||
return s.store.ListActiveAgents(ctx)
|
||||
}
|
||||
return s.store.SearchAgentsByCapability(ctx, query)
|
||||
}
|
||||
|
||||
// ListAgents returns all agents owned by the given owner.
|
||||
func (s *AgentService) ListAgents(ctx context.Context, ownerID int64) ([]*Agent, error) {
|
||||
return s.store.ListAgentsByOwner(ctx, ownerID)
|
||||
}
|
||||
|
||||
// generateAPIKey creates a cryptographically random API key (32 bytes, hex encoded).
|
||||
func generateAPIKey() (string, error) {
|
||||
b := make([]byte, 32)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hex.EncodeToString(b), nil
|
||||
}
|
||||
@@ -0,0 +1,248 @@
|
||||
package agents
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/trace"
|
||||
)
|
||||
|
||||
func newTestService(t *testing.T) *AgentService {
|
||||
t.Helper()
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteAgentStore(db)
|
||||
tracer := trace.NewTracer(db)
|
||||
t.Cleanup(func() { tracer.Close() })
|
||||
|
||||
return NewAgentService(store, tracer)
|
||||
}
|
||||
|
||||
func TestAgentService_Register(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
agentName string
|
||||
displayName string
|
||||
agentType string
|
||||
capabilities json.RawMessage
|
||||
ownerID int64
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "successful registration",
|
||||
agentName: "test-bot",
|
||||
displayName: "Test Bot",
|
||||
agentType: "ai",
|
||||
capabilities: json.RawMessage(`{"skills":["testing"]}`),
|
||||
ownerID: 1,
|
||||
},
|
||||
{
|
||||
name: "empty name fails",
|
||||
agentName: "",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "invalid type",
|
||||
agentName: "invalid-type",
|
||||
agentType: "robot",
|
||||
ownerID: 1,
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "invalid capabilities JSON",
|
||||
agentName: "bad-caps",
|
||||
agentType: "ai",
|
||||
capabilities: json.RawMessage("not json"),
|
||||
ownerID: 1,
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "default type is ai",
|
||||
agentName: "default-type",
|
||||
ownerID: 1,
|
||||
},
|
||||
{
|
||||
name: "human type is valid",
|
||||
agentName: "human-agent",
|
||||
agentType: "human",
|
||||
ownerID: 1,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
svc := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
agent, apiKey, err := svc.Register(ctx, tt.agentName, tt.displayName, tt.agentType, tt.capabilities, tt.ownerID)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Fatalf("Register() error = %v, wantErr %v", err, tt.wantErr)
|
||||
}
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
if agent.ID == 0 {
|
||||
t.Error("agent ID should not be 0")
|
||||
}
|
||||
if apiKey == "" {
|
||||
t.Error("API key should not be empty")
|
||||
}
|
||||
if len(apiKey) < 32 {
|
||||
t.Errorf("API key too short: %d chars", len(apiKey))
|
||||
}
|
||||
if agent.Status != AgentStatusActive {
|
||||
t.Errorf("status = %s, want %s", agent.Status, AgentStatusActive)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentService_Authenticate(t *testing.T) {
|
||||
svc := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
_, apiKey, err := svc.Register(ctx, "auth-bot", "Auth Bot", "ai", nil, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("Register: %v", err)
|
||||
}
|
||||
|
||||
t.Run("valid key", func(t *testing.T) {
|
||||
agent, err := svc.Authenticate(ctx, apiKey)
|
||||
if err != nil {
|
||||
t.Fatalf("Authenticate: %v", err)
|
||||
}
|
||||
if agent.Name != "auth-bot" {
|
||||
t.Errorf("Name = %s, want auth-bot", agent.Name)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("invalid key", func(t *testing.T) {
|
||||
_, err := svc.Authenticate(ctx, "invalid-key")
|
||||
if err == nil {
|
||||
t.Error("expected error for invalid key")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestAgentService_GetAgent(t *testing.T) {
|
||||
svc := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
svc.Register(ctx, "get-bot", "Get Bot", "ai", nil, 1)
|
||||
|
||||
t.Run("existing agent", func(t *testing.T) {
|
||||
agent, err := svc.GetAgent(ctx, "get-bot")
|
||||
if err != nil {
|
||||
t.Fatalf("GetAgent: %v", err)
|
||||
}
|
||||
if agent.Name != "get-bot" {
|
||||
t.Errorf("Name = %s, want get-bot", agent.Name)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("non-existing agent", func(t *testing.T) {
|
||||
_, err := svc.GetAgent(ctx, "ghost")
|
||||
if err == nil {
|
||||
t.Error("expected error for non-existing agent")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestAgentService_UpdateAgent(t *testing.T) {
|
||||
svc := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
svc.Register(ctx, "update-bot", "Update Bot", "ai", json.RawMessage(`{"skills":["v1"]}`), 1)
|
||||
|
||||
updated, err := svc.UpdateAgent(ctx, "update-bot", "Updated Bot", json.RawMessage(`{"skills":["v1","v2"]}`))
|
||||
if err != nil {
|
||||
t.Fatalf("UpdateAgent: %v", err)
|
||||
}
|
||||
if updated.DisplayName != "Updated Bot" {
|
||||
t.Errorf("DisplayName = %s, want Updated Bot", updated.DisplayName)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentService_Deregister(t *testing.T) {
|
||||
svc := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
svc.Register(ctx, "dereg-bot", "Dereg Bot", "ai", nil, 1)
|
||||
|
||||
t.Run("owner can deregister", func(t *testing.T) {
|
||||
err := svc.Deregister(ctx, "dereg-bot", 1)
|
||||
if err != nil {
|
||||
t.Fatalf("Deregister: %v", err)
|
||||
}
|
||||
|
||||
_, err = svc.GetAgent(ctx, "dereg-bot")
|
||||
if err == nil {
|
||||
t.Error("expected error for deregistered agent")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("wrong owner cannot deregister", func(t *testing.T) {
|
||||
svc.Register(ctx, "other-bot", "Other Bot", "ai", nil, 1)
|
||||
err := svc.Deregister(ctx, "other-bot", 999)
|
||||
if err == nil {
|
||||
t.Error("expected error for wrong owner")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestAgentService_DiscoverAgents(t *testing.T) {
|
||||
svc := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
svc.Register(ctx, "search-bot", "Search Bot", "ai", json.RawMessage(`{"skills":["web-search"]}`), 1)
|
||||
svc.Register(ctx, "analyze-bot", "Analyze Bot", "ai", json.RawMessage(`{"skills":["sentiment"]}`), 1)
|
||||
|
||||
t.Run("find by capability", func(t *testing.T) {
|
||||
agents, err := svc.DiscoverAgents(ctx, "web-search")
|
||||
if err != nil {
|
||||
t.Fatalf("DiscoverAgents: %v", err)
|
||||
}
|
||||
if len(agents) != 1 {
|
||||
t.Errorf("got %d agents, want 1", len(agents))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("empty query returns all", func(t *testing.T) {
|
||||
agents, err := svc.DiscoverAgents(ctx, "")
|
||||
if err != nil {
|
||||
t.Fatalf("DiscoverAgents: %v", err)
|
||||
}
|
||||
if len(agents) != 2 {
|
||||
t.Errorf("got %d agents, want 2", len(agents))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("no match returns empty", func(t *testing.T) {
|
||||
agents, err := svc.DiscoverAgents(ctx, "quantum")
|
||||
if err != nil {
|
||||
t.Fatalf("DiscoverAgents: %v", err)
|
||||
}
|
||||
if len(agents) != 0 {
|
||||
t.Errorf("got %d agents, want 0", len(agents))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestAgentService_ListAgents(t *testing.T) {
|
||||
svc := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
svc.Register(ctx, "list-a", "Bot A", "ai", nil, 1)
|
||||
svc.Register(ctx, "list-b", "Bot B", "ai", nil, 1)
|
||||
|
||||
agents, err := svc.ListAgents(ctx, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("ListAgents: %v", err)
|
||||
}
|
||||
if len(agents) != 2 {
|
||||
t.Errorf("got %d agents, want 2", len(agents))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
package agents
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// AgentStore defines the storage interface for agent operations.
|
||||
type AgentStore interface {
|
||||
CreateAgent(ctx context.Context, agent *Agent) error
|
||||
GetAgentByName(ctx context.Context, name string) (*Agent, error)
|
||||
GetAgentByID(ctx context.Context, id int64) (*Agent, error)
|
||||
UpdateAgent(ctx context.Context, agent *Agent) error
|
||||
DeactivateAgent(ctx context.Context, name string) error
|
||||
ListActiveAgents(ctx context.Context) ([]*Agent, error)
|
||||
ListAgentsByOwner(ctx context.Context, ownerID int64) ([]*Agent, error)
|
||||
SearchAgentsByCapability(ctx context.Context, query string) ([]*Agent, error)
|
||||
}
|
||||
|
||||
// SQLiteAgentStore implements AgentStore using SQLite.
|
||||
type SQLiteAgentStore struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
// NewSQLiteAgentStore creates a new SQLite-backed agent store.
|
||||
func NewSQLiteAgentStore(db *sql.DB) *SQLiteAgentStore {
|
||||
return &SQLiteAgentStore{db: db}
|
||||
}
|
||||
|
||||
func (s *SQLiteAgentStore) CreateAgent(ctx context.Context, agent *Agent) error {
|
||||
caps := string(agent.Capabilities)
|
||||
if caps == "" {
|
||||
caps = "{}"
|
||||
}
|
||||
|
||||
result, err := s.db.ExecContext(ctx,
|
||||
`INSERT INTO agents (name, display_name, type, capabilities, owner_id, api_key_hash, status, created_at, updated_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`,
|
||||
agent.Name, agent.DisplayName, agent.Type, caps, agent.OwnerID, agent.APIKeyHash, AgentStatusActive,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("insert agent: %w", err)
|
||||
}
|
||||
id, err := result.LastInsertId()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get agent id: %w", err)
|
||||
}
|
||||
agent.ID = id
|
||||
agent.Status = AgentStatusActive
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SQLiteAgentStore) GetAgentByName(ctx context.Context, name string) (*Agent, error) {
|
||||
return s.scanAgent(s.db.QueryRowContext(ctx,
|
||||
`SELECT id, name, display_name, type, capabilities, owner_id, api_key_hash, status, created_at, updated_at
|
||||
FROM agents WHERE name = ? AND status = 'active'`, name,
|
||||
))
|
||||
}
|
||||
|
||||
func (s *SQLiteAgentStore) GetAgentByID(ctx context.Context, id int64) (*Agent, error) {
|
||||
return s.scanAgent(s.db.QueryRowContext(ctx,
|
||||
`SELECT id, name, display_name, type, capabilities, owner_id, api_key_hash, status, created_at, updated_at
|
||||
FROM agents WHERE id = ? AND status = 'active'`, id,
|
||||
))
|
||||
}
|
||||
|
||||
func (s *SQLiteAgentStore) UpdateAgent(ctx context.Context, agent *Agent) error {
|
||||
caps := string(agent.Capabilities)
|
||||
if caps == "" {
|
||||
caps = "{}"
|
||||
}
|
||||
|
||||
_, err := s.db.ExecContext(ctx,
|
||||
`UPDATE agents SET display_name = ?, capabilities = ?, api_key_hash = ?, updated_at = CURRENT_TIMESTAMP
|
||||
WHERE name = ? AND status = 'active'`,
|
||||
agent.DisplayName, caps, agent.APIKeyHash, agent.Name,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *SQLiteAgentStore) DeactivateAgent(ctx context.Context, name string) error {
|
||||
result, err := s.db.ExecContext(ctx,
|
||||
`UPDATE agents SET status = 'inactive', updated_at = CURRENT_TIMESTAMP
|
||||
WHERE name = ? AND status = 'active'`,
|
||||
name,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("deactivate agent: %w", err)
|
||||
}
|
||||
rowsAffected, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if rowsAffected == 0 {
|
||||
return fmt.Errorf("agent not found: %s", name)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SQLiteAgentStore) ListActiveAgents(ctx context.Context) ([]*Agent, error) {
|
||||
rows, err := s.db.QueryContext(ctx,
|
||||
`SELECT id, name, display_name, type, capabilities, owner_id, api_key_hash, status, created_at, updated_at
|
||||
FROM agents WHERE status = 'active' ORDER BY name`,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
return s.scanAgents(rows)
|
||||
}
|
||||
|
||||
func (s *SQLiteAgentStore) ListAgentsByOwner(ctx context.Context, ownerID int64) ([]*Agent, error) {
|
||||
rows, err := s.db.QueryContext(ctx,
|
||||
`SELECT id, name, display_name, type, capabilities, owner_id, api_key_hash, status, created_at, updated_at
|
||||
FROM agents WHERE owner_id = ? AND status = 'active' ORDER BY name`,
|
||||
ownerID,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
return s.scanAgents(rows)
|
||||
}
|
||||
|
||||
func (s *SQLiteAgentStore) SearchAgentsByCapability(ctx context.Context, query string) ([]*Agent, error) {
|
||||
// Simple LIKE search on the capabilities JSON field
|
||||
rows, err := s.db.QueryContext(ctx,
|
||||
`SELECT id, name, display_name, type, capabilities, owner_id, api_key_hash, status, created_at, updated_at
|
||||
FROM agents WHERE status = 'active' AND capabilities LIKE ? ORDER BY name`,
|
||||
"%"+query+"%",
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
return s.scanAgents(rows)
|
||||
}
|
||||
|
||||
func (s *SQLiteAgentStore) scanAgent(row *sql.Row) (*Agent, error) {
|
||||
var agent Agent
|
||||
var caps string
|
||||
err := row.Scan(
|
||||
&agent.ID, &agent.Name, &agent.DisplayName, &agent.Type,
|
||||
&caps, &agent.OwnerID, &agent.APIKeyHash, &agent.Status,
|
||||
&agent.CreatedAt, &agent.UpdatedAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
agent.Capabilities = json.RawMessage(caps)
|
||||
return &agent, nil
|
||||
}
|
||||
|
||||
func (s *SQLiteAgentStore) scanAgents(rows *sql.Rows) ([]*Agent, error) {
|
||||
var agents []*Agent
|
||||
for rows.Next() {
|
||||
var agent Agent
|
||||
var caps string
|
||||
err := rows.Scan(
|
||||
&agent.ID, &agent.Name, &agent.DisplayName, &agent.Type,
|
||||
&caps, &agent.OwnerID, &agent.APIKeyHash, &agent.Status,
|
||||
&agent.CreatedAt, &agent.UpdatedAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
agent.Capabilities = json.RawMessage(caps)
|
||||
agents = append(agents, &agent)
|
||||
}
|
||||
if agents == nil {
|
||||
agents = []*Agent{}
|
||||
}
|
||||
return agents, rows.Err()
|
||||
}
|
||||
@@ -0,0 +1,250 @@
|
||||
package agents
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/storage"
|
||||
)
|
||||
|
||||
func newTestDB(t *testing.T) *sql.DB {
|
||||
t.Helper()
|
||||
dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())
|
||||
db, err := sql.Open("sqlite", dsn)
|
||||
if err != nil {
|
||||
t.Fatalf("open database: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { db.Close() })
|
||||
|
||||
if _, err := db.Exec("PRAGMA foreign_keys=ON"); err != nil {
|
||||
t.Fatalf("enable foreign keys: %v", err)
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
if err := storage.RunMigrations(ctx, db); err != nil {
|
||||
t.Fatalf("run migrations: %v", err)
|
||||
}
|
||||
|
||||
// Seed a test user for owner_id FK
|
||||
db.Exec(`INSERT OR IGNORE INTO users (id, username, password_hash, display_name) VALUES (1, 'testowner', 'hash', 'Test Owner')`)
|
||||
|
||||
return db
|
||||
}
|
||||
|
||||
func TestSQLiteAgentStore_CreateAndGet(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteAgentStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
agent := &Agent{
|
||||
Name: "test-bot",
|
||||
DisplayName: "Test Bot",
|
||||
Type: "ai",
|
||||
Capabilities: json.RawMessage(`{"skills":["testing"]}`),
|
||||
OwnerID: 1,
|
||||
APIKeyHash: "somehash",
|
||||
}
|
||||
|
||||
if err := store.CreateAgent(ctx, agent); err != nil {
|
||||
t.Fatalf("CreateAgent: %v", err)
|
||||
}
|
||||
|
||||
if agent.ID == 0 {
|
||||
t.Error("agent ID should not be 0")
|
||||
}
|
||||
|
||||
// Get by name
|
||||
got, err := store.GetAgentByName(ctx, "test-bot")
|
||||
if err != nil {
|
||||
t.Fatalf("GetAgentByName: %v", err)
|
||||
}
|
||||
if got.DisplayName != "Test Bot" {
|
||||
t.Errorf("DisplayName = %q, want %q", got.DisplayName, "Test Bot")
|
||||
}
|
||||
|
||||
// Get by ID
|
||||
got2, err := store.GetAgentByID(ctx, agent.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetAgentByID: %v", err)
|
||||
}
|
||||
if got2.Name != "test-bot" {
|
||||
t.Errorf("Name = %q, want %q", got2.Name, "test-bot")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteAgentStore_DuplicateName(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteAgentStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
agent := &Agent{
|
||||
Name: "dup-bot",
|
||||
DisplayName: "Dup Bot",
|
||||
Type: "ai",
|
||||
Capabilities: json.RawMessage("{}"),
|
||||
OwnerID: 1,
|
||||
APIKeyHash: "hash1",
|
||||
}
|
||||
|
||||
if err := store.CreateAgent(ctx, agent); err != nil {
|
||||
t.Fatalf("CreateAgent: %v", err)
|
||||
}
|
||||
|
||||
agent2 := &Agent{
|
||||
Name: "dup-bot",
|
||||
DisplayName: "Dup Bot 2",
|
||||
Type: "ai",
|
||||
Capabilities: json.RawMessage("{}"),
|
||||
OwnerID: 1,
|
||||
APIKeyHash: "hash2",
|
||||
}
|
||||
|
||||
err := store.CreateAgent(ctx, agent2)
|
||||
if err == nil {
|
||||
t.Error("expected error for duplicate name")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteAgentStore_Update(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteAgentStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
agent := &Agent{
|
||||
Name: "update-bot",
|
||||
DisplayName: "Update Bot",
|
||||
Type: "ai",
|
||||
Capabilities: json.RawMessage(`{"skills":["v1"]}`),
|
||||
OwnerID: 1,
|
||||
APIKeyHash: "hash",
|
||||
}
|
||||
|
||||
store.CreateAgent(ctx, agent)
|
||||
|
||||
agent.DisplayName = "Updated Bot"
|
||||
agent.Capabilities = json.RawMessage(`{"skills":["v1","v2"]}`)
|
||||
|
||||
if err := store.UpdateAgent(ctx, agent); err != nil {
|
||||
t.Fatalf("UpdateAgent: %v", err)
|
||||
}
|
||||
|
||||
got, err := store.GetAgentByName(ctx, "update-bot")
|
||||
if err != nil {
|
||||
t.Fatalf("GetAgentByName: %v", err)
|
||||
}
|
||||
if got.DisplayName != "Updated Bot" {
|
||||
t.Errorf("DisplayName = %q, want %q", got.DisplayName, "Updated Bot")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteAgentStore_Deactivate(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteAgentStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
agent := &Agent{
|
||||
Name: "deactivate-bot",
|
||||
DisplayName: "Deactivate Bot",
|
||||
Type: "ai",
|
||||
Capabilities: json.RawMessage("{}"),
|
||||
OwnerID: 1,
|
||||
APIKeyHash: "hash",
|
||||
}
|
||||
|
||||
store.CreateAgent(ctx, agent)
|
||||
|
||||
if err := store.DeactivateAgent(ctx, "deactivate-bot"); err != nil {
|
||||
t.Fatalf("DeactivateAgent: %v", err)
|
||||
}
|
||||
|
||||
// Should not be found (GetByName filters active only)
|
||||
_, err := store.GetAgentByName(ctx, "deactivate-bot")
|
||||
if err == nil {
|
||||
t.Error("expected error for deactivated agent")
|
||||
}
|
||||
|
||||
// Deactivate non-existent
|
||||
err = store.DeactivateAgent(ctx, "nonexistent")
|
||||
if err == nil {
|
||||
t.Error("expected error for non-existent agent")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteAgentStore_ListActive(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteAgentStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
for _, name := range []string{"bot-a", "bot-b", "bot-c"} {
|
||||
store.CreateAgent(ctx, &Agent{
|
||||
Name: name,
|
||||
DisplayName: name,
|
||||
Type: "ai",
|
||||
Capabilities: json.RawMessage("{}"),
|
||||
OwnerID: 1,
|
||||
APIKeyHash: "hash",
|
||||
})
|
||||
}
|
||||
|
||||
agents, err := store.ListActiveAgents(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("ListActiveAgents: %v", err)
|
||||
}
|
||||
if len(agents) != 3 {
|
||||
t.Errorf("got %d agents, want 3", len(agents))
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteAgentStore_SearchByCapability(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteAgentStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
store.CreateAgent(ctx, &Agent{
|
||||
Name: "searcher",
|
||||
DisplayName: "Searcher",
|
||||
Type: "ai",
|
||||
Capabilities: json.RawMessage(`{"skills":["web-search","summarization"]}`),
|
||||
OwnerID: 1,
|
||||
APIKeyHash: "hash",
|
||||
})
|
||||
|
||||
store.CreateAgent(ctx, &Agent{
|
||||
Name: "analyzer",
|
||||
DisplayName: "Analyzer",
|
||||
Type: "ai",
|
||||
Capabilities: json.RawMessage(`{"skills":["sentiment-analysis"]}`),
|
||||
OwnerID: 1,
|
||||
APIKeyHash: "hash",
|
||||
})
|
||||
|
||||
t.Run("match found", func(t *testing.T) {
|
||||
results, err := store.SearchAgentsByCapability(ctx, "web-search")
|
||||
if err != nil {
|
||||
t.Fatalf("SearchAgentsByCapability: %v", err)
|
||||
}
|
||||
if len(results) != 1 {
|
||||
t.Errorf("got %d results, want 1", len(results))
|
||||
}
|
||||
if len(results) > 0 && results[0].Name != "searcher" {
|
||||
t.Errorf("Name = %s, want searcher", results[0].Name)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("no match", func(t *testing.T) {
|
||||
results, err := store.SearchAgentsByCapability(ctx, "quantum-computing")
|
||||
if err != nil {
|
||||
t.Fatalf("SearchAgentsByCapability: %v", err)
|
||||
}
|
||||
if len(results) != 0 {
|
||||
t.Errorf("got %d results, want 0", len(results))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
var _ = storage.RunMigrations
|
||||
@@ -0,0 +1,27 @@
|
||||
// Package agents provides agent registry and authentication for SynapBus.
|
||||
package agents
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Agent status constants.
|
||||
const (
|
||||
AgentStatusActive = "active"
|
||||
AgentStatusInactive = "inactive"
|
||||
)
|
||||
|
||||
// Agent represents a registered entity that can send/receive messages.
|
||||
type Agent struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
DisplayName string `json:"display_name"`
|
||||
Type string `json:"type"`
|
||||
Capabilities json.RawMessage `json:"capabilities"`
|
||||
OwnerID int64 `json:"owner_id"`
|
||||
APIKeyHash string `json:"-"`
|
||||
Status string `json:"status"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/agents"
|
||||
)
|
||||
|
||||
type mcpContextKey string
|
||||
|
||||
const agentNameKey mcpContextKey = "mcp_agent_name"
|
||||
|
||||
// AgentNameFromContext extracts the agent name from the MCP context.
|
||||
func AgentNameFromContext(ctx context.Context) (string, bool) {
|
||||
name, ok := ctx.Value(agentNameKey).(string)
|
||||
return name, ok
|
||||
}
|
||||
|
||||
// ContextWithAgentName returns a new context with the agent name set.
|
||||
func ContextWithAgentName(ctx context.Context, name string) context.Context {
|
||||
return context.WithValue(ctx, agentNameKey, name)
|
||||
}
|
||||
|
||||
// extractAgentName gets the agent name from the request context.
|
||||
// It first checks for the MCP-level agent name, then falls back to the
|
||||
// agents package context (set by HTTP auth middleware).
|
||||
func extractAgentName(ctx context.Context) (string, bool) {
|
||||
if name, ok := AgentNameFromContext(ctx); ok {
|
||||
return name, true
|
||||
}
|
||||
if agent, ok := agents.AgentFromContext(ctx); ok {
|
||||
return agent.Name, true
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
@@ -0,0 +1,78 @@
|
||||
// Package mcp provides the MCP server implementation for SynapBus.
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Connection represents an active agent connection.
|
||||
type Connection struct {
|
||||
ID string `json:"id"`
|
||||
AgentName string `json:"agent_name"`
|
||||
Transport string `json:"transport"`
|
||||
ConnectedAt time.Time `json:"connected_at"`
|
||||
LastActivity time.Time `json:"last_activity"`
|
||||
}
|
||||
|
||||
// ConnectionManager tracks active MCP connections (thread-safe).
|
||||
type ConnectionManager struct {
|
||||
mu sync.RWMutex
|
||||
conns map[string]*Connection
|
||||
}
|
||||
|
||||
// NewConnectionManager creates a new ConnectionManager.
|
||||
func NewConnectionManager() *ConnectionManager {
|
||||
return &ConnectionManager{
|
||||
conns: make(map[string]*Connection),
|
||||
}
|
||||
}
|
||||
|
||||
// Add registers a new connection.
|
||||
func (cm *ConnectionManager) Add(conn *Connection) {
|
||||
cm.mu.Lock()
|
||||
defer cm.mu.Unlock()
|
||||
cm.conns[conn.ID] = conn
|
||||
}
|
||||
|
||||
// Remove unregisters a connection.
|
||||
func (cm *ConnectionManager) Remove(id string) {
|
||||
cm.mu.Lock()
|
||||
defer cm.mu.Unlock()
|
||||
delete(cm.conns, id)
|
||||
}
|
||||
|
||||
// Get returns a connection by ID.
|
||||
func (cm *ConnectionManager) Get(id string) (*Connection, bool) {
|
||||
cm.mu.RLock()
|
||||
defer cm.mu.RUnlock()
|
||||
conn, ok := cm.conns[id]
|
||||
return conn, ok
|
||||
}
|
||||
|
||||
// Count returns the number of active connections.
|
||||
func (cm *ConnectionManager) Count() int {
|
||||
cm.mu.RLock()
|
||||
defer cm.mu.RUnlock()
|
||||
return len(cm.conns)
|
||||
}
|
||||
|
||||
// List returns all active connections.
|
||||
func (cm *ConnectionManager) List() []*Connection {
|
||||
cm.mu.RLock()
|
||||
defer cm.mu.RUnlock()
|
||||
result := make([]*Connection, 0, len(cm.conns))
|
||||
for _, conn := range cm.conns {
|
||||
result = append(result, conn)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// UpdateActivity updates the last activity timestamp for a connection.
|
||||
func (cm *ConnectionManager) UpdateActivity(id string) {
|
||||
cm.mu.Lock()
|
||||
defer cm.mu.Unlock()
|
||||
if conn, ok := cm.conns[id]; ok {
|
||||
conn.LastActivity = time.Now()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"time"
|
||||
)
|
||||
|
||||
// HealthStatus represents the health check response.
|
||||
type HealthStatus struct {
|
||||
Status string `json:"status"`
|
||||
Version string `json:"version"`
|
||||
Uptime string `json:"uptime"`
|
||||
ActiveConnections int `json:"active_connections"`
|
||||
}
|
||||
|
||||
// NewHealthHandler creates an HTTP handler for health checks.
|
||||
func NewHealthHandler(connMgr *ConnectionManager, version string, startTime time.Time) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
status := HealthStatus{
|
||||
Status: "ok",
|
||||
Version: version,
|
||||
Uptime: time.Since(startTime).Round(time.Second).String(),
|
||||
ActiveConnections: connMgr.Count(),
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
json.NewEncoder(w).Encode(status)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,81 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
|
||||
"github.com/mark3labs/mcp-go/server"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/agents"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/messaging"
|
||||
)
|
||||
|
||||
// MCPServer wraps the mcp-go server with SynapBus services.
|
||||
type MCPServer struct {
|
||||
mcpServer *server.MCPServer
|
||||
sseServer *server.SSEServer
|
||||
connMgr *ConnectionManager
|
||||
agentService *agents.AgentService
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewMCPServer creates and configures a new MCP server with all tools registered.
|
||||
func NewMCPServer(
|
||||
msgService *messaging.MessagingService,
|
||||
agentService *agents.AgentService,
|
||||
) *MCPServer {
|
||||
logger := slog.Default().With("component", "mcp-server")
|
||||
|
||||
// Create the mcp-go server
|
||||
mcpSrv := server.NewMCPServer(
|
||||
"SynapBus",
|
||||
"0.1.0",
|
||||
server.WithToolCapabilities(true),
|
||||
)
|
||||
|
||||
// Register all tools
|
||||
registrar := NewToolRegistrar(msgService, agentService)
|
||||
registrar.RegisterAll(mcpSrv)
|
||||
|
||||
// Create SSE transport with context func for auth propagation
|
||||
sseServer := server.NewSSEServer(mcpSrv,
|
||||
server.WithSSEContextFunc(func(ctx context.Context, r *http.Request) context.Context {
|
||||
// Propagate agent identity from HTTP auth to MCP context
|
||||
if agent, ok := agents.AgentFromContext(r.Context()); ok {
|
||||
return ContextWithAgentName(ctx, agent.Name)
|
||||
}
|
||||
return ctx
|
||||
}),
|
||||
)
|
||||
|
||||
s := &MCPServer{
|
||||
mcpServer: mcpSrv,
|
||||
sseServer: sseServer,
|
||||
connMgr: NewConnectionManager(),
|
||||
agentService: agentService,
|
||||
logger: logger,
|
||||
}
|
||||
|
||||
logger.Info("MCP server initialized")
|
||||
return s
|
||||
}
|
||||
|
||||
// SSEHandler returns the SSE handler for mounting on a router.
|
||||
func (s *MCPServer) SSEHandler() http.Handler {
|
||||
return s.sseServer
|
||||
}
|
||||
|
||||
// ConnectionManager returns the connection manager.
|
||||
func (s *MCPServer) ConnectionManager() *ConnectionManager {
|
||||
return s.connMgr
|
||||
}
|
||||
|
||||
// Shutdown gracefully shuts down the MCP server.
|
||||
func (s *MCPServer) Shutdown(ctx context.Context) error {
|
||||
s.logger.Info("MCP server shutting down")
|
||||
if err := s.sseServer.Shutdown(ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,408 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
"github.com/mark3labs/mcp-go/server"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/agents"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/messaging"
|
||||
)
|
||||
|
||||
// ToolRegistrar registers all SynapBus MCP tools on the given server.
|
||||
type ToolRegistrar struct {
|
||||
msgService *messaging.MessagingService
|
||||
agentService *agents.AgentService
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewToolRegistrar creates a new tool registrar.
|
||||
func NewToolRegistrar(msgService *messaging.MessagingService, agentService *agents.AgentService) *ToolRegistrar {
|
||||
return &ToolRegistrar{
|
||||
msgService: msgService,
|
||||
agentService: agentService,
|
||||
logger: slog.Default().With("component", "mcp-tools"),
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterAll registers all tools on the MCP server.
|
||||
func (tr *ToolRegistrar) RegisterAll(s *server.MCPServer) {
|
||||
s.AddTool(tr.sendMessageTool(), tr.handleSendMessage)
|
||||
s.AddTool(tr.readInboxTool(), tr.handleReadInbox)
|
||||
s.AddTool(tr.claimMessagesTool(), tr.handleClaimMessages)
|
||||
s.AddTool(tr.markDoneTool(), tr.handleMarkDone)
|
||||
s.AddTool(tr.searchMessagesTool(), tr.handleSearchMessages)
|
||||
s.AddTool(tr.registerAgentTool(), tr.handleRegisterAgent)
|
||||
s.AddTool(tr.discoverAgentsTool(), tr.handleDiscoverAgents)
|
||||
s.AddTool(tr.updateAgentTool(), tr.handleUpdateAgent)
|
||||
s.AddTool(tr.deregisterAgentTool(), tr.handleDeregisterAgent)
|
||||
|
||||
tr.logger.Info("all MCP tools registered", "count", 9)
|
||||
}
|
||||
|
||||
// --- Tool Definitions ---
|
||||
|
||||
func (tr *ToolRegistrar) sendMessageTool() mcp.Tool {
|
||||
return mcp.NewTool("send_message",
|
||||
mcp.WithDescription("Send a direct message to another agent or to a channel"),
|
||||
mcp.WithString("to", mcp.Description("Name of the recipient agent"), mcp.Required()),
|
||||
mcp.WithString("body", mcp.Description("Message body text"), mcp.Required()),
|
||||
mcp.WithString("subject", mcp.Description("Conversation subject (optional)")),
|
||||
mcp.WithNumber("priority", mcp.Description("Message priority (1-10, default 5)"), mcp.Min(1), mcp.Max(10)),
|
||||
mcp.WithString("metadata", mcp.Description("JSON metadata object (optional)")),
|
||||
mcp.WithNumber("channel_id", mcp.Description("Channel ID for channel messages (optional)")),
|
||||
)
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) readInboxTool() mcp.Tool {
|
||||
return mcp.NewTool("read_inbox",
|
||||
mcp.WithDescription("Read messages from the authenticated agent's inbox"),
|
||||
mcp.WithNumber("limit", mcp.Description("Maximum number of messages to return (default 50)")),
|
||||
mcp.WithString("status_filter", mcp.Description("Filter by message status: pending, processing, done, failed")),
|
||||
mcp.WithBoolean("include_read", mcp.Description("Include previously read messages (default false)")),
|
||||
mcp.WithNumber("min_priority", mcp.Description("Minimum priority filter (1-10)")),
|
||||
mcp.WithString("from_agent", mcp.Description("Filter by sender agent name")),
|
||||
)
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) claimMessagesTool() mcp.Tool {
|
||||
return mcp.NewTool("claim_messages",
|
||||
mcp.WithDescription("Atomically claim pending messages for processing"),
|
||||
mcp.WithNumber("limit", mcp.Description("Maximum number of messages to claim (default 10)")),
|
||||
)
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) markDoneTool() mcp.Tool {
|
||||
return mcp.NewTool("mark_done",
|
||||
mcp.WithDescription("Mark a claimed message as done or failed"),
|
||||
mcp.WithNumber("message_id", mcp.Description("ID of the message to mark"), mcp.Required()),
|
||||
mcp.WithString("status", mcp.Description("New status: 'done' or 'failed' (default 'done')")),
|
||||
mcp.WithString("reason", mcp.Description("Failure reason (only for status='failed')")),
|
||||
)
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) searchMessagesTool() mcp.Tool {
|
||||
return mcp.NewTool("search_messages",
|
||||
mcp.WithDescription("Search messages using full-text search"),
|
||||
mcp.WithString("query", mcp.Description("Search query string")),
|
||||
mcp.WithNumber("limit", mcp.Description("Maximum results to return (default 20)")),
|
||||
mcp.WithNumber("min_priority", mcp.Description("Minimum priority filter (1-10)")),
|
||||
mcp.WithString("from_agent", mcp.Description("Filter by sender agent name")),
|
||||
mcp.WithString("status", mcp.Description("Filter by message status")),
|
||||
)
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) registerAgentTool() mcp.Tool {
|
||||
return mcp.NewTool("register_agent",
|
||||
mcp.WithDescription("Register a new agent and receive an API key"),
|
||||
mcp.WithString("name", mcp.Description("Unique agent name"), mcp.Required()),
|
||||
mcp.WithString("display_name", mcp.Description("Human-readable display name")),
|
||||
mcp.WithString("type", mcp.Description("Agent type: 'ai' or 'human' (default 'ai')")),
|
||||
mcp.WithString("capabilities", mcp.Description("JSON capabilities object")),
|
||||
)
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) discoverAgentsTool() mcp.Tool {
|
||||
return mcp.NewTool("discover_agents",
|
||||
mcp.WithDescription("Discover agents by capability keywords"),
|
||||
mcp.WithString("query", mcp.Description("Capability keyword to search for")),
|
||||
)
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) updateAgentTool() mcp.Tool {
|
||||
return mcp.NewTool("update_agent",
|
||||
mcp.WithDescription("Update the authenticated agent's display name or capabilities"),
|
||||
mcp.WithString("display_name", mcp.Description("New display name")),
|
||||
mcp.WithString("capabilities", mcp.Description("New JSON capabilities object")),
|
||||
)
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) deregisterAgentTool() mcp.Tool {
|
||||
return mcp.NewTool("deregister_agent",
|
||||
mcp.WithDescription("Deregister the authenticated agent (soft delete)"),
|
||||
)
|
||||
}
|
||||
|
||||
// --- Tool Handlers ---
|
||||
|
||||
func (tr *ToolRegistrar) handleSendMessage(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
to := req.GetString("to", "")
|
||||
body := req.GetString("body", "")
|
||||
subject := req.GetString("subject", "")
|
||||
priority := req.GetInt("priority", 5)
|
||||
metadataStr := req.GetString("metadata", "")
|
||||
|
||||
if to == "" {
|
||||
return mcp.NewToolResultError("'to' parameter is required"), nil
|
||||
}
|
||||
if body == "" {
|
||||
return mcp.NewToolResultError("'body' parameter is required"), nil
|
||||
}
|
||||
|
||||
var channelID *int64
|
||||
if cid := req.GetInt("channel_id", 0); cid > 0 {
|
||||
v := int64(cid)
|
||||
channelID = &v
|
||||
}
|
||||
|
||||
opts := messaging.SendOptions{
|
||||
Subject: subject,
|
||||
Priority: priority,
|
||||
Metadata: metadataStr,
|
||||
ChannelID: channelID,
|
||||
}
|
||||
|
||||
msg, err := tr.msgService.SendMessage(ctx, agentName, to, body, opts)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("send_message failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"message_id": msg.ID,
|
||||
"conversation_id": msg.ConversationID,
|
||||
"status": msg.Status,
|
||||
})
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) handleReadInbox(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
opts := messaging.ReadOptions{
|
||||
Limit: req.GetInt("limit", 50),
|
||||
Status: req.GetString("status_filter", ""),
|
||||
MinPriority: req.GetInt("min_priority", 0),
|
||||
FromAgent: req.GetString("from_agent", ""),
|
||||
}
|
||||
|
||||
// Handle include_read boolean
|
||||
args := req.GetArguments()
|
||||
if v, ok := args["include_read"]; ok {
|
||||
if b, ok := v.(bool); ok {
|
||||
opts.IncludeRead = b
|
||||
}
|
||||
}
|
||||
|
||||
messages, err := tr.msgService.ReadInbox(ctx, agentName, opts)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("read_inbox failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"messages": messages,
|
||||
"count": len(messages),
|
||||
})
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) handleClaimMessages(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
limit := req.GetInt("limit", 10)
|
||||
|
||||
messages, err := tr.msgService.ClaimMessages(ctx, agentName, limit)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("claim_messages failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"messages": messages,
|
||||
"count": len(messages),
|
||||
})
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) handleMarkDone(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
messageID, err := req.RequireInt("message_id")
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError("'message_id' parameter is required"), nil
|
||||
}
|
||||
|
||||
status := req.GetString("status", "done")
|
||||
reason := req.GetString("reason", "")
|
||||
|
||||
switch status {
|
||||
case "done":
|
||||
if err := tr.msgService.MarkDone(ctx, int64(messageID), agentName); err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("mark_done failed: %s", err)), nil
|
||||
}
|
||||
case "failed":
|
||||
if err := tr.msgService.MarkFailed(ctx, int64(messageID), agentName, reason); err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("mark_failed failed: %s", err)), nil
|
||||
}
|
||||
default:
|
||||
return mcp.NewToolResultError("status must be 'done' or 'failed'"), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"message_id": messageID,
|
||||
"status": status,
|
||||
})
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) handleSearchMessages(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
query := req.GetString("query", "")
|
||||
opts := messaging.SearchOptions{
|
||||
Limit: req.GetInt("limit", 20),
|
||||
MinPriority: req.GetInt("min_priority", 0),
|
||||
FromAgent: req.GetString("from_agent", ""),
|
||||
Status: req.GetString("status", ""),
|
||||
}
|
||||
|
||||
messages, err := tr.msgService.SearchMessages(ctx, agentName, query, opts)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("search_messages failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"messages": messages,
|
||||
"count": len(messages),
|
||||
})
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) handleRegisterAgent(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
// register_agent does not require authentication
|
||||
name := req.GetString("name", "")
|
||||
if name == "" {
|
||||
return mcp.NewToolResultError("'name' parameter is required"), nil
|
||||
}
|
||||
|
||||
displayName := req.GetString("display_name", name)
|
||||
agentType := req.GetString("type", "ai")
|
||||
capsStr := req.GetString("capabilities", "{}")
|
||||
|
||||
var caps json.RawMessage
|
||||
if capsStr != "" {
|
||||
if !json.Valid([]byte(capsStr)) {
|
||||
return mcp.NewToolResultError("capabilities must be valid JSON"), nil
|
||||
}
|
||||
caps = json.RawMessage(capsStr)
|
||||
}
|
||||
|
||||
// Use owner_id=1 as default (first user). In production, this would come from auth.
|
||||
agent, apiKey, err := tr.agentService.Register(ctx, name, displayName, agentType, caps, 1)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("register_agent failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"agent_id": agent.ID,
|
||||
"name": agent.Name,
|
||||
"api_key": apiKey,
|
||||
"created_at": agent.CreatedAt,
|
||||
})
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) handleDiscoverAgents(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
query := req.GetString("query", "")
|
||||
_ = agentName // just verifying auth
|
||||
|
||||
agentsList, err := tr.agentService.DiscoverAgents(ctx, query)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("discover_agents failed: %s", err)), nil
|
||||
}
|
||||
|
||||
// Strip sensitive fields
|
||||
result := make([]map[string]any, len(agentsList))
|
||||
for i, a := range agentsList {
|
||||
result[i] = map[string]any{
|
||||
"name": a.Name,
|
||||
"display_name": a.DisplayName,
|
||||
"type": a.Type,
|
||||
"capabilities": a.Capabilities,
|
||||
"status": a.Status,
|
||||
}
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"agents": result,
|
||||
"count": len(result),
|
||||
})
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) handleUpdateAgent(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
displayName := req.GetString("display_name", "")
|
||||
capsStr := req.GetString("capabilities", "")
|
||||
|
||||
var caps json.RawMessage
|
||||
if capsStr != "" {
|
||||
if !json.Valid([]byte(capsStr)) {
|
||||
return mcp.NewToolResultError("capabilities must be valid JSON"), nil
|
||||
}
|
||||
caps = json.RawMessage(capsStr)
|
||||
}
|
||||
|
||||
agent, err := tr.agentService.UpdateAgent(ctx, agentName, displayName, caps)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("update_agent failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"name": agent.Name,
|
||||
"display_name": agent.DisplayName,
|
||||
"capabilities": agent.Capabilities,
|
||||
})
|
||||
}
|
||||
|
||||
func (tr *ToolRegistrar) handleDeregisterAgent(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
// Get the agent to find owner_id
|
||||
agent, err := tr.agentService.GetAgent(ctx, agentName)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("deregister_agent failed: %s", err)), nil
|
||||
}
|
||||
|
||||
if err := tr.agentService.Deregister(ctx, agentName, agent.OwnerID); err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("deregister_agent failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"name": agentName,
|
||||
"status": "deregistered",
|
||||
})
|
||||
}
|
||||
|
||||
// resultJSON marshals data to a JSON text MCP result.
|
||||
func resultJSON(data any) (*mcp.CallToolResult, error) {
|
||||
b, err := json.Marshal(data)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to marshal response: %s", err)), nil
|
||||
}
|
||||
return mcp.NewToolResultText(string(b)), nil
|
||||
}
|
||||
@@ -0,0 +1,390 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
mcplib "github.com/mark3labs/mcp-go/mcp"
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/agents"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/messaging"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/storage"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/trace"
|
||||
)
|
||||
|
||||
func newTestDB(t *testing.T) *sql.DB {
|
||||
t.Helper()
|
||||
dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())
|
||||
db, err := sql.Open("sqlite", dsn)
|
||||
if err != nil {
|
||||
t.Fatalf("open database: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { db.Close() })
|
||||
|
||||
if _, err := db.Exec("PRAGMA foreign_keys=ON"); err != nil {
|
||||
t.Fatalf("enable foreign keys: %v", err)
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
if err := storage.RunMigrations(ctx, db); err != nil {
|
||||
t.Fatalf("run migrations: %v", err)
|
||||
}
|
||||
|
||||
// Seed test user
|
||||
db.Exec(`INSERT OR IGNORE INTO users (id, username, password_hash, display_name) VALUES (1, 'testowner', 'hash', 'Test Owner')`)
|
||||
|
||||
return db
|
||||
}
|
||||
|
||||
func newTestRegistrar(t *testing.T) (*ToolRegistrar, *messaging.MessagingService, *agents.AgentService, *sql.DB) {
|
||||
t.Helper()
|
||||
db := newTestDB(t)
|
||||
|
||||
tracer := trace.NewTracer(db)
|
||||
t.Cleanup(func() { tracer.Close() })
|
||||
|
||||
msgStore := messaging.NewSQLiteMessageStore(db)
|
||||
msgService := messaging.NewMessagingService(msgStore, tracer)
|
||||
|
||||
agentStore := agents.NewSQLiteAgentStore(db)
|
||||
agentService := agents.NewAgentService(agentStore, tracer)
|
||||
|
||||
registrar := NewToolRegistrar(msgService, agentService)
|
||||
return registrar, msgService, agentService, db
|
||||
}
|
||||
|
||||
func makeRequest(args map[string]any) mcplib.CallToolRequest {
|
||||
return mcplib.CallToolRequest{
|
||||
Params: mcplib.CallToolParams{
|
||||
Arguments: args,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolHandler_RegisterAgent(t *testing.T) {
|
||||
tr, _, _, _ := newTestRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("successful registration", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"name": "test-agent",
|
||||
"display_name": "Test Agent",
|
||||
"type": "ai",
|
||||
})
|
||||
|
||||
result, err := tr.handleRegisterAgent(ctx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleRegisterAgent: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
// Parse response
|
||||
var resp map[string]any
|
||||
text := result.Content[0].(mcplib.TextContent).Text
|
||||
if err := json.Unmarshal([]byte(text), &resp); err != nil {
|
||||
t.Fatalf("unmarshal response: %v", err)
|
||||
}
|
||||
if resp["api_key"] == nil || resp["api_key"] == "" {
|
||||
t.Error("expected api_key in response")
|
||||
}
|
||||
if resp["name"] != "test-agent" {
|
||||
t.Errorf("name = %v, want test-agent", resp["name"])
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing name", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{})
|
||||
|
||||
result, err := tr.handleRegisterAgent(ctx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleRegisterAgent: %v", err)
|
||||
}
|
||||
if !result.IsError {
|
||||
t.Error("expected error for missing name")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestToolHandler_SendMessage(t *testing.T) {
|
||||
tr, _, agentSvc, _ := newTestRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Register agents
|
||||
agentSvc.Register(ctx, "sender", "Sender", "ai", nil, 1)
|
||||
agentSvc.Register(ctx, "receiver", "Receiver", "ai", nil, 1)
|
||||
|
||||
// Set up authenticated context
|
||||
authCtx := ContextWithAgentName(ctx, "sender")
|
||||
|
||||
t.Run("successful send", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"to": "receiver",
|
||||
"body": "Hello from test",
|
||||
})
|
||||
|
||||
result, err := tr.handleSendMessage(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleSendMessage: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
var resp map[string]any
|
||||
text := result.Content[0].(mcplib.TextContent).Text
|
||||
json.Unmarshal([]byte(text), &resp)
|
||||
if resp["message_id"] == nil {
|
||||
t.Error("expected message_id in response")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing to", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"body": "no recipient",
|
||||
})
|
||||
|
||||
result, _ := tr.handleSendMessage(authCtx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error for missing 'to'")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing body", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"to": "receiver",
|
||||
})
|
||||
|
||||
result, _ := tr.handleSendMessage(authCtx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error for missing body")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("unauthenticated", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"to": "receiver",
|
||||
"body": "should fail",
|
||||
})
|
||||
|
||||
result, _ := tr.handleSendMessage(ctx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error for unauthenticated request")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestToolHandler_ReadInbox(t *testing.T) {
|
||||
tr, msgSvc, agentSvc, _ := newTestRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
agentSvc.Register(ctx, "sender", "Sender", "ai", nil, 1)
|
||||
agentSvc.Register(ctx, "reader", "Reader", "ai", nil, 1)
|
||||
|
||||
msgSvc.SendMessage(ctx, "sender", "reader", "test message", messaging.SendOptions{})
|
||||
|
||||
authCtx := ContextWithAgentName(ctx, "reader")
|
||||
|
||||
t.Run("read messages", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{})
|
||||
|
||||
result, err := tr.handleReadInbox(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleReadInbox: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
var resp map[string]any
|
||||
text := result.Content[0].(mcplib.TextContent).Text
|
||||
json.Unmarshal([]byte(text), &resp)
|
||||
count := resp["count"].(float64)
|
||||
if count != 1 {
|
||||
t.Errorf("count = %v, want 1", count)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("unauthenticated", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{})
|
||||
result, _ := tr.handleReadInbox(ctx, req)
|
||||
if !result.IsError {
|
||||
t.Error("expected error for unauthenticated request")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestToolHandler_ClaimMessages(t *testing.T) {
|
||||
tr, msgSvc, agentSvc, _ := newTestRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
agentSvc.Register(ctx, "sender", "Sender", "ai", nil, 1)
|
||||
agentSvc.Register(ctx, "claimer", "Claimer", "ai", nil, 1)
|
||||
|
||||
msgSvc.SendMessage(ctx, "sender", "claimer", "task 1", messaging.SendOptions{})
|
||||
msgSvc.SendMessage(ctx, "sender", "claimer", "task 2", messaging.SendOptions{})
|
||||
|
||||
authCtx := ContextWithAgentName(ctx, "claimer")
|
||||
|
||||
req := makeRequest(map[string]any{"limit": float64(1)})
|
||||
result, err := tr.handleClaimMessages(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleClaimMessages: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
var resp map[string]any
|
||||
text := result.Content[0].(mcplib.TextContent).Text
|
||||
json.Unmarshal([]byte(text), &resp)
|
||||
count := resp["count"].(float64)
|
||||
if count != 1 {
|
||||
t.Errorf("count = %v, want 1", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolHandler_MarkDone(t *testing.T) {
|
||||
tr, msgSvc, agentSvc, _ := newTestRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
agentSvc.Register(ctx, "sender", "Sender", "ai", nil, 1)
|
||||
agentSvc.Register(ctx, "worker", "Worker", "ai", nil, 1)
|
||||
|
||||
msg, _ := msgSvc.SendMessage(ctx, "sender", "worker", "do this", messaging.SendOptions{})
|
||||
msgSvc.ClaimMessages(ctx, "worker", 1)
|
||||
|
||||
authCtx := ContextWithAgentName(ctx, "worker")
|
||||
|
||||
req := makeRequest(map[string]any{
|
||||
"message_id": float64(msg.ID),
|
||||
})
|
||||
|
||||
result, err := tr.handleMarkDone(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleMarkDone: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolHandler_SearchMessages(t *testing.T) {
|
||||
tr, msgSvc, agentSvc, _ := newTestRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
agentSvc.Register(ctx, "sender", "Sender", "ai", nil, 1)
|
||||
agentSvc.Register(ctx, "searcher", "Searcher", "ai", nil, 1)
|
||||
|
||||
msgSvc.SendMessage(ctx, "sender", "searcher", "deployment failed", messaging.SendOptions{})
|
||||
msgSvc.SendMessage(ctx, "sender", "searcher", "all clear", messaging.SendOptions{})
|
||||
|
||||
authCtx := ContextWithAgentName(ctx, "searcher")
|
||||
|
||||
t.Run("keyword search", func(t *testing.T) {
|
||||
req := makeRequest(map[string]any{
|
||||
"query": "deployment",
|
||||
})
|
||||
|
||||
result, err := tr.handleSearchMessages(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleSearchMessages: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
var resp map[string]any
|
||||
text := result.Content[0].(mcplib.TextContent).Text
|
||||
json.Unmarshal([]byte(text), &resp)
|
||||
count := resp["count"].(float64)
|
||||
if count != 1 {
|
||||
t.Errorf("count = %v, want 1", count)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestToolHandler_DiscoverAgents(t *testing.T) {
|
||||
tr, _, agentSvc, _ := newTestRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
agentSvc.Register(ctx, "bot-a", "Bot A", "ai", json.RawMessage(`{"skills":["search"]}`), 1)
|
||||
agentSvc.Register(ctx, "bot-b", "Bot B", "ai", json.RawMessage(`{"skills":["analyze"]}`), 1)
|
||||
|
||||
authCtx := ContextWithAgentName(ctx, "bot-a")
|
||||
|
||||
req := makeRequest(map[string]any{
|
||||
"query": "search",
|
||||
})
|
||||
|
||||
result, err := tr.handleDiscoverAgents(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleDiscoverAgents: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
var resp map[string]any
|
||||
text := result.Content[0].(mcplib.TextContent).Text
|
||||
json.Unmarshal([]byte(text), &resp)
|
||||
count := resp["count"].(float64)
|
||||
if count != 1 {
|
||||
t.Errorf("count = %v, want 1", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolHandler_UpdateAgent(t *testing.T) {
|
||||
tr, _, agentSvc, _ := newTestRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
agentSvc.Register(ctx, "update-me", "Update Me", "ai", nil, 1)
|
||||
|
||||
authCtx := ContextWithAgentName(ctx, "update-me")
|
||||
|
||||
req := makeRequest(map[string]any{
|
||||
"display_name": "Updated Name",
|
||||
"capabilities": `{"skills":["new-skill"]}`,
|
||||
})
|
||||
|
||||
result, err := tr.handleUpdateAgent(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleUpdateAgent: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
|
||||
var resp map[string]any
|
||||
text := result.Content[0].(mcplib.TextContent).Text
|
||||
json.Unmarshal([]byte(text), &resp)
|
||||
if resp["display_name"] != "Updated Name" {
|
||||
t.Errorf("display_name = %v, want Updated Name", resp["display_name"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolHandler_DeregisterAgent(t *testing.T) {
|
||||
tr, _, agentSvc, _ := newTestRegistrar(t)
|
||||
ctx := context.Background()
|
||||
|
||||
agentSvc.Register(ctx, "bye-bot", "Bye Bot", "ai", nil, 1)
|
||||
|
||||
authCtx := ContextWithAgentName(ctx, "bye-bot")
|
||||
|
||||
req := makeRequest(map[string]any{})
|
||||
|
||||
result, err := tr.handleDeregisterAgent(authCtx, req)
|
||||
if err != nil {
|
||||
t.Fatalf("handleDeregisterAgent: %v", err)
|
||||
}
|
||||
if result.IsError {
|
||||
t.Fatalf("unexpected error: %v", result.Content)
|
||||
}
|
||||
}
|
||||
|
||||
var _ = storage.RunMigrations
|
||||
@@ -0,0 +1,30 @@
|
||||
package messaging
|
||||
|
||||
// SendOptions configures message sending behavior.
|
||||
type SendOptions struct {
|
||||
Subject string `json:"subject,omitempty"`
|
||||
Priority int `json:"priority,omitempty"`
|
||||
Metadata string `json:"metadata,omitempty"`
|
||||
ChannelID *int64 `json:"channel_id,omitempty"`
|
||||
ConversationID *int64 `json:"conversation_id,omitempty"`
|
||||
}
|
||||
|
||||
// ReadOptions configures inbox reading behavior.
|
||||
type ReadOptions struct {
|
||||
Status string `json:"status,omitempty"`
|
||||
FromAgent string `json:"from_agent,omitempty"`
|
||||
ConversationID *int64 `json:"conversation_id,omitempty"`
|
||||
MinPriority int `json:"min_priority,omitempty"`
|
||||
Limit int `json:"limit,omitempty"`
|
||||
IncludeRead bool `json:"include_read,omitempty"`
|
||||
}
|
||||
|
||||
// SearchOptions configures message search behavior.
|
||||
type SearchOptions struct {
|
||||
FromAgent string `json:"from_agent,omitempty"`
|
||||
ToAgent string `json:"to_agent,omitempty"`
|
||||
ChannelID *int64 `json:"channel_id,omitempty"`
|
||||
MinPriority int `json:"min_priority,omitempty"`
|
||||
Status string `json:"status,omitempty"`
|
||||
Limit int `json:"limit,omitempty"`
|
||||
}
|
||||
@@ -0,0 +1,330 @@
|
||||
package messaging
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strings"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/trace"
|
||||
)
|
||||
|
||||
// MessagingService provides business logic for messaging operations.
|
||||
type MessagingService struct {
|
||||
store MessageStore
|
||||
tracer *trace.Tracer
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewMessagingService creates a new messaging service.
|
||||
func NewMessagingService(store MessageStore, tracer *trace.Tracer) *MessagingService {
|
||||
return &MessagingService{
|
||||
store: store,
|
||||
tracer: tracer,
|
||||
logger: slog.Default().With("component", "messaging"),
|
||||
}
|
||||
}
|
||||
|
||||
// SendMessage creates a message, auto-creating conversations as needed.
|
||||
func (s *MessagingService) SendMessage(ctx context.Context, from, to, body string, opts SendOptions) (*Message, error) {
|
||||
// Validate inputs
|
||||
if strings.TrimSpace(body) == "" {
|
||||
return nil, fmt.Errorf("message body cannot be empty")
|
||||
}
|
||||
|
||||
if to == "" && opts.ChannelID == nil {
|
||||
return nil, fmt.Errorf("either to_agent or channel_id must be specified")
|
||||
}
|
||||
|
||||
if to != "" && opts.ChannelID != nil {
|
||||
return nil, fmt.Errorf("cannot specify both to_agent and channel_id")
|
||||
}
|
||||
|
||||
priority := opts.Priority
|
||||
if priority == 0 {
|
||||
priority = 5
|
||||
}
|
||||
if priority < 1 || priority > 10 {
|
||||
return nil, fmt.Errorf("priority must be between 1 and 10")
|
||||
}
|
||||
|
||||
var metadata json.RawMessage
|
||||
if opts.Metadata != "" {
|
||||
if !json.Valid([]byte(opts.Metadata)) {
|
||||
return nil, fmt.Errorf("metadata must be valid JSON")
|
||||
}
|
||||
metadata = json.RawMessage(opts.Metadata)
|
||||
} else {
|
||||
metadata = json.RawMessage("{}")
|
||||
}
|
||||
|
||||
// For DMs, verify recipient agent exists
|
||||
if to != "" {
|
||||
exists, err := s.store.AgentExists(ctx, to)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("check agent existence: %w", err)
|
||||
}
|
||||
if !exists {
|
||||
return nil, fmt.Errorf("agent not found: %s", to)
|
||||
}
|
||||
}
|
||||
|
||||
// Find or create conversation
|
||||
var conv *Conversation
|
||||
var err error
|
||||
|
||||
if opts.ConversationID != nil {
|
||||
conv, err = s.store.GetConversation(ctx, *opts.ConversationID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get conversation: %w", err)
|
||||
}
|
||||
} else if opts.Subject != "" && to != "" {
|
||||
conv, err = s.store.FindConversation(ctx, opts.Subject, from, to)
|
||||
if err != nil && err != sql.ErrNoRows {
|
||||
return nil, fmt.Errorf("find conversation: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
if conv == nil {
|
||||
conv = &Conversation{
|
||||
Subject: opts.Subject,
|
||||
CreatedBy: from,
|
||||
ChannelID: opts.ChannelID,
|
||||
}
|
||||
if err := s.store.InsertConversation(ctx, conv); err != nil {
|
||||
return nil, fmt.Errorf("create conversation: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
msg := &Message{
|
||||
ConversationID: conv.ID,
|
||||
FromAgent: from,
|
||||
ToAgent: to,
|
||||
ChannelID: opts.ChannelID,
|
||||
Body: body,
|
||||
Priority: priority,
|
||||
Status: StatusPending,
|
||||
Metadata: metadata,
|
||||
}
|
||||
|
||||
if err := s.store.InsertMessage(ctx, msg); err != nil {
|
||||
return nil, fmt.Errorf("insert message: %w", err)
|
||||
}
|
||||
|
||||
s.logger.Info("message sent",
|
||||
"from", from,
|
||||
"to", to,
|
||||
"message_id", msg.ID,
|
||||
"conversation_id", conv.ID,
|
||||
"priority", priority,
|
||||
)
|
||||
|
||||
// Record trace
|
||||
if s.tracer != nil {
|
||||
s.tracer.Record(ctx, from, "send_message", map[string]any{
|
||||
"message_id": msg.ID,
|
||||
"conversation_id": conv.ID,
|
||||
"to": to,
|
||||
"priority": priority,
|
||||
})
|
||||
}
|
||||
|
||||
return msg, nil
|
||||
}
|
||||
|
||||
// ReadInbox returns messages for an agent and advances the read position.
|
||||
func (s *MessagingService) ReadInbox(ctx context.Context, agentName string, opts ReadOptions) ([]*Message, error) {
|
||||
messages, err := s.store.GetInboxMessages(ctx, agentName, opts)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get inbox messages: %w", err)
|
||||
}
|
||||
|
||||
// Advance inbox state for each conversation
|
||||
conversationMaxID := make(map[int64]int64)
|
||||
for _, msg := range messages {
|
||||
if msg.ID > conversationMaxID[msg.ConversationID] {
|
||||
conversationMaxID[msg.ConversationID] = msg.ID
|
||||
}
|
||||
}
|
||||
|
||||
for convID, maxMsgID := range conversationMaxID {
|
||||
if err := s.store.UpdateInboxState(ctx, agentName, convID, maxMsgID); err != nil {
|
||||
s.logger.Error("failed to update inbox state",
|
||||
"agent", agentName,
|
||||
"conversation_id", convID,
|
||||
"error", err,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
s.logger.Info("inbox read",
|
||||
"agent", agentName,
|
||||
"message_count", len(messages),
|
||||
)
|
||||
|
||||
if s.tracer != nil {
|
||||
s.tracer.Record(ctx, agentName, "read_inbox", map[string]any{
|
||||
"message_count": len(messages),
|
||||
"include_read": opts.IncludeRead,
|
||||
})
|
||||
}
|
||||
|
||||
return messages, nil
|
||||
}
|
||||
|
||||
// ClaimMessages atomically claims pending messages for processing.
|
||||
func (s *MessagingService) ClaimMessages(ctx context.Context, agentName string, limit int) ([]*Message, error) {
|
||||
if limit <= 0 {
|
||||
limit = 10
|
||||
}
|
||||
|
||||
messages, err := s.store.ClaimMessages(ctx, agentName, limit)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("claim messages: %w", err)
|
||||
}
|
||||
|
||||
s.logger.Info("messages claimed",
|
||||
"agent", agentName,
|
||||
"count", len(messages),
|
||||
)
|
||||
|
||||
if s.tracer != nil {
|
||||
ids := make([]int64, len(messages))
|
||||
for i, m := range messages {
|
||||
ids[i] = m.ID
|
||||
}
|
||||
s.tracer.Record(ctx, agentName, "claim_messages", map[string]any{
|
||||
"claimed_count": len(messages),
|
||||
"message_ids": ids,
|
||||
})
|
||||
}
|
||||
|
||||
return messages, nil
|
||||
}
|
||||
|
||||
// MarkDone marks a message as done. Only the claiming agent can do this.
|
||||
func (s *MessagingService) MarkDone(ctx context.Context, messageID int64, agentName string) error {
|
||||
msg, err := s.store.GetMessageByID(ctx, messageID)
|
||||
if err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return fmt.Errorf("message not found: %d", messageID)
|
||||
}
|
||||
return fmt.Errorf("get message: %w", err)
|
||||
}
|
||||
|
||||
if msg.Status != StatusProcessing {
|
||||
return fmt.Errorf("message %d is not in processing status (current: %s)", messageID, msg.Status)
|
||||
}
|
||||
|
||||
if msg.ClaimedBy != agentName {
|
||||
return fmt.Errorf("message %d is not claimed by agent %s", messageID, agentName)
|
||||
}
|
||||
|
||||
if err := s.store.UpdateMessageStatus(ctx, messageID, StatusDone, agentName, nil); err != nil {
|
||||
return fmt.Errorf("mark done: %w", err)
|
||||
}
|
||||
|
||||
s.logger.Info("message marked done",
|
||||
"agent", agentName,
|
||||
"message_id", messageID,
|
||||
)
|
||||
|
||||
if s.tracer != nil {
|
||||
s.tracer.Record(ctx, agentName, "mark_done", map[string]any{
|
||||
"message_id": messageID,
|
||||
"status": StatusDone,
|
||||
})
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// MarkFailed marks a message as failed with a reason.
|
||||
func (s *MessagingService) MarkFailed(ctx context.Context, messageID int64, agentName, reason string) error {
|
||||
msg, err := s.store.GetMessageByID(ctx, messageID)
|
||||
if err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return fmt.Errorf("message not found: %d", messageID)
|
||||
}
|
||||
return fmt.Errorf("get message: %w", err)
|
||||
}
|
||||
|
||||
if msg.Status != StatusProcessing {
|
||||
return fmt.Errorf("message %d is not in processing status (current: %s)", messageID, msg.Status)
|
||||
}
|
||||
|
||||
if msg.ClaimedBy != agentName {
|
||||
return fmt.Errorf("message %d is not claimed by agent %s", messageID, agentName)
|
||||
}
|
||||
|
||||
// Merge failure reason into metadata
|
||||
var existingMeta map[string]any
|
||||
if err := json.Unmarshal(msg.Metadata, &existingMeta); err != nil {
|
||||
existingMeta = make(map[string]any)
|
||||
}
|
||||
existingMeta["error"] = reason
|
||||
mergedMeta, err := json.Marshal(existingMeta)
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshal metadata: %w", err)
|
||||
}
|
||||
|
||||
if err := s.store.UpdateMessageStatus(ctx, messageID, StatusFailed, agentName, mergedMeta); err != nil {
|
||||
return fmt.Errorf("mark failed: %w", err)
|
||||
}
|
||||
|
||||
s.logger.Info("message marked failed",
|
||||
"agent", agentName,
|
||||
"message_id", messageID,
|
||||
"reason", reason,
|
||||
)
|
||||
|
||||
if s.tracer != nil {
|
||||
s.tracer.Record(ctx, agentName, "mark_failed", map[string]any{
|
||||
"message_id": messageID,
|
||||
"status": StatusFailed,
|
||||
"reason": reason,
|
||||
})
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// SearchMessages performs full-text search on messages.
|
||||
func (s *MessagingService) SearchMessages(ctx context.Context, agentName, query string, opts SearchOptions) ([]*Message, error) {
|
||||
messages, err := s.store.SearchMessages(ctx, agentName, query, opts)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("search messages: %w", err)
|
||||
}
|
||||
|
||||
s.logger.Info("messages searched",
|
||||
"agent", agentName,
|
||||
"query", query,
|
||||
"result_count", len(messages),
|
||||
)
|
||||
|
||||
if s.tracer != nil {
|
||||
s.tracer.Record(ctx, agentName, "search_messages", map[string]any{
|
||||
"query": query,
|
||||
"result_count": len(messages),
|
||||
})
|
||||
}
|
||||
|
||||
return messages, nil
|
||||
}
|
||||
|
||||
// GetConversation returns a conversation and its messages.
|
||||
func (s *MessagingService) GetConversation(ctx context.Context, id int64) (*Conversation, []*Message, error) {
|
||||
conv, err := s.store.GetConversation(ctx, id)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("get conversation: %w", err)
|
||||
}
|
||||
|
||||
messages, err := s.store.GetConversationMessages(ctx, id)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("get conversation messages: %w", err)
|
||||
}
|
||||
|
||||
return conv, messages, nil
|
||||
}
|
||||
@@ -0,0 +1,526 @@
|
||||
package messaging
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/storage"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/trace"
|
||||
)
|
||||
|
||||
func newTestService(t *testing.T) (*MessagingService, *sql.DB) {
|
||||
t.Helper()
|
||||
db := newTestDB(t)
|
||||
|
||||
// Seed test agents
|
||||
seedAgent(t, db, "sender")
|
||||
seedAgent(t, db, "receiver")
|
||||
seedAgent(t, db, "worker")
|
||||
|
||||
store := NewSQLiteMessageStore(db)
|
||||
tracer := trace.NewTracer(db)
|
||||
t.Cleanup(func() { tracer.Close() })
|
||||
|
||||
svc := NewMessagingService(store, tracer)
|
||||
return svc, db
|
||||
}
|
||||
|
||||
func TestMessagingService_SendMessage(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
from string
|
||||
to string
|
||||
body string
|
||||
opts SendOptions
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "successful send",
|
||||
from: "sender",
|
||||
to: "receiver",
|
||||
body: "Hello, world!",
|
||||
opts: SendOptions{Subject: "Test", Priority: 5},
|
||||
},
|
||||
{
|
||||
name: "empty body fails",
|
||||
from: "sender",
|
||||
to: "receiver",
|
||||
body: "",
|
||||
opts: SendOptions{},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "whitespace body fails",
|
||||
from: "sender",
|
||||
to: "receiver",
|
||||
body: " ",
|
||||
opts: SendOptions{},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "invalid priority",
|
||||
from: "sender",
|
||||
to: "receiver",
|
||||
body: "test",
|
||||
opts: SendOptions{Priority: 15},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "priority too low",
|
||||
from: "sender",
|
||||
to: "receiver",
|
||||
body: "test",
|
||||
opts: SendOptions{Priority: -1},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "invalid metadata JSON",
|
||||
from: "sender",
|
||||
to: "receiver",
|
||||
body: "test",
|
||||
opts: SendOptions{Metadata: "not json{"},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "non-existent recipient",
|
||||
from: "sender",
|
||||
to: "ghost-agent",
|
||||
body: "test",
|
||||
opts: SendOptions{},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "no recipient specified",
|
||||
from: "sender",
|
||||
to: "",
|
||||
body: "test",
|
||||
opts: SendOptions{},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "valid metadata",
|
||||
from: "sender",
|
||||
to: "receiver",
|
||||
body: "with metadata",
|
||||
opts: SendOptions{Metadata: `{"key":"value"}`},
|
||||
},
|
||||
{
|
||||
name: "default priority",
|
||||
from: "sender",
|
||||
to: "receiver",
|
||||
body: "default priority",
|
||||
opts: SendOptions{},
|
||||
},
|
||||
{
|
||||
name: "priority 1 is valid",
|
||||
from: "sender",
|
||||
to: "receiver",
|
||||
body: "low priority",
|
||||
opts: SendOptions{Priority: 1},
|
||||
},
|
||||
{
|
||||
name: "priority 10 is valid",
|
||||
from: "sender",
|
||||
to: "receiver",
|
||||
body: "high priority",
|
||||
opts: SendOptions{Priority: 10},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
msg, err := svc.SendMessage(ctx, tt.from, tt.to, tt.body, tt.opts)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Fatalf("SendMessage() error = %v, wantErr %v", err, tt.wantErr)
|
||||
}
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
if msg.ID == 0 {
|
||||
t.Error("message ID should not be 0")
|
||||
}
|
||||
if msg.Status != StatusPending {
|
||||
t.Errorf("status = %s, want %s", msg.Status, StatusPending)
|
||||
}
|
||||
if msg.FromAgent != tt.from {
|
||||
t.Errorf("from = %s, want %s", msg.FromAgent, tt.from)
|
||||
}
|
||||
if msg.ToAgent != tt.to {
|
||||
t.Errorf("to = %s, want %s", msg.ToAgent, tt.to)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestMessagingService_ConversationReuse(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
msg1, err := svc.SendMessage(ctx, "sender", "receiver", "first message", SendOptions{Subject: "Reuse Topic"})
|
||||
if err != nil {
|
||||
t.Fatalf("SendMessage 1: %v", err)
|
||||
}
|
||||
|
||||
msg2, err := svc.SendMessage(ctx, "sender", "receiver", "second message", SendOptions{Subject: "Reuse Topic"})
|
||||
if err != nil {
|
||||
t.Fatalf("SendMessage 2: %v", err)
|
||||
}
|
||||
|
||||
if msg1.ConversationID != msg2.ConversationID {
|
||||
t.Errorf("conversations should be reused: msg1=%d, msg2=%d", msg1.ConversationID, msg2.ConversationID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMessagingService_ReadInbox(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
_, err := svc.SendMessage(ctx, "sender", "receiver", "msg 1", SendOptions{Priority: 3})
|
||||
if err != nil {
|
||||
t.Fatalf("SendMessage: %v", err)
|
||||
}
|
||||
_, err = svc.SendMessage(ctx, "sender", "receiver", "msg 2", SendOptions{Priority: 8})
|
||||
if err != nil {
|
||||
t.Fatalf("SendMessage: %v", err)
|
||||
}
|
||||
|
||||
t.Run("returns messages ordered by priority desc", func(t *testing.T) {
|
||||
messages, err := svc.ReadInbox(ctx, "receiver", ReadOptions{IncludeRead: true})
|
||||
if err != nil {
|
||||
t.Fatalf("ReadInbox: %v", err)
|
||||
}
|
||||
if len(messages) != 2 {
|
||||
t.Fatalf("got %d messages, want 2", len(messages))
|
||||
}
|
||||
if messages[0].Priority < messages[1].Priority {
|
||||
t.Errorf("messages not ordered by priority desc: %d, %d", messages[0].Priority, messages[1].Priority)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestMessagingService_ReadInbox_ReadUnread(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
_, err := svc.SendMessage(ctx, "sender", "receiver", "msg 1", SendOptions{Priority: 3})
|
||||
if err != nil {
|
||||
t.Fatalf("SendMessage: %v", err)
|
||||
}
|
||||
_, err = svc.SendMessage(ctx, "sender", "receiver", "msg 2", SendOptions{Priority: 8})
|
||||
if err != nil {
|
||||
t.Fatalf("SendMessage: %v", err)
|
||||
}
|
||||
|
||||
// First read
|
||||
messages, err := svc.ReadInbox(ctx, "receiver", ReadOptions{})
|
||||
if err != nil {
|
||||
t.Fatalf("ReadInbox: %v", err)
|
||||
}
|
||||
if len(messages) == 0 {
|
||||
t.Fatal("expected messages on first read")
|
||||
}
|
||||
|
||||
// Second read without include_read
|
||||
messages, err = svc.ReadInbox(ctx, "receiver", ReadOptions{})
|
||||
if err != nil {
|
||||
t.Fatalf("ReadInbox: %v", err)
|
||||
}
|
||||
if len(messages) != 0 {
|
||||
t.Errorf("got %d messages on second read (no include_read), want 0", len(messages))
|
||||
}
|
||||
|
||||
// With include_read
|
||||
messages, err = svc.ReadInbox(ctx, "receiver", ReadOptions{IncludeRead: true})
|
||||
if err != nil {
|
||||
t.Fatalf("ReadInbox: %v", err)
|
||||
}
|
||||
if len(messages) != 2 {
|
||||
t.Errorf("got %d messages with include_read, want 2", len(messages))
|
||||
}
|
||||
}
|
||||
|
||||
func TestMessagingService_ReadInbox_Filters(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
svc.SendMessage(ctx, "sender", "receiver", "low pri", SendOptions{Priority: 3})
|
||||
svc.SendMessage(ctx, "sender", "receiver", "high pri", SendOptions{Priority: 8})
|
||||
|
||||
t.Run("filter by from_agent", func(t *testing.T) {
|
||||
messages, err := svc.ReadInbox(ctx, "receiver", ReadOptions{
|
||||
FromAgent: "sender",
|
||||
IncludeRead: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("ReadInbox: %v", err)
|
||||
}
|
||||
if len(messages) != 2 {
|
||||
t.Errorf("got %d messages from sender, want 2", len(messages))
|
||||
}
|
||||
|
||||
messages, err = svc.ReadInbox(ctx, "receiver", ReadOptions{
|
||||
FromAgent: "nobody",
|
||||
IncludeRead: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("ReadInbox: %v", err)
|
||||
}
|
||||
if len(messages) != 0 {
|
||||
t.Errorf("got %d messages from nobody, want 0", len(messages))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("filter by min_priority", func(t *testing.T) {
|
||||
messages, err := svc.ReadInbox(ctx, "receiver", ReadOptions{
|
||||
MinPriority: 7,
|
||||
IncludeRead: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("ReadInbox: %v", err)
|
||||
}
|
||||
if len(messages) != 1 {
|
||||
t.Errorf("got %d messages with min_priority=7, want 1", len(messages))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("limit", func(t *testing.T) {
|
||||
messages, err := svc.ReadInbox(ctx, "receiver", ReadOptions{
|
||||
Limit: 1,
|
||||
IncludeRead: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("ReadInbox: %v", err)
|
||||
}
|
||||
if len(messages) != 1 {
|
||||
t.Errorf("got %d messages with limit=1, want 1", len(messages))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestMessagingService_ClaimMessages(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
for i := 0; i < 5; i++ {
|
||||
_, err := svc.SendMessage(ctx, "sender", "worker", "task", SendOptions{Priority: 5})
|
||||
if err != nil {
|
||||
t.Fatalf("SendMessage: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Claim 3
|
||||
claimed, err := svc.ClaimMessages(ctx, "worker", 3)
|
||||
if err != nil {
|
||||
t.Fatalf("ClaimMessages: %v", err)
|
||||
}
|
||||
if len(claimed) != 3 {
|
||||
t.Errorf("claimed %d, want 3", len(claimed))
|
||||
}
|
||||
for _, m := range claimed {
|
||||
if m.Status != StatusProcessing {
|
||||
t.Errorf("status = %s, want %s", m.Status, StatusProcessing)
|
||||
}
|
||||
}
|
||||
|
||||
// Claim remaining
|
||||
claimed2, err := svc.ClaimMessages(ctx, "worker", 10)
|
||||
if err != nil {
|
||||
t.Fatalf("ClaimMessages: %v", err)
|
||||
}
|
||||
if len(claimed2) != 2 {
|
||||
t.Errorf("claimed %d, want 2", len(claimed2))
|
||||
}
|
||||
|
||||
// No more pending
|
||||
claimed3, err := svc.ClaimMessages(ctx, "worker", 5)
|
||||
if err != nil {
|
||||
t.Fatalf("ClaimMessages: %v", err)
|
||||
}
|
||||
if len(claimed3) != 0 {
|
||||
t.Errorf("claimed %d, want 0", len(claimed3))
|
||||
}
|
||||
}
|
||||
|
||||
func TestMessagingService_MarkDone(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
msg, err := svc.SendMessage(ctx, "sender", "worker", "task", SendOptions{Priority: 5})
|
||||
if err != nil {
|
||||
t.Fatalf("SendMessage: %v", err)
|
||||
}
|
||||
|
||||
claimed, err := svc.ClaimMessages(ctx, "worker", 1)
|
||||
if err != nil {
|
||||
t.Fatalf("ClaimMessages: %v", err)
|
||||
}
|
||||
if len(claimed) != 1 {
|
||||
t.Fatalf("claimed %d, want 1", len(claimed))
|
||||
}
|
||||
|
||||
// Mark done
|
||||
err = svc.MarkDone(ctx, msg.ID, "worker")
|
||||
if err != nil {
|
||||
t.Fatalf("MarkDone: %v", err)
|
||||
}
|
||||
|
||||
// Verify status
|
||||
got, err := svc.store.GetMessageByID(ctx, msg.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetMessageByID: %v", err)
|
||||
}
|
||||
if got.Status != StatusDone {
|
||||
t.Errorf("status = %s, want %s", got.Status, StatusDone)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMessagingService_MarkDone_AlreadyDone(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
msg, _ := svc.SendMessage(ctx, "sender", "worker", "task", SendOptions{})
|
||||
svc.ClaimMessages(ctx, "worker", 1)
|
||||
svc.MarkDone(ctx, msg.ID, "worker")
|
||||
|
||||
// Try again
|
||||
err := svc.MarkDone(ctx, msg.ID, "worker")
|
||||
if err == nil {
|
||||
t.Error("expected error for marking already-done message")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMessagingService_MarkDone_WrongAgent(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
msg, _ := svc.SendMessage(ctx, "sender", "worker", "task", SendOptions{})
|
||||
svc.ClaimMessages(ctx, "worker", 1)
|
||||
|
||||
err := svc.MarkDone(ctx, msg.ID, "sender")
|
||||
if err == nil {
|
||||
t.Error("expected error for wrong agent")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMessagingService_MarkDone_NonExistent(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
err := svc.MarkDone(ctx, 99999, "worker")
|
||||
if err == nil {
|
||||
t.Error("expected error for non-existent message")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMessagingService_MarkFailed(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
msg, err := svc.SendMessage(ctx, "sender", "worker", "failing task", SendOptions{})
|
||||
if err != nil {
|
||||
t.Fatalf("SendMessage: %v", err)
|
||||
}
|
||||
|
||||
svc.ClaimMessages(ctx, "worker", 1)
|
||||
|
||||
err = svc.MarkFailed(ctx, msg.ID, "worker", "timeout error")
|
||||
if err != nil {
|
||||
t.Fatalf("MarkFailed: %v", err)
|
||||
}
|
||||
|
||||
got, err := svc.store.GetMessageByID(ctx, msg.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetMessageByID: %v", err)
|
||||
}
|
||||
if got.Status != StatusFailed {
|
||||
t.Errorf("status = %s, want %s", got.Status, StatusFailed)
|
||||
}
|
||||
if string(got.Metadata) == "{}" {
|
||||
t.Error("expected metadata to contain error info")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMessagingService_SearchMessages(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
svc.SendMessage(ctx, "sender", "receiver", "deployment failure in prod", SendOptions{Priority: 8})
|
||||
svc.SendMessage(ctx, "sender", "receiver", "deployment success in staging", SendOptions{Priority: 3})
|
||||
svc.SendMessage(ctx, "sender", "receiver", "security alert detected", SendOptions{Priority: 9})
|
||||
|
||||
t.Run("keyword search", func(t *testing.T) {
|
||||
results, err := svc.SearchMessages(ctx, "receiver", "deployment", SearchOptions{})
|
||||
if err != nil {
|
||||
t.Fatalf("SearchMessages: %v", err)
|
||||
}
|
||||
if len(results) != 2 {
|
||||
t.Errorf("got %d results, want 2", len(results))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("empty query returns recent", func(t *testing.T) {
|
||||
results, err := svc.SearchMessages(ctx, "receiver", "", SearchOptions{})
|
||||
if err != nil {
|
||||
t.Fatalf("SearchMessages: %v", err)
|
||||
}
|
||||
if len(results) != 3 {
|
||||
t.Errorf("got %d results, want 3", len(results))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestMessagingService_GetConversation(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
msg1, _ := svc.SendMessage(ctx, "sender", "receiver", "first in thread", SendOptions{Subject: "Thread"})
|
||||
svc.SendMessage(ctx, "sender", "receiver", "second in thread", SendOptions{Subject: "Thread"})
|
||||
|
||||
conv, messages, err := svc.GetConversation(ctx, msg1.ConversationID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetConversation: %v", err)
|
||||
}
|
||||
|
||||
if conv.Subject != "Thread" {
|
||||
t.Errorf("subject = %s, want Thread", conv.Subject)
|
||||
}
|
||||
|
||||
if len(messages) != 2 {
|
||||
t.Errorf("got %d messages, want 2", len(messages))
|
||||
}
|
||||
}
|
||||
|
||||
func TestMessagingService_TracesRecorded(t *testing.T) {
|
||||
svc, db := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
svc.SendMessage(ctx, "sender", "receiver", "traced message", SendOptions{})
|
||||
svc.ReadInbox(ctx, "receiver", ReadOptions{IncludeRead: true})
|
||||
|
||||
// Wait for async trace writes with retries
|
||||
var count int
|
||||
for i := 0; i < 20; i++ {
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
err := db.QueryRow("SELECT COUNT(*) FROM traces").Scan(&count)
|
||||
if err != nil {
|
||||
t.Fatalf("count traces: %v", err)
|
||||
}
|
||||
if count >= 2 {
|
||||
break
|
||||
}
|
||||
}
|
||||
if count < 2 {
|
||||
t.Errorf("expected at least 2 traces, got %d", count)
|
||||
}
|
||||
}
|
||||
|
||||
// suppress unused import warning for storage package
|
||||
var _ = storage.RunMigrations
|
||||
@@ -0,0 +1,519 @@
|
||||
package messaging
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// MessageStore defines the storage interface for messaging operations.
|
||||
type MessageStore interface {
|
||||
InsertMessage(ctx context.Context, msg *Message) error
|
||||
InsertConversation(ctx context.Context, conv *Conversation) error
|
||||
FindConversation(ctx context.Context, subject, fromAgent, toAgent string) (*Conversation, error)
|
||||
GetInboxMessages(ctx context.Context, agentName string, opts ReadOptions) ([]*Message, error)
|
||||
GetInboxState(ctx context.Context, agentName string, conversationID int64) (*InboxState, error)
|
||||
UpdateInboxState(ctx context.Context, agentName string, conversationID int64, lastReadMsgID int64) error
|
||||
ClaimMessages(ctx context.Context, agentName string, limit int) ([]*Message, error)
|
||||
UpdateMessageStatus(ctx context.Context, id int64, status, claimedBy string, metadata json.RawMessage) error
|
||||
GetMessageByID(ctx context.Context, id int64) (*Message, error)
|
||||
SearchMessages(ctx context.Context, agentName, query string, opts SearchOptions) ([]*Message, error)
|
||||
GetConversation(ctx context.Context, id int64) (*Conversation, error)
|
||||
GetConversationMessages(ctx context.Context, conversationID int64) ([]*Message, error)
|
||||
AgentExists(ctx context.Context, agentName string) (bool, error)
|
||||
}
|
||||
|
||||
// SQLiteMessageStore implements MessageStore using SQLite.
|
||||
type SQLiteMessageStore struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
// NewSQLiteMessageStore creates a new SQLite-backed message store.
|
||||
func NewSQLiteMessageStore(db *sql.DB) *SQLiteMessageStore {
|
||||
return &SQLiteMessageStore{db: db}
|
||||
}
|
||||
|
||||
func (s *SQLiteMessageStore) InsertConversation(ctx context.Context, conv *Conversation) error {
|
||||
result, err := s.db.ExecContext(ctx,
|
||||
`INSERT INTO conversations (subject, created_by, channel_id, created_at, updated_at)
|
||||
VALUES (?, ?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`,
|
||||
conv.Subject, conv.CreatedBy, conv.ChannelID,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("insert conversation: %w", err)
|
||||
}
|
||||
id, err := result.LastInsertId()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get conversation id: %w", err)
|
||||
}
|
||||
conv.ID = id
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SQLiteMessageStore) FindConversation(ctx context.Context, subject, fromAgent, toAgent string) (*Conversation, error) {
|
||||
var conv Conversation
|
||||
var channelID sql.NullInt64
|
||||
err := s.db.QueryRowContext(ctx,
|
||||
`SELECT c.id, c.subject, c.created_by, c.channel_id, c.created_at, c.updated_at
|
||||
FROM conversations c
|
||||
WHERE c.subject = ? AND c.channel_id IS NULL
|
||||
AND EXISTS (
|
||||
SELECT 1 FROM messages m WHERE m.conversation_id = c.id
|
||||
AND ((m.from_agent = ? AND m.to_agent = ?) OR (m.from_agent = ? AND m.to_agent = ?))
|
||||
)
|
||||
ORDER BY c.id DESC LIMIT 1`,
|
||||
subject, fromAgent, toAgent, toAgent, fromAgent,
|
||||
).Scan(&conv.ID, &conv.Subject, &conv.CreatedBy, &channelID, &conv.CreatedAt, &conv.UpdatedAt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if channelID.Valid {
|
||||
conv.ChannelID = &channelID.Int64
|
||||
}
|
||||
return &conv, nil
|
||||
}
|
||||
|
||||
func (s *SQLiteMessageStore) InsertMessage(ctx context.Context, msg *Message) error {
|
||||
metadata := msg.Metadata
|
||||
if metadata == nil {
|
||||
metadata = json.RawMessage("{}")
|
||||
}
|
||||
|
||||
var toAgent sql.NullString
|
||||
if msg.ToAgent != "" {
|
||||
toAgent = sql.NullString{String: msg.ToAgent, Valid: true}
|
||||
}
|
||||
|
||||
result, err := s.db.ExecContext(ctx,
|
||||
`INSERT INTO messages (conversation_id, from_agent, to_agent, channel_id, body, priority, status, metadata, created_at, updated_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`,
|
||||
msg.ConversationID, msg.FromAgent, toAgent, msg.ChannelID, msg.Body, msg.Priority, msg.Status, string(metadata),
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("insert message: %w", err)
|
||||
}
|
||||
id, err := result.LastInsertId()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get message id: %w", err)
|
||||
}
|
||||
msg.ID = id
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SQLiteMessageStore) GetInboxMessages(ctx context.Context, agentName string, opts ReadOptions) ([]*Message, error) {
|
||||
var conditions []string
|
||||
var args []any
|
||||
|
||||
// Direct messages to this agent
|
||||
conditions = append(conditions, "m.to_agent = ?")
|
||||
args = append(args, agentName)
|
||||
|
||||
// Filter by read/unread using inbox_state
|
||||
if !opts.IncludeRead {
|
||||
conditions = append(conditions,
|
||||
`m.id > COALESCE(
|
||||
(SELECT last_read_message_id FROM inbox_state
|
||||
WHERE agent_name = ? AND conversation_id = m.conversation_id), 0)`)
|
||||
args = append(args, agentName)
|
||||
}
|
||||
|
||||
if opts.Status != "" {
|
||||
conditions = append(conditions, "m.status = ?")
|
||||
args = append(args, opts.Status)
|
||||
}
|
||||
|
||||
if opts.FromAgent != "" {
|
||||
conditions = append(conditions, "m.from_agent = ?")
|
||||
args = append(args, opts.FromAgent)
|
||||
}
|
||||
|
||||
if opts.ConversationID != nil {
|
||||
conditions = append(conditions, "m.conversation_id = ?")
|
||||
args = append(args, *opts.ConversationID)
|
||||
}
|
||||
|
||||
if opts.MinPriority > 0 {
|
||||
conditions = append(conditions, "m.priority >= ?")
|
||||
args = append(args, opts.MinPriority)
|
||||
}
|
||||
|
||||
limit := opts.Limit
|
||||
if limit <= 0 {
|
||||
limit = 50
|
||||
}
|
||||
|
||||
query := fmt.Sprintf(
|
||||
`SELECT m.id, m.conversation_id, m.from_agent, m.to_agent, m.channel_id,
|
||||
m.body, m.priority, m.status, m.metadata, m.claimed_by, m.claimed_at,
|
||||
m.created_at, m.updated_at
|
||||
FROM messages m
|
||||
WHERE %s
|
||||
ORDER BY m.priority DESC, m.created_at ASC
|
||||
LIMIT ?`,
|
||||
strings.Join(conditions, " AND "),
|
||||
)
|
||||
args = append(args, limit)
|
||||
|
||||
rows, err := s.db.QueryContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("query inbox: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
return scanMessages(rows)
|
||||
}
|
||||
|
||||
func (s *SQLiteMessageStore) GetInboxState(ctx context.Context, agentName string, conversationID int64) (*InboxState, error) {
|
||||
var state InboxState
|
||||
err := s.db.QueryRowContext(ctx,
|
||||
`SELECT agent_name, conversation_id, last_read_message_id, updated_at
|
||||
FROM inbox_state WHERE agent_name = ? AND conversation_id = ?`,
|
||||
agentName, conversationID,
|
||||
).Scan(&state.AgentName, &state.ConversationID, &state.LastReadMessageID, &state.UpdatedAt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &state, nil
|
||||
}
|
||||
|
||||
func (s *SQLiteMessageStore) UpdateInboxState(ctx context.Context, agentName string, conversationID int64, lastReadMsgID int64) error {
|
||||
_, err := s.db.ExecContext(ctx,
|
||||
`INSERT INTO inbox_state (agent_name, conversation_id, last_read_message_id, updated_at)
|
||||
VALUES (?, ?, ?, CURRENT_TIMESTAMP)
|
||||
ON CONFLICT(agent_name, conversation_id) DO UPDATE SET
|
||||
last_read_message_id = MAX(last_read_message_id, excluded.last_read_message_id),
|
||||
updated_at = CURRENT_TIMESTAMP`,
|
||||
agentName, conversationID, lastReadMsgID,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *SQLiteMessageStore) ClaimMessages(ctx context.Context, agentName string, limit int) ([]*Message, error) {
|
||||
if limit <= 0 {
|
||||
limit = 10
|
||||
}
|
||||
|
||||
// First, find the IDs of pending messages to claim
|
||||
idRows, err := s.db.QueryContext(ctx,
|
||||
`SELECT id FROM messages
|
||||
WHERE to_agent = ? AND status = 'pending'
|
||||
ORDER BY priority DESC, created_at ASC
|
||||
LIMIT ?`,
|
||||
agentName, limit,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("find pending messages: %w", err)
|
||||
}
|
||||
|
||||
var ids []int64
|
||||
for idRows.Next() {
|
||||
var id int64
|
||||
if err := idRows.Scan(&id); err != nil {
|
||||
idRows.Close()
|
||||
return nil, fmt.Errorf("scan message id: %w", err)
|
||||
}
|
||||
ids = append(ids, id)
|
||||
}
|
||||
idRows.Close()
|
||||
|
||||
if len(ids) == 0 {
|
||||
return []*Message{}, nil
|
||||
}
|
||||
|
||||
// Build placeholders for IN clause
|
||||
placeholders := make([]string, len(ids))
|
||||
args := make([]any, 0, len(ids)+1)
|
||||
args = append(args, agentName)
|
||||
for i, id := range ids {
|
||||
placeholders[i] = "?"
|
||||
args = append(args, id)
|
||||
}
|
||||
|
||||
// Atomically claim these specific messages
|
||||
_, err = s.db.ExecContext(ctx,
|
||||
fmt.Sprintf(
|
||||
`UPDATE messages SET
|
||||
status = 'processing',
|
||||
claimed_by = ?,
|
||||
claimed_at = CURRENT_TIMESTAMP,
|
||||
updated_at = CURRENT_TIMESTAMP
|
||||
WHERE id IN (%s) AND status = 'pending'`,
|
||||
strings.Join(placeholders, ","),
|
||||
),
|
||||
args...,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("claim messages: %w", err)
|
||||
}
|
||||
|
||||
// Return the claimed messages by their specific IDs
|
||||
fetchArgs := make([]any, len(ids))
|
||||
for i, id := range ids {
|
||||
fetchArgs[i] = id
|
||||
}
|
||||
|
||||
fetchPlaceholders := make([]string, len(ids))
|
||||
for i := range ids {
|
||||
fetchPlaceholders[i] = "?"
|
||||
}
|
||||
|
||||
rows, err := s.db.QueryContext(ctx,
|
||||
fmt.Sprintf(
|
||||
`SELECT id, conversation_id, from_agent, to_agent, channel_id,
|
||||
body, priority, status, metadata, claimed_by, claimed_at,
|
||||
created_at, updated_at
|
||||
FROM messages
|
||||
WHERE id IN (%s)
|
||||
ORDER BY priority DESC, created_at ASC`,
|
||||
strings.Join(fetchPlaceholders, ","),
|
||||
),
|
||||
fetchArgs...,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("query claimed messages: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
return scanMessages(rows)
|
||||
}
|
||||
|
||||
func (s *SQLiteMessageStore) UpdateMessageStatus(ctx context.Context, id int64, status, claimedBy string, metadata json.RawMessage) error {
|
||||
var result sql.Result
|
||||
var err error
|
||||
|
||||
if metadata != nil {
|
||||
result, err = s.db.ExecContext(ctx,
|
||||
`UPDATE messages SET status = ?, metadata = ?, updated_at = CURRENT_TIMESTAMP
|
||||
WHERE id = ? AND claimed_by = ?`,
|
||||
status, string(metadata), id, claimedBy,
|
||||
)
|
||||
} else {
|
||||
result, err = s.db.ExecContext(ctx,
|
||||
`UPDATE messages SET status = ?, updated_at = CURRENT_TIMESTAMP
|
||||
WHERE id = ? AND claimed_by = ?`,
|
||||
status, id, claimedBy,
|
||||
)
|
||||
}
|
||||
if err != nil {
|
||||
return fmt.Errorf("update message status: %w", err)
|
||||
}
|
||||
|
||||
rowsAffected, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get rows affected: %w", err)
|
||||
}
|
||||
if rowsAffected == 0 {
|
||||
return fmt.Errorf("message not found or not claimed by agent")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SQLiteMessageStore) GetMessageByID(ctx context.Context, id int64) (*Message, error) {
|
||||
row := s.db.QueryRowContext(ctx,
|
||||
`SELECT id, conversation_id, from_agent, to_agent, channel_id,
|
||||
body, priority, status, metadata, claimed_by, claimed_at,
|
||||
created_at, updated_at
|
||||
FROM messages WHERE id = ?`, id,
|
||||
)
|
||||
return scanMessage(row)
|
||||
}
|
||||
|
||||
func (s *SQLiteMessageStore) SearchMessages(ctx context.Context, agentName, query string, opts SearchOptions) ([]*Message, error) {
|
||||
var conditions []string
|
||||
var args []any
|
||||
|
||||
// Scope to messages accessible by this agent
|
||||
conditions = append(conditions, "(m.to_agent = ? OR m.from_agent = ?)")
|
||||
args = append(args, agentName, agentName)
|
||||
|
||||
var joinClause string
|
||||
var orderClause string
|
||||
|
||||
if query != "" {
|
||||
joinClause = "JOIN messages_fts ON messages_fts.rowid = m.id"
|
||||
conditions = append(conditions, "messages_fts MATCH ?")
|
||||
args = append(args, query)
|
||||
orderClause = "ORDER BY rank"
|
||||
} else {
|
||||
orderClause = "ORDER BY m.created_at DESC"
|
||||
}
|
||||
|
||||
if opts.FromAgent != "" {
|
||||
conditions = append(conditions, "m.from_agent = ?")
|
||||
args = append(args, opts.FromAgent)
|
||||
}
|
||||
|
||||
if opts.ToAgent != "" {
|
||||
conditions = append(conditions, "m.to_agent = ?")
|
||||
args = append(args, opts.ToAgent)
|
||||
}
|
||||
|
||||
if opts.MinPriority > 0 {
|
||||
conditions = append(conditions, "m.priority >= ?")
|
||||
args = append(args, opts.MinPriority)
|
||||
}
|
||||
|
||||
if opts.Status != "" {
|
||||
conditions = append(conditions, "m.status = ?")
|
||||
args = append(args, opts.Status)
|
||||
}
|
||||
|
||||
limit := opts.Limit
|
||||
if limit <= 0 {
|
||||
limit = 20
|
||||
}
|
||||
|
||||
querySQL := fmt.Sprintf(
|
||||
`SELECT m.id, m.conversation_id, m.from_agent, m.to_agent, m.channel_id,
|
||||
m.body, m.priority, m.status, m.metadata, m.claimed_by, m.claimed_at,
|
||||
m.created_at, m.updated_at
|
||||
FROM messages m
|
||||
%s
|
||||
WHERE %s
|
||||
%s
|
||||
LIMIT ?`,
|
||||
joinClause,
|
||||
strings.Join(conditions, " AND "),
|
||||
orderClause,
|
||||
)
|
||||
args = append(args, limit)
|
||||
|
||||
rows, err := s.db.QueryContext(ctx, querySQL, args...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("search messages: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
return scanMessages(rows)
|
||||
}
|
||||
|
||||
func (s *SQLiteMessageStore) GetConversation(ctx context.Context, id int64) (*Conversation, error) {
|
||||
var conv Conversation
|
||||
var channelID sql.NullInt64
|
||||
err := s.db.QueryRowContext(ctx,
|
||||
`SELECT id, subject, created_by, channel_id, created_at, updated_at
|
||||
FROM conversations WHERE id = ?`, id,
|
||||
).Scan(&conv.ID, &conv.Subject, &conv.CreatedBy, &channelID, &conv.CreatedAt, &conv.UpdatedAt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if channelID.Valid {
|
||||
conv.ChannelID = &channelID.Int64
|
||||
}
|
||||
return &conv, nil
|
||||
}
|
||||
|
||||
func (s *SQLiteMessageStore) GetConversationMessages(ctx context.Context, conversationID int64) ([]*Message, error) {
|
||||
rows, err := s.db.QueryContext(ctx,
|
||||
`SELECT id, conversation_id, from_agent, to_agent, channel_id,
|
||||
body, priority, status, metadata, claimed_by, claimed_at,
|
||||
created_at, updated_at
|
||||
FROM messages WHERE conversation_id = ?
|
||||
ORDER BY created_at ASC`, conversationID,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get conversation messages: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
return scanMessages(rows)
|
||||
}
|
||||
|
||||
func (s *SQLiteMessageStore) AgentExists(ctx context.Context, agentName string) (bool, error) {
|
||||
var count int
|
||||
err := s.db.QueryRowContext(ctx,
|
||||
`SELECT COUNT(*) FROM agents WHERE name = ? AND status = 'active'`,
|
||||
agentName,
|
||||
).Scan(&count)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
// scanMessages scans multiple message rows.
|
||||
func scanMessages(rows *sql.Rows) ([]*Message, error) {
|
||||
var messages []*Message
|
||||
for rows.Next() {
|
||||
msg, err := scanMessageFromRows(rows)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
messages = append(messages, msg)
|
||||
}
|
||||
if messages == nil {
|
||||
messages = []*Message{}
|
||||
}
|
||||
return messages, rows.Err()
|
||||
}
|
||||
|
||||
// scanMessageFromRows scans a single message from sql.Rows.
|
||||
func scanMessageFromRows(rows *sql.Rows) (*Message, error) {
|
||||
var msg Message
|
||||
var toAgent, claimedBy sql.NullString
|
||||
var channelID sql.NullInt64
|
||||
var claimedAt sql.NullTime
|
||||
var metadata string
|
||||
|
||||
err := rows.Scan(
|
||||
&msg.ID, &msg.ConversationID, &msg.FromAgent, &toAgent, &channelID,
|
||||
&msg.Body, &msg.Priority, &msg.Status, &metadata, &claimedBy, &claimedAt,
|
||||
&msg.CreatedAt, &msg.UpdatedAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("scan message: %w", err)
|
||||
}
|
||||
|
||||
if toAgent.Valid {
|
||||
msg.ToAgent = toAgent.String
|
||||
}
|
||||
if channelID.Valid {
|
||||
msg.ChannelID = &channelID.Int64
|
||||
}
|
||||
if claimedBy.Valid {
|
||||
msg.ClaimedBy = claimedBy.String
|
||||
}
|
||||
if claimedAt.Valid {
|
||||
msg.ClaimedAt = &claimedAt.Time
|
||||
}
|
||||
msg.Metadata = json.RawMessage(metadata)
|
||||
|
||||
return &msg, nil
|
||||
}
|
||||
|
||||
// scanMessage scans a single message from sql.Row.
|
||||
func scanMessage(row *sql.Row) (*Message, error) {
|
||||
var msg Message
|
||||
var toAgent, claimedBy sql.NullString
|
||||
var channelID sql.NullInt64
|
||||
var claimedAt sql.NullTime
|
||||
var metadata string
|
||||
|
||||
err := row.Scan(
|
||||
&msg.ID, &msg.ConversationID, &msg.FromAgent, &toAgent, &channelID,
|
||||
&msg.Body, &msg.Priority, &msg.Status, &metadata, &claimedBy, &claimedAt,
|
||||
&msg.CreatedAt, &msg.UpdatedAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if toAgent.Valid {
|
||||
msg.ToAgent = toAgent.String
|
||||
}
|
||||
if channelID.Valid {
|
||||
msg.ChannelID = &channelID.Int64
|
||||
}
|
||||
if claimedBy.Valid {
|
||||
msg.ClaimedBy = claimedBy.String
|
||||
}
|
||||
if claimedAt.Valid {
|
||||
t := claimedAt.Time
|
||||
msg.ClaimedAt = &t
|
||||
}
|
||||
msg.Metadata = json.RawMessage(metadata)
|
||||
|
||||
return &msg, nil
|
||||
}
|
||||
@@ -0,0 +1,536 @@
|
||||
package messaging
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/storage"
|
||||
)
|
||||
|
||||
func newTestDB(t *testing.T) *sql.DB {
|
||||
t.Helper()
|
||||
// Use file::memory: with shared cache so all connections see the same database
|
||||
// Each test gets a unique name to avoid cross-test interference
|
||||
dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())
|
||||
db, err := sql.Open("sqlite", dsn)
|
||||
if err != nil {
|
||||
t.Fatalf("open database: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { db.Close() })
|
||||
|
||||
// Enable foreign keys
|
||||
if _, err := db.Exec("PRAGMA foreign_keys=ON"); err != nil {
|
||||
t.Fatalf("enable foreign keys: %v", err)
|
||||
}
|
||||
|
||||
// Run migrations
|
||||
ctx := context.Background()
|
||||
if err := storage.RunMigrations(ctx, db); err != nil {
|
||||
t.Fatalf("run migrations: %v", err)
|
||||
}
|
||||
|
||||
return db
|
||||
}
|
||||
|
||||
// seedAgent inserts a test agent and its owner.
|
||||
func seedAgent(t *testing.T, db *sql.DB, name string) {
|
||||
t.Helper()
|
||||
// Ensure a user exists for the owner_id
|
||||
db.Exec(`INSERT OR IGNORE INTO users (id, username, password_hash, display_name) VALUES (1, 'testowner', 'hash', 'Test Owner')`)
|
||||
_, err := db.Exec(
|
||||
`INSERT OR IGNORE INTO agents (name, display_name, type, capabilities, owner_id, api_key_hash, status) VALUES (?, ?, 'ai', '{}', 1, 'testhash', 'active')`,
|
||||
name, name,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("seed agent %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteMessageStore_InsertConversation(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteMessageStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
conv := &Conversation{
|
||||
Subject: "Test Subject",
|
||||
CreatedBy: "agent-a",
|
||||
}
|
||||
|
||||
if err := store.InsertConversation(ctx, conv); err != nil {
|
||||
t.Fatalf("InsertConversation: %v", err)
|
||||
}
|
||||
|
||||
if conv.ID == 0 {
|
||||
t.Error("conversation ID should not be 0")
|
||||
}
|
||||
|
||||
// Verify conversation exists
|
||||
got, err := store.GetConversation(ctx, conv.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetConversation: %v", err)
|
||||
}
|
||||
if got.Subject != "Test Subject" {
|
||||
t.Errorf("Subject = %q, want %q", got.Subject, "Test Subject")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteMessageStore_InsertAndGetMessage(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteMessageStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
seedAgent(t, db, "sender")
|
||||
seedAgent(t, db, "receiver")
|
||||
|
||||
conv := &Conversation{Subject: "test", CreatedBy: "sender"}
|
||||
if err := store.InsertConversation(ctx, conv); err != nil {
|
||||
t.Fatalf("InsertConversation: %v", err)
|
||||
}
|
||||
|
||||
msg := &Message{
|
||||
ConversationID: conv.ID,
|
||||
FromAgent: "sender",
|
||||
ToAgent: "receiver",
|
||||
Body: "Hello!",
|
||||
Priority: 5,
|
||||
Status: StatusPending,
|
||||
}
|
||||
|
||||
if err := store.InsertMessage(ctx, msg); err != nil {
|
||||
t.Fatalf("InsertMessage: %v", err)
|
||||
}
|
||||
|
||||
if msg.ID == 0 {
|
||||
t.Error("message ID should not be 0")
|
||||
}
|
||||
|
||||
got, err := store.GetMessageByID(ctx, msg.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetMessageByID: %v", err)
|
||||
}
|
||||
|
||||
if got.Body != "Hello!" {
|
||||
t.Errorf("Body = %q, want %q", got.Body, "Hello!")
|
||||
}
|
||||
if got.FromAgent != "sender" {
|
||||
t.Errorf("FromAgent = %q, want %q", got.FromAgent, "sender")
|
||||
}
|
||||
if got.ToAgent != "receiver" {
|
||||
t.Errorf("ToAgent = %q, want %q", got.ToAgent, "receiver")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteMessageStore_FindConversation(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteMessageStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
seedAgent(t, db, "agent-a")
|
||||
seedAgent(t, db, "agent-b")
|
||||
|
||||
// Create conversation and add a message
|
||||
conv := &Conversation{Subject: "Topic X", CreatedBy: "agent-a"}
|
||||
if err := store.InsertConversation(ctx, conv); err != nil {
|
||||
t.Fatalf("InsertConversation: %v", err)
|
||||
}
|
||||
|
||||
msg := &Message{
|
||||
ConversationID: conv.ID,
|
||||
FromAgent: "agent-a",
|
||||
ToAgent: "agent-b",
|
||||
Body: "test",
|
||||
Priority: 5,
|
||||
Status: StatusPending,
|
||||
}
|
||||
if err := store.InsertMessage(ctx, msg); err != nil {
|
||||
t.Fatalf("InsertMessage: %v", err)
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
subject string
|
||||
from string
|
||||
to string
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "find existing conversation",
|
||||
subject: "Topic X",
|
||||
from: "agent-a",
|
||||
to: "agent-b",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "find reverse direction",
|
||||
subject: "Topic X",
|
||||
from: "agent-b",
|
||||
to: "agent-a",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "no match for different subject",
|
||||
subject: "Topic Y",
|
||||
from: "agent-a",
|
||||
to: "agent-b",
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
found, err := store.FindConversation(ctx, tt.subject, tt.from, tt.to)
|
||||
if tt.wantErr {
|
||||
if err == nil {
|
||||
t.Error("expected error, got nil")
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("FindConversation: %v", err)
|
||||
}
|
||||
if found.ID != conv.ID {
|
||||
t.Errorf("found conversation ID = %d, want %d", found.ID, conv.ID)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteMessageStore_GetInboxMessages(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteMessageStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
seedAgent(t, db, "sender")
|
||||
seedAgent(t, db, "reader")
|
||||
|
||||
conv := &Conversation{Subject: "inbox test", CreatedBy: "sender"}
|
||||
if err := store.InsertConversation(ctx, conv); err != nil {
|
||||
t.Fatalf("InsertConversation: %v", err)
|
||||
}
|
||||
|
||||
// Insert messages with different priorities
|
||||
for _, p := range []int{3, 8, 5} {
|
||||
msg := &Message{
|
||||
ConversationID: conv.ID,
|
||||
FromAgent: "sender",
|
||||
ToAgent: "reader",
|
||||
Body: "msg",
|
||||
Priority: p,
|
||||
Status: StatusPending,
|
||||
}
|
||||
if err := store.InsertMessage(ctx, msg); err != nil {
|
||||
t.Fatalf("InsertMessage: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
t.Run("returns messages ordered by priority desc", func(t *testing.T) {
|
||||
messages, err := store.GetInboxMessages(ctx, "reader", ReadOptions{IncludeRead: true})
|
||||
if err != nil {
|
||||
t.Fatalf("GetInboxMessages: %v", err)
|
||||
}
|
||||
if len(messages) != 3 {
|
||||
t.Fatalf("got %d messages, want 3", len(messages))
|
||||
}
|
||||
if messages[0].Priority != 8 {
|
||||
t.Errorf("first message priority = %d, want 8", messages[0].Priority)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("respects limit", func(t *testing.T) {
|
||||
messages, err := store.GetInboxMessages(ctx, "reader", ReadOptions{Limit: 1, IncludeRead: true})
|
||||
if err != nil {
|
||||
t.Fatalf("GetInboxMessages: %v", err)
|
||||
}
|
||||
if len(messages) != 1 {
|
||||
t.Errorf("got %d messages, want 1", len(messages))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("filters by status", func(t *testing.T) {
|
||||
messages, err := store.GetInboxMessages(ctx, "reader", ReadOptions{
|
||||
Status: StatusProcessing,
|
||||
IncludeRead: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("GetInboxMessages: %v", err)
|
||||
}
|
||||
if len(messages) != 0 {
|
||||
t.Errorf("got %d messages, want 0 (no processing messages)", len(messages))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("read/unread tracking", func(t *testing.T) {
|
||||
// Read all messages (unread only)
|
||||
messages, err := store.GetInboxMessages(ctx, "reader", ReadOptions{})
|
||||
if err != nil {
|
||||
t.Fatalf("GetInboxMessages: %v", err)
|
||||
}
|
||||
if len(messages) != 3 {
|
||||
t.Fatalf("got %d messages, want 3", len(messages))
|
||||
}
|
||||
|
||||
// Advance inbox state
|
||||
maxID := messages[0].ID
|
||||
for _, m := range messages {
|
||||
if m.ID > maxID {
|
||||
maxID = m.ID
|
||||
}
|
||||
}
|
||||
if err := store.UpdateInboxState(ctx, "reader", conv.ID, maxID); err != nil {
|
||||
t.Fatalf("UpdateInboxState: %v", err)
|
||||
}
|
||||
|
||||
// Read again without include_read — should be empty
|
||||
messages, err = store.GetInboxMessages(ctx, "reader", ReadOptions{})
|
||||
if err != nil {
|
||||
t.Fatalf("GetInboxMessages: %v", err)
|
||||
}
|
||||
if len(messages) != 0 {
|
||||
t.Errorf("got %d messages after reading, want 0", len(messages))
|
||||
}
|
||||
|
||||
// Read again with include_read — should return all
|
||||
messages, err = store.GetInboxMessages(ctx, "reader", ReadOptions{IncludeRead: true})
|
||||
if err != nil {
|
||||
t.Fatalf("GetInboxMessages: %v", err)
|
||||
}
|
||||
if len(messages) != 3 {
|
||||
t.Errorf("got %d messages with include_read, want 3", len(messages))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestSQLiteMessageStore_ClaimMessages(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteMessageStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
seedAgent(t, db, "sender")
|
||||
seedAgent(t, db, "worker")
|
||||
|
||||
conv := &Conversation{Subject: "claim test", CreatedBy: "sender"}
|
||||
if err := store.InsertConversation(ctx, conv); err != nil {
|
||||
t.Fatalf("InsertConversation: %v", err)
|
||||
}
|
||||
|
||||
// Insert 5 pending messages
|
||||
for i := 0; i < 5; i++ {
|
||||
msg := &Message{
|
||||
ConversationID: conv.ID,
|
||||
FromAgent: "sender",
|
||||
ToAgent: "worker",
|
||||
Body: "task",
|
||||
Priority: 5,
|
||||
Status: StatusPending,
|
||||
}
|
||||
if err := store.InsertMessage(ctx, msg); err != nil {
|
||||
t.Fatalf("InsertMessage: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
t.Run("claim with limit", func(t *testing.T) {
|
||||
claimed, err := store.ClaimMessages(ctx, "worker", 3)
|
||||
if err != nil {
|
||||
t.Fatalf("ClaimMessages: %v", err)
|
||||
}
|
||||
if len(claimed) != 3 {
|
||||
t.Errorf("claimed %d messages, want 3", len(claimed))
|
||||
}
|
||||
for _, msg := range claimed {
|
||||
if msg.Status != StatusProcessing {
|
||||
t.Errorf("claimed message status = %s, want %s", msg.Status, StatusProcessing)
|
||||
}
|
||||
if msg.ClaimedBy != "worker" {
|
||||
t.Errorf("claimed_by = %s, want worker", msg.ClaimedBy)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("already claimed messages are skipped", func(t *testing.T) {
|
||||
// Only 2 remaining pending
|
||||
claimed, err := store.ClaimMessages(ctx, "worker", 10)
|
||||
if err != nil {
|
||||
t.Fatalf("ClaimMessages: %v", err)
|
||||
}
|
||||
if len(claimed) != 2 {
|
||||
t.Errorf("claimed %d messages, want 2", len(claimed))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("no pending returns empty", func(t *testing.T) {
|
||||
claimed, err := store.ClaimMessages(ctx, "worker", 5)
|
||||
if err != nil {
|
||||
t.Fatalf("ClaimMessages: %v", err)
|
||||
}
|
||||
if len(claimed) != 0 {
|
||||
t.Errorf("claimed %d messages, want 0", len(claimed))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestSQLiteMessageStore_UpdateMessageStatus(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteMessageStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
seedAgent(t, db, "sender")
|
||||
seedAgent(t, db, "worker")
|
||||
|
||||
conv := &Conversation{Subject: "status test", CreatedBy: "sender"}
|
||||
if err := store.InsertConversation(ctx, conv); err != nil {
|
||||
t.Fatalf("InsertConversation: %v", err)
|
||||
}
|
||||
|
||||
msg := &Message{
|
||||
ConversationID: conv.ID,
|
||||
FromAgent: "sender",
|
||||
ToAgent: "worker",
|
||||
Body: "test",
|
||||
Priority: 5,
|
||||
Status: StatusPending,
|
||||
}
|
||||
if err := store.InsertMessage(ctx, msg); err != nil {
|
||||
t.Fatalf("InsertMessage: %v", err)
|
||||
}
|
||||
|
||||
// Claim the message first
|
||||
claimed, err := store.ClaimMessages(ctx, "worker", 1)
|
||||
if err != nil {
|
||||
t.Fatalf("ClaimMessages: %v", err)
|
||||
}
|
||||
if len(claimed) != 1 {
|
||||
t.Fatalf("claimed %d, want 1", len(claimed))
|
||||
}
|
||||
|
||||
t.Run("mark done by correct agent", func(t *testing.T) {
|
||||
err := store.UpdateMessageStatus(ctx, claimed[0].ID, StatusDone, "worker", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("UpdateMessageStatus: %v", err)
|
||||
}
|
||||
|
||||
got, err := store.GetMessageByID(ctx, claimed[0].ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetMessageByID: %v", err)
|
||||
}
|
||||
if got.Status != StatusDone {
|
||||
t.Errorf("status = %s, want %s", got.Status, StatusDone)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("wrong agent cannot update", func(t *testing.T) {
|
||||
err := store.UpdateMessageStatus(ctx, claimed[0].ID, StatusDone, "other-agent", nil)
|
||||
if err == nil {
|
||||
t.Error("expected error for wrong agent, got nil")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestSQLiteMessageStore_SearchMessages(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteMessageStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
seedAgent(t, db, "sender")
|
||||
seedAgent(t, db, "searcher")
|
||||
|
||||
conv := &Conversation{Subject: "search test", CreatedBy: "sender"}
|
||||
if err := store.InsertConversation(ctx, conv); err != nil {
|
||||
t.Fatalf("InsertConversation: %v", err)
|
||||
}
|
||||
|
||||
msgs := []struct {
|
||||
body string
|
||||
priority int
|
||||
}{
|
||||
{"deployment failure in production", 8},
|
||||
{"deployment succeeded on staging", 3},
|
||||
{"security alert: unauthorized access", 9},
|
||||
}
|
||||
|
||||
for _, m := range msgs {
|
||||
msg := &Message{
|
||||
ConversationID: conv.ID,
|
||||
FromAgent: "sender",
|
||||
ToAgent: "searcher",
|
||||
Body: m.body,
|
||||
Priority: m.priority,
|
||||
Status: StatusPending,
|
||||
}
|
||||
if err := store.InsertMessage(ctx, msg); err != nil {
|
||||
t.Fatalf("InsertMessage: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
t.Run("FTS keyword match", func(t *testing.T) {
|
||||
results, err := store.SearchMessages(ctx, "searcher", "deployment", SearchOptions{})
|
||||
if err != nil {
|
||||
t.Fatalf("SearchMessages: %v", err)
|
||||
}
|
||||
if len(results) != 2 {
|
||||
t.Errorf("got %d results, want 2", len(results))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("empty query returns recent", func(t *testing.T) {
|
||||
results, err := store.SearchMessages(ctx, "searcher", "", SearchOptions{})
|
||||
if err != nil {
|
||||
t.Fatalf("SearchMessages: %v", err)
|
||||
}
|
||||
if len(results) != 3 {
|
||||
t.Errorf("got %d results, want 3", len(results))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("min_priority filter", func(t *testing.T) {
|
||||
results, err := store.SearchMessages(ctx, "searcher", "", SearchOptions{MinPriority: 7})
|
||||
if err != nil {
|
||||
t.Fatalf("SearchMessages: %v", err)
|
||||
}
|
||||
if len(results) != 2 {
|
||||
t.Errorf("got %d results, want 2", len(results))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("limit", func(t *testing.T) {
|
||||
results, err := store.SearchMessages(ctx, "searcher", "", SearchOptions{Limit: 1})
|
||||
if err != nil {
|
||||
t.Fatalf("SearchMessages: %v", err)
|
||||
}
|
||||
if len(results) != 1 {
|
||||
t.Errorf("got %d results, want 1", len(results))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestSQLiteMessageStore_AgentExists(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteMessageStore(db)
|
||||
ctx := context.Background()
|
||||
|
||||
seedAgent(t, db, "exists-agent")
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
agent string
|
||||
exists bool
|
||||
}{
|
||||
{"existing agent", "exists-agent", true},
|
||||
{"non-existing agent", "ghost-agent", false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
exists, err := store.AgentExists(ctx, tt.agent)
|
||||
if err != nil {
|
||||
t.Fatalf("AgentExists: %v", err)
|
||||
}
|
||||
if exists != tt.exists {
|
||||
t.Errorf("AgentExists(%s) = %v, want %v", tt.agent, exists, tt.exists)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
// Package messaging provides core messaging types and services for SynapBus.
|
||||
package messaging
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Message status constants.
|
||||
const (
|
||||
StatusPending = "pending"
|
||||
StatusProcessing = "processing"
|
||||
StatusDone = "done"
|
||||
StatusFailed = "failed"
|
||||
)
|
||||
|
||||
// Message represents a single message in the system.
|
||||
type Message struct {
|
||||
ID int64 `json:"id"`
|
||||
ConversationID int64 `json:"conversation_id"`
|
||||
FromAgent string `json:"from_agent"`
|
||||
ToAgent string `json:"to_agent,omitempty"`
|
||||
ChannelID *int64 `json:"channel_id,omitempty"`
|
||||
Body string `json:"body"`
|
||||
Priority int `json:"priority"`
|
||||
Status string `json:"status"`
|
||||
Metadata json.RawMessage `json:"metadata"`
|
||||
ClaimedBy string `json:"claimed_by,omitempty"`
|
||||
ClaimedAt *time.Time `json:"claimed_at,omitempty"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
// Conversation groups related messages into a thread.
|
||||
type Conversation struct {
|
||||
ID int64 `json:"id"`
|
||||
Subject string `json:"subject"`
|
||||
CreatedBy string `json:"created_by"`
|
||||
ChannelID *int64 `json:"channel_id,omitempty"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
// InboxState tracks per-agent, per-conversation read position.
|
||||
type InboxState struct {
|
||||
AgentName string `json:"agent_name"`
|
||||
ConversationID int64 `json:"conversation_id"`
|
||||
LastReadMessageID int64 `json:"last_read_message_id"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
@@ -0,0 +1,149 @@
|
||||
package storage
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"embed"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"log/slog"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
//go:embed schema/*.sql
|
||||
var embeddedSchema embed.FS
|
||||
|
||||
// RunMigrations applies unapplied SQL migrations from the embedded schema directory.
|
||||
// Migrations are tracked in the schema_migrations table.
|
||||
func RunMigrations(ctx context.Context, db *sql.DB) error {
|
||||
return runMigrationsFromFS(ctx, db, embeddedSchema, "schema")
|
||||
}
|
||||
|
||||
// runMigrationsFromFS applies migrations from a given filesystem and directory path.
|
||||
// Exported for testing with custom migration files.
|
||||
func runMigrationsFromFS(ctx context.Context, db *sql.DB, fsys fs.FS, dir string) error {
|
||||
// Create schema_migrations table if it doesn't exist
|
||||
_, err := db.ExecContext(ctx, `
|
||||
CREATE TABLE IF NOT EXISTS schema_migrations (
|
||||
version INTEGER PRIMARY KEY,
|
||||
applied_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
)
|
||||
`)
|
||||
if err != nil {
|
||||
return fmt.Errorf("create schema_migrations table: %w", err)
|
||||
}
|
||||
|
||||
// Get applied versions
|
||||
applied, err := getAppliedVersions(ctx, db)
|
||||
if err != nil {
|
||||
return fmt.Errorf("get applied versions: %w", err)
|
||||
}
|
||||
|
||||
// Read migration files
|
||||
entries, err := fs.ReadDir(fsys, dir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("read migration directory: %w", err)
|
||||
}
|
||||
|
||||
// Parse and sort migration files
|
||||
type migration struct {
|
||||
version int
|
||||
filename string
|
||||
}
|
||||
var migrations []migration
|
||||
|
||||
for _, entry := range entries {
|
||||
if entry.IsDir() || !strings.HasSuffix(entry.Name(), ".sql") {
|
||||
continue
|
||||
}
|
||||
version, err := parseMigrationVersion(entry.Name())
|
||||
if err != nil {
|
||||
slog.Warn("skipping non-migration file", "filename", entry.Name(), "error", err)
|
||||
continue
|
||||
}
|
||||
migrations = append(migrations, migration{version: version, filename: entry.Name()})
|
||||
}
|
||||
|
||||
sort.Slice(migrations, func(i, j int) bool {
|
||||
return migrations[i].version < migrations[j].version
|
||||
})
|
||||
|
||||
// Apply unapplied migrations
|
||||
for _, m := range migrations {
|
||||
if applied[m.version] {
|
||||
slog.Debug("migration already applied", "version", m.version, "filename", m.filename)
|
||||
continue
|
||||
}
|
||||
|
||||
content, err := fs.ReadFile(fsys, dir+"/"+m.filename)
|
||||
if err != nil {
|
||||
return fmt.Errorf("read migration %s: %w", m.filename, err)
|
||||
}
|
||||
|
||||
if err := applyMigration(ctx, db, m.version, string(content)); err != nil {
|
||||
return fmt.Errorf("apply migration %s: %w", m.filename, err)
|
||||
}
|
||||
|
||||
slog.Info("migration applied", "version", m.version, "filename", m.filename)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// parseMigrationVersion extracts the version number from a migration filename.
|
||||
// Expected format: NNN_description.sql (e.g., 001_initial.sql)
|
||||
func parseMigrationVersion(filename string) (int, error) {
|
||||
parts := strings.SplitN(filename, "_", 2)
|
||||
if len(parts) < 2 {
|
||||
return 0, fmt.Errorf("invalid migration filename: %s", filename)
|
||||
}
|
||||
version, err := strconv.Atoi(parts[0])
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("invalid version number in %s: %w", filename, err)
|
||||
}
|
||||
return version, nil
|
||||
}
|
||||
|
||||
// getAppliedVersions returns a set of already-applied migration versions.
|
||||
func getAppliedVersions(ctx context.Context, db *sql.DB) (map[int]bool, error) {
|
||||
rows, err := db.QueryContext(ctx, "SELECT version FROM schema_migrations")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
applied := make(map[int]bool)
|
||||
for rows.Next() {
|
||||
var version int
|
||||
if err := rows.Scan(&version); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
applied[version] = true
|
||||
}
|
||||
return applied, rows.Err()
|
||||
}
|
||||
|
||||
// applyMigration runs a migration inside a transaction and records it.
|
||||
func applyMigration(ctx context.Context, db *sql.DB, version int, content string) error {
|
||||
tx, err := db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return fmt.Errorf("begin transaction: %w", err)
|
||||
}
|
||||
defer tx.Rollback()
|
||||
|
||||
// Execute migration SQL (may contain multiple statements)
|
||||
if _, err := tx.ExecContext(ctx, content); err != nil {
|
||||
return fmt.Errorf("execute migration: %w", err)
|
||||
}
|
||||
|
||||
// Record migration
|
||||
if _, err := tx.ExecContext(ctx,
|
||||
"INSERT OR IGNORE INTO schema_migrations (version) VALUES (?)", version,
|
||||
); err != nil {
|
||||
return fmt.Errorf("record migration: %w", err)
|
||||
}
|
||||
|
||||
return tx.Commit()
|
||||
}
|
||||
@@ -0,0 +1,196 @@
|
||||
package storage
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"testing"
|
||||
"testing/fstest"
|
||||
)
|
||||
|
||||
func newTestDB(t *testing.T) *sql.DB {
|
||||
t.Helper()
|
||||
db, err := sql.Open("sqlite", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open database: %v", err)
|
||||
}
|
||||
// Enable foreign keys for test DB
|
||||
if _, err := db.Exec("PRAGMA foreign_keys=ON"); err != nil {
|
||||
t.Fatalf("enable foreign keys: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { db.Close() })
|
||||
return db
|
||||
}
|
||||
|
||||
func TestRunMigrations(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
files fstest.MapFS
|
||||
wantErr bool
|
||||
wantCount int
|
||||
}{
|
||||
{
|
||||
name: "apply single migration",
|
||||
files: fstest.MapFS{
|
||||
"migrations/001_create_test.sql": &fstest.MapFile{
|
||||
Data: []byte(`CREATE TABLE IF NOT EXISTS test_table (
|
||||
id INTEGER PRIMARY KEY,
|
||||
name TEXT NOT NULL
|
||||
);`),
|
||||
},
|
||||
},
|
||||
wantErr: false,
|
||||
wantCount: 1,
|
||||
},
|
||||
{
|
||||
name: "apply multiple migrations in order",
|
||||
files: fstest.MapFS{
|
||||
"migrations/001_first.sql": &fstest.MapFile{
|
||||
Data: []byte(`CREATE TABLE IF NOT EXISTS first_table (id INTEGER PRIMARY KEY);`),
|
||||
},
|
||||
"migrations/002_second.sql": &fstest.MapFile{
|
||||
Data: []byte(`CREATE TABLE IF NOT EXISTS second_table (id INTEGER PRIMARY KEY);`),
|
||||
},
|
||||
},
|
||||
wantErr: false,
|
||||
wantCount: 2,
|
||||
},
|
||||
{
|
||||
name: "empty directory succeeds",
|
||||
files: fstest.MapFS{
|
||||
"migrations/readme.txt": &fstest.MapFile{Data: []byte("not a sql file")},
|
||||
},
|
||||
wantErr: false,
|
||||
wantCount: 0,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
ctx := context.Background()
|
||||
|
||||
err := runMigrationsFromFS(ctx, db, tt.files, "migrations")
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Fatalf("RunMigrations() error = %v, wantErr %v", err, tt.wantErr)
|
||||
}
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
// Verify migrations were recorded
|
||||
var count int
|
||||
err = db.QueryRow("SELECT COUNT(*) FROM schema_migrations").Scan(&count)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to count migrations: %v", err)
|
||||
}
|
||||
if count != tt.wantCount {
|
||||
t.Errorf("migration count = %d, want %d", count, tt.wantCount)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunMigrations_Idempotent(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
ctx := context.Background()
|
||||
|
||||
files := fstest.MapFS{
|
||||
"migrations/001_create.sql": &fstest.MapFile{
|
||||
Data: []byte(`CREATE TABLE IF NOT EXISTS idempotent_test (id INTEGER PRIMARY KEY);`),
|
||||
},
|
||||
}
|
||||
|
||||
// Run migrations twice
|
||||
if err := runMigrationsFromFS(ctx, db, files, "migrations"); err != nil {
|
||||
t.Fatalf("first run: %v", err)
|
||||
}
|
||||
if err := runMigrationsFromFS(ctx, db, files, "migrations"); err != nil {
|
||||
t.Fatalf("second run: %v", err)
|
||||
}
|
||||
|
||||
// Verify only one migration recorded
|
||||
var count int
|
||||
if err := db.QueryRow("SELECT COUNT(*) FROM schema_migrations").Scan(&count); err != nil {
|
||||
t.Fatalf("count migrations: %v", err)
|
||||
}
|
||||
if count != 1 {
|
||||
t.Errorf("migration count = %d after two runs, want 1", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunMigrations_SequentialOrder(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
ctx := context.Background()
|
||||
|
||||
files := fstest.MapFS{
|
||||
"migrations/003_third.sql": &fstest.MapFile{
|
||||
Data: []byte(`CREATE TABLE IF NOT EXISTS third (id INTEGER PRIMARY KEY);`),
|
||||
},
|
||||
"migrations/001_first.sql": &fstest.MapFile{
|
||||
Data: []byte(`CREATE TABLE IF NOT EXISTS first (id INTEGER PRIMARY KEY);`),
|
||||
},
|
||||
"migrations/002_second.sql": &fstest.MapFile{
|
||||
Data: []byte(`CREATE TABLE IF NOT EXISTS second (id INTEGER PRIMARY KEY);`),
|
||||
},
|
||||
}
|
||||
|
||||
if err := runMigrationsFromFS(ctx, db, files, "migrations"); err != nil {
|
||||
t.Fatalf("RunMigrations: %v", err)
|
||||
}
|
||||
|
||||
// Verify all three migrations were applied
|
||||
var count int
|
||||
if err := db.QueryRow("SELECT COUNT(*) FROM schema_migrations").Scan(&count); err != nil {
|
||||
t.Fatalf("count migrations: %v", err)
|
||||
}
|
||||
if count != 3 {
|
||||
t.Errorf("migration count = %d, want 3", count)
|
||||
}
|
||||
|
||||
// Verify order by checking applied_at ordering matches version ordering
|
||||
rows, err := db.Query("SELECT version FROM schema_migrations ORDER BY applied_at")
|
||||
if err != nil {
|
||||
t.Fatalf("query versions: %v", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var versions []int
|
||||
for rows.Next() {
|
||||
var v int
|
||||
if err := rows.Scan(&v); err != nil {
|
||||
t.Fatalf("scan version: %v", err)
|
||||
}
|
||||
versions = append(versions, v)
|
||||
}
|
||||
|
||||
if len(versions) != 3 {
|
||||
t.Fatalf("got %d versions, want 3", len(versions))
|
||||
}
|
||||
|
||||
for i := 1; i < len(versions); i++ {
|
||||
if versions[i] <= versions[i-1] {
|
||||
t.Errorf("migrations not in order: %v", versions)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunMigrations_EmbeddedSchema(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Run the actual embedded migrations
|
||||
if err := RunMigrations(ctx, db); err != nil {
|
||||
t.Fatalf("RunMigrations (embedded): %v", err)
|
||||
}
|
||||
|
||||
// Verify key tables exist
|
||||
tables := []string{"messages", "conversations", "agents", "traces", "inbox_state", "channels"}
|
||||
for _, table := range tables {
|
||||
var name string
|
||||
err := db.QueryRow("SELECT name FROM sqlite_master WHERE type='table' AND name=?", table).Scan(&name)
|
||||
if err != nil {
|
||||
t.Errorf("table %s not found: %v", table, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,216 @@
|
||||
-- SynapBus initial schema
|
||||
-- All tables use INTEGER PRIMARY KEY for SQLite rowid alias
|
||||
|
||||
-- Human user accounts (OAuth 2.1)
|
||||
CREATE TABLE IF NOT EXISTS users (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
username TEXT NOT NULL UNIQUE,
|
||||
password_hash TEXT NOT NULL,
|
||||
display_name TEXT NOT NULL DEFAULT '',
|
||||
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
-- Registered agents
|
||||
CREATE TABLE IF NOT EXISTS agents (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
name TEXT NOT NULL UNIQUE,
|
||||
display_name TEXT NOT NULL DEFAULT '',
|
||||
type TEXT NOT NULL DEFAULT 'ai' CHECK (type IN ('ai', 'human')),
|
||||
capabilities TEXT NOT NULL DEFAULT '{}', -- JSON
|
||||
owner_id INTEGER NOT NULL REFERENCES users(id),
|
||||
api_key_hash TEXT NOT NULL,
|
||||
status TEXT NOT NULL DEFAULT 'active' CHECK (status IN ('active', 'inactive')),
|
||||
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX idx_agents_owner ON agents(owner_id);
|
||||
CREATE INDEX idx_agents_status ON agents(status);
|
||||
|
||||
-- Conversations (threads)
|
||||
CREATE TABLE IF NOT EXISTS conversations (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
subject TEXT NOT NULL DEFAULT '',
|
||||
created_by TEXT NOT NULL, -- agent name
|
||||
channel_id INTEGER REFERENCES channels(id),
|
||||
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
-- Messages
|
||||
CREATE TABLE IF NOT EXISTS messages (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
conversation_id INTEGER NOT NULL REFERENCES conversations(id),
|
||||
from_agent TEXT NOT NULL,
|
||||
to_agent TEXT, -- NULL for channel messages
|
||||
channel_id INTEGER REFERENCES channels(id),
|
||||
body TEXT NOT NULL,
|
||||
priority INTEGER NOT NULL DEFAULT 5 CHECK (priority BETWEEN 1 AND 10),
|
||||
status TEXT NOT NULL DEFAULT 'pending' CHECK (status IN ('pending', 'processing', 'done', 'failed')),
|
||||
metadata TEXT NOT NULL DEFAULT '{}', -- JSON
|
||||
claimed_by TEXT, -- agent that claimed for processing
|
||||
claimed_at TIMESTAMP,
|
||||
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX idx_messages_conversation ON messages(conversation_id);
|
||||
CREATE INDEX idx_messages_from ON messages(from_agent);
|
||||
CREATE INDEX idx_messages_to ON messages(to_agent);
|
||||
CREATE INDEX idx_messages_channel ON messages(channel_id);
|
||||
CREATE INDEX idx_messages_status ON messages(status);
|
||||
CREATE INDEX idx_messages_priority ON messages(priority);
|
||||
CREATE INDEX idx_messages_created ON messages(created_at);
|
||||
|
||||
-- Full-text search index for messages
|
||||
CREATE VIRTUAL TABLE IF NOT EXISTS messages_fts USING fts5(
|
||||
body,
|
||||
content='messages',
|
||||
content_rowid='id'
|
||||
);
|
||||
|
||||
-- Triggers to keep FTS in sync
|
||||
CREATE TRIGGER messages_ai AFTER INSERT ON messages BEGIN
|
||||
INSERT INTO messages_fts(rowid, body) VALUES (new.id, new.body);
|
||||
END;
|
||||
CREATE TRIGGER messages_ad AFTER DELETE ON messages BEGIN
|
||||
INSERT INTO messages_fts(messages_fts, rowid, body) VALUES('delete', old.id, old.body);
|
||||
END;
|
||||
CREATE TRIGGER messages_au AFTER UPDATE ON messages BEGIN
|
||||
INSERT INTO messages_fts(messages_fts, rowid, body) VALUES('delete', old.id, old.body);
|
||||
INSERT INTO messages_fts(rowid, body) VALUES (new.id, new.body);
|
||||
END;
|
||||
|
||||
-- Channels
|
||||
CREATE TABLE IF NOT EXISTS channels (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
name TEXT NOT NULL UNIQUE,
|
||||
description TEXT NOT NULL DEFAULT '',
|
||||
topic TEXT NOT NULL DEFAULT '',
|
||||
type TEXT NOT NULL DEFAULT 'standard' CHECK (type IN ('standard', 'blackboard', 'auction')),
|
||||
is_private INTEGER NOT NULL DEFAULT 0,
|
||||
created_by TEXT NOT NULL, -- agent name
|
||||
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
-- Channel membership
|
||||
CREATE TABLE IF NOT EXISTS channel_members (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
channel_id INTEGER NOT NULL REFERENCES channels(id) ON DELETE CASCADE,
|
||||
agent_name TEXT NOT NULL,
|
||||
role TEXT NOT NULL DEFAULT 'member' CHECK (role IN ('owner', 'member')),
|
||||
joined_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
UNIQUE(channel_id, agent_name)
|
||||
);
|
||||
|
||||
CREATE INDEX idx_channel_members_channel ON channel_members(channel_id);
|
||||
CREATE INDEX idx_channel_members_agent ON channel_members(agent_name);
|
||||
|
||||
-- Read/unread tracking per agent per conversation
|
||||
CREATE TABLE IF NOT EXISTS inbox_state (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
agent_name TEXT NOT NULL,
|
||||
conversation_id INTEGER NOT NULL REFERENCES conversations(id),
|
||||
last_read_message_id INTEGER NOT NULL DEFAULT 0,
|
||||
updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
UNIQUE(agent_name, conversation_id)
|
||||
);
|
||||
|
||||
CREATE INDEX idx_inbox_state_agent ON inbox_state(agent_name);
|
||||
|
||||
-- Agent activity traces
|
||||
CREATE TABLE IF NOT EXISTS traces (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
agent_name TEXT NOT NULL,
|
||||
action TEXT NOT NULL,
|
||||
details TEXT NOT NULL DEFAULT '{}', -- JSON
|
||||
error TEXT,
|
||||
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX idx_traces_agent ON traces(agent_name);
|
||||
CREATE INDEX idx_traces_action ON traces(action);
|
||||
CREATE INDEX idx_traces_created ON traces(created_at);
|
||||
|
||||
-- Attachments (content-addressable)
|
||||
CREATE TABLE IF NOT EXISTS attachments (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
hash TEXT NOT NULL, -- SHA-256
|
||||
original_filename TEXT NOT NULL,
|
||||
size INTEGER NOT NULL,
|
||||
mime_type TEXT NOT NULL DEFAULT 'application/octet-stream',
|
||||
message_id INTEGER REFERENCES messages(id),
|
||||
uploaded_by TEXT NOT NULL, -- agent name
|
||||
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX idx_attachments_hash ON attachments(hash);
|
||||
CREATE INDEX idx_attachments_message ON attachments(message_id);
|
||||
|
||||
-- OAuth 2.1 tokens
|
||||
CREATE TABLE IF NOT EXISTS oauth_tokens (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
client_id TEXT NOT NULL,
|
||||
user_id INTEGER REFERENCES users(id),
|
||||
access_token_hash TEXT NOT NULL UNIQUE,
|
||||
refresh_token_hash TEXT,
|
||||
scope TEXT NOT NULL DEFAULT '',
|
||||
expires_at TIMESTAMP NOT NULL,
|
||||
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX idx_oauth_tokens_client ON oauth_tokens(client_id);
|
||||
CREATE INDEX idx_oauth_tokens_user ON oauth_tokens(user_id);
|
||||
|
||||
-- OAuth 2.1 clients
|
||||
CREATE TABLE IF NOT EXISTS oauth_clients (
|
||||
id TEXT PRIMARY KEY,
|
||||
secret_hash TEXT NOT NULL,
|
||||
name TEXT NOT NULL,
|
||||
redirect_uris TEXT NOT NULL DEFAULT '[]', -- JSON array
|
||||
grant_types TEXT NOT NULL DEFAULT '[]', -- JSON array
|
||||
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
-- Task auction support
|
||||
CREATE TABLE IF NOT EXISTS tasks (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
channel_id INTEGER NOT NULL REFERENCES channels(id),
|
||||
posted_by TEXT NOT NULL,
|
||||
title TEXT NOT NULL,
|
||||
description TEXT NOT NULL DEFAULT '',
|
||||
requirements TEXT NOT NULL DEFAULT '{}', -- JSON
|
||||
deadline TIMESTAMP,
|
||||
status TEXT NOT NULL DEFAULT 'open' CHECK (status IN ('open', 'assigned', 'completed', 'cancelled')),
|
||||
assigned_to TEXT,
|
||||
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX idx_tasks_channel ON tasks(channel_id);
|
||||
CREATE INDEX idx_tasks_status ON tasks(status);
|
||||
|
||||
-- Task bids
|
||||
CREATE TABLE IF NOT EXISTS task_bids (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
task_id INTEGER NOT NULL REFERENCES tasks(id) ON DELETE CASCADE,
|
||||
agent_name TEXT NOT NULL,
|
||||
capabilities TEXT NOT NULL DEFAULT '{}', -- JSON
|
||||
time_estimate TEXT,
|
||||
message TEXT NOT NULL DEFAULT '',
|
||||
status TEXT NOT NULL DEFAULT 'pending' CHECK (status IN ('pending', 'accepted', 'rejected')),
|
||||
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
UNIQUE(task_id, agent_name)
|
||||
);
|
||||
|
||||
CREATE INDEX idx_task_bids_task ON task_bids(task_id);
|
||||
|
||||
-- Schema version tracking
|
||||
CREATE TABLE IF NOT EXISTS schema_migrations (
|
||||
version INTEGER PRIMARY KEY,
|
||||
applied_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
INSERT INTO schema_migrations (version) VALUES (1);
|
||||
@@ -0,0 +1,71 @@
|
||||
// Package storage provides SQLite storage layer for SynapBus.
|
||||
package storage
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
)
|
||||
|
||||
// DB wraps a *sql.DB with SynapBus-specific configuration.
|
||||
type DB struct {
|
||||
*sql.DB
|
||||
}
|
||||
|
||||
// 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.
|
||||
func New(ctx context.Context, dataDir string) (*DB, error) {
|
||||
var dsn string
|
||||
|
||||
if dataDir == "" || dataDir == ":memory:" {
|
||||
dsn = ":memory:"
|
||||
} else {
|
||||
if err := os.MkdirAll(dataDir, 0o755); err != nil {
|
||||
return nil, fmt.Errorf("create data directory: %w", err)
|
||||
}
|
||||
dsn = filepath.Join(dataDir, "synapbus.db")
|
||||
}
|
||||
|
||||
db, err := sql.Open("sqlite", dsn)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open database: %w", err)
|
||||
}
|
||||
|
||||
// Configure SQLite pragmas
|
||||
pragmas := []string{
|
||||
"PRAGMA journal_mode=WAL",
|
||||
"PRAGMA busy_timeout=5000",
|
||||
"PRAGMA foreign_keys=ON",
|
||||
}
|
||||
|
||||
for _, pragma := range pragmas {
|
||||
if _, err := db.ExecContext(ctx, pragma); err != nil {
|
||||
db.Close()
|
||||
return nil, fmt.Errorf("execute %s: %w", pragma, err)
|
||||
}
|
||||
}
|
||||
|
||||
// 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)
|
||||
}
|
||||
|
||||
slog.Info("database opened",
|
||||
"dsn", dsn,
|
||||
"journal_mode", journalMode,
|
||||
)
|
||||
|
||||
return &DB{DB: db}, nil
|
||||
}
|
||||
|
||||
// Close closes the database connection.
|
||||
func (db *DB) Close() error {
|
||||
return db.DB.Close()
|
||||
}
|
||||
@@ -0,0 +1,82 @@
|
||||
package storage
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestNew(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
dataDir string
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "in-memory database",
|
||||
dataDir: ":memory:",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "empty string creates in-memory",
|
||||
dataDir: "",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "temp directory",
|
||||
dataDir: t.TempDir(),
|
||||
wantErr: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
db, err := New(ctx, tt.dataDir)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Fatalf("New() error = %v, wantErr %v", err, tt.wantErr)
|
||||
}
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
// Verify WAL mode (in-memory uses "memory" journal mode)
|
||||
if tt.dataDir != "" && tt.dataDir != ":memory:" {
|
||||
var journalMode string
|
||||
err = db.QueryRow("PRAGMA journal_mode").Scan(&journalMode)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to query journal_mode: %v", err)
|
||||
}
|
||||
if journalMode != "wal" {
|
||||
t.Errorf("journal_mode = %s, want wal", journalMode)
|
||||
}
|
||||
}
|
||||
|
||||
// Verify foreign keys are enabled
|
||||
var fk int
|
||||
err = db.QueryRow("PRAGMA foreign_keys").Scan(&fk)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to query foreign_keys: %v", err)
|
||||
}
|
||||
if fk != 1 {
|
||||
t.Errorf("foreign_keys = %d, want 1", fk)
|
||||
}
|
||||
|
||||
// Verify busy_timeout
|
||||
var timeout int
|
||||
err = db.QueryRow("PRAGMA busy_timeout").Scan(&timeout)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to query busy_timeout: %v", err)
|
||||
}
|
||||
if timeout != 5000 {
|
||||
t.Errorf("busy_timeout = %d, want 5000", timeout)
|
||||
}
|
||||
|
||||
// Verify database is usable
|
||||
_, err = db.Exec("CREATE TABLE test (id INTEGER PRIMARY KEY)")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create test table: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,87 @@
|
||||
package trace
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// StoredTrace represents a trace entry read from the database.
|
||||
type StoredTrace struct {
|
||||
ID int64
|
||||
AgentName string
|
||||
Action string
|
||||
Details string
|
||||
Error sql.NullString
|
||||
CreatedAt time.Time
|
||||
}
|
||||
|
||||
// TraceStore defines the interface for reading trace entries.
|
||||
type TraceStore interface {
|
||||
GetTraces(ctx context.Context, agentName string, limit int) ([]*StoredTrace, error)
|
||||
GetTracesByAction(ctx context.Context, action string, limit int) ([]*StoredTrace, error)
|
||||
}
|
||||
|
||||
// SQLiteTraceStore implements TraceStore using SQLite.
|
||||
type SQLiteTraceStore struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
// NewSQLiteTraceStore creates a new SQLite-backed trace store.
|
||||
func NewSQLiteTraceStore(db *sql.DB) *SQLiteTraceStore {
|
||||
return &SQLiteTraceStore{db: db}
|
||||
}
|
||||
|
||||
func (s *SQLiteTraceStore) GetTraces(ctx context.Context, agentName string, limit int) ([]*StoredTrace, error) {
|
||||
if limit <= 0 {
|
||||
limit = 50
|
||||
}
|
||||
|
||||
rows, err := s.db.QueryContext(ctx,
|
||||
`SELECT id, agent_name, action, details, error, created_at
|
||||
FROM traces WHERE agent_name = ?
|
||||
ORDER BY created_at DESC LIMIT ?`,
|
||||
agentName, limit,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("query traces: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
return scanTraces(rows)
|
||||
}
|
||||
|
||||
func (s *SQLiteTraceStore) GetTracesByAction(ctx context.Context, action string, limit int) ([]*StoredTrace, error) {
|
||||
if limit <= 0 {
|
||||
limit = 50
|
||||
}
|
||||
|
||||
rows, err := s.db.QueryContext(ctx,
|
||||
`SELECT id, agent_name, action, details, error, created_at
|
||||
FROM traces WHERE action = ?
|
||||
ORDER BY created_at DESC LIMIT ?`,
|
||||
action, limit,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("query traces: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
return scanTraces(rows)
|
||||
}
|
||||
|
||||
func scanTraces(rows *sql.Rows) ([]*StoredTrace, error) {
|
||||
var traces []*StoredTrace
|
||||
for rows.Next() {
|
||||
var t StoredTrace
|
||||
if err := rows.Scan(&t.ID, &t.AgentName, &t.Action, &t.Details, &t.Error, &t.CreatedAt); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
traces = append(traces, &t)
|
||||
}
|
||||
if traces == nil {
|
||||
traces = []*StoredTrace{}
|
||||
}
|
||||
return traces, rows.Err()
|
||||
}
|
||||
@@ -0,0 +1,125 @@
|
||||
// Package trace provides agent activity trace recording for SynapBus.
|
||||
package trace
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
)
|
||||
|
||||
// TraceEntry represents a single trace record.
|
||||
type TraceEntry struct {
|
||||
AgentName string
|
||||
Action string
|
||||
Details any
|
||||
Error string
|
||||
}
|
||||
|
||||
// Tracer records agent actions to the traces table.
|
||||
// It uses a buffered channel for async recording.
|
||||
type Tracer struct {
|
||||
db *sql.DB
|
||||
logger *slog.Logger
|
||||
ch chan TraceEntry
|
||||
done chan struct{}
|
||||
}
|
||||
|
||||
// NewTracer creates a new Tracer with a buffered channel.
|
||||
func NewTracer(db *sql.DB) *Tracer {
|
||||
t := &Tracer{
|
||||
db: db,
|
||||
logger: slog.Default().With("component", "tracer"),
|
||||
ch: make(chan TraceEntry, 256),
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
go t.processLoop()
|
||||
return t
|
||||
}
|
||||
|
||||
// Record enqueues a trace entry for async storage.
|
||||
func (t *Tracer) Record(ctx context.Context, agentName, action string, details any) {
|
||||
entry := TraceEntry{
|
||||
AgentName: agentName,
|
||||
Action: action,
|
||||
Details: details,
|
||||
}
|
||||
|
||||
select {
|
||||
case t.ch <- entry:
|
||||
default:
|
||||
t.logger.Warn("trace channel full, dropping entry",
|
||||
"agent", agentName,
|
||||
"action", action,
|
||||
)
|
||||
}
|
||||
|
||||
t.logger.Info("trace recorded",
|
||||
"agent", agentName,
|
||||
"action", action,
|
||||
)
|
||||
}
|
||||
|
||||
// RecordError enqueues a trace entry with an error.
|
||||
func (t *Tracer) RecordError(ctx context.Context, agentName, action string, details any, traceErr error) {
|
||||
entry := TraceEntry{
|
||||
AgentName: agentName,
|
||||
Action: action,
|
||||
Details: details,
|
||||
}
|
||||
if traceErr != nil {
|
||||
entry.Error = traceErr.Error()
|
||||
}
|
||||
|
||||
select {
|
||||
case t.ch <- entry:
|
||||
default:
|
||||
t.logger.Warn("trace channel full, dropping entry",
|
||||
"agent", agentName,
|
||||
"action", action,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
// Close stops the tracer and flushes remaining entries.
|
||||
func (t *Tracer) Close() {
|
||||
close(t.ch)
|
||||
<-t.done
|
||||
}
|
||||
|
||||
func (t *Tracer) processLoop() {
|
||||
defer close(t.done)
|
||||
for entry := range t.ch {
|
||||
t.writeEntry(entry)
|
||||
}
|
||||
}
|
||||
|
||||
func (t *Tracer) writeEntry(entry TraceEntry) {
|
||||
detailsJSON, err := json.Marshal(entry.Details)
|
||||
if err != nil {
|
||||
t.logger.Error("failed to marshal trace details",
|
||||
"error", err,
|
||||
"agent", entry.AgentName,
|
||||
"action", entry.Action,
|
||||
)
|
||||
detailsJSON = []byte("{}")
|
||||
}
|
||||
|
||||
var traceErr sql.NullString
|
||||
if entry.Error != "" {
|
||||
traceErr = sql.NullString{String: entry.Error, Valid: true}
|
||||
}
|
||||
|
||||
_, err = t.db.Exec(
|
||||
`INSERT INTO traces (agent_name, action, details, error, created_at)
|
||||
VALUES (?, ?, ?, ?, CURRENT_TIMESTAMP)`,
|
||||
entry.AgentName, entry.Action, string(detailsJSON), traceErr,
|
||||
)
|
||||
if err != nil {
|
||||
t.logger.Error("failed to write trace entry",
|
||||
"error", err,
|
||||
"agent", entry.AgentName,
|
||||
"action", entry.Action,
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,166 @@
|
||||
package trace
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
)
|
||||
|
||||
func newTestDB(t *testing.T) *sql.DB {
|
||||
t.Helper()
|
||||
db, err := sql.Open("sqlite", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open database: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { db.Close() })
|
||||
|
||||
// Create traces table
|
||||
_, err = db.Exec(`
|
||||
CREATE TABLE IF NOT EXISTS traces (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
agent_name TEXT NOT NULL,
|
||||
action TEXT NOT NULL,
|
||||
details TEXT NOT NULL DEFAULT '{}',
|
||||
error TEXT,
|
||||
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
)
|
||||
`)
|
||||
if err != nil {
|
||||
t.Fatalf("create traces table: %v", err)
|
||||
}
|
||||
|
||||
return db
|
||||
}
|
||||
|
||||
func TestTracer_Record(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
agent string
|
||||
action string
|
||||
details any
|
||||
}{
|
||||
{
|
||||
name: "simple trace",
|
||||
agent: "test-agent",
|
||||
action: "send_message",
|
||||
details: map[string]any{
|
||||
"to": "other-agent",
|
||||
"message": "hello",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "trace with nil details",
|
||||
agent: "agent-a",
|
||||
action: "read_inbox",
|
||||
details: nil,
|
||||
},
|
||||
{
|
||||
name: "trace with string details",
|
||||
agent: "agent-b",
|
||||
action: "search",
|
||||
details: "query string",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
tracer := NewTracer(db)
|
||||
defer tracer.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
tracer.Record(ctx, tt.agent, tt.action, tt.details)
|
||||
|
||||
// Give async writer time to process
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
// Verify trace was written
|
||||
var count int
|
||||
err := db.QueryRow(
|
||||
"SELECT COUNT(*) FROM traces WHERE agent_name = ? AND action = ?",
|
||||
tt.agent, tt.action,
|
||||
).Scan(&count)
|
||||
if err != nil {
|
||||
t.Fatalf("query trace: %v", err)
|
||||
}
|
||||
if count != 1 {
|
||||
t.Errorf("trace count = %d, want 1", count)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestTracer_RecordError(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
tracer := NewTracer(db)
|
||||
defer tracer.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
tracer.RecordError(ctx, "test-agent", "failed_action",
|
||||
map[string]any{"key": "value"},
|
||||
fmt.Errorf("something went wrong"),
|
||||
)
|
||||
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
var errorText sql.NullString
|
||||
err := db.QueryRow(
|
||||
"SELECT error FROM traces WHERE agent_name = 'test-agent' AND action = 'failed_action'",
|
||||
).Scan(&errorText)
|
||||
if err != nil {
|
||||
t.Fatalf("query trace: %v", err)
|
||||
}
|
||||
if !errorText.Valid || errorText.String != "something went wrong" {
|
||||
t.Errorf("error = %v, want 'something went wrong'", errorText)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTraceStore(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
tracer := NewTracer(db)
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Record some traces
|
||||
tracer.Record(ctx, "agent-a", "send_message", map[string]any{"to": "agent-b"})
|
||||
tracer.Record(ctx, "agent-a", "read_inbox", map[string]any{"count": 5})
|
||||
tracer.Record(ctx, "agent-b", "send_message", map[string]any{"to": "agent-a"})
|
||||
|
||||
tracer.Close() // flush all entries
|
||||
|
||||
store := NewSQLiteTraceStore(db)
|
||||
|
||||
t.Run("get traces by agent", func(t *testing.T) {
|
||||
traces, err := store.GetTraces(ctx, "agent-a", 10)
|
||||
if err != nil {
|
||||
t.Fatalf("GetTraces: %v", err)
|
||||
}
|
||||
if len(traces) != 2 {
|
||||
t.Errorf("got %d traces, want 2", len(traces))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("get traces by action", func(t *testing.T) {
|
||||
traces, err := store.GetTracesByAction(ctx, "send_message", 10)
|
||||
if err != nil {
|
||||
t.Fatalf("GetTracesByAction: %v", err)
|
||||
}
|
||||
if len(traces) != 2 {
|
||||
t.Errorf("got %d traces, want 2", len(traces))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("limit results", func(t *testing.T) {
|
||||
traces, err := store.GetTraces(ctx, "agent-a", 1)
|
||||
if err != nil {
|
||||
t.Fatalf("GetTraces: %v", err)
|
||||
}
|
||||
if len(traces) != 1 {
|
||||
t.Errorf("got %d traces, want 1", len(traces))
|
||||
}
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user