diff --git a/internal/tool/batch_exec.go b/internal/tool/batch_exec.go index 743d09d7..d3c97d8c 100644 --- a/internal/tool/batch_exec.go +++ b/internal/tool/batch_exec.go @@ -8,6 +8,7 @@ import ( "io" "net/http" "os" + "strconv" "strings" "time" ) @@ -31,8 +32,8 @@ func (BatchExecTool) Parameters() map[string]interface{} { "properties": map[string]interface{}{ "action": map[string]interface{}{ "type": "string", - "enum": []string{"submit", "poll"}, - "description": "submit: send prompts; poll: check batch status.", + "enum": []string{"submit", "poll", "wait"}, + "description": "submit: send prompts; poll: single status check; wait: poll until the batch reaches a terminal state (with backoff + Retry-After honoring).", }, "prompts": map[string]interface{}{ "type": "array", @@ -45,12 +46,20 @@ func (BatchExecTool) Parameters() map[string]interface{} { }, "batch_id": map[string]interface{}{ "type": "string", - "description": "Batch ID to poll (action=poll).", + "description": "Batch ID to poll/wait on.", }, "max_tokens": map[string]interface{}{ "type": "integer", "description": "Max output tokens per request (default 4096).", }, + "timeout_seconds": map[string]interface{}{ + "type": "integer", + "description": "Max seconds to wait (default 600).", + }, + "poll_interval_seconds": map[string]interface{}{ + "type": "integer", + "description": "Initial poll interval in seconds (default 2).", + }, }, "required": []string{"action"}, } @@ -58,12 +67,25 @@ func (BatchExecTool) Parameters() map[string]interface{} { var batchHTTP = &http.Client{Timeout: 5 * time.Minute} +// batchExecParams are the parsed parameters for all BatchExec actions. +type batchExecParams struct { + Action string `json:"action"` + Prompts []string `json:"prompts"` + Model string `json:"model"` + BatchID string `json:"batch_id"` + MaxTokens int `json:"max_tokens"` + TimeoutSeconds int `json:"timeout_seconds"` + PollIntervalSec int `json:"poll_interval_seconds"` +} + // batchAPIKey reads the key from env. func batchAPIKey() string { return os.Getenv("ANTHROPIC_API_KEY") } func batchDefaultModel() string { return "claude-sonnet-4-20250514" } -const batchBaseURL = "https://api.anthropic.com" +// batchBaseURL is the Anthropic API host. A var (not const) so tests can +// redirect the client to an httptest server. +var batchBaseURL = "https://api.anthropic.com" func batchHeaders(req *http.Request, apiKey string) *http.Request { req.Header.Set("Content-Type", "application/json") @@ -80,13 +102,7 @@ type batchPollResult struct { } func (BatchExecTool) Execute(ctx context.Context, input json.RawMessage) (string, error) { - var p struct { - Action string `json:"action"` - Prompts []string `json:"prompts"` - Model string `json:"model"` - BatchID string `json:"batch_id"` - MaxTokens int `json:"max_tokens"` - } + var p batchExecParams if err := json.Unmarshal(input, &p); err != nil { return "", fmt.Errorf("invalid input: %w", err) } @@ -104,19 +120,17 @@ func (BatchExecTool) Execute(ctx context.Context, input json.RawMessage) (string return "", fmt.Errorf("batch_id is required for poll") } return batchPoll(ctx, apiKey, p.BatchID) + case "wait": + if p.BatchID == "" { + return "", fmt.Errorf("batch_id is required for wait") + } + return batchWait(ctx, apiKey, p.BatchID, p.TimeoutSeconds, p.PollIntervalSec) default: - return "", fmt.Errorf("unsupported action %q (use submit or poll)", p.Action) + return "", fmt.Errorf("unsupported action %q (use submit, poll, or wait)", p.Action) } } -func batchSubmit(ctx context.Context, apiKey string, p struct { - Action string `json:"action"` - Prompts []string `json:"prompts"` - Model string `json:"model"` - BatchID string `json:"batch_id"` - MaxTokens int `json:"max_tokens"` -}, -) (string, error) { +func batchSubmit(ctx context.Context, apiKey string, p batchExecParams) (string, error) { if len(p.Prompts) == 0 { return "", fmt.Errorf("at least one prompt is required") } @@ -201,3 +215,91 @@ func batchPoll(ctx context.Context, apiKey, batchID string) (string, error) { out, _ := json.MarshalIndent(result, "", " ") return string(out), nil } + +// batchTerminalStates are the statuses that mean the batch is done. +var batchTerminalStates = map[string]bool{ + "ended": true, "completed": true, "failed": true, + "expired": true, "canceled": true, "cancelled": true, +} + +func isBatchTerminal(s string) bool { return batchTerminalStates[strings.ToLower(s)] } + +// batchWait polls until the batch reaches a terminal state, with exponential +// backoff + jitter capped at 30s, honoring Retry-After on 429/5xx, bounded by +// timeoutSeconds (default 600). Mirrors eyrie's WaitUntilDone but inlined to +// stay boundary-compliant (no eyrie/client import). +func batchWait(ctx context.Context, apiKey, batchID string, timeoutSeconds, pollIntervalSec int) (string, error) { + if timeoutSeconds <= 0 { + timeoutSeconds = 600 + } + initial := time.Duration(pollIntervalSec) * time.Second + if initial <= 0 { + initial = 2 * time.Second + } + deadline := time.Now().Add(time.Duration(timeoutSeconds) * time.Second) + attempt := 0 + for { + status, retryAfter, err := batchStatus(ctx, apiKey, batchID) + if err != nil { + return "", err + } + if isBatchTerminal(status) { + out, _ := json.MarshalIndent(batchPollResult{BatchID: batchID, Status: status}, "", " ") + return string(out), nil + } + if time.Now().After(deadline) { + return "", fmt.Errorf("batch %s not terminal within %ds (last status %s)", batchID, timeoutSeconds, status) + } + delay := batchBackoffDelay(attempt, initial, retryAfter) + select { + case <-ctx.Done(): + return "", ctx.Err() + case <-time.After(delay): + } + attempt++ + } +} + +// batchStatus performs a single status fetch, returning (status, retryAfter, err). +func batchStatus(ctx context.Context, apiKey, batchID string) (string, string, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, + batchBaseURL+"/v1/messages/batches/"+batchID, nil) // #nosec G107 -- fixed API host + validated ID + if err != nil { + return "", "", err + } + batchHeaders(req, apiKey) + resp, err := batchHTTP.Do(req) + if err != nil { + return "", "", fmt.Errorf("batch status: %w", err) + } + defer func() { _ = resp.Body.Close() }() + raw, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<20)) + if resp.StatusCode == http.StatusTooManyRequests || resp.StatusCode >= 500 { + return "in_progress", resp.Header.Get("Retry-After"), nil // transient: keep polling + } + if resp.StatusCode != http.StatusOK { + return "", "", fmt.Errorf("batch API %d: %s", resp.StatusCode, strings.TrimSpace(string(raw))) + } + var result batchPollResult + if err := json.Unmarshal(raw, &result); err != nil { + return "", "", fmt.Errorf("batch status parse: %w", err) + } + return result.Status, "", nil +} + +// batchBackoffDelay computes attempt-th exponential delay with jitter, capped +// at 30s, overridden by Retry-After seconds when larger. +func batchBackoffDelay(attempt int, initial time.Duration, retryAfter string) time.Duration { + d := initial << uint(minInt(attempt, 6)) + if d > 30*time.Second || d <= 0 { + d = 30 * time.Second + } + // ±20% jitter. + d = time.Duration(float64(d) * (0.8 + float64(int(time.Now().UnixNano())%50)/100.0*0.4)) + if ra, err := strconv.Atoi(strings.TrimSpace(retryAfter)); err == nil && ra > 0 { + if rd := time.Duration(ra) * time.Second; rd > d { + return rd + } + } + return d +} diff --git a/internal/tool/batch_exec_test.go b/internal/tool/batch_exec_test.go index db736235..9814f969 100644 --- a/internal/tool/batch_exec_test.go +++ b/internal/tool/batch_exec_test.go @@ -3,8 +3,12 @@ package tool import ( "context" "encoding/json" + "fmt" + "net/http" + "net/http/httptest" "strings" "testing" + "time" ) func TestBatchExecRequiresAPIKey(t *testing.T) { @@ -46,3 +50,101 @@ func TestBatchExecInvalidAction(t *testing.T) { t.Fatal("expected error for unsupported action") } } + +func TestBatchWaitPollsUntilTerminal(t *testing.T) { + var polls int + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + polls++ + status := "in_progress" + if polls >= 3 { + status = "ended" + } + fmt.Fprintf(w, `{"id":"b1","status":%q}`, status) + })) + defer srv.Close() + oldBase := batchBaseURL + batchBaseURL = srv.URL + defer func() { batchBaseURL = oldBase }() + + out, err := batchWait(context.Background(), "k", "b1", 10, 1) + if err != nil { + t.Fatalf("batchWait: %v", err) + } + if !strings.Contains(out, `"status": "ended"`) { + t.Fatalf("out = %s", out) + } + if polls < 3 { + t.Fatalf("expected >=3 polls, got %d", polls) + } +} + +func TestBatchWaitHonorsRetryAfterOn429(t *testing.T) { + var saw429 bool + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if !saw429 { + saw429 = true + w.Header().Set("Retry-After", "1") + w.WriteHeader(http.StatusTooManyRequests) + return + } + fmt.Fprint(w, `{"id":"b1","status":"ended"}`) + })) + defer srv.Close() + oldBase := batchBaseURL + batchBaseURL = srv.URL + defer func() { batchBaseURL = oldBase }() + + start := time.Now() + out, err := batchWait(context.Background(), "k", "b1", 10, 1) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(out, `"status": "ended"`) { + t.Fatal("wrong status") + } + if elapsed := time.Since(start); elapsed < 900*time.Millisecond { + t.Fatalf("Retry-After not honored; elapsed=%v", elapsed) + } +} + +func TestBatchWaitTimeout(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + fmt.Fprint(w, `{"id":"b1","status":"in_progress"}`) + })) + defer srv.Close() + oldBase := batchBaseURL + batchBaseURL = srv.URL + defer func() { batchBaseURL = oldBase }() + + if _, err := batchWait(context.Background(), "k", "b1", 1, 1); err == nil || !strings.Contains(err.Error(), "not terminal") { + t.Fatalf("err = %v", err) + } +} + +func TestBatchWaitRequiresID(t *testing.T) { + t.Setenv("ANTHROPIC_API_KEY", "k") + _, err := (BatchExecTool{}).Execute(context.Background(), json.RawMessage(`{"action":"wait"}`)) + if err == nil || !strings.Contains(err.Error(), "batch_id is required for wait") { + t.Fatalf("err = %v", err) + } +} + +func TestIsBatchTerminal(t *testing.T) { + for _, ok := range []string{"ended", "completed", "failed", "expired", "canceled", "cancelled"} { + if !isBatchTerminal(ok) { + t.Fatalf("%q should be terminal", ok) + } + } + if isBatchTerminal("in_progress") { + t.Fatal("in_progress should not be terminal") + } +} + +func TestBatchBackoffDelayRetryAfter(t *testing.T) { + if d := batchBackoffDelay(0, time.Second, "60"); d < 59*time.Second { + t.Fatalf("Retry-After not applied: %v", d) + } + if d := batchBackoffDelay(30, time.Second, ""); d > 31*time.Second { + t.Fatalf("cap exceeded: %v", d) + } +}