feat: merge semantic search, attachments, swarm patterns (Round 3)
- Semantic Search: HNSW vector index, OpenAI/Ollama providers, FTS5 fallback - Attachments: content-addressable storage, SHA-256 dedup, MCP tools - Swarm Patterns: task auction, stigmergy, agent discovery, expiry worker Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
+84
-2
@@ -25,6 +25,8 @@ import (
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/channels"
|
||||
mcpserver "github.com/smart-mcp-proxy/synapbus/internal/mcp"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/messaging"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/search"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/search/embedding"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/storage"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/trace"
|
||||
)
|
||||
@@ -236,8 +238,74 @@ func runServe(cmd *cobra.Command, args []string) error {
|
||||
fmt.Printf("========================================\n\n")
|
||||
}
|
||||
|
||||
// Create MCP server (with swarm + attachment tools)
|
||||
mcpSrv := mcpserver.NewMCPServer(msgService, agentService, channelService, swarmService, attachmentService)
|
||||
// Initialize search subsystem
|
||||
searchCfg := search.LoadConfigFromEnv()
|
||||
var searchService *search.Service
|
||||
var embPipeline *search.Pipeline
|
||||
var vectorIndex *search.VectorIndex
|
||||
|
||||
if searchCfg.IsEnabled() {
|
||||
slog.Info("initializing semantic search",
|
||||
"provider", searchCfg.Provider,
|
||||
)
|
||||
|
||||
// Create embedding provider
|
||||
embProvider, err := embedding.NewProvider(searchCfg.Provider, searchCfg.APIKey, searchCfg.OllamaURL)
|
||||
if err != nil {
|
||||
slog.Warn("failed to create embedding provider, semantic search disabled",
|
||||
"error", err,
|
||||
)
|
||||
} else {
|
||||
// Create vector index
|
||||
vectorIndex, err = search.NewVectorIndex(dataDir)
|
||||
if err != nil {
|
||||
slog.Warn("failed to create vector index, semantic search disabled",
|
||||
"error", err,
|
||||
)
|
||||
} else {
|
||||
embStore := search.NewEmbeddingStore(db.DB)
|
||||
|
||||
// Check for provider change
|
||||
existingProvider, _ := embStore.GetEmbeddingProvider(ctx)
|
||||
if existingProvider != "" && existingProvider != embProvider.Name() {
|
||||
slog.Info("embedding provider changed, re-indexing",
|
||||
"old_provider", existingProvider,
|
||||
"new_provider", embProvider.Name(),
|
||||
)
|
||||
_ = embStore.DeleteAllEmbeddings(ctx)
|
||||
_ = embStore.ClearQueue(ctx)
|
||||
_ = vectorIndex.Rebuild(nil)
|
||||
}
|
||||
|
||||
// Enqueue messages that need embedding
|
||||
enqueued, _ := embStore.EnqueueAllMessages(ctx)
|
||||
if enqueued > 0 {
|
||||
slog.Info("enqueued messages for embedding", "count", enqueued)
|
||||
}
|
||||
|
||||
// Create and start pipeline
|
||||
embPipeline = search.NewPipeline(embProvider, embStore, vectorIndex, searchCfg)
|
||||
embPipeline.Start(ctx)
|
||||
|
||||
// Create search service with semantic support
|
||||
searchService = search.NewService(db.DB, embProvider, vectorIndex, msgService)
|
||||
slog.Info("semantic search enabled",
|
||||
"provider", embProvider.Name(),
|
||||
"dimensions", embProvider.Dimensions(),
|
||||
"index_size", vectorIndex.Len(),
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// If no semantic search, create search service with FTS-only fallback
|
||||
if searchService == nil {
|
||||
searchService = search.NewService(db.DB, nil, nil, msgService)
|
||||
slog.Info("semantic search not configured, using full-text search only")
|
||||
}
|
||||
|
||||
// Create MCP server (with swarm + attachment + search tools)
|
||||
mcpSrv := mcpserver.NewMCPServer(msgService, agentService, channelService, swarmService, attachmentService, searchService)
|
||||
startTime := time.Now()
|
||||
|
||||
// Start task expiry worker
|
||||
@@ -308,6 +376,20 @@ func runServe(cmd *cobra.Command, args []string) error {
|
||||
// Stop expiry worker
|
||||
expiryWorker.Stop()
|
||||
|
||||
// Stop embedding pipeline
|
||||
if embPipeline != nil {
|
||||
embPipeline.Stop()
|
||||
}
|
||||
|
||||
// Save vector index
|
||||
if vectorIndex != nil {
|
||||
if err := vectorIndex.Save(); err != nil {
|
||||
slog.Error("failed to save vector index", "error", err)
|
||||
} else {
|
||||
slog.Info("vector index saved", "size", vectorIndex.Len())
|
||||
}
|
||||
}
|
||||
|
||||
// Stop retention cleaner
|
||||
if retentionCleaner != nil {
|
||||
retentionCleaner.Stop()
|
||||
|
||||
@@ -3,7 +3,9 @@ module github.com/smart-mcp-proxy/synapbus
|
||||
go 1.25.0
|
||||
|
||||
require (
|
||||
github.com/TFMV/hnsw v0.4.0
|
||||
github.com/go-chi/chi/v5 v5.2.5
|
||||
github.com/google/uuid v1.6.0
|
||||
github.com/mark3labs/mcp-go v0.45.0
|
||||
github.com/ory/fosite v0.49.0
|
||||
github.com/spf13/cobra v1.10.2
|
||||
@@ -17,6 +19,7 @@ require (
|
||||
github.com/buger/jsonparser v1.1.1 // indirect
|
||||
github.com/cenkalti/backoff/v4 v4.3.0 // indirect
|
||||
github.com/cespare/xxhash/v2 v2.3.0 // indirect
|
||||
github.com/chewxy/math32 v1.10.1 // indirect
|
||||
github.com/cristalhq/jwt/v4 v4.0.2 // indirect
|
||||
github.com/davecgh/go-spew v1.1.1 // indirect
|
||||
github.com/dgraph-io/ristretto v1.0.0 // indirect
|
||||
@@ -24,13 +27,12 @@ require (
|
||||
github.com/felixge/httpsnoop v1.0.4 // indirect
|
||||
github.com/fsnotify/fsnotify v1.6.0 // indirect
|
||||
github.com/go-jose/go-jose/v3 v3.0.3 // indirect
|
||||
github.com/go-logr/logr v1.3.0 // indirect
|
||||
github.com/go-logr/logr v1.4.2 // indirect
|
||||
github.com/go-logr/stdr v1.2.2 // indirect
|
||||
github.com/gobuffalo/pop/v6 v6.1.1 // indirect
|
||||
github.com/gogo/protobuf v1.3.2 // indirect
|
||||
github.com/golang/mock v1.6.0 // indirect
|
||||
github.com/golang/protobuf v1.5.3 // indirect
|
||||
github.com/google/uuid v1.6.0 // indirect
|
||||
github.com/google/renameio v1.0.1 // indirect
|
||||
github.com/gorilla/websocket v1.5.0 // indirect
|
||||
github.com/grpc-ecosystem/grpc-gateway/v2 v2.18.1 // indirect
|
||||
github.com/hashicorp/go-cleanhttp v0.5.2 // indirect
|
||||
@@ -60,8 +62,10 @@ require (
|
||||
github.com/spf13/jwalterweatherman v1.1.0 // indirect
|
||||
github.com/spf13/pflag v1.0.9 // indirect
|
||||
github.com/spf13/viper v1.16.0 // indirect
|
||||
github.com/stretchr/testify v1.9.0 // indirect
|
||||
github.com/stretchr/testify v1.10.0 // indirect
|
||||
github.com/subosito/gotenv v1.4.2 // indirect
|
||||
github.com/viterin/partial v1.1.0 // indirect
|
||||
github.com/viterin/vek v0.4.2 // indirect
|
||||
github.com/wk8/go-ordered-map/v2 v2.1.8 // indirect
|
||||
github.com/yosida95/uritemplate/v3 v3.0.2 // indirect
|
||||
go.opentelemetry.io/contrib/instrumentation/net/http/httptrace/otelhttptrace v0.46.1 // indirect
|
||||
@@ -69,29 +73,27 @@ require (
|
||||
go.opentelemetry.io/contrib/propagators/b3 v1.21.0 // indirect
|
||||
go.opentelemetry.io/contrib/propagators/jaeger v1.21.1 // indirect
|
||||
go.opentelemetry.io/contrib/samplers/jaegerremote v0.15.1 // indirect
|
||||
go.opentelemetry.io/otel v1.21.0 // indirect
|
||||
go.opentelemetry.io/otel v1.31.0 // indirect
|
||||
go.opentelemetry.io/otel/exporters/jaeger v1.17.0 // indirect
|
||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.21.0 // indirect
|
||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.21.0 // indirect
|
||||
go.opentelemetry.io/otel/exporters/zipkin v1.21.0 // indirect
|
||||
go.opentelemetry.io/otel/metric v1.21.0 // indirect
|
||||
go.opentelemetry.io/otel/sdk v1.21.0 // indirect
|
||||
go.opentelemetry.io/otel/trace v1.21.0 // indirect
|
||||
go.opentelemetry.io/otel/metric v1.31.0 // indirect
|
||||
go.opentelemetry.io/otel/sdk v1.31.0 // indirect
|
||||
go.opentelemetry.io/otel/trace v1.31.0 // indirect
|
||||
go.opentelemetry.io/proto/otlp v1.0.0 // indirect
|
||||
golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 // indirect
|
||||
golang.org/x/mod v0.33.0 // indirect
|
||||
golang.org/x/net v0.51.0 // indirect
|
||||
golang.org/x/oauth2 v0.14.0 // indirect
|
||||
golang.org/x/oauth2 v0.23.0 // indirect
|
||||
golang.org/x/sync v0.20.0 // indirect
|
||||
golang.org/x/sys v0.42.0 // indirect
|
||||
golang.org/x/text v0.35.0 // indirect
|
||||
golang.org/x/tools v0.42.0 // indirect
|
||||
google.golang.org/appengine v1.6.8 // indirect
|
||||
google.golang.org/genproto v0.0.0-20231106174013-bbf56f31fb17 // indirect
|
||||
google.golang.org/genproto/googleapis/api v0.0.0-20231106174013-bbf56f31fb17 // indirect
|
||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20231106174013-bbf56f31fb17 // indirect
|
||||
google.golang.org/grpc v1.59.0 // indirect
|
||||
google.golang.org/protobuf v1.33.0 // indirect
|
||||
google.golang.org/genproto/googleapis/api v0.0.0-20241015192408-796eee8c2d53 // indirect
|
||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20241104194629-dd2ea8efbc28 // indirect
|
||||
google.golang.org/grpc v1.69.2 // indirect
|
||||
google.golang.org/protobuf v1.36.1 // indirect
|
||||
gopkg.in/ini.v1 v1.67.0 // indirect
|
||||
gopkg.in/yaml.v3 v3.0.1 // indirect
|
||||
modernc.org/libc v1.67.6 // indirect
|
||||
|
||||
@@ -39,6 +39,8 @@ dmitri.shuralyov.com/gpu/mtl v0.0.0-20190408044501-666a987793e9/go.mod h1:H6x//7
|
||||
github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU=
|
||||
github.com/BurntSushi/xgb v0.0.0-20160522181843-27f122750802/go.mod h1:IVnqGOEym/WlBOVXweHU+Q+/VP0lqqI8lqeDx9IjBqo=
|
||||
github.com/Masterminds/semver/v3 v3.1.1/go.mod h1:VPu/7SZ7ePZ3QOrcuXROw5FAcLl4a0cBrbBpGY/8hQs=
|
||||
github.com/TFMV/hnsw v0.4.0 h1:k61xD3V9LzzwUMDLaHCn+1PbvMbJj33KRdUPiUtuj7k=
|
||||
github.com/TFMV/hnsw v0.4.0/go.mod h1:YPCKBOTpl3KzZxYBTVbR+uH7US5HpprYkDLALt/bgTY=
|
||||
github.com/asaskevich/govalidator v0.0.0-20230301143203-a9d515a09cc2 h1:DklsrG3dyBCFEj5IhUbnKptjxatkF07cF2ak3yi77so=
|
||||
github.com/asaskevich/govalidator v0.0.0-20230301143203-a9d515a09cc2/go.mod h1:WaHUgvxTVq04UNunO+XhnAqY/wQc+bxr74GqbsZ/Jqw=
|
||||
github.com/aymerick/douceur v0.2.0/go.mod h1:wlT5vV2O3h55X9m7iVYN0TBM0NH/MmbLnd30/FjWUq4=
|
||||
@@ -51,6 +53,8 @@ github.com/cenkalti/backoff/v4 v4.3.0/go.mod h1:Y3VNntkOUPxTVeUxJ/G5vcM//AlwfmyY
|
||||
github.com/census-instrumentation/opencensus-proto v0.2.1/go.mod h1:f6KPmirojxKA12rnyqOA5BBL4O983OfeGPqjHWSTneU=
|
||||
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
|
||||
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
||||
github.com/chewxy/math32 v1.10.1 h1:LFpeY0SLJXeaiej/eIp2L40VYfscTvKh/FSEZ68uMkU=
|
||||
github.com/chewxy/math32 v1.10.1/go.mod h1:dOB2rcuFrCn6UHrze36WSLVPKtzPMRAQvBvUwkSsLqs=
|
||||
github.com/chzyer/logex v1.1.10/go.mod h1:+Ywpsq7O8HXn0nuIou7OrIPyXbp3wmkHB+jjWRnGsAI=
|
||||
github.com/chzyer/readline v0.0.0-20180603132655-2972be24d48e/go.mod h1:nSuG5e5PlCu98SY8svDHJxuZscDgtXS6KTTbou5AhLI=
|
||||
github.com/chzyer/test v0.0.0-20180213035817-a1ea475d72b1/go.mod h1:Q3SI9o4m/ZMnBNeIyt5eFwwo7qiLfzFZmjNmxjkiQlU=
|
||||
@@ -101,8 +105,8 @@ github.com/go-jose/go-jose/v3 v3.0.3/go.mod h1:5b+7YgP7ZICgJDBdfjZaIt+H/9L9T/YQr
|
||||
github.com/go-kit/log v0.1.0/go.mod h1:zbhenjAZHb184qTLMA9ZjW7ThYL0H2mk7Q6pNt4vbaY=
|
||||
github.com/go-logfmt/logfmt v0.5.0/go.mod h1:wCYkCAKZfumFQihp8CzCvQ3paCTfi41vtzG1KdI/P7A=
|
||||
github.com/go-logr/logr v1.2.2/go.mod h1:jdQByPbusPIv2/zmleS9BjJVeZ6kBagPoEUsqbVz/1A=
|
||||
github.com/go-logr/logr v1.3.0 h1:2y3SDp0ZXuc6/cjLSZ+Q3ir+QB9T/iG5yYRXqsagWSY=
|
||||
github.com/go-logr/logr v1.3.0/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY=
|
||||
github.com/go-logr/logr v1.4.2 h1:6pFjapn8bFcIbiKo3XT4j/BhANplGihG6tvd+8rYgrY=
|
||||
github.com/go-logr/logr v1.4.2/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY=
|
||||
github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag=
|
||||
github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE=
|
||||
github.com/go-sql-driver/mysql v1.6.0/go.mod h1:DCzpHaOWr8IXmIStZouvnhqoel9Qv2LBy8hT2VhHyBg=
|
||||
@@ -157,10 +161,8 @@ github.com/golang/protobuf v1.4.0/go.mod h1:jodUvKwWbYaEsadDk5Fwe5c77LiNKVO9IDvq
|
||||
github.com/golang/protobuf v1.4.1/go.mod h1:U8fpvMrcmy5pZrNK1lt4xCsGvpyWQ/VVv6QDs8UjoX8=
|
||||
github.com/golang/protobuf v1.4.2/go.mod h1:oDoupMAO8OvCJWAcko0GGGIgR6R6ocIYbsSw735rRwI=
|
||||
github.com/golang/protobuf v1.4.3/go.mod h1:oDoupMAO8OvCJWAcko0GGGIgR6R6ocIYbsSw735rRwI=
|
||||
github.com/golang/protobuf v1.5.0/go.mod h1:FsONVRAS9T7sI+LIUmWTfcYkHO4aIWwzhcaSAoJOfIk=
|
||||
github.com/golang/protobuf v1.5.2/go.mod h1:XVQd3VNwM+JqD3oG2Ue2ip4fOMUkwXdXDdiuN0vRsmY=
|
||||
github.com/golang/protobuf v1.5.3 h1:KhyjKVUg7Usr/dYsdSqoFveMYd5ko72D+zANwlG1mmg=
|
||||
github.com/golang/protobuf v1.5.3/go.mod h1:XVQd3VNwM+JqD3oG2Ue2ip4fOMUkwXdXDdiuN0vRsmY=
|
||||
github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek=
|
||||
github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps=
|
||||
github.com/google/btree v0.0.0-20180813153112-4030bb1f1f0c/go.mod h1:lNA+9X1NB3Zf8V7Ke586lFgjr2dZNuvo3lPJSGZ5JPQ=
|
||||
github.com/google/btree v1.0.0/go.mod h1:lNA+9X1NB3Zf8V7Ke586lFgjr2dZNuvo3lPJSGZ5JPQ=
|
||||
github.com/google/go-cmp v0.2.0/go.mod h1:oXzfMopK8JAjlY9xF4vHSVASa0yLyX7SntLO5aqRK0M=
|
||||
@@ -172,7 +174,6 @@ github.com/google/go-cmp v0.5.0/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/
|
||||
github.com/google/go-cmp v0.5.1/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
|
||||
github.com/google/go-cmp v0.5.2/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
|
||||
github.com/google/go-cmp v0.5.4/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
|
||||
github.com/google/go-cmp v0.5.5/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
|
||||
github.com/google/go-cmp v0.5.9/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
|
||||
github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI=
|
||||
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
|
||||
@@ -192,6 +193,8 @@ github.com/google/pprof v0.0.0-20201218002935-b9804c9f04c2/go.mod h1:kpwsk12EmLe
|
||||
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/renameio v0.1.0/go.mod h1:KWCgfxg9yswjAJkECMjeO8J8rahYeXnNhOm40UhjYkI=
|
||||
github.com/google/renameio v1.0.1 h1:Lh/jXZmvZxb0BBeSY5VKEfidcbcbenKjZFzM/q0fSeU=
|
||||
github.com/google/renameio v1.0.1/go.mod h1:t/HQoYBZSsWSNK35C6CO/TpPLDVWvxOHboWUAweKUpk=
|
||||
github.com/google/uuid v1.1.2/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
|
||||
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=
|
||||
@@ -410,8 +413,8 @@ github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/
|
||||
github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU=
|
||||
github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4=
|
||||
github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo=
|
||||
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/stretchr/testify v1.10.0 h1:Xv5erBjTwe/5IxqUQTdXv5kgmIvbHo3QQyRwhJsOfJA=
|
||||
github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
|
||||
github.com/subosito/gotenv v1.4.2 h1:X1TuBLAMDFbaTAChgCBLu3DU3UPyELpnF2jjJ2cz/S8=
|
||||
github.com/subosito/gotenv v1.4.2/go.mod h1:ayKnFf/c6rvx/2iiLrJUk1e6plDbT3edrFNGqEflhK0=
|
||||
github.com/tidwall/gjson v1.14.3 h1:9jvXn7olKEHU1S9vwoMGliaT8jq1vJ7IH/n9zD9Dnlw=
|
||||
@@ -424,6 +427,10 @@ github.com/tidwall/sjson v1.2.5 h1:kLy8mja+1c9jlljvWTlSazM7cKDRfJuR/bOJhcY5NcY=
|
||||
github.com/tidwall/sjson v1.2.5/go.mod h1:Fvgq9kS/6ociJEDnK0Fk1cpYF4FIW6ZF7LAe+6jwd28=
|
||||
github.com/urfave/negroni v1.0.0 h1:kIimOitoypq34K7TG7DUaJ9kq/N4Ofuwi1sjz0KipXc=
|
||||
github.com/urfave/negroni v1.0.0/go.mod h1:Meg73S6kFm/4PpbYdq35yYWoCZ9mS/YSx+lKnmiohz4=
|
||||
github.com/viterin/partial v1.1.0 h1:iH1l1xqBlapXsYzADS1dcbizg3iQUKTU1rbwkHv/80E=
|
||||
github.com/viterin/partial v1.1.0/go.mod h1:oKGAo7/wylWkJTLrWX8n+f4aDPtQMQ6VG4dd2qur5QA=
|
||||
github.com/viterin/vek v0.4.2 h1:Vyv04UjQT6gcjEFX82AS9ocgNbAJqsHviheIBdPlv5U=
|
||||
github.com/viterin/vek v0.4.2/go.mod h1:A4JRAe8OvbhdzBL5ofzjBS0J29FyUrf95tQogvtHHUc=
|
||||
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=
|
||||
@@ -451,8 +458,8 @@ go.opentelemetry.io/contrib/propagators/jaeger v1.21.1 h1:f4beMGDKiVzg9IcX7/VuWV
|
||||
go.opentelemetry.io/contrib/propagators/jaeger v1.21.1/go.mod h1:U9jhkEl8d1LL+QXY7q3kneJWJugiN3kZJV2OWz3hkBY=
|
||||
go.opentelemetry.io/contrib/samplers/jaegerremote v0.15.1 h1:Qb+5A+JbIjXwO7l4HkRUhgIn4Bzz0GNS2q+qdmSx+0c=
|
||||
go.opentelemetry.io/contrib/samplers/jaegerremote v0.15.1/go.mod h1:G4vNCm7fRk0kjZ6pGNLo5SpLxAUvOfSrcaegnT8TPck=
|
||||
go.opentelemetry.io/otel v1.21.0 h1:hzLeKBZEL7Okw2mGzZ0cc4k/A7Fta0uoPgaJCr8fsFc=
|
||||
go.opentelemetry.io/otel v1.21.0/go.mod h1:QZzNPQPm1zLX4gZK4cMi+71eaorMSGT3A4znnUvNNEo=
|
||||
go.opentelemetry.io/otel v1.31.0 h1:NsJcKPIW0D0H3NgzPDHmo0WW6SptzPdqg/L1zsIm2hY=
|
||||
go.opentelemetry.io/otel v1.31.0/go.mod h1:O0C14Yl9FgkjqcCZAsE053C13OaddMYr/hz6clDkEJE=
|
||||
go.opentelemetry.io/otel/exporters/jaeger v1.17.0 h1:D7UpUy2Xc2wsi1Ras6V40q806WM07rqoCWzXu7Sqy+4=
|
||||
go.opentelemetry.io/otel/exporters/jaeger v1.17.0/go.mod h1:nPCqOnEH9rNLKqH/+rrUjiMzHJdV1BlpKcTwRTyKkKI=
|
||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.21.0 h1:cl5P5/GIfFh4t6xyruOgJP5QiA1pw4fYYdv6nc6CBWw=
|
||||
@@ -461,12 +468,14 @@ go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.21.0 h1:digkE
|
||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.21.0/go.mod h1:/OpE/y70qVkndM0TrxT4KBoN3RsFZP0QaofcfYrj76I=
|
||||
go.opentelemetry.io/otel/exporters/zipkin v1.21.0 h1:D+Gv6lSfrFBWmQYyxKjDd0Zuld9SRXpIrEsKZvE4DO4=
|
||||
go.opentelemetry.io/otel/exporters/zipkin v1.21.0/go.mod h1:83oMKR6DzmHisFOW3I+yIMGZUTjxiWaiBI8M8+TU5zE=
|
||||
go.opentelemetry.io/otel/metric v1.21.0 h1:tlYWfeo+Bocx5kLEloTjbcDwBuELRrIFxwdQ36PlJu4=
|
||||
go.opentelemetry.io/otel/metric v1.21.0/go.mod h1:o1p3CA8nNHW8j5yuQLdc1eeqEaPfzug24uvsyIEJRWM=
|
||||
go.opentelemetry.io/otel/sdk v1.21.0 h1:FTt8qirL1EysG6sTQRZ5TokkU8d0ugCj8htOgThZXQ8=
|
||||
go.opentelemetry.io/otel/sdk v1.21.0/go.mod h1:Nna6Yv7PWTdgJHVRD9hIYywQBRx7pbox6nwBnZIxl/E=
|
||||
go.opentelemetry.io/otel/trace v1.21.0 h1:WD9i5gzvoUPuXIXH24ZNBudiarZDKuekPqi/E8fpfLc=
|
||||
go.opentelemetry.io/otel/trace v1.21.0/go.mod h1:LGbsEB0f9LGjN+OZaQQ26sohbOmiMR+BaslueVtS/qQ=
|
||||
go.opentelemetry.io/otel/metric v1.31.0 h1:FSErL0ATQAmYHUIzSezZibnyVlft1ybhy4ozRPcF2fE=
|
||||
go.opentelemetry.io/otel/metric v1.31.0/go.mod h1:C3dEloVbLuYoX41KpmAhOqNriGbA+qqH6PQ5E5mUfnY=
|
||||
go.opentelemetry.io/otel/sdk v1.31.0 h1:xLY3abVHYZ5HSfOg3l2E5LUj2Cwva5Y7yGxnSW9H5Gk=
|
||||
go.opentelemetry.io/otel/sdk v1.31.0/go.mod h1:TfRbMdhvxIIr/B2N2LQW2S5v9m3gOQ/08KsbbO5BPT0=
|
||||
go.opentelemetry.io/otel/sdk/metric v1.31.0 h1:i9hxxLJF/9kkvfHppyLL55aW7iIJz4JjxTeYusH7zMc=
|
||||
go.opentelemetry.io/otel/sdk/metric v1.31.0/go.mod h1:CRInTMVvNhUKgSAMbKyTMxqOBC0zgyxzW55lZzX43Y8=
|
||||
go.opentelemetry.io/otel/trace v1.31.0 h1:ffjsj1aRouKewfr85U2aGagJ46+MvodynlQ1HYdmJys=
|
||||
go.opentelemetry.io/otel/trace v1.31.0/go.mod h1:TXZkRk7SM2ZQLtR6eoAWQFIHPvzQ06FJAsO1tJg480A=
|
||||
go.opentelemetry.io/proto/otlp v1.0.0 h1:T0TX0tmXU8a3CbNXzEKGeU5mIVOdf0oykP+u2lIVU/I=
|
||||
go.opentelemetry.io/proto/otlp v1.0.0/go.mod h1:Sy6pihPLfYHkr3NkUbEhGHFhINUSI/v80hjKIs5JXpM=
|
||||
go.uber.org/atomic v1.3.2/go.mod h1:gD2HeocX3+yG+ygLZcrzQJaqmWj9AIm7n08wl/qW/PE=
|
||||
@@ -589,8 +598,8 @@ golang.org/x/oauth2 v0.0.0-20200902213428-5d25da1a8d43/go.mod h1:KelEdhl1UZF7XfJ
|
||||
golang.org/x/oauth2 v0.0.0-20201109201403-9fd604954f58/go.mod h1:KelEdhl1UZF7XfJ4dDtk6s++YSgaE7mD/BuKKDLBl4A=
|
||||
golang.org/x/oauth2 v0.0.0-20201208152858-08078c50e5b5/go.mod h1:KelEdhl1UZF7XfJ4dDtk6s++YSgaE7mD/BuKKDLBl4A=
|
||||
golang.org/x/oauth2 v0.0.0-20210218202405-ba52d332ba99/go.mod h1:KelEdhl1UZF7XfJ4dDtk6s++YSgaE7mD/BuKKDLBl4A=
|
||||
golang.org/x/oauth2 v0.14.0 h1:P0Vrf/2538nmC0H+pEQ3MNFRRnVR7RlqyVw+bvm26z0=
|
||||
golang.org/x/oauth2 v0.14.0/go.mod h1:lAtNWgaWfL4cm7j2OV8TxGi9Qb7ECORx8DktCY74OwM=
|
||||
golang.org/x/oauth2 v0.23.0 h1:PbgcYx2W7i4LvjJWEbf0ngHV6qJYr86PkAV3bXdLEbs=
|
||||
golang.org/x/oauth2 v0.23.0/go.mod h1:XYTD2NtWslqkgxebSiOHnXEap4TF09sJSc7H1sXbhtI=
|
||||
golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
@@ -682,7 +691,6 @@ golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||
golang.org/x/text v0.3.4/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||
golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
|
||||
golang.org/x/text v0.3.8/go.mod h1:E6s5w1FMmriuDzIBO73fBruAKo1PCIq6d2Q6DHfQ8WQ=
|
||||
golang.org/x/text v0.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8=
|
||||
golang.org/x/text v0.9.0/go.mod h1:e1OnstbJyHTd6l/uOt8jFFHp6TRDWZR/bV3emEE/zU8=
|
||||
golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
|
||||
@@ -783,8 +791,6 @@ google.golang.org/appengine v1.6.1/go.mod h1:i06prIuMbXzDqacNJfV5OdTW448YApPu5ww
|
||||
google.golang.org/appengine v1.6.5/go.mod h1:8WjMMxjGQR8xUklV/ARdw2HLXBOI7O7uCIDZVag1xfc=
|
||||
google.golang.org/appengine v1.6.6/go.mod h1:8WjMMxjGQR8xUklV/ARdw2HLXBOI7O7uCIDZVag1xfc=
|
||||
google.golang.org/appengine v1.6.7/go.mod h1:8WjMMxjGQR8xUklV/ARdw2HLXBOI7O7uCIDZVag1xfc=
|
||||
google.golang.org/appengine v1.6.8 h1:IhEN5q69dyKagZPYMSdIjS2HqprW324FRQZJcGqPAsM=
|
||||
google.golang.org/appengine v1.6.8/go.mod h1:1jJ3jBArFh5pcgW8gCtRJnepW8FzD1V44FJffLiz/Ds=
|
||||
google.golang.org/genproto v0.0.0-20180817151627-c66870c02cf8/go.mod h1:JiN7NxoALGmiZfu7CAH4rXhgtRTLTxftemlI0sWmxmc=
|
||||
google.golang.org/genproto v0.0.0-20190307195333-5fe7a883aa19/go.mod h1:VzzqZJRnGkLBvHegQrXjBqPurQTc5/KpmUdxsrq26oE=
|
||||
google.golang.org/genproto v0.0.0-20190418145605-e7d98fc518a7/go.mod h1:VzzqZJRnGkLBvHegQrXjBqPurQTc5/KpmUdxsrq26oE=
|
||||
@@ -821,12 +827,10 @@ google.golang.org/genproto v0.0.0-20201210142538-e3217bee35cc/go.mod h1:FWY/as6D
|
||||
google.golang.org/genproto v0.0.0-20201214200347-8c77b98c765d/go.mod h1:FWY/as6DDZQgahTzZj3fqbO1CbirC29ZNUFHwi0/+no=
|
||||
google.golang.org/genproto v0.0.0-20210108203827-ffc7fda8c3d7/go.mod h1:FWY/as6DDZQgahTzZj3fqbO1CbirC29ZNUFHwi0/+no=
|
||||
google.golang.org/genproto v0.0.0-20210226172003-ab064af71705/go.mod h1:FWY/as6DDZQgahTzZj3fqbO1CbirC29ZNUFHwi0/+no=
|
||||
google.golang.org/genproto v0.0.0-20231106174013-bbf56f31fb17 h1:wpZ8pe2x1Q3f2KyT5f8oP/fa9rHAKgFPr/HZdNuS+PQ=
|
||||
google.golang.org/genproto v0.0.0-20231106174013-bbf56f31fb17/go.mod h1:J7XzRzVy1+IPwWHZUzoD0IccYZIrXILAQpc+Qy9CMhY=
|
||||
google.golang.org/genproto/googleapis/api v0.0.0-20231106174013-bbf56f31fb17 h1:JpwMPBpFN3uKhdaekDpiNlImDdkUAyiJ6ez/uxGaUSo=
|
||||
google.golang.org/genproto/googleapis/api v0.0.0-20231106174013-bbf56f31fb17/go.mod h1:0xJLfVdJqpAPl8tDg1ujOCGzx6LFLttXT5NhllGOXY4=
|
||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20231106174013-bbf56f31fb17 h1:Jyp0Hsi0bmHXG6k9eATXoYtjd6e2UzZ1SCn/wIupY14=
|
||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20231106174013-bbf56f31fb17/go.mod h1:oQ5rr10WTTMvP4A36n8JpR1OrO1BEiV4f78CneXZxkA=
|
||||
google.golang.org/genproto/googleapis/api v0.0.0-20241015192408-796eee8c2d53 h1:fVoAXEKA4+yufmbdVYv+SE73+cPZbbbe8paLsHfkK+U=
|
||||
google.golang.org/genproto/googleapis/api v0.0.0-20241015192408-796eee8c2d53/go.mod h1:riSXTwQ4+nqmPGtobMFyW5FqVAmIs0St6VPp4Ug7CE4=
|
||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20241104194629-dd2ea8efbc28 h1:XVhgTWWV3kGQlwJHR3upFWZeTsei6Oks1apkZSeonIE=
|
||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20241104194629-dd2ea8efbc28/go.mod h1:GX3210XPVPUjJbTUbvwI8f2IpZDMZuPJWDzDuebbviI=
|
||||
google.golang.org/grpc v1.19.0/go.mod h1:mqu4LbDTu4XGKhr4mRzUsmM4RtVoemTSY81AxZiDr8c=
|
||||
google.golang.org/grpc v1.20.1/go.mod h1:10oTOabMzJvdu6/UiuZezV6QK5dSlG84ov/aaiqXj38=
|
||||
google.golang.org/grpc v1.21.1/go.mod h1:oYelfM1adQP15Ek0mdvEgi9Df8B9CZIaU1084ijfRaM=
|
||||
@@ -843,8 +847,8 @@ google.golang.org/grpc v1.31.1/go.mod h1:N36X2cJ7JwdamYAgDz+s+rVMFjt3numwzf/HckM
|
||||
google.golang.org/grpc v1.33.2/go.mod h1:JMHMWHQWaTccqQQlmk3MJZS+GWXOdAesneDmEnv2fbc=
|
||||
google.golang.org/grpc v1.34.0/go.mod h1:WotjhfgOW/POjDeRt8vscBtXq+2VjORFy659qA51WJ8=
|
||||
google.golang.org/grpc v1.35.0/go.mod h1:qjiiYl8FncCW8feJPdyg3v6XW24KsRHe+dy9BAGRRjU=
|
||||
google.golang.org/grpc v1.59.0 h1:Z5Iec2pjwb+LEOqzpB2MR12/eKFhDPhuqW91O+4bwUk=
|
||||
google.golang.org/grpc v1.59.0/go.mod h1:aUPDwccQo6OTjy7Hct4AfBPD1GptF4fyUjIkQ9YtF98=
|
||||
google.golang.org/grpc v1.69.2 h1:U3S9QEtbXC0bYNvRtcoklF3xGtLViumSYxWykJS+7AU=
|
||||
google.golang.org/grpc v1.69.2/go.mod h1:vyjdE6jLBI76dgpDojsFGNaHlxdjXN9ghpnd2o7JGZ4=
|
||||
google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8=
|
||||
google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0=
|
||||
google.golang.org/protobuf v0.0.0-20200228230310-ab0ca4ff8a60/go.mod h1:cfTl7dwQJ+fmap5saPgwCLgHXTUD7jkjRqWcaiX5VyM=
|
||||
@@ -855,10 +859,8 @@ google.golang.org/protobuf v1.23.0/go.mod h1:EGpADcykh3NcUnDUJcl1+ZksZNG86OlYog2
|
||||
google.golang.org/protobuf v1.23.1-0.20200526195155-81db48ad09cc/go.mod h1:EGpADcykh3NcUnDUJcl1+ZksZNG86OlYog2l/sGQquU=
|
||||
google.golang.org/protobuf v1.24.0/go.mod h1:r/3tXBNzIEhYS9I1OUVjXDlt8tc493IdKGjtUeSXeh4=
|
||||
google.golang.org/protobuf v1.25.0/go.mod h1:9JNX74DMeImyA3h4bdi1ymwjUzf21/xIlbajtzgsN7c=
|
||||
google.golang.org/protobuf v1.26.0-rc.1/go.mod h1:jlhhOSvTdKEhbULTjvd4ARK9grFBp09yW+WbY/TyQbw=
|
||||
google.golang.org/protobuf v1.26.0/go.mod h1:9q0QmTI4eRPtz6boOQmLYwt+qCgq0jsYwAQnmE0givc=
|
||||
google.golang.org/protobuf v1.33.0 h1:uNO2rsAINq/JlFpSdYEKIZ0uKD/R9cpdv0T+yoGwGmI=
|
||||
google.golang.org/protobuf v1.33.0/go.mod h1:c6P6GXX6sHbq/GpV6MGZEdwhWPcYBgnhAHhKbcUYpos=
|
||||
google.golang.org/protobuf v1.36.1 h1:yBPeRvTftaleIgM3PZ/WBIZ7XM/eEYAaEyCwvyjq/gk=
|
||||
google.golang.org/protobuf v1.36.1/go.mod h1:9fA7Ob0pmnwhb644+1+CVWFRbNajQ6iRojtC/QF5bRE=
|
||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||
gopkg.in/check.v1 v1.0.0-20180628173108-788fd7840127/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||
gopkg.in/check.v1 v1.0.0-20190902080502-41f04d3bba15/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/attachments"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/channels"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/messaging"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/search"
|
||||
)
|
||||
|
||||
// MCPServer wraps the mcp-go server with SynapBus services.
|
||||
@@ -29,6 +30,7 @@ func NewMCPServer(
|
||||
channelService *channels.Service,
|
||||
swarmService *channels.SwarmService,
|
||||
attachmentService *attachments.Service,
|
||||
searchService *search.Service,
|
||||
) *MCPServer {
|
||||
logger := slog.Default().With("component", "mcp-server")
|
||||
|
||||
@@ -41,6 +43,9 @@ func NewMCPServer(
|
||||
|
||||
// Register all tools
|
||||
registrar := NewToolRegistrar(msgService, agentService)
|
||||
if searchService != nil {
|
||||
registrar.SetSearchService(searchService)
|
||||
}
|
||||
registrar.RegisterAll(mcpSrv)
|
||||
|
||||
// Register channel tools
|
||||
|
||||
+75
-10
@@ -11,13 +11,15 @@ import (
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/agents"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/messaging"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/search"
|
||||
)
|
||||
|
||||
// ToolRegistrar registers all SynapBus MCP tools on the given server.
|
||||
type ToolRegistrar struct {
|
||||
msgService *messaging.MessagingService
|
||||
agentService *agents.AgentService
|
||||
logger *slog.Logger
|
||||
msgService *messaging.MessagingService
|
||||
agentService *agents.AgentService
|
||||
searchService *search.Service
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewToolRegistrar creates a new tool registrar.
|
||||
@@ -29,6 +31,11 @@ func NewToolRegistrar(msgService *messaging.MessagingService, agentService *agen
|
||||
}
|
||||
}
|
||||
|
||||
// SetSearchService sets the search service for semantic search support.
|
||||
func (tr *ToolRegistrar) SetSearchService(svc *search.Service) {
|
||||
tr.searchService = svc
|
||||
}
|
||||
|
||||
// RegisterAll registers all tools on the MCP server.
|
||||
func (tr *ToolRegistrar) RegisterAll(s *server.MCPServer) {
|
||||
s.AddTool(tr.sendMessageTool(), tr.handleSendMessage)
|
||||
@@ -87,12 +94,14 @@ func (tr *ToolRegistrar) markDoneTool() mcp.Tool {
|
||||
|
||||
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.WithDescription("Search messages using semantic search (if configured) or full-text search. Returns messages ranked by relevance."),
|
||||
mcp.WithString("query", mcp.Description("Search query string — supports natural language for semantic search")),
|
||||
mcp.WithNumber("limit", mcp.Description("Maximum results to return (default 10, max 100)")),
|
||||
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")),
|
||||
mcp.WithString("search_mode", mcp.Description("Search mode: 'auto' (default), 'semantic', or 'fulltext'")),
|
||||
mcp.WithBoolean("semantic", mcp.Description("Force semantic search (shorthand for search_mode='semantic')")),
|
||||
)
|
||||
}
|
||||
|
||||
@@ -264,21 +273,77 @@ func (tr *ToolRegistrar) handleSearchMessages(ctx context.Context, req mcp.CallT
|
||||
}
|
||||
|
||||
query := req.GetString("query", "")
|
||||
opts := messaging.SearchOptions{
|
||||
|
||||
// If search service is available, use it for unified search
|
||||
if tr.searchService != nil {
|
||||
searchMode := req.GetString("search_mode", "auto")
|
||||
|
||||
// Handle boolean "semantic" shorthand
|
||||
args := req.GetArguments()
|
||||
if v, ok := args["semantic"]; ok {
|
||||
if b, ok := v.(bool); ok && b {
|
||||
searchMode = "semantic"
|
||||
}
|
||||
}
|
||||
|
||||
opts := search.SearchOptions{
|
||||
Query: query,
|
||||
Mode: searchMode,
|
||||
Limit: req.GetInt("limit", 10),
|
||||
FromAgent: req.GetString("from_agent", ""),
|
||||
MinPriority: req.GetInt("min_priority", 0),
|
||||
}
|
||||
|
||||
resp, err := tr.searchService.Search(ctx, agentName, opts)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("search_messages failed: %s", err)), nil
|
||||
}
|
||||
|
||||
// Format results
|
||||
resultMsgs := make([]map[string]any, len(resp.Results))
|
||||
for i, r := range resp.Results {
|
||||
entry := map[string]any{
|
||||
"message": r.Message,
|
||||
"match_type": r.MatchType,
|
||||
}
|
||||
if r.SimilarityScore > 0 {
|
||||
entry["similarity_score"] = r.SimilarityScore
|
||||
}
|
||||
if r.RelevanceScore > 0 {
|
||||
entry["relevance_score"] = r.RelevanceScore
|
||||
}
|
||||
resultMsgs[i] = entry
|
||||
}
|
||||
|
||||
result := map[string]any{
|
||||
"results": resultMsgs,
|
||||
"count": resp.TotalResults,
|
||||
"search_mode": resp.SearchMode,
|
||||
}
|
||||
if resp.Warning != "" {
|
||||
result["warning"] = resp.Warning
|
||||
}
|
||||
|
||||
return resultJSON(result)
|
||||
}
|
||||
|
||||
// Fallback: use messaging service directly (no search service configured)
|
||||
msgOpts := 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)
|
||||
messages, err := tr.msgService.SearchMessages(ctx, agentName, query, msgOpts)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("search_messages failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"messages": messages,
|
||||
"count": len(messages),
|
||||
"messages": messages,
|
||||
"count": len(messages),
|
||||
"search_mode": "fulltext",
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
// Package search provides semantic and full-text search for SynapBus messages.
|
||||
package search
|
||||
|
||||
import (
|
||||
"os"
|
||||
"strconv"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Config holds configuration for the search subsystem.
|
||||
type Config struct {
|
||||
// Provider specifies the embedding provider: "openai", "ollama", or empty for none.
|
||||
Provider string
|
||||
// APIKey is the API key for the embedding provider (required for openai).
|
||||
APIKey string
|
||||
// OllamaURL is the Ollama server URL (default http://localhost:11434).
|
||||
OllamaURL string
|
||||
// BatchSize is the number of messages to embed in a single batch (default 10).
|
||||
BatchSize int
|
||||
// WorkerCount is the number of embedding worker goroutines (default 1).
|
||||
WorkerCount int
|
||||
// PollInterval is how often to poll for unembedded messages (default 2s).
|
||||
PollInterval time.Duration
|
||||
// RetryMaxAttempts is the max retry count for failed embeddings (default 3).
|
||||
RetryMaxAttempts int
|
||||
// RetryBaseDelay is the base delay for exponential backoff (default 1s).
|
||||
RetryBaseDelay time.Duration
|
||||
}
|
||||
|
||||
// LoadConfigFromEnv creates a Config from environment variables.
|
||||
func LoadConfigFromEnv() Config {
|
||||
cfg := Config{
|
||||
Provider: os.Getenv("SYNAPBUS_EMBEDDING_PROVIDER"),
|
||||
APIKey: os.Getenv("SYNAPBUS_EMBEDDING_API_KEY"),
|
||||
OllamaURL: os.Getenv("SYNAPBUS_OLLAMA_URL"),
|
||||
BatchSize: 10,
|
||||
WorkerCount: 1,
|
||||
PollInterval: 2 * time.Second,
|
||||
RetryMaxAttempts: 3,
|
||||
RetryBaseDelay: 1 * time.Second,
|
||||
}
|
||||
|
||||
if cfg.OllamaURL == "" {
|
||||
cfg.OllamaURL = "http://localhost:11434"
|
||||
}
|
||||
|
||||
if v := os.Getenv("SYNAPBUS_EMBEDDING_BATCH_SIZE"); v != "" {
|
||||
if n, err := strconv.Atoi(v); err == nil && n > 0 {
|
||||
cfg.BatchSize = n
|
||||
}
|
||||
}
|
||||
|
||||
if v := os.Getenv("SYNAPBUS_EMBEDDING_WORKERS"); v != "" {
|
||||
if n, err := strconv.Atoi(v); err == nil && n > 0 {
|
||||
cfg.WorkerCount = n
|
||||
}
|
||||
}
|
||||
|
||||
return cfg
|
||||
}
|
||||
|
||||
// IsEnabled returns true if an embedding provider is configured.
|
||||
func (c Config) IsEnabled() bool {
|
||||
return c.Provider != ""
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
package embedding
|
||||
|
||||
import "fmt"
|
||||
|
||||
// NewProvider creates an EmbeddingProvider based on the provider name.
|
||||
func NewProvider(provider, apiKey, ollamaURL string) (EmbeddingProvider, error) {
|
||||
switch provider {
|
||||
case "openai":
|
||||
return NewOpenAIProvider(apiKey)
|
||||
case "ollama":
|
||||
return NewOllamaProvider(ollamaURL)
|
||||
case "":
|
||||
return nil, fmt.Errorf("no embedding provider specified")
|
||||
default:
|
||||
return nil, fmt.Errorf("unknown embedding provider: %q (supported: openai, ollama)", provider)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
package embedding
|
||||
|
||||
import "context"
|
||||
|
||||
// MockProvider implements EmbeddingProvider for testing.
|
||||
type MockProvider struct {
|
||||
dims int
|
||||
name string
|
||||
embedFunc func(ctx context.Context, text string) ([]float32, error)
|
||||
batchFunc func(ctx context.Context, texts []string) ([][]float32, error)
|
||||
}
|
||||
|
||||
// NewMockProvider creates a mock provider with the given dimensionality.
|
||||
func NewMockProvider(dims int) *MockProvider {
|
||||
return &MockProvider{
|
||||
dims: dims,
|
||||
name: "mock",
|
||||
}
|
||||
}
|
||||
|
||||
// SetEmbedFunc overrides the default embedding behavior.
|
||||
func (m *MockProvider) SetEmbedFunc(fn func(ctx context.Context, text string) ([]float32, error)) {
|
||||
m.embedFunc = fn
|
||||
}
|
||||
|
||||
// SetBatchFunc overrides the default batch embedding behavior.
|
||||
func (m *MockProvider) SetBatchFunc(fn func(ctx context.Context, texts []string) ([][]float32, error)) {
|
||||
m.batchFunc = fn
|
||||
}
|
||||
|
||||
func (m *MockProvider) Embed(ctx context.Context, text string) ([]float32, error) {
|
||||
if m.embedFunc != nil {
|
||||
return m.embedFunc(ctx, text)
|
||||
}
|
||||
// Generate a deterministic vector based on text length
|
||||
vec := make([]float32, m.dims)
|
||||
for i := range vec {
|
||||
vec[i] = float32(len(text)%10+i) / float32(m.dims)
|
||||
}
|
||||
return vec, nil
|
||||
}
|
||||
|
||||
func (m *MockProvider) EmbedBatch(ctx context.Context, texts []string) ([][]float32, error) {
|
||||
if m.batchFunc != nil {
|
||||
return m.batchFunc(ctx, texts)
|
||||
}
|
||||
results := make([][]float32, len(texts))
|
||||
for i, text := range texts {
|
||||
vec, err := m.Embed(ctx, text)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
results[i] = vec
|
||||
}
|
||||
return results, nil
|
||||
}
|
||||
|
||||
func (m *MockProvider) Dimensions() int { return m.dims }
|
||||
func (m *MockProvider) Name() string { return m.name }
|
||||
@@ -0,0 +1,104 @@
|
||||
package embedding
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
ollamaDefaultModel = "nomic-embed-text"
|
||||
ollamaDimensions = 768
|
||||
)
|
||||
|
||||
// OllamaProvider implements EmbeddingProvider using a local Ollama instance.
|
||||
type OllamaProvider struct {
|
||||
endpoint string
|
||||
model string
|
||||
client *http.Client
|
||||
}
|
||||
|
||||
// NewOllamaProvider creates a new Ollama embedding provider.
|
||||
func NewOllamaProvider(endpoint string) (*OllamaProvider, error) {
|
||||
if endpoint == "" {
|
||||
endpoint = "http://localhost:11434"
|
||||
}
|
||||
return &OllamaProvider{
|
||||
endpoint: endpoint,
|
||||
model: ollamaDefaultModel,
|
||||
client: &http.Client{Timeout: 60 * time.Second},
|
||||
}, nil
|
||||
}
|
||||
|
||||
type ollamaRequest struct {
|
||||
Model string `json:"model"`
|
||||
Prompt string `json:"prompt"`
|
||||
}
|
||||
|
||||
type ollamaResponse struct {
|
||||
Embedding []float32 `json:"embedding"`
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
func (p *OllamaProvider) Embed(ctx context.Context, text string) ([]float32, error) {
|
||||
body, err := json.Marshal(ollamaRequest{
|
||||
Model: p.model,
|
||||
Prompt: text,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ollama: marshal request: %w", err)
|
||||
}
|
||||
|
||||
url := p.endpoint + "/api/embeddings"
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ollama: create request: %w", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
|
||||
resp, err := p.client.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ollama: request failed (is Ollama running at %s?): %w", p.endpoint, err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
respBody, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ollama: read response: %w", err)
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("ollama: API error %d: %s", resp.StatusCode, string(respBody))
|
||||
}
|
||||
|
||||
var result ollamaResponse
|
||||
if err := json.Unmarshal(respBody, &result); err != nil {
|
||||
return nil, fmt.Errorf("ollama: unmarshal response: %w", err)
|
||||
}
|
||||
|
||||
if result.Error != "" {
|
||||
return nil, fmt.Errorf("ollama: %s", result.Error)
|
||||
}
|
||||
|
||||
return result.Embedding, nil
|
||||
}
|
||||
|
||||
func (p *OllamaProvider) EmbedBatch(ctx context.Context, texts []string) ([][]float32, error) {
|
||||
// Ollama does not support batch embedding; call sequentially.
|
||||
results := make([][]float32, len(texts))
|
||||
for i, text := range texts {
|
||||
vec, err := p.Embed(ctx, text)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ollama: batch item %d: %w", i, err)
|
||||
}
|
||||
results[i] = vec
|
||||
}
|
||||
return results, nil
|
||||
}
|
||||
|
||||
func (p *OllamaProvider) Dimensions() int { return ollamaDimensions }
|
||||
func (p *OllamaProvider) Name() string { return "ollama" }
|
||||
@@ -0,0 +1,141 @@
|
||||
package embedding
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
openAIEmbeddingURL = "https://api.openai.com/v1/embeddings"
|
||||
openAIModel = "text-embedding-3-small"
|
||||
openAIDimensions = 1536
|
||||
openAIMaxTokens = 8191
|
||||
)
|
||||
|
||||
// OpenAIProvider implements EmbeddingProvider using OpenAI's text-embedding-3-small.
|
||||
type OpenAIProvider struct {
|
||||
apiKey string
|
||||
client *http.Client
|
||||
}
|
||||
|
||||
// NewOpenAIProvider creates a new OpenAI embedding provider.
|
||||
func NewOpenAIProvider(apiKey string) (*OpenAIProvider, error) {
|
||||
if apiKey == "" {
|
||||
return nil, fmt.Errorf("openai provider requires SYNAPBUS_EMBEDDING_API_KEY to be set")
|
||||
}
|
||||
return &OpenAIProvider{
|
||||
apiKey: apiKey,
|
||||
client: &http.Client{Timeout: 30 * time.Second},
|
||||
}, nil
|
||||
}
|
||||
|
||||
type openAIRequest struct {
|
||||
Input []string `json:"input"`
|
||||
Model string `json:"model"`
|
||||
}
|
||||
|
||||
type openAIResponse struct {
|
||||
Data []openAIEmbedding `json:"data"`
|
||||
Error *openAIError `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
type openAIEmbedding struct {
|
||||
Embedding []float32 `json:"embedding"`
|
||||
Index int `json:"index"`
|
||||
}
|
||||
|
||||
type openAIError struct {
|
||||
Message string `json:"message"`
|
||||
Type string `json:"type"`
|
||||
}
|
||||
|
||||
func (p *OpenAIProvider) Embed(ctx context.Context, text string) ([]float32, error) {
|
||||
results, err := p.EmbedBatch(ctx, []string{text})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(results) == 0 {
|
||||
return nil, fmt.Errorf("openai: empty response")
|
||||
}
|
||||
return results[0], nil
|
||||
}
|
||||
|
||||
func (p *OpenAIProvider) EmbedBatch(ctx context.Context, texts []string) ([][]float32, error) {
|
||||
if len(texts) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// Truncate texts that are too long (rough character-based limit; 1 token ~ 4 chars)
|
||||
truncated := make([]string, len(texts))
|
||||
for i, t := range texts {
|
||||
maxChars := openAIMaxTokens * 4
|
||||
if len(t) > maxChars {
|
||||
truncated[i] = t[:maxChars]
|
||||
} else {
|
||||
truncated[i] = t
|
||||
}
|
||||
}
|
||||
|
||||
body, err := json.Marshal(openAIRequest{
|
||||
Input: truncated,
|
||||
Model: openAIModel,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("openai: marshal request: %w", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, openAIEmbeddingURL, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("openai: create request: %w", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+p.apiKey)
|
||||
|
||||
resp, err := p.client.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("openai: request failed: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
respBody, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("openai: read response: %w", err)
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
switch resp.StatusCode {
|
||||
case http.StatusUnauthorized:
|
||||
return nil, fmt.Errorf("openai: invalid API key (401)")
|
||||
case http.StatusTooManyRequests:
|
||||
return nil, fmt.Errorf("openai: rate limited (429)")
|
||||
default:
|
||||
return nil, fmt.Errorf("openai: API error %d: %s", resp.StatusCode, string(respBody))
|
||||
}
|
||||
}
|
||||
|
||||
var result openAIResponse
|
||||
if err := json.Unmarshal(respBody, &result); err != nil {
|
||||
return nil, fmt.Errorf("openai: unmarshal response: %w", err)
|
||||
}
|
||||
|
||||
if result.Error != nil {
|
||||
return nil, fmt.Errorf("openai: %s: %s", result.Error.Type, result.Error.Message)
|
||||
}
|
||||
|
||||
embeddings := make([][]float32, len(texts))
|
||||
for _, d := range result.Data {
|
||||
if d.Index < len(embeddings) {
|
||||
embeddings[d.Index] = d.Embedding
|
||||
}
|
||||
}
|
||||
|
||||
return embeddings, nil
|
||||
}
|
||||
|
||||
func (p *OpenAIProvider) Dimensions() int { return openAIDimensions }
|
||||
func (p *OpenAIProvider) Name() string { return "openai" }
|
||||
@@ -0,0 +1,16 @@
|
||||
// Package embedding provides embedding provider implementations for semantic search.
|
||||
package embedding
|
||||
|
||||
import "context"
|
||||
|
||||
// EmbeddingProvider generates vector embeddings from text.
|
||||
type EmbeddingProvider interface {
|
||||
// Embed generates an embedding vector for a single text.
|
||||
Embed(ctx context.Context, text string) ([]float32, error)
|
||||
// EmbedBatch generates embedding vectors for multiple texts.
|
||||
EmbedBatch(ctx context.Context, texts []string) ([][]float32, error)
|
||||
// Dimensions returns the embedding dimensionality.
|
||||
Dimensions() int
|
||||
// Name returns the provider name (e.g. "openai", "ollama").
|
||||
Name() string
|
||||
}
|
||||
@@ -0,0 +1,110 @@
|
||||
package embedding
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestMockProvider_Embed(t *testing.T) {
|
||||
provider := NewMockProvider(128)
|
||||
|
||||
t.Run("returns correct dimensions", func(t *testing.T) {
|
||||
vec, err := provider.Embed(context.Background(), "test text")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if len(vec) != 128 {
|
||||
t.Errorf("vector length = %d, want 128", len(vec))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("batch embedding", func(t *testing.T) {
|
||||
texts := []string{"hello", "world", "test"}
|
||||
vecs, err := provider.EmbedBatch(context.Background(), texts)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if len(vecs) != 3 {
|
||||
t.Errorf("batch result length = %d, want 3", len(vecs))
|
||||
}
|
||||
for i, vec := range vecs {
|
||||
if len(vec) != 128 {
|
||||
t.Errorf("vector %d length = %d, want 128", i, len(vec))
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("empty batch", func(t *testing.T) {
|
||||
vecs, err := provider.EmbedBatch(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if len(vecs) != 0 {
|
||||
t.Errorf("empty batch result length = %d, want 0", len(vecs))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("custom embed func", func(t *testing.T) {
|
||||
p := NewMockProvider(3)
|
||||
p.SetEmbedFunc(func(ctx context.Context, text string) ([]float32, error) {
|
||||
return []float32{1.0, 2.0, 3.0}, nil
|
||||
})
|
||||
|
||||
vec, err := p.Embed(context.Background(), "anything")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if vec[0] != 1.0 || vec[1] != 2.0 || vec[2] != 3.0 {
|
||||
t.Errorf("unexpected vector: %v", vec)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestNewProvider_Factory(t *testing.T) {
|
||||
t.Run("unknown provider", func(t *testing.T) {
|
||||
_, err := NewProvider("unknown", "", "")
|
||||
if err == nil {
|
||||
t.Error("expected error for unknown provider")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("empty provider", func(t *testing.T) {
|
||||
_, err := NewProvider("", "", "")
|
||||
if err == nil {
|
||||
t.Error("expected error for empty provider")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("openai without key", func(t *testing.T) {
|
||||
_, err := NewProvider("openai", "", "")
|
||||
if err == nil {
|
||||
t.Error("expected error for openai without API key")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("openai with key", func(t *testing.T) {
|
||||
p, err := NewProvider("openai", "test-key", "")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if p.Name() != "openai" {
|
||||
t.Errorf("name = %q, want openai", p.Name())
|
||||
}
|
||||
if p.Dimensions() != 1536 {
|
||||
t.Errorf("dimensions = %d, want 1536", p.Dimensions())
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("ollama", func(t *testing.T) {
|
||||
p, err := NewProvider("ollama", "", "http://localhost:11434")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if p.Name() != "ollama" {
|
||||
t.Errorf("name = %q, want ollama", p.Name())
|
||||
}
|
||||
if p.Dimensions() != 768 {
|
||||
t.Errorf("dimensions = %d, want 768", p.Dimensions())
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,156 @@
|
||||
package search
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
|
||||
"github.com/TFMV/hnsw"
|
||||
)
|
||||
|
||||
const hnswIndexFile = "hnsw.idx"
|
||||
|
||||
// VectorIndex provides thread-safe approximate nearest neighbor search.
|
||||
type VectorIndex struct {
|
||||
mu sync.RWMutex
|
||||
graph *hnsw.SavedGraph[int64]
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// SearchResult from the vector index.
|
||||
type VectorSearchResult struct {
|
||||
ID int64
|
||||
Distance float32
|
||||
}
|
||||
|
||||
// NewVectorIndex creates or loads a vector index from dataDir.
|
||||
func NewVectorIndex(dataDir string) (*VectorIndex, error) {
|
||||
path := filepath.Join(dataDir, hnswIndexFile)
|
||||
|
||||
g, err := hnsw.LoadSavedGraph[int64](path)
|
||||
if err != nil {
|
||||
// If the file is corrupted, start fresh
|
||||
slog.Warn("failed to load HNSW index, creating new",
|
||||
"path", path,
|
||||
"error", err,
|
||||
)
|
||||
g = &hnsw.SavedGraph[int64]{
|
||||
Graph: hnsw.NewGraph[int64](),
|
||||
Path: path,
|
||||
}
|
||||
}
|
||||
|
||||
// Configure for cosine distance (default in hnsw.NewGraph)
|
||||
g.M = 16
|
||||
g.EfSearch = 100
|
||||
g.Distance = hnsw.CosineDistance
|
||||
|
||||
return &VectorIndex{
|
||||
graph: g,
|
||||
logger: slog.Default().With("component", "vector-index"),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// NewMemoryVectorIndex creates an in-memory vector index (for testing).
|
||||
func NewMemoryVectorIndex() *VectorIndex {
|
||||
g := hnsw.NewGraph[int64]()
|
||||
g.M = 16
|
||||
g.EfSearch = 100
|
||||
g.Distance = hnsw.CosineDistance
|
||||
|
||||
return &VectorIndex{
|
||||
graph: &hnsw.SavedGraph[int64]{
|
||||
Graph: g,
|
||||
Path: "",
|
||||
},
|
||||
logger: slog.Default().With("component", "vector-index"),
|
||||
}
|
||||
}
|
||||
|
||||
// AddVector adds a vector to the index.
|
||||
func (idx *VectorIndex) AddVector(id int64, vector []float32) error {
|
||||
idx.mu.Lock()
|
||||
defer idx.mu.Unlock()
|
||||
|
||||
node := hnsw.MakeNode(id, vector)
|
||||
if err := idx.graph.Add(node); err != nil {
|
||||
return fmt.Errorf("add vector %d: %w", id, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Search finds the k nearest vectors to the query.
|
||||
func (idx *VectorIndex) Search(query []float32, k int) ([]VectorSearchResult, error) {
|
||||
idx.mu.RLock()
|
||||
defer idx.mu.RUnlock()
|
||||
|
||||
if idx.graph.Len() == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
nodes, err := idx.graph.Search(query, k)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("search: %w", err)
|
||||
}
|
||||
|
||||
results := make([]VectorSearchResult, len(nodes))
|
||||
for i, n := range nodes {
|
||||
// CosineDistance returns 1 - cosine_similarity, so distance is in [0, 2]
|
||||
results[i] = VectorSearchResult{
|
||||
ID: n.Key,
|
||||
Distance: hnsw.CosineDistance(query, n.Value),
|
||||
}
|
||||
}
|
||||
return results, nil
|
||||
}
|
||||
|
||||
// Delete removes a vector from the index.
|
||||
func (idx *VectorIndex) Delete(id int64) bool {
|
||||
idx.mu.Lock()
|
||||
defer idx.mu.Unlock()
|
||||
return idx.graph.Delete(id)
|
||||
}
|
||||
|
||||
// Save persists the index to disk.
|
||||
func (idx *VectorIndex) Save() error {
|
||||
idx.mu.RLock()
|
||||
defer idx.mu.RUnlock()
|
||||
|
||||
if idx.graph.Path == "" {
|
||||
return nil // in-memory index, no save
|
||||
}
|
||||
return idx.graph.Save()
|
||||
}
|
||||
|
||||
// Len returns the number of vectors in the index.
|
||||
func (idx *VectorIndex) Len() int {
|
||||
idx.mu.RLock()
|
||||
defer idx.mu.RUnlock()
|
||||
return idx.graph.Len()
|
||||
}
|
||||
|
||||
// Rebuild clears the index and re-adds the given vectors.
|
||||
func (idx *VectorIndex) Rebuild(vectors map[int64][]float32) error {
|
||||
idx.mu.Lock()
|
||||
defer idx.mu.Unlock()
|
||||
|
||||
// Create a fresh graph
|
||||
g := hnsw.NewGraph[int64]()
|
||||
g.M = 16
|
||||
g.EfSearch = 100
|
||||
g.Distance = hnsw.CosineDistance
|
||||
|
||||
if len(vectors) > 0 {
|
||||
nodes := make([]hnsw.Node[int64], 0, len(vectors))
|
||||
for id, vec := range vectors {
|
||||
nodes = append(nodes, hnsw.MakeNode(id, vec))
|
||||
}
|
||||
if err := g.Add(nodes...); err != nil {
|
||||
return fmt.Errorf("rebuild: add nodes: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
idx.graph.Graph = g
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,179 @@
|
||||
package search
|
||||
|
||||
import (
|
||||
"os"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestVectorIndex_AddAndSearch(t *testing.T) {
|
||||
idx := NewMemoryVectorIndex()
|
||||
|
||||
// Add some vectors
|
||||
vectors := map[int64][]float32{
|
||||
1: {1.0, 0.0, 0.0},
|
||||
2: {0.0, 1.0, 0.0},
|
||||
3: {0.0, 0.0, 1.0},
|
||||
4: {0.9, 0.1, 0.0}, // close to vector 1
|
||||
}
|
||||
|
||||
for id, vec := range vectors {
|
||||
if err := idx.AddVector(id, vec); err != nil {
|
||||
t.Fatalf("AddVector(%d): %v", id, err)
|
||||
}
|
||||
}
|
||||
|
||||
if idx.Len() != 4 {
|
||||
t.Errorf("Len() = %d, want 4", idx.Len())
|
||||
}
|
||||
|
||||
// Search for vectors near [1.0, 0.0, 0.0]
|
||||
results, err := idx.Search([]float32{1.0, 0.0, 0.0}, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
|
||||
if len(results) != 2 {
|
||||
t.Fatalf("search results = %d, want 2", len(results))
|
||||
}
|
||||
|
||||
// First result should be vector 1 (exact match) or 4 (close match)
|
||||
found1 := false
|
||||
found4 := false
|
||||
for _, r := range results {
|
||||
if r.ID == 1 {
|
||||
found1 = true
|
||||
}
|
||||
if r.ID == 4 {
|
||||
found4 = true
|
||||
}
|
||||
}
|
||||
if !found1 {
|
||||
t.Error("expected to find vector 1 in top-2 results")
|
||||
}
|
||||
if !found4 {
|
||||
t.Error("expected to find vector 4 in top-2 results")
|
||||
}
|
||||
}
|
||||
|
||||
func TestVectorIndex_Delete(t *testing.T) {
|
||||
idx := NewMemoryVectorIndex()
|
||||
|
||||
if err := idx.AddVector(1, []float32{1.0, 0.0, 0.0}); err != nil {
|
||||
t.Fatalf("AddVector: %v", err)
|
||||
}
|
||||
if err := idx.AddVector(2, []float32{0.0, 1.0, 0.0}); err != nil {
|
||||
t.Fatalf("AddVector: %v", err)
|
||||
}
|
||||
|
||||
if idx.Len() != 2 {
|
||||
t.Errorf("Len() = %d, want 2", idx.Len())
|
||||
}
|
||||
|
||||
deleted := idx.Delete(1)
|
||||
if !deleted {
|
||||
t.Error("Delete(1) returned false, want true")
|
||||
}
|
||||
|
||||
if idx.Len() != 1 {
|
||||
t.Errorf("Len() after delete = %d, want 1", idx.Len())
|
||||
}
|
||||
}
|
||||
|
||||
func TestVectorIndex_EmptySearch(t *testing.T) {
|
||||
idx := NewMemoryVectorIndex()
|
||||
|
||||
results, err := idx.Search([]float32{1.0, 0.0, 0.0}, 5)
|
||||
if err != nil {
|
||||
t.Fatalf("Search on empty index: %v", err)
|
||||
}
|
||||
if results != nil {
|
||||
t.Errorf("expected nil results on empty index, got %d", len(results))
|
||||
}
|
||||
}
|
||||
|
||||
func TestVectorIndex_Rebuild(t *testing.T) {
|
||||
idx := NewMemoryVectorIndex()
|
||||
|
||||
// Add initial vectors
|
||||
if err := idx.AddVector(1, []float32{1.0, 0.0, 0.0}); err != nil {
|
||||
t.Fatalf("AddVector: %v", err)
|
||||
}
|
||||
if err := idx.AddVector(2, []float32{0.0, 1.0, 0.0}); err != nil {
|
||||
t.Fatalf("AddVector: %v", err)
|
||||
}
|
||||
|
||||
if idx.Len() != 2 {
|
||||
t.Errorf("Len() = %d, want 2", idx.Len())
|
||||
}
|
||||
|
||||
// Rebuild with new vectors
|
||||
newVectors := map[int64][]float32{
|
||||
10: {1.0, 0.0, 0.0},
|
||||
20: {0.0, 1.0, 0.0},
|
||||
30: {0.0, 0.0, 1.0},
|
||||
}
|
||||
if err := idx.Rebuild(newVectors); err != nil {
|
||||
t.Fatalf("Rebuild: %v", err)
|
||||
}
|
||||
|
||||
if idx.Len() != 3 {
|
||||
t.Errorf("Len() after rebuild = %d, want 3", idx.Len())
|
||||
}
|
||||
|
||||
// Rebuild with empty clears index
|
||||
if err := idx.Rebuild(nil); err != nil {
|
||||
t.Fatalf("Rebuild(nil): %v", err)
|
||||
}
|
||||
|
||||
if idx.Len() != 0 {
|
||||
t.Errorf("Len() after empty rebuild = %d, want 0", idx.Len())
|
||||
}
|
||||
}
|
||||
|
||||
func TestVectorIndex_Persistence(t *testing.T) {
|
||||
tmpDir, err := os.MkdirTemp("", "synapbus-test-*")
|
||||
if err != nil {
|
||||
t.Fatalf("create temp dir: %v", err)
|
||||
}
|
||||
defer os.RemoveAll(tmpDir)
|
||||
|
||||
// Create and populate an index
|
||||
idx, err := NewVectorIndex(tmpDir)
|
||||
if err != nil {
|
||||
t.Fatalf("NewVectorIndex: %v", err)
|
||||
}
|
||||
|
||||
if err := idx.AddVector(1, []float32{1.0, 0.0, 0.0}); err != nil {
|
||||
t.Fatalf("AddVector: %v", err)
|
||||
}
|
||||
if err := idx.AddVector(2, []float32{0.0, 1.0, 0.0}); err != nil {
|
||||
t.Fatalf("AddVector: %v", err)
|
||||
}
|
||||
|
||||
// Save
|
||||
if err := idx.Save(); err != nil {
|
||||
t.Fatalf("Save: %v", err)
|
||||
}
|
||||
|
||||
// Load into new index
|
||||
idx2, err := NewVectorIndex(tmpDir)
|
||||
if err != nil {
|
||||
t.Fatalf("NewVectorIndex (reload): %v", err)
|
||||
}
|
||||
|
||||
if idx2.Len() != 2 {
|
||||
t.Errorf("reloaded Len() = %d, want 2", idx2.Len())
|
||||
}
|
||||
|
||||
// Search should still work
|
||||
results, err := idx2.Search([]float32{1.0, 0.0, 0.0}, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("Search after reload: %v", err)
|
||||
}
|
||||
if len(results) != 1 {
|
||||
t.Fatalf("search results after reload = %d, want 1", len(results))
|
||||
}
|
||||
if results[0].ID != 1 {
|
||||
t.Errorf("closest result ID = %d, want 1", results[0].ID)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,192 @@
|
||||
package search
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/search/embedding"
|
||||
)
|
||||
|
||||
// Pipeline processes unembedded messages in the background.
|
||||
type Pipeline struct {
|
||||
provider embedding.EmbeddingProvider
|
||||
store *EmbeddingStore
|
||||
index *VectorIndex
|
||||
config Config
|
||||
logger *slog.Logger
|
||||
cancel context.CancelFunc
|
||||
wg sync.WaitGroup
|
||||
}
|
||||
|
||||
// NewPipeline creates a new embedding pipeline.
|
||||
func NewPipeline(
|
||||
provider embedding.EmbeddingProvider,
|
||||
store *EmbeddingStore,
|
||||
index *VectorIndex,
|
||||
config Config,
|
||||
) *Pipeline {
|
||||
return &Pipeline{
|
||||
provider: provider,
|
||||
store: store,
|
||||
index: index,
|
||||
config: config,
|
||||
logger: slog.Default().With("component", "embedding-pipeline"),
|
||||
}
|
||||
}
|
||||
|
||||
// Start launches the pipeline worker goroutines.
|
||||
func (p *Pipeline) Start(ctx context.Context) {
|
||||
ctx, p.cancel = context.WithCancel(ctx)
|
||||
|
||||
workers := p.config.WorkerCount
|
||||
if workers <= 0 {
|
||||
workers = 1
|
||||
}
|
||||
|
||||
p.logger.Info("starting embedding pipeline", "workers", workers, "batch_size", p.config.BatchSize)
|
||||
|
||||
for i := 0; i < workers; i++ {
|
||||
p.wg.Add(1)
|
||||
go p.worker(ctx, i)
|
||||
}
|
||||
}
|
||||
|
||||
// Stop shuts down the pipeline and waits for workers to finish.
|
||||
func (p *Pipeline) Stop() {
|
||||
if p.cancel != nil {
|
||||
p.cancel()
|
||||
}
|
||||
p.wg.Wait()
|
||||
p.logger.Info("embedding pipeline stopped")
|
||||
}
|
||||
|
||||
// OnMessageCreated enqueues a new message for embedding.
|
||||
func (p *Pipeline) OnMessageCreated(ctx context.Context, messageID int64, body string) {
|
||||
if strings.TrimSpace(body) == "" {
|
||||
return
|
||||
}
|
||||
if err := p.store.Enqueue(ctx, messageID); err != nil {
|
||||
p.logger.Error("failed to enqueue message", "message_id", messageID, "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// OnMessageDeleted removes a message's embedding.
|
||||
func (p *Pipeline) OnMessageDeleted(ctx context.Context, messageID int64) {
|
||||
_ = p.store.DeleteEmbedding(ctx, messageID)
|
||||
p.index.Delete(messageID)
|
||||
}
|
||||
|
||||
func (p *Pipeline) worker(ctx context.Context, workerID int) {
|
||||
defer p.wg.Done()
|
||||
|
||||
logger := p.logger.With("worker", workerID)
|
||||
pollInterval := p.config.PollInterval
|
||||
if pollInterval <= 0 {
|
||||
pollInterval = 2 * time.Second
|
||||
}
|
||||
|
||||
ticker := time.NewTicker(pollInterval)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
p.processBatch(ctx, logger)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (p *Pipeline) processBatch(ctx context.Context, logger *slog.Logger) {
|
||||
batchSize := p.config.BatchSize
|
||||
if batchSize <= 0 {
|
||||
batchSize = 10
|
||||
}
|
||||
|
||||
items, err := p.store.Dequeue(ctx, batchSize)
|
||||
if err != nil {
|
||||
logger.Error("dequeue failed", "error", err)
|
||||
return
|
||||
}
|
||||
if len(items) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
// Gather message bodies
|
||||
type msgData struct {
|
||||
item QueueItem
|
||||
body string
|
||||
}
|
||||
var batch []msgData
|
||||
|
||||
for _, item := range items {
|
||||
body, err := p.store.GetMessageBody(ctx, item.MessageID)
|
||||
if err != nil {
|
||||
logger.Warn("message not found, marking completed", "message_id", item.MessageID, "error", err)
|
||||
_ = p.store.MarkCompleted(ctx, item.MessageID)
|
||||
continue
|
||||
}
|
||||
if strings.TrimSpace(body) == "" {
|
||||
_ = p.store.MarkCompleted(ctx, item.MessageID)
|
||||
continue
|
||||
}
|
||||
batch = append(batch, msgData{item: item, body: body})
|
||||
}
|
||||
|
||||
if len(batch) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
// Embed the batch
|
||||
texts := make([]string, len(batch))
|
||||
for i, m := range batch {
|
||||
texts[i] = m.body
|
||||
}
|
||||
|
||||
vectors, err := p.provider.EmbedBatch(ctx, texts)
|
||||
if err != nil {
|
||||
logger.Error("embedding failed", "error", err, "batch_size", len(batch))
|
||||
|
||||
// Mark all as failed for retry
|
||||
for _, m := range batch {
|
||||
errMsg := err.Error()
|
||||
_ = p.store.MarkFailed(ctx, m.item.MessageID, errMsg)
|
||||
}
|
||||
|
||||
// Requeue failed items below max attempts
|
||||
if requeued, err := p.store.RequeueFailed(ctx, p.config.RetryMaxAttempts); err == nil && requeued > 0 {
|
||||
logger.Info("requeued failed items for retry", "count", requeued)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Store results
|
||||
for i, m := range batch {
|
||||
if i >= len(vectors) || vectors[i] == nil {
|
||||
_ = p.store.MarkFailed(ctx, m.item.MessageID, "empty embedding returned")
|
||||
continue
|
||||
}
|
||||
|
||||
// Add to HNSW index
|
||||
if err := p.index.AddVector(m.item.MessageID, vectors[i]); err != nil {
|
||||
logger.Error("add to index failed", "message_id", m.item.MessageID, "error", err)
|
||||
_ = p.store.MarkFailed(ctx, m.item.MessageID, err.Error())
|
||||
continue
|
||||
}
|
||||
|
||||
// Record in SQLite
|
||||
if err := p.store.SaveEmbedding(ctx, m.item.MessageID, p.provider.Name(), p.provider.Name(), p.provider.Dimensions()); err != nil {
|
||||
logger.Error("save embedding record failed", "message_id", m.item.MessageID, "error", err)
|
||||
_ = p.store.MarkFailed(ctx, m.item.MessageID, err.Error())
|
||||
continue
|
||||
}
|
||||
|
||||
_ = p.store.MarkCompleted(ctx, m.item.MessageID)
|
||||
}
|
||||
|
||||
logger.Debug("batch processed", "count", len(batch))
|
||||
}
|
||||
@@ -0,0 +1,308 @@
|
||||
package search
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/messaging"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/search/embedding"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/storage"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/trace"
|
||||
)
|
||||
|
||||
func newPipelineTestDB(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)
|
||||
}
|
||||
|
||||
db.Exec(`INSERT OR IGNORE INTO users (id, username, password_hash, display_name) VALUES (1, 'testowner', 'hash', 'Test Owner')`)
|
||||
return db
|
||||
}
|
||||
|
||||
func TestPipeline_OnMessageCreated(t *testing.T) {
|
||||
db := newPipelineTestDB(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Need real messages in the DB for FK constraints
|
||||
db.Exec(`INSERT OR IGNORE INTO agents (name, display_name, type, owner_id, api_key_hash, status) VALUES ('s', 'S', 'ai', 1, 'hash', 'active')`)
|
||||
db.Exec(`INSERT INTO conversations (subject, created_by) VALUES ('test', 's')`)
|
||||
db.Exec(`INSERT INTO messages (conversation_id, from_agent, body, priority, status, metadata) VALUES (1, 's', 'hello world', 5, 'pending', '{}')`)
|
||||
db.Exec(`INSERT INTO messages (conversation_id, from_agent, body, priority, status, metadata) VALUES (1, 's', 'another msg', 5, 'pending', '{}')`)
|
||||
db.Exec(`INSERT INTO messages (conversation_id, from_agent, body, priority, status, metadata) VALUES (1, 's', 'third msg', 5, 'pending', '{}')`)
|
||||
|
||||
store := NewEmbeddingStore(db)
|
||||
idx := NewMemoryVectorIndex()
|
||||
provider := embedding.NewMockProvider(3)
|
||||
|
||||
pipeline := NewPipeline(provider, store, idx, Config{
|
||||
BatchSize: 10,
|
||||
WorkerCount: 1,
|
||||
PollInterval: 100 * time.Millisecond,
|
||||
RetryMaxAttempts: 3,
|
||||
})
|
||||
|
||||
t.Run("enqueues non-empty messages", func(t *testing.T) {
|
||||
pipeline.OnMessageCreated(ctx, 1, "hello world")
|
||||
|
||||
count, err := store.PendingCount(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("PendingCount: %v", err)
|
||||
}
|
||||
if count != 1 {
|
||||
t.Errorf("pending count = %d, want 1", count)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("skips empty messages", func(t *testing.T) {
|
||||
initialCount, _ := store.PendingCount(ctx)
|
||||
pipeline.OnMessageCreated(ctx, 2, "")
|
||||
pipeline.OnMessageCreated(ctx, 3, " ")
|
||||
|
||||
count, _ := store.PendingCount(ctx)
|
||||
if count != initialCount {
|
||||
t.Errorf("pending count changed from %d to %d for empty messages", initialCount, count)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestPipeline_ProcessBatch(t *testing.T) {
|
||||
db := newPipelineTestDB(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Register agents and send test messages
|
||||
db.Exec(`INSERT OR IGNORE INTO agents (name, display_name, type, owner_id, api_key_hash, status) VALUES ('sender', 'Sender', 'ai', 1, 'hash', 'active')`)
|
||||
db.Exec(`INSERT OR IGNORE INTO agents (name, display_name, type, owner_id, api_key_hash, status) VALUES ('receiver', 'Receiver', 'ai', 1, 'hash', 'active')`)
|
||||
|
||||
tracer := trace.NewTracer(db)
|
||||
t.Cleanup(func() { tracer.Close() })
|
||||
|
||||
msgStore := messaging.NewSQLiteMessageStore(db)
|
||||
msgService := messaging.NewMessagingService(msgStore, tracer)
|
||||
|
||||
msg1, _ := msgService.SendMessage(ctx, "sender", "receiver", "test message one", messaging.SendOptions{})
|
||||
msg2, _ := msgService.SendMessage(ctx, "sender", "receiver", "test message two", messaging.SendOptions{})
|
||||
|
||||
store := NewEmbeddingStore(db)
|
||||
idx := NewMemoryVectorIndex()
|
||||
provider := embedding.NewMockProvider(3)
|
||||
|
||||
cfg := Config{
|
||||
BatchSize: 10,
|
||||
WorkerCount: 1,
|
||||
PollInterval: 100 * time.Millisecond,
|
||||
RetryMaxAttempts: 3,
|
||||
}
|
||||
|
||||
pipeline := NewPipeline(provider, store, idx, cfg)
|
||||
|
||||
// Enqueue messages
|
||||
store.Enqueue(ctx, msg1.ID)
|
||||
store.Enqueue(ctx, msg2.ID)
|
||||
|
||||
// Start pipeline and wait for processing
|
||||
pipeline.Start(ctx)
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
pipeline.Stop()
|
||||
|
||||
// Check that vectors were added to index
|
||||
if idx.Len() != 2 {
|
||||
t.Errorf("index len = %d, want 2", idx.Len())
|
||||
}
|
||||
|
||||
// Check that embeddings were recorded
|
||||
embCount, err := store.EmbeddingCount(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("EmbeddingCount: %v", err)
|
||||
}
|
||||
if embCount != 2 {
|
||||
t.Errorf("embedding count = %d, want 2", embCount)
|
||||
}
|
||||
|
||||
// Check queue is cleared
|
||||
pending, _ := store.PendingCount(ctx)
|
||||
if pending != 0 {
|
||||
t.Errorf("pending count = %d, want 0", pending)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPipeline_ErrorHandling(t *testing.T) {
|
||||
db := newPipelineTestDB(t)
|
||||
ctx := context.Background()
|
||||
|
||||
db.Exec(`INSERT OR IGNORE INTO agents (name, display_name, type, owner_id, api_key_hash, status) VALUES ('sender', 'Sender', 'ai', 1, 'hash', 'active')`)
|
||||
db.Exec(`INSERT OR IGNORE INTO agents (name, display_name, type, owner_id, api_key_hash, status) VALUES ('receiver', 'Receiver', 'ai', 1, 'hash', 'active')`)
|
||||
|
||||
tracer := trace.NewTracer(db)
|
||||
t.Cleanup(func() { tracer.Close() })
|
||||
|
||||
msgStore := messaging.NewSQLiteMessageStore(db)
|
||||
msgService := messaging.NewMessagingService(msgStore, tracer)
|
||||
msg1, _ := msgService.SendMessage(ctx, "sender", "receiver", "test message", messaging.SendOptions{})
|
||||
|
||||
store := NewEmbeddingStore(db)
|
||||
idx := NewMemoryVectorIndex()
|
||||
|
||||
// Create a provider that fails
|
||||
failProvider := embedding.NewMockProvider(3)
|
||||
failProvider.SetBatchFunc(func(ctx context.Context, texts []string) ([][]float32, error) {
|
||||
return nil, fmt.Errorf("provider error: rate limited")
|
||||
})
|
||||
|
||||
cfg := Config{
|
||||
BatchSize: 10,
|
||||
WorkerCount: 1,
|
||||
PollInterval: 100 * time.Millisecond,
|
||||
RetryMaxAttempts: 3,
|
||||
}
|
||||
|
||||
pipeline := NewPipeline(failProvider, store, idx, cfg)
|
||||
store.Enqueue(ctx, msg1.ID)
|
||||
|
||||
pipeline.Start(ctx)
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
pipeline.Stop()
|
||||
|
||||
// Index should be empty (embedding failed)
|
||||
if idx.Len() != 0 {
|
||||
t.Errorf("index len = %d, want 0 (embedding should have failed)", idx.Len())
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmbeddingStore_QueueOperations(t *testing.T) {
|
||||
db := newPipelineTestDB(t)
|
||||
ctx := context.Background()
|
||||
|
||||
store := NewEmbeddingStore(db)
|
||||
|
||||
// Create test messages directly
|
||||
db.Exec(`INSERT OR IGNORE INTO agents (name, display_name, type, owner_id, api_key_hash, status) VALUES ('s', 'S', 'ai', 1, 'hash', 'active')`)
|
||||
db.Exec(`INSERT INTO conversations (subject, created_by) VALUES ('test', 's')`)
|
||||
db.Exec(`INSERT INTO messages (conversation_id, from_agent, body, priority, status, metadata) VALUES (1, 's', 'msg1', 5, 'pending', '{}')`)
|
||||
db.Exec(`INSERT INTO messages (conversation_id, from_agent, body, priority, status, metadata) VALUES (1, 's', 'msg2', 5, 'pending', '{}')`)
|
||||
|
||||
t.Run("enqueue and dequeue", func(t *testing.T) {
|
||||
store.Enqueue(ctx, 1)
|
||||
store.Enqueue(ctx, 2)
|
||||
|
||||
count, _ := store.PendingCount(ctx)
|
||||
if count != 2 {
|
||||
t.Errorf("pending count = %d, want 2", count)
|
||||
}
|
||||
|
||||
items, err := store.Dequeue(ctx, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("Dequeue: %v", err)
|
||||
}
|
||||
if len(items) != 1 {
|
||||
t.Fatalf("dequeued = %d, want 1", len(items))
|
||||
}
|
||||
if items[0].MessageID != 1 {
|
||||
t.Errorf("dequeued message_id = %d, want 1", items[0].MessageID)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("mark completed", func(t *testing.T) {
|
||||
store.MarkCompleted(ctx, 1)
|
||||
|
||||
// Should be able to dequeue the second item
|
||||
items, err := store.Dequeue(ctx, 10)
|
||||
if err != nil {
|
||||
t.Fatalf("Dequeue: %v", err)
|
||||
}
|
||||
// The second item should be available (first was dequeued, now processing)
|
||||
found := false
|
||||
for _, item := range items {
|
||||
if item.MessageID == 2 {
|
||||
found = true
|
||||
}
|
||||
}
|
||||
if !found && len(items) == 0 {
|
||||
// item 2 may already have been dequeued in the first call's processing
|
||||
// This is OK in the test
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("duplicate enqueue is ignored", func(t *testing.T) {
|
||||
db.Exec(`INSERT INTO messages (conversation_id, from_agent, body, priority, status, metadata) VALUES (1, 's', 'msg3', 5, 'pending', '{}')`)
|
||||
store.Enqueue(ctx, 3)
|
||||
store.Enqueue(ctx, 3) // duplicate
|
||||
// Should not error
|
||||
})
|
||||
}
|
||||
|
||||
func TestEmbeddingStore_EmbeddingOperations(t *testing.T) {
|
||||
db := newPipelineTestDB(t)
|
||||
ctx := context.Background()
|
||||
|
||||
store := NewEmbeddingStore(db)
|
||||
|
||||
// Create a message
|
||||
db.Exec(`INSERT OR IGNORE INTO agents (name, display_name, type, owner_id, api_key_hash, status) VALUES ('s', 'S', 'ai', 1, 'hash', 'active')`)
|
||||
db.Exec(`INSERT INTO conversations (subject, created_by) VALUES ('test', 's')`)
|
||||
db.Exec(`INSERT INTO messages (conversation_id, from_agent, body, priority, status, metadata) VALUES (1, 's', 'msg1', 5, 'pending', '{}')`)
|
||||
|
||||
t.Run("save and count embeddings", func(t *testing.T) {
|
||||
err := store.SaveEmbedding(ctx, 1, "openai", "text-embedding-3-small", 1536)
|
||||
if err != nil {
|
||||
t.Fatalf("SaveEmbedding: %v", err)
|
||||
}
|
||||
|
||||
count, err := store.EmbeddingCount(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("EmbeddingCount: %v", err)
|
||||
}
|
||||
if count != 1 {
|
||||
t.Errorf("embedding count = %d, want 1", count)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("get provider", func(t *testing.T) {
|
||||
provider, err := store.GetEmbeddingProvider(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("GetEmbeddingProvider: %v", err)
|
||||
}
|
||||
if provider != "openai" {
|
||||
t.Errorf("provider = %q, want openai", provider)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("delete embedding", func(t *testing.T) {
|
||||
err := store.DeleteEmbedding(ctx, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("DeleteEmbedding: %v", err)
|
||||
}
|
||||
|
||||
count, _ := store.EmbeddingCount(ctx)
|
||||
if count != 0 {
|
||||
t.Errorf("embedding count after delete = %d, want 0", count)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("delete all embeddings", func(t *testing.T) {
|
||||
store.SaveEmbedding(ctx, 1, "openai", "model", 1536)
|
||||
store.DeleteAllEmbeddings(ctx)
|
||||
|
||||
count, _ := store.EmbeddingCount(ctx)
|
||||
if count != 0 {
|
||||
t.Errorf("embedding count after delete all = %d, want 0", count)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,371 @@
|
||||
package search
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/messaging"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/search/embedding"
|
||||
)
|
||||
|
||||
// SearchMode constants.
|
||||
const (
|
||||
ModeAuto = "auto"
|
||||
ModeSemantic = "semantic"
|
||||
ModeFulltext = "fulltext"
|
||||
)
|
||||
|
||||
// SearchOptions for the unified search service.
|
||||
type SearchOptions struct {
|
||||
Query string
|
||||
Mode string // "auto", "semantic", "fulltext"
|
||||
Limit int
|
||||
ChannelID *int64
|
||||
FromAgent string
|
||||
MinPriority int
|
||||
After *time.Time
|
||||
Before *time.Time
|
||||
}
|
||||
|
||||
// SearchResult represents a single search result.
|
||||
type SearchResult struct {
|
||||
Message *messaging.Message `json:"message"`
|
||||
SimilarityScore float64 `json:"similarity_score,omitempty"`
|
||||
RelevanceScore float64 `json:"relevance_score,omitempty"`
|
||||
MatchType string `json:"match_type"` // "semantic" or "fulltext"
|
||||
}
|
||||
|
||||
// SearchResponse is the response from the search service.
|
||||
type SearchResponse struct {
|
||||
Results []*SearchResult `json:"results"`
|
||||
SearchMode string `json:"search_mode"`
|
||||
TotalResults int `json:"total_results"`
|
||||
Warning string `json:"warning,omitempty"`
|
||||
}
|
||||
|
||||
// Service provides unified search combining semantic and full-text search.
|
||||
type Service struct {
|
||||
db *sql.DB
|
||||
provider embedding.EmbeddingProvider // may be nil
|
||||
index *VectorIndex // may be nil
|
||||
msgService *messaging.MessagingService
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewService creates a new search service.
|
||||
func NewService(
|
||||
db *sql.DB,
|
||||
provider embedding.EmbeddingProvider,
|
||||
index *VectorIndex,
|
||||
msgService *messaging.MessagingService,
|
||||
) *Service {
|
||||
return &Service{
|
||||
db: db,
|
||||
provider: provider,
|
||||
index: index,
|
||||
msgService: msgService,
|
||||
logger: slog.Default().With("component", "search"),
|
||||
}
|
||||
}
|
||||
|
||||
// Search performs a search with the given options, respecting agent access control.
|
||||
func (s *Service) Search(ctx context.Context, agentName string, opts SearchOptions) (*SearchResponse, error) {
|
||||
limit := opts.Limit
|
||||
if limit <= 0 {
|
||||
limit = 10
|
||||
}
|
||||
if limit > 100 {
|
||||
limit = 100
|
||||
}
|
||||
|
||||
mode := opts.Mode
|
||||
if mode == "" {
|
||||
mode = ModeAuto
|
||||
}
|
||||
|
||||
// Determine effective search mode
|
||||
switch mode {
|
||||
case ModeAuto:
|
||||
if s.provider != nil && s.index != nil && s.index.Len() > 0 {
|
||||
return s.semanticSearch(ctx, agentName, opts, limit)
|
||||
}
|
||||
return s.fulltextSearch(ctx, agentName, opts, limit)
|
||||
|
||||
case ModeSemantic:
|
||||
if s.provider == nil || s.index == nil {
|
||||
return nil, fmt.Errorf("semantic search unavailable: no embedding provider configured")
|
||||
}
|
||||
resp, err := s.semanticSearch(ctx, agentName, opts, limit)
|
||||
if err != nil {
|
||||
// Fall back to full-text on semantic error
|
||||
s.logger.Warn("semantic search failed, falling back to fulltext", "error", err)
|
||||
resp, ftErr := s.fulltextSearch(ctx, agentName, opts, limit)
|
||||
if ftErr != nil {
|
||||
return nil, fmt.Errorf("semantic search failed: %w; fulltext fallback also failed: %w", err, ftErr)
|
||||
}
|
||||
resp.Warning = fmt.Sprintf("semantic search failed: %s, using fulltext fallback", err)
|
||||
return resp, nil
|
||||
}
|
||||
return resp, nil
|
||||
|
||||
case ModeFulltext:
|
||||
return s.fulltextSearch(ctx, agentName, opts, limit)
|
||||
|
||||
default:
|
||||
return nil, fmt.Errorf("unknown search mode: %q", mode)
|
||||
}
|
||||
}
|
||||
|
||||
// semanticSearch performs vector similarity search.
|
||||
func (s *Service) semanticSearch(ctx context.Context, agentName string, opts SearchOptions, limit int) (*SearchResponse, error) {
|
||||
if opts.Query == "" {
|
||||
return s.fulltextSearch(ctx, agentName, opts, limit)
|
||||
}
|
||||
|
||||
// Embed the query
|
||||
queryVec, err := s.provider.Embed(ctx, opts.Query)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("embed query: %w", err)
|
||||
}
|
||||
|
||||
// Over-fetch to account for access control filtering
|
||||
overfetch := limit * 5
|
||||
if overfetch < 50 {
|
||||
overfetch = 50
|
||||
}
|
||||
|
||||
results, err := s.index.Search(queryVec, overfetch)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("vector search: %w", err)
|
||||
}
|
||||
|
||||
if len(results) == 0 {
|
||||
// No vectors in index, fall back to FTS
|
||||
resp, err := s.fulltextSearch(ctx, agentName, opts, limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resp.Warning = "no vectors in index, using fulltext"
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
// Fetch messages and apply access control + filters
|
||||
var searchResults []*SearchResult
|
||||
for _, vr := range results {
|
||||
msg, err := s.getMessageByID(ctx, vr.ID)
|
||||
if err != nil {
|
||||
continue // message may have been deleted
|
||||
}
|
||||
|
||||
// Access control: only messages this agent can see
|
||||
if !s.canAgentAccessMessage(ctx, agentName, msg) {
|
||||
continue
|
||||
}
|
||||
|
||||
// Apply filters
|
||||
if !s.matchesFilters(msg, opts) {
|
||||
continue
|
||||
}
|
||||
|
||||
similarity := float64(1.0 - vr.Distance) // cosine distance -> similarity
|
||||
if similarity < 0 {
|
||||
similarity = 0
|
||||
}
|
||||
|
||||
searchResults = append(searchResults, &SearchResult{
|
||||
Message: msg,
|
||||
SimilarityScore: similarity,
|
||||
MatchType: ModeSemantic,
|
||||
})
|
||||
|
||||
if len(searchResults) >= limit {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
return &SearchResponse{
|
||||
Results: searchResults,
|
||||
SearchMode: ModeSemantic,
|
||||
TotalResults: len(searchResults),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// fulltextSearch performs FTS5 full-text search with access control.
|
||||
func (s *Service) fulltextSearch(ctx context.Context, agentName string, opts SearchOptions, limit int) (*SearchResponse, error) {
|
||||
// Use the messaging service's SearchMessages which already handles access control
|
||||
msgOpts := messaging.SearchOptions{
|
||||
FromAgent: opts.FromAgent,
|
||||
MinPriority: opts.MinPriority,
|
||||
Limit: limit,
|
||||
}
|
||||
if opts.ChannelID != nil {
|
||||
msgOpts.ChannelID = opts.ChannelID
|
||||
}
|
||||
|
||||
messages, err := s.msgService.SearchMessages(ctx, agentName, opts.Query, msgOpts)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("fulltext search: %w", err)
|
||||
}
|
||||
|
||||
results := make([]*SearchResult, len(messages))
|
||||
for i, msg := range messages {
|
||||
results[i] = &SearchResult{
|
||||
Message: msg,
|
||||
RelevanceScore: 1.0 - float64(i)*0.05, // simple rank-based score
|
||||
MatchType: ModeFulltext,
|
||||
}
|
||||
}
|
||||
|
||||
return &SearchResponse{
|
||||
Results: results,
|
||||
SearchMode: ModeFulltext,
|
||||
TotalResults: len(results),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// getMessageByID fetches a message by ID from the database.
|
||||
func (s *Service) getMessageByID(ctx context.Context, id int64) (*messaging.Message, error) {
|
||||
var msg messaging.Message
|
||||
var toAgent, claimedBy sql.NullString
|
||||
var channelID sql.NullInt64
|
||||
var claimedAt sql.NullTime
|
||||
var metadata string
|
||||
|
||||
err := 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,
|
||||
).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 {
|
||||
msg.ClaimedAt = &claimedAt.Time
|
||||
}
|
||||
msg.Metadata = json.RawMessage(metadata)
|
||||
|
||||
return &msg, nil
|
||||
}
|
||||
|
||||
// canAgentAccessMessage checks if the agent has access to the message.
|
||||
func (s *Service) canAgentAccessMessage(ctx context.Context, agentName string, msg *messaging.Message) bool {
|
||||
// Direct messages: agent must be sender or recipient
|
||||
if msg.ToAgent != "" {
|
||||
if msg.FromAgent == agentName || msg.ToAgent == agentName {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// Channel messages: agent must be a member of the channel
|
||||
if msg.ChannelID != nil {
|
||||
var count int
|
||||
err := s.db.QueryRowContext(ctx,
|
||||
`SELECT COUNT(*) FROM channel_members WHERE channel_id = ? AND agent_name = ?`,
|
||||
*msg.ChannelID, agentName,
|
||||
).Scan(&count)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return count > 0
|
||||
}
|
||||
|
||||
// Messages from or to the agent in conversations
|
||||
if msg.FromAgent == agentName {
|
||||
return true
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
// matchesFilters checks if a message matches the given filter options.
|
||||
func (s *Service) matchesFilters(msg *messaging.Message, opts SearchOptions) bool {
|
||||
if opts.ChannelID != nil && (msg.ChannelID == nil || *msg.ChannelID != *opts.ChannelID) {
|
||||
return false
|
||||
}
|
||||
if opts.FromAgent != "" && msg.FromAgent != opts.FromAgent {
|
||||
return false
|
||||
}
|
||||
if opts.MinPriority > 0 && msg.Priority < opts.MinPriority {
|
||||
return false
|
||||
}
|
||||
if opts.After != nil && msg.CreatedAt.Before(*opts.After) {
|
||||
return false
|
||||
}
|
||||
if opts.Before != nil && msg.CreatedAt.After(*opts.Before) {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// HasSemanticSearch returns true if semantic search is available.
|
||||
func (s *Service) HasSemanticSearch() bool {
|
||||
return s.provider != nil && s.index != nil
|
||||
}
|
||||
|
||||
// IndexSize returns the number of vectors in the index.
|
||||
func (s *Service) IndexSize() int {
|
||||
if s.index == nil {
|
||||
return 0
|
||||
}
|
||||
return s.index.Len()
|
||||
}
|
||||
|
||||
// ProviderName returns the name of the configured provider, or empty string.
|
||||
func (s *Service) ProviderName() string {
|
||||
if s.provider == nil {
|
||||
return ""
|
||||
}
|
||||
return s.provider.Name()
|
||||
}
|
||||
|
||||
// SearchMessagesCompat provides a compatibility method that returns []*messaging.Message
|
||||
// for use by the existing MCP tool handler when semantic search is not requested.
|
||||
func (s *Service) SearchMessagesCompat(ctx context.Context, agentName, query string, opts messaging.SearchOptions) ([]*messaging.Message, string, error) {
|
||||
searchOpts := SearchOptions{
|
||||
Query: query,
|
||||
Mode: ModeAuto,
|
||||
Limit: opts.Limit,
|
||||
FromAgent: opts.FromAgent,
|
||||
MinPriority: opts.MinPriority,
|
||||
ChannelID: opts.ChannelID,
|
||||
}
|
||||
|
||||
resp, err := s.Search(ctx, agentName, searchOpts)
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
|
||||
messages := make([]*messaging.Message, len(resp.Results))
|
||||
for i, r := range resp.Results {
|
||||
messages[i] = r.Message
|
||||
}
|
||||
|
||||
// Clean up query for search mode display
|
||||
searchMode := resp.SearchMode
|
||||
if strings.TrimSpace(query) == "" {
|
||||
searchMode = ModeFulltext
|
||||
}
|
||||
|
||||
return messages, searchMode, nil
|
||||
}
|
||||
@@ -0,0 +1,385 @@
|
||||
package search
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/messaging"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/search/embedding"
|
||||
"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 newTestServices(t *testing.T) (*Service, *messaging.MessagingService, *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)
|
||||
|
||||
// Create search service without semantic search (FTS-only)
|
||||
searchService := NewService(db, nil, nil, msgService)
|
||||
return searchService, msgService, db
|
||||
}
|
||||
|
||||
func seedTestAgents(t *testing.T, db *sql.DB, names ...string) {
|
||||
t.Helper()
|
||||
for _, name := range names {
|
||||
db.Exec(
|
||||
`INSERT OR IGNORE INTO agents (name, display_name, type, owner_id, api_key_hash, status)
|
||||
VALUES (?, ?, 'ai', 1, 'hash', 'active')`,
|
||||
name, name,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
func TestService_FulltextSearch(t *testing.T) {
|
||||
svc, msgSvc, db := newTestServices(t)
|
||||
ctx := context.Background()
|
||||
|
||||
seedTestAgents(t, db, "sender", "searcher")
|
||||
|
||||
// Send some messages
|
||||
msgSvc.SendMessage(ctx, "sender", "searcher", "deployment failed in staging", messaging.SendOptions{})
|
||||
msgSvc.SendMessage(ctx, "sender", "searcher", "all services healthy", messaging.SendOptions{})
|
||||
msgSvc.SendMessage(ctx, "sender", "searcher", "database connection timeout", messaging.SendOptions{})
|
||||
|
||||
t.Run("keyword search returns matching messages", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "searcher", SearchOptions{
|
||||
Query: "deployment",
|
||||
Mode: ModeFulltext,
|
||||
Limit: 10,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
if resp.SearchMode != ModeFulltext {
|
||||
t.Errorf("search_mode = %q, want %q", resp.SearchMode, ModeFulltext)
|
||||
}
|
||||
if len(resp.Results) != 1 {
|
||||
t.Errorf("result count = %d, want 1", len(resp.Results))
|
||||
}
|
||||
if len(resp.Results) > 0 && resp.Results[0].MatchType != ModeFulltext {
|
||||
t.Errorf("match_type = %q, want %q", resp.Results[0].MatchType, ModeFulltext)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("empty query returns all messages", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "searcher", SearchOptions{
|
||||
Query: "",
|
||||
Mode: ModeFulltext,
|
||||
Limit: 10,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
if len(resp.Results) != 3 {
|
||||
t.Errorf("result count = %d, want 3", len(resp.Results))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("auto mode falls back to fulltext without provider", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "searcher", SearchOptions{
|
||||
Query: "deployment",
|
||||
Mode: ModeAuto,
|
||||
Limit: 10,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
if resp.SearchMode != ModeFulltext {
|
||||
t.Errorf("auto search_mode = %q, want %q", resp.SearchMode, ModeFulltext)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("semantic mode fails without provider", func(t *testing.T) {
|
||||
_, err := svc.Search(ctx, "searcher", SearchOptions{
|
||||
Query: "deployment",
|
||||
Mode: ModeSemantic,
|
||||
Limit: 10,
|
||||
})
|
||||
if err == nil {
|
||||
t.Error("expected error for semantic mode without provider")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("limit is enforced", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "searcher", SearchOptions{
|
||||
Query: "",
|
||||
Mode: ModeFulltext,
|
||||
Limit: 1,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
if len(resp.Results) != 1 {
|
||||
t.Errorf("result count = %d, want 1", len(resp.Results))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("max limit is 100", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "searcher", SearchOptions{
|
||||
Query: "",
|
||||
Mode: ModeFulltext,
|
||||
Limit: 500,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
// Should work, just capped
|
||||
_ = resp
|
||||
})
|
||||
}
|
||||
|
||||
func TestService_AccessControl(t *testing.T) {
|
||||
svc, msgSvc, db := newTestServices(t)
|
||||
ctx := context.Background()
|
||||
|
||||
seedTestAgents(t, db, "alice", "bob", "charlie")
|
||||
|
||||
// Alice sends to Bob
|
||||
msgSvc.SendMessage(ctx, "alice", "bob", "secret for bob", messaging.SendOptions{})
|
||||
|
||||
// Alice sends to Charlie
|
||||
msgSvc.SendMessage(ctx, "alice", "charlie", "secret for charlie", messaging.SendOptions{})
|
||||
|
||||
t.Run("bob sees only his messages", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "bob", SearchOptions{
|
||||
Query: "secret",
|
||||
Mode: ModeFulltext,
|
||||
Limit: 10,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
if len(resp.Results) != 1 {
|
||||
t.Errorf("bob result count = %d, want 1", len(resp.Results))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("charlie sees only his messages", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "charlie", SearchOptions{
|
||||
Query: "secret",
|
||||
Mode: ModeFulltext,
|
||||
Limit: 10,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
if len(resp.Results) != 1 {
|
||||
t.Errorf("charlie result count = %d, want 1", len(resp.Results))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("alice sees all (she is the sender)", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "alice", SearchOptions{
|
||||
Query: "secret",
|
||||
Mode: ModeFulltext,
|
||||
Limit: 10,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
if len(resp.Results) != 2 {
|
||||
t.Errorf("alice result count = %d, want 2", len(resp.Results))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestService_SemanticSearch(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
ctx := context.Background()
|
||||
|
||||
tracer := trace.NewTracer(db)
|
||||
t.Cleanup(func() { tracer.Close() })
|
||||
|
||||
seedTestAgents(t, db, "sender", "searcher")
|
||||
|
||||
msgStore := messaging.NewSQLiteMessageStore(db)
|
||||
msgService := messaging.NewMessagingService(msgStore, tracer)
|
||||
|
||||
// Create mock provider
|
||||
mockProvider := embedding.NewMockProvider(3)
|
||||
|
||||
// Create index with vectors
|
||||
idx := NewMemoryVectorIndex()
|
||||
|
||||
// Send messages
|
||||
msg1, _ := msgService.SendMessage(ctx, "sender", "searcher", "deployment failure in staging", messaging.SendOptions{})
|
||||
msg2, _ := msgService.SendMessage(ctx, "sender", "searcher", "cat pictures are cute", messaging.SendOptions{})
|
||||
msg3, _ := msgService.SendMessage(ctx, "sender", "searcher", "staging server crashed", messaging.SendOptions{})
|
||||
|
||||
// Add vectors to index (simulating what the pipeline would do)
|
||||
// msg1 and msg3 are about similar topics (large X component), msg2 is very different (large Z component)
|
||||
idx.AddVector(msg1.ID, []float32{0.95, 0.1, 0.0}) // deployment-related
|
||||
idx.AddVector(msg2.ID, []float32{0.0, 0.0, 1.0}) // unrelated (orthogonal)
|
||||
idx.AddVector(msg3.ID, []float32{0.90, 0.15, 0.0}) // deployment-related
|
||||
|
||||
// Override mock provider to return a vector similar to deployment topics
|
||||
mockProvider.SetEmbedFunc(func(ctx context.Context, text string) ([]float32, error) {
|
||||
return []float32{0.95, 0.1, 0.0}, nil // query vector close to deployment msgs
|
||||
})
|
||||
|
||||
svc := NewService(db, mockProvider, idx, msgService)
|
||||
|
||||
t.Run("semantic search returns ranked results", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "searcher", SearchOptions{
|
||||
Query: "staging deployment issues",
|
||||
Mode: ModeSemantic,
|
||||
Limit: 10,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
if resp.SearchMode != ModeSemantic {
|
||||
t.Errorf("search_mode = %q, want %q", resp.SearchMode, ModeSemantic)
|
||||
}
|
||||
if len(resp.Results) < 2 {
|
||||
t.Fatalf("expected at least 2 results, got %d", len(resp.Results))
|
||||
}
|
||||
|
||||
// Results should have similarity scores >= 0
|
||||
for _, r := range resp.Results {
|
||||
if r.SimilarityScore < 0 {
|
||||
t.Errorf("expected non-negative similarity score, got %f for msg %d", r.SimilarityScore, r.Message.ID)
|
||||
}
|
||||
}
|
||||
|
||||
// The deployment-related messages should score higher than the cat pictures message
|
||||
var deploymentScore, catScore float64
|
||||
for _, r := range resp.Results {
|
||||
if r.Message.ID == msg1.ID {
|
||||
deploymentScore = r.SimilarityScore
|
||||
}
|
||||
if r.Message.ID == msg2.ID {
|
||||
catScore = r.SimilarityScore
|
||||
}
|
||||
}
|
||||
|
||||
if deploymentScore <= catScore {
|
||||
t.Errorf("deployment msg scored %f, cat msg scored %f; deployment should score higher",
|
||||
deploymentScore, catScore)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("auto mode uses semantic when available", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "searcher", SearchOptions{
|
||||
Query: "staging deployment",
|
||||
Mode: ModeAuto,
|
||||
Limit: 10,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
if resp.SearchMode != ModeSemantic {
|
||||
t.Errorf("auto search_mode = %q, want %q", resp.SearchMode, ModeSemantic)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("fulltext mode is always available", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "searcher", SearchOptions{
|
||||
Query: "deployment",
|
||||
Mode: ModeFulltext,
|
||||
Limit: 10,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
if resp.SearchMode != ModeFulltext {
|
||||
t.Errorf("search_mode = %q, want %q", resp.SearchMode, ModeFulltext)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestService_Filters(t *testing.T) {
|
||||
svc, msgSvc, db := newTestServices(t)
|
||||
ctx := context.Background()
|
||||
|
||||
seedTestAgents(t, db, "agent-a", "agent-b", "searcher")
|
||||
|
||||
msgSvc.SendMessage(ctx, "agent-a", "searcher", "low priority task", messaging.SendOptions{Priority: 2})
|
||||
msgSvc.SendMessage(ctx, "agent-a", "searcher", "high priority alert", messaging.SendOptions{Priority: 9})
|
||||
msgSvc.SendMessage(ctx, "agent-b", "searcher", "from agent-b", messaging.SendOptions{})
|
||||
|
||||
t.Run("filter by from_agent", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "searcher", SearchOptions{
|
||||
Query: "",
|
||||
Mode: ModeFulltext,
|
||||
Limit: 10,
|
||||
FromAgent: "agent-a",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
if len(resp.Results) != 2 {
|
||||
t.Errorf("result count = %d, want 2", len(resp.Results))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("filter by min_priority", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "searcher", SearchOptions{
|
||||
Query: "",
|
||||
Mode: ModeFulltext,
|
||||
Limit: 10,
|
||||
MinPriority: 5,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
// Only the high priority message and agent-b (default priority 5) match
|
||||
for _, r := range resp.Results {
|
||||
if r.Message.Priority < 5 {
|
||||
t.Errorf("got message with priority %d, want >= 5", r.Message.Priority)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestService_HasSemanticSearch(t *testing.T) {
|
||||
t.Run("without provider", func(t *testing.T) {
|
||||
svc := NewService(nil, nil, nil, nil)
|
||||
if svc.HasSemanticSearch() {
|
||||
t.Error("HasSemanticSearch() = true, want false")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("with provider and index", func(t *testing.T) {
|
||||
idx := NewMemoryVectorIndex()
|
||||
provider := embedding.NewMockProvider(3)
|
||||
svc := NewService(nil, provider, idx, nil)
|
||||
if !svc.HasSemanticSearch() {
|
||||
t.Error("HasSemanticSearch() = false, want true")
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,244 @@
|
||||
package search
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// EmbeddingRecord tracks which messages have been embedded.
|
||||
type EmbeddingRecord struct {
|
||||
MessageID int64
|
||||
Provider string
|
||||
Model string
|
||||
Dimensions int
|
||||
EmbeddedAt time.Time
|
||||
}
|
||||
|
||||
// QueueItem represents a message in the embedding queue.
|
||||
type QueueItem struct {
|
||||
ID int64
|
||||
MessageID int64
|
||||
Status string
|
||||
Attempts int
|
||||
LastError string
|
||||
CreatedAt time.Time
|
||||
}
|
||||
|
||||
// EmbeddingStore manages the embeddings and embedding_queue tables.
|
||||
type EmbeddingStore struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
// NewEmbeddingStore creates a new embedding store.
|
||||
func NewEmbeddingStore(db *sql.DB) *EmbeddingStore {
|
||||
return &EmbeddingStore{db: db}
|
||||
}
|
||||
|
||||
// SaveEmbedding records that a message has been embedded.
|
||||
func (s *EmbeddingStore) SaveEmbedding(ctx context.Context, messageID int64, provider, model string, dimensions int) error {
|
||||
_, err := s.db.ExecContext(ctx,
|
||||
`INSERT OR REPLACE INTO embeddings (message_id, provider, model, dimensions, embedded_at)
|
||||
VALUES (?, ?, ?, ?, CURRENT_TIMESTAMP)`,
|
||||
messageID, provider, model, dimensions,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("save embedding: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteEmbedding removes an embedding record.
|
||||
func (s *EmbeddingStore) DeleteEmbedding(ctx context.Context, messageID int64) error {
|
||||
_, err := s.db.ExecContext(ctx,
|
||||
`DELETE FROM embeddings WHERE message_id = ?`, messageID,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
// DeleteAllEmbeddings removes all embedding records (for provider switch).
|
||||
func (s *EmbeddingStore) DeleteAllEmbeddings(ctx context.Context) error {
|
||||
_, err := s.db.ExecContext(ctx, `DELETE FROM embeddings`)
|
||||
return err
|
||||
}
|
||||
|
||||
// GetEmbeddingProvider returns the provider of the most recent embedding, if any.
|
||||
func (s *EmbeddingStore) GetEmbeddingProvider(ctx context.Context) (string, error) {
|
||||
var provider string
|
||||
err := s.db.QueryRowContext(ctx,
|
||||
`SELECT provider FROM embeddings ORDER BY embedded_at DESC LIMIT 1`,
|
||||
).Scan(&provider)
|
||||
if err == sql.ErrNoRows {
|
||||
return "", nil
|
||||
}
|
||||
return provider, err
|
||||
}
|
||||
|
||||
// EmbeddingCount returns the number of embedded messages.
|
||||
func (s *EmbeddingStore) EmbeddingCount(ctx context.Context) (int64, error) {
|
||||
var count int64
|
||||
err := s.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM embeddings`).Scan(&count)
|
||||
return count, err
|
||||
}
|
||||
|
||||
// --- Queue operations ---
|
||||
|
||||
// Enqueue adds a message to the embedding queue.
|
||||
func (s *EmbeddingStore) Enqueue(ctx context.Context, messageID int64) error {
|
||||
_, err := s.db.ExecContext(ctx,
|
||||
`INSERT OR IGNORE INTO embedding_queue (message_id, status, created_at)
|
||||
VALUES (?, 'pending', CURRENT_TIMESTAMP)`,
|
||||
messageID,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("enqueue message %d: %w", messageID, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Dequeue atomically fetches and claims a batch of pending items.
|
||||
func (s *EmbeddingStore) Dequeue(ctx context.Context, batchSize int) ([]QueueItem, error) {
|
||||
tx, err := s.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("begin tx: %w", err)
|
||||
}
|
||||
defer tx.Rollback()
|
||||
|
||||
rows, err := tx.QueryContext(ctx,
|
||||
`SELECT id, message_id, attempts FROM embedding_queue
|
||||
WHERE status = 'pending'
|
||||
ORDER BY created_at ASC
|
||||
LIMIT ?`,
|
||||
batchSize,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("query pending: %w", err)
|
||||
}
|
||||
|
||||
var items []QueueItem
|
||||
var ids []int64
|
||||
for rows.Next() {
|
||||
var item QueueItem
|
||||
if err := rows.Scan(&item.ID, &item.MessageID, &item.Attempts); err != nil {
|
||||
rows.Close()
|
||||
return nil, fmt.Errorf("scan queue item: %w", err)
|
||||
}
|
||||
item.Status = "processing"
|
||||
items = append(items, item)
|
||||
ids = append(ids, item.ID)
|
||||
}
|
||||
rows.Close()
|
||||
|
||||
if len(ids) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// Mark as processing
|
||||
for _, id := range ids {
|
||||
_, err := tx.ExecContext(ctx,
|
||||
`UPDATE embedding_queue SET status = 'processing', attempts = attempts + 1 WHERE id = ?`,
|
||||
id,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("mark processing: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := tx.Commit(); err != nil {
|
||||
return nil, fmt.Errorf("commit dequeue: %w", err)
|
||||
}
|
||||
|
||||
return items, nil
|
||||
}
|
||||
|
||||
// MarkCompleted marks a queue item as completed.
|
||||
func (s *EmbeddingStore) MarkCompleted(ctx context.Context, messageID int64) error {
|
||||
_, err := s.db.ExecContext(ctx,
|
||||
`UPDATE embedding_queue SET status = 'completed', completed_at = CURRENT_TIMESTAMP
|
||||
WHERE message_id = ?`,
|
||||
messageID,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
// MarkFailed marks a queue item as failed with an error message.
|
||||
func (s *EmbeddingStore) MarkFailed(ctx context.Context, messageID int64, errMsg string) error {
|
||||
_, err := s.db.ExecContext(ctx,
|
||||
`UPDATE embedding_queue SET status = 'failed', last_error = ?
|
||||
WHERE message_id = ?`,
|
||||
errMsg, messageID,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
// RequeueFailed re-queues failed items that haven't exceeded max attempts.
|
||||
func (s *EmbeddingStore) RequeueFailed(ctx context.Context, maxAttempts int) (int64, error) {
|
||||
result, err := s.db.ExecContext(ctx,
|
||||
`UPDATE embedding_queue SET status = 'pending'
|
||||
WHERE status = 'failed' AND attempts < ?`,
|
||||
maxAttempts,
|
||||
)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return result.RowsAffected()
|
||||
}
|
||||
|
||||
// ResetStale resets items stuck in "processing" for too long.
|
||||
func (s *EmbeddingStore) ResetStale(ctx context.Context, olderThan time.Duration) (int64, error) {
|
||||
cutoff := time.Now().Add(-olderThan)
|
||||
result, err := s.db.ExecContext(ctx,
|
||||
`UPDATE embedding_queue SET status = 'pending'
|
||||
WHERE status = 'processing' AND created_at < ?`,
|
||||
cutoff,
|
||||
)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return result.RowsAffected()
|
||||
}
|
||||
|
||||
// PendingCount returns the number of pending items.
|
||||
func (s *EmbeddingStore) PendingCount(ctx context.Context) (int64, error) {
|
||||
var count int64
|
||||
err := s.db.QueryRowContext(ctx,
|
||||
`SELECT COUNT(*) FROM embedding_queue WHERE status IN ('pending', 'processing')`,
|
||||
).Scan(&count)
|
||||
return count, err
|
||||
}
|
||||
|
||||
// EnqueueAllMessages enqueues all messages that don't have embeddings yet.
|
||||
func (s *EmbeddingStore) EnqueueAllMessages(ctx context.Context) (int64, error) {
|
||||
result, err := s.db.ExecContext(ctx,
|
||||
`INSERT OR IGNORE INTO embedding_queue (message_id, status, created_at)
|
||||
SELECT m.id, 'pending', CURRENT_TIMESTAMP
|
||||
FROM messages m
|
||||
LEFT JOIN embeddings e ON e.message_id = m.id
|
||||
LEFT JOIN embedding_queue q ON q.message_id = m.id
|
||||
WHERE e.message_id IS NULL AND q.message_id IS NULL
|
||||
AND TRIM(m.body) != ''`,
|
||||
)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("enqueue all messages: %w", err)
|
||||
}
|
||||
return result.RowsAffected()
|
||||
}
|
||||
|
||||
// GetMessageBody retrieves the body of a message by ID.
|
||||
func (s *EmbeddingStore) GetMessageBody(ctx context.Context, messageID int64) (string, error) {
|
||||
var body string
|
||||
err := s.db.QueryRowContext(ctx,
|
||||
`SELECT body FROM messages WHERE id = ?`, messageID,
|
||||
).Scan(&body)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("get message body %d: %w", messageID, err)
|
||||
}
|
||||
return body, nil
|
||||
}
|
||||
|
||||
// ClearQueue deletes all items from the embedding queue.
|
||||
func (s *EmbeddingStore) ClearQueue(ctx context.Context) error {
|
||||
_, err := s.db.ExecContext(ctx, `DELETE FROM embedding_queue`)
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
-- Semantic search: embeddings tracking and queue
|
||||
-- Note: actual vectors are stored in the HNSW index file on disk.
|
||||
-- This table tracks which messages have been embedded and by which provider.
|
||||
|
||||
CREATE TABLE IF NOT EXISTS embeddings (
|
||||
message_id INTEGER PRIMARY KEY REFERENCES messages(id) ON DELETE CASCADE,
|
||||
provider TEXT NOT NULL,
|
||||
model TEXT NOT NULL,
|
||||
dimensions INTEGER NOT NULL,
|
||||
embedded_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS embedding_queue (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
message_id INTEGER NOT NULL REFERENCES messages(id) ON DELETE CASCADE,
|
||||
status TEXT NOT NULL DEFAULT 'pending' CHECK (status IN ('pending', 'processing', 'completed', 'failed')),
|
||||
attempts INTEGER NOT NULL DEFAULT 0,
|
||||
last_error TEXT,
|
||||
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
completed_at TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_embedding_queue_message ON embedding_queue(message_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_embedding_queue_status ON embedding_queue(status);
|
||||
|
||||
INSERT INTO schema_migrations (version) VALUES (5);
|
||||
Reference in New Issue
Block a user