feat: hybrid search (RRF fusion) + minimum similarity threshold
- Auto mode now runs both semantic and fulltext searches, merging results using Reciprocal Rank Fusion (RRF, k=60) for best of both - New min_similarity parameter (default 0.25) filters semantic noise - Results that match both sources are marked as "hybrid" match_type - New getFloat bridge helper for MCP min_similarity parameter Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.6
parent
ed33da1093
commit
53c8ba5bb0
+29
-5
@@ -285,11 +285,12 @@ func (b *ServiceBridge) callSearchMessages(ctx context.Context, args map[string]
|
||||
searchMode := getString(args, "search_mode", "auto")
|
||||
|
||||
opts := search.SearchOptions{
|
||||
Query: query,
|
||||
Mode: searchMode,
|
||||
Limit: getInt(args, "limit", 10),
|
||||
FromAgent: getString(args, "from_agent", ""),
|
||||
MinPriority: getInt(args, "min_priority", 0),
|
||||
Query: query,
|
||||
Mode: searchMode,
|
||||
Limit: getInt(args, "limit", 10),
|
||||
FromAgent: getString(args, "from_agent", ""),
|
||||
MinPriority: getInt(args, "min_priority", 0),
|
||||
MinSimilarity: getFloat(args, "min_similarity", 0),
|
||||
}
|
||||
|
||||
resp, err := b.searchService.Search(ctx, b.agentName, opts)
|
||||
@@ -1301,6 +1302,29 @@ func getInt(args map[string]any, key string, defaultVal int) int {
|
||||
return defaultVal
|
||||
}
|
||||
|
||||
// getFloat extracts a float64 value from args with a default.
|
||||
func getFloat(args map[string]any, key string, defaultVal float64) float64 {
|
||||
v, ok := args[key]
|
||||
if !ok {
|
||||
return defaultVal
|
||||
}
|
||||
switch n := v.(type) {
|
||||
case float64:
|
||||
return n
|
||||
case int:
|
||||
return float64(n)
|
||||
case int64:
|
||||
return float64(n)
|
||||
case json.Number:
|
||||
f, err := n.Float64()
|
||||
if err != nil {
|
||||
return defaultVal
|
||||
}
|
||||
return f
|
||||
}
|
||||
return defaultVal
|
||||
}
|
||||
|
||||
// getBool extracts a bool value from args with a default.
|
||||
func getBool(args map[string]any, key string, defaultVal bool) bool {
|
||||
v, ok := args[key]
|
||||
|
||||
+128
-10
@@ -22,16 +22,23 @@ const (
|
||||
|
||||
// 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
|
||||
Query string
|
||||
Mode string // "auto", "semantic", "fulltext"
|
||||
Limit int
|
||||
ChannelID *int64
|
||||
FromAgent string
|
||||
MinPriority int
|
||||
After *time.Time
|
||||
Before *time.Time
|
||||
MinSimilarity float64 // minimum semantic similarity threshold (default 0.25)
|
||||
}
|
||||
|
||||
// DefaultMinSimilarity is the noise floor for semantic results.
|
||||
const DefaultMinSimilarity = 0.25
|
||||
|
||||
// rrfK is the Reciprocal Rank Fusion constant.
|
||||
const rrfK = 60
|
||||
|
||||
// SearchResult represents a single search result.
|
||||
type SearchResult struct {
|
||||
Message *messaging.Message `json:"message"`
|
||||
@@ -91,8 +98,8 @@ func (s *Service) Search(ctx context.Context, agentName string, opts SearchOptio
|
||||
// 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)
|
||||
if s.provider != nil && s.index != nil && s.index.Len() > 0 && opts.Query != "" {
|
||||
return s.hybridSearch(ctx, agentName, opts, limit)
|
||||
}
|
||||
return s.fulltextSearch(ctx, agentName, opts, limit)
|
||||
|
||||
@@ -177,6 +184,15 @@ func (s *Service) semanticSearch(ctx context.Context, agentName string, opts Sea
|
||||
similarity = 0
|
||||
}
|
||||
|
||||
// Apply minimum similarity threshold
|
||||
minSim := opts.MinSimilarity
|
||||
if minSim <= 0 {
|
||||
minSim = DefaultMinSimilarity
|
||||
}
|
||||
if similarity < minSim {
|
||||
continue
|
||||
}
|
||||
|
||||
searchResults = append(searchResults, &SearchResult{
|
||||
Message: msg,
|
||||
SimilarityScore: similarity,
|
||||
@@ -229,6 +245,108 @@ func (s *Service) fulltextSearch(ctx context.Context, agentName string, opts Sea
|
||||
}, nil
|
||||
}
|
||||
|
||||
// hybridSearch runs both semantic and fulltext searches, then merges results using
|
||||
// Reciprocal Rank Fusion (RRF). This gives the best of both worlds: semantic
|
||||
// understanding for conceptual queries and exact-match precision for keyword queries.
|
||||
func (s *Service) hybridSearch(ctx context.Context, agentName string, opts SearchOptions, limit int) (*SearchResponse, error) {
|
||||
// Run both searches, collecting errors but not failing if one source works.
|
||||
semResp, semErr := s.semanticSearch(ctx, agentName, opts, limit)
|
||||
ftResp, ftErr := s.fulltextSearch(ctx, agentName, opts, limit)
|
||||
|
||||
if semErr != nil && ftErr != nil {
|
||||
return nil, fmt.Errorf("hybrid search: semantic: %w; fulltext: %w", semErr, ftErr)
|
||||
}
|
||||
|
||||
var semResults, ftResults []*SearchResult
|
||||
if semErr == nil && semResp != nil {
|
||||
semResults = semResp.Results
|
||||
}
|
||||
if ftErr == nil && ftResp != nil {
|
||||
ftResults = ftResp.Results
|
||||
}
|
||||
|
||||
// If only one source returned results, use that directly
|
||||
if len(semResults) == 0 {
|
||||
if ftResp != nil {
|
||||
ftResp.SearchMode = ModeFulltext
|
||||
return ftResp, nil
|
||||
}
|
||||
return &SearchResponse{SearchMode: ModeAuto}, nil
|
||||
}
|
||||
if len(ftResults) == 0 {
|
||||
if semResp != nil {
|
||||
semResp.SearchMode = ModeAuto
|
||||
return semResp, nil
|
||||
}
|
||||
return &SearchResponse{SearchMode: ModeAuto}, nil
|
||||
}
|
||||
|
||||
merged := mergeRRF(semResults, ftResults, limit)
|
||||
|
||||
return &SearchResponse{
|
||||
Results: merged,
|
||||
SearchMode: ModeAuto,
|
||||
TotalResults: len(merged),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// mergeRRF merges two ranked result lists using Reciprocal Rank Fusion.
|
||||
// RRF score for a document d across rankings R1..Rn: sum(1 / (k + rank_i(d))).
|
||||
// This is rank-based and robust to score-scale differences between semantic and fulltext.
|
||||
func mergeRRF(semantic, fulltext []*SearchResult, limit int) []*SearchResult {
|
||||
type scored struct {
|
||||
result *SearchResult
|
||||
score float64
|
||||
}
|
||||
|
||||
// Map message ID -> accumulated RRF score + best result
|
||||
byID := make(map[int64]*scored)
|
||||
|
||||
for rank, r := range semantic {
|
||||
id := r.Message.ID
|
||||
rrfScore := 1.0 / (float64(rrfK) + float64(rank+1))
|
||||
if existing, ok := byID[id]; ok {
|
||||
existing.score += rrfScore
|
||||
} else {
|
||||
byID[id] = &scored{result: r, score: rrfScore}
|
||||
}
|
||||
}
|
||||
|
||||
for rank, r := range fulltext {
|
||||
id := r.Message.ID
|
||||
rrfScore := 1.0 / (float64(rrfK) + float64(rank+1))
|
||||
if existing, ok := byID[id]; ok {
|
||||
existing.score += rrfScore
|
||||
// If both matched, mark as hybrid
|
||||
existing.result.MatchType = "hybrid"
|
||||
} else {
|
||||
byID[id] = &scored{result: r, score: rrfScore}
|
||||
}
|
||||
}
|
||||
|
||||
// Sort by RRF score descending
|
||||
sorted := make([]*scored, 0, len(byID))
|
||||
for _, s := range byID {
|
||||
sorted = append(sorted, s)
|
||||
}
|
||||
// Simple insertion sort — result set is small (<200)
|
||||
for i := 1; i < len(sorted); i++ {
|
||||
for j := i; j > 0 && sorted[j].score > sorted[j-1].score; j-- {
|
||||
sorted[j], sorted[j-1] = sorted[j-1], sorted[j]
|
||||
}
|
||||
}
|
||||
|
||||
if len(sorted) > limit {
|
||||
sorted = sorted[:limit]
|
||||
}
|
||||
|
||||
results := make([]*SearchResult, len(sorted))
|
||||
for i, s := range sorted {
|
||||
results[i] = s.result
|
||||
}
|
||||
return results
|
||||
}
|
||||
|
||||
// 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
|
||||
|
||||
@@ -293,7 +293,7 @@ func TestService_SemanticSearch(t *testing.T) {
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("auto mode uses semantic when available", func(t *testing.T) {
|
||||
t.Run("auto mode uses hybrid when semantic available", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "searcher", SearchOptions{
|
||||
Query: "staging deployment",
|
||||
Mode: ModeAuto,
|
||||
@@ -302,8 +302,8 @@ func TestService_SemanticSearch(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
if resp.SearchMode != ModeSemantic {
|
||||
t.Errorf("auto search_mode = %q, want %q", resp.SearchMode, ModeSemantic)
|
||||
if resp.SearchMode != ModeAuto {
|
||||
t.Errorf("auto search_mode = %q, want %q", resp.SearchMode, ModeAuto)
|
||||
}
|
||||
})
|
||||
|
||||
@@ -366,6 +366,225 @@ func TestService_Filters(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestService_HybridSearch(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)
|
||||
|
||||
mockProvider := embedding.NewMockProvider(3)
|
||||
idx := NewMemoryVectorIndex()
|
||||
|
||||
// Send messages — some match fulltext, some match semantically, some both
|
||||
msg1, _ := msgService.SendMessage(ctx, "sender", "searcher", "deployment failure in staging environment", messaging.SendOptions{})
|
||||
msg2, _ := msgService.SendMessage(ctx, "sender", "searcher", "cat pictures are cute and funny", messaging.SendOptions{})
|
||||
msg3, _ := msgService.SendMessage(ctx, "sender", "searcher", "staging server crashed unexpectedly", messaging.SendOptions{})
|
||||
|
||||
// msg1 & msg3 are semantically similar (deployment topic), msg2 is different
|
||||
idx.AddVector(msg1.ID, []float32{0.95, 0.1, 0.0})
|
||||
idx.AddVector(msg2.ID, []float32{0.0, 0.0, 1.0})
|
||||
idx.AddVector(msg3.ID, []float32{0.90, 0.15, 0.0})
|
||||
|
||||
mockProvider.SetEmbedFunc(func(ctx context.Context, text string) ([]float32, error) {
|
||||
return []float32{0.95, 0.1, 0.0}, nil
|
||||
})
|
||||
|
||||
svc := NewService(db, mockProvider, idx, msgService)
|
||||
|
||||
t.Run("auto mode uses hybrid when semantic available", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "searcher", SearchOptions{
|
||||
Query: "staging",
|
||||
Mode: ModeAuto,
|
||||
Limit: 10,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
if resp.SearchMode != ModeAuto {
|
||||
t.Errorf("search_mode = %q, want %q", resp.SearchMode, ModeAuto)
|
||||
}
|
||||
// Should have results from both semantic and fulltext
|
||||
if len(resp.Results) == 0 {
|
||||
t.Fatal("expected results from hybrid search")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("hybrid merges unique results from both sources", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "searcher", SearchOptions{
|
||||
Query: "staging",
|
||||
Mode: ModeAuto,
|
||||
Limit: 10,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
|
||||
// Messages that appear in both semantic and fulltext should be marked "hybrid"
|
||||
hasHybrid := false
|
||||
for _, r := range resp.Results {
|
||||
if r.MatchType == "hybrid" {
|
||||
hasHybrid = true
|
||||
break
|
||||
}
|
||||
}
|
||||
// msg1 and msg3 contain "staging" (fulltext hit) AND are deployment-related (semantic hit)
|
||||
if !hasHybrid {
|
||||
t.Error("expected at least one hybrid match type from overlapping results")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("auto falls back to fulltext for empty query", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "searcher", SearchOptions{
|
||||
Query: "",
|
||||
Mode: ModeAuto,
|
||||
Limit: 10,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
if resp.SearchMode != ModeFulltext {
|
||||
t.Errorf("search_mode = %q, want %q for empty query", resp.SearchMode, ModeFulltext)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestService_MinSimilarityThreshold(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)
|
||||
|
||||
mockProvider := embedding.NewMockProvider(3)
|
||||
idx := NewMemoryVectorIndex()
|
||||
|
||||
// Create messages with varying similarity to query
|
||||
msgHigh, _ := msgService.SendMessage(ctx, "sender", "searcher", "high relevance deployment", messaging.SendOptions{})
|
||||
msgLow, _ := msgService.SendMessage(ctx, "sender", "searcher", "completely unrelated topic", messaging.SendOptions{})
|
||||
|
||||
// msgHigh is very similar (close vector), msgLow is orthogonal (low similarity)
|
||||
idx.AddVector(msgHigh.ID, []float32{0.95, 0.1, 0.0})
|
||||
idx.AddVector(msgLow.ID, []float32{0.1, 0.1, 0.95}) // very different direction
|
||||
|
||||
mockProvider.SetEmbedFunc(func(ctx context.Context, text string) ([]float32, error) {
|
||||
return []float32{0.95, 0.1, 0.0}, nil
|
||||
})
|
||||
|
||||
svc := NewService(db, mockProvider, idx, msgService)
|
||||
|
||||
t.Run("default threshold filters low-similarity noise", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "searcher", SearchOptions{
|
||||
Query: "deployment topic",
|
||||
Mode: ModeSemantic,
|
||||
Limit: 10,
|
||||
// MinSimilarity defaults to 0.25
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
for _, r := range resp.Results {
|
||||
if r.SimilarityScore < DefaultMinSimilarity {
|
||||
t.Errorf("result with similarity %f below default threshold %f",
|
||||
r.SimilarityScore, DefaultMinSimilarity)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("custom high threshold filters more results", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "searcher", SearchOptions{
|
||||
Query: "deployment topic",
|
||||
Mode: ModeSemantic,
|
||||
Limit: 10,
|
||||
MinSimilarity: 0.80,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
// Only the high-similarity message should pass
|
||||
for _, r := range resp.Results {
|
||||
if r.SimilarityScore < 0.80 {
|
||||
t.Errorf("result with similarity %f below custom threshold 0.80",
|
||||
r.SimilarityScore)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("very low threshold returns all results", func(t *testing.T) {
|
||||
resp, err := svc.Search(ctx, "searcher", SearchOptions{
|
||||
Query: "deployment topic",
|
||||
Mode: ModeSemantic,
|
||||
Limit: 10,
|
||||
MinSimilarity: 0.01,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Search: %v", err)
|
||||
}
|
||||
if len(resp.Results) < 2 {
|
||||
t.Errorf("expected both messages with low threshold, got %d", len(resp.Results))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestMergeRRF(t *testing.T) {
|
||||
mkResult := func(id int64, matchType string) *SearchResult {
|
||||
return &SearchResult{
|
||||
Message: &messaging.Message{ID: id},
|
||||
MatchType: matchType,
|
||||
}
|
||||
}
|
||||
|
||||
t.Run("deduplicates and boosts overlapping results", func(t *testing.T) {
|
||||
sem := []*SearchResult{mkResult(1, ModeSemantic), mkResult(2, ModeSemantic), mkResult(3, ModeSemantic)}
|
||||
ft := []*SearchResult{mkResult(2, ModeFulltext), mkResult(4, ModeFulltext), mkResult(1, ModeFulltext)}
|
||||
|
||||
merged := mergeRRF(sem, ft, 10)
|
||||
|
||||
// IDs 1 and 2 appear in both — should be ranked higher
|
||||
if len(merged) != 4 {
|
||||
t.Fatalf("expected 4 merged results, got %d", len(merged))
|
||||
}
|
||||
// Top results should be the overlapping ones (higher RRF score)
|
||||
topIDs := map[int64]bool{merged[0].Message.ID: true, merged[1].Message.ID: true}
|
||||
if !topIDs[1] || !topIDs[2] {
|
||||
t.Errorf("expected IDs 1 and 2 at top, got %d and %d", merged[0].Message.ID, merged[1].Message.ID)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("respects limit", func(t *testing.T) {
|
||||
sem := []*SearchResult{mkResult(1, ModeSemantic), mkResult(2, ModeSemantic)}
|
||||
ft := []*SearchResult{mkResult(3, ModeFulltext), mkResult(4, ModeFulltext)}
|
||||
|
||||
merged := mergeRRF(sem, ft, 2)
|
||||
if len(merged) != 2 {
|
||||
t.Fatalf("expected 2 results with limit=2, got %d", len(merged))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("marks overlapping results as hybrid", func(t *testing.T) {
|
||||
sem := []*SearchResult{mkResult(1, ModeSemantic)}
|
||||
ft := []*SearchResult{mkResult(1, ModeFulltext)}
|
||||
|
||||
merged := mergeRRF(sem, ft, 10)
|
||||
if len(merged) != 1 {
|
||||
t.Fatalf("expected 1 merged result, got %d", len(merged))
|
||||
}
|
||||
if merged[0].MatchType != "hybrid" {
|
||||
t.Errorf("match_type = %q, want %q", merged[0].MatchType, "hybrid")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestService_HasSemanticSearch(t *testing.T) {
|
||||
t.Run("without provider", func(t *testing.T) {
|
||||
svc := NewService(nil, nil, nil, nil)
|
||||
|
||||
Reference in New Issue
Block a user