From 2190ff39cbd9741cad00090375ab4c73cf06d0e6 Mon Sep 17 00:00:00 2001 From: Rene Zander Date: Sat, 3 Oct 2026 11:21:23 +0000 Subject: [PATCH] Prepare v0.9.0: durable budgets and tool execution lifecycle --- .gitignore | 2 + CHANGELOG.md | 14 + README.md | 4 +- approval_storage.go | 18 + approval_storage_test.go | 95 ++++ budget.go | 312 +++++++++++++ budget_cmd.go | 160 +++++++ budget_cmd_test.go | 68 +++ budget_governance_test.go | 355 +++++++++++++++ docs/governance-lifecycle.md | 75 ++++ docs/releasing.md | 19 + internal/config/config.go | 23 +- internal/outputschema/schema.go | 210 +++++++++ internal/outputschema/schema_test.go | 62 +++ internal/outputschema/yaml.go | 163 +++++++ internal/outputschema/yaml_test.go | 61 +++ internal/skills/skills.go | 11 +- internal/state/budget.go | 121 +++++ internal/state/budget_recovery.go | 35 ++ internal/state/budget_test.go | 146 ++++++ internal/state/state.go | 30 +- internal/state/tool_lifecycle.go | 188 ++++++++ internal/state/tool_lifecycle_test.go | 159 +++++++ internal/validate/validate.go | 34 +- .../validate/validate_output_schema_test.go | 36 ++ llm_client_test.go | 6 +- main.go | 406 +++++------------ model_policy.go | 21 +- output_validation.go | 91 ++++ output_validation_test.go | 152 +++++++ package-lock.json | 4 +- package.json | 2 +- provider_bounds_test.go | 94 ++++ provider_call.go | 127 ++++++ tool_gate.go | 421 ++++++++---------- tool_lifecycle.go | 211 +++++++++ tool_lifecycle_test.go | 291 ++++++++++++ version.go | 2 +- 38 files changed, 3626 insertions(+), 603 deletions(-) create mode 100644 approval_storage.go create mode 100644 approval_storage_test.go create mode 100644 budget.go create mode 100644 budget_cmd.go create mode 100644 budget_cmd_test.go create mode 100644 budget_governance_test.go create mode 100644 docs/governance-lifecycle.md create mode 100644 docs/releasing.md create mode 100644 internal/outputschema/schema.go create mode 100644 internal/outputschema/schema_test.go create mode 100644 internal/outputschema/yaml.go create mode 100644 internal/outputschema/yaml_test.go create mode 100644 internal/state/budget.go create mode 100644 internal/state/budget_recovery.go create mode 100644 internal/state/budget_test.go create mode 100644 internal/state/tool_lifecycle.go create mode 100644 internal/state/tool_lifecycle_test.go create mode 100644 internal/validate/validate_output_schema_test.go create mode 100644 output_validation.go create mode 100644 output_validation_test.go create mode 100644 provider_bounds_test.go create mode 100644 provider_call.go create mode 100644 tool_lifecycle.go create mode 100644 tool_lifecycle_test.go diff --git a/.gitignore b/.gitignore index 53a8735..adda706 100644 --- a/.gitignore +++ b/.gitignore @@ -13,3 +13,5 @@ state.db-wal aiops-architecture-diagram.py draftyard /npm/bin/ + +.build/ diff --git a/CHANGELOG.md b/CHANGELOG.md index 79f8315..a2dc211 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,19 @@ # Changelog +## 0.9.0 + +- Isolate pipeline token and cost totals and serialize model admission against settled daily usage. +- Persist daily usage by UTC date and preserve uncertain provider calls across restarts. +- Govern every engine model request, including classification and rewrites, and charge responses before output policy review. +- Require durable approval and receipt writes before releasing pipeline, model, or tool actions. +- Record immutable, action-bound external execution outcomes with authenticated completion and polling. +- Apply expiration and current policy checks consistently to live and recovered tool permits. +- Add authenticated revocation of pending and allowed tool actions with conditional transitions against consumption. +- Validate structured output and configured scalar schemas with exact numeric comparisons and explicit integer semantics. +- Bound provider response reads and stop automatic retries when billing is uncertain. + +`draftcat budget status` inspects daily usage and unsettled calls; `draftcat budget reconcile` records verified provider usage after an interrupted call. See the [budget and lifecycle guide](docs/governance-lifecycle.md) for upgrade behavior, recovery, and caller-attested outcome semantics. + ## 0.8.0 - Add durable, pipeline-scoped `Idempotency-Key` webhook retries. Matching bodies return the original admission; changed bodies are rejected. diff --git a/README.md b/README.md index 9e15eb7..c4c1ae2 100644 --- a/README.md +++ b/README.md @@ -65,6 +65,8 @@ Sometimes a customer, auditor, or partner needs evidence that a human approved a This is an **experimental cryptographic preview**, not a production compliance claim. It uses an embedded BN254/Groth16 circuit and a development single-party setup; the circuit has not received an independent audit. Use it to evaluate the disclosure model, then replace the setup through a ceremony before relying on it in production. See [zero-knowledge approval proofs](docs/zk-approval-proofs.md) for the trust model, exact statement, and limitations. +> **New in v0.9.0:** daily model usage persists in SQLite, parallel pipelines keep separate budget totals, and every engine model call passes one admission gate. Approval records commit before work is released. Unconsumed tool actions can be revoked, consumed actions accept durable caller-reported outcomes, and live or recovered permits share expiration and policy checks. Exact scalar output validation and bounded provider reads complete the release. See the [budget and lifecycle guide](docs/governance-lifecycle.md). +> > **New in v0.8.0:** webhook retries can carry a durable `Idempotency-Key`, signed replay identities are claimed atomically, and tool permits recheck their policy before execution. Requests reject oversized or ambiguous data and retain exact numbers. Older state stores upgrade safely; completed and failed runs carry exact approval identities. Audit commands read without modifying the database, and `draftcat receipts verify` checks exported JSONL offline. See the [upgrade and reliability guide](docs/reliability.md). > > **New in v0.7.0:** execution decisions now carry their proof. Every tool-gate route is authenticated, each request has a stable action identity and exact policy binding, and an allowed decision becomes an atomic consume-once permit before the side effect runs. Webhook acceptance is durable before HTTP 202 and can be polled after handoff. Versioned receipts bind action, payload, policy, and expiry, with `draftcat receipts list|show|export` for verification-ready JSONL. Ordered `model_policy` rules can deny or send matching model input/output to a human, while `/healthz` and `/readyz` give orchestrators a safe listener contract. @@ -295,7 +297,7 @@ An approval step can narrow who may decide it: approvers: [111111, 222222] # which ones (subset of allowed_users) ``` -Cost caps are checked between calls: a call is refused once spend has reached the cap. Pair them with `per_step_tokens` to bound the size of any single call. A transient provider failure (429, 408, 5xx) is retried with backoff — honouring `Retry-After` — before it fails a step; `provider.max_retries` sets the budget. +Cost caps are checked between calls: a call is refused once spend has reached the cap. Pair them with `per_step_tokens` to bound the size of any single call. Explicit rate-limit rejections (HTTP 429 without declared usage) retry with backoff and honour `Retry-After`; `provider.max_retries` sets the attempt limit. Transport errors, timeouts, server errors, and responses with uncertain billing halt the call and require usage reconciliation before another dispatch. See the [budget recovery guide](docs/governance-lifecycle.md). The tool-call gate is configured the same way, per tool: diff --git a/approval_storage.go b/approval_storage.go new file mode 100644 index 0000000..ac4b794 --- /dev/null +++ b/approval_storage.go @@ -0,0 +1,18 @@ +package main + +import ( + "fmt" + "os" + + statestore "github.com/renezander030/draftcat/internal/state" +) + +func persistApprovalReceipt(envelope statestore.ApprovalEnvelope) error { + if state == nil { + return nil + } + if os.Getenv("DRAFTCAT_APPROVAL_SECRET") != "" && (envelope.Nonce == "" || envelope.Signature == "") { + return fmt.Errorf("signed approval receipt could not be created") + } + return state.RecordApprovalV2(envelope) +} diff --git a/approval_storage_test.go b/approval_storage_test.go new file mode 100644 index 0000000..4a58954 --- /dev/null +++ b/approval_storage_test.go @@ -0,0 +1,95 @@ +package main + +import ( + "context" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/renezander030/draftcat/internal/config" + statestore "github.com/renezander030/draftcat/internal/state" +) + +type storageApprovalChannel struct { + stubApprovalChannel + prompts, sends int +} + +func (s *storageApprovalChannel) Send(string) error { s.sends++; return nil } +func (s *storageApprovalChannel) SendForApproval(ctx context.Context, draft string, approvers []int64) (OperatorDecision, error) { + s.prompts++ + return s.stubApprovalChannel.SendForApproval(ctx, draft, approvers) +} + +func TestPipelineStorageFailureStopsRelease(t *testing.T) { + for _, table := range []string{"pending_approvals", "action_approvals"} { + t.Run(table, func(t *testing.T) { + st, err := statestore.OpenStateStore(filepath.Join(t.TempDir(), "state.db")) + if err != nil { + t.Fatal(err) + } + old := state + state = st + t.Cleanup(func() { state = old; _ = st.Close() }) + _, err = st.DB().ExecContext(context.Background(), "CREATE TRIGGER reject_storage BEFORE INSERT ON "+table+" BEGIN SELECT RAISE(ABORT,'storage unavailable'); END") + if err != nil { + t.Fatal(err) + } + cfg := &config.Config{Timeouts: config.TimeoutConfig{OperatorApproval: "1s"}} + ch := &storageApprovalChannel{stubApprovalChannel: stubApprovalChannel{action: "approve", id: 7}} + p := config.PipelineConfig{Name: "pipeline", Steps: []config.StepConfig{{Name: "review", Type: "approval"}, {Name: "notify", Type: "deterministic", Action: "notify"}}} + err = runPipeline(cfg, p, &BudgetTracker{dayStart: time.Now()}, ch, nil, nil) + if err == nil || !strings.Contains(err.Error(), "unavailable") { + t.Fatalf("failure must stop release: %v", err) + } + if ch.sends != 0 { + t.Fatal("pipeline executed after storage failure") + } + if table == "pending_approvals" && ch.prompts != 0 { + t.Fatal("prompt sent before durable pending write") + } + }) + } +} + +func TestAutomaticApprovalRequiresReceiptStorage(t *testing.T) { + st, err := statestore.OpenStateStore(filepath.Join(t.TempDir(), "state.db")) + if err != nil { + t.Fatal(err) + } + old := state + state = st + t.Cleanup(func() { state = old; _ = st.Close() }) + _, err = st.DB().ExecContext(context.Background(), "CREATE TRIGGER reject_receipts BEFORE INSERT ON action_approvals BEGIN SELECT RAISE(ABORT,'storage unavailable'); END") + if err != nil { + t.Fatal(err) + } + cfg := &config.Config{Policy: config.ApprovalPolicy{AutoApprove: []config.AutoApproveRule{{Risk: config.RiskLow}}}} + ch := &storageApprovalChannel{} + p := config.PipelineConfig{Name: "pipeline", Steps: []config.StepConfig{{Name: "review", Type: "approval", Risk: config.RiskLow}, {Name: "notify", Type: "deterministic", Action: "notify"}}} + if err := runPipeline(cfg, p, &BudgetTracker{dayStart: time.Now()}, ch, nil, nil); err == nil { + t.Fatal("automatic approval ignored receipt failure") + } + if ch.sends != 0 { + t.Fatal("automatic approval released next step") + } +} + +func TestModelReviewRequiresReceiptStorage(t *testing.T) { + st, err := statestore.OpenStateStore(filepath.Join(t.TempDir(), "state.db")) + if err != nil { + t.Fatal(err) + } + old, oldCh := state, opChan + state, opChan = st, &stubApprovalChannel{action: "approve", id: 7} + t.Cleanup(func() { state, opChan = old, oldCh; _ = st.Close() }) + _, err = st.DB().ExecContext(context.Background(), "CREATE TRIGGER reject_receipts BEFORE INSERT ON action_approvals BEGIN SELECT RAISE(ABORT,'storage unavailable'); END") + if err != nil { + t.Fatal(err) + } + cfg := &config.Config{Timeouts: config.TimeoutConfig{OperatorApproval: "1s"}, ModelPolicy: config.ModelPolicyConfig{Rules: []config.ModelPolicyRule{{ID: "review", Phase: "output", Pattern: "claim", Action: "review"}}}} + if err := enforceModelPolicy(context.Background(), cfg, "drafter", "output", "claim"); err == nil { + t.Fatal("model output released without durable approval receipt") + } +} diff --git a/budget.go b/budget.go new file mode 100644 index 0000000..e036d5c --- /dev/null +++ b/budget.go @@ -0,0 +1,312 @@ +package main + +import ( + "context" + "fmt" + "math" + "sync" + "time" + + "github.com/renezander030/draftcat/internal/config" + statestore "github.com/renezander030/draftcat/internal/state" +) + +// BudgetTracker shares the daily ledger and admission gate, while each run owns +// its pipeline totals. Cost caps remain stop thresholds checked between calls. +type BudgetTracker struct { + mu sync.Mutex + gateOnce sync.Once + gate chan struct{} + parent *BudgetTracker + store *statestore.StateStore + stateErr error + unsettled int + tokensUsedToday int + tokensUsedPipeline int + callsToday int + callMinutesToday int + costToday float64 + costPipeline float64 + dayStart time.Time + dayCostLimit float64 + pipelineCostLimit float64 + pipelineTokenLimit int +} + +func (b *BudgetTracker) root() *BudgetTracker { + if b.parent != nil { + return b.parent.root() + } + return b +} + +func (b *BudgetTracker) newRun(tokenLimit int) *BudgetTracker { + r := b.root() + r.mu.Lock() + defer r.mu.Unlock() + return &BudgetTracker{parent: r, dayCostLimit: r.dayCostLimit, pipelineCostLimit: r.pipelineCostLimit, pipelineTokenLimit: tokenLimit} +} + +func (b *BudgetTracker) attachStore(store *statestore.StateStore) error { + r := b.root() + r.mu.Lock() + defer r.mu.Unlock() + r.store = store + r.refreshLocked(time.Now()) + return r.stateErr +} + +func (b *BudgetTracker) refreshLocked(now time.Time) { + r := b.root() + now = now.UTC() + if r.dayStart.UTC().Format("2006-01-02") != now.Format("2006-01-02") { + r.tokensUsedToday = 0 + r.costToday = 0 + r.callsToday = 0 + r.callMinutesToday = 0 + r.dayStart = time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, time.UTC) + } + if r.store != nil { + d, err := r.store.BudgetDay(context.Background(), now) + if err != nil { + r.stateErr = err + return + } + r.tokensUsedToday = d.Tokens + r.costToday = d.Cost + r.callsToday = d.Calls + r.callMinutesToday = d.CallMinutes + r.unsettled = d.Unsettled + } +} + +func (b *BudgetTracker) resetIfNewDay() { + r := b.root() + r.mu.Lock() + defer r.mu.Unlock() + r.refreshLocked(time.Now()) +} + +func (b *BudgetTracker) checkLocked(dayLimit, requested int, dayCost, pipelineCost float64) error { + r := b.root() + if r.stateErr != nil { + return fmt.Errorf("BUDGET_BLOCKED: usage ledger unavailable: %w", r.stateErr) + } + if r.unsettled > 0 { + return fmt.Errorf("BUDGET_BLOCKED: unresolved provider usage; reconcile pending calls before retrying") + } + if requested < 0 || r.tokensUsedToday < 0 || b.tokensUsedPipeline < 0 || !validCost(r.costToday) || !validCost(b.costPipeline) || !validCost(dayCost) || !validCost(pipelineCost) { + return fmt.Errorf("BUDGET_BLOCKED: invalid usage or budget limits") + } + if dayLimit > 0 && requested > dayLimit-r.tokensUsedToday { + return fmt.Errorf("BUDGET_BLOCKED: daily token limit %d would be exceeded (used: %d, requested: %d)", dayLimit, r.tokensUsedToday, requested) + } + if b.pipelineTokenLimit > 0 && requested > b.pipelineTokenLimit-b.tokensUsedPipeline { + return fmt.Errorf("BUDGET_BLOCKED: per-pipeline token limit %d would be exceeded (used: %d, requested: %d)", b.pipelineTokenLimit, b.tokensUsedPipeline, requested) + } + if dayCost > 0 && r.costToday >= dayCost { + return fmt.Errorf("BUDGET_BLOCKED: daily cost limit %.4f reached (spent: %.4f)", dayCost, r.costToday) + } + if pipelineCost > 0 && b.costPipeline >= pipelineCost { + return fmt.Errorf("BUDGET_BLOCKED: per-pipeline cost limit %.4f reached (spent: %.4f)", pipelineCost, b.costPipeline) + } + return nil +} + +func validCost(cost float64) bool { return cost >= 0 && !math.IsNaN(cost) && !math.IsInf(cost, 0) } + +func (b *BudgetTracker) check(limit, requested int) error { + r := b.root() + r.mu.Lock() + defer r.mu.Unlock() + r.refreshLocked(time.Now()) + pipelineCost := 0.0 + if b.parent != nil { + pipelineCost = b.pipelineCostLimit + } + return b.checkLocked(limit, requested, r.dayCostLimit, pipelineCost) +} + +func (b *BudgetTracker) CheckCost(dayLimit, pipelineLimit float64) error { + r := b.root() + r.mu.Lock() + defer r.mu.Unlock() + r.refreshLocked(time.Now()) + return b.checkLocked(0, 0, dayLimit, pipelineLimit) +} + +func (b *BudgetTracker) recordUsageLocked(tokens int, cost float64) { + r := b.root() + if tokens < 0 || !validCost(cost) || tokens > int(^uint(0)>>1)-r.tokensUsedToday || tokens > int(^uint(0)>>1)-b.tokensUsedPipeline || !validCost(r.costToday+cost) || !validCost(b.costPipeline+cost) { + r.stateErr = fmt.Errorf("invalid provider usage") + return + } + r.tokensUsedToday += tokens + r.costToday += cost + b.tokensUsedPipeline += tokens + b.costPipeline += cost +} + +func (b *BudgetTracker) record(tokens int) { + r := b.root() + r.mu.Lock() + defer r.mu.Unlock() + r.refreshLocked(time.Now()) + b.recordUsageLocked(tokens, 0) + if r.stateErr == nil && r.store != nil { + r.stateErr = r.store.AddBudgetUsage(context.Background(), time.Now(), tokens, 0, 0, 0) + } +} + +func (b *BudgetTracker) RecordCost(cost float64) { + r := b.root() + r.mu.Lock() + defer r.mu.Unlock() + r.refreshLocked(time.Now()) + b.recordUsageLocked(0, cost) + if r.stateErr == nil && r.store != nil { + r.stateErr = r.store.AddBudgetUsage(context.Background(), time.Now(), 0, cost, 0, 0) + } +} + +func (b *BudgetTracker) CheckCalls(limit int) error { + r := b.root() + r.mu.Lock() + defer r.mu.Unlock() + r.refreshLocked(time.Now()) + if err := b.checkLocked(0, 0, 0, 0); err != nil { + return err + } + if limit > 0 && r.callsToday >= limit { + return fmt.Errorf("BUDGET_BLOCKED: daily call limit %d would be exceeded (used: %d)", limit, r.callsToday) + } + return nil +} + +func (b *BudgetTracker) CheckCallMinutes(limit, requested int) error { + r := b.root() + r.mu.Lock() + defer r.mu.Unlock() + r.refreshLocked(time.Now()) + if err := b.checkLocked(0, 0, 0, 0); err != nil { + return err + } + if requested < 0 { + return fmt.Errorf("BUDGET_BLOCKED: invalid call minutes") + } + if limit > 0 && requested > limit-r.callMinutesToday { + return fmt.Errorf("BUDGET_BLOCKED: daily call-minute limit %d would be exceeded (used: %d, requested: %d)", limit, r.callMinutesToday, requested) + } + return nil +} + +func (b *BudgetTracker) RecordCall(minutes int) { + r := b.root() + r.mu.Lock() + defer r.mu.Unlock() + r.refreshLocked(time.Now()) + if minutes < 0 { + r.stateErr = fmt.Errorf("invalid call duration") + return + } + r.callsToday++ + r.callMinutesToday += minutes + if r.store != nil { + r.stateErr = r.store.AddBudgetUsage(context.Background(), time.Now(), 0, 0, 1, minutes) + } +} + +func (b *BudgetTracker) snapshot() *BudgetTracker { return b.snapshotAt(time.Now()) } + +func (b *BudgetTracker) snapshotAt(now time.Time) *BudgetTracker { + r := b.root() + r.mu.Lock() + defer r.mu.Unlock() + r.refreshLocked(now) + return &BudgetTracker{tokensUsedToday: r.tokensUsedToday, tokensUsedPipeline: b.tokensUsedPipeline, costToday: r.costToday, costPipeline: b.costPipeline, callsToday: r.callsToday, callMinutesToday: r.callMinutesToday, dayStart: r.dayStart, dayCostLimit: r.dayCostLimit, pipelineCostLimit: b.pipelineCostLimit, stateErr: r.stateErr, unsettled: r.unsettled} +} + +func (b *BudgetTracker) acquire(ctx context.Context) error { + r := b.root() + r.gateOnce.Do(func() { r.gate = make(chan struct{}, 1) }) + if err := ctx.Err(); err != nil { + return err + } + select { + case r.gate <- struct{}{}: + return nil + case <-ctx.Done(): + return ctx.Err() + } +} +func (b *BudgetTracker) release() { <-b.root().gate } + +func (b *BudgetTracker) admit(ctx context.Context, cfg *config.Config, requested int) (string, error) { + if err := ctx.Err(); err != nil { + return "", err + } + r := b.root() + r.mu.Lock() + defer r.mu.Unlock() + r.refreshLocked(time.Now()) + r.dayCostLimit = cfg.Budgets.PerDayCost + r.pipelineCostLimit = cfg.Budgets.PerPipelineCost + pipelineCost := 0.0 + if b.parent != nil { + b.pipelineCostLimit = cfg.Budgets.PerPipelineCost + pipelineCost = b.pipelineCostLimit + } + if err := b.checkLocked(cfg.Budgets.PerDayTokens, requested, r.dayCostLimit, pipelineCost); err != nil { + return "", err + } + id := newRunID(time.Now()) + if r.store != nil { + if err := r.store.BeginBudgetCall(ctx, id, time.Now(), requested, cfg.Budgets.PerDayTokens, r.dayCostLimit); err != nil { + return "", err + } + } + r.unsettled++ + return id, nil +} + +func (b *BudgetTracker) finish(id string, resp *CompletionResponse, callErr error, cfg *config.Config) error { + r := b.root() + r.mu.Lock() + defer r.mu.Unlock() + settleCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + if resp != nil { + tokens := resp.InputTokens + resp.OutputTokens + if resp.InputTokens < 0 || resp.OutputTokens < 0 || tokens < 0 || !validCost(resp.CostUSD) { + return fmt.Errorf("BUDGET_BLOCKED: invalid provider usage; reconciliation required") + } + if r.store != nil { + if err := r.store.SettleBudgetCall(settleCtx, id, tokens, resp.CostUSD); err != nil { + r.stateErr = err + return fmt.Errorf("BUDGET_BLOCKED: could not settle provider usage: %w", err) + } + } + b.recordUsageLocked(tokens, resp.CostUSD) + r.unsettled-- + if r.stateErr != nil { + return fmt.Errorf("BUDGET_BLOCKED: %w", r.stateErr) + } + if cfg.Budgets.PerDayTokens > 0 && r.tokensUsedToday > cfg.Budgets.PerDayTokens { + return fmt.Errorf("BUDGET_BLOCKED: response exceeded daily token limit %d (used: %d)", cfg.Budgets.PerDayTokens, r.tokensUsedToday) + } + if b.pipelineTokenLimit > 0 && b.tokensUsedPipeline > b.pipelineTokenLimit { + return fmt.Errorf("BUDGET_BLOCKED: response exceeded per-pipeline token limit %d (used: %d)", b.pipelineTokenLimit, b.tokensUsedPipeline) + } + return nil + } + if !uncertainProviderError(callErr) { + if r.store != nil { + if err := r.store.ReleaseBudgetCall(settleCtx, id); err != nil { + r.stateErr = err + return fmt.Errorf("BUDGET_BLOCKED: could not release rejected request: %w", err) + } + } + r.unsettled-- + } + return nil +} diff --git a/budget_cmd.go b/budget_cmd.go new file mode 100644 index 0000000..5842964 --- /dev/null +++ b/budget_cmd.go @@ -0,0 +1,160 @@ +package main + +import ( + "context" + "encoding/json" + "fmt" + "math" + "os" + "strconv" + "strings" + "time" + + "github.com/renezander030/draftcat/internal/config" + statestore "github.com/renezander030/draftcat/internal/state" + "gopkg.in/yaml.v3" +) + +func runBudgetCmd(args []string) int { + if len(args) == 0 || args[0] == "help" || args[0] == "--help" || args[0] == "-h" { + fmt.Println("Usage: draftcat budget status [--json] [--config path]") + fmt.Println(" draftcat budget reconcile --tokens N --cost USD [--config path]") + fmt.Println("Reconcile only after stopping the engine and confirming actual provider usage.") + return 0 + } + cmd := args[0] + if cmd != "status" && cmd != "reconcile" { + fmt.Fprintf(os.Stderr, "budget: unknown command %q\n", cmd) + return 2 + } + path, id := "config.yaml", "" + tokens, cost := 0, 0.0 + hasTokens, hasCost, jsonOut := false, false, false + for i := 1; i < len(args); i++ { + arg := args[i] + switch arg { + case "--json": + jsonOut = true + case "--config", "--tokens", "--cost": + if i+1 >= len(args) { + fmt.Fprintf(os.Stderr, "budget: %s requires a value\n", arg) + return 2 + } + i++ + value := args[i] + switch arg { + case "--config": + path = value + case "--tokens": + n, err := strconv.Atoi(value) + if err != nil || n < 0 { + fmt.Fprintln(os.Stderr, "budget: tokens must be a nonnegative integer") + return 2 + } + tokens, hasTokens = n, true + case "--cost": + n, err := strconv.ParseFloat(value, 64) + if err != nil || n < 0 || math.IsNaN(n) || math.IsInf(n, 0) { + fmt.Fprintln(os.Stderr, "budget: cost must be a finite nonnegative amount") + return 2 + } + cost, hasCost = n, true + } + default: + if cmd != "reconcile" || id != "" || strings.HasPrefix(arg, "-") || !validActionID(arg) { + fmt.Fprintf(os.Stderr, "budget: unexpected argument %q\n", arg) + return 2 + } + id = arg + } + } + if cmd == "status" && (hasTokens || hasCost) { + fmt.Fprintln(os.Stderr, "budget: usage flags require reconcile") + return 2 + } + if cmd == "reconcile" && (id == "" || !hasTokens || !hasCost || jsonOut) { + fmt.Fprintln(os.Stderr, "budget: reconcile requires one call-id, --tokens and --cost") + return 2 + } + if cmd == "status" { + st, closeStore, code := openStateForCmd(path) + if code != 0 { + return code + } + defer closeStore() + return printBudgetStatus(st, jsonOut) + } + // This administrative write requires an existing regular database. It must + // never turn a misspelled path into a fresh store with an empty spend ledger. + statePath := strings.TrimSpace(os.Getenv("DRAFTCAT_STATE_PATH")) + if statePath == "" { + // #nosec G304 -- the local operator supplies their config path. + if data, err := os.ReadFile(path); err == nil { + var cfg config.Config + if yaml.Unmarshal(data, &cfg) == nil { + statePath = cfg.State.Path + } + } + } + if statePath == "" { + statePath = "./state.db" + } + info, err := os.Stat(statePath) + if err != nil || !info.Mode().IsRegular() { + fmt.Fprintln(os.Stderr, "budget: existing state database required") + return 1 + } + st, err := statestore.OpenStateStore(statePath) + if err != nil { + fmt.Fprintf(os.Stderr, "budget: open state: %v\n", err) + return 1 + } + defer func() { _ = st.Close() }() + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + if err := st.SettleBudgetCall(ctx, id, tokens, cost); err != nil { + fmt.Fprintf(os.Stderr, "budget: reconcile: %v\n", err) + return 1 + } + fmt.Printf("Reconciled %s: %d tokens, %.6f cost\n", id, tokens, cost) + return 0 +} + +func printBudgetStatus(st *statestore.StateStore, jsonOut bool) int { + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + now := time.Now().UTC() + day, err := st.BudgetDay(ctx, now) + if err != nil { + fmt.Fprintf(os.Stderr, "budget: read usage: %v\n", err) + return 1 + } + pending, err := st.PendingBudgetCalls(ctx) + if err != nil { + fmt.Fprintf(os.Stderr, "budget: read pending calls: %v\n", err) + return 1 + } + if jsonOut { + report := struct { + Day string `json:"day"` + Tokens int `json:"tokens"` + Cost float64 `json:"cost"` + Calls int `json:"calls"` + CallMinutes int `json:"call_minutes"` + Pending []statestore.PendingBudgetCall `json:"pending_calls"` + }{now.Format("2006-01-02"), day.Tokens, day.Cost, day.Calls, day.CallMinutes, pending} + if err := json.NewEncoder(os.Stdout).Encode(report); err != nil { + fmt.Fprintf(os.Stderr, "budget: output: %v\n", err) + return 1 + } + return 0 + } + fmt.Printf("%s UTC: %d tokens, %.6f cost, %d voice calls, %d call minutes\n", now.Format("2006-01-02"), day.Tokens, day.Cost, day.Calls, day.CallMinutes) + for _, call := range pending { + fmt.Printf("Unsettled provider call %s (day %s, started %s)\n", call.CallID, call.Day, call.CreatedAt.Format(time.RFC3339)) + } + if len(pending) > 0 { + fmt.Println("Stop the engine, confirm provider usage, then use budget reconcile before resuming.") + } + return 0 +} diff --git a/budget_cmd_test.go b/budget_cmd_test.go new file mode 100644 index 0000000..bd2f660 --- /dev/null +++ b/budget_cmd_test.go @@ -0,0 +1,68 @@ +package main + +import ( + "context" + "os" + "path/filepath" + "testing" + "time" + + statestore "github.com/renezander030/draftcat/internal/state" +) + +func TestBudgetCommandRecoveryChargesOriginalDayOnce(t *testing.T) { + path := filepath.Join(t.TempDir(), "state.db") + st, err := statestore.OpenStateStore(path) + if err != nil { + t.Fatal(err) + } + day := time.Date(2026, 9, 30, 23, 59, 0, 0, time.UTC) + if err := st.BeginBudgetCall(context.Background(), "call-uncertain", day, 100, 1000, 1); err != nil { + t.Fatal(err) + } + _ = st.Close() + t.Setenv("DRAFTCAT_STATE_PATH", path) + if code := runBudgetCmd([]string{"status", "--json"}); code != 0 { + t.Fatalf("status exit %d", code) + } + if code := runBudgetCmd([]string{"reconcile", "call-uncertain", "--tokens", "120", "--cost", "0.12"}); code != 0 { + t.Fatalf("reconcile exit %d", code) + } + if code := runBudgetCmd([]string{"reconcile", "call-uncertain", "--tokens", "120", "--cost", "0.12"}); code == 0 { + t.Fatal("duplicate reconciliation succeeded") + } + st, err = statestore.OpenStateStoreReadOnly(path) + if err != nil { + t.Fatal(err) + } + defer func() { _ = st.Close() }() + got, err := st.BudgetDay(context.Background(), day) + if err != nil || got.Tokens != 120 || got.Cost != 0.12 || got.Unsettled != 0 { + t.Fatalf("usage %+v error %v", got, err) + } + calls, err := st.PendingBudgetCalls(context.Background()) + if err != nil || len(calls) != 0 { + t.Fatalf("pending %+v error %v", calls, err) + } +} + +func TestBudgetCommandsDoNotCreateMissingState(t *testing.T) { + path := filepath.Join(t.TempDir(), "missing.db") + t.Setenv("DRAFTCAT_STATE_PATH", path) + for _, args := range [][]string{{"status"}, {"reconcile", "call", "--tokens", "0", "--cost", "0"}} { + if code := runBudgetCmd(args); code == 0 { + t.Fatalf("missing database accepted: %v", args) + } + if _, err := os.Stat(path); !os.IsNotExist(err) { + t.Fatalf("command created database: %v", err) + } + } +} + +func TestBudgetReconcileRejectsInvalidUsage(t *testing.T) { + for _, args := range [][]string{{"reconcile", "call", "--tokens", "-1", "--cost", "0"}, {"reconcile", "call", "--tokens", "1", "--cost", "NaN"}, {"reconcile", "call", "--tokens", "1"}, {"status", "--tokens", "1"}} { + if code := runBudgetCmd(args); code != 2 { + t.Fatalf("invalid arguments %v exit %d", args, code) + } + } +} diff --git a/budget_governance_test.go b/budget_governance_test.go new file mode 100644 index 0000000..caa5b5d --- /dev/null +++ b/budget_governance_test.go @@ -0,0 +1,355 @@ +package main + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "math" + "net/http" + "net/http/httptest" + "path/filepath" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/renezander030/draftcat/internal/config" + statestore "github.com/renezander030/draftcat/internal/state" +) + +func providerBudgetConfig(url string) *config.Config { + cfg := llmCfg(url, "openrouter") + cfg.Models["m"] = config.ModelConfig{Model: "x", MaxTokens: 1} + cfg.Roles["classifier"] = "m" + return cfg +} + +func providerBudgetStore(t *testing.T) *statestore.StateStore { + t.Helper() + st, err := statestore.OpenStateStore(filepath.Join(t.TempDir(), "state.db")) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = st.Close() }) + return st +} + +const smallPaidCompletion = `{"choices":[{"message":{"content":"hi"}}],"usage":{"prompt_tokens":1,"completion_tokens":1,"cost":0.6}}` + +func TestBudgetConcurrentAdmissionStopsAfterSettledCap(t *testing.T) { + var hits atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hits.Add(1) + _, _ = io.WriteString(w, smallPaidCompletion) + })) + defer srv.Close() + cfg := providerBudgetConfig(srv.URL) + cfg.Budgets.PerDayCost = .5 + b := &BudgetTracker{dayStart: time.Now()} + if err := b.attachStore(providerBudgetStore(t)); err != nil { + t.Fatal(err) + } + var wg sync.WaitGroup + var succeeded atomic.Int32 + for i := 0; i < 16; i++ { + wg.Add(1) + go func() { + defer wg.Done() + if _, err := callLLM(context.Background(), cfg, "drafter", "p", b.newRun(0)); err == nil { + succeeded.Add(1) + } + }() + } + wg.Wait() + if hits.Load() != 1 || succeeded.Load() != 1 { + t.Fatalf("provider hits=%d successes=%d, want exactly one paid call", hits.Load(), succeeded.Load()) + } + snapshot := b.snapshot() + if snapshot.costToday != .6 || snapshot.tokensUsedToday != 2 || snapshot.unsettled != 0 { + t.Fatalf("settled snapshot: %+v", snapshot) + } +} + +func TestBudgetRunCountersStayIndependent(t *testing.T) { + var hits atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hits.Add(1) + _, _ = io.WriteString(w, smallPaidCompletion) + })) + defer srv.Close() + cfg := providerBudgetConfig(srv.URL) + cfg.Budgets.PerPipelineCost = .5 + b := &BudgetTracker{dayStart: time.Now()} + a, c := b.newRun(10), b.newRun(10) + for _, run := range []*BudgetTracker{a, c} { + if _, err := callLLM(context.Background(), cfg, "drafter", "p", run); err != nil { + t.Fatal(err) + } + if _, err := callLLM(context.Background(), cfg, "drafter", "p", run); err == nil { + t.Fatal("run exceeded its own cost threshold") + } + } + if hits.Load() != 2 || a.snapshot().tokensUsedPipeline != 2 || c.snapshot().tokensUsedPipeline != 2 || b.snapshot().costToday != 1.2 { + t.Fatalf("hits=%d a=%+v c=%+v day=%+v", hits.Load(), a.snapshot(), c.snapshot(), b.snapshot()) + } +} + +func TestBudgetPipelineTokenOverrunChargesAndHaltsOutput(t *testing.T) { + var hits atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hits.Add(1) + _, _ = io.WriteString(w, `{"choices":[{"message":{"content":"hi"}}],"usage":{"prompt_tokens":5,"completion_tokens":1,"cost":0.6}}`) + })) + defer srv.Close() + cfg := providerBudgetConfig(srv.URL) + b := &BudgetTracker{dayStart: time.Now()} + run := b.newRun(5) + response, err := callLLM(context.Background(), cfg, "drafter", "p", run) + if response != nil || err == nil || !strings.Contains(err.Error(), "per-pipeline token limit") { + t.Fatalf("response=%+v err=%v", response, err) + } + if _, err := callLLM(context.Background(), cfg, "drafter", "p", run); err == nil { + t.Fatal("overrun was admitted again") + } + if hits.Load() != 1 || b.snapshot().tokensUsedToday != 6 || run.snapshot().tokensUsedPipeline != 6 { + t.Fatalf("lost charged usage: %+v", b.snapshot()) + } +} + +func TestBudgetDailyTokenOverrunChargesAndHaltsOutput(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { _, _ = io.WriteString(w, smallPaidCompletion) })) + defer srv.Close() + cfg := providerBudgetConfig(srv.URL) + cfg.Budgets.PerDayTokens = 1 + b := &BudgetTracker{dayStart: time.Now()} + response, err := callLLM(context.Background(), cfg, "classifier", "p", b) + if response != nil || err == nil || !strings.Contains(err.Error(), "daily token limit") || b.snapshot().tokensUsedToday != 2 { + t.Fatalf("response=%+v err=%v usage=%+v", response, err, b.snapshot()) + } +} + +func TestBudgetDeniedOutputStillChargesPaidUsage(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { _, _ = io.WriteString(w, smallPaidCompletion) })) + defer srv.Close() + cfg := providerBudgetConfig(srv.URL) + cfg.Budgets.PerDayCost = .5 + cfg.ModelPolicy.Rules = []config.ModelPolicyRule{{ID: "output-denied", Phase: "output", Pattern: "hi", Action: "deny", Reason: "blocked claim"}} + b := &BudgetTracker{dayStart: time.Now()} + st := providerBudgetStore(t) + if err := b.attachStore(st); err != nil { + t.Fatal(err) + } + if response, err := callLLM(context.Background(), cfg, "drafter", "p", b); response != nil || err == nil || !strings.Contains(err.Error(), "blocked by policy") { + t.Fatalf("response=%+v err=%v", response, err) + } + day, err := st.BudgetDay(context.Background(), time.Now()) + if err != nil || day.Cost != .6 || day.Tokens != 2 || day.Unsettled != 0 { + t.Fatalf("paid denied usage=%+v err=%v", day, err) + } + restarted := &BudgetTracker{dayStart: time.Now()} + if err := restarted.attachStore(st); err != nil { + t.Fatal(err) + } + if _, err := callLLM(context.Background(), cfg, "classifier", "intent", restarted); err == nil || !strings.Contains(err.Error(), "cost limit") { + t.Fatalf("restart allowed paid intent classification: %v", err) + } +} + +type budgetReviewChannel struct { + stubApprovalChannel + entered chan struct{} + done chan struct{} +} + +func (s *budgetReviewChannel) SendForApproval(ctx context.Context, _ string, _ []int64) (OperatorDecision, error) { + close(s.entered) + select { + case <-s.done: + return OperatorDecision{Action: "approve", ApproverID: 7}, nil + case <-ctx.Done(): + return OperatorDecision{}, ctx.Err() + } +} + +func TestBudgetPaidReviewReleasesAdmissionBeforeHumanDecision(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var request struct { + Messages []struct { + Content string `json:"content"` + } `json:"messages"` + } + _ = json.NewDecoder(r.Body).Decode(&request) + text := "hi" + if request.Messages[0].Content == "review" { + text = "guarantee" + } + _, _ = fmt.Fprintf(w, `{"choices":[{"message":{"content":%q}}],"usage":{"prompt_tokens":1,"completion_tokens":1,"cost":0.1}}`, text) + })) + defer srv.Close() + cfg := providerBudgetConfig(srv.URL) + cfg.Timeouts.OperatorApproval = "5s" + cfg.ModelPolicy.Rules = []config.ModelPolicyRule{{ID: "review", Phase: "output", Pattern: "guarantee", Action: "review"}} + previousState, previousChannel := state, opChan + state = nil + review := &budgetReviewChannel{entered: make(chan struct{}), done: make(chan struct{})} + opChan = review + t.Cleanup(func() { state, opChan = previousState, previousChannel }) + b := &BudgetTracker{dayStart: time.Now()} + first := make(chan error, 1) + go func() { _, err := callLLM(context.Background(), cfg, "drafter", "review", b); first <- err }() + select { + case <-review.entered: + case <-time.After(time.Second): + t.Fatal("output review did not start") + } + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + if _, err := callLLM(ctx, cfg, "classifier", "another run", b); err != nil { + close(review.done) + t.Fatalf("human review held model admission: %v", err) + } + if b.snapshot().costToday != .2 { + t.Fatalf("usage not settled before review: %+v", b.snapshot()) + } + close(review.done) + if err := <-first; err != nil { + t.Fatal(err) + } +} + +func TestBudgetInvalidUsageRequiresReconciliationAcrossRestart(t *testing.T) { + for name, usage := range map[string]string{"negative": `{"prompt_tokens":-10,"completion_tokens":1,"cost":0.1}`, "missing": `{"completion_tokens":1,"cost":0.1}`, "negative-cost": `{"prompt_tokens":1,"completion_tokens":1,"cost":-0.1}`} { + t.Run(name, func(t *testing.T) { + var hits atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hits.Add(1) + _, _ = fmt.Fprintf(w, `{"choices":[{"message":{"content":"hi"}}],"usage":%s}`, usage) + })) + defer srv.Close() + st := providerBudgetStore(t) + cfg := providerBudgetConfig(srv.URL) + b := &BudgetTracker{dayStart: time.Now()} + if err := b.attachStore(st); err != nil { + t.Fatal(err) + } + if _, err := callLLM(context.Background(), cfg, "drafter", "p", b); err == nil { + t.Fatal("invalid usage accepted") + } + restarted := &BudgetTracker{dayStart: time.Now()} + if err := restarted.attachStore(st); err != nil { + t.Fatal(err) + } + if _, err := callLLM(context.Background(), cfg, "classifier", "p", restarted); err == nil || !strings.Contains(err.Error(), "unresolved") { + t.Fatalf("unresolved call reopened admission: %v", err) + } + if hits.Load() != 1 { + t.Fatalf("invalid response retried %d times", hits.Load()) + } + }) + } +} + +func TestBudgetAdmissionWaitHonorsContext(t *testing.T) { + b := &BudgetTracker{} + if err := b.acquire(context.Background()); err != nil { + t.Fatal(err) + } + defer b.release() + ctx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond) + defer cancel() + if err := b.acquire(ctx); !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("waiting admission ignored context: %v", err) + } +} + +func TestBudgetInvalidLegacyUsageCannotReopenCap(t *testing.T) { + for _, cost := range []float64{-1, math.NaN(), math.Inf(1)} { + b := &BudgetTracker{dayStart: time.Now(), costToday: 1} + b.RecordCost(cost) + if err := b.CheckCost(2, 0); err == nil { + t.Fatalf("invalid cost %v reopened cap", cost) + } + } + b := &BudgetTracker{dayStart: time.Now(), tokensUsedToday: 10} + b.record(-10) + if err := b.check(100, 1); err == nil { + t.Fatal("negative tokens reopened token budget") + } +} + +func TestBudgetDayResetUsesFullUTCDate(t *testing.T) { + b := &BudgetTracker{dayStart: time.Date(2026, 9, 3, 0, 0, 0, 0, time.UTC), tokensUsedToday: 9, costToday: 1} + b.snapshotAt(time.Date(2026, 10, 3, 0, 0, 0, 0, time.UTC)) + if b.tokensUsedToday != 0 || b.costToday != 0 { + t.Fatal("same day-of-month in a new month kept old usage") + } +} + +func TestCallLLMMaxTokensRespectsPerStepCap(t *testing.T) { + var requested int + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var req struct { + MaxTokens int `json:"max_tokens"` + } + _ = json.NewDecoder(r.Body).Decode(&req) + requested = req.MaxTokens + _, _ = io.WriteString(w, smallPaidCompletion) + })) + defer srv.Close() + cfg := providerBudgetConfig(srv.URL) + cfg.Models["m"] = config.ModelConfig{Model: "x", MaxTokens: 50} + cfg.Budgets.PerStepTokens = 7 + if _, err := callLLM(context.Background(), cfg, "drafter", "p", &BudgetTracker{dayStart: time.Now()}); err != nil { + t.Fatal(err) + } + if requested != 7 { + t.Fatalf("max_tokens=%d, want per-step cap 7", requested) + } +} + +func TestBudgetPipelineUsesModelAllowanceWithoutSharedRunReset(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { _, _ = io.WriteString(w, smallPaidCompletion) })) + defer srv.Close() + cfg := providerBudgetConfig(srv.URL) + cfg.Budgets.PerStepTokens = 10 + cfg.Budgets.PerPipelineTokens = 3 + b := &BudgetTracker{dayStart: time.Now()} + pipeline := config.PipelineConfig{Name: "scoped", Steps: []config.StepConfig{{Name: "draft", Type: "ai", Role: "drafter", Prompt: "p"}}} + previous := state + state = nil + t.Cleanup(func() { state = previous }) + for i := 0; i < 2; i++ { + if err := runPipeline(cfg, pipeline, b, nil, nil, nil); err != nil { + t.Fatalf("run %d rejected actual model allowance: %v", i, err) + } + } + if b.snapshot().tokensUsedToday != 4 { + t.Fatalf("run restart lost shared daily usage: %+v", b.snapshot()) + } +} + +func TestBudgetInputPolicyRejectionDoesNotAdmitProvider(t *testing.T) { + var hits atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hits.Add(1) + _, _ = io.WriteString(w, smallPaidCompletion) + })) + defer srv.Close() + cfg := providerBudgetConfig(srv.URL) + cfg.ModelPolicy.Rules = []config.ModelPolicyRule{{ID: "deny-input", Phase: "input", Pattern: "private", Action: "deny"}} + b := &BudgetTracker{dayStart: time.Now()} + st := providerBudgetStore(t) + if err := b.attachStore(st); err != nil { + t.Fatal(err) + } + if _, err := callLLM(context.Background(), cfg, "drafter", "private", b); err == nil { + t.Fatal("input policy did not deny") + } + day, err := st.BudgetDay(context.Background(), time.Now()) + if err != nil || hits.Load() != 0 || day.Tokens != 0 || day.Cost != 0 || day.Unsettled != 0 { + t.Fatalf("rejected input spent/admitted usage: hits=%d usage=%+v err=%v", hits.Load(), day, err) + } +} diff --git a/docs/governance-lifecycle.md b/docs/governance-lifecycle.md new file mode 100644 index 0000000..916974b --- /dev/null +++ b/docs/governance-lifecycle.md @@ -0,0 +1,75 @@ +# Budgets and tool execution lifecycle + +Draftcat keeps model usage in SQLite, isolates each pipeline run's token and cost totals, and records tool authorization separately from external execution. These controls use the existing service and CLI. + +## Daily and per-run budgets + +Daily usage is keyed by the full UTC calendar date and survives engine restarts. Model calls enter one context-aware admission gate, check settled usage, and write a pending usage row before contacting the provider. Separate instances sharing the same database cannot admit a second model call while one has unsettled usage. Each pipeline run has its own `per_pipeline_tokens` and `per_pipeline_cost` counters. + +Every engine model request passes this boundary, including intent classification and operator-requested rewrites. Responses are charged before model output policy is applied, so a paid response remains counted when its content is denied or sent for human review. Human output review does not hold the provider admission gate. + +`per_step_tokens` bounds the requested completion length together with the selected model's `max_tokens`. Daily and pipeline token totals count both input and output usage reported by the provider. Token limits and money caps stop further calls based on settled usage; an admitted call can cross a threshold because the final prompt usage and charge are only known afterward. A money cap is a stop threshold rather than a prepaid dollar reservation. Configure small completion limits to bound generation work. + +Provider response bodies have a 4 MiB limit. Unreadable, oversized, malformed, or invalid-usage responses stop processing. Rate-limit responses can retry within the configured retry count and request deadline. Transport errors and responses with uncertain billing do not automatically retry as another paid request. + +## Inspect and reconcile uncertain usage + +A crash during a provider call, or a response without trustworthy usage, leaves a pending usage row. Subsequent model calls remain blocked until the operator verifies the actual usage. This applies even when the pending call began on an earlier UTC date. + +```bash +draftcat budget status --json +``` + +Stop the engine before reconciliation. Check the provider's usage record, then charge the actual total tokens and cost for the listed call: + +```bash +draftcat budget reconcile --tokens 1200 --cost 0.0042 +``` + +The charge belongs to the original call's UTC date. Reconciliation settles the call once; it cannot charge a completed call again. Use zero only when the provider confirms that no tokens or money were billed. Both commands accept `--config path`, and `DRAFTCAT_STATE_PATH` takes precedence. Inspection is read-only and a misspelled state path does not create a new database. Start the engine once to migrate an older database before inspecting the new budget tables. + +## Approval storage + +Pipeline approvals, automatic approval policies, and model policy reviews require their durable records to succeed before they release work. Tool decisions and permit consumption commit their receipt and state transition in one transaction. An unavailable database or failed audit write blocks release. + +Existing receipt versions and signatures remain valid. New tool consumption receipts retain their existing action, payload, policy, and expiration binding. This release adds budget and outcome tables without replacing historical receipts. + +## Revoke an unconsumed tool action + +Use the same authenticated tool-gate listener and the exact `binding_hash` returned with the action: + +```http +POST /gate/tool-call//revoke +Authorization: Bearer +Content-Type: application/json + +{"binding_hash":"sha256:"} +``` + +A pending or allowed action becomes `revoked`. A matching retry returns the same state. Revocation and consumption race through conditional database transitions: only one can win. A late human approval cannot restore a revoked action. Revocation cannot interrupt an external action whose permit was already consumed; that request receives HTTP 409. + +The service's bearer credential authorizes this route. Keep it in the trusted harness or operator application. When `webhook.require_signature` is enabled, lifecycle POST requests also need the existing signed-body header described in the [tool gate guide](tool-gate.md). + +## Record an external execution outcome + +After the harness successfully consumes a permit and executes its bound action, it can attest to the result: + +```http +POST /gate/tool-call//complete +Authorization: Bearer +Content-Type: application/json + +{"binding_hash":"sha256:","status":"succeeded","result_hash":"sha256:"} +``` + +`status` is `succeeded` or `failed`; `result_hash` is optional and must be a SHA-256 digest when supplied. Send a hash of the result rather than the result itself. The record is durable, immutable, and bound to the consumed action. Matching retries succeed; a different status or hash receives HTTP 409. Reporting an outcome never grants another execution permit. + +Authenticated polling includes `execution_status`: `not_started`, `unreported`, `succeeded`, or `failed`. Reported outcomes also include `completed_at`, optional `result_hash`, and `execution_evidence: caller_attestation`. This records the harness's assertion; Draftcat does not independently observe or prove the external side effect. A missing acknowledgment remains `unreported`, preserving the uncertainty after a harness crash. + +Live and recovered permit status use the same expiration and current policy checks. A permit expires at its recorded expiration instant. A policy change invalidates unconsumed authorization and requires a new action approval. + +## Exact structured output + +The existing flat `output_schema` supports `int`, `number`, `bool`, and `string` types, numeric `min`/`max`, and scalar `enum` values. Every declared field is required. Additional model output fields remain accepted for compatibility. + +Integer fields require mathematical integers, so `1.0` and `1e3` are valid while `1.5` is rejected. JSON numeric tokens and YAML schema numeric values are compared without a float64 round trip. Duplicate keys, trailing content, non-object output, unsupported constraints, malformed definitions, and non-scalar enum values are rejected. Startup validation checks schemas from both prompt skills and inline pipeline steps. This is Draftcat's flat schema contract; full JSON Schema keywords are not supported. diff --git a/docs/releasing.md b/docs/releasing.md new file mode 100644 index 0000000..2025f15 --- /dev/null +++ b/docs/releasing.md @@ -0,0 +1,19 @@ +# Releasing Draftcat + +Draftcat is a Go CLI and service. Version tags publish six native binaries with checksums to GitHub Releases, a container image to GHCR, and an npm installer that downloads the matching binary. The Go module uses the same version tag. This repository does not provide a Python package or an MCP protocol server to publish to PyPI or the MCP Registry. + +## Review and version + +Merge the reviewed release PR only after its checks pass. Keep `version.go`, `package.json`, and both version entries in `package-lock.json` consistent with the intended tag. The release workflow verifies this before building. Create the version tag on the reviewed merge commit; pushing `v*` triggers `.github/workflows/release.yml`. + +The workflow runs lean and voice tests, validates the configuration, tests the npm installer, and checks the package contents. It builds Linux, macOS, and Windows binaries for x64 and arm64. The npm job waits for the GitHub assets and smoke-tests their installer before publishing. + +## npm authentication + +The npm job has `id-token: write` and runs on a GitHub-hosted runner with Node 24. For trusted publishing, the `draftcat` package settings on npm must authorize GitHub owner `renezander030`, repository `draftcat`, workflow filename `release.yml`, and direct `npm publish`. No GitHub environment is configured in this job. npm requires CLI 11.5.1 or newer for this route. See the [npm trusted publishing documentation](https://docs.npmjs.com/trusted-publishers/). + +The workflow also accepts a configured `NPM_TOKEN` repository secret as a token fallback. A successful build or GitHub release does not establish npm publication: inspect the npm job separately. If it fails with `ENEEDAUTH`, verify the exact package-side trusted publisher settings or use an authenticated maintainer account. Do not publish a second copy when the version already exists. + +## Verify the published version + +Check that the release has all six compressed assets and `SHA256SUMS`, the container version tag exists, and npm reports the intended version. Install the exact npm version in a clean prefix and run `draftcat --version` to verify the downloaded binary. Confirm the Go module tag is discoverable through the Go proxy. Retain the release workflow URL and report any destination that failed independently. diff --git a/internal/config/config.go b/internal/config/config.go index 71daf61..9a63e6c 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -17,6 +17,7 @@ import ( "strings" "time" + "github.com/renezander030/draftcat/internal/outputschema" "gopkg.in/yaml.v3" ghlapi "github.com/renezander030/draftcat/internal/ghl" @@ -157,7 +158,7 @@ type ProviderConfig struct { // OpenAI-compatible endpoint (those reject unknown request fields). UsageAccounting *bool `yaml:"usage_accounting"` // MaxRetries bounds how many times one call is retried on a transient - // failure (429, 408, 5xx, network error). 0 = default (3). + // rate-limit rejection (429 without declared usage). 0 = default (3). MaxRetries int `yaml:"max_retries"` apiKey string } @@ -352,16 +353,16 @@ type PipelineConfig struct { } type StepConfig struct { - Name string `yaml:"name"` - Type string `yaml:"type"` // deterministic, ai, approval - Action string `yaml:"action"` // deterministic action name - Role string `yaml:"role"` - Skill string `yaml:"skill"` // reference to skills/.yaml - Prompt string `yaml:"prompt"` - Vars map[string]string `yaml:"vars"` // variables injected into skill prompt - Mode string `yaml:"mode"` - Channel string `yaml:"channel"` - OutputSchema map[string]interface{} `yaml:"output_schema"` + Name string `yaml:"name"` + Type string `yaml:"type"` // deterministic, ai, approval + Action string `yaml:"action"` // deterministic action name + Role string `yaml:"role"` + Skill string `yaml:"skill"` // reference to skills/.yaml + Prompt string `yaml:"prompt"` + Vars map[string]string `yaml:"vars"` // variables injected into skill prompt + Mode string `yaml:"mode"` + Channel string `yaml:"channel"` + OutputSchema outputschema.Schema `yaml:"output_schema"` // Quorum is the number of distinct human operators that must approve this // approval step before the action is released. 0 or 1 = single approver // (default, unchanged behavior). Only the telegram channel implements N>=2. diff --git a/internal/outputschema/schema.go b/internal/outputschema/schema.go new file mode 100644 index 0000000..1ad9763 --- /dev/null +++ b/internal/outputschema/schema.go @@ -0,0 +1,210 @@ +// Package outputschema validates Draftcat's flat output schema definitions and +// compares numeric values without rounding their decimal representation. +package outputschema + +import ( + "encoding/json" + "fmt" + "math" + "math/big" + "sort" + "strconv" + "strings" +) + +// Finding identifies an invalid field definition. Findings are sorted by field +// and constraint so startup and runtime report the same first failure. +type Finding struct { + Field string + Message string +} + +// Check validates the existing flat field schema: every declared field is +// required; type, numeric min/max, and scalar enum are its supported constraints. +func Check(schema map[string]interface{}) []Finding { + var findings []Finding + for _, field := range Fields(schema) { + add := func(format string, args ...interface{}) { + findings = append(findings, Finding{field, fmt.Sprintf(format, args...)}) + } + def, ok := schema[field].(map[string]interface{}) + if !ok { + add("definition must be a map") + continue + } + for _, key := range Fields(def) { + if key != "type" && key != "min" && key != "max" && key != "enum" { + add("unsupported constraint %q", key) + } + } + typeName := "" + if raw, present := def["type"]; present { + var valid bool + typeName, valid = raw.(string) + if !valid || (typeName != "int" && typeName != "number" && typeName != "bool" && typeName != "string") { + add("type must be int, number, bool, or string") + } + } else if _, present := def["enum"]; !present { + add("missing type or enum") + } + min, hasMin := def["min"] + max, hasMax := def["max"] + if hasMin || hasMax { + if typeName != "int" && typeName != "number" { + add("min/max require an int or number type") + } + minNum, minOK := Number(min) + maxNum, maxOK := Number(max) + if hasMin && !minOK { + add("min must be a finite number") + } + if hasMax && !maxOK { + add("max must be a finite number") + } + if hasMin && hasMax && minOK && maxOK && minNum.Cmp(maxNum) > 0 { + add("min exceeds max") + } + } + if raw, present := def["enum"]; present { + allowed, valid := raw.([]interface{}) + if !valid || len(allowed) == 0 { + add("enum must be a nonempty list of scalar values") + continue + } + for i, value := range allowed { + if !Scalar(value) { + add("enum[%d] must be a finite scalar value", i) + } else if typeName != "" && !MatchesType(value, typeName) { + add("enum[%d] does not match the declared type", i) + } + } + } + } + return findings +} + +// Fields returns map keys in a stable order. +func Fields(fields map[string]interface{}) []string { + keys := make([]string, 0, len(fields)) + for key := range fields { + keys = append(keys, key) + } + sort.Strings(keys) + return keys +} + +// Number accepts numeric values from YAML and JSON, retaining all integer and +// json.Number digits. Bounds on token size and exponent bound rational work. +func Number(value interface{}) (*big.Rat, bool) { + var token string + switch n := value.(type) { + case json.Number: + token = n.String() + case int: + token = strconv.FormatInt(int64(n), 10) + case int8: + token = strconv.FormatInt(int64(n), 10) + case int16: + token = strconv.FormatInt(int64(n), 10) + case int32: + token = strconv.FormatInt(int64(n), 10) + case int64: + token = strconv.FormatInt(n, 10) + case uint: + token = strconv.FormatUint(uint64(n), 10) + case uint8: + token = strconv.FormatUint(uint64(n), 10) + case uint16: + token = strconv.FormatUint(uint64(n), 10) + case uint32: + token = strconv.FormatUint(uint64(n), 10) + case uint64: + token = strconv.FormatUint(n, 10) + case float32: + if math.IsNaN(float64(n)) || math.IsInf(float64(n), 0) { + return nil, false + } + token = strconv.FormatFloat(float64(n), 'g', -1, 32) + case float64: + if math.IsNaN(n) || math.IsInf(n, 0) { + return nil, false + } + token = strconv.FormatFloat(n, 'g', -1, 64) + default: + return nil, false + } + if len(token) == 0 || len(token) > 4096 || !json.Valid([]byte(token)) { + return nil, false + } + // json.Valid also accepts literals and containers. Numeric inputs must + // start with a sign or digit before big.Rat sees them. + if token[0] != '-' && (token[0] < '0' || token[0] > '9') { + return nil, false + } + if pos := strings.IndexAny(token, "eE"); pos >= 0 { + exponent, err := strconv.Atoi(token[pos+1:]) + if err != nil || exponent < -4096 || exponent > 4096 { + return nil, false + } + } + return new(big.Rat).SetString(token) +} + +// MatchesType treats 1.0 and 1e3 as integers by value, and 1.5 as fractional. +func MatchesType(value interface{}, typeName string) bool { + switch typeName { + case "int", "number": + number, ok := Number(value) + return ok && (typeName == "number" || number.IsInt()) + case "bool": + _, ok := value.(bool) + return ok + case "string": + _, ok := value.(string) + return ok + } + return false +} + +func Scalar(value interface{}) bool { + if value == nil { + return true + } + switch value.(type) { + case bool, string: + return true + default: + _, ok := Number(value) + return ok + } +} + +// EnumContains compares only scalar values, so a model returning an object or +// array where an enum is expected is a validation failure, never a panic. +func EnumContains(allowed []interface{}, value interface{}) bool { + if number, ok := Number(value); ok { + for _, member := range allowed { + if other, ok := Number(member); ok && number.Cmp(other) == 0 { + return true + } + } + return false + } + for _, member := range allowed { + switch v := value.(type) { + case nil: + if member == nil { + return true + } + case string: + if other, ok := member.(string); ok && v == other { + return true + } + case bool: + if other, ok := member.(bool); ok && v == other { + return true + } + } + } + return false +} diff --git a/internal/outputschema/schema_test.go b/internal/outputschema/schema_test.go new file mode 100644 index 0000000..b45c5ac --- /dev/null +++ b/internal/outputschema/schema_test.go @@ -0,0 +1,62 @@ +package outputschema + +import ( + "encoding/json" + "math" + "strings" + "testing" +) + +func TestCheckRejectsUncheckedDefinitions(t *testing.T) { + for _, tc := range []struct { + name string + definition interface{} + message string + }{ + {"non-map", "int", "definition must be a map"}, + {"empty", map[string]interface{}{}, "missing type or enum"}, + {"unsupported type", map[string]interface{}{"type": "array"}, "type must be"}, + {"malformed type", map[string]interface{}{"type": 1}, "type must be"}, + {"unsupported constraint", map[string]interface{}{"type": "int", "minimum": 1}, "unsupported constraint"}, + {"wrong bound type", map[string]interface{}{"type": "string", "min": 1}, "min/max require"}, + {"string bound", map[string]interface{}{"type": "int", "min": "1"}, "min must be"}, + {"infinite bound", map[string]interface{}{"type": "number", "max": math.Inf(1)}, "max must be"}, + {"inverted bounds", map[string]interface{}{"type": "int", "min": 2, "max": 1}, "min exceeds max"}, + {"non-list enum", map[string]interface{}{"enum": "safe"}, "enum must be"}, + {"empty enum", map[string]interface{}{"enum": []interface{}{}}, "enum must be"}, + {"composite enum", map[string]interface{}{"enum": []interface{}{[]interface{}{1}}}, "finite scalar"}, + {"inconsistent enum", map[string]interface{}{"type": "int", "enum": []interface{}{1.5}}, "does not match"}, + } { + t.Run(tc.name, func(t *testing.T) { + findings := Check(map[string]interface{}{"value": tc.definition}) + for _, f := range findings { + if strings.Contains(f.Message, tc.message) { + return + } + } + t.Fatalf("wanted %q, got %+v", tc.message, findings) + }) + } +} + +func TestCheckAcceptsFlatScalarContract(t *testing.T) { + schema := map[string]interface{}{ + "score": map[string]interface{}{"type": "int", "min": 1, "max": 5, "enum": []interface{}{1, 3, 5}}, + "number": map[string]interface{}{"type": "number", "min": 0.1, "max": 1}, + "choice": map[string]interface{}{"enum": []interface{}{"yes", false, nil}}, + } + if f := Check(schema); len(f) != 0 { + t.Fatalf("valid flat schema: %+v", f) + } +} + +func TestExactNumberWorkIsBounded(t *testing.T) { + for _, value := range []interface{}{json.Number("1e999999999"), json.Number(strings.Repeat("9", 4097)), json.Number("1/2"), "1", math.NaN(), math.Inf(1)} { + if _, ok := Number(value); ok { + t.Errorf("accepted invalid or excessive numeric token %v", value) + } + } + if n, ok := Number(json.Number("9007199254740993")); !ok || n.Num().String() != "9007199254740993" { + t.Fatalf("large integer rounded: %v", n) + } +} diff --git a/internal/outputschema/yaml.go b/internal/outputschema/yaml.go new file mode 100644 index 0000000..219a8c0 --- /dev/null +++ b/internal/outputschema/yaml.go @@ -0,0 +1,163 @@ +package outputschema + +import ( + "encoding/json" + "fmt" + "math/big" + "regexp" + "strings" + + "gopkg.in/yaml.v3" +) + +// Schema is the existing flat field map. Its YAML decoder retains numeric +// constraint and enum digits instead of converting decimal scalars to float64. +type Schema map[string]interface{} + +func (schema *Schema) UnmarshalYAML(node *yaml.Node) error { + value, err := decodeYAMLValue(node, 0) + if err != nil { + return err + } + if value == nil { + *schema = nil + return nil + } + fields, ok := value.(map[string]interface{}) + if !ok { + return fmt.Errorf("output_schema must be a field map") + } + *schema = Schema(fields) + return nil +} + +func decodeYAMLValue(node *yaml.Node, depth int) (interface{}, error) { + if depth > 128 { + return nil, fmt.Errorf("output_schema YAML nesting exceeds 128 levels") + } + switch node.Kind { + case yaml.AliasNode: + if node.Alias == nil { + return nil, fmt.Errorf("invalid output_schema YAML alias") + } + return decodeYAMLValue(node.Alias, depth+1) + case yaml.MappingNode: + result := map[string]interface{}{} + seen := map[string]bool{} + var merges []*yaml.Node + for i := 0; i < len(node.Content); i += 2 { + keyNode := node.Content[i] + var key string + if err := keyNode.Decode(&key); err != nil { + return nil, fmt.Errorf("output_schema keys must be strings") + } + if seen[key] { + return nil, fmt.Errorf("duplicate output_schema YAML key") + } + seen[key] = true + if keyNode.Tag == "!!merge" { + merges = append(merges, node.Content[i+1]) + continue + } + value, err := decodeYAMLValue(node.Content[i+1], depth+1) + if err != nil { + return nil, err + } + result[key] = value + } + for _, merge := range merges { + value, err := decodeYAMLValue(merge, depth+1) + if err != nil { + return nil, err + } + members, sequence := value.([]interface{}) + if !sequence { + members = []interface{}{value} + } + for _, member := range members { + fields, ok := member.(map[string]interface{}) + if !ok { + return nil, fmt.Errorf("output_schema YAML merge must contain maps") + } + for key, value := range fields { + if _, present := result[key]; !present { + result[key] = value + } + } + } + } + return result, nil + case yaml.SequenceNode: + result := make([]interface{}, 0, len(node.Content)) + for _, child := range node.Content { + value, err := decodeYAMLValue(child, depth+1) + if err != nil { + return nil, err + } + result = append(result, value) + } + return result, nil + case yaml.ScalarNode: + if node.Tag == "!!int" || node.Tag == "!!float" { + return yamlNumber(node) + } + var value interface{} + if err := node.Decode(&value); err != nil { + return nil, fmt.Errorf("invalid output_schema YAML scalar") + } + return value, nil + default: + return nil, fmt.Errorf("invalid output_schema YAML value") + } +} + +var yamlDecimal = regexp.MustCompile(`^[+-]?(?:[0-9]+(?:\.[0-9]*)?|\.[0-9]+)(?:[eE][+-]?[0-9]+)?$`) + +func yamlNumber(node *yaml.Node) (json.Number, error) { + token := strings.ReplaceAll(node.Value, "_", "") + if len(token) > 4096 { + return "", fmt.Errorf("output_schema numeric scalar exceeds 4096 characters") + } + if node.Tag == "!!int" { + // Base zero preserves YAML's hexadecimal, binary, and octal integers. + number, ok := new(big.Int).SetString(token, 0) + if !ok { + return "", fmt.Errorf("invalid output_schema integer") + } + token = number.String() + } else { + if !yamlDecimal.MatchString(token) { + return "", fmt.Errorf("output_schema numbers must be finite decimals") + } + token = strings.TrimPrefix(token, "+") + sign := "" + if strings.HasPrefix(token, "-") { + sign = "-" + token = token[1:] + } + exponent := "" + if pos := strings.IndexAny(token, "eE"); pos >= 0 { + exponent = token[pos:] + token = token[:pos] + } + whole, fraction, decimal := strings.Cut(token, ".") + whole = strings.TrimLeft(whole, "0") + if whole == "" { + whole = "0" + } + if decimal { + if fraction == "" { + fraction = "0" + } + token = whole + "." + fraction + } else { + token = whole + } + token = sign + token + exponent + } + value := json.Number(token) + if _, ok := Number(value); !ok { + return "", fmt.Errorf("invalid or excessive output_schema numeric scalar") + } + return value, nil +} diff --git a/internal/outputschema/yaml_test.go b/internal/outputschema/yaml_test.go new file mode 100644 index 0000000..cb56f38 --- /dev/null +++ b/internal/outputschema/yaml_test.go @@ -0,0 +1,61 @@ +package outputschema + +import ( + "encoding/json" + "testing" + + "gopkg.in/yaml.v3" +) + +func TestSchemaRetainsYAMLNumericDigits(t *testing.T) { + for _, tc := range []struct{ input, want string }{ + {"0.100000000000000001", "0.100000000000000001"}, + {"9007199254740993", "9007199254740993"}, + {"0x20", "32"}, {"0b100000", "32"}, {"040", "32"}, {"0o40", "32"}, + {"+32", "32"}, {"01.2", "1.2"}, {".5", "0.5"}, {"1.", "1.0"}, + {"1.e2", "1.0e2"}, {"1E+03", "1E+03"}, {"-.5", "-0.5"}, + } { + t.Run(tc.input, func(t *testing.T) { + var schema Schema + if err := yaml.Unmarshal([]byte("value: {type: number, min: "+tc.input+", enum: ["+tc.input+"]}\n"), &schema); err != nil { + t.Fatal(err) + } + def := schema["value"].(map[string]interface{}) + for _, value := range []interface{}{def["min"], def["enum"].([]interface{})[0]} { + if value != json.Number(tc.want) { + t.Fatalf("YAML token %s changed to %#v", tc.input, value) + } + } + }) + } +} + +func TestSchemaAliasesAndMergesKeepExactNumbers(t *testing.T) { + input := []byte("first: &bounds {type: number, min: 0.100000000000000001}\nsecond: {<<: *bounds, max: 1e0}\n") + var schema Schema + if err := yaml.Unmarshal(input, &schema); err != nil { + t.Fatal(err) + } + second := schema["second"].(map[string]interface{}) + if second["min"] != json.Number("0.100000000000000001") || second["max"] != json.Number("1e0") { + t.Fatalf("merged bound digits changed: %+v", second) + } + if f := Check(schema); len(f) != 0 { + t.Fatalf("merged scalar contract: %+v", f) + } +} + +func TestSchemaRejectsDuplicateAndNonfiniteYAMLScalars(t *testing.T) { + for _, input := range []string{ + "value: {type: number, min: .inf}\n", + "value: {type: number, enum: [.nan]}\n", + "value: {type: int, min: 1, min: 2}\n", + "value: {type: int}\nvalue: {type: string}\n", + "value: &recursive {enum: [*recursive]}\n", + } { + var schema Schema + if err := yaml.Unmarshal([]byte(input), &schema); err == nil { + t.Errorf("accepted duplicate, nonfinite, or recursive schema: %q", input) + } + } +} diff --git a/internal/skills/skills.go b/internal/skills/skills.go index 6cadb69..039434a 100644 --- a/internal/skills/skills.go +++ b/internal/skills/skills.go @@ -8,15 +8,16 @@ import ( "os" "path/filepath" + "github.com/renezander030/draftcat/internal/outputschema" "gopkg.in/yaml.v3" ) type SkillDef struct { - Name string `yaml:"name"` - Description string `yaml:"description"` - Role string `yaml:"role"` - Prompt string `yaml:"prompt"` - OutputSchema map[string]interface{} `yaml:"output_schema"` + Name string `yaml:"name"` + Description string `yaml:"description"` + Role string `yaml:"role"` + Prompt string `yaml:"prompt"` + OutputSchema outputschema.Schema `yaml:"output_schema"` } // SkillRegistry loads and holds all skills from the skills/ directory. diff --git a/internal/state/budget.go b/internal/state/budget.go new file mode 100644 index 0000000..d3947f1 --- /dev/null +++ b/internal/state/budget.go @@ -0,0 +1,121 @@ +package state + +import ( + "context" + "database/sql" + "fmt" + "math" + "time" +) + +// BudgetDay holds settled usage for a full UTC calendar date. +type BudgetDay struct { + Tokens int + Cost float64 + Calls int + CallMinutes int + Unsettled int +} + +func budgetDate(at time.Time) string { return at.UTC().Format("2006-01-02") } + +func (s *StateStore) BudgetDay(ctx context.Context, at time.Time) (BudgetDay, error) { + var d BudgetDay + if s == nil || s.db == nil { + return d, fmt.Errorf("budget store unavailable") + } + err := s.db.QueryRowContext(ctx, `SELECT COALESCE((SELECT tokens FROM budget_days WHERE day=?),0), COALESCE((SELECT cost FROM budget_days WHERE day=?),0), COALESCE((SELECT calls FROM budget_days WHERE day=?),0), COALESCE((SELECT call_minutes FROM budget_days WHERE day=?),0), (SELECT COUNT(*) FROM budget_calls WHERE status='pending')`, budgetDate(at), budgetDate(at), budgetDate(at), budgetDate(at)).Scan(&d.Tokens, &d.Cost, &d.Calls, &d.CallMinutes, &d.Unsettled) + if err == nil && (d.Tokens < 0 || d.Cost < 0 || math.IsNaN(d.Cost) || math.IsInf(d.Cost, 0)) { + err = fmt.Errorf("invalid persisted budget usage") + } + return d, err +} + +// BeginBudgetCall atomically admits one provider call across store handles. +// A pending call survives a crash; no subsequent call can assume it was free. +// Cost is a stop threshold on settled spend, not an estimated dollar ceiling. +func (s *StateStore) BeginBudgetCall(ctx context.Context, id string, at time.Time, requested, tokenLimit int, costLimit float64) error { + if s == nil || s.db == nil { + return fmt.Errorf("budget store unavailable") + } + if id == "" || requested < 0 || math.IsNaN(costLimit) || math.IsInf(costLimit, 0) { + return fmt.Errorf("invalid budget admission") + } + result, err := s.db.ExecContext(ctx, `INSERT INTO budget_calls (call_id, day, status, created_at) + SELECT ?, ?, 'pending', ? WHERE NOT EXISTS (SELECT 1 FROM budget_calls WHERE status='pending') + AND (? <= 0 OR COALESCE((SELECT tokens FROM budget_days WHERE day=?),0) <= ? - ?) + AND (? <= 0 OR COALESCE((SELECT cost FROM budget_days WHERE day=?),0) < ?)`, id, budgetDate(at), at.Unix(), tokenLimit, budgetDate(at), tokenLimit, requested, costLimit, budgetDate(at), costLimit) + if err != nil { + return err + } + n, err := result.RowsAffected() + if err != nil { + return err + } + if n != 1 { + return fmt.Errorf("BUDGET_BLOCKED: daily limit reached or unresolved provider usage; reconcile the pending call before retrying") + } + return nil +} + +// SettleBudgetCall charges an admitted call exactly once, including denied +// output. Reconciliation may use this same method after verifying actual usage. +func (s *StateStore) SettleBudgetCall(ctx context.Context, id string, tokens int, cost float64) error { + if tokens < 0 || cost < 0 || math.IsNaN(cost) || math.IsInf(cost, 0) { + return fmt.Errorf("invalid provider usage") + } + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return err + } + defer func() { _ = tx.Rollback() }() + var day, status string + if err := tx.QueryRowContext(ctx, `SELECT day,status FROM budget_calls WHERE call_id=?`, id).Scan(&day, &status); err != nil { + return err + } + if status != "pending" { + return fmt.Errorf("provider call is already settled") + } + if err := addBudgetUsageTx(ctx, tx, day, tokens, cost, 0, 0); err != nil { + return err + } + if _, err := tx.ExecContext(ctx, `UPDATE budget_calls SET status='settled',tokens=?,cost=?,settled_at=? WHERE call_id=? AND status='pending'`, tokens, cost, time.Now().Unix(), id); err != nil { + return err + } + return tx.Commit() +} + +// ReleaseBudgetCall closes a request known not to have produced billable usage. +func (s *StateStore) ReleaseBudgetCall(ctx context.Context, id string) error { + _, err := s.db.ExecContext(ctx, `UPDATE budget_calls SET status='rejected',settled_at=? WHERE call_id=? AND status='pending'`, time.Now().Unix(), id) + return err +} + +func (s *StateStore) AddBudgetUsage(ctx context.Context, at time.Time, tokens int, cost float64, calls, minutes int) error { + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return err + } + defer func() { _ = tx.Rollback() }() + if err := addBudgetUsageTx(ctx, tx, budgetDate(at), tokens, cost, calls, minutes); err != nil { + return err + } + return tx.Commit() +} + +func addBudgetUsageTx(ctx context.Context, tx *sql.Tx, day string, tokens int, cost float64, calls, minutes int) error { + if tokens < 0 || cost < 0 || calls < 0 || minutes < 0 || math.IsNaN(cost) || math.IsInf(cost, 0) { + return fmt.Errorf("invalid budget usage") + } + var current BudgetDay + err := tx.QueryRowContext(ctx, `SELECT COALESCE((SELECT tokens FROM budget_days WHERE day=?),0), COALESCE((SELECT cost FROM budget_days WHERE day=?),0), COALESCE((SELECT calls FROM budget_days WHERE day=?),0), COALESCE((SELECT call_minutes FROM budget_days WHERE day=?),0)`, day, day, day, day).Scan(¤t.Tokens, ¤t.Cost, ¤t.Calls, ¤t.CallMinutes) + if err != nil { + return err + } + maxInt := int(^uint(0) >> 1) + if current.Tokens < 0 || current.Calls < 0 || current.CallMinutes < 0 || current.Cost < 0 || tokens > maxInt-current.Tokens || calls > maxInt-current.Calls || minutes > maxInt-current.CallMinutes || math.IsNaN(current.Cost+cost) || math.IsInf(current.Cost+cost, 0) { + return fmt.Errorf("budget usage overflow or invalid stored totals") + } + _, err = tx.ExecContext(ctx, `INSERT INTO budget_days (day,tokens,cost,calls,call_minutes) VALUES (?,?,?,?,?) ON CONFLICT(day) DO UPDATE SET tokens=excluded.tokens,cost=excluded.cost,calls=excluded.calls,call_minutes=excluded.call_minutes`, day, current.Tokens+tokens, current.Cost+cost, current.Calls+calls, current.CallMinutes+minutes) + return err +} diff --git a/internal/state/budget_recovery.go b/internal/state/budget_recovery.go new file mode 100644 index 0000000..502165b --- /dev/null +++ b/internal/state/budget_recovery.go @@ -0,0 +1,35 @@ +package state + +import ( + "context" + "fmt" + "time" +) + +type PendingBudgetCall struct { + CallID string `json:"call_id"` + Day string `json:"day"` + CreatedAt time.Time `json:"created_at"` +} + +func (s *StateStore) PendingBudgetCalls(ctx context.Context) ([]PendingBudgetCall, error) { + if s == nil || s.db == nil { + return nil, fmt.Errorf("budget store unavailable") + } + rows, err := s.db.QueryContext(ctx, `SELECT call_id,day,created_at FROM budget_calls WHERE status='pending' ORDER BY created_at,call_id`) + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + out := make([]PendingBudgetCall, 0) + for rows.Next() { + var row PendingBudgetCall + var at int64 + if err := rows.Scan(&row.CallID, &row.Day, &at); err != nil { + return nil, err + } + row.CreatedAt = time.Unix(at, 0).UTC() + out = append(out, row) + } + return out, rows.Err() +} diff --git a/internal/state/budget_test.go b/internal/state/budget_test.go new file mode 100644 index 0000000..e5e4950 --- /dev/null +++ b/internal/state/budget_test.go @@ -0,0 +1,146 @@ +package state + +import ( + "context" + "math" + "path/filepath" + "sync" + "sync/atomic" + "testing" + "time" +) + +func TestBudgetDurableUsageAndUTCFullDates(t *testing.T) { + path := filepath.Join(t.TempDir(), "budget.db") + ctx := context.Background() + st, err := OpenStateStore(path) + if err != nil { + t.Fatal(err) + } + at := time.Date(2026, 9, 30, 23, 59, 0, 0, time.UTC) + if err := st.BeginBudgetCall(ctx, "paid", at, 10, 100, .5); err != nil { + t.Fatal(err) + } + if err := st.SettleBudgetCall(ctx, "paid", 20, .6); err != nil { + t.Fatal(err) + } + if err := st.SettleBudgetCall(ctx, "paid", 20, .6); err == nil { + t.Fatal("same call charged twice") + } + _ = st.Close() + st, err = OpenStateStore(path) + if err != nil { + t.Fatal(err) + } + defer st.Close() + day, err := st.BudgetDay(ctx, at) + if err != nil || day.Tokens != 20 || day.Cost != .6 || day.Unsettled != 0 { + t.Fatalf("durable usage=%+v err=%v", day, err) + } + if err := st.BeginBudgetCall(ctx, "same-day", at, 1, 100, .5); err == nil { + t.Fatal("restart reopened settled cost cap") + } + next := at.Add(2 * time.Minute) + if err := st.BeginBudgetCall(ctx, "next-day", next, 1, 100, .5); err != nil { + t.Fatal(err) + } + if err := st.SettleBudgetCall(ctx, "next-day", 2, .1); err != nil { + t.Fatal(err) + } + for date, want := range map[time.Time]int{at: 20, next: 2, at.AddDate(1, 0, 0): 0, at.AddDate(0, 1, 0): 0} { + d, err := st.BudgetDay(ctx, date) + if err != nil || d.Tokens != want { + t.Fatalf("date=%s usage=%+v err=%v", date, d, err) + } + } + local := time.Date(2026, 10, 1, 1, 59, 0, 0, time.FixedZone("plus-two", 2*3600)) + d, err := st.BudgetDay(ctx, local) + if err != nil || d.Tokens != 20 { + t.Fatalf("local time was not grouped by UTC date: %+v %v", d, err) + } +} + +func TestBudgetPendingCallSurvivesRestartAndBlocksUntilReconciled(t *testing.T) { + path := filepath.Join(t.TempDir(), "budget.db") + ctx := context.Background() + at := time.Now() + st, err := OpenStateStore(path) + if err != nil { + t.Fatal(err) + } + if err := st.BeginBudgetCall(ctx, "unknown", at, 1, 0, 0); err != nil { + t.Fatal(err) + } + _ = st.Close() + st, err = OpenStateStore(path) + if err != nil { + t.Fatal(err) + } + defer st.Close() + if err := st.BeginBudgetCall(ctx, "next", at.AddDate(0, 0, 1), 1, 0, 0); err == nil { + t.Fatal("pending usage was forgotten on restart/new day") + } + if err := st.SettleBudgetCall(ctx, "unknown", 12, .2); err != nil { + t.Fatal(err) + } + if err := st.BeginBudgetCall(ctx, "next", at, 1, 100, .5); err != nil { + t.Fatal(err) + } +} + +func TestBudgetAdmissionIsAtomicAcrossStoreHandles(t *testing.T) { + path := filepath.Join(t.TempDir(), "budget.db") + ctx := context.Background() + a, err := OpenStateStore(path) + if err != nil { + t.Fatal(err) + } + defer a.Close() + b, err := OpenStateStore(path) + if err != nil { + t.Fatal(err) + } + defer b.Close() + var wg sync.WaitGroup + var admitted atomic.Int32 + for id, st := range map[string]*StateStore{"a": a, "b": b} { + wg.Add(1) + go func(id string, st *StateStore) { + defer wg.Done() + if st.BeginBudgetCall(ctx, id, time.Now(), 1, 100, 1) == nil { + admitted.Add(1) + } + }(id, st) + } + wg.Wait() + if admitted.Load() != 1 { + t.Fatalf("admitted=%d, want exactly one outstanding provider call", admitted.Load()) + } +} + +func TestBudgetRejectsInvalidAndOverflowUsageWithoutReleasingPending(t *testing.T) { + st := newTempStore(t) + ctx := context.Background() + at := time.Now() + if err := st.BeginBudgetCall(ctx, "pending", at, 1, 0, 0); err != nil { + t.Fatal(err) + } + for _, cost := range []float64{-1, math.NaN(), math.Inf(1)} { + if err := st.SettleBudgetCall(ctx, "pending", 1, cost); err == nil { + t.Fatalf("invalid cost %v settled", cost) + } + } + if err := st.SettleBudgetCall(ctx, "pending", -1, 0); err == nil { + t.Fatal("negative tokens settled") + } + if err := st.AddBudgetUsage(ctx, at, int(^uint(0)>>1), 0, 0, 0); err != nil { + t.Fatal(err) + } + if err := st.SettleBudgetCall(ctx, "pending", 1, 0); err == nil { + t.Fatal("token arithmetic overflow accepted") + } + day, err := st.BudgetDay(ctx, at) + if err != nil || day.Unsettled != 1 || day.Tokens != int(^uint(0)>>1) { + t.Fatalf("invalid accounting changed ledger: %+v %v", day, err) + } +} diff --git a/internal/state/state.go b/internal/state/state.go index d0baa46..20f9642 100644 --- a/internal/state/state.go +++ b/internal/state/state.go @@ -148,6 +148,30 @@ CREATE TABLE IF NOT EXISTS tool_actions ( consumed_at INTEGER NOT NULL DEFAULT 0 ); CREATE INDEX IF NOT EXISTS idx_tool_actions_status ON tool_actions(status, updated_at); +CREATE TABLE IF NOT EXISTS tool_execution_outcomes ( + action_id TEXT PRIMARY KEY REFERENCES tool_actions(action_id), + binding_hash TEXT NOT NULL, + status TEXT NOT NULL CHECK(status IN ('succeeded','failed')), + result_hash TEXT NOT NULL DEFAULT '', + completed_at INTEGER NOT NULL +); +CREATE TABLE IF NOT EXISTS budget_days ( + day TEXT PRIMARY KEY, + tokens INTEGER NOT NULL DEFAULT 0, + cost REAL NOT NULL DEFAULT 0, + calls INTEGER NOT NULL DEFAULT 0, + call_minutes INTEGER NOT NULL DEFAULT 0 +); +CREATE TABLE IF NOT EXISTS budget_calls ( + call_id TEXT PRIMARY KEY, + day TEXT NOT NULL, + status TEXT NOT NULL, + tokens INTEGER NOT NULL DEFAULT 0, + cost REAL NOT NULL DEFAULT 0, + created_at INTEGER NOT NULL, + settled_at INTEGER NOT NULL DEFAULT 0 +); +CREATE INDEX IF NOT EXISTS idx_budget_calls_status ON budget_calls(status); CREATE TABLE IF NOT EXISTS webhook_admissions ( id TEXT PRIMARY KEY, pipeline TEXT NOT NULL, @@ -853,7 +877,7 @@ func (s *StateStore) ConsumeToolAction(id, bindingHash string, at time.Time) (To } res, err := s.db.ExecContext(context.Background(), `UPDATE tool_actions SET status='consumed', consumed_at=?, updated_at=? - WHERE action_id=? AND binding_hash=? AND status='allowed' AND expires_at>=?`, + WHERE action_id=? AND binding_hash=? AND status='allowed' AND expires_at>?`, at.Unix(), at.Unix(), id, bindingHash, at.Unix()) if err != nil { return ToolAction{}, false, err @@ -863,7 +887,7 @@ func (s *StateStore) ConsumeToolAction(id, bindingHash string, at time.Time) (To if _, err := s.db.ExecContext(context.Background(), `UPDATE tool_actions SET status='expired', decision='deny', reason='permit expired before consumption', updated_at=? - WHERE action_id=? AND binding_hash=? AND status='allowed' AND expires_at?`, at.Unix(), at.Unix(), id, binding, at.Unix()) + if err != nil { + return ToolAction{}, false, err + } + n, err := res.RowsAffected() + if err != nil { + return ToolAction{}, false, err + } + if n == 1 { + if err := appendToolReceipt(tx, e); err != nil { + return ToolAction{}, false, err + } + } + a, err := toolActionTx(tx, id) + if err != nil { + return ToolAction{}, false, err + } + if err := tx.Commit(); err != nil { + return ToolAction{}, false, err + } + return a, n == 1, nil +} + +// NormalizeToolAction closes invalid unconsumed permits. Current policy and +// expiry apply equally to live tickets and records recovered after restart. +func (s *StateStore) NormalizeToolAction(id, policy string, at time.Time) (ToolAction, error) { + if s == nil || s.db == nil { + return ToolAction{}, fmt.Errorf("state store unavailable") + } + _, err := s.db.ExecContext(context.Background(), `UPDATE tool_actions SET status='expired',decision='deny',reason=CASE WHEN policy_hash<>? THEN 'policy changed; request a new approval' ELSE 'permit expired before consumption' END,updated_at=? WHERE action_id=? AND status IN ('pending','allowed') AND (policy_hash<>? OR expires_at<=?)`, policy, at.Unix(), id, policy, at.Unix()) + if err != nil { + return ToolAction{}, err + } + return s.ToolAction(id) +} + +// CancelToolAction revokes only pending or allowed permits with this exact +// binding. The database serializes it against consumption. Matching retries +// of an already revoked action succeed; consumed actions cannot be revoked. +func (s *StateStore) CancelToolAction(id, binding string, at time.Time) (ToolAction, bool, error) { + if s == nil || s.db == nil { + return ToolAction{}, false, fmt.Errorf("state store unavailable") + } + _, err := s.db.ExecContext(context.Background(), `UPDATE tool_actions SET status='revoked',decision='deny',reason='action revoked before consumption',decided_by='caller',updated_at=? WHERE action_id=? AND binding_hash=? AND status IN ('pending','allowed')`, at.Unix(), id, binding) + if err != nil { + return ToolAction{}, false, err + } + a, err := s.ToolAction(id) + return a, err == nil && a.BindingHash == binding && a.Status == "revoked", err +} + +func (s *StateStore) ToolOutcome(id string) (ToolOutcome, error) { + if s == nil || s.db == nil { + return ToolOutcome{}, fmt.Errorf("state store unavailable") + } + var o ToolOutcome + var at int64 + err := s.db.QueryRowContext(context.Background(), `SELECT action_id,binding_hash,status,result_hash,completed_at FROM tool_execution_outcomes WHERE action_id=?`, id).Scan(&o.ActionID, &o.BindingHash, &o.Status, &o.ResultHash, &at) + if err == nil { + o.CompletedAt = time.Unix(at, 0) + } + return o, err +} + +// CompleteToolAction records a single immutable caller-reported result after +// consumption. Same-outcome retries succeed; conflicting reports do not. +func (s *StateStore) CompleteToolAction(o ToolOutcome) (ToolOutcome, bool, error) { + if s == nil || s.db == nil { + return ToolOutcome{}, false, fmt.Errorf("state store unavailable") + } + if o.Status != "succeeded" && o.Status != "failed" { + return ToolOutcome{}, false, fmt.Errorf("invalid execution status") + } + tx, err := s.db.BeginTx(context.Background(), nil) + if err != nil { + return ToolOutcome{}, false, err + } + defer func() { _ = tx.Rollback() }() + a, err := toolActionTx(tx, o.ActionID) + if err != nil { + return ToolOutcome{}, false, err + } + if a.Status != "consumed" || a.BindingHash != o.BindingHash { + return ToolOutcome{}, false, nil + } + _, err = tx.ExecContext(context.Background(), `INSERT OR IGNORE INTO tool_execution_outcomes(action_id,binding_hash,status,result_hash,completed_at) VALUES(?,?,?,?,?)`, o.ActionID, o.BindingHash, o.Status, o.ResultHash, o.CompletedAt.Unix()) + if err != nil { + return ToolOutcome{}, false, err + } + var saved ToolOutcome + var at int64 + err = tx.QueryRowContext(context.Background(), `SELECT action_id,binding_hash,status,result_hash,completed_at FROM tool_execution_outcomes WHERE action_id=?`, o.ActionID).Scan(&saved.ActionID, &saved.BindingHash, &saved.Status, &saved.ResultHash, &at) + if err != nil { + return ToolOutcome{}, false, err + } + saved.CompletedAt = time.Unix(at, 0) + if err := tx.Commit(); err != nil { + return ToolOutcome{}, false, err + } + return saved, saved.BindingHash == o.BindingHash && saved.Status == o.Status && saved.ResultHash == o.ResultHash, nil +} diff --git a/internal/state/tool_lifecycle_test.go b/internal/state/tool_lifecycle_test.go new file mode 100644 index 0000000..944eb5a --- /dev/null +++ b/internal/state/tool_lifecycle_test.go @@ -0,0 +1,159 @@ +package state + +import ( + "context" + "fmt" + "path/filepath" + "sync" + "testing" + "time" +) + +func lifecycleAction(t *testing.T, s *StateStore, id string, now time.Time) ToolAction { + t.Helper() + a := ToolAction{ActionID: id, Tool: "send", ArgsHash: "args", PolicyHash: "policy", BindingHash: "binding", CreatedAt: now, UpdatedAt: now, ExpiresAt: now.Add(time.Hour)} + if _, _, err := s.ReserveToolAction(a); err != nil { + t.Fatal(err) + } + return a +} +func lifecycleReceipt(a ToolAction, at time.Time, decision string) ApprovalEnvelope { + return ApprovalEnvelope{ReceiptID: "receipt-" + a.ActionID + "-" + decision, ActionID: a.ActionID, Pipeline: "tool-gate", Step: a.Tool, PayloadHash: a.ArgsHash, PolicyHash: a.PolicyHash, BindingHash: a.BindingHash, ExpiresAt: a.ExpiresAt, DecidedAt: at, Decision: decision, Lifecycle: decision} +} + +func TestToolLifecycleRevokeConsumeRace(t *testing.T) { + s := openTestStore(t) + now := time.Now().Truncate(time.Second) + for i := 0; i < 30; i++ { + a := lifecycleAction(t, s, fmt.Sprintf("race-%d", i), now) + if err := s.DecideToolAction(a.ActionID, "allowed", "allow", "approved", "operator", now); err != nil { + t.Fatal(err) + } + start := make(chan struct{}) + var wg sync.WaitGroup + wg.Add(2) + var consumed, revoked bool + var consumeErr, revokeErr error + go func() { + defer wg.Done() + <-start + _, consumed, consumeErr = s.ConsumeToolActionWithReceipt(a.ActionID, a.BindingHash, now, lifecycleReceipt(a, now, "consume")) + }() + go func() { + defer wg.Done() + <-start + _, revoked, revokeErr = s.CancelToolAction(a.ActionID, a.BindingHash, now) + }() + close(start) + wg.Wait() + if consumeErr != nil || revokeErr != nil || consumed == revoked { + t.Fatalf("consume=%v revoke=%v errors=%v/%v", consumed, revoked, consumeErr, revokeErr) + } + got, err := s.ToolAction(a.ActionID) + if err != nil { + t.Fatal(err) + } + if revoked && got.Status != "revoked" || consumed && got.Status != "consumed" { + t.Fatalf("race state %+v", got) + } + if revoked { + if _, ok, err := s.CancelToolAction(a.ActionID, a.BindingHash, now); err != nil || !ok { + t.Fatalf("revoke retry=%v %v", ok, err) + } + if got, err := s.DecideToolActionWithReceipt(a.ActionID, "allowed", "allow", "late approval", "operator", now, lifecycleReceipt(a, now, "approve")); err != nil || got.Status != "revoked" { + t.Fatalf("late approval resurrected: %+v %v", got, err) + } + } + } +} + +func TestToolLifecycleAuditFailureRollsBackDecisionAndConsumption(t *testing.T) { + s := openTestStore(t) + now := time.Now().Truncate(time.Second) + a := lifecycleAction(t, s, "audit", now) + if _, err := s.db.ExecContext(context.Background(), `CREATE TRIGGER reject_tool_receipt BEFORE INSERT ON action_approvals BEGIN SELECT RAISE(FAIL,'audit unavailable'); END`); err != nil { + t.Fatal(err) + } + if _, err := s.DecideToolActionWithReceipt(a.ActionID, "allowed", "allow", "approve", "operator", now, lifecycleReceipt(a, now, "approve")); err == nil { + t.Fatal("failed audit permitted decision") + } + if got, _ := s.ToolAction(a.ActionID); got.Status != "pending" { + t.Fatalf("decision escaped rollback: %+v", got) + } + if err := s.DecideToolAction(a.ActionID, "allowed", "allow", "approved", "operator", now); err != nil { + t.Fatal(err) + } + if _, ok, err := s.ConsumeToolActionWithReceipt(a.ActionID, a.BindingHash, now, lifecycleReceipt(a, now, "consume")); err == nil || ok { + t.Fatal("failed audit issued execution permit") + } + if got, _ := s.ToolAction(a.ActionID); got.Status != "allowed" { + t.Fatalf("consumption escaped rollback: %+v", got) + } +} + +func TestToolLifecycleOutcomesImmutableAndSurviveRestart(t *testing.T) { + path := filepath.Join(t.TempDir(), "state.db") + s, err := OpenStateStore(path) + if err != nil { + t.Fatal(err) + } + now := time.Now().Truncate(time.Second) + a := lifecycleAction(t, s, "outcome", now) + o := ToolOutcome{ActionID: a.ActionID, BindingHash: a.BindingHash, Status: "succeeded", ResultHash: "sha256:result", CompletedAt: now} + if _, ok, err := s.CompleteToolAction(o); err != nil || ok { + t.Fatalf("unconsumed completion=%v %v", ok, err) + } + if err := s.DecideToolAction(a.ActionID, "allowed", "allow", "approved", "operator", now); err != nil { + t.Fatal(err) + } + if _, ok, err := s.ConsumeToolActionWithReceipt(a.ActionID, a.BindingHash, now, lifecycleReceipt(a, now, "consume")); err != nil || !ok { + t.Fatalf("consume=%v %v", ok, err) + } + if _, ok, err := s.CompleteToolAction(o); err != nil || !ok { + t.Fatalf("complete=%v %v", ok, err) + } + if err := s.Close(); err != nil { + t.Fatal(err) + } + s, err = OpenStateStore(path) + if err != nil { + t.Fatal(err) + } + defer s.Close() + if got, err := s.ToolOutcome(a.ActionID); err != nil || got.Status != "succeeded" || got.ResultHash != o.ResultHash { + t.Fatalf("restored outcome=%+v %v", got, err) + } + o.CompletedAt = now.Add(time.Minute) + if got, ok, err := s.CompleteToolAction(o); err != nil || !ok || !got.CompletedAt.Equal(now) { + t.Fatalf("retry mutated outcome=%+v %v %v", got, ok, err) + } + o.Status = "failed" + if _, ok, err := s.CompleteToolAction(o); err != nil || ok { + t.Fatalf("conflicting result accepted=%v %v", ok, err) + } + if _, ok, err := s.ConsumeToolActionWithReceipt(a.ActionID, a.BindingHash, now, lifecycleReceipt(a, now, "consume-again")); err != nil || ok { + t.Fatal("outcome reopened execution") + } +} + +func TestToolLifecycleExactExpiryAndPolicyNormalization(t *testing.T) { + s := openTestStore(t) + now := time.Now().Truncate(time.Second) + for _, kind := range []string{"expiry", "policy"} { + a := lifecycleAction(t, s, kind, now) + if err := s.DecideToolAction(a.ActionID, "allowed", "allow", "approved", "operator", now); err != nil { + t.Fatal(err) + } + at, policy := a.ExpiresAt, a.PolicyHash + if kind == "policy" { + at, policy = now, "new-policy" + } + got, err := s.NormalizeToolAction(a.ActionID, policy, at) + if err != nil || got.Status != "expired" || got.Decision != "deny" { + t.Fatalf("normalization=%+v %v", got, err) + } + if _, ok, err := s.ConsumeToolActionWithReceipt(a.ActionID, a.BindingHash, at, lifecycleReceipt(a, at, "consume")); err != nil || ok { + t.Fatalf("invalid permit consumed=%v %v", ok, err) + } + } +} diff --git a/internal/validate/validate.go b/internal/validate/validate.go index 3eb9211..de2b14a 100644 --- a/internal/validate/validate.go +++ b/internal/validate/validate.go @@ -15,6 +15,7 @@ import ( "github.com/renezander030/draftcat/internal/channels" "github.com/renezander030/draftcat/internal/config" + "github.com/renezander030/draftcat/internal/outputschema" skillsapi "github.com/renezander030/draftcat/internal/skills" ) @@ -195,29 +196,7 @@ func loadSkillsForValidate(skillsDir string, rep *validateReport) map[string]*sk if s.Prompt == "" { rep.errf("skills/"+s.Name, "missing 'prompt'") } - for field, def := range s.OutputSchema { - dm, ok := def.(map[string]interface{}) - if !ok { - rep.warnf("skills/"+s.Name, "output_schema.%s: definition is not a map", field) - continue - } - _, hasEnum := dm["enum"] - t, _ := dm["type"].(string) - switch t { - case "int", "number", "bool", "string": - case "": - if !hasEnum { - rep.warnf("skills/"+s.Name, "output_schema.%s: missing 'type'", field) - } - default: - rep.warnf("skills/"+s.Name, "output_schema.%s: unsupported type %q (validator handles int|number|bool|string)", field, t) - } - if hasEnum { - if _, ok := dm["enum"].([]interface{}); !ok { - rep.warnf("skills/"+s.Name, "output_schema.%s: 'enum' must be a list", field) - } - } - } + checkOutputSchema(s.OutputSchema, "skills/"+s.Name, rep) if _, dup := skills[s.Name]; dup { rep.errf("skills/"+s.Name, "duplicate skill name (also defined in another file)") } @@ -383,6 +362,7 @@ func checkPipelines(cfg *config.Config, skills map[string]*skillsapi.SkillDef, s } case "ai": aiSteps++ + checkOutputSchema(st.OutputSchema, spath, rep) if st.Skill == "" && st.Prompt == "" { rep.errf(spath, "ai step needs either 'skill' or inline 'prompt'") } @@ -935,3 +915,11 @@ func min3(a, b, c int) int { } return a } + +// checkOutputSchema shares the runtime's flat schema contract. An unsupported +// definition must fail before a pipeline can accept unchecked model output. +func checkOutputSchema(schema map[string]interface{}, path string, rep *validateReport) { + for _, finding := range outputschema.Check(schema) { + rep.errf(path+".output_schema."+finding.Field, "%s", finding.Message) + } +} diff --git a/internal/validate/validate_output_schema_test.go b/internal/validate/validate_output_schema_test.go new file mode 100644 index 0000000..440f731 --- /dev/null +++ b/internal/validate/validate_output_schema_test.go @@ -0,0 +1,36 @@ +package validate + +import ( + "os" + "path/filepath" + "strings" + "testing" + + "github.com/renezander030/draftcat/internal/config" +) + +func TestInlineOutputSchemaErrorsBeforeRun(t *testing.T) { + cfg := &config.Config{Pipelines: []config.PipelineConfig{{Name: "p", Steps: []config.StepConfig{{ + Name: "classify", Type: "ai", Prompt: "classify", OutputSchema: map[string]interface{}{ + "score": map[string]interface{}{"type": "int", "min": 5, "max": 1}, + }, + }}}}} + errs := findingsAt(runCheckPipelines(cfg), "error", "output_schema.score") + if len(errs) != 1 || !strings.Contains(errs[0].Message, "min exceeds max") { + t.Fatalf("bad inline schema didn't fail startup: %+v", errs) + } +} + +func TestSkillOutputSchemaErrorsBeforeRun(t *testing.T) { + dir := t.TempDir() + content := []byte("name: classify\nprompt: classify\noutput_schema:\n score: {type: unsupported}\n choice: {enum: wrong}\n") + if err := os.WriteFile(filepath.Join(dir, "classify.yaml"), content, 0600); err != nil { + t.Fatal(err) + } + rep := &validateReport{} + loadSkillsForValidate(dir, rep) + errs := findingsAt(rep, "error", "output_schema") + if len(errs) != 2 || !strings.HasSuffix(errs[0].Path, "choice") || !strings.HasSuffix(errs[1].Path, "score") { + t.Fatalf("bad skill schema didn't return ordered startup errors: %+v", rep.Findings) + } +} diff --git a/llm_client_test.go b/llm_client_test.go index ad012c8..ad25228 100644 --- a/llm_client_test.go +++ b/llm_client_test.go @@ -141,8 +141,8 @@ func TestCallLLM_ClientErrorIsNotRetried(t *testing.T) { } } -func TestCallLLM_GivesUpAfterTheRetryBudget(t *testing.T) { - s := &llmServer{steps: []func(http.ResponseWriter){status(503, "0")}} +func TestCallLLM_GivesUpAfterTheRateLimitRetryBudget(t *testing.T) { + s := &llmServer{steps: []func(http.ResponseWriter){status(429, "0")}} srv := httptest.NewServer(http.HandlerFunc(s.handler)) defer srv.Close() cfg := llmCfg(srv.URL, "openrouter") @@ -200,7 +200,7 @@ func TestRetryDelay_HonorsRetryAfterAndCaps(t *testing.T) { } func TestLLMRetryable(t *testing.T) { - for code, want := range map[int]bool{429: true, 408: true, 500: true, 503: true, 400: false, 401: false, 404: false} { + for code, want := range map[int]bool{429: true, 408: false, 500: false, 503: false, 400: false, 401: false, 404: false} { if got := llmRetryable(code); got != want { t.Errorf("llmRetryable(%d) = %v, want %v", code, got, want) } diff --git a/main.go b/main.go index 84f9e38..73b59c1 100644 --- a/main.go +++ b/main.go @@ -189,98 +189,6 @@ func (s *Scheduler) GetDue() []string { // --- Guardrails --- -type BudgetTracker struct { - tokensUsedToday int - tokensUsedPipeline int - callsToday int - callMinutesToday int - costToday float64 - costPipeline float64 - dayStart time.Time - // Cost caps are held on the tracker rather than passed per call so that - // check() enforces them everywhere it is already called. Wiring them at - // each of the eight LLM call sites instead would mean a new call site added - // later silently spends without a cap — the failure mode is invisible until - // the bill arrives. 0 = no cap. - dayCostLimit float64 - pipelineCostLimit float64 -} - -func (b *BudgetTracker) resetIfNewDay() { - if b.dayStart.Day() != time.Now().Day() { - b.tokensUsedToday = 0 - b.callsToday = 0 - b.callMinutesToday = 0 - b.costToday = 0 - b.dayStart = time.Now() - } -} - -func (b *BudgetTracker) check(limit int, requested int) error { - b.resetIfNewDay() - if b.tokensUsedToday+requested > limit { - return fmt.Errorf("BUDGET_BLOCKED: daily token limit %d would be exceeded (used: %d, requested: %d)", limit, b.tokensUsedToday, requested) - } - // Money caps ride the same pre-call gate as token caps. - return b.CheckCost(b.dayCostLimit, b.pipelineCostLimit) -} - -func (b *BudgetTracker) record(tokens int) { - b.tokensUsedToday += tokens - b.tokensUsedPipeline += tokens -} - -func (b *BudgetTracker) CheckCalls(limit int) error { - b.resetIfNewDay() - if limit > 0 && b.callsToday+1 > limit { - return fmt.Errorf("BUDGET_BLOCKED: daily call limit %d would be exceeded (used: %d)", limit, b.callsToday) - } - return nil -} - -func (b *BudgetTracker) CheckCallMinutes(limit, requestedMinutes int) error { - b.resetIfNewDay() - if limit > 0 && b.callMinutesToday+requestedMinutes > limit { - return fmt.Errorf("BUDGET_BLOCKED: daily call-minute limit %d would be exceeded (used: %d, requested: %d)", limit, b.callMinutesToday, requestedMinutes) - } - return nil -} - -func (b *BudgetTracker) RecordCall(durationMinutes int) { - b.resetIfNewDay() - b.callsToday++ - b.callMinutesToday += durationMinutes -} - -// CheckCost blocks the next LLM call once spend has already reached a cap. -// Token caps answer "how much did it think"; this answers the question the -// person paying actually asks — "what is this costing me today". The per-call -// cost was already computed from the model rates (see callLLM) and then thrown -// away; now it accumulates and enforces. -// -// Deliberately checked BETWEEN calls rather than estimated ahead of one: a -// pre-call estimate needs the response token count, which does not exist yet, -// and guessing it either blocks legitimate work or lets the real overshoot -// through anyway. So a single in-flight call may exceed the cap by at most one -// step, bounded by per_step_tokens. Limits <= 0 mean "no cap". -func (b *BudgetTracker) CheckCost(dayLimit, pipelineLimit float64) error { - b.resetIfNewDay() - if dayLimit > 0 && b.costToday >= dayLimit { - return fmt.Errorf("BUDGET_BLOCKED: daily cost limit %.4f reached (spent: %.4f)", dayLimit, b.costToday) - } - if pipelineLimit > 0 && b.costPipeline >= pipelineLimit { - return fmt.Errorf("BUDGET_BLOCKED: per-pipeline cost limit %.4f reached (spent: %.4f)", pipelineLimit, b.costPipeline) - } - return nil -} - -// RecordCost accumulates the cost of one completed LLM call. -func (b *BudgetTracker) RecordCost(cost float64) { - b.resetIfNewDay() - b.costToday += cost - b.costPipeline += cost -} - // --- Input Security --- // Applied to ALL operator input before it reaches any AI step. // This is not optional — the engine validates channel security config at startup. @@ -485,11 +393,10 @@ var httpClient = &http.Client{ Timeout: 30 * time.Second, } -// llmRetryable reports whether a status is worth another attempt: rate limits, -// request timeouts and server-side failures. Any other 4xx is the caller's -// mistake, and retrying it only spends the budget again. +// llmRetryable identifies an explicit rate-limit rejection. Timeouts and server +// errors can follow a paid response, so they require reconciliation first. func llmRetryable(status int) bool { - return status == http.StatusTooManyRequests || status == http.StatusRequestTimeout || status >= 500 + return status == http.StatusTooManyRequests } // retryDelay is how long to wait before the given retry (1-based). A @@ -532,7 +439,10 @@ func retryDelay(retry int, retryAfter string, now time.Time) time.Duration { return base + jitter } -func callLLM(ctx context.Context, cfg *config.Config, role string, prompt string) (*CompletionResponse, error) { +func providerCallLLM(ctx context.Context, cfg *config.Config, role string, prompt string, maxTokens int) (*CompletionResponse, error) { + if err := ctx.Err(); err != nil { + return nil, providerError(err, false) + } modelName, ok := cfg.Roles[role] if !ok { return nil, fmt.Errorf("unknown role: %s", role) @@ -541,16 +451,13 @@ func callLLM(ctx context.Context, cfg *config.Config, role string, prompt string if !ok { return nil, fmt.Errorf("unknown model: %s", modelName) } - if err := enforceModelPolicy(ctx, cfg, role, "input", prompt); err != nil { - return nil, err - } reqBody := map[string]interface{}{ "model": modelCfg.Model, "messages": []map[string]string{ {"role": "user", "content": prompt}, }, - "max_tokens": modelCfg.MaxTokens, + "max_tokens": maxTokens, } // OpenRouter reports the real charge (usage.cost) and the cached/reasoning // token breakdown when asked. Other OpenAI-compatible servers reject @@ -576,7 +483,7 @@ func callLLM(ctx context.Context, cfg *config.Config, role string, prompt string select { case <-time.After(delay): case <-ctx.Done(): - return nil, fmt.Errorf("LLM call abandoned during backoff: %w (last error: %w)", ctx.Err(), lastErr) + return nil, providerError(fmt.Errorf("LLM call abandoned during backoff: %w", ctx.Err()), false) } } attempts++ @@ -584,7 +491,7 @@ func callLLM(ctx context.Context, cfg *config.Config, role string, prompt string req, err := http.NewRequestWithContext(ctx, "POST", cfg.Provider.BaseURL+"/chat/completions", bytes.NewReader(body)) if err != nil { - return nil, err + return nil, providerError(err, false) } req.Header.Set("Content-Type", "application/json") req.Header.Set("Authorization", "Bearer "+cfg.Provider.APIKey()) @@ -592,28 +499,31 @@ func callLLM(ctx context.Context, cfg *config.Config, role string, prompt string start := time.Now() resp, err := httpClient.Do(req) if err != nil { - lastErr = fmt.Errorf("LLM request failed: %w", err) - if ctx.Err() != nil { - return nil, lastErr - } - continue + // Dispatch may have reached the provider. Repeating an ambiguous + // request would spend again without accounting for the first one. + return nil, providerError(fmt.Errorf("LLM request failed; usage may require reconciliation: %w", err), true) } latency = time.Since(start).Milliseconds() - respBody, _ = io.ReadAll(resp.Body) - resp.Body.Close() + respBody, err = readProviderResponse(ctx, resp.Body) + _ = resp.Body.Close() + if err != nil { + return nil, providerError(err, true) + } if resp.StatusCode == http.StatusOK { lastErr = nil break } - lastErr = fmt.Errorf("LLM API error %d: %s", resp.StatusCode, string(respBody)) - if !llmRetryable(resp.StatusCode) { - return nil, lastErr + lastErr = fmt.Errorf("LLM API error %d", resp.StatusCode) + // Only an explicit rate-limit rejection is safe to retry. A timeout + // or server error can follow a paid completion. + if !llmRetryable(resp.StatusCode) || providerDeclaresUsage(respBody) { + return nil, providerError(lastErr, providerDeclaresUsage(respBody) || resp.StatusCode == http.StatusRequestTimeout || resp.StatusCode >= 500) } retryAfter = resp.Header.Get("Retry-After") } if lastErr != nil { - return nil, fmt.Errorf("%w (gave up after %d attempt(s))", lastErr, attempts) + return nil, providerError(fmt.Errorf("%w (gave up after %d attempt(s))", lastErr, attempts), false) } var result struct { @@ -623,8 +533,8 @@ func callLLM(ctx context.Context, cfg *config.Config, role string, prompt string } `json:"message"` } `json:"choices"` Usage struct { - PromptTokens int `json:"prompt_tokens"` - CompletionTokens int `json:"completion_tokens"` + PromptTokens *int `json:"prompt_tokens"` + CompletionTokens *int `json:"completion_tokens"` Cost *float64 `json:"cost"` PromptTokensDetails struct { CachedTokens int `json:"cached_tokens"` @@ -636,30 +546,35 @@ func callLLM(ctx context.Context, cfg *config.Config, role string, prompt string Model string `json:"model"` } if err := json.Unmarshal(respBody, &result); err != nil { - return nil, fmt.Errorf("failed to parse LLM response: %w", err) + return nil, providerError(fmt.Errorf("failed to parse LLM response: %w", err), true) } - if len(result.Choices) == 0 { - return nil, fmt.Errorf("LLM returned no choices") + if result.Usage.PromptTokens == nil || result.Usage.CompletionTokens == nil { + return nil, providerError(fmt.Errorf("LLM response is missing token usage; reconciliation required"), true) } - if err := enforceModelPolicy(ctx, cfg, role, "output", result.Choices[0].Message.Content); err != nil { - return nil, err + inputTokens, outputTokens := *result.Usage.PromptTokens, *result.Usage.CompletionTokens + if inputTokens < 0 || outputTokens < 0 || inputTokens > int(^uint(0)>>1)-outputTokens || + result.Usage.PromptTokensDetails.CachedTokens < 0 || result.Usage.CompletionTokensDetails.ReasoningTokens < 0 || + (result.Usage.Cost != nil && !validCost(*result.Usage.Cost)) { + return nil, providerError(fmt.Errorf("LLM response contains invalid usage; reconciliation required"), true) } // The configured rates are the fallback. When the provider states what it // actually charged, the caps are enforced on that number — rates alone // undercount reasoning tokens and overcount cached ones. - cost := float64(result.Usage.PromptTokens)/1000*modelCfg.CostIn + - float64(result.Usage.CompletionTokens)/1000*modelCfg.CostOut + cost := float64(inputTokens)/1000*modelCfg.CostIn + + float64(outputTokens)/1000*modelCfg.CostOut source := "rates" if usageAccounting && result.Usage.Cost != nil { cost = *result.Usage.Cost source = "provider" } - return &CompletionResponse{ - Text: result.Choices[0].Message.Content, - InputTokens: result.Usage.PromptTokens, - OutputTokens: result.Usage.CompletionTokens, + if !validCost(cost) { + return nil, providerError(fmt.Errorf("LLM response cost is invalid; reconciliation required"), true) + } + completion := &CompletionResponse{ + InputTokens: inputTokens, + OutputTokens: outputTokens, CachedTokens: result.Usage.PromptTokensDetails.CachedTokens, ReasoningTokens: result.Usage.CompletionTokensDetails.ReasoningTokens, LatencyMs: latency, @@ -667,114 +582,12 @@ func callLLM(ctx context.Context, cfg *config.Config, role string, prompt string CostSource: source, Model: result.Model, Attempts: attempts, - }, nil -} - -// --- Output Validation --- - -func toFloat64(v interface{}) *float64 { - switch n := v.(type) { - case float64: - return &n - case int: - f := float64(n) - return &f - case int64: - f := float64(n) - return &f - default: - return nil - } -} - -// enumContains reports whether val is a member of the allowed set. Numbers are -// compared by value (YAML decodes enum ints as int, JSON decodes the field as -// float64), everything else by interface equality. -func enumContains(allowed []interface{}, val interface{}) bool { - vNum := toFloat64(val) - for _, a := range allowed { - if vNum != nil { - if aNum := toFloat64(a); aNum != nil && *aNum == *vNum { - return true - } - continue - } - if a == val { - return true - } } - return false -} - -func validateOutput(text string, schema map[string]interface{}) (map[string]interface{}, error) { - if len(schema) == 0 { - return nil, nil - } - - // Strip markdown code fences if present - cleaned := strings.TrimSpace(text) - if strings.HasPrefix(cleaned, "```") { - lines := strings.Split(cleaned, "\n") - if len(lines) >= 3 { - cleaned = strings.Join(lines[1:len(lines)-1], "\n") - } - } - - var parsed map[string]interface{} - if err := json.Unmarshal([]byte(cleaned), &parsed); err != nil { - return nil, fmt.Errorf("output is not valid JSON: %w\nRaw: %s", err, text) - } - - for key, schemaDef := range schema { - val, exists := parsed[key] - if !exists { - return nil, fmt.Errorf("missing required field: %s", key) - } - - if defMap, ok := schemaDef.(map[string]interface{}); ok { - if typeName, ok := defMap["type"].(string); ok { - switch typeName { - case "int", "number": - num := toFloat64(val) - if num == nil { - return nil, fmt.Errorf("field %s: expected number, got %T", key, val) - } - if minVal, ok := defMap["min"]; ok { - if mv := toFloat64(minVal); mv != nil && *num < *mv { - return nil, fmt.Errorf("field %s: value %v below min %v", key, *num, *mv) - } - } - if maxVal, ok := defMap["max"]; ok { - if mv := toFloat64(maxVal); mv != nil && *num > *mv { - return nil, fmt.Errorf("field %s: value %v above max %v", key, *num, *mv) - } - } - case "bool": - if _, ok := val.(bool); !ok { - return nil, fmt.Errorf("field %s: expected bool, got %T", key, val) - } - case "string": - if _, ok := val.(string); !ok { - return nil, fmt.Errorf("field %s: expected string, got %T", key, val) - } - } - } - - // enum: value must be one of the allowed members (works for any type). - // Compared after JSON decode, so numbers are float64 on both sides. - if rawEnum, ok := defMap["enum"]; ok { - allowed, ok := rawEnum.([]interface{}) - if !ok { - return nil, fmt.Errorf("field %s: schema 'enum' must be a list", key) - } - if !enumContains(allowed, val) { - return nil, fmt.Errorf("field %s: value %v not in allowed set %v", key, val, allowed) - } - } - } + if len(result.Choices) == 0 { + return completion, fmt.Errorf("LLM returned no choices") } - - return parsed, nil + completion.Text = result.Choices[0].Message.Content + return completion, nil } // --- Operator Channel Interface --- @@ -1266,8 +1079,7 @@ func runPipeline(cfg *config.Config, pipeline config.PipelineConfig, budget *Bud } }() log.Printf("[pipeline:%s] starting (run %s)", pipeline.Name, runID) - budget.tokensUsedPipeline = 0 - budget.costPipeline = 0 + budget = budget.newRun(cfg.Budgets.PerPipelineTokens) // Observability: one pipeline span (always) + one span per step. Steps that // run to completion emit at the bottom of the loop; a step that halts the @@ -1285,7 +1097,7 @@ func runPipeline(cfg *config.Config, pipeline config.PipelineConfig, budget *Bud status := "ok" fields := map[string]interface{}{ "steps_completed": stepsCompleted, - "tokens": budget.tokensUsedPipeline, + "tokens": budget.snapshot().tokensUsedPipeline, } if err != nil { status = "error" @@ -1559,11 +1371,6 @@ Description: We need an experienced LLM engineer to build a retrieval-augmented } case "ai": - // Budget pre-flight - if err := budget.check(cfg.Budgets.PerDayTokens, cfg.Budgets.PerStepTokens); err != nil { - log.Printf("[pipeline:%s][step:%s] %s", pipeline.Name, step.Name, err) - return err - } // Resolve skill or use inline prompt prompt := step.Prompt @@ -1571,6 +1378,9 @@ Description: We need an experienced LLM engineer to build a retrieval-augmented schema := step.OutputSchema if step.Skill != "" { + if skills == nil { + return fmt.Errorf("[step:%s] skill registry unavailable: %s", step.Name, step.Skill) + } skill, ok := skills.Get(step.Skill) if !ok { return fmt.Errorf("[step:%s] unknown skill: %s", step.Name, step.Skill) @@ -1600,15 +1410,13 @@ Description: We need an experienced LLM engineer to build a retrieval-augmented } aiCtx, aiCancel := context.WithTimeout(ctx, aiTimeout) - resp, err := callLLM(aiCtx, cfg, role, prompt) + resp, err := callLLM(aiCtx, cfg, role, prompt, budget) aiCancel() if err != nil { return fmt.Errorf("[step:%s] LLM call failed: %w", step.Name, err) } // Record token usage - budget.record(resp.InputTokens + resp.OutputTokens) - budget.RecordCost(resp.CostUSD) log.Printf("[pipeline:%s][step:%s] model=%s tokens=%d+%d cost=$%.4f latency=%dms", pipeline.Name, step.Name, resp.Model, resp.InputTokens, resp.OutputTokens, resp.CostUSD, resp.LatencyMs) @@ -1624,7 +1432,7 @@ Description: We need an experienced LLM engineer to build a retrieval-augmented return fmt.Errorf("[step:%s] output validation failed: %w", step.Name, err) } data["ai_output"] = parsed - log.Printf("[pipeline:%s][step:%s] output validated: %v", pipeline.Name, step.Name, parsed) + log.Printf("[pipeline:%s][step:%s] output validated (%d fields)", pipeline.Name, step.Name, len(parsed)) } else { data["ai_output"] = resp.Text } @@ -1637,15 +1445,15 @@ Description: We need an experienced LLM engineer to build a retrieval-augmented switch v := aiOutput.(type) { case map[string]interface{}: - score, _ := v["score"].(float64) + score := outputScore(v["score"]) reason, _ := v["reason"].(string) reject, _ := v["reject"].(bool) status := "MATCH" if reject { status = "REJECT" } - draftMsg = fmt.Sprintf("[draftcat] %s - Score: %d/5\n\n%s\n\n%v", - status, int(score), reason, data["input"]) + draftMsg = fmt.Sprintf("[draftcat] %s - Score: %s/5\n\n%s\n\n%v", + status, score, reason, data["input"]) default: draftMsg = fmt.Sprintf("[draftcat] Draft for review:\n\n%v", v) } @@ -1704,10 +1512,10 @@ Description: We need an experienced LLM engineer to build a retrieval-augmented // only the sha256 of the exact draft shown — never the draft itself. // When a signing secret is set, each row also carries an HMAC receipt // so a later reader can prove the row wasn't altered after the fact. - recordAudit := func(decision, payload string, approvers []int64) { + recordAudit := func(decision, payload string, approvers []int64) error { obs.RecordApproval(pipeline.Name, step.Name, decision) if state == nil { - return + return nil } var opID int64 got := 0 @@ -1724,9 +1532,7 @@ Description: We need an experienced LLM engineer to build a retrieval-augmented envelope := newApprovalEnvelope(approvalSecret, runID, pipeline.Name, step, decidedAt, decidedAt.Add(approvalTimeout), decision, opID, payloadHash, quorumN, got, "human-approval") - if e := state.RecordApprovalV2(envelope); e != nil { - log.Printf("[pipeline:%s][step:%s] audit write failed: %v", pipeline.Name, step.Name, e) - } + return persistApprovalReceipt(envelope) } // maxAdjust bounds rewrite cycles: single-approver keeps today's single @@ -1747,7 +1553,7 @@ Description: We need an experienced LLM engineer to build a retrieval-augmented // exemption is written to the audit trail as "policy_approve" with the // rule that fired, so the trail never conflates a policy release with a // human decision. - if rule := cfg.Policy.AutoApproves(pipeline.Name, step, budget.costPipeline); rule != nil { + if rule := cfg.Policy.AutoApproves(pipeline.Name, step, budget.snapshot().costPipeline); rule != nil { sum := sha256.Sum256([]byte(currentDraft)) ph := hex.EncodeToString(sum[:]) reason := rule.Reason @@ -1761,8 +1567,8 @@ Description: We need an experienced LLM engineer to build a retrieval-augmented envelope := newApprovalEnvelope(approvalSecret, runID, pipeline.Name, step, decidedAt, decidedAt.Add(approvalTimeout), "policy_approve", 0, ph, quorumN, quorumN, reason) - if e := state.RecordApprovalV2(envelope); e != nil { - log.Printf("[pipeline:%s][step:%s] audit write failed: %v", pipeline.Name, step.Name, e) + if e := persistApprovalReceipt(envelope); e != nil { + return fmt.Errorf("[step:%s] approval receipt unavailable: %w", step.Name, e) } } data["approved"] = true @@ -1780,7 +1586,7 @@ Description: We need an experienced LLM engineer to build a retrieval-augmented pendingID, perr := state.BeginApproval(pipeline.Name, step.Name, hex.EncodeToString(pendingSum[:]), quorumN, openedAt, openedAt.Add(approvalTimeout)) if perr != nil { - log.Printf("[pipeline:%s][step:%s] pending-approval write failed: %v", pipeline.Name, step.Name, perr) + return fmt.Errorf("[step:%s] pending approval unavailable: %w", step.Name, perr) } // withStep puts the step name on the context so a channel that @@ -1793,50 +1599,57 @@ Description: We need an experienced LLM engineer to build a retrieval-augmented // The gate reached a terminal state in-process, whatever it was — // close the pending row so boot does not later call it interrupted. if rerr := state.ResolveApproval(pendingID); rerr != nil { - log.Printf("[pipeline:%s][step:%s] pending-approval resolve failed: %v", pipeline.Name, step.Name, rerr) + return fmt.Errorf("[step:%s] approval settlement unavailable: %w", step.Name, rerr) } if aerr != nil { - recordAudit("timeout", currentDraft, nil) + if auditErr := recordAudit("timeout", currentDraft, nil); auditErr != nil { + return fmt.Errorf("[step:%s] approval receipt unavailable: %w", step.Name, auditErr) + } return fmt.Errorf("[step:%s] %w", step.Name, aerr) } log.Printf("[pipeline:%s][step:%s] operator decision: %s", pipeline.Name, step.Name, action) if action == "approve" { - recordAudit("approve", currentDraft, approvers) + if auditErr := recordAudit("approve", currentDraft, approvers); auditErr != nil { + return fmt.Errorf("[step:%s] approval receipt unavailable: %w", step.Name, auditErr) + } data["approved"] = true break } if action == "skip" { - recordAudit("skip", currentDraft, nil) + if auditErr := recordAudit("skip", currentDraft, nil); auditErr != nil { + return fmt.Errorf("[step:%s] approval receipt unavailable: %w", step.Name, auditErr) + } data["approved"] = false return nil } if action != "adjust" { // e.g. a quorum "timeout" returned without an error - recordAudit("timeout", currentDraft, nil) + if auditErr := recordAudit("timeout", currentDraft, nil); auditErr != nil { + return fmt.Errorf("[step:%s] approval receipt unavailable: %w", step.Name, auditErr) + } return fmt.Errorf("[step:%s] approval not completed (%s)", step.Name, action) } // action == "adjust" - recordAudit("adjust", currentDraft, nil) + if auditErr := recordAudit("adjust", currentDraft, nil); auditErr != nil { + return fmt.Errorf("[step:%s] approval receipt unavailable: %w", step.Name, auditErr) + } adjustCycles++ log.Printf("[pipeline:%s][step:%s] adjustment: %q", pipeline.Name, step.Name, adjustText) if adjustCycles > maxAdjust { if step.Quorum >= 2 { _ = ch.Send(fmt.Sprintf("Adjustment limit (%d) reached without quorum approval — halting.", maxAdjust)) - recordAudit("quorum_fail", currentDraft, nil) + if auditErr := recordAudit("quorum_fail", currentDraft, nil); auditErr != nil { + return fmt.Errorf("[step:%s] approval receipt unavailable: %w", step.Name, auditErr) + } } data["approved"] = false return nil } - if err := budget.check(cfg.Budgets.PerDayTokens, cfg.Budgets.PerStepTokens); err != nil { - ch.Send(fmt.Sprintf("Budget exceeded, cannot rewrite: %s", err)) - return err - } - adjustPrompt := fmt.Sprintf("Original output:\n%s\n\nOperator feedback:\n%s\n\nRewrite incorporating the feedback. Respond with ONLY valid JSON in the same format.", data["ai_raw"], adjustText) aiTimeout, _ := time.ParseDuration(cfg.Timeouts.AICall) @@ -1844,14 +1657,12 @@ Description: We need an experienced LLM engineer to build a retrieval-augmented aiTimeout = 30 * time.Second } aiCtx, aiCancel := context.WithTimeout(ctx, aiTimeout) - resp, rerr := callLLM(aiCtx, cfg, "drafter", adjustPrompt) + resp, rerr := callLLM(aiCtx, cfg, "drafter", adjustPrompt, budget) aiCancel() if rerr != nil { _ = ch.Send(fmt.Sprintf("Rewrite failed: %s", rerr)) return rerr } - budget.record(resp.InputTokens + resp.OutputTokens) - budget.RecordCost(resp.CostUSD) // Revised draft re-enters the gate (at 0/N for quorum steps). currentDraft = fmt.Sprintf("[draftcat] Revised:\n\n%s", resp.Text) @@ -1867,7 +1678,7 @@ Description: We need an experienced LLM engineer to build a retrieval-augmented } allDone = true - log.Printf("[pipeline:%s] completed. tokens_used=%d", pipeline.Name, budget.tokensUsedPipeline) + log.Printf("[pipeline:%s] completed. tokens_used=%d", pipeline.Name, budget.snapshot().tokensUsedPipeline) return nil } @@ -2115,20 +1926,14 @@ func handleEmails(args string, bot *TGBot, cfg *config.Config, budget *BudgetTra bot.Send(header + formatted) } else { // Use LLM to summarize - if err := budget.check(cfg.Budgets.PerDayTokens, 1024); err != nil { - bot.Send(header + formatted[:2000] + "\n\n[truncated]") - return - } bot.sendTyping() ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) resp, err := callLLM(ctx, cfg, "classifier", fmt.Sprintf( - "Summarize these emails in a brief list. For each: sender, subject, 1-line summary. Be concise.\n\n%s", formatted)) + "Summarize these emails in a brief list. For each: sender, subject, 1-line summary. Be concise.\n\n%s", formatted), budget) cancel() if err != nil { bot.Send(header + formatted[:2000] + "\n\n[truncated]") } else { - budget.record(resp.InputTokens + resp.OutputTokens) - budget.RecordCost(resp.CostUSD) bot.Send(header + resp.Text) } } @@ -2204,10 +2009,6 @@ func handleReply(args string, bot *TGBot, cfg *config.Config, budget *BudgetTrac } else { // AI drafts a reply bot.sendTyping() - if err := budget.check(cfg.Budgets.PerDayTokens, 1024); err != nil { - bot.Send("[reply] Budget limit reached.") - return - } ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) prompt := fmt.Sprintf(`Draft a brief, professional reply to this email. Just the reply body, no subject line or headers. @@ -2215,14 +2016,12 @@ From: %s Subject: %s Body: %s`, fullEmail.From, fullEmail.Subject, fullEmail.Body) - resp, err := callLLM(ctx, cfg, "drafter", prompt) + resp, err := callLLM(ctx, cfg, "drafter", prompt, budget) cancel() if err != nil { bot.Send(fmt.Sprintf("[reply] Draft failed: %s", err)) return } - budget.record(resp.InputTokens + resp.OutputTokens) - budget.RecordCost(resp.CostUSD) replyBody = resp.Text } @@ -2343,25 +2142,18 @@ func handleThread(args string, bot *TGBot, cfg *config.Config, budget *BudgetTra threadText := gmailapi.FormatThreadForPrompt(threadEmails, "rio@ramaris.app") // Summarize with LLM - if err := budget.check(cfg.Budgets.PerDayTokens, 2048); err != nil { - // No budget — send raw - bot.Send(fmt.Sprintf("[thread] %d messages in thread:\n\n%s", len(threadEmails), threadText[:min(3500, len(threadText))])) - return - } ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) prompt := fmt.Sprintf(`Summarize this email thread. Show the back-and-forth between sent and received messages. Include key points, decisions, and any action items. Be concise. Thread (%d messages): %s`, len(threadEmails), threadText) - resp, err := callLLM(ctx, cfg, "classifier", prompt) + resp, err := callLLM(ctx, cfg, "classifier", prompt, budget) cancel() if err != nil { bot.Send(fmt.Sprintf("[thread] %d messages:\n\n%s", len(threadEmails), threadText[:min(3500, len(threadText))])) return } - budget.record(resp.InputTokens + resp.OutputTokens) - budget.RecordCost(resp.CostUSD) bot.Send(fmt.Sprintf("[thread] %d messages — %s\n\n%s", len(threadEmails), target.Subject, resp.Text)) } @@ -2544,10 +2336,16 @@ func handleStatus(bot *TGBot, budget *BudgetTracker, sched *Scheduler, cfg *conf // without waiting for the next approval prompt to tell them — and open gates // are counted so a stuck approval is one line away instead of invisible. func statusReport(active, paused int, budget *BudgetTracker, cfg *config.Config, open []statestore.PendingApproval, openErr error, now time.Time) string { + budget = budget.snapshotAt(now) lines := []string{ "[status] Engine running", fmt.Sprintf("Pipelines: %d active, %d paused", active, paused), } + if budget.stateErr != nil { + lines = append(lines, "Budget ledger unavailable; model calls blocked.") + } else if budget.unsettled > 0 { + lines = append(lines, fmt.Sprintf("Provider usage unresolved: %d call(s); reconcile before retrying.", budget.unsettled)) + } if cap := cfg.Budgets.PerDayTokens; cap > 0 { lines = append(lines, fmt.Sprintf("Tokens today: %d / %d (%.0f%% left)", budget.tokensUsedToday, cap, pctLeft(float64(budget.tokensUsedToday), float64(cap)))) @@ -2675,6 +2473,8 @@ func main() { os.Exit(runRunsCmd(os.Args[2:])) case "pending": os.Exit(runPendingCmd(os.Args[2:])) + case "budget": + os.Exit(runBudgetCmd(os.Args[2:])) case "receipts": os.Exit(runReceiptsCmd(os.Args[2:])) case "hitl": @@ -2687,6 +2487,7 @@ func main() { fmt.Println(" draftcat validate [--strict] lint config + skills, exit non-zero on errors") fmt.Println(" draftcat test dry-run a pipeline using fixtures//") fmt.Println(" draftcat runs [pipeline] [--json] recent runs + the approval decisions in each") + fmt.Println(" draftcat budget inspect spend or reconcile uncertain provider usage") fmt.Println(" draftcat pending [--json] approval gates waiting on a human right now") fmt.Println(" draftcat receipts inspect, export, or verify receipts") fmt.Println(" draftcat audit-verify check approval-receipt signatures (needs DRAFTCAT_APPROVAL_SECRET)") @@ -2869,6 +2670,9 @@ func main() { dayCostLimit: cfg.Budgets.PerDayCost, pipelineCostLimit: cfg.Budgets.PerPipelineCost, } + if err := budget.attachStore(state); err != nil { + log.Printf("[budget] usage ledger unavailable; model calls blocked: %v", err) + } chatHistory := newChatHistory(20) // keep last 20 turns // Any approval gate still marked pending was open when this process last @@ -3022,14 +2826,12 @@ Rules: - Questions about what you can do = {"intent":"chat"} - IMPORTANT: "from:" means received FROM someone. "to:" means sent TO someone. If the user asks what THEY sent/replied to someone, use "to: in:sent"`, emailCtx, text) - intentResp, err := callLLM(intentCtx, &cfg, "classifier", intentPrompt) + intentResp, err := callLLM(intentCtx, &cfg, "classifier", intentPrompt, budget) intentCancel() if err != nil { log.Printf("[intent] classifier error: %v", err) } else { - budget.record(intentResp.InputTokens + intentResp.OutputTokens) - budget.RecordCost(intentResp.CostUSD) // Parse intent cleaned := strings.TrimSpace(intentResp.Text) if strings.HasPrefix(cleaned, "```") { @@ -3079,9 +2881,7 @@ Rules: // Regular message — respond via LLM with conversation history chatHistory.Add("user", text) - if err := budget.check(cfg.Budgets.PerDayTokens, 512); err != nil { - bot.Send("Budget limit reached.") - } else { + { aiCtx, aiCancel := context.WithTimeout(context.Background(), 15*time.Second) var skillList, pipelineList string for _, s := range skillReg.List() { @@ -3117,14 +2917,12 @@ When the operator asks about emails, you can fetch them directly. When they ask Conversation so far: %s`, pipelineList, skillList, gmailStatus, history) - resp, err := callLLM(aiCtx, &cfg, "drafter", sysPrompt) + resp, err := callLLM(aiCtx, &cfg, "drafter", sysPrompt, budget) aiCancel() if err != nil { log.Printf("[msg] LLM error: %v", err) bot.Send("Commands: /help /cron /skills /run /status") } else { - budget.record(resp.InputTokens + resp.OutputTokens) - budget.RecordCost(resp.CostUSD) log.Printf("[msg] LLM reply (%d tokens, %dms): %s", resp.InputTokens+resp.OutputTokens, resp.LatencyMs, resp.Text[:min(80, len(resp.Text))]) chatHistory.Add("assistant", resp.Text) if err := bot.Send(resp.Text); err != nil { diff --git a/model_policy.go b/model_policy.go index 3ff5d85..f653b03 100644 --- a/model_policy.go +++ b/model_policy.go @@ -35,7 +35,10 @@ func enforceModelPolicy(ctx context.Context, cfg *config.Config, role, phase, te expires := time.Now() if action == "review" { if opChan == nil { - recordModelPolicyDecision(ctx, cfg, role, phase, text, rule, decision, operatorID, expires) + recordErr := recordModelPolicyDecision(ctx, cfg, role, phase, text, rule, decision, operatorID, expires) + if recordErr != nil { + return fmt.Errorf("model policy %q requires review but no operator channel is running (receipt: %w)", rule.ID, recordErr) + } return fmt.Errorf("model policy %q requires review but no operator channel is running", rule.ID) } timeout, _ := time.ParseDuration(cfg.Timeouts.OperatorApproval) @@ -50,7 +53,9 @@ func enforceModelPolicy(ctx context.Context, cfg *config.Config, role, phase, te cancel() if reviewErr == nil && dec.Action == "approve" { decision, operatorID = "approve", dec.ApproverID - recordModelPolicyDecision(ctx, cfg, role, phase, text, rule, decision, operatorID, expires) + if err := recordModelPolicyDecision(ctx, cfg, role, phase, text, rule, decision, operatorID, expires); err != nil { + return fmt.Errorf("model policy approval receipt unavailable: %w", err) + } return nil } if reviewErr != nil { @@ -59,14 +64,16 @@ func enforceModelPolicy(ctx context.Context, cfg *config.Config, role, phase, te decision = dec.Action } } - recordModelPolicyDecision(ctx, cfg, role, phase, text, rule, decision, operatorID, expires) + if err := recordModelPolicyDecision(ctx, cfg, role, phase, text, rule, decision, operatorID, expires); err != nil { + return fmt.Errorf("model policy decision receipt unavailable: %w", err) + } return fmt.Errorf("model %s blocked by policy %q: %s", phase, rule.ID, rule.Reason) } func recordModelPolicyDecision(ctx context.Context, cfg *config.Config, role, phase, text string, - rule *config.ModelPolicyRule, decision string, operatorID int64, expires time.Time) { + rule *config.ModelPolicyRule, decision string, operatorID int64, expires time.Time) error { if state == nil { - return + return nil } sum := sha256.Sum256([]byte(text)) payloadHash := hex.EncodeToString(sum[:]) @@ -75,9 +82,7 @@ func recordModelPolicyDecision(ctx context.Context, cfg *config.Config, role, ph envelope := newApprovalEnvelope([]byte(os.Getenv("DRAFTCAT_APPROVAL_SECRET")), runIDFromContext(ctx), pipelineFromContext(ctx), step, time.Now(), expires, decision, operatorID, payloadHash, 1, boolCount(decision == "approve"), policy) - if err := state.RecordApprovalV2(envelope); err != nil { - return - } + return persistApprovalReceipt(envelope) } func boolCount(v bool) int { diff --git a/output_validation.go b/output_validation.go new file mode 100644 index 0000000..95a72e3 --- /dev/null +++ b/output_validation.go @@ -0,0 +1,91 @@ +package main + +import ( + "fmt" + "strings" + + "github.com/renezander030/draftcat/internal/outputschema" +) + +// enumContains retains the flat schema helper's exact scalar comparisons. +func enumContains(allowed []interface{}, value interface{}) bool { + return outputschema.EnumContains(allowed, value) +} + +func validateOutput(text string, schema map[string]interface{}) (map[string]interface{}, error) { + if len(schema) == 0 { + return nil, nil + } + if findings := outputschema.Check(schema); len(findings) > 0 { + return nil, fmt.Errorf("field %s: invalid output schema: %s", findings[0].Field, findings[0].Message) + } + + cleaned := strings.TrimSpace(text) + if strings.HasPrefix(cleaned, "```") { + lines := strings.Split(cleaned, "\n") + if len(lines) < 3 || strings.TrimSpace(lines[len(lines)-1]) != "```" || + (strings.TrimSpace(lines[0]) != "```" && strings.TrimSpace(lines[0]) != "```json") { + return nil, fmt.Errorf("output has an invalid JSON code fence") + } + cleaned = strings.Join(lines[1:len(lines)-1], "\n") + } + + var parsed map[string]interface{} + if err := decodeStrictJSON([]byte(cleaned), &parsed); err != nil || parsed == nil { + // Decoder details can include model-controlled field names. Keep model + // output out of errors, which flow into logs and operator notifications. + return nil, fmt.Errorf("output must be one unambiguous JSON object") + } + + for _, key := range outputschema.Fields(schema) { + value, exists := parsed[key] + if !exists { + return nil, fmt.Errorf("missing required field: %s", key) + } + def, ok := schema[key].(map[string]interface{}) + if !ok { + return nil, fmt.Errorf("field %s: invalid output schema definition", key) + } + if typeName, hasType := def["type"].(string); hasType { + if !outputschema.MatchesType(value, typeName) { + return nil, fmt.Errorf("field %s: expected %s", key, typeName) + } + if typeName == "int" || typeName == "number" { + number, _ := outputschema.Number(value) + if min, present := def["min"]; present { + bound, _ := outputschema.Number(min) + if number.Cmp(bound) < 0 { + return nil, fmt.Errorf("field %s: value below min", key) + } + } + if max, present := def["max"]; present { + bound, _ := outputschema.Number(max) + if number.Cmp(bound) > 0 { + return nil, fmt.Errorf("field %s: value above max", key) + } + } + } + } + if enum, present := def["enum"]; present { + allowed, ok := enum.([]interface{}) + if !ok { + return nil, fmt.Errorf("field %s: invalid output schema enum", key) + } + if !enumContains(allowed, value) { + return nil, fmt.Errorf("field %s: value not in allowed set", key) + } + } + } + return parsed, nil +} + +// outputScore formats integer scores without converting through float64. +func outputScore(value interface{}) string { + if number, ok := outputschema.Number(value); ok { + if number.IsInt() { + return number.Num().String() + } + return fmt.Sprint(value) + } + return "0" +} diff --git a/output_validation_test.go b/output_validation_test.go new file mode 100644 index 0000000..b1668d0 --- /dev/null +++ b/output_validation_test.go @@ -0,0 +1,152 @@ +package main + +import ( + "encoding/json" + "strings" + "testing" + + "github.com/renezander030/draftcat/internal/config" + skillsapi "github.com/renezander030/draftcat/internal/skills" + "gopkg.in/yaml.v3" +) + +func TestValidateOutputIntegerAndExactNumbers(t *testing.T) { + for _, tc := range []struct { + name, input string + definition map[string]interface{} + wantError string + }{ + {"fractional integer", `{"value":3.5}`, map[string]interface{}{"type": "int"}, "expected int"}, + {"integral decimal", `{"value":3.0}`, map[string]interface{}{"type": "int"}, ""}, + {"integral exponent", `{"value":3e2}`, map[string]interface{}{"type": "int"}, ""}, + {"fractional exponent", `{"value":3e-2}`, map[string]interface{}{"type": "int"}, "expected int"}, + {"exact large bound", `{"value":9007199254740993}`, map[string]interface{}{"type": "int", "max": int64(9007199254740992)}, "above max"}, + {"exact decimal bound", `{"value":0.100000000000000001}`, map[string]interface{}{"type": "number", "max": 0.1}, "above max"}, + {"exact large enum", `{"value":9007199254740993}`, map[string]interface{}{"type": "int", "enum": []interface{}{int64(9007199254740993)}}, ""}, + {"large enum neighbor", `{"value":9007199254740992}`, map[string]interface{}{"type": "int", "enum": []interface{}{int64(9007199254740993)}}, "not in allowed set"}, + {"numeric string", `{"value":"3"}`, map[string]interface{}{"type": "number"}, "expected number"}, + } { + t.Run(tc.name, func(t *testing.T) { + parsed, err := validateOutput(tc.input, map[string]interface{}{"value": tc.definition}) + if tc.wantError != "" { + if err == nil || !strings.Contains(err.Error(), tc.wantError) { + t.Fatalf("wanted %q, got %v", tc.wantError, err) + } + return + } + if err != nil { + t.Fatal(err) + } + if _, ok := parsed["value"].(json.Number); !ok { + t.Fatalf("numeric output lost exact JSON representation: %T", parsed["value"]) + } + encoded, err := json.Marshal(parsed) + if err != nil || string(encoded) != tc.input { + t.Fatalf("numeric token changed: %s, %v", encoded, err) + } + }) + } +} + +func TestValidateOutputRejectsAmbiguousObjects(t *testing.T) { + schema := map[string]interface{}{"value": map[string]interface{}{"type": "int"}} + for _, input := range []string{ + `{"value":1,"value":2}`, + `{"value":1,"extra":{"secret":"first","secret":"second"}}`, + `{"value":1} {"value":2}`, + `null`, `[]`, `1`, `"text"`, + "```json\n{\"value\":1}\nnot a closing fence", + "```json\n{\"value\":1}\n```\ntrailing prose", + } { + if _, err := validateOutput(input, schema); err == nil { + t.Errorf("accepted ambiguous/non-object output %q", input) + } + } +} + +func TestValidateOutputPreservesExtraFieldsAndNumericTokens(t *testing.T) { + input := `{"value":1,"extra":{"large":9007199254740993}}` + parsed, err := validateOutput(input, map[string]interface{}{"value": map[string]interface{}{"type": "int"}}) + if err != nil { + t.Fatal(err) + } + extra, ok := parsed["extra"].(map[string]interface{}) + if !ok || extra["large"] != json.Number("9007199254740993") { + t.Fatalf("extra fields changed: %+v", parsed) + } +} + +func TestValidateOutputCompositeEnumFailsWithoutPanic(t *testing.T) { + schema := map[string]interface{}{"value": map[string]interface{}{"enum": []interface{}{"safe"}}} + for _, input := range []string{`{"value":{}}`, `{"value":[]}`} { + if _, err := validateOutput(input, schema); err == nil { + t.Errorf("accepted composite enum value %q", input) + } + } + // A malformed configured member is rejected before value comparison. + bad := map[string]interface{}{"value": map[string]interface{}{"enum": []interface{}{map[string]interface{}{"secret": "hidden"}}}} + if _, err := validateOutput(`{"value":{}}`, bad); err == nil { + t.Fatal("accepted composite enum definition") + } +} + +func TestValidateOutputErrorsOmitModelData(t *testing.T) { + schema := map[string]interface{}{"value": map[string]interface{}{"type": "string", "enum": []interface{}{"safe"}}} + for _, input := range []string{ + `{"value":"sensitive-customer-data"}`, + `{"value":"sensitive-customer-data", broken}`, + `{"sensitive-customer-data":1,"sensitive-customer-data":2}`, + } { + _, err := validateOutput(input, schema) + if err == nil || strings.Contains(err.Error(), "sensitive-customer-data") { + t.Fatalf("error disclosed model data: %v", err) + } + } +} + +func TestValidateOutputReportsFieldsDeterministically(t *testing.T) { + schema := map[string]interface{}{"z": map[string]interface{}{"type": "int"}, "a": map[string]interface{}{"type": "int"}} + for i := 0; i < 20; i++ { + _, err := validateOutput(`{}`, schema) + if err == nil || err.Error() != "missing required field: a" { + t.Fatalf("unstable error: %v", err) + } + } +} + +func TestOutputScoreKeepsExactIntegers(t *testing.T) { + for _, tc := range []struct { + value interface{} + want string + }{ + {json.Number("3"), "3"}, {json.Number("3.0"), "3"}, + {json.Number("9007199254740993"), "9007199254740993"}, {float64(4), "4"}, {nil, "0"}, + } { + if got := outputScore(tc.value); got != tc.want { + t.Errorf("score %v rendered as %q, want %q", tc.value, got, tc.want) + } + } +} + +func TestValidateOutputUsesExactYAMLStepAndSkillSchemas(t *testing.T) { + var step config.StepConfig + var skill skillsapi.SkillDef + definition := []byte("output_schema:\n decimal: {type: number, max: 0.100000000000000001}\n large: {type: int, enum: [9007199254740993]}\n") + if err := yaml.Unmarshal(definition, &step); err != nil { + t.Fatal(err) + } + if err := yaml.Unmarshal(definition, &skill); err != nil { + t.Fatal(err) + } + for _, schema := range []map[string]interface{}{step.OutputSchema, skill.OutputSchema} { + if _, err := validateOutput(`{"decimal":0.100000000000000001,"large":9007199254740993}`, schema); err != nil { + t.Fatalf("exact configured decimals/integers changed: %v", err) + } + if _, err := validateOutput(`{"decimal":0.100000000000000002,"large":9007199254740993}`, schema); err == nil { + t.Fatal("accepted value one decimal unit above exact YAML max") + } + if _, err := validateOutput(`{"decimal":0.1,"large":9007199254740992}`, schema); err == nil { + t.Fatal("accepted rounded neighbor of exact YAML enum") + } + } +} diff --git a/package-lock.json b/package-lock.json index 5a4e570..e90a482 100644 --- a/package-lock.json +++ b/package-lock.json @@ -1,12 +1,12 @@ { "name": "draftcat", - "version": "0.8.0", + "version": "0.9.0", "lockfileVersion": 3, "requires": true, "packages": { "": { "name": "draftcat", - "version": "0.8.0", + "version": "0.9.0", "cpu": [ "x64", "arm64" diff --git a/package.json b/package.json index f714cc3..db2ad0d 100644 --- a/package.json +++ b/package.json @@ -1,6 +1,6 @@ { "name": "draftcat", - "version": "0.8.0", + "version": "0.9.0", "description": "Governed AI pipelines with human approval gates", "license": "MIT", "author": "Rene Zander", diff --git a/provider_bounds_test.go b/provider_bounds_test.go new file mode 100644 index 0000000..a1153f5 --- /dev/null +++ b/provider_bounds_test.go @@ -0,0 +1,94 @@ +package main + +import ( + "context" + "errors" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" + "time" +) + +type interruptedProviderReader struct{} + +func (interruptedProviderReader) Read(p []byte) (int, error) { + copy(p, "partial provider data") + return 21, io.ErrUnexpectedEOF +} + +func TestProviderResponseReadBoundsAndErrors(t *testing.T) { + if _, err := readProviderResponse(context.Background(), strings.NewReader(strings.Repeat("x", int(maxProviderResponseBytes+1)))); err == nil || !strings.Contains(err.Error(), "exceeds") { + t.Fatalf("oversized response accepted: %v", err) + } + if _, err := readProviderResponse(context.Background(), interruptedProviderReader{}); !errors.Is(err, io.ErrUnexpectedEOF) { + t.Fatalf("body read error ignored: %v", err) + } + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if _, err := readProviderResponse(ctx, strings.NewReader("data")); !errors.Is(err, context.Canceled) { + t.Fatalf("canceled body read accepted: %v", err) + } +} + +func TestBudgetOversizedResponseRemainsUnsettledAndIsNotRetried(t *testing.T) { + var hits atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hits.Add(1) + _, _ = io.WriteString(w, strings.Repeat("secret-provider-data", int(maxProviderResponseBytes)/20+100)) + })) + defer srv.Close() + cfg := providerBudgetConfig(srv.URL) + st := providerBudgetStore(t) + b := &BudgetTracker{dayStart: time.Now()} + if err := b.attachStore(st); err != nil { + t.Fatal(err) + } + _, err := callLLM(context.Background(), cfg, "drafter", "p", b) + if err == nil || !strings.Contains(err.Error(), "exceeds") || strings.Contains(err.Error(), "secret-provider-data") { + t.Fatalf("unsafe oversized response error: %v", err) + } + if _, err := callLLM(context.Background(), cfg, "drafter", "p", b); err == nil || !strings.Contains(err.Error(), "unresolved") { + t.Fatalf("uncertain bill reopened budget: %v", err) + } + if hits.Load() != 1 { + t.Fatalf("oversized paid response retried %d times", hits.Load()) + } +} + +func TestCallLLMServerFailureIsNotRetriedOrDumped(t *testing.T) { + var hits atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hits.Add(1) + w.WriteHeader(503) + _, _ = io.WriteString(w, `{"error":"secret-provider-data"}`) + })) + defer srv.Close() + cfg := providerBudgetConfig(srv.URL) + b := &BudgetTracker{dayStart: time.Now()} + _, err := callLLM(context.Background(), cfg, "drafter", "p", b) + if err == nil || hits.Load() != 1 || strings.Contains(err.Error(), "secret-provider-data") { + t.Fatalf("hits=%d unsafe error=%v", hits.Load(), err) + } + if b.snapshot().unsettled != 1 { + t.Fatal("uncertain server failure was treated as free") + } +} + +func TestBudgetTruncatedPaidResponseIsNotRetried(t *testing.T) { + var hits atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hits.Add(1) + w.Header().Set("Content-Length", "1000") + _, _ = io.WriteString(w, `{"choices":[`) + })) + defer srv.Close() + cfg := providerBudgetConfig(srv.URL) + b := &BudgetTracker{dayStart: time.Now()} + _, err := callLLM(context.Background(), cfg, "drafter", "p", b) + if !errors.Is(err, io.ErrUnexpectedEOF) || hits.Load() != 1 || b.snapshot().unsettled != 1 { + t.Fatalf("truncated response: hits=%d error=%v snapshot=%+v", hits.Load(), err, b.snapshot()) + } +} diff --git a/provider_call.go b/provider_call.go new file mode 100644 index 0000000..cf0f3e8 --- /dev/null +++ b/provider_call.go @@ -0,0 +1,127 @@ +package main + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + + "github.com/renezander030/draftcat/internal/config" +) + +const maxProviderResponseBytes int64 = 4 << 20 + +type providerCallError struct { + err error + uncertain bool +} + +func (e *providerCallError) Error() string { return e.err.Error() } +func (e *providerCallError) Unwrap() error { return e.err } +func providerError(err error, uncertain bool) error { + return &providerCallError{err: err, uncertain: uncertain} +} +func uncertainProviderError(err error) bool { + var e *providerCallError + if errors.As(err, &e) { + return e.uncertain + } + return true +} + +type contextReader struct { + ctx context.Context + r io.Reader +} + +func (r contextReader) Read(p []byte) (int, error) { + if err := r.ctx.Err(); err != nil { + return 0, err + } + return r.r.Read(p) +} + +func readProviderResponse(ctx context.Context, body io.Reader) ([]byte, error) { + data, err := io.ReadAll(io.LimitReader(contextReader{ctx: ctx, r: body}, maxProviderResponseBytes+1)) + if err != nil { + return nil, fmt.Errorf("read provider response: %w", err) + } + if err := ctx.Err(); err != nil { + return nil, fmt.Errorf("read provider response: %w", err) + } + if int64(len(data)) > maxProviderResponseBytes { + return nil, fmt.Errorf("provider response exceeds %d-byte limit", maxProviderResponseBytes) + } + return data, nil +} + +// callLLM is the single governance boundary for engine model calls. A tracker +// is optional for standalone client tests; every engine caller supplies one. +func callLLM(ctx context.Context, cfg *config.Config, role, prompt string, trackers ...*BudgetTracker) (*CompletionResponse, error) { + modelName, ok := cfg.Roles[role] + if !ok { + return nil, fmt.Errorf("unknown role: %s", role) + } + model, ok := cfg.Models[modelName] + if !ok { + return nil, fmt.Errorf("unknown model: %s", modelName) + } + if err := ctx.Err(); err != nil { + return nil, err + } + if err := enforceModelPolicy(ctx, cfg, role, "input", prompt); err != nil { + return nil, err + } + maxTokens := model.MaxTokens + if maxTokens <= 0 { + maxTokens = 2048 + } + if cap := cfg.Budgets.PerStepTokens; cap > 0 && maxTokens > cap { + maxTokens = cap + } + var b *BudgetTracker + if len(trackers) > 0 { + b = trackers[0] + } + var resp *CompletionResponse + var callErr error + if b == nil { + resp, callErr = providerCallLLM(ctx, cfg, role, prompt, maxTokens) + } else { + // The admission gate protects provider dispatch and settlement only. + // Human review of a paid response must not hold another run's budget. + resp, callErr = func() (*CompletionResponse, error) { + if err := b.acquire(ctx); err != nil { + return nil, fmt.Errorf("budget admission canceled: %w", err) + } + defer b.release() + id, err := b.admit(ctx, cfg, maxTokens) + if err != nil { + return nil, err + } + result, providerErr := providerCallLLM(ctx, cfg, role, prompt, maxTokens) + if err := b.finish(id, result, providerErr, cfg); err != nil { + return nil, err + } + return result, providerErr + }() + } + if callErr != nil { + return nil, callErr + } + if err := enforceModelPolicy(ctx, cfg, role, "output", resp.Text); err != nil { + return nil, err + } + return resp, nil +} + +func providerDeclaresUsage(body []byte) bool { + var object map[string]json.RawMessage + if json.Unmarshal(body, &object) != nil { + return false + } + usage, ok := object["usage"] + return ok && !bytes.Equal(bytes.TrimSpace(usage), []byte("null")) +} diff --git a/tool_gate.go b/tool_gate.go index 6dd1432..23dbe3c 100644 --- a/tool_gate.go +++ b/tool_gate.go @@ -39,7 +39,6 @@ import ( "context" "crypto/rand" "crypto/sha256" - "database/sql" "encoding/hex" "encoding/json" "errors" @@ -88,13 +87,14 @@ type ToolCallRequest struct { bindingHash string expires time.Time providedActionID bool + approverID int64 } // ToolCallResponse is the gate's answer. type ToolCallResponse struct { ActionID string `json:"action_id,omitempty"` Decision string `json:"decision"` // "allow" | "deny" | "pending" - State string `json:"state,omitempty"` // pending | allowed | denied | expired | consumed + State string `json:"state,omitempty"` // pending | allowed | denied | expired | consumed | revoked Permit string `json:"permit,omitempty"` // "execute" only on the first successful consume Reason string `json:"reason"` ArgsHash string `json:"args_hash"` @@ -107,10 +107,15 @@ type ToolCallResponse struct { // ApprovalID identifies a decision that needs a human. It is set on every // human-path response, so a sync caller that later loses the connection // can still collect the decision. - ApprovalID string `json:"approval_id,omitempty"` - Poll string `json:"poll,omitempty"` - Consume string `json:"consume,omitempty"` - ExpiresAt string `json:"expires_at,omitempty"` + ApprovalID string `json:"approval_id,omitempty"` + Poll string `json:"poll,omitempty"` + Consume string `json:"consume,omitempty"` + ExpiresAt string `json:"expires_at,omitempty"` + ExecutionStatus string `json:"execution_status,omitempty"` + ResultHash string `json:"result_hash,omitempty"` + CompletedAt string `json:"completed_at,omitempty"` + ExecutionEvidence string `json:"execution_evidence,omitempty"` + unavailable bool } const toolGatePath = "/gate/tool-call" @@ -141,6 +146,7 @@ type toolTicket struct { decidedAt time.Time consumedAt time.Time resp ToolCallResponse + cancel context.CancelFunc } func (t *toolTicket) result() (ToolCallResponse, bool) { @@ -149,34 +155,6 @@ func (t *toolTicket) result() (ToolCallResponse, bool) { return t.resp, t.decided } -func (t *toolTicket) resolve(resp ToolCallResponse, at time.Time) { - t.mu.Lock() - defer t.mu.Unlock() - if t.decided { - return - } - t.decided = true - t.decidedAt = at - t.resp = resp - close(t.done) -} - -func (t *toolTicket) consume(binding string, at time.Time) (ToolCallResponse, bool) { - t.mu.Lock() - defer t.mu.Unlock() - if binding == "" || binding != t.BindingHash || !t.decided || t.resp.State != "allowed" || at.After(t.Expires) { - return t.resp, false - } - t.consumedAt = at - t.resp.State = "consumed" - t.resp.Permit = "" - t.resp.Reason = "permit already consumed" - out := t.resp - out.Permit = "execute" - out.Reason = "permit consumed; execute this bound action once" - return out, true -} - // pendingResponse is what a caller sees while the human has not decided. func (t *toolTicket) pendingResponse() ToolCallResponse { return ToolCallResponse{ @@ -305,6 +283,9 @@ func (g *toolGate) HandleCall(w http.ResponseWriter, r *http.Request) { rule, listed := g.cfg.ToolGate.Lookup(req.Tool) policyHash := g.policyHash(req.Tool, rule, listed) expires := g.now().Add(g.approvalWindow()) + if state != nil { + expires = expires.Truncate(time.Second) + } if strings.TrimSpace(req.ExpiresAt) != "" { requested, err := time.Parse(time.RFC3339, req.ExpiresAt) if err != nil || !requested.After(g.now()) { @@ -337,7 +318,7 @@ func (g *toolGate) HandleCall(w http.ResponseWriter, r *http.Request) { // Default deny. An unlisted tool is refused whatever else is true. if !listed { log.Printf("[tool-gate] DENY %s (not in the allowlist) agent=%q", req.Tool, req.Agent) - recordToolDecision(g.cfg, req, argsHash, "deny", 0, "unlisted") + reason := fmt.Sprintf("tool %q is not in tool_gate.tools — the gate denies by default", req.Tool) g.notifyDenial(req, argsHash, "unlisted", reason) resp := g.finishAction(req, expires, ToolCallResponse{ @@ -359,7 +340,7 @@ func (g *toolGate) HandleCall(w http.ResponseWriter, r *http.Request) { if rule.DeniesOnMismatch() { reason := "arguments outside the rule — " + why log.Printf("[tool-gate] DENY %s (args mismatch: %s) agent=%q", req.Tool, why, req.Agent) - recordToolDecision(g.cfg, req, argsHash, "deny", 0, "args mismatch: "+why) + g.notifyDenial(req, argsHash, "mismatch", reason) resp := g.finishAction(req, expires, ToolCallResponse{ Decision: "deny", Reason: reason, @@ -377,7 +358,7 @@ func (g *toolGate) HandleCall(w http.ResponseWriter, r *http.Request) { // config and reviewable, and it is recorded as one. if !needHuman { log.Printf("[tool-gate] ALLOW %s (allowlisted, risk=%s) agent=%q", req.Tool, rule.RiskOf(), req.Agent) - recordToolDecision(g.cfg, req, argsHash, "policy_approve", 0, "allowlisted risk="+rule.RiskOf()) + resp := g.finishAction(req, expires, ToolCallResponse{ Decision: "allow", Reason: "allowlisted in tool_gate", ArgsHash: argsHash, DecidedBy: "allowlist", Rule: "listed", @@ -408,7 +389,10 @@ func (g *toolGate) HandleCall(w http.ResponseWriter, r *http.Request) { if tk != nil && tk.ID != req.ActionID { if !req.providedActionID { if state != nil { - _ = state.DecideToolAction(req.ActionID, "denied", "deny", "joined identical pending action", "repeat-guard", g.now()) + if err := state.DecideToolAction(req.ActionID, "denied", "deny", "joined identical pending action", "repeat-guard", g.now()); err != nil { + writeToolDecision(w, http.StatusServiceUnavailable, unavailableToolResponse(req.ActionID)) + return + } } writeToolDecision(w, http.StatusAccepted, tk.pendingResponse()) return @@ -481,11 +465,20 @@ func (g *toolGate) HandleStatus(w http.ResponseWriter, r *http.Request) { } id := parts[0] if len(parts) == 2 { - if parts[1] != "consume" || r.Method != http.MethodPost { + if r.Method != http.MethodPost { w.WriteHeader(http.StatusMethodNotAllowed) return } - g.handleConsume(w, r, id) + switch parts[1] { + case "consume": + g.handleConsume(w, r, id) + case "revoke": + g.handleRevoke(w, r, id) + case "complete": + g.handleComplete(w, r, id) + default: + w.WriteHeader(http.StatusNotFound) + } return } if r.Method != http.MethodGet { @@ -495,19 +488,6 @@ func (g *toolGate) HandleStatus(w http.ResponseWriter, r *http.Request) { g.mu.Lock() tk := g.tickets[id] g.mu.Unlock() - if tk == nil { - if state != nil { - if a, err := state.ToolAction(id); err == nil { - writeToolDecision(w, http.StatusOK, responseFromToolAction(a)) - return - } - } - writeToolDecision(w, http.StatusNotFound, ToolCallResponse{ - ActionID: id, Decision: "deny", State: "denied", - Reason: "unknown action id - the gate has no decision for it; ask again", ApprovalID: id, - }) - return - } if q := strings.TrimSpace(r.URL.Query().Get("wait")); q != "" { d, err := time.ParseDuration(q) if err != nil || d < 0 { @@ -517,22 +497,17 @@ func (g *toolGate) HandleStatus(w http.ResponseWriter, r *http.Request) { if d > toolGateStatusWaitCap { d = toolGateStatusWaitCap } - select { - case <-tk.done: - case <-time.After(d): - case <-r.Context().Done(): - return - } - } - if resp, decided := tk.result(); decided { - if resp.State == "allowed" && g.now().After(tk.Expires) { - resp.Decision, resp.State, resp.Consume, resp.Permit = "deny", "expired", "", "" - resp.Reason = "permit expired before consumption" + if tk != nil { + select { + case <-tk.done: + case <-time.After(d): + case <-r.Context().Done(): + return + } } - writeToolDecision(w, http.StatusOK, resp) - return } - writeToolDecision(w, http.StatusOK, tk.pendingResponse()) + resp, code := g.actionResponse(id) + writeToolDecision(w, code, resp) } func (g *toolGate) handleConsume(w http.ResponseWriter, r *http.Request, id string) { @@ -547,77 +522,45 @@ func (g *toolGate) handleConsume(w http.ResponseWriter, r *http.Request, id stri http.Error(w, "malformed json", http.StatusBadRequest) return } - now := g.now() - if state != nil { - saved, err := state.ToolAction(id) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - writeToolDecision(w, http.StatusNotFound, ToolCallResponse{ActionID: id, Decision: "deny", State: "denied", Reason: "unknown action id"}) - } else { - http.Error(w, "state unavailable", http.StatusServiceUnavailable) - } - return - } - if saved.BindingHash != body.BindingHash { - writeToolDecision(w, http.StatusConflict, responseFromToolAction(saved)) - return - } - rule, listed := g.cfg.ToolGate.Lookup(saved.Tool) - if saved.PolicyHash != g.policyHash(saved.Tool, rule, listed) { - if err := state.RevokeToolAction(id, saved.PolicyHash, now); err != nil { - http.Error(w, "state unavailable", http.StatusServiceUnavailable) - return - } - resp := responseFromToolAction(saved) - resp.Decision, resp.State, resp.Permit, resp.Consume, resp.Reason = "deny", "expired", "", "", "policy changed; request a new approval" - writeToolDecision(w, http.StatusConflict, resp) - return - } - a, consumed, err := state.ConsumeToolAction(id, body.BindingHash, now) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - writeToolDecision(w, http.StatusNotFound, ToolCallResponse{ActionID: id, Decision: "deny", State: "denied", Reason: "unknown action id"}) - } else { - http.Error(w, "state unavailable", http.StatusServiceUnavailable) - } - return - } - if consumed { - recordToolConsumption(a, now) - resp := responseFromToolAction(a) - resp.Decision, resp.Permit = "allow", "execute" - resp.Reason = "permit consumed; execute this bound action once" - g.mu.Lock() - if tk := g.tickets[id]; tk != nil { - _, _ = tk.consume(body.BindingHash, now) - } - g.mu.Unlock() - writeToolDecision(w, http.StatusOK, resp) - return - } - writeToolDecision(w, http.StatusConflict, responseFromToolAction(a)) + if state == nil { + writeToolDecision(w, http.StatusServiceUnavailable, unavailableToolResponse(id)) return } - g.mu.Lock() - tk := g.tickets[id] - g.mu.Unlock() - if tk == nil { - writeToolDecision(w, http.StatusNotFound, ToolCallResponse{ActionID: id, Decision: "deny", State: "denied", Reason: "unknown action id"}) + resp, code := g.actionResponse(id) + if code != http.StatusOK { + writeToolDecision(w, code, resp) return } - rule, listed := g.cfg.ToolGate.Lookup(tk.Tool) - if tk.PolicyHash != g.policyHash(tk.Tool, rule, listed) { - resp, _ := tk.result() - resp.Decision, resp.State, resp.Permit, resp.Consume, resp.Reason = "deny", "expired", "", "", "policy changed; request a new approval" + if body.BindingHash == "" || body.BindingHash != resp.BindingHash || resp.State != "allowed" { + resp.Permit = "" writeToolDecision(w, http.StatusConflict, resp) return } - resp, ok := tk.consume(body.BindingHash, now) - if !ok { - resp.Permit = "" + a, err := state.ToolAction(id) + if err != nil { + writeToolDecision(w, http.StatusServiceUnavailable, unavailableToolResponse(id)) + return + } + now := g.now() + e, err := toolConsumptionReceipt(a, now) + if err != nil { + writeToolDecision(w, http.StatusServiceUnavailable, unavailableToolResponse(id)) + return + } + a, consumed, err := state.ConsumeToolActionWithReceipt(id, body.BindingHash, now, e) + if err != nil { + writeToolDecision(w, http.StatusServiceUnavailable, unavailableToolResponse(id)) + return + } + resp = responseFromToolAction(a) + g.updateTicket(id, resp) + if !consumed { writeToolDecision(w, http.StatusConflict, resp) return } + resp.Decision, resp.Permit = "allow", "execute" + resp.ExecutionStatus = "unreported" + resp.Reason = "permit consumed; execute this bound action once" writeToolDecision(w, http.StatusOK, resp) } @@ -716,7 +659,7 @@ func (g *toolGate) repeatCheck(req ToolCallRequest, rule config.ToolRule, argsHa case "deny": reason := fmt.Sprintf("repeat of a call denied %s ago (%s) — not asking again inside repeat_window", ago, resp.Reason) log.Printf("[tool-gate] DENY %s (repeat guard: denied %s ago) agent=%q", req.Tool, ago, req.Agent) - recordToolDecision(g.cfg, req, argsHash, "deny", 0, "repeat-guard: denied "+ago.String()+" ago") + g.notifyDenial(req, argsHash, "repeat", reason) return nil, ToolCallResponse{Decision: "deny", Reason: reason, ArgsHash: argsHash, DecidedBy: "repeat-guard", Rule: "repeat_window", ApprovalID: tk.ID}, true @@ -724,7 +667,7 @@ func (g *toolGate) repeatCheck(req ToolCallRequest, rule config.ToolRule, argsHa if rule.RememberApproval { reason := fmt.Sprintf("identical call approved %s ago — remember_approval reuses it inside repeat_window", ago) log.Printf("[tool-gate] ALLOW %s (repeat guard: approved %s ago) agent=%q", req.Tool, ago, req.Agent) - recordToolDecision(g.cfg, req, argsHash, "policy_approve", 0, "repeat-guard: remembered approval from "+ago.String()+" ago") + return nil, ToolCallResponse{Decision: "allow", Reason: reason, ArgsHash: argsHash, DecidedBy: "repeat-guard", Rule: "remember_approval", ApprovalID: tk.ID}, true } @@ -732,7 +675,7 @@ func (g *toolGate) repeatCheck(req ToolCallRequest, rule config.ToolRule, argsHa if capN := g.cfg.ToolGate.MaxRepeats; capN > 0 && asks >= capN { reason := fmt.Sprintf("asked %d times inside repeat_window — max_repeats reached, not asking again", asks) log.Printf("[tool-gate] DENY %s (repeat guard: %d asks) agent=%q", req.Tool, asks, req.Agent) - recordToolDecision(g.cfg, req, argsHash, "deny", 0, fmt.Sprintf("repeat-guard: max_repeats %d reached", capN)) + g.notifyDenial(req, argsHash, "max_repeats", reason) return nil, ToolCallResponse{Decision: "deny", Reason: reason, ArgsHash: argsHash, DecidedBy: "repeat-guard", Rule: "max_repeats"}, true @@ -744,61 +687,46 @@ func (g *toolGate) repeatCheck(req ToolCallRequest, rule config.ToolRule, argsHa // The gate is written to pending_approvals before the prompt leaves, so a // process that dies mid-wait is reconciled at next boot like a pipeline gate. func (g *toolGate) askHuman(parent context.Context, tk *toolTicket, req ToolCallRequest, rule config.ToolRule, mismatch string) { - timeout := g.approvalWindow() - pendingID, perr := state.BeginApproval("tool-gate", req.Tool, tk.ArgsHash, 1, tk.Created, tk.Expires) - if perr != nil { - log.Printf("[tool-gate] pending-approval write failed: %v", perr) + pendingID, err := state.BeginApproval("tool-gate", req.Tool, tk.ArgsHash, 1, tk.Created, tk.Expires) + if err != nil { + g.updateTicket(tk.ID, unavailableToolResponse(tk.ID)) + return } - - draft := fmt.Sprintf("[draftcat] Tool call awaiting approval\n\nagent: %s\ntool: %s\nrisk: %s\nargs: %s", - orDash(req.Agent), req.Tool, rule.RiskOf(), prettyArgs(req.Args)) + draft := fmt.Sprintf("[draftcat] Tool call awaiting approval\n\nagent: %s\ntool: %s\nrisk: %s\nargs: %s", orDash(req.Agent), req.Tool, rule.RiskOf(), prettyArgs(req.Args)) if mismatch != "" { draft += "\nrule: arguments outside the rule — " + mismatch } draft += "\n\nargs sha256: " + tk.ArgsHash - + timeout := tk.Expires.Sub(g.now()) ctx, cancel := withGateMetaTimeout(parent, req.RunID, req.Tool, rule.RiskOf(), timeout) defer cancel() - dec, derr := g.ch.SendForApproval(ctx, draft, nil) - - if rerr := state.ResolveApproval(pendingID); rerr != nil { - log.Printf("[tool-gate] pending-approval resolve failed: %v", rerr) + tk.mu.Lock() + if tk.decided { + tk.mu.Unlock() + _ = state.ResolveApproval(pendingID) + return } - - now := g.now() - if derr != nil || dec.Action != "approve" { - reason := "operator did not approve" - if derr != nil { - reason = "approval failed: " + derr.Error() - } else if dec.Action == "timeout" { - reason = "approval timed out — the gate denies rather than assumes yes" - } - log.Printf("[tool-gate] DENY %s (%s) agent=%q", req.Tool, reason, req.Agent) - recordToolDecision(g.cfg, req, tk.ArgsHash, "deny", dec.ApproverID, reason) - tk.resolve(ToolCallResponse{ - ActionID: tk.ID, Decision: "deny", State: "denied", Reason: reason, - ArgsHash: tk.ArgsHash, PolicyHash: tk.PolicyHash, BindingHash: tk.BindingHash, - DecidedBy: "operator", ApprovalID: tk.ID, - }, now) - if state != nil { - _ = state.DecideToolAction(tk.ID, "denied", "deny", reason, "operator", now) - } + tk.cancel = cancel + tk.mu.Unlock() + dec, derr := g.ch.SendForApproval(ctx, draft, nil) + if err := state.ResolveApproval(pendingID); err != nil { + g.updateTicket(tk.ID, unavailableToolResponse(tk.ID)) return } - - log.Printf("[tool-gate] ALLOW %s (operator %d) agent=%q", req.Tool, dec.ApproverID, req.Agent) - recordToolDecision(g.cfg, req, tk.ArgsHash, "approve", dec.ApproverID, "operator approved") - tk.resolve(ToolCallResponse{ - ActionID: tk.ID, Decision: "allow", State: "allowed", - Reason: "operator approved; consume the permit before executing", - ArgsHash: tk.ArgsHash, PolicyHash: tk.PolicyHash, BindingHash: tk.BindingHash, - DecidedBy: "operator", ApprovalID: tk.ID, - Consume: toolGatePath + "/" + tk.ID + "/consume", - ExpiresAt: tk.Expires.UTC().Format(time.RFC3339), - }, now) - if state != nil { - _ = state.DecideToolAction(tk.ID, "allowed", "allow", "operator approved", "operator", now) + if resp, decided := tk.result(); decided && resp.State != "pending" { + return } + req.approverID = dec.ApproverID + resp := ToolCallResponse{Decision: "deny", Reason: "operator did not approve", ArgsHash: tk.ArgsHash, PolicyHash: tk.PolicyHash, BindingHash: tk.BindingHash, DecidedBy: "operator"} + switch { + case derr != nil: + resp.Reason = "approval failed: " + derr.Error() + case dec.Action == "timeout": + resp.Reason = "approval timed out — the gate denies rather than assumes yes" + case dec.Action == "approve": + resp.Decision, resp.Reason = "allow", "operator approved" + } + g.finishAction(req, tk.Expires, resp) } // notifyDenial tells the operator about a refusal the gate made on its own. @@ -924,21 +852,19 @@ func responseFromToolAction(a statestore.ToolAction) ToolCallResponse { resp.Reason = "awaiting operator decision" } if a.Status == "allowed" { - if time.Now().After(a.ExpiresAt) { - resp.Decision, resp.State = "deny", "expired" - resp.Reason = "permit expired before consumption" - } else { - resp.Consume = toolGatePath + "/" + a.ActionID + "/consume" - if resp.Reason == "" { - resp.Reason = "decision allows this action; consume the permit before executing" - } + resp.Consume = toolGatePath + "/" + a.ActionID + "/consume" + if resp.Reason == "" { + resp.Reason = "decision allows this action; consume the permit before executing" } } + + resp.ExecutionStatus = "not_started" if a.Status == "consumed" { + resp.ExecutionStatus = "unreported" resp.Permit = "" resp.Reason = "permit already consumed" } - if a.Status == "expired" { + if a.Status == "expired" || a.Status == "revoked" { resp.Decision = "deny" } return resp @@ -947,78 +873,90 @@ func responseFromToolAction(a statestore.ToolAction) ToolCallResponse { // reserveAction is the idempotency gate. The first request reserves its action // id; matching retries return the existing state and drift is rejected. func (g *toolGate) reserveAction(req ToolCallRequest, argsHash, policyHash, bindingHash string, expires time.Time) (ToolCallResponse, int, bool) { - g.mu.Lock() - tk := g.tickets[req.ActionID] - g.mu.Unlock() - if tk != nil { - if tk.BindingHash != bindingHash { - return ToolCallResponse{ActionID: req.ActionID, Decision: "deny", State: "denied", - Reason: "action_id is already bound to different action data"}, http.StatusConflict, true + if state != nil { + a, created, err := state.ReserveToolAction(statestore.ToolAction{ActionID: req.ActionID, Tool: req.Tool, Agent: req.Agent, RunID: req.RunID, ArgsHash: argsHash, PolicyHash: policyHash, BindingHash: bindingHash, CreatedAt: g.now(), UpdatedAt: g.now(), ExpiresAt: expires}) + if err != nil { + if a.ActionID != "" && a.BindingHash != bindingHash { + return ToolCallResponse{ActionID: req.ActionID, Decision: "deny", State: "denied", Reason: "action_id is already bound to different action data"}, http.StatusConflict, true + } + return unavailableToolResponse(req.ActionID), http.StatusServiceUnavailable, true } - if resp, decided := tk.result(); decided { - if resp.State == "allowed" && g.now().After(tk.Expires) { - resp.Decision, resp.State, resp.Consume, resp.Permit = "deny", "expired", "", "" - resp.Reason = "permit expired before consumption" + if !created { + resp, code := g.actionResponse(req.ActionID) + if resp.State == "pending" && code == http.StatusOK { + code = http.StatusAccepted } - return resp, http.StatusOK, true + return resp, code, true } - return tk.pendingResponse(), http.StatusAccepted, true + return ToolCallResponse{}, 0, false } - if state == nil { + g.mu.Lock() + tk := g.tickets[req.ActionID] + g.mu.Unlock() + if tk == nil { return ToolCallResponse{}, 0, false } - a, created, err := state.ReserveToolAction(statestore.ToolAction{ - ActionID: req.ActionID, Tool: req.Tool, Agent: req.Agent, RunID: req.RunID, - ArgsHash: argsHash, PolicyHash: policyHash, BindingHash: bindingHash, - CreatedAt: g.now(), UpdatedAt: g.now(), ExpiresAt: expires, - }) - if err != nil { - return ToolCallResponse{ActionID: req.ActionID, Decision: "deny", State: "denied", Reason: err.Error()}, http.StatusConflict, true + if tk.BindingHash != bindingHash { + return ToolCallResponse{ActionID: req.ActionID, Decision: "deny", State: "denied", Reason: "action_id is already bound to different action data"}, http.StatusConflict, true } - if !created { - code := http.StatusOK - if a.Status == "pending" { - code = http.StatusAccepted - } - return responseFromToolAction(a), code, true + resp, code := g.actionResponse(req.ActionID) + if resp.State == "pending" && code == http.StatusOK { + code = http.StatusAccepted } - return ToolCallResponse{}, 0, false + return resp, code, true } func (g *toolGate) finishAction(req ToolCallRequest, expires time.Time, resp ToolCallResponse) ToolCallResponse { now := g.now() - resp.ActionID = req.ActionID - resp.ApprovalID = req.ActionID + resp.ActionID, resp.ApprovalID = req.ActionID, req.ActionID resp.ExpiresAt = expires.UTC().Format(time.RFC3339) resp.Poll = toolGatePath + "/" + req.ActionID - status := "denied" + resp.State = "denied" + decision := "deny" if resp.Decision == "allow" { - status = "allowed" + resp.State = "allowed" resp.Consume = toolGatePath + "/" + req.ActionID + "/consume" resp.Reason += "; consume the permit before executing" + decision = "policy_approve" + if resp.DecidedBy == "operator" { + decision = "approve" + } } - resp.State = status - tk := &toolTicket{ - ID: req.ActionID, Tool: req.Tool, Agent: req.Agent, RunID: req.RunID, - ArgsHash: resp.ArgsHash, PolicyHash: resp.PolicyHash, BindingHash: resp.BindingHash, - Created: now, Expires: expires, done: make(chan struct{}), + if !now.Before(expires) { + resp.Decision, resp.State, resp.Consume = "deny", "expired", "" + resp.Reason = "permit expired before consumption" + decision = "deny" } - tk.resolve(resp, now) - g.mu.Lock() - g.tickets[tk.ID] = tk - g.mu.Unlock() if state != nil { - _ = state.DecideToolAction(req.ActionID, status, resp.Decision, resp.Reason, resp.DecidedBy, now) + e, err := toolDecisionReceipt(req, resp.ArgsHash, decision, req.approverID, resp.Reason, now) + if err != nil { + resp = unavailableToolResponse(req.ActionID) + } else { + a, err := state.DecideToolActionWithReceipt(req.ActionID, map[string]string{"allow": "allowed", "deny": "denied"}[resp.Decision], resp.Decision, resp.Reason, resp.DecidedBy, now, e) + if err != nil { + resp = unavailableToolResponse(req.ActionID) + } else { + old := resp + resp = responseFromToolAction(a) + if a.Status == old.State { + resp.Rule = old.Rule + } + } + } + } + obs.RecordApproval("tool-gate", req.Tool, decision) + g.mu.Lock() + tk := g.tickets[req.ActionID] + if tk == nil { + tk = &toolTicket{ID: req.ActionID, Tool: req.Tool, Agent: req.Agent, RunID: req.RunID, ArgsHash: hashToolArgs(req.Args), PolicyHash: req.policyHash, BindingHash: req.bindingHash, Created: now, Expires: expires, done: make(chan struct{})} + g.tickets[tk.ID] = tk } + g.mu.Unlock() + g.updateTicket(req.ActionID, resp) return resp } -func recordToolDecision(cfg *config.Config, req ToolCallRequest, argsHash, decision string, operator int64, reason string) { - obs.RecordApproval("tool-gate", req.Tool, decision) - if state == nil { - return - } - decidedAt := time.Now() +func toolDecisionReceipt(req ToolCallRequest, argsHash, decision string, operator int64, reason string, decidedAt time.Time) (statestore.ApprovalEnvelope, error) { receiptID := "rcpt_" + strings.TrimPrefix(newToolTicketID(), "tc_") expires := req.expires if expires.IsZero() { @@ -1049,16 +987,13 @@ func recordToolDecision(cfg *config.Config, req ToolCallRequest, argsHash, decis }, nonce) } } - if err := state.RecordApprovalV2(envelope); err != nil { - log.Printf("[tool-gate] audit write failed: %v", err) + if len(secret) > 0 && (envelope.Nonce == "" || envelope.Signature == "") { + return envelope, errors.New("approval receipt signing failed") } - _ = cfg + return envelope, nil } -func recordToolConsumption(a statestore.ToolAction, at time.Time) { - if state == nil { - return - } +func toolConsumptionReceipt(a statestore.ToolAction, at time.Time) (statestore.ApprovalEnvelope, error) { e := statestore.ApprovalEnvelope{ ReceiptID: "rcpt_" + strings.TrimPrefix(newToolTicketID(), "tc_"), RunID: a.RunID, ActionID: a.ActionID, Pipeline: "tool-gate", Step: a.Tool, @@ -1079,12 +1014,16 @@ func recordToolConsumption(a statestore.ToolAction, at time.Time) { }, nonce) } } - if err := state.RecordApprovalV2(e); err != nil { - log.Printf("[tool-gate] consumption receipt write failed: %v", err) + if len(secret) > 0 && (e.Nonce == "" || e.Signature == "") { + return e, errors.New("approval receipt signing failed") } + return e, nil } func writeToolDecision(w http.ResponseWriter, code int, resp ToolCallResponse) { + if resp.unavailable { + code = http.StatusServiceUnavailable + } w.Header().Set("Content-Type", "application/json") w.WriteHeader(code) _ = json.NewEncoder(w).Encode(resp) diff --git a/tool_lifecycle.go b/tool_lifecycle.go new file mode 100644 index 0000000..c8ac309 --- /dev/null +++ b/tool_lifecycle.go @@ -0,0 +1,211 @@ +package main + +import ( + "database/sql" + "encoding/hex" + "errors" + "net/http" + "strings" + "time" + + statestore "github.com/renezander030/draftcat/internal/state" +) + +func unavailableToolResponse(id string) ToolCallResponse { + return ToolCallResponse{ActionID: id, ApprovalID: id, Decision: "deny", State: "denied", Reason: "durable approval state or audit unavailable; no execution authorized", unavailable: true} +} + +// updateTicket wakes all waiters on a terminal state and interrupts the +// operator prompt. Persisted state decides the race, not prompt completion. +func (g *toolGate) updateTicket(id string, resp ToolCallResponse) { + g.mu.Lock() + tk := g.tickets[id] + g.mu.Unlock() + if tk == nil { + return + } + tk.mu.Lock() + if (tk.resp.State == "revoked" || tk.resp.State == "consumed" || tk.resp.State == "expired") && (resp.State == "allowed" || resp.State == "pending") { + tk.mu.Unlock() + return + } + if resp.ArgsHash == "" { + resp.ArgsHash = tk.ArgsHash + } + if resp.BindingHash == "" { + resp.BindingHash = tk.BindingHash + } + if resp.PolicyHash == "" { + resp.PolicyHash = tk.PolicyHash + } + tk.resp = resp + terminal := resp.State != "pending" + if terminal && !tk.decided { + tk.decided = true + tk.decidedAt = g.now() + close(tk.done) + } + cancel := tk.cancel + tk.mu.Unlock() + if terminal && cancel != nil { + cancel() + } +} + +func (g *toolGate) actionResponse(id string) (ToolCallResponse, int) { + if state != nil { + a, err := state.ToolAction(id) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return unknownToolResponse(id), http.StatusNotFound + } + return unavailableToolResponse(id), http.StatusServiceUnavailable + } + rule, listed := g.cfg.ToolGate.Lookup(a.Tool) + a, err = state.NormalizeToolAction(id, g.policyHash(a.Tool, rule, listed), g.now()) + if err != nil { + return unavailableToolResponse(id), http.StatusServiceUnavailable + } + if a.Status == "pending" { + g.mu.Lock() + tk := g.tickets[id] + g.mu.Unlock() + if tk != nil { + if failed, _ := tk.result(); failed.unavailable { + return failed, http.StatusServiceUnavailable + } + } + } + resp := responseFromToolAction(a) + resp.ExecutionStatus = "not_started" + if a.Status == "consumed" { + resp.ExecutionStatus = "unreported" + } + o, err := state.ToolOutcome(id) + if err == nil { + resp.ExecutionStatus = o.Status + resp.ResultHash = o.ResultHash + resp.CompletedAt = o.CompletedAt.UTC().Format(time.RFC3339) + resp.ExecutionEvidence = "caller_attestation" + } else if !errors.Is(err, sql.ErrNoRows) { + return unavailableToolResponse(id), http.StatusServiceUnavailable + } + g.updateTicket(id, resp) + return resp, http.StatusOK + } + g.mu.Lock() + tk := g.tickets[id] + g.mu.Unlock() + if tk == nil { + return unknownToolResponse(id), http.StatusNotFound + } + resp, decided := tk.result() + if !decided { + resp = tk.pendingResponse() + } + if resp.State == "pending" || resp.State == "allowed" { + rule, listed := g.cfg.ToolGate.Lookup(tk.Tool) + reason := "" + if tk.PolicyHash != g.policyHash(tk.Tool, rule, listed) { + reason = "policy changed; request a new approval" + } else if !g.now().Before(tk.Expires) { + reason = "permit expired before consumption" + } + if reason != "" { + resp.Decision, resp.State, resp.Consume, resp.Permit, resp.Reason = "deny", "expired", "", "", reason + g.updateTicket(id, resp) + } + } + return resp, http.StatusOK +} + +func unknownToolResponse(id string) ToolCallResponse { + return ToolCallResponse{ActionID: id, ApprovalID: id, Decision: "deny", State: "denied", Reason: "unknown action id - the gate has no decision for it; ask again"} +} + +func (g *toolGate) handleRevoke(w http.ResponseWriter, r *http.Request, id string) { + var body struct { + BindingHash string `json:"binding_hash"` + } + raw, ok := readRequestBody(w, r, 4096) + if !ok { + return + } + if err := decodeStrictJSON(raw, &body); err != nil { + http.Error(w, "malformed json", http.StatusBadRequest) + return + } + if state == nil { + writeToolDecision(w, http.StatusServiceUnavailable, unavailableToolResponse(id)) + return + } + if body.BindingHash == "" { + http.Error(w, "binding_hash required", http.StatusBadRequest) + return + } + a, revoked, err := state.CancelToolAction(id, body.BindingHash, g.now()) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + writeToolDecision(w, http.StatusNotFound, unknownToolResponse(id)) + } else { + writeToolDecision(w, http.StatusServiceUnavailable, unavailableToolResponse(id)) + } + return + } + resp := responseFromToolAction(a) + g.updateTicket(id, resp) + if !revoked { + writeToolDecision(w, http.StatusConflict, resp) + return + } + writeToolDecision(w, http.StatusOK, resp) +} + +func validResultHash(hash string) bool { + if hash == "" { + return true + } + if !strings.HasPrefix(hash, "sha256:") || len(hash) != 71 { + return false + } + _, err := hex.DecodeString(strings.TrimPrefix(hash, "sha256:")) + return err == nil +} + +func (g *toolGate) handleComplete(w http.ResponseWriter, r *http.Request, id string) { + var body struct { + BindingHash string `json:"binding_hash"` + Status string `json:"status"` + ResultHash string `json:"result_hash"` + } + raw, ok := readRequestBody(w, r, 4096) + if !ok { + return + } + if err := decodeStrictJSON(raw, &body); err != nil { + http.Error(w, "malformed json", http.StatusBadRequest) + return + } + if body.BindingHash == "" || (body.Status != "succeeded" && body.Status != "failed") || !validResultHash(body.ResultHash) { + http.Error(w, "binding_hash and status succeeded|failed required; result_hash must be SHA-256", http.StatusBadRequest) + return + } + if state == nil { + writeToolDecision(w, http.StatusServiceUnavailable, unavailableToolResponse(id)) + return + } + _, accepted, err := state.CompleteToolAction(statestore.ToolOutcome{ActionID: id, BindingHash: body.BindingHash, Status: body.Status, ResultHash: body.ResultHash, CompletedAt: g.now()}) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + writeToolDecision(w, http.StatusNotFound, unknownToolResponse(id)) + } else { + writeToolDecision(w, http.StatusServiceUnavailable, unavailableToolResponse(id)) + } + return + } + resp, code := g.actionResponse(id) + if !accepted && code == http.StatusOK { + code = http.StatusConflict + } + writeToolDecision(w, code, resp) +} diff --git a/tool_lifecycle_test.go b/tool_lifecycle_test.go new file mode 100644 index 0000000..e1dc573 --- /dev/null +++ b/tool_lifecycle_test.go @@ -0,0 +1,291 @@ +package main + +import ( + "context" + "database/sql" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/renezander030/draftcat/internal/config" + statestore "github.com/renezander030/draftcat/internal/state" +) + +func lifecycleHTTPStore(t *testing.T) (string, *statestore.StateStore) { + t.Helper() + path := filepath.Join(t.TempDir(), "state.db") + st, err := statestore.OpenStateStore(path) + if err != nil { + t.Fatal(err) + } + old := state + state = st + t.Cleanup(func() { _ = state.Close(); state = old }) + return path, st +} +func lifecyclePost(t *testing.T, g *toolGate, id, operation, body string) (int, ToolCallResponse) { + t.Helper() + rec := httptest.NewRecorder() + g.HandleStatus(rec, httptest.NewRequest(http.MethodPost, toolGatePath+"/"+id+"/"+operation, strings.NewReader(body))) + var resp ToolCallResponse + if rec.Code != http.StatusBadRequest { + if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil { + t.Fatalf("decode %q: %v", rec.Body.String(), err) + } + } + return rec.Code, resp +} +func revokeBody(binding string) string { return fmt.Sprintf(`{"binding_hash":%q}`, binding) } +func completeBody(binding, status, hash string) string { + return fmt.Sprintf(`{"binding_hash":%q,"status":%q,"result_hash":%q}`, binding, status, hash) +} + +func TestToolLifecycleHTTPCompletionAndRestart(t *testing.T) { + path, st := lifecycleHTTPStore(t) + cfg := gateCfg(config.ToolRule{Name: "send"}) + g := newToolGate(cfg, nil) + _, allowed := postGate(t, g, `{"action_id":"result","tool":"send"}`) + hash := "sha256:" + strings.Repeat("a", 64) + body := completeBody(allowed.BindingHash, "succeeded", hash) + if code, _ := lifecyclePost(t, g, "result", "complete", body); code != 409 { + t.Fatalf("unconsumed result accepted: %d", code) + } + if code, resp := consumeGate(t, g, "result", allowed.BindingHash); code != 200 || resp.Permit != "execute" || resp.ExecutionStatus != "unreported" { + t.Fatalf("consume=%d %+v", code, resp) + } + code, completed := lifecyclePost(t, g, "result", "complete", body) + if code != 200 || completed.State != "consumed" || completed.ExecutionStatus != "succeeded" || completed.ResultHash != hash || completed.Permit != "" || completed.ExecutionEvidence != "caller_attestation" { + t.Fatalf("completion=%d %+v", code, completed) + } + if code, _ := lifecyclePost(t, g, "result", "complete", completeBody(allowed.BindingHash, "failed", hash)); code != 409 { + t.Fatalf("conflicting outcome accepted: %d", code) + } + if code, _ := lifecyclePost(t, g, "result", "complete", completeBody("wrong", "succeeded", hash)); code != 409 { + t.Fatalf("changed binding accepted: %d", code) + } + if code, _ := lifecyclePost(t, g, "result", "complete", `{"binding_hash":"x","status":"succeeded","result":"private payload"}`); code != 400 { + t.Fatalf("payload accepted: %d", code) + } + if err := st.Close(); err != nil { + t.Fatal(err) + } + var err error + state, err = statestore.OpenStateStore(path) + if err != nil { + t.Fatal(err) + } + g = newToolGate(cfg, nil) + if code, resp := getGate(t, g, "result", ""); code != 200 || resp.ExecutionStatus != "succeeded" || resp.ResultHash != hash || resp.CompletedAt != completed.CompletedAt || resp.Permit != "" || resp.Consume != "" { + t.Fatalf("restored=%d %+v", code, resp) + } + if code, resp := lifecyclePost(t, g, "result", "complete", body); code != 200 || resp.CompletedAt != completed.CompletedAt { + t.Fatalf("retry=%d %+v", code, resp) + } + if code, resp := consumeGate(t, g, "result", allowed.BindingHash); code != 409 || resp.Permit != "" { + t.Fatalf("result reopened permit: %d %+v", code, resp) + } + if code, resp := lifecyclePost(t, g, "result", "revoke", revokeBody(allowed.BindingHash)); code != 409 || resp.State != "consumed" { + t.Fatalf("consumed revoke=%d %+v", code, resp) + } +} + +func TestToolLifecycleHTTPPollNormalizesLiveAndRecovered(t *testing.T) { + lifecycleHTTPStore(t) + for _, change := range []string{"expiry", "policy"} { + for _, recovered := range []bool{false, true} { + t.Run(fmt.Sprintf("%s/recovered=%v", change, recovered), func(t *testing.T) { + cfg := gateCfg(config.ToolRule{Name: "send"}) + g := newToolGate(cfg, nil) + _, allowed := postGate(t, g, fmt.Sprintf(`{"action_id":%q,"tool":"send"}`, fmt.Sprintf("%s-%v", change, recovered))) + expires, err := time.Parse(time.RFC3339, allowed.ExpiresAt) + if err != nil { + t.Fatal(err) + } + if change == "policy" { + cfg.ToolGate.Tools[0].RequireApproval = true + } + if recovered { + g = newToolGate(cfg, nil) + } + if change == "expiry" { + g.now = func() time.Time { return expires } + } + code, resp := getGate(t, g, allowed.ActionID, "") + if code != 200 || resp.Decision != "deny" || resp.State != "expired" || resp.Consume != "" || resp.Permit != "" { + t.Fatalf("invalid poll=%d %+v", code, resp) + } + if code, resp := consumeGate(t, g, allowed.ActionID, allowed.BindingHash); code != 409 || resp.Permit != "" { + t.Fatalf("invalid consume=%d %+v", code, resp) + } + }) + } + } +} + +type revokedPromptChannel struct { + stubApprovalChannel + began chan struct{} + cancelled chan struct{} + release chan struct{} + returned chan struct{} +} + +func (ch *revokedPromptChannel) SendForApproval(ctx context.Context, _ string, _ []int64) (OperatorDecision, error) { + close(ch.began) + <-ctx.Done() + close(ch.cancelled) + <-ch.release + close(ch.returned) + return OperatorDecision{Action: "approve", ApproverID: 42}, nil +} + +func TestToolLifecycleHTTPRevokeWakesPollAndLateApprovalCannotResurrect(t *testing.T) { + _, st := lifecycleHTTPStore(t) + ch := &revokedPromptChannel{began: make(chan struct{}), cancelled: make(chan struct{}), release: make(chan struct{}), returned: make(chan struct{})} + cfg := gateCfg(config.ToolRule{Name: "send", RequireApproval: true}) + cfg.Timeouts.OperatorApproval = "30s" + g := newToolGate(cfg, ch) + _, pending := postGate(t, g, `{"action_id":"cancel","tool":"send","mode":"async"}`) + select { + case <-ch.began: + case <-time.After(time.Second): + t.Fatal("prompt never started") + } + poll := make(chan ToolCallResponse, 1) + go func() { + rec := httptest.NewRecorder() + g.HandleStatus(rec, httptest.NewRequest(http.MethodGet, toolGatePath+"/cancel?wait=20s", nil)) + var resp ToolCallResponse + _ = json.Unmarshal(rec.Body.Bytes(), &resp) + poll <- resp + }() + if code, resp := lifecyclePost(t, g, "cancel", "revoke", revokeBody(pending.BindingHash)); code != 200 || resp.State != "revoked" || resp.Decision != "deny" { + t.Fatalf("revoke=%d %+v", code, resp) + } + select { + case resp := <-poll: + if resp.State != "revoked" { + t.Fatalf("poll=%+v", resp) + } + case <-time.After(time.Second): + t.Fatal("revocation did not wake poll") + } + select { + case <-ch.cancelled: + case <-time.After(time.Second): + t.Fatal("prompt was not canceled") + } + close(ch.release) + <-ch.returned + deadline := time.Now().Add(time.Second) + for time.Now().Before(deadline) { + open, err := st.OpenApprovals() + if err != nil { + t.Fatal(err) + } + if len(open) == 0 { + break + } + time.Sleep(time.Millisecond) + } + if code, resp := getGate(t, g, "cancel", ""); code != 200 || resp.State != "revoked" || resp.Decision != "deny" { + t.Fatalf("late approval resurrected: %d %+v", code, resp) + } + if code, _ := lifecyclePost(t, g, "cancel", "revoke", revokeBody(pending.BindingHash)); code != 200 { + t.Fatalf("revocation retry=%d", code) + } + if code, _ := lifecyclePost(t, g, "cancel", "revoke", revokeBody("different")); code != 409 { + t.Fatalf("changed binding accepted=%d", code) + } + if code, resp := consumeGate(t, g, "cancel", pending.BindingHash); code != 409 || resp.Permit != "" { + t.Fatalf("revoked execution=%d %+v", code, resp) + } +} + +func TestToolLifecycleHTTPAuditAndStoreFailuresDeny(t *testing.T) { + path, st := lifecycleHTTPStore(t) + g := newToolGate(gateCfg(config.ToolRule{Name: "send"}), nil) + _, allowed := postGate(t, g, `{"action_id":"consume-audit","tool":"send"}`) + db, err := sql.Open("sqlite", path) + if err != nil { + t.Fatal(err) + } + defer db.Close() + if _, err := db.ExecContext(context.Background(), `CREATE TRIGGER reject_tool_audit BEFORE INSERT ON action_approvals BEGIN SELECT RAISE(FAIL,'receipt unavailable'); END`); err != nil { + t.Fatal(err) + } + if code, resp := postGate(t, g, `{"action_id":"decision-audit","tool":"send"}`); code != 503 || resp.Decision != "deny" || resp.Permit != "" { + t.Fatalf("audit decision=%d %+v", code, resp) + } + if got, err := st.ToolAction("decision-audit"); err != nil || got.Status != "pending" { + t.Fatalf("audit decision persisted=%+v %v", got, err) + } + if code, resp := getGate(t, g, "decision-audit", ""); code != 503 || resp.Decision != "deny" { + t.Fatalf("failed audit was reported pending: %d %+v", code, resp) + } + if code, resp := consumeGate(t, g, allowed.ActionID, allowed.BindingHash); code != 503 || resp.Permit != "" { + t.Fatalf("audit consume=%d %+v", code, resp) + } + if got, err := st.ToolAction(allowed.ActionID); err != nil || got.Status != "allowed" { + t.Fatalf("audit consumed permit=%+v %v", got, err) + } + if err := st.Close(); err != nil { + t.Fatal(err) + } + if code, resp := getGate(t, g, allowed.ActionID, ""); code != 503 || resp.Decision != "deny" { + t.Fatalf("cached allow escaped store failure=%d %+v", code, resp) + } + if code, resp := consumeGate(t, g, allowed.ActionID, allowed.BindingHash); code != 503 || resp.Permit != "" { + t.Fatalf("closed store execution=%d %+v", code, resp) + } +} + +func TestToolLifecycleHTTPRequiresDurableStoreAndExistingAuth(t *testing.T) { + old := state + state = nil + t.Cleanup(func() { state = old }) + cfg := gateCfg(config.ToolRule{Name: "send"}) + g := newToolGate(cfg, nil) + for _, operation := range []string{"consume", "revoke", "complete"} { + body := revokeBody("binding") + if operation == "complete" { + body = completeBody("binding", "succeeded", "") + } + if code, resp := lifecyclePost(t, g, "missing", operation, body); code != 503 || resp.Permit != "" { + t.Fatalf("%s without store=%d %+v", operation, code, resp) + } + } + cfg.Webhook.SetSecret("secret") + h := newWebhookHandler(cfg, newScheduler(nil), &BudgetTracker{}, &TGBot{}, nil) + for _, operation := range []string{"consume", "revoke", "complete"} { + rec := httptest.NewRecorder() + h.ServeHTTP(rec, httptest.NewRequest(http.MethodPost, toolGatePath+"/missing/"+operation, strings.NewReader(revokeBody("binding")))) + if rec.Code != http.StatusUnauthorized { + t.Fatalf("%s bypassed authentication: %d", operation, rec.Code) + } + } +} + +func TestToolLifecycleHTTPPendingWriteFailureDoesNotPrompt(t *testing.T) { + path, _ := lifecycleHTTPStore(t) + db, err := sql.Open("sqlite", path) + if err != nil { + t.Fatal(err) + } + defer db.Close() + if _, err := db.ExecContext(context.Background(), `CREATE TRIGGER reject_tool_pending BEFORE INSERT ON pending_approvals BEGIN SELECT RAISE(FAIL,'pending unavailable'); END`); err != nil { + t.Fatal(err) + } + ch := &countingChannel{stubApprovalChannel: stubApprovalChannel{action: "approve"}} + g := newToolGate(gateCfg(config.ToolRule{Name: "send", RequireApproval: true}), ch) + _, pending := postGate(t, g, `{"action_id":"pending-write","tool":"send","mode":"async"}`) + code, resp := getGate(t, g, pending.ActionID, "?wait=1s") + if code != 503 || resp.Decision != "deny" || ch.askCount() != 0 { + t.Fatalf("pending failure=%d %+v prompts=%d", code, resp, ch.askCount()) + } +} diff --git a/version.go b/version.go index 35739ef..236eec5 100644 --- a/version.go +++ b/version.go @@ -1,4 +1,4 @@ package main // version is overridden by the native release build. -var version = "0.8.0" +var version = "0.9.0"