From 509c98de85d9cb0a7e334c8ba170e4db3e116a41 Mon Sep 17 00:00:00 2001 From: shenlan Date: Wed, 13 Aug 2025 11:10:04 +0800 Subject: [PATCH] Add ragbench evaluation tool and workflow --- .github/workflows/ragbench.yml | 36 +++++ ragbench/cmd/ragbench/main.go | 46 ++++++ ragbench/go.mod | 5 + ragbench/go.sum | 4 + ragbench/internal/bench/bench.go | 249 +++++++++++++++++++++++++++++ ragbench/internal/bench/metrics.go | 130 +++++++++++++++ ragbench/queries.example.yaml | 9 ++ 7 files changed, 479 insertions(+) create mode 100644 .github/workflows/ragbench.yml create mode 100644 ragbench/cmd/ragbench/main.go create mode 100644 ragbench/go.mod create mode 100644 ragbench/go.sum create mode 100644 ragbench/internal/bench/bench.go create mode 100644 ragbench/internal/bench/metrics.go create mode 100644 ragbench/queries.example.yaml diff --git a/.github/workflows/ragbench.yml b/.github/workflows/ragbench.yml new file mode 100644 index 0000000..e7c564f --- /dev/null +++ b/.github/workflows/ragbench.yml @@ -0,0 +1,36 @@ +name: RAG Benchmark +on: + pull_request: + push: + branches: [ main, 'release/**' ] + +jobs: + bench: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + + - name: Setup Go + uses: actions/setup-go@v5 + with: { go-version: '1.22' } + + - name: Run ragbench + env: + API_BASE: ${{ secrets.RAG_API_BASE }} + run: | + cd ragbench + if [ -n "${API_BASE}" ]; then API_FLAG="-api ${API_BASE}"; fi + go run ./cmd/ragbench ${API_FLAG} -in queries.yaml -out report.md + + - name: Upload report artifact + uses: actions/upload-artifact@v4 + with: + name: rag-benchmark-report + path: ragbench/report.md + if-no-files-found: error + + - name: Comment report to PR + if: github.event_name == 'pull_request' + uses: marocchino/sticky-pull-request-comment@v2 + with: + path: ragbench/report.md diff --git a/ragbench/cmd/ragbench/main.go b/ragbench/cmd/ragbench/main.go new file mode 100644 index 0000000..f7e1bf1 --- /dev/null +++ b/ragbench/cmd/ragbench/main.go @@ -0,0 +1,46 @@ +package main + +import ( + "flag" + "fmt" + "log" + "os" + + "ragbench/internal/bench" +) + +func main() { + in := flag.String("in", "queries.yaml", "path to queries.yaml") + api := flag.String("api", "", "override API base (optional, otherwise from yaml)") + k := flag.Int("k", 0, "override K (optional, otherwise from yaml)") + out := flag.String("out", "report.md", "output markdown report") + parallel := flag.Int("parallel", 16, "concurrency") + timeoutMS := flag.Int("timeout_ms", 8000, "per request timeout ms") + flag.Parse() + + cfg, err := bench.LoadConfig(*in) + if err != nil { + log.Fatal(err) + } + if *api != "" { + cfg.APIBase = *api + } + if *k > 0 { + cfg.K = *k + } + + res := bench.RunSuite(cfg, bench.Options{ + Parallel: *parallel, + TimeoutMS: *timeoutMS, + }) + + md := bench.RenderMarkdown(cfg, res) + if err := os.WriteFile(*out, []byte(md), 0644); err != nil { + log.Fatal(err) + } + + fmt.Println("✅ RAG benchmark done. Report:", *out) + fmt.Printf("Hit@%d: %.2f%% Recall@%d: %.2f%% MRR: %.4f nDCG@%d: %.4f P95: %dms Errors: %d\n", + cfg.K, 100*res.Metrics.HitAtK, cfg.K, 100*res.Metrics.RecallAtK, + res.Metrics.MRR, cfg.K, res.Metrics.NDCGAtK, res.Latency.P95, res.Errors) +} diff --git a/ragbench/go.mod b/ragbench/go.mod new file mode 100644 index 0000000..dd5bc55 --- /dev/null +++ b/ragbench/go.mod @@ -0,0 +1,5 @@ +module ragbench + +go 1.22 + +require gopkg.in/yaml.v3 v3.0.1 diff --git a/ragbench/go.sum b/ragbench/go.sum new file mode 100644 index 0000000..a62c313 --- /dev/null +++ b/ragbench/go.sum @@ -0,0 +1,4 @@ +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/ragbench/internal/bench/bench.go b/ragbench/internal/bench/bench.go new file mode 100644 index 0000000..199d2e6 --- /dev/null +++ b/ragbench/internal/bench/bench.go @@ -0,0 +1,249 @@ +package bench + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "os" + "sync" + "time" + + "gopkg.in/yaml.v3" +) + +type Config struct { + APIBase string `yaml:"api_base"` + K int `yaml:"k"` + Queries []QItem `yaml:"queries"` +} + +type QItem struct { + ID string `yaml:"id"` + Query string `yaml:"query"` + ExpectedDocIDs []string `yaml:"expected_doc_ids"` +} + +func LoadConfig(path string) (*Config, error) { + b, err := osReadFile(path) + if err != nil { + return nil, err + } + var c Config + if err := yaml.Unmarshal(b, &c); err != nil { + return nil, err + } + if c.K <= 0 { + c.K = 5 + } + return &c, nil +} + +// Abstraction of /api/rag/query response +type QueryRequest struct { + Question string `json:"question"` + K int `json:"k"` +} +type Hit struct { + ID string `json:"id"` + Score float64 `json:"score"` + Snippet string `json:"snippet,omitempty"` +} +type QueryResponse struct { + Answer string `json:"answer"` + Hits []Hit `json:"hits"` +} + +type Options struct { + Parallel int + TimeoutMS int +} + +type PerCase struct { + ID string + Query string + K int + Hits []Hit + Latency int64 // ms + Error error + Expected map[string]struct{} +} + +type SuiteResult struct { + Cases []PerCase + Metrics Metrics + Latency LatencyStats + Errors int +} + +func RunSuite(cfg *Config, opt Options) *SuiteResult { + if opt.Parallel <= 0 { + opt.Parallel = 16 + } + if opt.TimeoutMS <= 0 { + opt.TimeoutMS = 8000 + } + + wg := sync.WaitGroup{} + sem := make(chan struct{}, opt.Parallel) + out := make([]PerCase, len(cfg.Queries)) + + for i := range cfg.Queries { + wg.Add(1) + sem <- struct{}{} + go func(i int) { + defer wg.Done() + defer func() { <-sem }() + q := cfg.Queries[i] + out[i] = runOne(cfg.APIBase, q, cfg.K, opt.TimeoutMS) + }(i) + } + wg.Wait() + + m := CalcMetrics(out, cfg.K) + lat := calcLatency(out) + errs := 0 + for _, c := range out { + if c.Error != nil { + errs++ + } + } + + return &SuiteResult{Cases: out, Metrics: m, Latency: lat, Errors: errs} +} + +func runOne(apiBase string, q QItem, k, timeoutMS int) PerCase { + ctx, cancel := context.WithTimeout(context.Background(), time.Duration(timeoutMS)*time.Millisecond) + defer cancel() + + expected := make(map[string]struct{}, len(q.ExpectedDocIDs)) + for _, id := range q.ExpectedDocIDs { + expected[id] = struct{}{} + } + + reqBody, _ := json.Marshal(QueryRequest{Question: q.Query, K: k}) + url := fmt.Sprintf("%s/api/rag/query", apiBase) + + t0 := time.Now() + resp, err := httpDo(ctx, "POST", url, "application/json", bytes.NewReader(reqBody)) + lat := time.Since(t0).Milliseconds() + if err != nil { + return PerCase{ID: q.ID, Query: q.Query, K: k, Latency: lat, Error: err, Expected: expected} + } + defer resp.Body.Close() + b, _ := io.ReadAll(resp.Body) + if resp.StatusCode >= 300 { + return PerCase{ID: q.ID, Query: q.Query, K: k, Latency: lat, Error: fmt.Errorf("http %d: %s", resp.StatusCode, truncate(string(b), 200)), Expected: expected} + } + var qr QueryResponse + if err := json.Unmarshal(b, &qr); err != nil { + return PerCase{ID: q.ID, Query: q.Query, K: k, Latency: lat, Error: err, Expected: expected} + } + // normalize: trim to K and remove dups preserving order + hits := dedupTopK(qr.Hits, k) + + return PerCase{ID: q.ID, Query: q.Query, K: k, Hits: hits, Latency: lat, Expected: expected} +} + +func dedupTopK(h []Hit, k int) []Hit { + seen := map[string]struct{}{} + out := make([]Hit, 0, k) + for _, x := range h { + if x.ID == "" { + continue + } + if _, ok := seen[x.ID]; ok { + continue + } + seen[x.ID] = struct{}{} + out = append(out, x) + if len(out) >= k { + break + } + } + return out +} + +// tiny helpers (no external deps) +func osReadFile(p string) ([]byte, error) { return os.ReadFile(p) } +func httpDo(ctx context.Context, method, url, ctype string, body io.Reader) (*http.Response, error) { + req, _ := http.NewRequestWithContext(ctx, method, url, body) + if ctype != "" { + req.Header.Set("Content-Type", ctype) + } + return http.DefaultClient.Do(req) +} +func truncate(s string, n int) string { + if len(s) <= n { + return s + } + return s[:n] + "..." +} + +// Markdown rendering +func RenderMarkdown(cfg *Config, res *SuiteResult) string { + // summary + md := &bytes.Buffer{} + fmt.Fprintf(md, "# RAG Benchmark Report\n\n") + fmt.Fprintf(md, "- API: `%s`\n- K: `%d`\n- Cases: `%d`\n- Errors: `%d`\n\n", cfg.APIBase, cfg.K, len(res.Cases), res.Errors) + + fmt.Fprintf(md, "## Summary Metrics\n\n") + fmt.Fprintf(md, "| Metric | Value |\n|---|---|\n") + fmt.Fprintf(md, "| Hit@%d | %.2f%% |\n", cfg.K, 100*res.Metrics.HitAtK) + fmt.Fprintf(md, "| Recall@%d | %.2f%% |\n", cfg.K, 100*res.Metrics.RecallAtK) + fmt.Fprintf(md, "| MRR | %.4f |\n", res.Metrics.MRR) + fmt.Fprintf(md, "| nDCG@%d | %.4f |\n", cfg.K, res.Metrics.NDCGAtK) + fmt.Fprintf(md, "| P50 latency | %d ms |\n", res.Latency.P50) + fmt.Fprintf(md, "| P95 latency | %d ms |\n\n", res.Latency.P95) + + // failures table + fail := failedCases(res.Cases) + if len(fail) > 0 { + fmt.Fprintf(md, "## Failures (%d)\n\n", len(fail)) + fmt.Fprintf(md, "| ID | Error |\n|---|---|\n") + for _, f := range fail { + fmt.Fprintf(md, "| %s | %s |\n", f.ID, truncate(fmt.Sprintf("%v", f.Error), 180)) + } + fmt.Fprintln(md) + } + + // per-case top-K (optional; keep short) + fmt.Fprintf(md, "## Per-case (Top-%d IDs)\n\n", cfg.K) + fmt.Fprintf(md, "| ID | Hit@K | Latency(ms) | TopIDs |\n|---|---:|---:|---|\n") + for _, c := range res.Cases { + hit := caseHitAtK(c) + topIDs := make([]string, 0, len(c.Hits)) + for _, h := range c.Hits { + topIDs = append(topIDs, h.ID) + } + fmt.Fprintf(md, "| %s | %t | %d | %s |\n", c.ID, hit, c.Latency, backtick(join(topIDs, ", "))) + } + return md.String() +} + +func failedCases(cs []PerCase) []PerCase { + out := make([]PerCase, 0) + for _, c := range cs { + if c.Error != nil { + out = append(out, c) + } + } + return out +} +func backtick(s string) string { return "`" + s + "`" } +func join(a []string, sep string) string { + switch len(a) { + case 0: + return "" + case 1: + return a[0] + default: + b := a[0] + for i := 1; i < len(a); i++ { + b += sep + a[i] + } + return b + } +} diff --git a/ragbench/internal/bench/metrics.go b/ragbench/internal/bench/metrics.go new file mode 100644 index 0000000..4981c8c --- /dev/null +++ b/ragbench/internal/bench/metrics.go @@ -0,0 +1,130 @@ +package bench + +import ( + "math" + "sort" +) + +type Metrics struct { + HitAtK float64 + RecallAtK float64 + MRR float64 + NDCGAtK float64 +} + +type LatencyStats struct{ P50, P95 int } + +func CalcMetrics(cases []PerCase, k int) Metrics { + var hitN, total int + var rrSum float64 + var dcgSum, idcgSum float64 + var recallSum float64 + + for _, c := range cases { + if c.Error != nil { + continue + } + total++ + + // expected set + exp := c.Expected + if len(exp) == 0 { + continue + } + + // Hit@K + hit := caseHitAtK(c) + if hit { + hitN++ + } + + // Recall@K: (# 前K内命中标准片段) / (标准片段总数) + matched := 0 + for _, h := range c.Hits { + if _, ok := exp[h.ID]; ok { + matched++ + } + } + recall := float64(matched) / float64(len(exp)) + if recall > 1 { + recall = 1 + } + recallSum += recall + + // Reciprocal Rank + rank := firstHitRank(c) + if rank > 0 { + rrSum += 1.0 / float64(rank) + } + + // nDCG@K (二值相关性) + dcg := 0.0 + for i, h := range c.Hits { + if _, ok := exp[h.ID]; ok { + rel := 1.0 + dcg += (math.Pow(2, rel) - 1) / math.Log2(float64(i+2)) // i从0 + } + } + // IDCG:把rel=1的若干个放在最前面 + g := minInt(len(exp), k) + idcg := 0.0 + for i := 0; i < g; i++ { + idcg += (math.Pow(2, 1.0) - 1) / math.Log2(float64(i+2)) + } + if idcg == 0 { + idcg = 1 + } // 防零 + dcgSum += dcg + idcgSum += idcg + } + + if total == 0 { + return Metrics{} + } + return Metrics{ + HitAtK: float64(hitN) / float64(total), + RecallAtK: recallSum / float64(total), + MRR: rrSum / float64(total), + NDCGAtK: (dcgSum / idcgSum), + } +} + +func calcLatency(cases []PerCase) LatencyStats { + var xs []int + for _, c := range cases { + if c.Error == nil { + xs = append(xs, int(c.Latency)) + } + } + if len(xs) == 0 { + return LatencyStats{} + } + sort.Ints(xs) + p50 := xs[len(xs)*50/100] + p95 := xs[len(xs)*95/100] + return LatencyStats{P50: p50, P95: p95} +} + +func caseHitAtK(c PerCase) bool { + for _, h := range c.Hits { + if _, ok := c.Expected[h.ID]; ok { + return true + } + } + return false +} +func firstHitRank(c PerCase) int { + for i, h := range c.Hits { + if _, ok := c.Expected[h.ID]; ok { + return i + 1 + } + } + return 0 +} + +func minInt(a, b int) int { + if a < b { + return a + } + return b +} diff --git a/ragbench/queries.example.yaml b/ragbench/queries.example.yaml new file mode 100644 index 0000000..249415b --- /dev/null +++ b/ragbench/queries.example.yaml @@ -0,0 +1,9 @@ +api_base: "http://127.0.0.1:8080" +k: 10 +queries: + - id: Q001 + query: "What is Hybrid Search in RAG?" + expected_doc_ids: ["doc12#p3", "doc5#p7"] + - id: Q002 + query: "How to migrate embedding dimensions in pgvector?" + expected_doc_ids: ["doc20#p1"]