package cmd import ( "encoding/json" "fmt" "os" "strings" "time" "github.com/kb-search/kb/internal/api" "github.com/kb-search/kb/internal/output" "github.com/spf13/cobra" ) var benchCmd = &cobra.Command{ Use: "bench ", Short: "Benchmark search quality against a query fixture", Long: `Run a fixture of queries with known-relevant documents against each search backend (fts, vec, hybrid, rerank) and report precision@k, recall and MRR per backend. Use this to baseline search quality before ranking changes and to measure their effect. See docs/bench-example.json for the fixture format.`, Args: cobra.ExactArgs(1), RunE: runBench, } func init() { benchCmd.Flags().Int("top", 0, "override result count per query (k)") benchCmd.Flags().String("backends", "fts,vec,hybrid,rerank", "comma-separated backends to run") rootCmd.AddCommand(benchCmd) } type benchSelector struct { DocumentID int64 `json:"document_id"` SourcePath string `json:"source_path"` TitleContains string `json:"title_contains"` } type benchQuery struct { ID string `json:"id"` Query string `json:"query"` Tags []string `json:"tags"` DocType string `json:"doc_type"` Top int `json:"top"` Relevant []benchSelector `json:"relevant"` Notes string `json:"notes"` } type benchFixture struct { Description string `json:"description"` Top int `json:"top"` Queries []benchQuery `json:"queries"` } type benchDoc struct { DocumentID int64 Title string SourcePath string } type queryResult struct { ID string `json:"id"` Precision float64 `json:"precision"` Recall float64 `json:"recall"` MRR float64 `json:"mrr"` LatencyMS float64 `json:"latency_ms"` Returned int `json:"returned"` } type backendResult struct { Precision float64 `json:"precision"` Recall float64 `json:"recall"` MRR float64 `json:"mrr"` AvgLatencyMS float64 `json:"avg_latency_ms"` Queries []queryResult `json:"queries"` } func loadFixture(path string) (*benchFixture, error) { data, err := os.ReadFile(path) if err != nil { return nil, fmt.Errorf("cannot read fixture: %w", err) } var fx benchFixture if err := json.Unmarshal(data, &fx); err != nil { return nil, fmt.Errorf("invalid fixture JSON: %w", err) } if len(fx.Queries) == 0 { return nil, fmt.Errorf("fixture has no queries") } for i, q := range fx.Queries { if q.Query == "" { return nil, fmt.Errorf("query %d: missing query text", i) } if len(q.Relevant) == 0 { return nil, fmt.Errorf("query %q: no relevant selectors", q.Query) } for j, sel := range q.Relevant { set := 0 if sel.DocumentID != 0 { set++ } if sel.SourcePath != "" { set++ } if sel.TitleContains != "" { set++ } if set != 1 { return nil, fmt.Errorf("query %q selector %d: exactly one of document_id, source_path, title_contains must be set", q.Query, j) } } } return &fx, nil } func (s benchSelector) matches(doc benchDoc) bool { switch { case s.DocumentID != 0: return doc.DocumentID == s.DocumentID case s.SourcePath != "": return doc.SourcePath == s.SourcePath case s.TitleContains != "": return strings.Contains(strings.ToLower(doc.Title), strings.ToLower(s.TitleContains)) } return false } // dedupeByDocument collapses ranked chunks into ranked documents, keeping the // best (first) position for each document. func dedupeByDocument(docs []benchDoc) []benchDoc { seen := map[int64]bool{} var out []benchDoc for _, d := range docs { if seen[d.DocumentID] { continue } seen[d.DocumentID] = true out = append(out, d) } return out } // scoreQuery computes document-level precision, recall and MRR for one // query's ranked document list against the relevant-document selectors. func scoreQuery(ranked []benchDoc, relevant []benchSelector) (precision, recall, mrr float64) { if len(ranked) == 0 { return 0, 0, 0 } matchedDocs := 0 firstMatch := 0 selectorHit := make([]bool, len(relevant)) for i, doc := range ranked { docMatched := false for j, sel := range relevant { if sel.matches(doc) { docMatched = true selectorHit[j] = true } } if docMatched { matchedDocs++ if firstMatch == 0 { firstMatch = i + 1 } } } selectorsMatched := 0 for _, hit := range selectorHit { if hit { selectorsMatched++ } } precision = float64(matchedDocs) / float64(len(ranked)) recall = float64(selectorsMatched) / float64(len(relevant)) if firstMatch > 0 { mrr = 1.0 / float64(firstMatch) } return precision, recall, mrr } // rerankAvailable probes engine status for a loaded reranker. func rerankAvailable(client *api.Client) bool { resp, err := client.Get("/api/v1/status") if err != nil { return false } var status struct { Rerank struct { Enabled bool `json:"enabled"` Loaded bool `json:"loaded"` } `json:"rerank"` } if err := api.DecodeJSON(resp, &status); err != nil { return false } return status.Rerank.Enabled && status.Rerank.Loaded } func benchSearch(client *api.Client, q benchQuery, backend string, top int) ([]benchDoc, float64, error) { body := map[string]interface{}{ "query": q.Query, "top": top, } if len(q.Tags) > 0 { body["tags"] = q.Tags } if q.DocType != "" { body["doc_type"] = q.DocType } switch backend { case "fts": body["fts_only"] = true case "vec": body["vec_only"] = true case "hybrid": body["rerank"] = false case "rerank": body["rerank"] = true } start := time.Now() resp, err := client.Post("/api/v1/search", body) if err != nil { return nil, 0, err } if err := api.CheckError(resp); err != nil { return nil, 0, err } var result struct { Results []struct { DocumentID int64 `json:"document_id"` Title string `json:"title"` SourcePath string `json:"source_path"` } `json:"results"` } if err := api.DecodeJSON(resp, &result); err != nil { return nil, 0, err } latency := float64(time.Since(start).Microseconds()) / 1000.0 var docs []benchDoc for _, r := range result.Results { docs = append(docs, benchDoc{DocumentID: r.DocumentID, Title: r.Title, SourcePath: r.SourcePath}) } return dedupeByDocument(docs), latency, nil } func runBench(cmd *cobra.Command, args []string) error { topFlag, _ := cmd.Flags().GetInt("top") backendsFlag, _ := cmd.Flags().GetString("backends") fx, err := loadFixture(args[0]) if err != nil { return err } var backends []string for _, b := range strings.Split(backendsFlag, ",") { b = strings.TrimSpace(b) if b == "" { continue } switch b { case "fts", "vec", "hybrid", "rerank": backends = append(backends, b) default: return fmt.Errorf("unknown backend %q (valid: fts, vec, hybrid, rerank)", b) } } if len(backends) == 0 { return fmt.Errorf("no backends selected") } client := api.NewClient() rerankSkipped := false if contains(backends, "rerank") && !rerankAvailable(client) { backends = remove(backends, "rerank") rerankSkipped = true if len(backends) == 0 { return fmt.Errorf("reranking is not available on this engine (requires engine with reranker enabled)") } } results := map[string]*backendResult{} for _, backend := range backends { br := &backendResult{} for _, q := range fx.Queries { top := 10 if fx.Top > 0 { top = fx.Top } if q.Top > 0 { top = q.Top } if topFlag > 0 { top = topFlag } ranked, latency, err := benchSearch(client, q, backend, top) if err != nil { return fmt.Errorf("backend %s, query %q: %w", backend, q.Query, err) } p, r, m := scoreQuery(ranked, q.Relevant) id := q.ID if id == "" { id = q.Query } br.Queries = append(br.Queries, queryResult{ ID: id, Precision: p, Recall: r, MRR: m, LatencyMS: latency, Returned: len(ranked), }) } n := float64(len(br.Queries)) for _, qr := range br.Queries { br.Precision += qr.Precision / n br.Recall += qr.Recall / n br.MRR += qr.MRR / n br.AvgLatencyMS += qr.LatencyMS / n } results[backend] = br } if output.IsJSON() { output.PrintJSON(map[string]interface{}{ "description": fx.Description, "query_count": len(fx.Queries), "backends": results, "rerank_skipped": rerankSkipped, }) return nil } if fx.Description != "" { fmt.Printf("%s (%d queries)\n\n", fx.Description, len(fx.Queries)) } headers := []string{"Backend", "Precision", "Recall", "MRR", "Avg ms"} var rows [][]string for _, backend := range backends { br := results[backend] rows = append(rows, []string{ backend, fmt.Sprintf("%.3f", br.Precision), fmt.Sprintf("%.3f", br.Recall), fmt.Sprintf("%.3f", br.MRR), fmt.Sprintf("%.0f", br.AvgLatencyMS), }) } output.PrintTable(headers, rows) if rerankSkipped { fmt.Println("\nrerank: n/a (reranker not enabled on this engine)") } return nil } func contains(ss []string, s string) bool { for _, v := range ss { if v == s { return true } } return false } func remove(ss []string, s string) []string { var out []string for _, v := range ss { if v != s { out = append(out, v) } } return out }