From b5bf69df658f28df10a1df6e7c673f9f8a922b20 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jesus=20Nu=C3=B1ez?= Date: Mon, 31 Aug 2026 20:12:39 -0400 Subject: [PATCH 1/2] CI perf & resilience: parallel runner, dynamic proxy timeout, model-alias fallback Three opt-in improvements aimed at CI adoption, none changing default behavior: - runner: add Options.Concurrency and `augur run -concurrency/-j N` to run agent invocations in parallel. Each invocation's stdout/stderr is buffered and flushed as one block so parallel output never interleaves; the Summary is assembled in deterministic task order regardless of completion order. Defaults to 1 (the original strictly-sequential path is unchanged). The proxy's per-(scenario,run) state, the trace writer, and the cassette are already mutex-guarded, so parallel invocations record correctly. - proxy: replace the hardcoded 10m client timeout with a per-request context deadline governed by Server.Timeout (flag -timeout on run/proxy) and overridable per call via the X-Augur-Timeout header, for slow reasoning models. Deriving the deadline from the inbound request context keeps client cancellation propagating to the upstream call (no orphaned tokens). -timeout 0 disables the deadline. - cost: add opt-in Pricing.PrefixFallback + Resolve, matching a model absent from the snapshot to its longest dash-delimited prefix (gpt-4o-2024-08-06 -> gpt-4o) with a dash-boundary rule so a base name can't swallow an unrelated one. aggregate surfaces every fallback in Result.ModelAliases; `-normalize-models` on aggregate/gate enables it and warns per normalized model so it is never silent. Off by default: an un-priced call is still a hard ErrUnknownModel. Tests: runner concurrency (real overlap, ordering, non-interleaved output, abort/continue-on-error), proxy timeout (fires, header override, header stripped), cost Resolve (exact/fallback/longest-match/dash-boundary), aggregate alias surfacing. README: new "Tuning for CI" section. Full suite green. --- README.md | 20 +++ aggregate/aggregate.go | 22 +++- aggregate/alias_test.go | 67 ++++++++++ aggregate_cmd.go | 3 + cost/cost.go | 59 ++++++++- cost/resolve_test.go | 78 ++++++++++++ gate_cmd.go | 3 + pricing_source.go | 20 +++ proxy/proxy.go | 61 ++++++++- proxy/timeout_test.go | 126 +++++++++++++++++++ proxy_cmd.go | 2 + run_cmd.go | 12 +- runner/runner.go | 172 +++++++++++++++++++++++--- runner/runner_concurrency_test.go | 199 ++++++++++++++++++++++++++++++ 14 files changed, 811 insertions(+), 33 deletions(-) create mode 100644 aggregate/alias_test.go create mode 100644 cost/resolve_test.go create mode 100644 proxy/timeout_test.go create mode 100644 runner/runner_concurrency_test.go diff --git a/README.md b/README.md index df7fee8..eaf9e15 100644 --- a/README.md +++ b/README.md @@ -145,6 +145,26 @@ augur gate --traffic traffic.yaml --budget budget.yaml `augur gate` is the one you wire into CI. +### Tuning for CI + +- **Parallel scenarios.** `augur run --concurrency N` (alias `-j N`) runs up to N + agent invocations at once. With ~20 repetitions over network-bound LLM calls + this cuts wall-clock time sharply; the proxy and trace ledger are + concurrency-safe, and each invocation's stdout/stderr is flushed as one block + so parallel output never interleaves. Defaults to `1` (sequential). +- **Per-call timeout.** `augur run --timeout 5m` / `augur proxy --timeout 5m` + bounds how long the proxy waits on the provider (default 10m). A single call + can override it with the `X-Augur-Timeout` header (a Go duration such as + `300s`) — handy for slow reasoning models (o1/o3/r1). The deadline is a context + deadline, so a client that disconnects cancels the upstream call immediately + (no orphaned tokens). `--timeout 0` disables the proxy-imposed deadline. +- **Model-alias fallback.** `augur aggregate --normalize-models` / + `augur gate --normalize-models` prices a model that is absent from the snapshot + via its longest dash-delimited prefix (e.g. `gpt-4o-2024-08-06` → `gpt-4o`), so + a provider bumping a dated suffix doesn't fail the build with + `unknown model`. It is **opt-in** and prints a warning naming every model it + normalized, because a silent fallback can mis-bill. + ### Record once, replay for free Running the agent against the real provider on every CI push spends real tokens. diff --git a/aggregate/aggregate.go b/aggregate/aggregate.go index 5232b38..13728b3 100644 --- a/aggregate/aggregate.go +++ b/aggregate/aggregate.go @@ -67,6 +67,12 @@ type Result struct { // Runs lists every run (sorted by scenario then run id) for hand // reconciliation against the raw trace. Runs []Run `json:"runs"` + // ModelAliases records every model in the trace that was priced via the + // pricing snapshot's prefix fallback (Pricing.PrefixFallback), mapping the + // requested model to the snapshot key it was billed against. It is empty when + // no normalization happened (the default). Callers surface it as a warning so + // a fallback is never silent. + ModelAliases map[string]string `json:"model_aliases,omitempty"` } // Knobs are what-if multipliers applied to every call's cost, for sensitivity @@ -132,6 +138,8 @@ func AggregateWithKnobs(records []trace.Record, pricing cost.Pricing, knobs Knob var runOrder []runKey // scenario -> model -> accumulating usage scenarioModels := make(map[string]map[string]*ModelUsage) + // requested model -> canonical snapshot key, for models priced via fallback. + var aliases map[string]string for _, rec := range records { u := cost.Usage{ @@ -140,7 +148,18 @@ func AggregateWithKnobs(records []trace.Record, pricing cost.Pricing, knobs Knob CachedTokens: rec.CachedTokens, CacheWriteTokens: rec.CacheWriteTokens, } - b, err := pricing.Breakdown(rec.Model, u) + canonical, mp, ok := pricing.Resolve(rec.Model) + if !ok { + return Result{}, fmt.Errorf("aggregate: scenario %q run %q seq %d: %w: %q", + rec.ScenarioID, rec.RunID, rec.Seq, cost.ErrUnknownModel, rec.Model) + } + if canonical != rec.Model { + if aliases == nil { + aliases = make(map[string]string) + } + aliases[rec.Model] = canonical + } + b, err := mp.Breakdown(u) if err != nil { return Result{}, fmt.Errorf("aggregate: scenario %q run %q seq %d: %w", rec.ScenarioID, rec.RunID, rec.Seq, err) @@ -226,6 +245,7 @@ func AggregateWithKnobs(records []trace.Record, pricing cost.Pricing, knobs Knob SnapshotDate: pricing.SnapshotDate, Scenarios: scenarios, Runs: allRuns, + ModelAliases: aliases, }, nil } diff --git a/aggregate/alias_test.go b/aggregate/alias_test.go new file mode 100644 index 0000000..dde27da --- /dev/null +++ b/aggregate/alias_test.go @@ -0,0 +1,67 @@ +package aggregate + +import ( + "errors" + "testing" + + "augur/cost" + "augur/trace" +) + +// TestAggregateUnknownModelStillErrors confirms the default (no fallback) +// behavior is unchanged: an un-priced model is a hard error. +func TestAggregateUnknownModelStillErrors(t *testing.T) { + records := []trace.Record{rec("s", "r", 0, "gpt-4o-2024-08-06", 1000, 500, 0)} + _, err := Aggregate(records, testPricing()) + if !errors.Is(err, cost.ErrUnknownModel) { + t.Fatalf("err = %v, want ErrUnknownModel", err) + } +} + +// TestAggregateModelAliasesSurfaced checks that with PrefixFallback on, a dated +// model is priced via its base entry AND reported in Result.ModelAliases so the +// caller can warn. +func TestAggregateModelAliasesSurfaced(t *testing.T) { + pricing := testPricing() + pricing.PrefixFallback = true + + records := []trace.Record{ + rec("checkout", "run-1", 0, "gpt-4o-2024-08-06", 1_000_000, 0, 0), // -> gpt-4o + rec("checkout", "run-1", 1, "gpt-4o", 0, 0, 0), // exact, no alias + } + res, err := AggregateWithKnobs(records, pricing, Knobs{}) + if err != nil { + t.Fatalf("AggregateWithKnobs: %v", err) + } + + if got := res.ModelAliases["gpt-4o-2024-08-06"]; got != "gpt-4o" { + t.Errorf("ModelAliases[gpt-4o-2024-08-06] = %q, want gpt-4o", got) + } + if _, aliased := res.ModelAliases["gpt-4o"]; aliased { + t.Errorf("exact match gpt-4o must not appear in ModelAliases: %v", res.ModelAliases) + } + + // The dated call was priced at gpt-4o's input rate (2.50 / Mtok on 1M tokens). + var total float64 + for _, r := range res.Runs { + total += r.CostUSD + } + if total != 2.50 { + t.Errorf("total cost = %v, want 2.50 (billed via gpt-4o fallback)", total) + } +} + +// TestAggregateNoAliasesWhenAllExact checks ModelAliases stays nil when nothing +// was normalized, so the warning path is a genuine no-op on the common case. +func TestAggregateNoAliasesWhenAllExact(t *testing.T) { + pricing := testPricing() + pricing.PrefixFallback = true + records := []trace.Record{rec("s", "r", 0, "gpt-4o", 1000, 500, 0)} + res, err := Aggregate(records, pricing) + if err != nil { + t.Fatalf("Aggregate: %v", err) + } + if res.ModelAliases != nil { + t.Errorf("ModelAliases = %v, want nil when no fallback happened", res.ModelAliases) + } +} diff --git a/aggregate_cmd.go b/aggregate_cmd.go index cb6a288..5695b83 100644 --- a/aggregate_cmd.go +++ b/aggregate_cmd.go @@ -19,6 +19,7 @@ func runAggregate(args []string) error { pricingPath := fs.String("pricing", "pricing.yaml", "path to the pricing snapshot") tcoPath := fs.String("tco", "", "derive pricing from a self-hosted TCO config instead of -pricing") asJSON := fs.Bool("json", false, "emit the aggregation as JSON instead of a table") + normalizeModels := fs.Bool("normalize-models", false, "price a model absent from the snapshot via its longest dash-delimited prefix (e.g. gpt-4o-2024-08-06 -> gpt-4o); warns per fallback") if err := fs.Parse(args); err != nil { return err } @@ -38,11 +39,13 @@ func runAggregate(args []string) error { if err != nil { return err } + pricing.PrefixFallback = *normalizeModels res, err := aggregate.Aggregate(records, pricing) if err != nil { return err } + warnModelAliases(os.Stderr, res.ModelAliases) if *asJSON { enc := json.NewEncoder(os.Stdout) diff --git a/cost/cost.go b/cost/cost.go index be37e1b..3c548aa 100644 --- a/cost/cost.go +++ b/cost/cost.go @@ -10,6 +10,7 @@ package cost import ( "errors" "fmt" + "strings" ) // tokensPerMtok is the denominator that turns a per-million-token price into a @@ -132,18 +133,63 @@ func (p ModelPrice) Cost(u Usage) (float64, error) { type Pricing struct { SnapshotDate string Models map[string]ModelPrice + // PrefixFallback, when true, lets Resolve match a model that is absent from + // the snapshot to the longest snapshot key that is a dash-delimited prefix of + // it — so a dated or suffixed variant (e.g. "gpt-4o-2024-08-06") falls back to + // its base entry ("gpt-4o") instead of failing. It is OFF by default because a + // silent fallback can mis-bill; callers that enable it should surface which + // models were normalized (Resolve reports the matched key so they can warn). + PrefixFallback bool } -// Price returns the price for a model and whether it is known. +// Price returns the exact price for a model and whether it is known. It does not +// apply PrefixFallback — use Resolve for that. func (p Pricing) Price(model string) (ModelPrice, bool) { mp, ok := p.Models[model] return mp, ok } -// Cost computes the cost of a single call for the named model. It wraps -// ErrUnknownModel (so callers can errors.Is it) when the model is absent. +// Resolve looks up the price for a model, returning the snapshot key it matched +// (the canonical model), its price, and whether a match was found. An exact +// match returns the model unchanged. When PrefixFallback is on and there is no +// exact match, it returns the longest snapshot key that is a dash-delimited +// prefix of the model (so callers can detect normalization by comparing the +// returned canonical to the requested model). +func (p Pricing) Resolve(model string) (canonical string, mp ModelPrice, ok bool) { + if mp, ok := p.Models[model]; ok { + return model, mp, true + } + if !p.PrefixFallback { + return "", ModelPrice{}, false + } + best := "" + for key := range p.Models { + if isModelPrefix(model, key) && len(key) > len(best) { + best = key + } + } + if best == "" { + return "", ModelPrice{}, false + } + return best, p.Models[best], true +} + +// isModelPrefix reports whether key is a prefix of model at a dash boundary, so +// "gpt-4o" matches "gpt-4o" and "gpt-4o-2024-08-06" but never "gpt-4omini". The +// boundary rule keeps a base name from swallowing an unrelated one that merely +// shares a leading substring. +func isModelPrefix(model, key string) bool { + if !strings.HasPrefix(model, key) { + return false + } + return len(model) == len(key) || model[len(key)] == '-' +} + +// Cost computes the cost of a single call for the named model. It applies +// PrefixFallback via Resolve and wraps ErrUnknownModel (so callers can +// errors.Is it) when no model matches. func (p Pricing) Cost(model string, u Usage) (float64, error) { - mp, ok := p.Models[model] + _, mp, ok := p.Resolve(model) if !ok { return 0, fmt.Errorf("%w: %q", ErrUnknownModel, model) } @@ -151,9 +197,10 @@ func (p Pricing) Cost(model string, u Usage) (float64, error) { } // Breakdown computes the per-component cost of a single call for the named -// model, wrapping ErrUnknownModel when the model is absent. +// model, applying PrefixFallback via Resolve and wrapping ErrUnknownModel when +// no model matches. func (p Pricing) Breakdown(model string, u Usage) (Breakdown, error) { - mp, ok := p.Models[model] + _, mp, ok := p.Resolve(model) if !ok { return Breakdown{}, fmt.Errorf("%w: %q", ErrUnknownModel, model) } diff --git a/cost/resolve_test.go b/cost/resolve_test.go new file mode 100644 index 0000000..2afb3c1 --- /dev/null +++ b/cost/resolve_test.go @@ -0,0 +1,78 @@ +package cost + +import ( + "errors" + "testing" +) + +func fallbackPricing() Pricing { + return Pricing{ + SnapshotDate: "2026-08-31", + Models: map[string]ModelPrice{ + "gpt-4o": {Input: 2.5, Output: 10}, + "gpt-4o-mini": {Input: 0.15, Output: 0.6}, + }, + } +} + +func TestResolveExactMatch(t *testing.T) { + p := fallbackPricing() + canonical, mp, ok := p.Resolve("gpt-4o") + if !ok || canonical != "gpt-4o" || mp.Input != 2.5 { + t.Fatalf("Resolve(gpt-4o) = %q,%+v,%v; want gpt-4o with Input 2.5", canonical, mp, ok) + } +} + +func TestResolveFallbackOffByDefault(t *testing.T) { + p := fallbackPricing() + if _, _, ok := p.Resolve("gpt-4o-2024-08-06"); ok { + t.Error("Resolve should not fall back when PrefixFallback is off") + } + if _, err := p.Cost("gpt-4o-2024-08-06", Usage{InputTokens: 1000}); !errors.Is(err, ErrUnknownModel) { + t.Errorf("Cost error = %v, want ErrUnknownModel with fallback off", err) + } +} + +func TestResolvePrefixFallback(t *testing.T) { + p := fallbackPricing() + p.PrefixFallback = true + + // Dated variant falls back to its base entry. + canonical, mp, ok := p.Resolve("gpt-4o-2024-08-06") + if !ok || canonical != "gpt-4o" || mp.Input != 2.5 { + t.Errorf("Resolve(gpt-4o-2024-08-06) = %q,%+v,%v; want gpt-4o", canonical, mp, ok) + } + + // Longest prefix wins: the mini variant must not collapse to gpt-4o. + canonical, mp, ok = p.Resolve("gpt-4o-mini-2024-07-18") + if !ok || canonical != "gpt-4o-mini" || mp.Input != 0.15 { + t.Errorf("Resolve(gpt-4o-mini-2024-07-18) = %q,%+v,%v; want gpt-4o-mini", canonical, mp, ok) + } + + // Cost now succeeds via the fallback price. + got, err := p.Cost("gpt-4o-2024-08-06", Usage{InputTokens: 1_000_000}) + if err != nil { + t.Fatalf("Cost with fallback: %v", err) + } + if got != 2.5 { + t.Errorf("Cost = %v, want 2.5 (gpt-4o input price)", got) + } +} + +func TestResolveDashBoundary(t *testing.T) { + p := fallbackPricing() + p.PrefixFallback = true + // "gpt-4omini" shares a leading substring with "gpt-4o" but not at a dash + // boundary — it must NOT match, or a base name would swallow unrelated ones. + if canonical, _, ok := p.Resolve("gpt-4omini"); ok { + t.Errorf("Resolve(gpt-4omini) matched %q; want no match (dash boundary)", canonical) + } +} + +func TestResolveNoMatchStillFails(t *testing.T) { + p := fallbackPricing() + p.PrefixFallback = true + if _, _, ok := p.Resolve("claude-3-5-sonnet"); ok { + t.Error("Resolve matched an unrelated model family; want no match") + } +} diff --git a/gate_cmd.go b/gate_cmd.go index 4035cae..a790dc2 100644 --- a/gate_cmd.go +++ b/gate_cmd.go @@ -28,6 +28,7 @@ func runGate(args []string) error { ciLevel := fs.Float64("ci", 0.95, "confidence level for bootstrap intervals (0..1)") bootstrap := fs.Int("bootstrap", 2000, "number of bootstrap resamples") seed := fs.Uint64("seed", 1, "PRNG seed for reproducible bootstrap intervals") + normalizeModels := fs.Bool("normalize-models", false, "price a model absent from the snapshot via its longest dash-delimited prefix (e.g. gpt-4o-2024-08-06 -> gpt-4o); warns per fallback") knobFs := addKnobFlags(fs) if err := fs.Parse(args); err != nil { return err @@ -42,6 +43,7 @@ func runGate(args []string) error { if err != nil { return err } + pricing.PrefixFallback = *normalizeModels traffic, err := project.LoadTraffic(*trafficPath) if err != nil { return err @@ -55,6 +57,7 @@ func runGate(args []string) error { if err != nil { return err } + warnModelAliases(os.Stderr, res.ModelAliases) proj, err := project.Project(res, traffic, project.Options{ CILevel: *ciLevel, BootstrapSamples: *bootstrap, Seed: *seed, }) diff --git a/pricing_source.go b/pricing_source.go index cec0004..b3180da 100644 --- a/pricing_source.go +++ b/pricing_source.go @@ -2,6 +2,8 @@ package main import ( "fmt" + "io" + "sort" "augur/cost" "augur/tco" @@ -21,3 +23,21 @@ func resolvePricing(pricingPath, tcoPath string) (cost.Pricing, error) { } return cost.LoadPricing(pricingPath) } + +// warnModelAliases prints, to w, one warning per model that was priced via the +// snapshot's prefix fallback (see aggregate.Result.ModelAliases). A fallback is +// a convenience that can mis-bill if the base entry's price differs from the +// requested variant's, so it is never silent. Keys are sorted for stable output. +func warnModelAliases(w io.Writer, aliases map[string]string) { + if len(aliases) == 0 { + return + } + requested := make([]string, 0, len(aliases)) + for r := range aliases { + requested = append(requested, r) + } + sort.Strings(requested) + for _, r := range requested { + fmt.Fprintf(w, "augur: WARNING model %q not in pricing snapshot; billed at %q via prefix fallback\n", r, aliases[r]) + } +} diff --git a/proxy/proxy.go b/proxy/proxy.go index 03546b1..59ccbfa 100644 --- a/proxy/proxy.go +++ b/proxy/proxy.go @@ -15,6 +15,7 @@ package proxy import ( "bytes" + "context" "encoding/json" "fmt" "hash/fnv" @@ -35,8 +36,18 @@ import ( const ( HeaderScenarioID = "X-Augur-Scenario-Id" HeaderRunID = "X-Augur-Run-Id" + // HeaderTimeout lets a single call override the proxy's default upstream + // timeout (a Go duration string, e.g. "300s" or "5m"). Reasoning models + // (o1/o3/r1) can legitimately run for minutes; a fast call can be capped + // tighter. It is Augur's header and is stripped before forwarding. + HeaderTimeout = "X-Augur-Timeout" ) +// DefaultTimeout is the upstream request deadline applied when neither a +// -timeout flag nor a per-request X-Augur-Timeout header narrows it. It bounds +// how long the proxy waits on the provider before giving up. +const DefaultTimeout = 10 * time.Minute + // nowFunc returns the current time. It is a field on Server (defaulting to // time.Now) so tests can stamp deterministic timestamps. type nowFunc func() time.Time @@ -67,6 +78,13 @@ type Server struct { // Set false to forward the request byte-for-byte. InjectUsage bool + // Timeout bounds how long the proxy waits on the upstream provider for a + // single call, applied as a context deadline so client cancellation is still + // honored (a cancelled inbound request cancels the outbound one immediately). + // A per-call X-Augur-Timeout header overrides it; a value <= 0 means no + // proxy-imposed deadline. Defaults to DefaultTimeout from New. + Timeout time.Duration + mode Mode cassette *cassette.Cassette @@ -81,7 +99,12 @@ type Server struct { // change mode. func New(upstream *url.URL, tracer *trace.Writer, client *http.Client) *Server { if client == nil { - client = &http.Client{Timeout: 10 * time.Minute} + // No client-level Timeout: the deadline is governed per-request by + // Server.Timeout (and the X-Augur-Timeout override) via the request + // context, so a slow reasoning model can be granted more time without + // re-constructing the client, and a cancelled inbound request cancels + // the outbound one at once. + client = &http.Client{} } return &Server{ upstream: upstream, @@ -89,6 +112,7 @@ func New(upstream *url.URL, tracer *trace.Writer, client *http.Client) *Server { client: client, now: time.Now, InjectUsage: true, + Timeout: DefaultTimeout, seq: make(map[string]int), seen: make(map[string]map[uint64]bool), } @@ -183,7 +207,13 @@ func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) { if s.InjectUsage { fwdBody = maybeInjectIncludeUsage(r.URL.Path, reqBody) } - outReq, err := s.buildUpstreamRequest(r, fwdBody) + // Bound the upstream call by the effective timeout (per-request header, else + // the server default), derived from the inbound request context so client + // cancellation still propagates. cancel is released once the response is + // fully handled — including the streaming path, which reads the body below. + ctx, cancel := s.upstreamContext(r) + defer cancel() + outReq, err := s.buildUpstreamRequest(ctx, r, fwdBody) if err != nil { http.Error(w, "augur proxy: building upstream request: "+err.Error(), http.StatusBadGateway) return @@ -287,16 +317,36 @@ func pickModel(reqModel, respModel string) string { return respModel } +// upstreamContext derives the context that bounds the outbound call. It starts +// from the inbound request context (so a cancelled client cancels the provider +// call) and layers a deadline: the X-Augur-Timeout header if present and +// parseable, otherwise s.Timeout. A non-positive effective timeout means no +// proxy-imposed deadline — the returned cancel is then a no-op but is always +// safe (and required) to call. +func (s *Server) upstreamContext(r *http.Request) (context.Context, context.CancelFunc) { + timeout := s.Timeout + if h := r.Header.Get(HeaderTimeout); h != "" { + if d, err := time.ParseDuration(h); err == nil && d > 0 { + timeout = d + } + } + if timeout <= 0 { + return context.WithCancel(r.Context()) + } + return context.WithTimeout(r.Context(), timeout) +} + // buildUpstreamRequest clones the inbound request onto the upstream base URL, // preserving method, path, query, and headers (including Authorization) while // stripping Augur's own headers and Accept-Encoding (so Go's transport handles -// compression transparently and we get a decoded body to parse usage from). -func (s *Server) buildUpstreamRequest(r *http.Request, body []byte) (*http.Request, error) { +// compression transparently and we get a decoded body to parse usage from). The +// request carries ctx so its deadline and cancellation govern the call. +func (s *Server) buildUpstreamRequest(ctx context.Context, r *http.Request, body []byte) (*http.Request, error) { out := *s.upstream out.Path = singleJoiningSlash(s.upstream.Path, r.URL.Path) out.RawQuery = r.URL.RawQuery - req, err := http.NewRequestWithContext(r.Context(), r.Method, out.String(), bytes.NewReader(body)) + req, err := http.NewRequestWithContext(ctx, r.Method, out.String(), bytes.NewReader(body)) if err != nil { return nil, err } @@ -304,6 +354,7 @@ func (s *Server) buildUpstreamRequest(r *http.Request, body []byte) (*http.Reque stripHopByHop(req.Header) req.Header.Del(HeaderScenarioID) req.Header.Del(HeaderRunID) + req.Header.Del(HeaderTimeout) // Let the Go transport negotiate and transparently decode compression so the // response body we read (and parse usage from) is the decoded JSON. req.Header.Del("Accept-Encoding") diff --git a/proxy/timeout_test.go b/proxy/timeout_test.go new file mode 100644 index 0000000..4746840 --- /dev/null +++ b/proxy/timeout_test.go @@ -0,0 +1,126 @@ +package proxy + +import ( + "bytes" + "io" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "testing" + "time" + + "augur/trace" +) + +// slowUpstream replies with chatResponse only after the given delay, unless the +// caller (the proxy) disconnects first — mirroring a slow reasoning model while +// still honoring cancellation. +func slowUpstream(delay time.Duration) *httptest.Server { + return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + select { + case <-time.After(delay): + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, chatResponse) + case <-r.Context().Done(): + } + })) +} + +func postChat(t *testing.T, proxyURL string, hdr map[string]string) *http.Response { + t.Helper() + body := `{"model":"gpt-4o","messages":[{"role":"user","content":"hi"}]}` + req, _ := http.NewRequest(http.MethodPost, proxyURL+"/v1/chat/completions", strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + req.Header.Set(HeaderScenarioID, "s") + req.Header.Set(HeaderRunID, "r") + for k, v := range hdr { + req.Header.Set(k, v) + } + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("request: %v", err) + } + return resp +} + +// TestProxyTimeoutFires checks that a Server.Timeout shorter than the upstream's +// response time aborts the call promptly with a 502 instead of hanging. +func TestProxyTimeoutFires(t *testing.T) { + up := slowUpstream(2 * time.Second) + defer up.Close() + u, _ := url.Parse(up.URL) + + var buf bytes.Buffer + s := New(u, trace.NewWriter(&buf), up.Client()) + s.Timeout = 50 * time.Millisecond + srv := httptest.NewServer(s) + defer srv.Close() + + start := time.Now() + resp := postChat(t, srv.URL, nil) + _, _ = io.ReadAll(resp.Body) + resp.Body.Close() + elapsed := time.Since(start) + + if resp.StatusCode != http.StatusBadGateway { + t.Errorf("status = %d, want 502 on timeout", resp.StatusCode) + } + if elapsed > time.Second { + t.Errorf("call took %v, want it to abort near the 50ms timeout", elapsed) + } +} + +// TestProxyTimeoutHeaderOverride checks a per-request X-Augur-Timeout narrows the +// deadline below a generous Server.Timeout. +func TestProxyTimeoutHeaderOverride(t *testing.T) { + up := slowUpstream(2 * time.Second) + defer up.Close() + u, _ := url.Parse(up.URL) + + var buf bytes.Buffer + s := New(u, trace.NewWriter(&buf), up.Client()) + s.Timeout = 10 * time.Minute // generous default + srv := httptest.NewServer(s) + defer srv.Close() + + start := time.Now() + resp := postChat(t, srv.URL, map[string]string{HeaderTimeout: "50ms"}) + _, _ = io.ReadAll(resp.Body) + resp.Body.Close() + elapsed := time.Since(start) + + if resp.StatusCode != http.StatusBadGateway { + t.Errorf("status = %d, want 502 when the header timeout fires", resp.StatusCode) + } + if elapsed > time.Second { + t.Errorf("call took %v, want the 50ms header timeout to win over Server.Timeout", elapsed) + } +} + +// TestProxyTimeoutHeaderStripped checks the X-Augur-Timeout header is Augur's own +// and never forwarded to the provider, and that a normal call still succeeds. +func TestProxyTimeoutHeaderStripped(t *testing.T) { + up := newFakeUpstream() + defer up.close() + up.respBody = chatResponse + + var buf bytes.Buffer + s := newTestProxy(t, up, &buf) + srv := httptest.NewServer(s) + defer srv.Close() + + resp := postChat(t, srv.URL, map[string]string{HeaderTimeout: "5m"}) + _, _ = io.ReadAll(resp.Body) + resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + t.Fatalf("status = %d, want 200", resp.StatusCode) + } + if len(up.gotReqs) != 1 { + t.Fatalf("upstream got %d requests, want 1", len(up.gotReqs)) + } + if got := up.gotReqs[0].header.Get(HeaderTimeout); got != "" { + t.Errorf("X-Augur-Timeout leaked upstream: %q", got) + } +} diff --git a/proxy_cmd.go b/proxy_cmd.go index 5d29c67..47599d2 100644 --- a/proxy_cmd.go +++ b/proxy_cmd.go @@ -19,6 +19,7 @@ func runProxy(args []string) error { upstream := fs.String("upstream", "https://api.openai.com", "base URL of the real provider (OpenAI-, Anthropic-, or Gemini-compatible)") tracePath := fs.String("trace", "trace.jsonl", "path to append the cost trace to (JSONL)") injectUsage := fs.Bool("inject-usage", true, "auto-set stream_options.include_usage on OpenAI streaming requests so usage is captured exactly") + timeout := fs.Duration("timeout", proxy.DefaultTimeout, "per-call upstream timeout; a request may override it with the X-Augur-Timeout header (0 = no proxy-imposed deadline)") if err := fs.Parse(args); err != nil { return err } @@ -39,6 +40,7 @@ func runProxy(args []string) error { srv := proxy.New(up, tracer, nil) srv.InjectUsage = *injectUsage + srv.Timeout = *timeout fmt.Printf("augur proxy: listening on %s → forwarding to %s, tracing to %s\n", *listen, up.String(), *tracePath) diff --git a/run_cmd.go b/run_cmd.go index 849609e..fbefe3f 100644 --- a/run_cmd.go +++ b/run_cmd.go @@ -37,9 +37,15 @@ func runRun(args []string) error { record := fs.String("record", "", "record every response to this cassette file (real provider calls)") replay := fs.String("replay", "", "replay responses from this cassette file (no provider calls, no tokens)") injectUsage := fs.Bool("inject-usage", true, "auto-set stream_options.include_usage on OpenAI streaming requests so usage is captured exactly") + timeout := fs.Duration("timeout", proxy.DefaultTimeout, "per-call upstream timeout; an agent may override it per request with the X-Augur-Timeout header (0 = no proxy-imposed deadline)") + concurrency := fs.Int("concurrency", 1, "number of agent invocations to run in parallel (alias: -j)") + fs.IntVar(concurrency, "j", 1, "shorthand for -concurrency") if err := fs.Parse(args); err != nil { return err } + if *concurrency < 1 { + return fmt.Errorf("-concurrency must be >= 1, got %d", *concurrency) + } if *record != "" && *replay != "" { return fmt.Errorf("-record and -replay are mutually exclusive") } @@ -68,6 +74,7 @@ func runRun(args []string) error { pxy := proxy.New(up, tracer, nil) pxy.InjectUsage = *injectUsage + pxy.Timeout = *timeout cass, err := configureCassette(pxy, *record, *replay) if err != nil { return err @@ -102,14 +109,15 @@ func runRun(args []string) error { } fmt.Printf("augur run: proxy on %s (%s), tracing to %s\n", baseURL, modeLabel(*record, *replay, up), *tracePath) - fmt.Printf("augur run: %d scenario(s), %d run(s) each, session %q\n", - len(cfg.Scenarios), effectiveRuns(cfg.Runs, *runs), sess) + fmt.Printf("augur run: %d scenario(s), %d run(s) each, concurrency %d, session %q\n", + len(cfg.Scenarios), effectiveRuns(cfg.Runs, *runs), *concurrency, sess) sum, runErr := runner.Run(ctx, cfg, runner.Options{ BaseURL: baseURL, Runs: *runs, Session: sess, ContinueOnError: *continueOnError, + Concurrency: *concurrency, }) // Shut the proxy down and surface any serve error that isn't the expected diff --git a/runner/runner.go b/runner/runner.go index 6032d95..fff480e 100644 --- a/runner/runner.go +++ b/runner/runner.go @@ -1,6 +1,7 @@ package runner import ( + "bytes" "context" "fmt" "io" @@ -8,6 +9,7 @@ import ( "os/exec" "slices" "strings" + "sync" ) // Environment-variable names the runner injects per invocation. See the package @@ -52,6 +54,14 @@ type Options struct { // ContinueOnError keeps running after an invocation fails instead of // aborting at the first failure. ContinueOnError bool + // Concurrency is how many invocations run in parallel. Values <= 1 (the + // default) preserve the original strictly-sequential behavior, streaming each + // invocation's output directly. With >1, invocations run through a worker pool + // and each one's stdout/stderr is buffered and flushed as a block so parallel + // output never interleaves. The proxy and trace writer are concurrency-safe, + // so parallel invocations record correctly; the Summary is always assembled + // in scenario/run order regardless of completion order. + Concurrency int Stdout io.Writer Stderr io.Writer @@ -99,34 +109,158 @@ func Run(ctx context.Context, cfg Config, opts Options) (Summary, error) { stderr = os.Stderr } - var sum Summary + // Flatten scenarios × repetitions into an ordered task list. The Summary is + // always assembled in this order, so parallel completion order never changes + // the output. + var tasks []task for _, sc := range cfg.Scenarios { for i := range runs { - if err := ctx.Err(); err != nil { - return sum, err - } + tasks = append(tasks, task{sc: sc, index: i}) + } + } - runID := makeRunID(opts.Session, sc.ID, i) - env := buildEnv(baseEnv, sc, runID, opts.BaseURL) - args := substituteInput(cfg.Command, sc.Input) + concurrency := max(opts.Concurrency, 1) - err := execFn(ctx, args, env, stdout, stderr) - sum.Total++ - if err != nil { - sum.Failed++ - err = fmt.Errorf("scenario %q run %q (index %d): %w", sc.ID, runID, i, err) - } - sum.Invocations = append(sum.Invocations, Invocation{ - ScenarioID: sc.ID, RunID: runID, Index: i, Err: err, - }) - if err != nil && !opts.ContinueOnError { - return sum, err - } + r := invoker{ + cfg: cfg, opts: opts, baseEnv: baseEnv, + execFn: execFn, stdout: stdout, stderr: stderr, + } + if concurrency == 1 { + return r.runSequential(ctx, tasks) + } + return r.runParallel(ctx, tasks, concurrency) +} + +// task is one scheduled invocation: a scenario and which repetition (index) of +// it to run. +type task struct { + sc Scenario + index int +} + +// invoker holds the resolved per-Run configuration so the sequential and +// parallel paths share one place that builds each invocation. +type invoker struct { + cfg Config + opts Options + baseEnv []string + execFn ExecFunc + stdout io.Writer + stderr io.Writer +} + +// invoke runs one task, writing the agent's output to the given writers, and +// returns the Invocation (with any exec error already wrapped). +func (r invoker) invoke(ctx context.Context, t task, stdout, stderr io.Writer) Invocation { + runID := makeRunID(r.opts.Session, t.sc.ID, t.index) + env := buildEnv(r.baseEnv, t.sc, runID, r.opts.BaseURL) + args := substituteInput(r.cfg.Command, t.sc.Input) + + err := r.execFn(ctx, args, env, stdout, stderr) + if err != nil { + err = fmt.Errorf("scenario %q run %q (index %d): %w", t.sc.ID, runID, t.index, err) + } + return Invocation{ScenarioID: t.sc.ID, RunID: runID, Index: t.index, Err: err} +} + +// runSequential executes tasks one at a time, streaming each invocation's output +// directly. This is the original behavior and the default (Concurrency <= 1): +// unless ContinueOnError is set, the first failing invocation aborts the run. +func (r invoker) runSequential(ctx context.Context, tasks []task) (Summary, error) { + var sum Summary + for _, t := range tasks { + if err := ctx.Err(); err != nil { + return sum, err + } + inv := r.invoke(ctx, t, r.stdout, r.stderr) + sum.Total++ + if inv.Err != nil { + sum.Failed++ + } + sum.Invocations = append(sum.Invocations, inv) + if inv.Err != nil && !r.opts.ContinueOnError { + return sum, inv.Err } } return sum, nil } +// runParallel executes up to `concurrency` tasks at once. Each invocation's +// stdout/stderr is buffered and flushed as a single block under a mutex, so +// parallel output never interleaves. Without ContinueOnError, the first genuine +// failure cancels the run's context: in-flight invocations are signalled to +// stop and no further tasks are launched, and that triggering error is returned. +func (r invoker) runParallel(ctx context.Context, tasks []task, concurrency int) (Summary, error) { + results := make([]Invocation, len(tasks)) + ran := make([]bool, len(tasks)) + + runCtx, cancel := context.WithCancel(ctx) + defer cancel() + + sem := make(chan struct{}, concurrency) + var wg sync.WaitGroup + var flushMu sync.Mutex // serializes output flushes + var errOnce sync.Once + var triggerErr error // the failure that aborted the run (if any) + + for idx, t := range tasks { + if runCtx.Err() != nil { + break // aborted or cancelled: stop launching new work + } + wg.Add(1) + sem <- struct{}{} + go func(idx int, t task) { + defer wg.Done() + defer func() { <-sem }() + + var ob, eb bytes.Buffer + inv := r.invoke(runCtx, t, &ob, &eb) + + flushMu.Lock() + _, _ = io.Copy(r.stdout, &ob) + _, _ = io.Copy(r.stderr, &eb) + flushMu.Unlock() + + // Distinct index per goroutine: writes to different slice elements + // don't race, and the wg.Wait below establishes the read barrier. + results[idx] = inv + ran[idx] = true + if inv.Err != nil && !r.opts.ContinueOnError { + errOnce.Do(func() { triggerErr = inv.Err }) + cancel() + } + }(idx, t) + } + wg.Wait() + + sum := summarize(tasks, results, ran) + if triggerErr != nil { + return sum, triggerErr + } + if err := ctx.Err(); err != nil { + return sum, err + } + return sum, nil +} + +// summarize assembles the Summary in task order from the parallel results, +// including only invocations that actually ran. +func summarize(tasks []task, results []Invocation, ran []bool) Summary { + var sum Summary + for idx := range tasks { + if !ran[idx] { + continue + } + inv := results[idx] + sum.Total++ + if inv.Err != nil { + sum.Failed++ + } + sum.Invocations = append(sum.Invocations, inv) + } + return sum +} + // makeRunID builds a per-invocation run id. With a session it is // "--", otherwise "-", zero-padded // so lexical and numeric ordering agree for up to 1000 runs. diff --git a/runner/runner_concurrency_test.go b/runner/runner_concurrency_test.go new file mode 100644 index 0000000..0cddc8a --- /dev/null +++ b/runner/runner_concurrency_test.go @@ -0,0 +1,199 @@ +package runner + +import ( + "bytes" + "context" + "errors" + "fmt" + "io" + "strings" + "sync" + "testing" + "time" +) + +// TestRunParallelRunsEverything checks the parallel path executes every +// invocation and reports a Summary assembled in task order regardless of +// completion order. +func TestRunParallelRunsEverything(t *testing.T) { + var mu sync.Mutex + var n int + opts := Options{ + Concurrency: 4, + Stdout: io.Discard, Stderr: io.Discard, + Exec: func(_ context.Context, _, _ []string, _, _ io.Writer) error { + mu.Lock() + n++ + mu.Unlock() + return nil + }, + } + sum, err := Run(context.Background(), twoScenarios(), opts) + if err != nil { + t.Fatalf("Run: %v", err) + } + if n != 6 || sum.Total != 6 || sum.Failed != 0 { + t.Fatalf("n=%d Summary=%+v, want 6 invocations, Total 6 Failed 0", n, sum) + } + + // Summary order must be the deterministic task order (scenario, then index), + // not whatever order the goroutines happened to finish in. + want := []string{ + "checkout-000", "checkout-001", "checkout-002", + "faq-000", "faq-001", "faq-002", + } + if len(sum.Invocations) != len(want) { + t.Fatalf("got %d invocations, want %d", len(sum.Invocations), len(want)) + } + for i, inv := range sum.Invocations { + if inv.RunID != want[i] { + t.Errorf("Invocations[%d].RunID = %q, want %q", i, inv.RunID, want[i]) + } + } +} + +// TestRunParallelIsActuallyConcurrent proves invocations really overlap: with +// concurrency N and N tasks, all N must be in flight at once before any is +// released. If the runner were still sequential this blocks and the test times +// out. +func TestRunParallelIsActuallyConcurrent(t *testing.T) { + const n = 3 + cfg := Config{Runs: n, Command: []string{"x"}, Scenarios: []Scenario{{ID: "a", Input: "i"}}} + + var mu sync.Mutex + inFlight, peak := 0, 0 + var arrived sync.WaitGroup + arrived.Add(n) + gate := make(chan struct{}) + + opts := Options{ + Concurrency: n, + Stdout: io.Discard, Stderr: io.Discard, + Exec: func(_ context.Context, _, _ []string, _, _ io.Writer) error { + mu.Lock() + inFlight++ + if inFlight > peak { + peak = inFlight + } + mu.Unlock() + arrived.Done() + <-gate // hold until every invocation has arrived + mu.Lock() + inFlight-- + mu.Unlock() + return nil + }, + } + + done := make(chan Summary, 1) + go func() { + sum, _ := Run(context.Background(), cfg, opts) + done <- sum + }() + + allArrived := make(chan struct{}) + go func() { arrived.Wait(); close(allArrived) }() + select { + case <-allArrived: + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for concurrent invocations; runner is not parallel") + } + close(gate) + + sum := <-done + if sum.Total != n { + t.Errorf("Total = %d, want %d", sum.Total, n) + } + if peak != n { + t.Errorf("peak concurrency = %d, want %d", peak, n) + } +} + +// TestRunParallelOutputNotInterleaved checks each invocation's output is flushed +// as one contiguous block even when many run at once. +func TestRunParallelOutputNotInterleaved(t *testing.T) { + const runs = 20 + cfg := Config{Runs: runs, Command: []string{"x"}, Scenarios: []Scenario{{ID: "a", Input: "i"}}} + + var out bytes.Buffer // runner serializes flushes into this shared writer + opts := Options{ + Concurrency: 8, + Stdout: &out, Stderr: io.Discard, + Exec: func(_ context.Context, _, env []string, stdout, _ io.Writer) error { + id := envMap(env)[EnvRunID] + // Two separate writes: if flushing weren't buffered per invocation, + // another goroutine's block could land between them. + fmt.Fprintf(stdout, "<%s>", id) + fmt.Fprintf(stdout, "", id) + return nil + }, + } + if _, err := Run(context.Background(), cfg, opts); err != nil { + t.Fatalf("Run: %v", err) + } + + got := out.String() + for i := range runs { + block := fmt.Sprintf("", i, i) + if !strings.Contains(got, block) { + t.Errorf("invocation %d output not contiguous; missing %q in:\n%s", i, block, got) + } + } +} + +// TestRunParallelAbortsOnError checks that without ContinueOnError the first +// genuine failure aborts the run and is the error returned. +func TestRunParallelAbortsOnError(t *testing.T) { + cfg := Config{Runs: 1, Command: []string{"x"}, Scenarios: []Scenario{ + {ID: "boom", Input: "i"}, + }} + opts := Options{ + Concurrency: 2, + Stdout: io.Discard, Stderr: io.Discard, + Exec: func(_ context.Context, _, _ []string, _, _ io.Writer) error { + return errors.New("agent crashed") + }, + } + sum, err := Run(context.Background(), cfg, opts) + if err == nil { + t.Fatal("expected the failure to be returned") + } + if !strings.Contains(err.Error(), "agent crashed") { + t.Errorf("error = %v, want it to wrap the agent failure", err) + } + if sum.Failed == 0 { + t.Errorf("Failed = 0, want >= 1") + } +} + +// TestRunParallelContinueOnError checks every invocation runs and failures are +// counted, not aborted, under concurrency. +func TestRunParallelContinueOnError(t *testing.T) { + var mu sync.Mutex + var n int + opts := Options{ + Concurrency: 4, + ContinueOnError: true, + Stdout: io.Discard, Stderr: io.Discard, + Exec: func(_ context.Context, _, env []string, _, _ io.Writer) error { + mu.Lock() + n++ + mu.Unlock() + // Fail exactly the third repetition of each scenario. + if strings.HasSuffix(envMap(env)[EnvRunID], "-002") { + return errors.New("flaky") + } + return nil + }, + } + sum, err := Run(context.Background(), twoScenarios(), opts) + if err != nil { + t.Fatalf("ContinueOnError should not return error: %v", err) + } + if n != 6 || sum.Total != 6 { + t.Errorf("n=%d Total=%d, want 6 (kept going)", n, sum.Total) + } + if sum.Failed != 2 { // checkout-002 and faq-002 + t.Errorf("Failed = %d, want 2", sum.Failed) + } +} From a7fef9dd5a5175fef6fd0391d694c789c6a86dde Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jesus=20Nu=C3=B1ez?= Date: Mon, 31 Aug 2026 20:26:31 -0400 Subject: [PATCH 2/2] dist: prebuilt binaries via GoReleaser + zero-toolchain GitHub Action Ship augur as prebuilt, checksum-verified binaries so agent repos (Python/TS/ Node) run the cost gate without a Go toolchain in CI. - .goreleaser.yaml: cross-compile linux/darwin/windows x amd64/arm64 (static, CGO disabled), stable version-free archive names (augur__.tar.gz), sha256 checksums.txt, and a GitHub release per v* tag. Build metadata is injected via -ldflags into main.version/commit/date. - version_cmd.go: `augur version` reports the injected build info (from-source builds honestly say "dev"); it's also the handle the Action uses to confirm a downloaded binary runs. Registered in the dispatch table + usage order. - .github/workflows/release.yml: on a v* tag, run tests then GoReleaser (release --clean) to publish the binaries. CI gains a goreleaser-check job so the config is validated on every push. - action.yml: new "Resolve augur binary" step downloads the release asset matching the pinned tag (github.action_ref) and runner os/arch, verifies its sha256 against checksums.txt (tolerant of coreutils/BSD format variants), and extracts it. Falls back to the existing source build for branch/SHA pins, unbuilt platforms, or a download failure. New `version` input forces a tag or `source`. The from-source path is unchanged, just gated behind need-build. - README: document prebuilt distribution, the version input, and how to cut a release. .gitignore: ignore GoReleaser's /dist/. The download/checksum/os-arch/version-resolution logic was simulated locally (valid + corrupt archive, binary- and two-space checksum formats, tag/branch/ source refs). GoReleaser config is schema-checked in CI. Full suite green. --- .github/workflows/ci.yml | 14 ++++++ .github/workflows/release.yml | 36 ++++++++++++++ .gitignore | 2 + .goreleaser.yaml | 74 ++++++++++++++++++++++++++++ README.md | 14 +++++- action.yml | 90 ++++++++++++++++++++++++++++++++++- main.go | 3 +- version_cmd.go | 33 +++++++++++++ version_cmd_test.go | 44 +++++++++++++++++ 9 files changed, 306 insertions(+), 4 deletions(-) create mode 100644 .github/workflows/release.yml create mode 100644 .goreleaser.yaml create mode 100644 version_cmd.go create mode 100644 version_cmd_test.go diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index c53b7c9..a27e318 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -19,3 +19,17 @@ jobs: run: go vet ./... - name: Test run: go test ./... -count=1 + + goreleaser-check: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-go@v5 + with: + go-version: "1.25" + - name: Validate GoReleaser config + uses: goreleaser/goreleaser-action@v6 + with: + distribution: goreleaser + version: "~> v2" + args: check diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml new file mode 100644 index 0000000..b0a5505 --- /dev/null +++ b/.github/workflows/release.yml @@ -0,0 +1,36 @@ +name: Release + +# Cut a release by pushing a semver tag, e.g.: +# git tag v1.2.0 && git push origin v1.2.0 +# GoReleaser then cross-compiles augur and uploads the binaries + checksums to +# the GitHub release, which action.yml downloads instead of building from source. +on: + push: + tags: + - "v*" + +permissions: + contents: write # required to create the release and upload assets + +jobs: + goreleaser: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + with: + fetch-depth: 0 # GoReleaser needs full history for the changelog + + - uses: actions/setup-go@v5 + with: + go-version: "1.25" + + - name: Run tests before releasing + run: go test ./... -count=1 + + - uses: goreleaser/goreleaser-action@v6 + with: + distribution: goreleaser + version: "~> v2" + args: release --clean + env: + GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} diff --git a/.gitignore b/.gitignore index 11fa760..4f3ce19 100644 --- a/.gitignore +++ b/.gitignore @@ -30,6 +30,8 @@ go.work.sum # Augur build + generated artifacts (keep the repo clean of local runs) /augur +# GoReleaser build output +/dist/ /augur-report.md /augur-report.json /trace.jsonl diff --git a/.goreleaser.yaml b/.goreleaser.yaml new file mode 100644 index 0000000..8a2dea5 --- /dev/null +++ b/.goreleaser.yaml @@ -0,0 +1,74 @@ +# GoReleaser builds Augur for every platform CI agents run on and publishes the +# binaries (plus a checksums file) to the GitHub release for each `v*` tag. The +# GitHub Action (action.yml) then downloads the matching prebuilt binary instead +# of compiling from source, so agent repos written in Python/TS/Node don't need +# a Go toolchain in CI. +# +# Validate locally with `goreleaser check`; it runs in CI on every push. +version: 2 + +project_name: augur + +before: + hooks: + - go mod tidy + +builds: + - id: augur + main: . + binary: augur + env: + - CGO_ENABLED=0 # a static binary needs no libc on the runner + goos: + - linux + - darwin + - windows + goarch: + - amd64 + - arm64 + # Inject build metadata so `augur version` reports the release (see + # version_cmd.go). -s -w strip the symbol table to shrink the binary. + ldflags: + - -s -w + - -X main.version={{ .Version }} + - -X main.commit={{ .Commit }} + - -X main.date={{ .Date }} + +archives: + - id: augur + ids: + - augur + # Stable, version-free archive names: the release tag is already in the + # download URL path, so the Action can predict the asset name from just the + # runner's os/arch (see action.yml). + name_template: "augur_{{ .Os }}_{{ .Arch }}" + formats: + - tar.gz + files: + - LICENSE + - README.md + +checksum: + name_template: "checksums.txt" + algorithm: sha256 + +snapshot: + version_template: "{{ incpatch .Version }}-snapshot" + +release: + github: + owner: Cro22 + name: Augur + draft: false + prerelease: auto + +changelog: + sort: asc + filters: + exclude: + - "^docs:" + - "^test:" + - "^chore:" + - "^ci:" + - "Merge pull request" + - "Merge branch" diff --git a/README.md b/README.md index eaf9e15..c92ceb0 100644 --- a/README.md +++ b/README.md @@ -266,17 +266,29 @@ jobs: pull-requests: write # to post the report comment steps: - uses: actions/checkout@v4 - - uses: Cro22/augur@v1 + - uses: Cro22/augur@v1.1.0 # pin a release tag → prebuilt binary with: cassette: cassette.jsonl # replay → zero tokens traffic: traffic.yaml budget: budget.yaml ``` +**No Go toolchain required.** When you pin the action to a release tag +(`@v1.1.0`), it downloads the matching prebuilt binary from the release and +checksum-verifies it — so a Python/TS/Node agent repo runs the gate without a Go +build. Pin to a branch or SHA (`@master`) and it transparently falls back to +building from source; force either mode with the `version` input (`source`, or a +specific tag). Binaries are produced by [GoReleaser](.goreleaser.yaml) on every +`v*` tag ([release workflow](.github/workflows/release.yml)). + See [`action.yml`](action.yml) for all inputs and [`examples/github-workflow.yml`](examples/github-workflow.yml) for a fuller example. +**Cutting a release** (maintainers): `git tag v1.2.0 && git push origin v1.2.0` +triggers GoReleaser to cross-compile (linux/darwin/windows × amd64/arm64) and +attach the binaries + `checksums.txt` to the GitHub release. + --- ## Design decisions diff --git a/action.yml b/action.yml index 6ea257b..3a1d1f9 100644 --- a/action.yml +++ b/action.yml @@ -35,8 +35,11 @@ inputs: working-directory: description: "Directory the config files live in." default: "." + version: + description: "Which augur release to use. 'auto' downloads the prebuilt binary matching the tag you pinned the action to (e.g. @v1.2.0); a specific tag (v1.2.0) forces that release; 'source' always builds from source." + default: "auto" go-version: - description: "Go version used to build augur." + description: "Go version for the from-source fallback build (used when no prebuilt binary is available)." default: "1.25" outputs: @@ -50,12 +53,95 @@ outputs: runs: using: "composite" steps: + # Prefer a prebuilt binary from the GitHub release so agent repos (Python/TS/ + # Node) don't need a Go toolchain in CI. Falls back to a source build when no + # matching release asset exists (a branch/SHA pin, an unbuilt platform, or a + # download failure) — need-build drives the two steps that follow. + - name: Resolve augur binary + id: resolve + shell: bash + run: | + set -euo pipefail + want='${{ inputs.version }}' + ref='${{ github.action_ref }}' + + # Decide which release tag to fetch (empty => build from source). + version="" + if [ "$want" = "source" ]; then + version="" + elif [ "$want" = "auto" ] || [ -z "$want" ]; then + case "$ref" in + v[0-9]*) version="$ref" ;; # pinned to a release tag + *) version="" ;; # branch/SHA => no release asset + esac + else + version="$want" # explicit tag + fi + + if [ -z "$version" ]; then + echo "augur: no prebuilt release for ref='$ref' (version input '$want'); building from source." + echo "need-build=true" >> "$GITHUB_OUTPUT" + exit 0 + fi + + # Map the runner to the os/arch naming GoReleaser used for the assets. + case "${RUNNER_OS}" in + Linux) os=linux ;; + macOS) os=darwin ;; + Windows) os=windows ;; + *) echo "augur: unsupported RUNNER_OS=$RUNNER_OS; building from source."; echo "need-build=true" >> "$GITHUB_OUTPUT"; exit 0 ;; + esac + case "${RUNNER_ARCH}" in + X64) arch=amd64 ;; + ARM64) arch=arm64 ;; + *) echo "augur: unsupported RUNNER_ARCH=$RUNNER_ARCH; building from source."; echo "need-build=true" >> "$GITHUB_OUTPUT"; exit 0 ;; + esac + + base="https://github.com/Cro22/Augur/releases/download/${version}" + archive="augur_${os}_${arch}.tar.gz" + dest="${RUNNER_TEMP}/augur-dl" + mkdir -p "$dest" + + echo "augur: fetching ${archive} from release ${version}" + if ! curl -fsSL "${base}/${archive}" -o "${dest}/${archive}" \ + || ! curl -fsSL "${base}/checksums.txt" -o "${dest}/checksums.txt"; then + echo "augur: prebuilt download failed; falling back to a source build." + echo "need-build=true" >> "$GITHUB_OUTPUT" + exit 0 + fi + + # Verify the checksum before trusting the binary. Extract the expected + # hash with awk (tolerating a leading '*' binary-mode marker) and compare + # to the computed one, so it works with either sha256sum or shasum and is + # insensitive to the one-/two-space separator variants. + cd "$dest" + want_sum="$(awk -v f="${archive}" '{ n=$2; sub(/^\*/,"",n); if (n==f) print $1 }' checksums.txt)" + if command -v sha256sum >/dev/null 2>&1; then + got_sum="$(sha256sum "${archive}" | awk '{print $1}')" + else + got_sum="$(shasum -a 256 "${archive}" | awk '{print $1}')" + fi + if [ -z "$want_sum" ] || [ "$want_sum" != "$got_sum" ]; then + echo "::error::augur: checksum mismatch for ${archive} (expected '${want_sum}', got '${got_sum}')" + exit 1 + fi + tar -xzf "${archive}" + + bin="${dest}/augur" + [ -f "${dest}/augur.exe" ] && bin="${dest}/augur.exe" + chmod +x "$bin" 2>/dev/null || true + echo "AUGUR=${bin}" >> "$GITHUB_ENV" + echo "need-build=false" >> "$GITHUB_OUTPUT" + echo "augur: using prebuilt $("$bin" version | head -1)" + - name: Set up Go + if: ${{ steps.resolve.outputs.need-build == 'true' }} uses: actions/setup-go@v5 with: go-version: ${{ inputs.go-version }} - - name: Build augur + - name: Build augur (source fallback) + if: ${{ steps.resolve.outputs.need-build == 'true' }} shell: bash run: | cd "${{ github.action_path }}" diff --git a/main.go b/main.go index c05973c..1e11cee 100644 --- a/main.go +++ b/main.go @@ -41,10 +41,11 @@ var commands = map[string]command{ "project": {runProject, "project a trace to production unit economics with CIs"}, "gate": {runGate, "check a projection against budget.yaml (exit 1 if over)"}, "tco": {runTCO, "show effective $/Mtok for self-hosted models (TCO)"}, + "version": {runVersion, "print the augur version"}, } // order fixes the usage listing (maps don't iterate deterministically). -var order = []string{"proxy", "run", "aggregate", "project", "gate", "tco"} +var order = []string{"proxy", "run", "aggregate", "project", "gate", "tco", "version"} func main() { if len(os.Args) < 2 { diff --git a/version_cmd.go b/version_cmd.go new file mode 100644 index 0000000..3fd0f08 --- /dev/null +++ b/version_cmd.go @@ -0,0 +1,33 @@ +package main + +import ( + "fmt" + "runtime" +) + +// Build metadata, injected at release time via -ldflags (see .goreleaser.yaml): +// +// -X main.version=... -X main.commit=... -X main.date=... +// +// A plain `go build` leaves them at these defaults, so a from-source binary +// honestly reports itself as "dev". +var ( + version = "dev" + commit = "" + date = "" +) + +// runVersion prints the build version and, when the release build injected them, +// the commit and date, plus the Go toolchain and target platform. It is the +// handle the GitHub Action uses to confirm a downloaded prebuilt binary runs. +func runVersion(_ []string) error { + fmt.Printf("augur %s\n", version) + if commit != "" { + fmt.Printf(" commit: %s\n", commit) + } + if date != "" { + fmt.Printf(" built: %s\n", date) + } + fmt.Printf(" go: %s %s/%s\n", runtime.Version(), runtime.GOOS, runtime.GOARCH) + return nil +} diff --git a/version_cmd_test.go b/version_cmd_test.go new file mode 100644 index 0000000..a9b5bf3 --- /dev/null +++ b/version_cmd_test.go @@ -0,0 +1,44 @@ +package main + +import ( + "slices" + "strings" + "testing" +) + +// TestVersionCommandRegistered guards the dispatch wiring: `augur version` must +// resolve to runVersion and appear in the usage order. +func TestVersionCommandRegistered(t *testing.T) { + if _, ok := commands["version"]; !ok { + t.Fatal("version command not registered in the dispatch table") + } + if !slices.Contains(order, "version") { + t.Error("version missing from the usage order slice") + } +} + +// TestVersionDefault checks a from-source build reports itself as "dev" (the +// release build overrides this via -ldflags). +func TestVersionDefault(t *testing.T) { + if version != "dev" { + t.Errorf("default version = %q, want dev", version) + } +} + +// TestRunVersion checks the command runs cleanly with any args. +func TestRunVersion(t *testing.T) { + if err := runVersion(nil); err != nil { + t.Fatalf("runVersion: %v", err) + } + if err := runVersion([]string{"ignored"}); err != nil { + t.Fatalf("runVersion with args: %v", err) + } +} + +// TestVersionSummaryNoTabs keeps the usage table aligned: the summary must be a +// single line. +func TestVersionSummaryNoTabs(t *testing.T) { + if s := commands["version"].summary; strings.ContainsAny(s, "\n\t") { + t.Errorf("version summary has newline/tab: %q", s) + } +}