From 53c8ba5bb0c7c8bae185fabf28fc07067db1b3c2 Mon Sep 17 00:00:00 2001 From: Algis Dumbris Date: Wed, 1 Apr 2026 16:56:26 +0300 Subject: [PATCH] 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) --- internal/mcp/bridge.go | 34 ++++- internal/search/service.go | 138 ++++++++++++++++++-- internal/search/service_test.go | 225 +++++++++++++++++++++++++++++++- 3 files changed, 379 insertions(+), 18 deletions(-) diff --git a/internal/mcp/bridge.go b/internal/mcp/bridge.go index 2402575..93ed2a7 100644 --- a/internal/mcp/bridge.go +++ b/internal/mcp/bridge.go @@ -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] diff --git a/internal/search/service.go b/internal/search/service.go index 6c8165d..93f5578 100644 --- a/internal/search/service.go +++ b/internal/search/service.go @@ -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 diff --git a/internal/search/service_test.go b/internal/search/service_test.go index 5eff699..4089f1c 100644 --- a/internal/search/service_test.go +++ b/internal/search/service_test.go @@ -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)