diff --git a/Makefile b/Makefile index bdbce5d..422ef70 100644 --- a/Makefile +++ b/Makefile @@ -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 diff --git a/cmd/synapbus/main.go b/cmd/synapbus/main.go index 721df85..f16cccd 100644 --- a/cmd/synapbus/main.go +++ b/cmd/synapbus/main.go @@ -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 } diff --git a/go.mod b/go.mod index 598320d..45671c6 100644 --- a/go.mod +++ b/go.mod @@ -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 ) diff --git a/go.sum b/go.sum index a6ee3e0..8823752 100644 --- a/go.sum +++ b/go.sum @@ -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= diff --git a/internal/agents/middleware.go b/internal/agents/middleware.go new file mode 100644 index 0000000..30f50a6 --- /dev/null +++ b/internal/agents/middleware.go @@ -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)) + }) + } +} diff --git a/internal/agents/middleware_test.go b/internal/agents/middleware_test.go new file mode 100644 index 0000000..8995368 --- /dev/null +++ b/internal/agents/middleware_test.go @@ -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) + } + }) + } +} diff --git a/internal/agents/service.go b/internal/agents/service.go new file mode 100644 index 0000000..0a8679d --- /dev/null +++ b/internal/agents/service.go @@ -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 +} diff --git a/internal/agents/service_test.go b/internal/agents/service_test.go new file mode 100644 index 0000000..380d6d1 --- /dev/null +++ b/internal/agents/service_test.go @@ -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)) + } +} diff --git a/internal/agents/store.go b/internal/agents/store.go new file mode 100644 index 0000000..7701802 --- /dev/null +++ b/internal/agents/store.go @@ -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() +} diff --git a/internal/agents/store_test.go b/internal/agents/store_test.go new file mode 100644 index 0000000..1bcb4c2 --- /dev/null +++ b/internal/agents/store_test.go @@ -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 diff --git a/internal/agents/types.go b/internal/agents/types.go new file mode 100644 index 0000000..9d4935b --- /dev/null +++ b/internal/agents/types.go @@ -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"` +} diff --git a/internal/mcp/auth.go b/internal/mcp/auth.go new file mode 100644 index 0000000..5210d08 --- /dev/null +++ b/internal/mcp/auth.go @@ -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 +} diff --git a/internal/mcp/connection.go b/internal/mcp/connection.go new file mode 100644 index 0000000..6f3c338 --- /dev/null +++ b/internal/mcp/connection.go @@ -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() + } +} diff --git a/internal/mcp/health.go b/internal/mcp/health.go new file mode 100644 index 0000000..01c23e4 --- /dev/null +++ b/internal/mcp/health.go @@ -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) + } +} diff --git a/internal/mcp/server.go b/internal/mcp/server.go new file mode 100644 index 0000000..0c8866c --- /dev/null +++ b/internal/mcp/server.go @@ -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 +} diff --git a/internal/mcp/tools.go b/internal/mcp/tools.go new file mode 100644 index 0000000..1ba7d63 --- /dev/null +++ b/internal/mcp/tools.go @@ -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 +} diff --git a/internal/mcp/tools_test.go b/internal/mcp/tools_test.go new file mode 100644 index 0000000..3554f6d --- /dev/null +++ b/internal/mcp/tools_test.go @@ -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 diff --git a/internal/messaging/options.go b/internal/messaging/options.go new file mode 100644 index 0000000..38c9404 --- /dev/null +++ b/internal/messaging/options.go @@ -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"` +} diff --git a/internal/messaging/service.go b/internal/messaging/service.go new file mode 100644 index 0000000..5b09078 --- /dev/null +++ b/internal/messaging/service.go @@ -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 +} diff --git a/internal/messaging/service_test.go b/internal/messaging/service_test.go new file mode 100644 index 0000000..3b59eac --- /dev/null +++ b/internal/messaging/service_test.go @@ -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 diff --git a/internal/messaging/store.go b/internal/messaging/store.go new file mode 100644 index 0000000..3abe296 --- /dev/null +++ b/internal/messaging/store.go @@ -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 +} diff --git a/internal/messaging/store_test.go b/internal/messaging/store_test.go new file mode 100644 index 0000000..4318826 --- /dev/null +++ b/internal/messaging/store_test.go @@ -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) + } + }) + } +} diff --git a/internal/messaging/types.go b/internal/messaging/types.go new file mode 100644 index 0000000..14f33dc --- /dev/null +++ b/internal/messaging/types.go @@ -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"` +} diff --git a/internal/storage/migrations.go b/internal/storage/migrations.go new file mode 100644 index 0000000..e37a5cd --- /dev/null +++ b/internal/storage/migrations.go @@ -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() +} diff --git a/internal/storage/migrations_test.go b/internal/storage/migrations_test.go new file mode 100644 index 0000000..43a9956 --- /dev/null +++ b/internal/storage/migrations_test.go @@ -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) + } + } +} diff --git a/internal/storage/schema/001_initial.sql b/internal/storage/schema/001_initial.sql new file mode 100644 index 0000000..5d9ab03 --- /dev/null +++ b/internal/storage/schema/001_initial.sql @@ -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); diff --git a/internal/storage/sqlite.go b/internal/storage/sqlite.go new file mode 100644 index 0000000..fbe1c30 --- /dev/null +++ b/internal/storage/sqlite.go @@ -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() +} diff --git a/internal/storage/sqlite_test.go b/internal/storage/sqlite_test.go new file mode 100644 index 0000000..2406cfc --- /dev/null +++ b/internal/storage/sqlite_test.go @@ -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) + } + }) + } +} diff --git a/internal/trace/store.go b/internal/trace/store.go new file mode 100644 index 0000000..ebf596c --- /dev/null +++ b/internal/trace/store.go @@ -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() +} diff --git a/internal/trace/tracer.go b/internal/trace/tracer.go new file mode 100644 index 0000000..a4c9371 --- /dev/null +++ b/internal/trace/tracer.go @@ -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, + ) + } +} diff --git a/internal/trace/tracer_test.go b/internal/trace/tracer_test.go new file mode 100644 index 0000000..3f872fd --- /dev/null +++ b/internal/trace/tracer_test.go @@ -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)) + } + }) +}