diff --git a/internal/daemon/broker/broker.go b/internal/daemon/broker/broker.go index 9b46560..76c83fa 100644 --- a/internal/daemon/broker/broker.go +++ b/internal/daemon/broker/broker.go @@ -11,21 +11,25 @@ import ( var ErrClosed = errors.New("the broker has been closed") -// Defines main state where we fannin out snapshot of runs to the subscribers +// Broker centralised places for fan-out all snapshots to clients type Broker struct { + // Core mu sync.Mutex - subs map[uint64]*Subscriber + subs map[uint64]*subscriber snaps <-chan api.Snapshot + closed bool + + // Subscribers management currID uint64 - last *api.Snapshot - closed bool + // State handling + lastSnap *api.Snapshot } func New(snaps <-chan api.Snapshot) *Broker { return &Broker{ mu: sync.Mutex{}, - subs: make(map[uint64]*Subscriber), + subs: make(map[uint64]*subscriber), snaps: snaps, } } @@ -49,6 +53,7 @@ func (b *Broker) Run(ctx context.Context) { } } +// Close removes all connections to the broker subsequently closing all subscribers func (b *Broker) Close() { b.mu.Lock() defer b.mu.Unlock() @@ -61,7 +66,8 @@ func (b *Broker) Close() { clear(b.subs) } -func (b *Broker) Follow() (*Subscriber, error) { +// Follow creates a subscriber with unique id for the client +func (b *Broker) Follow() (*subscriber, error) { b.mu.Lock() defer b.mu.Unlock() @@ -70,19 +76,20 @@ func (b *Broker) Follow() (*Subscriber, error) { } b.currID++ - sub := &Subscriber{ + sub := &subscriber{ ID: b.currID, Snapshot: make(chan api.Snapshot, 1), } b.subs[sub.ID] = sub - if b.last != nil { - sub.set(*b.last) + if b.lastSnap != nil { + sub.set(*b.lastSnap) } return sub, nil } +// Unfollow removes subscriber for a broker func (b *Broker) Unfollow(id uint64) { b.mu.Lock() defer b.mu.Unlock() @@ -107,5 +114,5 @@ func (b *Broker) fanout(snap api.Snapshot) { sub.set(snap) } - b.last = &snap + b.lastSnap = &snap } diff --git a/internal/daemon/broker/broker_test.go b/internal/daemon/broker/broker_test.go new file mode 100644 index 0000000..135039c --- /dev/null +++ b/internal/daemon/broker/broker_test.go @@ -0,0 +1,277 @@ +package broker + +import ( + "context" + "errors" + "sync" + "testing" + "time" + + "github.com/HJyup/patchdock/internal/daemon/api" +) + +const waitLimit = 3 * time.Second + +func snapshotWith(runIDs ...string) api.Snapshot { + runs := make([]api.Run, 0, len(runIDs)) + for _, id := range runIDs { + runs = append(runs, api.Run{ID: id, Status: api.StatusQueued}) + } + return api.Snapshot{At: time.Now(), Runs: runs} +} + +func firstRunID(snap api.Snapshot) string { + if len(snap.Runs) == 0 { + return "" + } + return snap.Runs[0].ID +} + +func receive(t *testing.T, sub *subscriber) api.Snapshot { + t.Helper() + + select { + case snap, ok := <-sub.Snapshot: + if !ok { + t.Fatal("subscriber channel was closed, want a snapshot") + } + return snap + case <-time.After(waitLimit): + t.Fatal("no snapshot delivered") + return api.Snapshot{} + } +} + +func assertNoDelivery(t *testing.T, sub *subscriber, within time.Duration) { + t.Helper() + + select { + case snap, ok := <-sub.Snapshot: + if ok { + t.Fatalf("received %v, want no delivery", firstRunID(snap)) + } + case <-time.After(within): + } +} + +func awaitClosed(t *testing.T, sub *subscriber) { + t.Helper() + + deadline := time.After(waitLimit) + for { + select { + case _, ok := <-sub.Snapshot: + if !ok { + return + } + case <-deadline: + t.Fatal("subscriber channel was never closed") + } + } +} + +func mustFollow(t *testing.T, b *Broker) *subscriber { + t.Helper() + + sub, err := b.Follow() + if err != nil { + t.Fatalf("follow: %v", err) + } + return sub +} + +// Fan-out & subscribers + +func TestFanoutReachesEverySubscriber(t *testing.T) { + b := New(nil) + + first := mustFollow(t, b) + second := mustFollow(t, b) + + if first.ID == second.ID { + t.Errorf("subscriber IDs collided at %d", first.ID) + } + + b.fanout(snapshotWith("run-1")) + + if got := firstRunID(receive(t, first)); got != "run-1" { + t.Errorf("first subscriber got %q, want run-1", got) + } + if got := firstRunID(receive(t, second)); got != "run-1" { + t.Errorf("second subscriber got %q, want run-1", got) + } +} + +func TestFollowReplaysTheLastSnapshot(t *testing.T) { + b := New(nil) + + b.fanout(snapshotWith("run-1")) + b.fanout(snapshotWith("run-2")) + + late := mustFollow(t, b) + if got := firstRunID(receive(t, late)); got != "run-2" { + t.Errorf("late subscriber got %q, want the most recent snapshot run-2", got) + } +} + +func TestFollowBeforeAnySnapshotDeliversNothing(t *testing.T) { + b := New(nil) + + sub := mustFollow(t, b) + assertNoDelivery(t, sub, 50*time.Millisecond) +} + +func TestSlowSubscriberOnlySeesTheLatestSnapshot(t *testing.T) { + b := New(nil) + sub := mustFollow(t, b) + + b.fanout(snapshotWith("run-1")) + b.fanout(snapshotWith("run-2")) + b.fanout(snapshotWith("run-3")) + + if got := firstRunID(receive(t, sub)); got != "run-3" { + t.Errorf("got %q, want the newest snapshot run-3", got) + } + + assertNoDelivery(t, sub, 50*time.Millisecond) +} + +func TestUnfollowClosesTheChannelAndStopsDelivery(t *testing.T) { + b := New(nil) + + staying := mustFollow(t, b) + leaving := mustFollow(t, b) + + b.Unfollow(leaving.ID) + awaitClosed(t, leaving) + + b.fanout(snapshotWith("run-1")) + if got := firstRunID(receive(t, staying)); got != "run-1" { + t.Errorf("remaining subscriber got %q, want run-1", got) + } +} + +// Shutdown + +func TestCloseClosesEverySubscriber(t *testing.T) { + b := New(nil) + + first := mustFollow(t, b) + second := mustFollow(t, b) + + b.Close() + + awaitClosed(t, first) + awaitClosed(t, second) +} + +func TestFollowAfterCloseReturnsErrClosed(t *testing.T) { + b := New(nil) + b.Close() + + sub, err := b.Follow() + if !errors.Is(err, ErrClosed) { + t.Fatalf("Follow after Close = %v, want %v", err, ErrClosed) + } + if sub != nil { + t.Error("Follow returned a subscriber alongside an error") + } +} + +func TestCloseIsIdempotent(t *testing.T) { + b := New(nil) + sub := mustFollow(t, b) + + b.Close() + b.Close() // must not panic + + awaitClosed(t, sub) +} + +func TestUnfollowAfterCloseIsANoop(t *testing.T) { + b := New(nil) + sub := mustFollow(t, b) + + b.Close() + b.Unfollow(sub.ID) // must not double-close + + awaitClosed(t, sub) +} + +// Run loop + +func TestRunFansOutFromItsSource(t *testing.T) { + source := make(chan api.Snapshot) + b := New(source) + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + go b.Run(ctx) + + sub := mustFollow(t, b) + source <- snapshotWith("run-1") + + if got := firstRunID(receive(t, sub)); got != "run-1" { + t.Errorf("got %q, want run-1", got) + } +} + +func TestRunClosesSubscribersOnContextCancel(t *testing.T) { + source := make(chan api.Snapshot) + b := New(source) + + ctx, cancel := context.WithCancel(context.Background()) + go b.Run(ctx) + + sub := mustFollow(t, b) + cancel() + + awaitClosed(t, sub) + if _, err := b.Follow(); !errors.Is(err, ErrClosed) { + t.Fatalf("Follow after shutdown = %v, want %v", err, ErrClosed) + } +} + +func TestConcurrentFollowUnfollowAndFanout(t *testing.T) { + source := make(chan api.Snapshot) + b := New(source) + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + go b.Run(ctx) + + stop := make(chan struct{}) + producerDone := make(chan struct{}) + go func() { + defer close(producerDone) + for { + select { + case <-stop: + return + case source <- snapshotWith("run-1"): + } + } + }() + + var followers sync.WaitGroup + for range 8 { + followers.Go(func() { + for range 20 { + sub, err := b.Follow() + if err != nil { + return + } + select { + case <-sub.Snapshot: + default: + } + b.Unfollow(sub.ID) + } + }) + } + + followers.Wait() + + close(stop) + <-producerDone +} diff --git a/internal/daemon/broker/subscriber.go b/internal/daemon/broker/subscriber.go index ca1c658..3784962 100644 --- a/internal/daemon/broker/subscriber.go +++ b/internal/daemon/broker/subscriber.go @@ -4,16 +4,18 @@ import ( "sync" "github.com/HJyup/patchdock/internal/daemon/api" + "github.com/HJyup/patchdock/internal/utils" ) -type Subscriber struct { +// Subscriber represents a one-to-one connection to broker for a client +type subscriber struct { ID uint64 Snapshot chan api.Snapshot closed bool mu sync.Mutex } -func (s *Subscriber) set(snap api.Snapshot) bool { +func (s *subscriber) set(snap api.Snapshot) bool { s.mu.Lock() defer s.mu.Unlock() @@ -21,22 +23,10 @@ func (s *Subscriber) set(snap api.Snapshot) bool { return false } - // Drain the channel, since a new snapshot incoming - select { - case <-s.Snapshot: - default: - } - - // Place the value inside - select { - case s.Snapshot <- snap: - return true - default: - return false - } + return utils.SendLatest(s.Snapshot, snap) } -func (s *Subscriber) close() { +func (s *subscriber) close() { s.mu.Lock() defer s.mu.Unlock() diff --git a/internal/daemon/config/config.go b/internal/daemon/config/config.go index 3afb45b..c7a4c2a 100644 --- a/internal/daemon/config/config.go +++ b/internal/daemon/config/config.go @@ -14,17 +14,23 @@ import ( const ( DefaultMaxContainers = 3 DefaultRetention = utils.Duration(15 * time.Minute) + DefaultSnapshotTick = utils.Duration(200 * time.Millisecond) + DefaultInboxSize = 256 ) type Config struct { MaxContainers int `json:"max_containers"` Retention utils.Duration `json:"retention"` + SnapshotTick utils.Duration `json:"snapshot_tick"` + InboxSize int `json:"inbox_size"` } func Defaults() Config { return Config{ MaxContainers: DefaultMaxContainers, Retention: DefaultRetention, + SnapshotTick: DefaultSnapshotTick, + InboxSize: DefaultInboxSize, } } @@ -81,6 +87,12 @@ func (c *Config) Validate() error { if c.Retention <= 0 { errs = append(errs, errors.New("config.retention: must be > 0")) } + if c.InboxSize <= 0 { + errs = append(errs, errors.New("config.inbox_size: must be >= 1")) + } + if c.SnapshotTick <= 0 { + errs = append(errs, errors.New("config.snapshot_tick: must be > 0")) + } return errors.Join(errs...) } diff --git a/internal/daemon/config/config_test.go b/internal/daemon/config/config_test.go new file mode 100644 index 0000000..b18b13f --- /dev/null +++ b/internal/daemon/config/config_test.go @@ -0,0 +1,187 @@ +package config + +import ( + "os" + "path/filepath" + "strings" + "testing" + "time" +) + +func configPath(t *testing.T) string { + t.Helper() + return filepath.Join(t.TempDir(), "config.json") +} + +func writeConfig(t *testing.T, contents string) string { + t.Helper() + + path := configPath(t) + if err := os.WriteFile(path, []byte(contents), 0o600); err != nil { + t.Fatalf("write config: %v", err) + } + return path +} + +func mustLoad(t *testing.T, path string) *Config { + t.Helper() + + cfg, err := Load(path) + if err != nil { + t.Fatalf("load config: %v", err) + } + return cfg +} + +func TestDefaultsAreValid(t *testing.T) { + cfg := Defaults() + + if err := cfg.Validate(); err != nil { + t.Fatalf("Defaults() does not satisfy Validate: %v", err) + } +} + +func TestLoadCreatesTheFileWhenMissing(t *testing.T) { + path := configPath(t) + + cfg := mustLoad(t, path) + + if *cfg != Defaults() { + t.Errorf("config = %+v, want %+v", *cfg, Defaults()) + } + + info, err := os.Stat(path) + if err != nil { + t.Fatalf("config file was not created: %v", err) + } + + // The runtime directory is per-user and this file configures a daemon, so + // it must not be world- or group-readable. + if perm := info.Mode().Perm(); perm != 0o600 { + t.Errorf("permissions = %04o, want 0600", perm) + } +} + +func TestCreatedFileWritesDurationsAsStrings(t *testing.T) { + path := configPath(t) + mustLoad(t, path) + + raw, err := os.ReadFile(path) + if err != nil { + t.Fatalf("read config: %v", err) + } + contents := string(raw) + + for _, want := range []string{`"retention": "15m0s"`, `"snapshot_tick": "200ms"`} { + if !strings.Contains(contents, want) { + t.Errorf("config file does not contain %s\ngot:\n%s", want, contents) + } + } +} + +func TestCreatedFileReloadsUnchanged(t *testing.T) { + path := configPath(t) + + created := mustLoad(t, path) + reloaded := mustLoad(t, path) + + if *created != *reloaded { + t.Errorf("reloaded = %+v, want %+v", *reloaded, *created) + } +} + +func TestLoadParsesDurationStrings(t *testing.T) { + path := writeConfig(t, `{"retention": "90s", "snapshot_tick": "1s"}`) + + cfg := mustLoad(t, path) + + if got := cfg.Retention.Duration(); got != 90*time.Second { + t.Errorf("retention = %v, want 90s", got) + } + if got := cfg.SnapshotTick.Duration(); got != time.Second { + t.Errorf("snapshot_tick = %v, want 1s", got) + } +} + +func TestLoadRejectsUnknownFields(t *testing.T) { + path := writeConfig(t, `{"max_container": 7}`) + + _, err := Load(path) + if err == nil { + t.Fatal("a misspelled field was accepted") + } + if !strings.Contains(err.Error(), "max_container") { + t.Errorf("error = %v, want it to name the unknown field", err) + } +} + +func TestLoadRejectsBadInput(t *testing.T) { + tests := []struct { + name string + contents string + wantErr string + }{ + { + name: "empty file", + contents: "", + wantErr: "delete it to recreate it with defaults", + }, + { + name: "truncated json", + contents: `{"max_containers":`, + wantErr: "decode config", + }, + { + name: "wrong type for a duration", + contents: `{"retention": true}`, + wantErr: "decode config", + }, + { + name: "unparseable duration", + contents: `{"retention": "soon"}`, + wantErr: "decode config", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + path := writeConfig(t, tt.contents) + + _, err := Load(path) + if err == nil { + t.Fatalf("Load(%q) succeeded, want an error", tt.contents) + } + if !strings.Contains(err.Error(), tt.wantErr) { + t.Errorf("error = %v, want it to contain %q", err, tt.wantErr) + } + }) + } +} + +func TestValidateReportsEveryViolationAtOnce(t *testing.T) { + cfg := Config{} + + err := cfg.Validate() + if err == nil { + t.Fatal("the zero config was accepted") + } + + for _, want := range []string{ + "config.max_containers", + "config.retention", + "config.inbox_size", + "config.snapshot_tick", + } { + if !strings.Contains(err.Error(), want) { + t.Errorf("error is missing %s\ngot: %v", want, err) + } + } +} + +func TestLoadFailsWhenTheDirectoryDoesNotExist(t *testing.T) { + path := filepath.Join(t.TempDir(), "missing-dir", "config.json") + + if _, err := Load(path); err == nil { + t.Fatal("Load succeeded with no directory to write into") + } +} diff --git a/internal/daemon/daemon.go b/internal/daemon/daemon.go index 063fe22..bf548a1 100644 --- a/internal/daemon/daemon.go +++ b/internal/daemon/daemon.go @@ -64,8 +64,8 @@ func RunServer(ctx context.Context, dir runtimedir.Dir) error { } defer cli.Close() - q := queue.New(ctx, queue.Config{Runner: pipelineRunner(cli), Retention: cfg.Retention.Duration(), MaxContainers: cfg.MaxContainers}) - ch := q.Snaps() + q := queue.New(ctx, pipelineRunner(cli), cfg) + ch := q.Snapshots() b := broker.New(ch) diff --git a/internal/daemon/queue/runner.go b/internal/daemon/queue/contract.go similarity index 64% rename from internal/daemon/queue/runner.go rename to internal/daemon/queue/contract.go index d7afca0..dba52b0 100644 --- a/internal/daemon/queue/runner.go +++ b/internal/daemon/queue/contract.go @@ -7,17 +7,19 @@ import ( "github.com/HJyup/patchdock/internal/types" ) -// RunSpec is everything a Runner needs to execute one run +// RunSpec declares information that needed for one run in a queue type RunSpec struct { RunID string Repo string Task types.Task } +// Outcome is what a finished Runner returns back type Outcome struct { Accepted bool Branch string Patch auditlog.PatchStat } +// Runner executes one run to completion, reporting stage progress through reporter type Runner func(ctx context.Context, spec RunSpec, rep Reporter) (Outcome, error) diff --git a/internal/daemon/queue/event.go b/internal/daemon/queue/event.go index a6a082e..1ed27af 100644 --- a/internal/daemon/queue/event.go +++ b/internal/daemon/queue/event.go @@ -4,35 +4,42 @@ import ( "github.com/HJyup/patchdock/internal/types" ) +// event is the central communication object with a queue. type event interface{ queueEvent() } +// addEvent represents a new task that needs to be accepted to the queue type addEvent struct { repo string task types.Task res chan<- string } +// cancelEvent represents a queued task (by id) which we want to cancel type cancelEvent struct { runID string - err chan<- error + res chan<- error } +// stageEvent changes the state of the run type stageEvent struct { runID string stage types.StageName attempt int } +// acitivityEvent changes the text of the one-line activity per run type activityEvent struct { runID string text string } +// summaryEvent represents a summary defined by agents (in current implementation used by planner) type summaryEvent struct { runID string text string } +// doneEvent represents a full finished run from runner and we want to report back outcome type doneEvent struct { runID string out Outcome diff --git a/internal/daemon/queue/queue.go b/internal/daemon/queue/queue.go index 1ccff12..a5a3e47 100644 --- a/internal/daemon/queue/queue.go +++ b/internal/daemon/queue/queue.go @@ -7,78 +7,76 @@ import ( "time" "github.com/HJyup/patchdock/internal/daemon/api" + "github.com/HJyup/patchdock/internal/daemon/config" "github.com/HJyup/patchdock/internal/types" "github.com/HJyup/patchdock/internal/utils" ) -const ( - publishInterval = 200 * time.Millisecond - inboxSize = 256 -) - var ( ErrNotFound = errors.New("run not found") ErrFinished = errors.New("run has already finished") - ErrRepoPath = errors.New("repo must be an absolute path") ) -type Config struct { - Runner Runner - Retention time.Duration - MaxContainers int -} - +// run represents a live state published to watchers type run struct { state *api.Run task types.Task } -type queuedTasks struct { +// queuedRun is an admission ticket for a run that has not started yet +type queuedRun struct { ctx context.Context runID string } type Queue struct { + // Core inbox chan event snaps chan api.Snapshot runner Runner - // defines retention policy for finilised runs + // Config for making queue work retention time.Duration maxContainers int + snapshotTick time.Duration ctx context.Context - runs map[string]*run - // define all nesseary context cancel function so it's easy to cancel certain runs + // State information + runs map[string]*run cancels map[string]context.CancelFunc + dirty bool - // cloning runs are the most expensive operation in the Queue. Dirty will guard of cloning up-to-date data - dirty bool - - // scheduler implementation arrays - waiting []queuedTasks + // Scheduling + queuedRuns []queuedRun } -func New(ctx context.Context, cfg Config) *Queue { +func New(ctx context.Context, runner Runner, cfg *config.Config) *Queue { return &Queue{ - inbox: make(chan event, inboxSize), - snaps: make(chan api.Snapshot, 1), - runner: cfg.Runner, + inbox: make(chan event, cfg.InboxSize), + snaps: make(chan api.Snapshot, 1), + runner: runner, + maxContainers: cfg.MaxContainers, + snapshotTick: cfg.SnapshotTick.Duration(), + retention: cfg.Retention.Duration(), + ctx: ctx, + + runs: make(map[string]*run), + cancels: make(map[string]context.CancelFunc), - retention: cfg.Retention, - ctx: ctx, - runs: make(map[string]*run), - cancels: make(map[string]context.CancelFunc), + queuedRuns: make([]queuedRun, 0), } } func (q *Queue) Run() { - ticker := time.NewTicker(publishInterval) + ticker := time.NewTicker(q.snapshotTick) defer ticker.Stop() + defer close(q.snaps) + // Publish empty state to the queue q.publish() + for { select { case <-q.ctx.Done(): @@ -99,15 +97,66 @@ func (q *Queue) Run() { } } -func (q *Queue) Snaps() <-chan api.Snapshot { +// Snapshots returns a channel as a single point of recieving updates from the queue +func (q *Queue) Snapshots() <-chan api.Snapshot { return q.snaps } +// Core queue functions (remove, schedule, publish) + +func (q *Queue) evict() { + cutoff := time.Now().Add(-q.retention) + + for id, r := range q.runs { + if r.state.FinishedAt != nil && r.state.FinishedAt.Before(cutoff) { + delete(q.runs, id) + q.dirty = true + } + } +} + +func (q *Queue) publish() { + snap := q.snapshot() + + utils.SendLatest(q.snaps, snap) +} + +func (q *Queue) schedule() { + for len(q.queuedRuns) > 0 && q.activeCount() < q.maxContainers { + queued := q.queuedRuns[0] + q.queuedRuns = q.queuedRuns[1:] + + r, ok := q.runs[queued.runID] + if !ok || api.IsFinilised(r.state.Status) { + continue + } + + // Has been cannceled before (invalidate queued slice) + if queued.ctx.Err() != nil { + q.done(doneEvent{runID: queued.runID, cancelled: true}) + continue + } + + now := time.Now() + r.state.Status = api.StatusStarted + r.state.StartedAt = &now + q.dirty = true + + go q.execute(queued.ctx, RunSpec{ + RunID: r.state.ID, + Repo: r.state.Repo, + Task: r.task, + }) + } +} + +// Tasks public methods + +// Add queues one task and blocks until the queue assigns it a run ID func (q *Queue) Add(repo string, task types.Task) (string, error) { res := make(chan string, 1) - select { - case q.inbox <- addEvent{repo: filepath.Clean(repo), task: task, res: res}: - case <-q.ctx.Done(): + + if !q.send(q.ctx, addEvent{repo: filepath.Clean(repo), task: task, res: res}) { return "", q.ctx.Err() } @@ -119,12 +168,11 @@ func (q *Queue) Add(repo string, task types.Task) (string, error) { } } +// Cancel stops a run and reports whether the queue accepted the request func (q *Queue) Cancel(ctx context.Context, runID string) error { reply := make(chan error, 1) - select { - case q.inbox <- cancelEvent{runID: runID, err: reply}: - case <-ctx.Done(): + if !q.send(ctx, cancelEvent{runID: runID, res: reply}) { return ctx.Err() } @@ -136,6 +184,8 @@ func (q *Queue) Cancel(ctx context.Context, runID string) error { } } +// Message handling + func (q *Queue) handle(ev event) { switch e := ev.(type) { case addEvent: @@ -173,19 +223,18 @@ func (q *Queue) add(e addEvent) { e.res <- id q.cancels[r.state.ID] = cancel - - q.waiting = append(q.waiting, queuedTasks{ctx: ctx, runID: id}) + q.queuedRuns = append(q.queuedRuns, queuedRun{ctx: ctx, runID: id}) } func (q *Queue) cancel(e cancelEvent) { r, ok := q.runs[e.runID] if !ok { - e.err <- ErrNotFound + e.res <- ErrNotFound return } if api.IsFinilised(r.state.Status) { - e.err <- ErrFinished + e.res <- ErrFinished return } @@ -193,7 +242,16 @@ func (q *Queue) cancel(e cancelEvent) { cancel() } - e.err <- nil + // A run that never started has no pipeline to notice the cancelled context + // and report back, so retire it here. Waiting for the scheduler to do it is + // not enough: that only happens when a container slot is free, so on a busy + // queue the run would sit as queued until one opened up. Its ticket stays in + // queuedRuns and is skipped at admission by the finalised check. + if r.state.StartedAt == nil { + q.done(doneEvent{runID: e.runID, cancelled: true}) + } + + e.res <- nil } func (q *Queue) stage(e stageEvent) { @@ -214,17 +272,6 @@ func (q *Queue) stage(e stageEvent) { q.dirty = true } -func (q *Queue) active() int { - n := 0 - for _, r := range q.runs { - if r.state.StartedAt != nil && !api.IsFinilised(r.state.Status) { - n++ - } - } - - return n -} - func (q *Queue) activity(e activityEvent) { r, ok := q.runs[e.runID] if !ok || api.IsFinilised(r.state.Status) { @@ -284,75 +331,42 @@ func (q *Queue) done(e doneEvent) { q.dirty = true } -func (q *Queue) schedule() { - for len(q.waiting) > 0 && q.active() < q.maxContainers { - queued := q.waiting[0] - q.waiting = q.waiting[1:] - - r, ok := q.runs[queued.runID] - if !ok || api.IsFinilised(r.state.Status) { - continue - } - - if queued.ctx.Err() != nil { - q.done(doneEvent{runID: queued.runID, cancelled: true}) - continue - } - - now := time.Now() - r.state.Status = api.StatusStarted - r.state.StartedAt = &now - q.dirty = true - - go q.execute(queued.ctx, RunSpec{ - RunID: r.state.ID, - Repo: r.state.Repo, - Task: r.task, - }) - } -} +// Execution & Additional methods func (q *Queue) execute(ctx context.Context, spec RunSpec) { out, err := q.runner(ctx, spec, &reporter{queue: q, runID: spec.RunID}) - event := doneEvent{ + q.send(q.ctx, doneEvent{ runID: spec.RunID, out: out, err: err, cancelled: ctx.Err() != nil, - } + }) +} - select { - case q.inbox <- event: - case <-q.ctx.Done(): +func (q *Queue) snapshot() api.Snapshot { + runs := make([]api.Run, 0, len(q.runs)) + for _, r := range q.runs { + runs = append(runs, r.state.Clone()) } + return api.Snapshot{At: time.Now(), Runs: runs} } -func (q *Queue) evict() { - cutoff := time.Now().Add(-q.retention) - - for id, r := range q.runs { - if r.state.FinishedAt != nil && r.state.FinishedAt.Before(cutoff) { - delete(q.runs, id) - q.dirty = true +func (q *Queue) activeCount() int { + n := 0 + for _, r := range q.runs { + if r.state.StartedAt != nil && !api.IsFinilised(r.state.Status) { + n++ } } + return n } -func (q *Queue) publish() { - snap := q.snapshot() - +// send hands event to the queue loop, abandoning the attempt if ctx is cancelled +func (q *Queue) send(ctx context.Context, ev event) bool { select { - case <-q.snaps: - default: - } - - q.snaps <- snap -} - -func (q *Queue) snapshot() api.Snapshot { - runs := make([]api.Run, 0, len(q.runs)) - for _, r := range q.runs { - runs = append(runs, r.state.Clone()) + case q.inbox <- ev: + return true + case <-ctx.Done(): + return false } - return api.Snapshot{At: time.Now(), Runs: runs} } diff --git a/internal/daemon/queue/queue_test.go b/internal/daemon/queue/queue_test.go new file mode 100644 index 0000000..01e61d7 --- /dev/null +++ b/internal/daemon/queue/queue_test.go @@ -0,0 +1,648 @@ +package queue + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/HJyup/patchdock/internal/auditlog" + "github.com/HJyup/patchdock/internal/daemon/api" + "github.com/HJyup/patchdock/internal/daemon/config" + "github.com/HJyup/patchdock/internal/types" + "github.com/HJyup/patchdock/internal/utils" +) + +const testTick = 1 * time.Millisecond +const waitLimit = 3 * time.Second + +func testConfig(maxContainers int) config.Config { + cfg := config.Defaults() + cfg.MaxContainers = maxContainers + cfg.SnapshotTick = utils.Duration(testTick) + return cfg +} + +func newQueue(t *testing.T, cfg config.Config) *Queue { + t.Helper() + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + return New(ctx, nil, &cfg) +} + +func newRunningQueue(t *testing.T, cfg config.Config, runner Runner) *Queue { + t.Helper() + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + q := New(ctx, runner, &cfg) + go q.Run() + + return q +} + +func mustTask(t *testing.T, description string) types.Task { + t.Helper() + + task, err := types.NewTask(types.Task{Description: description}) + if err != nil { + t.Fatalf("build task: %v", err) + } + return task +} + +func seedRun(t *testing.T, q *Queue, description string) string { + t.Helper() + + res := make(chan string, 1) + q.add(addEvent{repo: "/repo", task: mustTask(t, description), res: res}) + + select { + case id := <-res: + return id + default: + t.Fatal("add did not reply with a run id") + return "" + } +} + +func awaitRun(t *testing.T, q *Queue, runID string, want api.Status) api.Run { + t.Helper() + + deadline := time.After(waitLimit) + var last api.Status + + for { + select { + case snap := <-q.Snapshots(): + for _, r := range snap.Runs { + if r.ID != runID { + continue + } + last = r.Status + if r.Status == want { + return r + } + } + case <-deadline: + t.Fatalf("run %s reached %q, want %q", runID, last, want) + return api.Run{} + } + } +} + +type blockingRunner struct { + started chan string + release chan struct{} +} + +func newBlockingRunner(capacity int) *blockingRunner { + return &blockingRunner{ + started: make(chan string, capacity), + release: make(chan struct{}), + } +} + +func (b *blockingRunner) run(ctx context.Context, spec RunSpec, _ Reporter) (Outcome, error) { + b.started <- spec.RunID + + select { + case <-b.release: + case <-ctx.Done(): + } + return Outcome{}, nil +} + +func (b *blockingRunner) awaitStart(t *testing.T) string { + t.Helper() + + select { + case id := <-b.started: + return id + case <-time.After(waitLimit): + t.Fatal("no run was admitted") + return "" + } +} + +func (b *blockingRunner) assertNoStart(t *testing.T, within time.Duration) { + t.Helper() + + select { + case id := <-b.started: + t.Fatalf("run %s was admitted when it should not have been", id) + case <-time.After(within): + } +} + +// States + +func TestAddRegistersAQueuedRun(t *testing.T) { + q := newQueue(t, testConfig(1)) + + id := seedRun(t, q, "first line\nsecond line") + + r, ok := q.runs[id] + if !ok { + t.Fatalf("run %s is not in the run map", id) + } + + if got := r.state.Status; got != api.StatusQueued { + t.Errorf("status = %q, want %q", got, api.StatusQueued) + } + if got := r.state.Title; got != "first line" { + t.Errorf("title = %q, want the first line of the description", got) + } + if r.state.QueuedAt.IsZero() { + t.Error("QueuedAt was never stamped") + } + if r.state.StartedAt != nil { + t.Error("a queued run must not have StartedAt") + } + if _, ok := q.cancels[id]; !ok { + t.Error("no cancel func registered, so the run can never be cancelled") + } + if len(q.queuedRuns) != 1 || q.queuedRuns[0].runID != id { + t.Errorf("queuedRuns = %+v, want one ticket for %s", q.queuedRuns, id) + } + if !q.dirty { + t.Error("adding a run must mark the queue dirty so the change is published") + } +} + +func TestCancelUnknownRun(t *testing.T) { + q := newQueue(t, testConfig(1)) + + res := make(chan error, 1) + q.cancel(cancelEvent{runID: "run-nope", res: res}) + + if err := <-res; !errors.Is(err, ErrNotFound) { + t.Fatalf("cancel(unknown) = %v, want %v", err, ErrNotFound) + } +} + +func TestCancelFinishedRun(t *testing.T) { + q := newQueue(t, testConfig(1)) + id := seedRun(t, q, "already done") + q.done(doneEvent{runID: id, out: Outcome{Accepted: true}}) + + res := make(chan error, 1) + q.cancel(cancelEvent{runID: id, res: res}) + + if err := <-res; !errors.Is(err, ErrFinished) { + t.Fatalf("cancel(finished) = %v, want %v", err, ErrFinished) + } +} + +func TestCancelQueuedRunRetiresItImmediately(t *testing.T) { + q := newQueue(t, testConfig(1)) + id := seedRun(t, q, "never started") + runCtx := q.queuedRuns[0].ctx + + res := make(chan error, 1) + q.cancel(cancelEvent{runID: id, res: res}) + + if err := <-res; err != nil { + t.Fatalf("cancel(queued) = %v, want nil", err) + } + if runCtx.Err() == nil { + t.Error("the run's context was not cancelled") + } + + state := q.runs[id].state + if state.Status != api.StatusCancelled { + t.Errorf("status = %q, want %q", state.Status, api.StatusCancelled) + } + if state.FinishedAt == nil { + t.Error("FinishedAt was never stamped, so evict will never reap it") + } + if state.StartedAt != nil { + t.Error("a run cancelled before admission must never have started") + } +} + +func TestCancelRunningRunOnlyTripsTheContext(t *testing.T) { + q := newQueue(t, testConfig(1)) + id := seedRun(t, q, "in flight") + runCtx := q.queuedRuns[0].ctx + + now := time.Now() + q.runs[id].state.StartedAt = &now + q.runs[id].state.Status = api.StatusCoding + + res := make(chan error, 1) + q.cancel(cancelEvent{runID: id, res: res}) + + if err := <-res; err != nil { + t.Fatalf("cancel(running) = %v, want nil", err) + } + if runCtx.Err() == nil { + t.Error("the run's context was not cancelled") + } + if got := q.runs[id].state.Status; got != api.StatusCoding { + t.Errorf("status = %q, want it to stay %q until execute reports done", got, api.StatusCoding) + } + if q.runs[id].state.FinishedAt != nil { + t.Error("a still-running run was marked finished before its pipeline reported") + } +} + +func TestStageUpdatesStatusAndAttempt(t *testing.T) { + tests := []struct { + stage types.StageName + want api.Status + }{ + {types.StagePlanner, api.StatusPlanning}, + {types.StageExecutor, api.StatusCoding}, + {types.StageReviewer, api.StatusReviewing}, + } + + for _, tt := range tests { + t.Run(string(tt.stage), func(t *testing.T) { + q := newQueue(t, testConfig(1)) + id := seedRun(t, q, "work") + q.runs[id].state.Activity = "stale activity" + + q.stage(stageEvent{runID: id, stage: tt.stage, attempt: 2}) + + state := q.runs[id].state + if state.Status != tt.want { + t.Errorf("status = %q, want %q", state.Status, tt.want) + } + if state.Attempt != 2 { + t.Errorf("attempt = %d, want 2", state.Attempt) + } + if state.StageStartedAt == nil { + t.Error("StageStartedAt was never stamped") + } + if state.Activity != "" { + t.Errorf("activity = %q, want it cleared on a stage change", state.Activity) + } + }) + } +} + +func TestSummaryIsIgnoredWhenUnchanged(t *testing.T) { + q := newQueue(t, testConfig(1)) + id := seedRun(t, q, "work") + + q.summary(summaryEvent{runID: id, text: "a plan"}) + if got := q.runs[id].state.Summary; got != "a plan" { + t.Fatalf("summary = %q, want %q", got, "a plan") + } + + q.dirty = false + q.summary(summaryEvent{runID: id, text: "a plan"}) + if q.dirty { + t.Error("repeating the same summary marked the queue dirty") + } +} + +func TestDoneTerminalStates(t *testing.T) { + patch := auditlog.PatchStat{Files: 2, Additions: 10, Deletions: 3} + + tests := []struct { + name string + event doneEvent + wantStatus api.Status + wantBranch string + wantSummary string + wantPatch bool + }{ + { + name: "accepted", + event: doneEvent{out: Outcome{Accepted: true, Branch: "patchdock/run-1", Patch: patch}}, + wantStatus: api.StatusSucceeded, + wantBranch: "patchdock/run-1", + wantPatch: true, + }, + { + name: "reviewer rejected every attempt", + event: doneEvent{out: Outcome{Accepted: false, Patch: patch}}, + wantStatus: api.StatusRejected, + wantPatch: true, + }, + { + name: "stage errored", + event: doneEvent{err: errors.New("planner stage: boom")}, + wantStatus: api.StatusFailed, + wantSummary: "planner stage: boom", + }, + { + name: "cancelled", + event: doneEvent{cancelled: true}, + wantStatus: api.StatusCancelled, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + q := newQueue(t, testConfig(1)) + id := seedRun(t, q, "work") + q.runs[id].state.Activity = "mid-flight activity" + + tt.event.runID = id + q.done(tt.event) + + state := q.runs[id].state + if state.Status != tt.wantStatus { + t.Errorf("status = %q, want %q", state.Status, tt.wantStatus) + } + if state.Branch != tt.wantBranch { + t.Errorf("branch = %q, want %q", state.Branch, tt.wantBranch) + } + if tt.wantSummary != "" && state.Summary != tt.wantSummary { + t.Errorf("summary = %q, want %q", state.Summary, tt.wantSummary) + } + if got := state.Patch != nil; got != tt.wantPatch { + t.Errorf("patch recorded = %v, want %v", got, tt.wantPatch) + } + + if state.FinishedAt == nil { + t.Error("FinishedAt was never stamped, so this run will never be evicted") + } + if state.Activity != "" { + t.Errorf("activity = %q, want it cleared on a finished run", state.Activity) + } + if !api.IsFinilised(state.Status) { + t.Errorf("status %q is not treated as finalised", state.Status) + } + }) + } +} + +func TestDoneReleasesTheCancelFunc(t *testing.T) { + q := newQueue(t, testConfig(1)) + id := seedRun(t, q, "work") + + q.done(doneEvent{runID: id, out: Outcome{Accepted: true}}) + + if _, ok := q.cancels[id]; ok { + t.Error("cancel func was not released, so the map grows for every run") + } +} + +func TestDoneIgnoresUnknownRuns(t *testing.T) { + q := newQueue(t, testConfig(1)) + + q.done(doneEvent{runID: "run-evicted", out: Outcome{Accepted: true}}) + + if len(q.runs) != 0 { + t.Errorf("runs = %+v, want a done for an unknown run to be dropped", q.runs) + } +} + +func TestEvictReapsOnlyFinishedRunsPastRetention(t *testing.T) { + cfg := testConfig(1) + cfg.Retention = utils.Duration(time.Minute) + q := newQueue(t, cfg) + + stale := seedRun(t, q, "finished long ago") + recent := seedRun(t, q, "just finished") + live := seedRun(t, q, "still running") + + q.done(doneEvent{runID: stale, out: Outcome{Accepted: true}}) + q.done(doneEvent{runID: recent, out: Outcome{Accepted: true}}) + + // Backdate rather than wait: retention is measured, not slept through. + old := time.Now().Add(-2 * time.Minute) + q.runs[stale].state.FinishedAt = &old + + q.evict() + + if _, ok := q.runs[stale]; ok { + t.Error("a run finished past the retention window was not evicted") + } + if _, ok := q.runs[recent]; !ok { + t.Error("a recently finished run was evicted too early") + } + if _, ok := q.runs[live]; !ok { + t.Error("an unfinished run was evicted") + } +} + +func TestActiveCountOnlyCountsRunningRuns(t *testing.T) { + q := newQueue(t, testConfig(3)) + + queued := seedRun(t, q, "queued") + running := seedRun(t, q, "running") + finished := seedRun(t, q, "finished") + + now := time.Now() + q.runs[running].state.StartedAt = &now + q.runs[running].state.Status = api.StatusCoding + + q.runs[finished].state.StartedAt = &now + q.done(doneEvent{runID: finished, out: Outcome{Accepted: true}}) + + if got := q.activeCount(); got != 1 { + t.Fatalf("activeCount = %d, want 1 (queued=%s, finished excluded)", got, queued) + } +} + +func TestSnapshotIsADeepCopy(t *testing.T) { + q := newQueue(t, testConfig(1)) + id := seedRun(t, q, "work") + + original := time.Now() + startedAt := original + q.runs[id].state.StartedAt = &startedAt + q.runs[id].state.Status = api.StatusCoding + + snap := q.snapshot() + if len(snap.Runs) != 1 { + t.Fatalf("snapshot has %d runs, want 1", len(snap.Runs)) + } + + q.runs[id].state.Status = api.StatusSucceeded + *q.runs[id].state.StartedAt = original.Add(time.Hour) + + if got := snap.Runs[0].Status; got != api.StatusCoding { + t.Errorf("snapshot status = %q, want it frozen at %q", got, api.StatusCoding) + } + if !snap.Runs[0].StartedAt.Equal(original) { + t.Error("snapshot shares its StartedAt pointer with the live run") + } +} + +// Loop + +func TestRunPublishesAnInitialSnapshot(t *testing.T) { + q := newRunningQueue(t, testConfig(1), nil) + + select { + case snap := <-q.Snapshots(): + if snap.At.IsZero() { + t.Error("published a snapshot with no timestamp") + } + case <-time.After(waitLimit): + t.Fatal("no snapshot published: the scheduling loop is not running") + } +} + +func TestScheduleFillsEveryFreeSlot(t *testing.T) { + runner := newBlockingRunner(3) + t.Cleanup(func() { close(runner.release) }) + + q := newRunningQueue(t, testConfig(3), runner.run) + + want := map[string]bool{} + for range 3 { + id, err := q.Add("/repo", mustTask(t, "do the thing")) + if err != nil { + t.Fatalf("add run: %v", err) + } + want[id] = true + } + + got := map[string]bool{} + for range 3 { + got[runner.awaitStart(t)] = true + } + + for id := range want { + if !got[id] { + t.Errorf("run %s was never admitted", id) + } + } +} + +func TestScheduleRespectsMaxContainers(t *testing.T) { + runner := newBlockingRunner(2) + t.Cleanup(func() { close(runner.release) }) + + q := newRunningQueue(t, testConfig(1), runner.run) + + first, err := q.Add("/repo", mustTask(t, "occupies the only slot")) + if err != nil { + t.Fatalf("add run: %v", err) + } + if _, err := q.Add("/repo", mustTask(t, "must wait")); err != nil { + t.Fatalf("add run: %v", err) + } + + if got := runner.awaitStart(t); got != first { + t.Fatalf("admitted %s first, want %s", got, first) + } + runner.assertNoStart(t, 20*testTick) +} + +func TestCancelBeforeAdmissionRetiresTheRun(t *testing.T) { + runner := newBlockingRunner(2) + t.Cleanup(func() { close(runner.release) }) + + q := newRunningQueue(t, testConfig(1), runner.run) + + first, err := q.Add("/repo", mustTask(t, "occupies the only slot")) + if err != nil { + t.Fatalf("add run: %v", err) + } + queued, err := q.Add("/repo", mustTask(t, "cancelled while waiting")) + if err != nil { + t.Fatalf("add run: %v", err) + } + + if got := runner.awaitStart(t); got != first { + t.Fatalf("admitted %s first, want %s", got, first) + } + + if err := q.Cancel(context.Background(), queued); err != nil { + t.Fatalf("cancel queued run: %v", err) + } + + run := awaitRun(t, q, queued, api.StatusCancelled) + if run.FinishedAt == nil { + t.Error("cancelled run has no FinishedAt, so evict will never reap it") + } + if run.StartedAt != nil { + t.Error("a run cancelled before admission must never have started") + } + runner.assertNoStart(t, 20*testTick) +} + +func TestRunReachesSucceeded(t *testing.T) { + outcome := Outcome{ + Accepted: true, + Branch: "patchdock/run-1", + Patch: auditlog.PatchStat{Files: 1, Additions: 4}, + } + + q := newRunningQueue(t, testConfig(1), func(context.Context, RunSpec, Reporter) (Outcome, error) { + return outcome, nil + }) + + id, err := q.Add("/repo", mustTask(t, "ship it")) + if err != nil { + t.Fatalf("add run: %v", err) + } + + run := awaitRun(t, q, id, api.StatusSucceeded) + if run.Branch != outcome.Branch { + t.Errorf("branch = %q, want %q", run.Branch, outcome.Branch) + } + if run.Patch == nil || run.Patch.Files != 1 { + t.Errorf("patch = %+v, want the runner's stat", run.Patch) + } + if run.StartedAt == nil { + t.Error("a run that executed has no StartedAt") + } +} + +func TestRunReachesFailed(t *testing.T) { + q := newRunningQueue(t, testConfig(1), func(context.Context, RunSpec, Reporter) (Outcome, error) { + return Outcome{}, errors.New("planner stage: no such image") + }) + + id, err := q.Add("/repo", mustTask(t, "will fail")) + if err != nil { + t.Fatalf("add run: %v", err) + } + + run := awaitRun(t, q, id, api.StatusFailed) + if run.Summary != "planner stage: no such image" { + t.Errorf("summary = %q, want the runner's error", run.Summary) + } +} + +func TestAddCleansTheRepoPath(t *testing.T) { + runner := newBlockingRunner(1) + t.Cleanup(func() { close(runner.release) }) + + q := newRunningQueue(t, testConfig(1), runner.run) + + id, err := q.Add("/repo/nested/../", mustTask(t, "work")) + if err != nil { + t.Fatalf("add run: %v", err) + } + + runner.awaitStart(t) + + deadline := time.After(waitLimit) + for { + select { + case snap := <-q.Snapshots(): + for _, r := range snap.Runs { + if r.ID != id { + continue + } + if r.Repo != "/repo" { + t.Fatalf("repo = %q, want %q", r.Repo, "/repo") + } + return + } + case <-deadline: + t.Fatalf("run %s never appeared in a snapshot", id) + } + } +} + +func TestCancelUnknownRunThroughThePublicAPI(t *testing.T) { + q := newRunningQueue(t, testConfig(1), nil) + + if err := q.Cancel(context.Background(), "run-does-not-exist"); !errors.Is(err, ErrNotFound) { + t.Fatalf("Cancel(unknown) = %v, want %v", err, ErrNotFound) + } +} diff --git a/internal/daemon/queue/reporter.go b/internal/daemon/queue/reporter.go index fece00d..57d2ed4 100644 --- a/internal/daemon/queue/reporter.go +++ b/internal/daemon/queue/reporter.go @@ -18,18 +18,16 @@ type reporter struct { } func (r *reporter) StageChange(stage types.StageName, attempt int) { - event := stageEvent{ + r.queue.send(r.queue.ctx, stageEvent{ runID: r.runID, stage: stage, attempt: attempt, - } - - select { - case r.queue.inbox <- event: - case <-r.queue.ctx.Done(): - } + }) } +// StageActivity deliberately does not use send: activity text is best-effort +// telemetry for the dashboard, so a full inbox drops it rather than making the +// pipeline wait on the queue loop. func (r *reporter) StageActivity(activity string) { select { case r.queue.inbox <- activityEvent{runID: r.runID, text: activity}: @@ -43,8 +41,5 @@ func (r *reporter) StageSummary(summary string) { return } - select { - case r.queue.inbox <- summaryEvent{runID: r.runID, text: summary}: - case <-r.queue.ctx.Done(): - } + r.queue.send(r.queue.ctx, summaryEvent{runID: r.runID, text: summary}) } diff --git a/internal/daemon/router_test.go b/internal/daemon/router_test.go new file mode 100644 index 0000000..4dff7f5 --- /dev/null +++ b/internal/daemon/router_test.go @@ -0,0 +1,240 @@ +package daemon + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/HJyup/patchdock/internal/daemon/api" +) + +type stubService struct { + health api.HealthResponse + runResp api.RunResponse + runErr error + cancelErr error + + snapshots chan api.Snapshot + snapErrs chan error + + gotRepo string + gotPrompt string + gotRunID string +} + +func newFakeService() *stubService { + return &stubService{ + health: api.HealthResponse{Status: "ok", Uptime: "1m0s", PID: 4242}, + runResp: api.RunResponse{RunID: "run-abc123"}, + snapshots: make(chan api.Snapshot), + snapErrs: make(chan error, 1), + } +} + +func (f *stubService) Health(context.Context) api.HealthResponse { return f.health } + +func (f *stubService) Run(_ context.Context, repo, prompt string) (api.RunResponse, error) { + f.gotRepo, f.gotPrompt = repo, prompt + return f.runResp, f.runErr +} + +func (f *stubService) Cancel(_ context.Context, runID string) error { + f.gotRunID = runID + return f.cancelErr +} + +func (f *stubService) Snapshot(context.Context) (<-chan api.Snapshot, <-chan error) { + return f.snapshots, f.snapErrs +} + +func do(t *testing.T, h http.Handler, method, target string, body string) *httptest.ResponseRecorder { + t.Helper() + + var reader *strings.Reader + if body != "" { + reader = strings.NewReader(body) + } else { + reader = strings.NewReader("") + } + + rec := httptest.NewRecorder() + h.ServeHTTP(rec, httptest.NewRequest(method, target, reader)) + return rec +} + +func TestRunAcceptsAPrompt(t *testing.T) { + svc := newFakeService() + rec := do(t, NewRouter(svc), http.MethodPost, "/run", `{"repo":"/abs/repo","prompt":"fix the bug"}`) + + if rec.Code != http.StatusCreated { + t.Fatalf("status = %d, want 201: %s", rec.Code, rec.Body) + } + + var got api.RunResponse + if err := json.Unmarshal(rec.Body.Bytes(), &got); err != nil { + t.Fatalf("decode body %q: %v", rec.Body.String(), err) + } + if got.RunID != "run-abc123" { + t.Errorf("run_id = %q, want run-abc123", got.RunID) + } + + if svc.gotRepo != "/abs/repo" || svc.gotPrompt != "fix the bug" { + t.Errorf("service received repo=%q prompt=%q, want the decoded payload", svc.gotRepo, svc.gotPrompt) + } +} + +func TestRunStatusCodes(t *testing.T) { + tests := []struct { + name string + body string + runErr error + wantCode int + }{ + { + name: "malformed body", + body: `{"repo":`, + wantCode: http.StatusBadRequest, + }, + { + name: "not json at all", + body: `hello`, + wantCode: http.StatusBadRequest, + }, + { + name: "service rejects the payload", + body: `{"repo":"relative/path","prompt":"x"}`, + runErr: ErrInvalidUserPayload, + wantCode: http.StatusBadRequest, + }, + { + name: "service fails internally", + body: `{"repo":"/abs/repo","prompt":"x"}`, + runErr: errors.New("queue is gone"), + wantCode: http.StatusInternalServerError, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + svc := newFakeService() + svc.runErr = tt.runErr + + rec := do(t, NewRouter(svc), http.MethodPost, "/run", tt.body) + if rec.Code != tt.wantCode { + t.Errorf("status = %d, want %d: %s", rec.Code, tt.wantCode, rec.Body) + } + }) + } +} + +func TestCancelWithoutARunIDIsNotRouted(t *testing.T) { + svc := newFakeService() + + rec := do(t, NewRouter(svc), http.MethodDelete, "/run/", "") + + if rec.Code != http.StatusNotFound { + t.Errorf("status = %d, want 404", rec.Code) + } +} + +// Snapshot stream + +func TestStreamEmitsSnapshotEvents(t *testing.T) { + svc := newFakeService() + rt := NewRouter(svc) + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/run", nil) + + done := make(chan struct{}) + go func() { + defer close(done) + rt.ServeHTTP(rec, req) + }() + + svc.snapshots <- api.Snapshot{ + At: time.Now(), + Runs: []api.Run{{ID: "run-1", Status: api.StatusCoding}}, + } + close(svc.snapshots) + + select { + case <-done: + case <-time.After(3 * time.Second): + t.Fatal("stream handler never returned after its source closed") + } + + if rec.Code != http.StatusOK { + t.Errorf("status = %d, want 200", rec.Code) + } + if got := rec.Header().Get("Content-Type"); got != "text/event-stream" { + t.Errorf("content-type = %q, want text/event-stream", got) + } + + body := rec.Body.String() + if !strings.Contains(body, "event: "+api.EventSnapshot) { + t.Errorf("body has no snapshot event:\n%s", body) + } + + if !strings.Contains(body, `"id":"run-1"`) { + t.Errorf("body does not carry the run payload:\n%s", body) + } +} + +func TestStreamEmitsErrorEvents(t *testing.T) { + svc := newFakeService() + rt := NewRouter(svc) + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/run", nil) + + done := make(chan struct{}) + go func() { + defer close(done) + rt.ServeHTTP(rec, req) + }() + + svc.snapErrs <- errors.New("follow broker: the broker has been closed") + + select { + case <-done: + case <-time.After(3 * time.Second): + t.Fatal("stream handler did not return after reporting an error") + } + + body := rec.Body.String() + if !strings.Contains(body, "event: "+api.EventError) { + t.Errorf("body has no error event:\n%s", body) + } + + if !strings.Contains(body, "the broker has been closed") { + t.Errorf("error event does not carry the message:\n%s", body) + } +} + +func TestStreamStopsWhenTheClientDisconnects(t *testing.T) { + svc := newFakeService() + rt := NewRouter(svc) + + ctx, cancel := context.WithCancel(context.Background()) + req := httptest.NewRequest(http.MethodGet, "/run", nil).WithContext(ctx) + + done := make(chan struct{}) + go func() { + defer close(done) + rt.ServeHTTP(httptest.NewRecorder(), req) + }() + + cancel() + + select { + case <-done: + case <-time.After(3 * time.Second): + t.Fatal("stream handler outlived its request context") + } +} diff --git a/internal/daemon/service.go b/internal/daemon/service.go index dbef5b8..a94a9bd 100644 --- a/internal/daemon/service.go +++ b/internal/daemon/service.go @@ -50,12 +50,12 @@ func (s *Service) Run(ctx context.Context, repo string, prompt string) (api.RunR task, err := types.NewTask(types.Task{Description: prompt}) if err != nil { - return api.RunResponse{}, errors.New("failed to create a task") + return api.RunResponse{}, fmt.Errorf("%w: %w", ErrInvalidUserPayload, err) } id, err := s.queue.Add(repo, task) if err != nil { - return api.RunResponse{}, errors.New("failed to add task to the queue") + return api.RunResponse{}, fmt.Errorf("queue task: %w", err) } return api.RunResponse{RunID: id}, nil diff --git a/internal/daemon/service_test.go b/internal/daemon/service_test.go new file mode 100644 index 0000000..969c708 --- /dev/null +++ b/internal/daemon/service_test.go @@ -0,0 +1,198 @@ +package daemon + +import ( + "context" + "errors" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/HJyup/patchdock/internal/daemon/broker" + "github.com/HJyup/patchdock/internal/daemon/config" + "github.com/HJyup/patchdock/internal/daemon/queue" + "github.com/HJyup/patchdock/internal/runtimedir" + "github.com/HJyup/patchdock/internal/utils" +) + +func idleRunner(context.Context, queue.RunSpec, queue.Reporter) (queue.Outcome, error) { + return queue.Outcome{Accepted: true, Branch: "patchdock/run-1"}, nil +} + +func blockedRunner(ctx context.Context, _ queue.RunSpec, _ queue.Reporter) (queue.Outcome, error) { + <-ctx.Done() + return queue.Outcome{}, ctx.Err() +} + +func newTestService(t *testing.T, runner queue.Runner) *Service { + t.Helper() + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + cfg := config.Defaults() + cfg.SnapshotTick = utils.Duration(2 * time.Millisecond) + + q := queue.New(ctx, runner, &cfg) + go q.Run() + + br := broker.New(q.Snapshots()) + go br.Run(ctx) + + dir, err := runtimedir.Resolve(filepath.Join(t.TempDir(), "rt")) + if err != nil { + t.Fatalf("resolve runtime dir: %v", err) + } + + return NewService(q, dir, br) +} + +func TestHealthReportsThisProcess(t *testing.T) { + svc := newTestService(t, idleRunner) + + got := svc.Health(context.Background()) + + if got.Status != "ok" { + t.Errorf("status = %q, want ok", got.Status) + } + + if got.PID != os.Getpid() { + t.Errorf("pid = %d, want %d", got.PID, os.Getpid()) + } + + if _, err := time.ParseDuration(got.Uptime); err != nil { + t.Errorf("uptime %q is not a duration: %v", got.Uptime, err) + } +} + +func TestRunQueuesTheTask(t *testing.T) { + svc := newTestService(t, blockedRunner) + + resp, err := svc.Run(context.Background(), "/abs/repo", "fix the bug") + if err != nil { + t.Fatalf("run: %v", err) + } + + // Kinda overlaying on the how we represent IDs but used because had some mixed things with ids + if !strings.HasPrefix(resp.RunID, "run-") { + t.Errorf("run id = %q, want a run- prefix", resp.RunID) + } +} + +func TestRunRejectsARelativeRepoPath(t *testing.T) { + svc := newTestService(t, idleRunner) + + _, err := svc.Run(context.Background(), "relative/repo", "fix the bug") + if !errors.Is(err, ErrInvalidUserPayload) { + t.Fatalf("error = %v, want it to wrap %v", err, ErrInvalidUserPayload) + } +} + +func TestRunRejectsAnEmptyPrompt(t *testing.T) { + svc := newTestService(t, idleRunner) + + _, err := svc.Run(context.Background(), "/abs/repo", "") + if !errors.Is(err, ErrInvalidUserPayload) { + t.Fatalf("error = %v, want it to wrap %v", err, ErrInvalidUserPayload) + } +} + +func TestCancelQueuedRun(t *testing.T) { + svc := newTestService(t, blockedRunner) + + resp, err := svc.Run(context.Background(), "/abs/repo", "cancel me") + if err != nil { + t.Fatalf("run: %v", err) + } + + if err := svc.Cancel(context.Background(), resp.RunID); err != nil { + t.Fatalf("cancel: %v", err) + } +} + +// Stream + +func TestSnapshotStreamsQueueState(t *testing.T) { + svc := newTestService(t, blockedRunner) + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + data, errs := svc.Snapshot(ctx) + resp, err := svc.Run(context.Background(), "/abs/repo", "watch me") + if err != nil { + t.Fatalf("run: %v", err) + } + + deadline := time.After(3 * time.Second) + for { + select { + case snap, ok := <-data: + if !ok { + t.Fatal("snapshot stream closed before the run appeared") + } + for _, run := range snap.Runs { + if run.ID == resp.RunID { + return + } + } + case err := <-errs: + t.Fatalf("stream error: %v", err) + case <-deadline: + t.Fatalf("run %s never appeared in a snapshot", resp.RunID) + } + } +} + +// Every SSE request cancels its context on disconnect. If Snapshot leaked its +// goroutine or its broker subscription, a busy dashboard would accumulate both. +func TestSnapshotStopsAndClosesOnContextCancel(t *testing.T) { + svc := newTestService(t, idleRunner) + + ctx, cancel := context.WithCancel(context.Background()) + data, errs := svc.Snapshot(ctx) + + cancel() + + deadline := time.After(3 * time.Second) + for data != nil || errs != nil { + select { + case _, ok := <-data: + if !ok { + data = nil + } + case _, ok := <-errs: + if !ok { + errs = nil + } + case <-deadline: + t.Fatal("Snapshot did not close its channels after its context was cancelled") + } + } +} + +func TestSnapshotReportsAClosedBroker(t *testing.T) { + svc := newTestService(t, idleRunner) + svc.br.Close() + + data, errs := svc.Snapshot(context.Background()) + + select { + case err := <-errs: + if !errors.Is(err, broker.ErrClosed) { + t.Fatalf("error = %v, want it to wrap %v", err, broker.ErrClosed) + } + case <-time.After(3 * time.Second): + t.Fatal("no error reported for a closed broker") + } + + select { + case _, ok := <-data: + if ok { + t.Error("a snapshot was delivered despite the broker being closed") + } + case <-time.After(3 * time.Second): + t.Fatal("the data channel was never closed") + } +} diff --git a/internal/runtimedir/runtimedir_test.go b/internal/runtimedir/runtimedir_test.go new file mode 100644 index 0000000..a57f86e --- /dev/null +++ b/internal/runtimedir/runtimedir_test.go @@ -0,0 +1,96 @@ +package runtimedir + +import ( + "os" + "path/filepath" + "testing" +) + +func mustResolve(t *testing.T, root string) Dir { + t.Helper() + + dir, err := Resolve(root) + if err != nil { + t.Fatalf("resolve %s: %v", root, err) + } + return dir +} + +func assertPerm(t *testing.T, path string, want os.FileMode) { + t.Helper() + + info, err := os.Stat(path) + if err != nil { + t.Fatalf("stat %s: %v", path, err) + } + if got := info.Mode().Perm(); got != want { + t.Errorf("%s has mode %04o, want %04o", path, got, want) + } +} + +func TestResolveCreatesThePrivateDirectory(t *testing.T) { + root := filepath.Join(t.TempDir(), "patchdock") + + dir := mustResolve(t, root) + + if dir.Root() != root { + t.Errorf("Root() = %q, want %q", dir.Root(), root) + } + assertPerm(t, root, 0o700) +} + +func TestResolveIsIdempotent(t *testing.T) { + root := filepath.Join(t.TempDir(), "patchdock") + + first := mustResolve(t, root) + second := mustResolve(t, root) + + if first != second { + t.Errorf("second resolve = %+v, want %+v", second, first) + } + assertPerm(t, root, 0o700) +} + +func TestPathsLiveUnderRoot(t *testing.T) { + root := filepath.Join(t.TempDir(), "patchdock") + dir := mustResolve(t, root) + + paths := map[string]string{ + "Socket": dir.Socket(), + "Config": dir.Config(), + "Lock": dir.Lock(), + "Log": dir.Log(), + } + + seen := make(map[string]string, len(paths)) + for name, path := range paths { + if got := filepath.Dir(path); got != root { + t.Errorf("%s() = %q, want it inside %q", name, path, root) + } + if other, clash := seen[path]; clash { + t.Errorf("%s() and %s() both resolve to %q", name, other, path) + } + seen[path] = name + } +} + +func TestPathNames(t *testing.T) { + dir := mustResolve(t, filepath.Join(t.TempDir(), "patchdock")) + + tests := []struct { + name string + got string + want string + }{ + {"Socket", dir.Socket(), "dock.sock"}, + {"Config", dir.Config(), "config.json"}, + {"Lock", dir.Lock(), "dock.lock"}, + {"Log", dir.Log(), "dock.log"}, + } + + for _, tt := range tests { + if base := filepath.Base(tt.got); base != tt.want { + t.Errorf("%s() base = %q, want %q", tt.name, base, tt.want) + } + } +} diff --git a/internal/utils/selects.go b/internal/utils/selects.go new file mode 100644 index 0000000..322584b --- /dev/null +++ b/internal/utils/selects.go @@ -0,0 +1,17 @@ +package utils + +// SendLatest replaces whatever is buffered in ch with v and reports whether v +// was delivered. It never blocks (should use capacity of one for channels) +func SendLatest[T any](ch chan T, v T) bool { + select { + case <-ch: + default: + } + + select { + case ch <- v: + return true + default: + return false + } +}