diff --git a/cmd/engine/main.go b/cmd/engine/main.go index 498f07d..dc126d6 100644 --- a/cmd/engine/main.go +++ b/cmd/engine/main.go @@ -21,6 +21,7 @@ import ( "github.com/hallelx2/llmgate" "github.com/hallelx2/llmgate/judge/typesafe" + "github.com/hallelx2/llmgate/middleware/limit" "github.com/hallelx2/llmgate/middleware/retry" "github.com/hallelx2/llmgate/pricing" "github.com/hallelx2/llmgate/provider/anthropic" @@ -127,7 +128,10 @@ func run() error { return fmt.Errorf("init llm: %w", err) } } - judge, err := buildJudge(cfg.LLM.Judge) + if llmClient != nil { + llmClient = limit.Client(newLimiter("llm", cfg.LLM.Concurrency, logger))(llmClient) + } + judge, err := buildJudge(cfg.LLM.Judge, newLimiter("judge", cfg.LLM.Concurrency, logger)) if err != nil { logger.Error("judge: config invalid", "err", err) os.Exit(1) @@ -416,7 +420,23 @@ func modelFor(c config.LLMConfig) string { // call; a Judge request that fails past them is handled by the TOC // builder, which keeps extraction's pages rather than degrading // silently (HAL-1369). -func buildJudge(c config.JudgeBlock) (llmgate.Judge, error) { +// newLimiter builds one adaptive limiter for a provider and logs every +// change it makes: a throttled run must be visible, never silent. +func newLimiter(name string, c config.ConcurrencyBlock, logger *slog.Logger) *limit.Limiter { + return limit.New(limit.Config{ + Initial: c.Initial, + Max: c.Max, + OnChange: func(e limit.Event) { + if e.Cause == "success" { + logger.Info("limiter: widened", "provider", name, "from", e.From, "to", e.To) + return + } + logger.Warn("limiter: narrowed", "provider", name, "cause", e.Cause, "from", e.From, "to", e.To, "paused_for", e.PausedFor, "err", e.Err) + }, + }) +} + +func buildJudge(c config.JudgeBlock, lim *limit.Limiter) (llmgate.Judge, error) { if c.TypeSafe.APIKey == "" { return nil, nil } @@ -428,7 +448,9 @@ func buildJudge(c config.JudgeBlock) (llmgate.Judge, error) { if err != nil { return nil, err } - return retry.NewJudge(retry.Config{MaxRetries: 3})(j), nil + // The limiter sits inside retry: each attempt takes a slot, and the + // failure that triggers a retry has already narrowed the limit. + return retry.NewJudge(retry.Config{MaxRetries: 3})(limit.Judge(lim)(j)), nil } func buildLLM(c config.LLMConfig) (llmgate.Client, error) { diff --git a/cmd/navbench/main.go b/cmd/navbench/main.go index c6cee7c..be38a86 100644 --- a/cmd/navbench/main.go +++ b/cmd/navbench/main.go @@ -19,9 +19,11 @@ import ( "path/filepath" "sort" "strings" + "sync" "time" "github.com/hallelx2/llmgate/judge/typesafe" + "github.com/hallelx2/llmgate/middleware/limit" "github.com/hallelx2/llmgate/middleware/retry" "github.com/hallelx2/vectorless-engine/pkg/ingest" @@ -71,7 +73,8 @@ func main() { out := flag.String("out", "", "JSONL of per-question outcomes") maxLeaves := flag.Int("leaves", 3, "sections read per question") maxPages := flag.Int("pages", 40, "pages judged per question") - limit := flag.Int("limit", 0, "stop after this many questions (0 = all)") + limitQ := flag.Int("limit", 0, "stop after this many questions (0 = all)") + parallel := flag.Int("parallel", 1, "questions in flight at once; the provider's adaptive limiter governs requests") flag.Parse() if *qPath == "" || *trees == "" || *pdfs == "" { fmt.Fprintln(os.Stderr, "usage: navbench -questions q.jsonl -trees dir -pdfs dir [-out o.jsonl]") @@ -84,11 +87,14 @@ func main() { fmt.Fprintln(os.Stderr, "judge:", err) os.Exit(1) } - nav := &retrieval.JudgeNavigator{Judge: retry.NewJudge(retry.Config{MaxRetries: 3})(tj), MaxLeaves: *maxLeaves, MaxPages: *maxPages} + lim := limit.New(limit.Config{Initial: 4, OnChange: func(e limit.Event) { + fmt.Fprintf(os.Stderr, " limiter %s %d -> %d %v\n", e.Cause, e.From, e.To, e.Err) + }}) + nav := &retrieval.JudgeNavigator{Judge: retry.NewJudge(retry.Config{MaxRetries: 3})(limit.Judge(lim)(tj)), MaxLeaves: *maxLeaves, MaxPages: *maxPages} qs := readQuestions(*qPath) - if *limit > 0 && len(qs) > *limit { - qs = qs[:*limit] + if *limitQ > 0 && len(qs) > *limitQ { + qs = qs[:*limitQ] } var of *os.File if *out != "" { @@ -98,60 +104,84 @@ func main() { pageCache := map[string][]ingest.PageText{} leafCache := map[string][]retrieval.NavLeaf{} - var results []outcome + // Documents are parsed once, up front and sequentially, so the + // parallel part is only Judge traffic. for _, q := range qs { - o := outcome{ID: q.ID, Doc: q.Doc, Question: q.Question, Gold: q.Evidence} - leaves, pages, err := load(q.Doc, *trees, *pdfs, leafCache, pageCache) - if err != nil { - o.Err = err.Error() - results = append(results, o) - report(o) - continue - } - byNum := map[int]string{} - for _, p := range pages { - byNum[p.PageNumber] = p.Text + if _, ok := leafCache[q.Doc]; !ok { + _, _, _ = load(q.Doc, *trees, *pdfs, leafCache, pageCache) } - loadPages := func(_ context.Context, l retrieval.NavLeaf) ([]retrieval.NavPage, error) { - var ps []retrieval.NavPage - for n := l.Start; n <= l.End; n++ { - if t, ok := byNum[n]; ok { - ps = append(ps, retrieval.NavPage{Number: n, Text: t}) + } + runStart := time.Now() + results := make([]outcome, len(qs)) + if *parallel < 1 { + *parallel = 1 + } + sem := make(chan struct{}, *parallel) + var wg sync.WaitGroup + var outMu sync.Mutex + for i, q := range qs { + wg.Add(1) + sem <- struct{}{} + go func(i int, q question) { + defer wg.Done() + defer func() { <-sem }() + o := outcome{ID: q.ID, Doc: q.Doc, Question: q.Question, Gold: q.Evidence} + leaves, pages, err := load(q.Doc, *trees, *pdfs, leafCache, pageCache) + if err != nil { + o.Err = err.Error() + results[i] = o + outMu.Lock() + report(o) + outMu.Unlock() + return + } + byNum := map[int]string{} + for _, p := range pages { + byNum[p.PageNumber] = p.Text + } + loadPages := func(_ context.Context, l retrieval.NavLeaf) ([]retrieval.NavPage, error) { + var ps []retrieval.NavPage + for n := l.Start; n <= l.End; n++ { + if t, ok := byNum[n]; ok { + ps = append(ps, retrieval.NavPage{Number: n, Text: t}) + } } + return ps, nil } - return ps, nil - } - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute) - start := time.Now() - res, err := nav.Navigate(ctx, q.Question, leaves, loadPages) - cancel() - o.Seconds = time.Since(start).Seconds() - if err != nil { - o.Err = err.Error() - results = append(results, o) + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute) + start := time.Now() + res, err := nav.Navigate(ctx, q.Question, leaves, loadPages) + cancel() + o.Seconds = time.Since(start).Seconds() + if err != nil { + o.Err = err.Error() + } else { + for _, l := range res.Selected { + o.Selected = append(o.Selected, l.Title) + o.SelectedPages = append(o.SelectedPages, [2]int{l.Start, l.End}) + } + for _, e := range res.Evidence { + o.Evidence = append(o.Evidence, e.Page.Number) + o.EvidenceP = append(o.EvidenceP, e.P) + } + o.PagesRead = len(res.Pages) + o.Requests, o.InTokens, o.CostUSD = res.Requests, res.Usage.InputTokens, res.Usage.CostUSD + o.LeafHit = allInRanges(q.Evidence, o.SelectedPages) + o.Recall = recall(q.Evidence, o.Evidence) + o.Hit = o.Recall == 1 + } + results[i] = o + outMu.Lock() report(o) - continue - } - for _, l := range res.Selected { - o.Selected = append(o.Selected, l.Title) - o.SelectedPages = append(o.SelectedPages, [2]int{l.Start, l.End}) - } - for _, e := range res.Evidence { - o.Evidence = append(o.Evidence, e.Page.Number) - o.EvidenceP = append(o.EvidenceP, e.P) - } - o.PagesRead = len(res.Pages) - o.Requests, o.InTokens, o.CostUSD = res.Requests, res.Usage.InputTokens, res.Usage.CostUSD - o.LeafHit = allInRanges(q.Evidence, o.SelectedPages) - o.Recall = recall(q.Evidence, o.Evidence) - o.Hit = o.Recall == 1 - results = append(results, o) - report(o) - if of != nil { - b, _ := json.Marshal(o) - of.Write(append(b, '\n')) - } + if of != nil { + b, _ := json.Marshal(o) + of.Write(append(b, '\n')) + } + outMu.Unlock() + }(i, q) } + wg.Wait() + fmt.Printf("\nwall %.1fs for %d questions at parallel=%d; limiter now %d\n", time.Since(runStart).Seconds(), len(qs), *parallel, lim.Limit()) summarise(results) } diff --git a/cmd/server/main.go b/cmd/server/main.go index 7f104b8..48fe422 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -29,6 +29,7 @@ import ( "github.com/hallelx2/llmgate" "github.com/hallelx2/llmgate/judge/typesafe" + "github.com/hallelx2/llmgate/middleware/limit" "github.com/hallelx2/llmgate/middleware/retry" "github.com/hallelx2/llmgate/pricing" "github.com/hallelx2/llmgate/provider/anthropic" @@ -144,7 +145,8 @@ func run() error { if err != nil { return fmt.Errorf("init llm: %w", err) } - judge, err := buildJudge(cfg.Engine.LLM.Judge) + llmClient = limit.Client(newLimiter("llm", cfg.Engine.LLM.Concurrency, logger))(llmClient) + judge, err := buildJudge(cfg.Engine.LLM.Judge, newLimiter("judge", cfg.Engine.LLM.Concurrency, logger)) if err != nil { logger.Error("judge: config invalid", "err", err) os.Exit(1) @@ -416,7 +418,23 @@ func modelFor(c enginecfg.LLMConfig) string { // call; a Judge request that fails past them is handled by the TOC // builder, which keeps extraction's pages rather than degrading // silently (HAL-1369). -func buildJudge(c enginecfg.JudgeBlock) (llmgate.Judge, error) { +// newLimiter builds one adaptive limiter for a provider and logs every +// change it makes: a throttled run must be visible, never silent. +func newLimiter(name string, c enginecfg.ConcurrencyBlock, logger *slog.Logger) *limit.Limiter { + return limit.New(limit.Config{ + Initial: c.Initial, + Max: c.Max, + OnChange: func(e limit.Event) { + if e.Cause == "success" { + logger.Info("limiter: widened", "provider", name, "from", e.From, "to", e.To) + return + } + logger.Warn("limiter: narrowed", "provider", name, "cause", e.Cause, "from", e.From, "to", e.To, "paused_for", e.PausedFor, "err", e.Err) + }, + }) +} + +func buildJudge(c enginecfg.JudgeBlock, lim *limit.Limiter) (llmgate.Judge, error) { if c.TypeSafe.APIKey == "" { return nil, nil } @@ -428,7 +446,9 @@ func buildJudge(c enginecfg.JudgeBlock) (llmgate.Judge, error) { if err != nil { return nil, err } - return retry.NewJudge(retry.Config{MaxRetries: 3})(j), nil + // The limiter sits inside retry: each attempt takes a slot, and the + // failure that triggers a retry has already narrowed the limit. + return retry.NewJudge(retry.Config{MaxRetries: 3})(limit.Judge(lim)(j)), nil } func buildLLM(c enginecfg.LLMConfig) (llmgate.Client, error) { diff --git a/cmd/tocdump/main.go b/cmd/tocdump/main.go index 43e0d8e..fac564e 100644 --- a/cmd/tocdump/main.go +++ b/cmd/tocdump/main.go @@ -23,10 +23,12 @@ import ( "os" "path/filepath" "strings" + "sync" "time" "github.com/hallelx2/llmgate" "github.com/hallelx2/llmgate/judge/typesafe" + "github.com/hallelx2/llmgate/middleware/limit" "github.com/hallelx2/llmgate/middleware/retry" "github.com/hallelx2/llmgate/provider/anthropic" @@ -58,7 +60,14 @@ func main() { // z.ai gateway routinely exceeds 90s on a 100+ page filing, and a // timeout there silently drops the whole tree. Measured 2026-09-18. callTimeout := flag.Duration("timeout", 300*time.Second, "per LLM call timeout") + parallel := flag.Int("parallel", 1, "documents in flight at once; the provider's adaptive limiter governs requests") flag.Parse() + if *parallel < 1 { + *parallel = 1 + } + lim = limit.New(limit.Config{Initial: *parallel, OnChange: func(e limit.Event) { + fmt.Fprintf(os.Stderr, " limiter %s %d -> %d %v\n", e.Cause, e.From, e.To, e.Err) + }}) if *docs == "" || *out == "" { fmt.Fprintln(os.Stderr, "usage: tocdump -docs -out [-no-judge]") os.Exit(2) @@ -81,42 +90,58 @@ func main() { } } + runStart := time.Now() pdfs, _ := filepath.Glob(filepath.Join(*docs, "*.pdf")) + sem := make(chan struct{}, *parallel) + var wg sync.WaitGroup + var printMu sync.Mutex for _, path := range pdfs { - name := strings.TrimSuffix(filepath.Base(path), ".pdf") - d := dump{Doc: name} + wg.Add(1) + sem <- struct{}{} + go func(path string) { + defer wg.Done() + defer func() { <-sem }() + name := strings.TrimSuffix(filepath.Base(path), ".pdf") + d := dump{Doc: name} - pages, err := readPages(path) - if err != nil { - d.Err = "parse: " + err.Error() - write(*out, d) - fmt.Printf(" %-28s parse FAILED\n", name) - continue - } - d.Pages = len(pages) + pages, err := readPages(path) + if err != nil { + d.Err = "parse: " + err.Error() + write(*out, d) + printMu.Lock() + fmt.Printf(" %-28s parse FAILED\n", name) + printMu.Unlock() + return + } + d.Pages = len(pages) - llm := client - if *judgeOnly { - llm = refusingClient{} - } - b := &ingest.TOCBuilder{LLM: llm, Judge: judge, LLMCallTimeout: *callTimeout, MinimalContext: *minimal} - ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute) - start := time.Now() - nodes, usage, err := b.Build(ctx, pages) - cancel() - d.Seconds = time.Since(start).Seconds() - d.Requests, d.InTokens, d.CostUSD = usage.LLMCalls, usage.InputTokens, usage.CostUSD - d.Generative, d.Degraded = usage.GenerativeCalls, usage.Degraded - if err != nil { - d.Err = err.Error() - } - d.Nodes = nodes - write(*out, d) + llm := client + if *judgeOnly { + llm = refusingClient{} + } + b := &ingest.TOCBuilder{LLM: llm, Judge: judge, LLMCallTimeout: *callTimeout, MinimalContext: *minimal} + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute) + start := time.Now() + nodes, usage, err := b.Build(ctx, pages) + cancel() + d.Seconds = time.Since(start).Seconds() + d.Requests, d.InTokens, d.CostUSD = usage.LLMCalls, usage.InputTokens, usage.CostUSD + d.Generative, d.Degraded = usage.GenerativeCalls, usage.Degraded + if err != nil { + d.Err = err.Error() + } + d.Nodes = nodes + write(*out, d) - leaves := countLeaves(nodes) - fmt.Printf(" %-28s %4d pages %6.1fs %3d req %d gen %3d leaves $%.4f %s %s\n", - name, d.Pages, d.Seconds, d.Requests, d.Generative, leaves, d.CostUSD, d.Err, strings.Join(d.Degraded, "; ")) + leaves := countLeaves(nodes) + printMu.Lock() + fmt.Printf(" %-28s %4d pages %6.1fs %3d req %d gen %3d leaves $%.4f %s %s\n", + name, d.Pages, d.Seconds, d.Requests, d.Generative, leaves, d.CostUSD, d.Err, strings.Join(d.Degraded, "; ")) + printMu.Unlock() + }(path) } + wg.Wait() + fmt.Printf(" wall %.1fs for %d documents at parallel=%d; limiter now %d\n", time.Since(runStart).Seconds(), len(pdfs), *parallel, lim.Limit()) } func write(dir string, d dump) { @@ -158,9 +183,13 @@ func buildJudge() (llmgate.Judge, error) { if err != nil { return nil, err } - return retry.NewJudge(retry.Config{MaxRetries: 3})(j), nil + // The limiter sits inside retry so each attempt takes a slot. + return retry.NewJudge(retry.Config{MaxRetries: 6, BaseDelay: 5 * time.Second, MaxDelay: 60 * time.Second})(limit.Judge(lim)(j)), nil } +// lim is the one adaptive limiter every document's Judge traffic shares. +var lim *limit.Limiter + func buildClient() (llmgate.Client, error) { get := func(k string) string { if v := os.Getenv(k); v != "" { diff --git a/config.example.yaml b/config.example.yaml index c19c523..be9e3f3 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -107,6 +107,15 @@ llm: model: "gemini-2.0-flash" reasoning_model: "gemini-2.5-pro" + # Concurrency: one adaptive limiter per provider (chat model, Judge). + # The limit starts at initial and moves — halved on a 429 or transport + # failure (honouring Retry-After), widened by one after a run of + # successes, never above max — so the provider's real capacity sets + # the pace, not a guess. Every change is logged. + concurrency: + initial: 4 + max: 64 + # Judge: a System One model (TypeSafe Jev) that answers the ingest # pipeline's judgements — contents-page detection and page resolution — # in one batched request per document instead of a generative call per diff --git a/config.server.example.yaml b/config.server.example.yaml index 260935f..488d7d7 100644 --- a/config.server.example.yaml +++ b/config.server.example.yaml @@ -90,6 +90,11 @@ engine: # api_key: "" # model: "gemini-2.0-flash" # reasoning_model: "" + # Adaptive per-provider concurrency (HAL-1372): starts at initial, + # halves on 429/transport failure, widens on success, capped at max. + concurrency: + initial: 4 + max: 64 # Judge (TypeSafe Jev): batched contents-page detection and page # resolution during ingest. Enabled when api_key is set — also read # from VLE_TYPESAFE_API_KEY / TYPESAFE_API_KEY. See HAL-1367. diff --git a/go.mod b/go.mod index 2e1ad6e..ddd69bc 100644 --- a/go.mod +++ b/go.mod @@ -14,7 +14,7 @@ require ( github.com/aws/smithy-go v1.25.0 github.com/go-chi/chi/v5 v5.2.5 github.com/google/uuid v1.6.0 - github.com/hallelx2/llmgate v0.4.0 + github.com/hallelx2/llmgate v0.5.0 github.com/hallelx2/pdftable v0.4.0 github.com/hibiken/asynq v0.26.0 github.com/jackc/pgx/v5 v5.9.2 diff --git a/go.sum b/go.sum index 106ba6e..58c27dc 100644 --- a/go.sum +++ b/go.sum @@ -134,6 +134,10 @@ github.com/grpc-ecosystem/grpc-gateway/v2 v2.27.2 h1:8Tjv8EJ+pM1xP8mK6egEbD1OgnV github.com/grpc-ecosystem/grpc-gateway/v2 v2.27.2/go.mod h1:pkJQ2tZHJ0aFOVEEot6oZmaVEZcRme73eIFmhiVuRWs= github.com/hallelx2/llmgate v0.4.0 h1:LvaRt1PWEiR0DcNd3j8qX0GEO0oh/rKPHqfL1RABzII= github.com/hallelx2/llmgate v0.4.0/go.mod h1:WpKwV/utKOmb+G5DwezSGjmGtz9D5I6MuS4MqAs7rkA= +github.com/hallelx2/llmgate v0.4.1-0.20260918182842-b2ec96425ecc h1:4b6tUdYQANqf5iCTnDZzY5dWHbRElVznCH4LiAovBkM= +github.com/hallelx2/llmgate v0.4.1-0.20260918182842-b2ec96425ecc/go.mod h1:WpKwV/utKOmb+G5DwezSGjmGtz9D5I6MuS4MqAs7rkA= +github.com/hallelx2/llmgate v0.5.0 h1:yNBm7NOtDrCnpeMPz2BqbNcCJeUTw0WN6oiIuEFankg= +github.com/hallelx2/llmgate v0.5.0/go.mod h1:WpKwV/utKOmb+G5DwezSGjmGtz9D5I6MuS4MqAs7rkA= github.com/hallelx2/pdftable v0.4.0 h1:ldF8qQrUbejWsbB/JSIqliEXF2jF8z8kjEsbKS8r2BU= github.com/hallelx2/pdftable v0.4.0/go.mod h1:pxNlc4D43wjzis7M6EfgQZvHOsQ4okggm+xqUu+OokI= github.com/hhrutter/lzw v1.0.0 h1:laL89Llp86W3rRs83LvKbwYRx6INE8gDn0XNb1oXtm0= diff --git a/pkg/config/config.go b/pkg/config/config.go index 238c878..b334469 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -377,6 +377,14 @@ type LLMConfig struct { OpenAI OpenAIBlock `yaml:"openai"` Gemini GeminiBlock `yaml:"gemini"` + // Concurrency bounds calls in flight per provider — one adaptive + // limiter for the chat model, one for the Judge. The limit starts at + // Initial and moves: halved on a 429 or transport failure, honouring + // Retry-After; widened by one after a run of successes; never above + // Max. Zero selects llmgate's defaults (4, 64). Replaces guessing a + // fixed number (HAL-1372). + Concurrency ConcurrencyBlock `yaml:"concurrency"` + // Judge configures the System One model that answers the pipeline's // judgements — contents-page detection and page resolution — in one // batched request each, instead of a generative call per page. Left @@ -385,6 +393,12 @@ type LLMConfig struct { Judge JudgeBlock `yaml:"judge"` } +// ConcurrencyBlock configures the adaptive per-provider limiter. +type ConcurrencyBlock struct { + Initial int `yaml:"initial"` + Max int `yaml:"max"` +} + // JudgeBlock configures the Judge. Only TypeSafe is supported today; the // Judge is enabled exactly when an API key is present. type JudgeBlock struct { diff --git a/pkg/retrieval/judgewalk.go b/pkg/retrieval/judgewalk.go index 56acb8f..1744122 100644 --- a/pkg/retrieval/judgewalk.go +++ b/pkg/retrieval/judgewalk.go @@ -5,6 +5,7 @@ import ( "fmt" "sort" "strings" + "sync" "regexp" @@ -262,7 +263,14 @@ func (n *JudgeNavigator) rankPages(ctx context.Context, query string, pages []Na for i, p := range pages { scores[i] = PageScore{Page: p} } - requests := 0 + // Build every batch first, then send them all at once. They do not + // depend on each other, and the provider's limiter — not a loop — + // decides how many are in flight (HAL-1372). + type batch struct { + state map[string]any + questions map[string]llmgate.Question + } + var batches []batch budget := n.reqTokens() for start := 0; start < len(pages); { state := map[string]any{"question": query} @@ -307,22 +315,48 @@ func (n *JudgeNavigator) rankPages(ctx context.Context, query string, pages []Na used += cost end++ } - res, err := n.Judge.Judge(ctx, llmgate.JudgeRequest{State: state, Questions: questions}) - if err != nil { - return nil, usage, requests, err - } - requests++ - usage.Add(judgeUsage(res)) - for qk := range questions { - p, err := res.Noul(qk) + batches = append(batches, batch{state, questions}) + start = end + } + + ctx, cancel := context.WithCancel(ctx) + defer cancel() + var ( + wg sync.WaitGroup + mu sync.Mutex + firstErr error + requests int + ) + for _, b := range batches { + wg.Add(1) + go func(b batch) { + defer wg.Done() + res, err := n.Judge.Judge(ctx, llmgate.JudgeRequest{State: b.state, Questions: b.questions}) + mu.Lock() + defer mu.Unlock() if err != nil { - continue + if firstErr == nil { + firstErr = err + cancel() + } + return } - var i int - fmt.Sscanf(qk, "p_%d", &i) - scores[i].P = p - } - start = end + requests++ + usage.Add(judgeUsage(res)) + for qk := range b.questions { + p, err := res.Noul(qk) + if err != nil { + continue + } + var i int + fmt.Sscanf(qk, "p_%d", &i) + scores[i].P = p + } + }(b) + } + wg.Wait() + if firstErr != nil { + return nil, usage, requests, firstErr } sort.SliceStable(scores, func(i, j int) bool { return scores[i].P > scores[j].P }) return scores, usage, requests, nil diff --git a/pkg/retrieval/judgewalk_test.go b/pkg/retrieval/judgewalk_test.go index 8999cbd..a6fab82 100644 --- a/pkg/retrieval/judgewalk_test.go +++ b/pkg/retrieval/judgewalk_test.go @@ -4,6 +4,7 @@ import ( "context" "errors" "strings" + "sync/atomic" "testing" "github.com/hallelx2/llmgate" @@ -13,10 +14,10 @@ import ( // navJudge answers leaf questions by title keyword and page questions // by text keyword, and counts requests. -func navJudge(leafHit, pageHit string) (*llmgate.MockJudge, *int) { - calls := 0 +func navJudge(leafHit, pageHit string) (*llmgate.MockJudge, *atomic.Int32) { + calls := &atomic.Int32{} j := &llmgate.MockJudge{Respond: func(_ context.Context, req llmgate.JudgeRequest) (*llmgate.Judgment, error) { - calls++ + calls.Add(1) st := req.State.(map[string]any) ans := map[string]llmgate.Answer{} for id := range req.Questions { @@ -32,7 +33,7 @@ func navJudge(leafHit, pageHit string) (*llmgate.MockJudge, *int) { } return &llmgate.Judgment{Model: "mock", Answers: ans, Usage: llmgate.Usage{InputTokens: 10, TotalTokens: 10, TokensReported: true}}, nil }} - return j, &calls + return j, calls } func tenKLeaves() []NavLeaf { @@ -72,8 +73,8 @@ func TestNavigateReadsTheBestLeafAndFindsTheEvidencePage(t *testing.T) { if res.Evidence[0].Page.LeafID != "8" { t.Errorf("evidence page should carry its leaf: %+v", res.Evidence[0].Page) } - if *calls != 2 || res.Requests != 2 { - t.Errorf("requests: mock saw %d, result says %d; want 2 (leaves, pages)", *calls, res.Requests) + if calls.Load() != 2 || res.Requests != 2 { + t.Errorf("requests: mock saw %d, result says %d; want 2 (leaves, pages)", calls.Load(), res.Requests) } if res.Coarse != nil { t.Errorf("six pages under a 40-page budget need no coarse pass") @@ -110,8 +111,8 @@ func TestNavigateCoarsePassReachesDeepIntoABigLeaf(t *testing.T) { if len(res.Evidence) == 0 || res.Evidence[0].Page.Number != 85 { t.Fatalf("page 85 not found: %+v", res.Evidence) } - if *calls < 3 { - t.Errorf("want leaves + coarse + full requests, got %d", *calls) + if calls.Load() < 3 { + t.Errorf("want leaves + coarse + full requests, got %d", calls.Load()) } } @@ -173,8 +174,8 @@ func TestRankPagesBatchesUnderTheRequestBudget(t *testing.T) { } // 1900 chars of "y" is ~475 tokens; two fit under 1200 with the // query, a third does not. - if reqs < 4 || *calls != reqs { - t.Errorf("10 pages of ~475 tokens under a 1200-token budget should take ≥4 requests, took %d (mock saw %d)", reqs, *calls) + if reqs < 4 || int(calls.Load()) != reqs { + t.Errorf("10 pages of ~475 tokens under a 1200-token budget should take ≥4 requests, took %d (mock saw %d)", reqs, calls.Load()) } if scored[0].Page.Number != 8 { t.Errorf("best page should be the one with the needle, got %d", scored[0].Page.Number) @@ -261,8 +262,8 @@ func TestNavigateFollowsACrossReference(t *testing.T) { if len(res.Evidence) == 0 || res.Evidence[0].Page.Number != 113 { t.Fatalf("evidence %+v, want page 113 first", res.Evidence) } - if *calls != 3 { - t.Errorf("requests: %d, want 3 (leaves, Item 3 page, Note 21 pages)", *calls) + if calls.Load() != 3 { + t.Errorf("requests: %d, want 3 (leaves, Item 3 page, Note 21 pages)", calls.Load()) } // Turned off, it stays on Item 3. n.NoFollowReferences = true