diff --git a/internal/engine/peer.go b/internal/engine/peer.go new file mode 100644 index 0000000..4f71aa9 --- /dev/null +++ b/internal/engine/peer.go @@ -0,0 +1,106 @@ +// Package engine — peer.go implements the peer registry for multi-node +// scheduling (v0.2 P02). A peer is a remote orca node reachable over +// mTLS. The registry is in-memory plus optionally SQLite-persisted; +// for P02 the in-memory map is the source of truth and persistence +// is best-effort. +package engine + +import ( + "context" + "fmt" + "sort" + "sync" + "time" + + "git.cloudinit.dev/coreci/orca/internal/store" +) + +// Peer is a remote orca node reachable over mTLS. +type Peer struct { + NodeID string + Address string // host:port (the peer's daemon listener) + ServerName string // expected SAN on the peer's cert + CAPath string // path to the CA cert this peer validates against + LastSeen time.Time + Capacity *store.NodeCapacity +} + +// PeerRegistry tracks known peers. Methods are safe for concurrent +// use; the underlying map is guarded by a sync.RWMutex. +type PeerRegistry struct { + mu sync.RWMutex + peers map[string]*Peer + // optional persistence (not required for P02; can be added later) + persist PeerPersister +} + +// PeerPersister is an optional callback for persisting peer records. +// P02 doesn't use it; it's here for the P03 audit log integration. +type PeerPersister interface { + SavePeer(ctx context.Context, p *Peer) error +} + +// NewPeerRegistry returns an empty registry. +func NewPeerRegistry() *PeerRegistry { + return &PeerRegistry{peers: make(map[string]*Peer)} +} + +// Add inserts or updates a peer record. +func (r *PeerRegistry) Add(p *Peer) error { + if p == nil { + return fmt.Errorf("PeerRegistry.Add: nil peer") + } + if p.NodeID == "" { + return fmt.Errorf("PeerRegistry.Add: NodeID is required") + } + r.mu.Lock() + r.peers[p.NodeID] = p + r.mu.Unlock() + return nil +} + +// Remove deletes a peer by ID. Returns true if a peer was removed. +func (r *PeerRegistry) Remove(nodeID string) bool { + r.mu.Lock() + defer r.mu.Unlock() + _, ok := r.peers[nodeID] + if ok { + delete(r.peers, nodeID) + } + return ok +} + +// Get returns the peer with the given ID, or nil. +func (r *PeerRegistry) Get(nodeID string) *Peer { + r.mu.RLock() + defer r.mu.RUnlock() + return r.peers[nodeID] +} + +// All returns a snapshot of all peers, sorted by NodeID for determinism. +func (r *PeerRegistry) All(_ context.Context) ([]*Peer, error) { + r.mu.RLock() + out := make([]*Peer, 0, len(r.peers)) + for _, p := range r.peers { + out = append(out, p) + } + r.mu.RUnlock() + sort.Slice(out, func(i, j int) bool { return out[i].NodeID < out[j].NodeID }) + return out, nil +} + +// Len returns the number of registered peers. +func (r *PeerRegistry) Len() int { + r.mu.RLock() + defer r.mu.RUnlock() + return len(r.peers) +} + +// UpdateLastSeen bumps the LastSeen timestamp on a peer. +func (r *PeerRegistry) UpdateLastSeen(nodeID string) { + r.mu.Lock() + if p, ok := r.peers[nodeID]; ok { + p.LastSeen = time.Now().UTC() + } + r.mu.Unlock() +} diff --git a/internal/engine/scheduler.go b/internal/engine/scheduler.go new file mode 100644 index 0000000..a0e6dcd --- /dev/null +++ b/internal/engine/scheduler.go @@ -0,0 +1,117 @@ +// Package engine — scheduler.go implements best-fit bin-packing for +// the multi-node scheduler (v0.2 P02, REQ-028). The scheduler +// receives a JobSpec, looks at the local NodeCapacity, and either +// runs locally or falls through to a remote peer via the dispatcher. +// +// The bin-pack scoring is intentionally simple: pick the node with +// the most free capacity (cpu_millicores + memory_mib weighted 1:1 +// after normalization). This is deterministic and easy to test. +package engine + +import ( + "context" + "fmt" + "sort" + + "git.cloudinit.dev/coreci/orca/internal/model" + "git.cloudinit.dev/coreci/orca/internal/store" +) + +// JobSpec is a minimal projection of the spec needed for scheduling +// decisions. The full spec parsing is in internal/jobspec; this is +// just enough to ask "does this fit?" and "where should it go?". +type JobSpec struct { + CPUMillicores int64 + MemoryMiB int64 + DiskMiB int64 +} + +// Fits reports whether the local node has enough free capacity to +// run the spec. Capacity accounting is conservative: a job is allowed +// to run only if cpu + memory + disk are all >= the spec. +func (s JobSpec) Fits(c *store.NodeCapacity) bool { + if c == nil { + return false + } + return c.CPUMillicores >= s.CPUMillicores && + c.MemoryMiB >= s.MemoryMiB && + c.DiskMiB >= s.DiskMiB +} + +// Score returns a sortable score for bin-packing; higher = more free +// capacity. Weighted roughly toward CPU (which is usually the +// constraint) but normalized so the test isn't fragile. +func (s JobSpec) Score(c *store.NodeCapacity) int64 { + if c == nil { + return -1 + } + // Use 1:1 weighting in normalized units (millicores vs MiB) to + // keep the score monotonic. This isn't physically meaningful + // (mixing units) but it gives a stable ordering for tests. + freeCPU := c.CPUMillicores - s.CPUMillicores + freeMem := c.MemoryMiB - s.MemoryMiB + if freeCPU < 0 || freeMem < 0 { + return -1 + } + return freeCPU + freeMem +} + +// PickNode selects the best-fit node from a slice of capacities. +// Returns the chosen *store.NodeCapacity and its index, or an error +// if none can fit. Ties are broken by NodeID (lexicographic) for +// determinism. +func PickNode(spec JobSpec, capacities []*store.NodeCapacity) (*store.NodeCapacity, int, error) { + if len(capacities) == 0 { + return nil, -1, fmt.Errorf("PickNode: no nodes available") + } + type scored struct { + c *store.NodeCapacity + idx int + score int64 + } + var fits []scored + for i, c := range capacities { + if !spec.Fits(c) { + continue + } + fits = append(fits, scored{c: c, idx: i, score: spec.Score(c)}) + } + if len(fits) == 0 { + return nil, -1, fmt.Errorf("PickNode: no node can fit the spec (cpu=%d mem=%d disk=%d)", + spec.CPUMillicores, spec.MemoryMiB, spec.DiskMiB) + } + sort.SliceStable(fits, func(i, j int) bool { + if fits[i].score != fits[j].score { + return fits[i].score > fits[j].score + } + return fits[i].c.NodeID < fits[j].c.NodeID + }) + return fits[0].c, fits[0].idx, nil +} + +// LocalNode is a minimal abstraction of the local node for the +// scheduler. The concrete implementation reads from the +// store.CapacityRepo. +type LocalNode interface { + Capacity(ctx context.Context) (*store.NodeCapacity, error) +} + +// memLocalNode returns capacity from a fixed *store.NodeCapacity. +// Useful for tests; production code wraps CapacityRepo. +type memLocalNode struct{ c *store.NodeCapacity } + +// MemLocalNode returns a LocalNode backed by a fixed capacity. Test-only. +func MemLocalNode(c *store.NodeCapacity) LocalNode { + return &memLocalNode{c: c} +} + +func (m *memLocalNode) Capacity(_ context.Context) (*store.NodeCapacity, error) { + if m.c == nil { + return nil, store.ErrNotFound + } + return m.c, nil +} + +// ensure model import compiles even if unused above (placeholder for +// future scheduler fields that take *model.Node). +var _ = model.NodeStateReady diff --git a/internal/engine/scheduler_test.go b/internal/engine/scheduler_test.go new file mode 100644 index 0000000..ce51f28 --- /dev/null +++ b/internal/engine/scheduler_test.go @@ -0,0 +1,66 @@ +package engine + +import ( + "testing" + + "git.cloudinit.dev/coreci/orca/internal/store" +) + +func TestPickNodeBestFit(t *testing.T) { + caps := []*store.NodeCapacity{ + {NodeID: "node-b", CPUMillicores: 1000, MemoryMiB: 1024, DiskMiB: 1024}, + {NodeID: "node-a", CPUMillicores: 4000, MemoryMiB: 4096, DiskMiB: 4096}, + {NodeID: "node-c", CPUMillicores: 500, MemoryMiB: 512, DiskMiB: 512}, + } + spec := JobSpec{CPUMillicores: 1000, MemoryMiB: 1024, DiskMiB: 1024} + got, idx, err := PickNode(spec, caps) + if err != nil { + t.Fatalf("PickNode: %v", err) + } + if got.NodeID != "node-a" { + t.Errorf("PickNode: got %s, want node-a (most free capacity)", got.NodeID) + } + if idx != 1 { + t.Errorf("PickNode: got idx %d, want 1", idx) + } +} + +func TestPickNodeNoFit(t *testing.T) { + caps := []*store.NodeCapacity{ + {NodeID: "node-a", CPUMillicores: 100, MemoryMiB: 100, DiskMiB: 100}, + } + spec := JobSpec{CPUMillicores: 1000, MemoryMiB: 1024, DiskMiB: 1024} + _, _, err := PickNode(spec, caps) + if err == nil { + t.Fatal("expected PickNode to fail when no node can fit") + } +} + +func TestPickNodeTieDeterministic(t *testing.T) { + // Two nodes with identical free capacity. Tie broken by NodeID + // (lexicographic) for determinism. + caps := []*store.NodeCapacity{ + {NodeID: "node-z", CPUMillicores: 4000, MemoryMiB: 4096, DiskMiB: 4096}, + {NodeID: "node-a", CPUMillicores: 4000, MemoryMiB: 4096, DiskMiB: 4096}, + } + spec := JobSpec{CPUMillicores: 1000, MemoryMiB: 1024, DiskMiB: 1024} + got, _, err := PickNode(spec, caps) + if err != nil { + t.Fatalf("PickNode: %v", err) + } + if got.NodeID != "node-a" { + t.Errorf("PickNode tie-break: got %s, want node-a (lexicographic)", got.NodeID) + } +} + +func TestJobSpecFits(t *testing.T) { + spec := JobSpec{CPUMillicores: 1000, MemoryMiB: 1024, DiskMiB: 1024} + c := &store.NodeCapacity{CPUMillicores: 2000, MemoryMiB: 2048, DiskMiB: 2048} + if !spec.Fits(c) { + t.Error("Fits: should fit") + } + c.CPUMillicores = 500 + if spec.Fits(c) { + t.Error("Fits: should not fit (CPU too low)") + } +} diff --git a/internal/store/capacity_repo.go b/internal/store/capacity_repo.go new file mode 100644 index 0000000..dd7210b --- /dev/null +++ b/internal/store/capacity_repo.go @@ -0,0 +1,124 @@ +// Package store — capacity_repo.go implements persistence for NodeCapacity +// declarations (v0.2 P02). Capacity is declared per node via +// `orca node capacity --set` (or from `~/.orca/node.hcl` at join time). +// The dispatcher reads capacity rows to bin-pack jobs across nodes. +package store + +import ( + "context" + "database/sql" + "errors" + "fmt" + "time" +) + +// NodeCapacity is the per-node resource declaration consumed by the +// scheduler. Units: +// - CPUMillicores: 1000 = 1 vCPU +// - MemoryMiB: mebibytes of RAM +// - DiskMiB: mebibytes of scratch disk +type NodeCapacity struct { + NodeID string + CPUMillicores int64 + MemoryMiB int64 + DiskMiB int64 + UpdatedAt time.Time +} + +// CapacityRepo is the persistence layer for NodeCapacity rows. +type CapacityRepo struct { + db *sql.DB +} + +// NewCapacityRepo returns a CapacityRepo backed by the given DB. +func NewCapacityRepo(db *sql.DB) *CapacityRepo { + return &CapacityRepo{db: db} +} + +// Upsert writes the capacity row for nodeID, replacing any prior row. +// The UpdatedAt column is set to time.Now().UTC() unless the caller +// supplied a non-zero value. +func (r *CapacityRepo) Upsert(ctx context.Context, c *NodeCapacity) error { + if c == nil { + return errors.New("CapacityRepo.Upsert: nil capacity") + } + if c.NodeID == "" { + return errors.New("CapacityRepo.Upsert: NodeID is required") + } + if c.UpdatedAt.IsZero() { + c.UpdatedAt = time.Now().UTC() + } + _, err := r.db.ExecContext(ctx, ` + INSERT INTO node_capacity (node_id, cpu_millicores, memory_mib, disk_mib, updated_at) + VALUES (?, ?, ?, ?, ?) + ON CONFLICT(node_id) DO UPDATE SET + cpu_millicores = excluded.cpu_millicores, + memory_mib = excluded.memory_mib, + disk_mib = excluded.disk_mib, + updated_at = excluded.updated_at + `, c.NodeID, c.CPUMillicores, c.MemoryMiB, c.DiskMiB, c.UpdatedAt) + if err != nil { + return fmt.Errorf("CapacityRepo.Upsert: %w", err) + } + return nil +} + +// Get returns the capacity for nodeID or ErrNotFound. +func (r *CapacityRepo) Get(ctx context.Context, nodeID string) (*NodeCapacity, error) { + if nodeID == "" { + return nil, errors.New("CapacityRepo.Get: nodeID is required") + } + row := r.db.QueryRowContext(ctx, ` + SELECT node_id, cpu_millicores, memory_mib, disk_mib, updated_at + FROM node_capacity WHERE node_id = ? + `, nodeID) + var c NodeCapacity + if err := row.Scan(&c.NodeID, &c.CPUMillicores, &c.MemoryMiB, &c.DiskMiB, &c.UpdatedAt); err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, ErrNotFound + } + return nil, fmt.Errorf("CapacityRepo.Get: %w", err) + } + return &c, nil +} + +// List returns all capacity rows ordered by node_id. +func (r *CapacityRepo) List(ctx context.Context) ([]*NodeCapacity, error) { + rows, err := r.db.QueryContext(ctx, ` + SELECT node_id, cpu_millicores, memory_mib, disk_mib, updated_at + FROM node_capacity ORDER BY node_id + `) + if err != nil { + return nil, fmt.Errorf("CapacityRepo.List: %w", err) + } + defer rows.Close() + var out []*NodeCapacity + for rows.Next() { + var c NodeCapacity + if err := rows.Scan(&c.NodeID, &c.CPUMillicores, &c.MemoryMiB, &c.DiskMiB, &c.UpdatedAt); err != nil { + return nil, fmt.Errorf("CapacityRepo.List: scan: %w", err) + } + out = append(out, &c) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("CapacityRepo.List: rows: %w", err) + } + return out, nil +} + +// Delete removes the capacity row for nodeID. Returns ErrNotFound if +// the row doesn't exist. +func (r *CapacityRepo) Delete(ctx context.Context, nodeID string) error { + res, err := r.db.ExecContext(ctx, `DELETE FROM node_capacity WHERE node_id = ?`, nodeID) + if err != nil { + return fmt.Errorf("CapacityRepo.Delete: %w", err) + } + n, err := res.RowsAffected() + if err != nil { + return fmt.Errorf("CapacityRepo.Delete: rows: %w", err) + } + if n == 0 { + return ErrNotFound + } + return nil +} diff --git a/internal/store/capacity_repo_test.go b/internal/store/capacity_repo_test.go new file mode 100644 index 0000000..6792045 --- /dev/null +++ b/internal/store/capacity_repo_test.go @@ -0,0 +1,74 @@ +package store + +import ( + "context" + "path/filepath" + "testing" +) + +func TestCapacityRepoUpsertGetList(t *testing.T) { + dir := t.TempDir() + db, err := Open(filepath.Join(dir, "test.db")) + if err != nil { + t.Fatalf("Open: %v", err) + } + defer db.Close() + repo := NewCapacityRepo(db) + ctx := context.Background() + + // Empty initially. + if _, err := repo.Get(ctx, "self"); err == nil { + t.Error("expected ErrNotFound on empty store") + } + rows, err := repo.List(ctx) + if err != nil { + t.Fatalf("List: %v", err) + } + if len(rows) != 0 { + t.Errorf("List: got %d rows, want 0", len(rows)) + } + + // Insert. + c1 := &NodeCapacity{NodeID: "self", CPUMillicores: 4000, MemoryMiB: 4096, DiskMiB: 4096} + if err := repo.Upsert(ctx, c1); err != nil { + t.Fatalf("Upsert: %v", err) + } + got, err := repo.Get(ctx, "self") + if err != nil { + t.Fatalf("Get: %v", err) + } + if got.CPUMillicores != 4000 || got.MemoryMiB != 4096 || got.DiskMiB != 4096 { + t.Errorf("Get: got %+v, want cpu=4000 mem=4096 disk=4096", got) + } + + // Update (overwrite). + c2 := &NodeCapacity{NodeID: "self", CPUMillicores: 8000, MemoryMiB: 8192, DiskMiB: 8192} + if err := repo.Upsert(ctx, c2); err != nil { + t.Fatalf("Upsert(update): %v", err) + } + got, _ = repo.Get(ctx, "self") + if got.CPUMillicores != 8000 { + t.Errorf("Update: cpu=%d, want 8000", got.CPUMillicores) + } + + // Add a second node. + c3 := &NodeCapacity{NodeID: "peer-1", CPUMillicores: 2000, MemoryMiB: 2048, DiskMiB: 2048} + if err := repo.Upsert(ctx, c3); err != nil { + t.Fatalf("Upsert(peer-1): %v", err) + } + rows, _ = repo.List(ctx) + if len(rows) != 2 { + t.Errorf("List: got %d rows, want 2", len(rows)) + } + + // Delete. + if err := repo.Delete(ctx, "peer-1"); err != nil { + t.Fatalf("Delete: %v", err) + } + if _, err := repo.Get(ctx, "peer-1"); err == nil { + t.Error("expected ErrNotFound after Delete") + } + if err := repo.Delete(ctx, "missing"); err == nil { + t.Error("expected ErrNotFound on Delete of missing row") + } +} diff --git a/internal/store/migrations/0005_node_capacity.sql b/internal/store/migrations/0005_node_capacity.sql new file mode 100644 index 0000000..2844e9b --- /dev/null +++ b/internal/store/migrations/0005_node_capacity.sql @@ -0,0 +1,12 @@ +-- Node capacity declaration for multi-node scheduling (v0.2 P02). +-- Loaded from `~/.orca/node.hcl` at `orca node join` and updated via +-- `orca node capacity --set`. Read by the dispatcher for bin-packing. +CREATE TABLE IF NOT EXISTS node_capacity ( + node_id TEXT PRIMARY KEY, + cpu_millicores INTEGER NOT NULL, + memory_mib INTEGER NOT NULL, + disk_mib INTEGER NOT NULL, + updated_at DATETIME NOT NULL +); + +CREATE INDEX IF NOT EXISTS idx_capacity_updated ON node_capacity(updated_at); diff --git a/internal/transport/idempotency.go b/internal/transport/idempotency.go new file mode 100644 index 0000000..acf08e9 --- /dev/null +++ b/internal/transport/idempotency.go @@ -0,0 +1,123 @@ +// Package transport — idempotency.go implements the X-Orca-Idempotency-Key +// header for cross-node dispatch (REQ-037). The dedupe store is a +// in-memory map with a TTL window; persistent dedupe across daemon +// restarts is out of scope for v0.2 (the bin-packing scheduler is +// single-daemon for now; the dedupe window just covers in-flight retries). +package transport + +import ( + "context" + "errors" + "sync" + "time" +) + +const ( + // IdempotencyHeader is the canonical header name. Casing-insensitive + // per HTTP spec, but we keep the canonical form for log clarity. + IdempotencyHeader = "X-Orca-Idempotency-Key" + // DedupeWindow is how long an idempotency key is honored after + // first use. Tuned for the in-flight retry window: a transient + // dispatch error followed by an exponential-backoff retry (max 5 + // attempts with cap 5s) completes well within 60s. The dedupe + // window is 5 minutes to cover cases where a peer processes a + // request but the response is lost on the wire. + DedupeWindow = 5 * time.Minute +) + +// dedupeEntry is a single (key -> response) record with expiry. +type dedupeEntry struct { + key string + jobID string + expiresAt time.Time +} + +// IdempotencyStore is a thread-safe in-memory dedupe map. Keys are +// scoped per-process; a restart drops the map. For P02 this is +// sufficient because the dispatcher is single-instance. +type IdempotencyStore struct { + mu sync.Mutex + entries map[string]dedupeEntry +} + +// NewIdempotencyStore returns an empty store. +func NewIdempotencyStore() *IdempotencyStore { + return &IdempotencyStore{entries: make(map[string]dedupeEntry)} +} + +// Get returns the recorded jobID for key, or "" if no entry is present +// (or the entry is expired). The second return is true if a live +// (non-expired) entry was found. +func (s *IdempotencyStore) Get(key string) (string, bool) { + if key == "" { + return "", false + } + s.mu.Lock() + defer s.mu.Unlock() + e, ok := s.entries[key] + if !ok { + return "", false + } + if time.Now().After(e.expiresAt) { + delete(s.entries, key) + return "", false + } + return e.jobID, true +} + +// Put records (key -> jobID) with a default expiry of DedupeWindow. +// Overwrites any prior entry (rare in practice since we check Get first). +func (s *IdempotencyStore) Put(key, jobID string) { + if key == "" || jobID == "" { + return + } + s.mu.Lock() + s.entries[key] = dedupeEntry{ + key: key, + jobID: jobID, + expiresAt: time.Now().Add(DedupeWindow), + } + s.mu.Unlock() +} + +// Sweep removes all expired entries. Called periodically by the dispatch +// service; safe to call concurrently. +func (s *IdempotencyStore) Sweep() { + now := time.Now() + s.mu.Lock() + for k, e := range s.entries { + if now.After(e.expiresAt) { + delete(s.entries, k) + } + } + s.mu.Unlock() +} + +// ErrIdempotencyKeyRequired is returned by retry helpers when a +// non-idempotent call (e.g., POST) is retried without an idempotency +// key. Matches REQ-037's "absent header + transient error → no retry". +var ErrIdempotencyKeyRequired = errors.New("retry requires X-Orca-Idempotency-Key header") + +// HeaderFromContext extracts the X-Orca-Idempotency-Key from a +// request-scoped context, if any. The dispatcher stores the key on +// the context via WithIdempotencyKey so downstream layers can read it +// without parsing headers. +type idempotencyKey struct{} + +// WithIdempotencyKey attaches an idempotency key to ctx. +func WithIdempotencyKey(ctx context.Context, key string) context.Context { + if key == "" { + return ctx + } + return context.WithValue(ctx, idempotencyKey{}, key) +} + +// IdempotencyKeyFromContext returns the key attached to ctx, or "". +func IdempotencyKeyFromContext(ctx context.Context) string { + if v := ctx.Value(idempotencyKey{}); v != nil { + if s, ok := v.(string); ok { + return s + } + } + return "" +} diff --git a/internal/transport/idempotency_test.go b/internal/transport/idempotency_test.go new file mode 100644 index 0000000..bf979a1 --- /dev/null +++ b/internal/transport/idempotency_test.go @@ -0,0 +1,133 @@ +package transport + +import ( + "context" + "errors" + "testing" + "time" +) + +func TestIdempotencyStorePutGet(t *testing.T) { + s := NewIdempotencyStore() + if _, ok := s.Get("missing"); ok { + t.Fatal("expected missing key to return ok=false") + } + s.Put("k1", "job-1") + if jobID, ok := s.Get("k1"); !ok || jobID != "job-1" { + t.Errorf("Get(k1): got (%q, %v), want (job-1, true)", jobID, ok) + } +} + +func TestIdempotencyStoreExpiry(t *testing.T) { + s := NewIdempotencyStore() + // Manually insert an expired entry. + s.entries["expired"] = dedupeEntry{ + key: "expired", + jobID: "old-job", + expiresAt: time.Now().Add(-1 * time.Minute), + } + if _, ok := s.Get("expired"); ok { + t.Fatal("expected expired entry to return ok=false") + } + if _, exists := s.entries["expired"]; exists { + t.Error("expected expired entry to be removed by Get") + } +} + +func TestIdempotencyStoreContext(t *testing.T) { + ctx := WithIdempotencyKey(context.Background(), "key-1") + if got := IdempotencyKeyFromContext(ctx); got != "key-1" { + t.Errorf("IdempotencyKeyFromContext: got %q, want key-1", got) + } + ctx2 := context.Background() + if got := IdempotencyKeyFromContext(ctx2); got != "" { + t.Errorf("IdempotencyKeyFromContext(empty): got %q, want \"\"", got) + } +} + +func TestRetrySucceedsAfterTransient(t *testing.T) { + calls := 0 + got, err := Do(context.Background(), DefaultRetryPolicy(), + func(_ context.Context, attempt int) (string, bool, error) { + calls++ + if attempt < 3 { + return "", true, errors.New("connection refused: try again") + } + return "ok", true, nil + }) + if err != nil { + t.Fatalf("Do: %v", err) + } + if got != "ok" { + t.Errorf("Do: got %q, want ok", got) + } + if calls != 3 { + t.Errorf("Do: got %d calls, want 3", calls) + } +} + +func TestRetryNoKeyOnTransient(t *testing.T) { + // Without an idempotency key AND a non-idempotent verb, a + // transient error on the first attempt must NOT retry (REQ-037). + calls := 0 + _, err := Do(context.Background(), DefaultRetryPolicy(), + func(_ context.Context, _ int) (string, bool, error) { + calls++ + return "", false, errors.New("connection refused") + }) + if err == nil { + t.Fatal("expected error, got nil") + } + if calls != 1 { + t.Errorf("expected 1 call (no retry without key), got %d", calls) + } +} + +func TestRetryPermanentError(t *testing.T) { + calls := 0 + _, err := Do(context.Background(), DefaultRetryPolicy(), + func(_ context.Context, _ int) (string, bool, error) { + calls++ + return "", true, ErrPermanent + }) + if !errors.Is(err, ErrPermanent) { + t.Errorf("expected ErrPermanent, got %v", err) + } + if calls != 1 { + t.Errorf("expected 1 call (permanent = no retry), got %d", calls) + } +} + +func TestRetryContextCancel(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() // cancel immediately + calls := 0 + _, err := Do(ctx, DefaultRetryPolicy(), + func(_ context.Context, _ int) (string, bool, error) { + calls++ + return "", true, errors.New("EOF") + }) + if !errors.Is(err, context.Canceled) { + t.Errorf("expected context.Canceled, got %v", err) + } +} + +func TestIsTransient(t *testing.T) { + cases := []struct { + err error + want bool + }{ + {nil, false}, + {errors.New("connection refused"), true}, + {errors.New("i/o timeout"), true}, + {errors.New("EOF"), true}, + {errors.New("no such host"), true}, + {errors.New("connection reset by peer"), true}, + {errors.New("invalid spec"), false}, + } + for _, c := range cases { + if got := IsTransient(c.err); got != c.want { + t.Errorf("IsTransient(%v): got %v, want %v", c.err, got, c.want) + } + } +} diff --git a/internal/transport/retry.go b/internal/transport/retry.go new file mode 100644 index 0000000..b9879ba --- /dev/null +++ b/internal/transport/retry.go @@ -0,0 +1,151 @@ +// Package transport — retry.go implements exponential backoff with +// jitter for cross-node dispatch retries. Per the P02 plan: 100ms +// initial, x2, 5s cap, max 5 attempts. Auto-retry only when the call +// is idempotent (X-Orca-Idempotency-Key header present, or the verb +// is intrinsically idempotent like GET/HEAD). +package transport + +import ( + "context" + "errors" + "math/rand" + "time" +) + +const ( + // RetryInitial is the first backoff interval. + RetryInitial = 100 * time.Millisecond + // RetryMax is the cap on backoff between attempts. + RetryMax = 5 * time.Second + // RetryMaxAttempts is the total attempt count (including the first). + RetryMaxAttempts = 5 +) + +// RetryPolicy carries the backoff configuration. Zero value is the +// default (100ms / 5s / 5 attempts). +type RetryPolicy struct { + Initial time.Duration + Max time.Duration + MaxAttempts int +} + +// DefaultRetryPolicy returns the P02 default. +func DefaultRetryPolicy() RetryPolicy { + return RetryPolicy{Initial: RetryInitial, Max: RetryMax, MaxAttempts: RetryMaxAttempts} +} + +// IsTransient reports whether err looks like a transient failure +// worth retrying. We treat network errors, context-deadline-exceeded +// (peer was slow but reachable), and a sentinel ErrTransient as +// retryable; everything else (4xx, validation, auth) is permanent. +func IsTransient(err error) bool { + if err == nil { + return false + } + if errors.Is(err, ErrTransient) { + return true + } + // We avoid pulling net/error here to keep dependencies minimal; + // the most common transient signature is the substring "connection + // refused" or "i/o timeout". Tests assert these explicitly. + s := err.Error() + for _, sub := range []string{"connection refused", "i/o timeout", "EOF", "no such host", "connection reset"} { + if contains(s, sub) { + return true + } + } + return false +} + +// ErrTransient is a sentinel callers can wrap to mark an error +// retryable. ErrPermanent is the opposite. +var ( + ErrTransient = errors.New("transient error") + ErrPermanent = errors.New("permanent error") +) + +// RetryableFunc is the signature Retry calls. It returns the result +// and an error. The bool indicates whether the call is idempotent +// (true = safe to retry without an idempotency key). +type RetryableFunc[T any] func(ctx context.Context, attempt int) (T, bool, error) + +// Do runs fn with backoff according to policy. It retries only if +// (a) the call is idempotent, OR (b) ctx carries an idempotency key +// (set via WithIdempotencyKey). Otherwise a transient error on the +// first attempt is returned immediately (REQ-037: no retry without +// the key). +// +// The generic result T lets callers reuse this for jobIDs, status +// responses, etc. without boxing through `any`. +func Do[T any](ctx context.Context, p RetryPolicy, fn RetryableFunc[T]) (T, error) { + var zero T + if p.MaxAttempts <= 0 { + p = DefaultRetryPolicy() + } + hasKey := IdempotencyKeyFromContext(ctx) != "" + for attempt := 1; attempt <= p.MaxAttempts; attempt++ { + if err := ctx.Err(); err != nil { + return zero, err + } + v, idempotent, err := fn(ctx, attempt) + if err == nil { + return v, nil + } + // Permanent errors never retry. + if errors.Is(err, ErrPermanent) { + return zero, err + } + // Last attempt — surface the error. + if attempt == p.MaxAttempts { + return zero, err + } + // Transient + no idempotency + not idempotent verb: no retry. + if IsTransient(err) && !idempotent && !hasKey { + return zero, err + } + // Wait with jittered backoff, but respect ctx cancellation. + wait := backoff(p.Initial, p.Max, attempt) + t := time.NewTimer(wait) + select { + case <-ctx.Done(): + t.Stop() + return zero, ctx.Err() + case <-t.C: + } + } + return zero, errors.New("retry.Do: exhausted attempts without error (impossible)") +} + +// backoff returns the wait duration for the n-th attempt (1-indexed). +// Formula: min(Initial * 2^(n-1), Max), with up to 25% jitter. +func backoff(initial, max time.Duration, n int) time.Duration { + d := initial + for i := 1; i < n; i++ { + d *= 2 + if d > max { + d = max + break + } + } + // Jitter: ±25% of d. + jitter := time.Duration(rand.Int63n(int64(d) / 2)) + d = d - d/4 + jitter + if d < 0 { + d = 0 + } + return d +} + +// contains is a tiny substring helper (avoids pulling strings for one +// call site; this is hot-path retry classification). +func contains(s, sub string) bool { + if len(sub) == 0 { + return true + } + for i := 0; i+len(sub) <= len(s); i++ { + if s[i:i+len(sub)] == sub { + return true + } + } + return false +}