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:
Algis Dumbris
2026-03-13 11:45:22 +02:00
co-authored by Claude Opus 4.6
parent f1bbb86a67
commit 2f55ce87c3
31 changed files with 5395 additions and 8 deletions
+1 -1
View File
@@ -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
View File
@@ -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
}
+29 -5
View File
@@ -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
)
+93
View File
@@ -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=
+61
View File
@@ -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))
})
}
}
+89
View File
@@ -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)
}
})
}
}
+219
View File
@@ -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
}
+248
View File
@@ -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))
}
}
+176
View File
@@ -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()
}
+250
View File
@@ -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
+27
View File
@@ -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"`
}
+35
View File
@@ -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
}
+78
View File
@@ -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()
}
}
+31
View File
@@ -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)
}
}
+81
View File
@@ -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
}
+408
View File
@@ -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
}
+390
View File
@@ -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
+30
View File
@@ -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"`
}
+330
View File
@@ -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
}
+526
View File
@@ -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
+519
View File
@@ -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
}
+536
View File
@@ -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)
}
})
}
}
+50
View File
@@ -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"`
}
+149
View File
@@ -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()
}
+196
View File
@@ -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)
}
}
}
+216
View File
@@ -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);
+71
View File
@@ -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()
}
+82
View File
@@ -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)
}
})
}
}
+87
View File
@@ -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()
}
+125
View File
@@ -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,
)
}
}
+166
View File
@@ -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))
}
})
}