package distworker import ( "encoding/json" "net/http" "net/http/httptest" "strconv" "sync" "sync/atomic" "testing" "time" "github.com/Jeffail/tunny" "github.com/gorilla/websocket" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "rocketgit.ru/rsmon/worker/internal/wire" ) // newTestRunner builds a runner with a deterministic executor and the // fixed dispatcher/pool layout. It registers a cleanup hook that drains // the runner so tests can leak-free. func newTestRunner(t *testing.T, maxConc, poolSize int, fn func(payload interface{}) interface{}) *Runner { t.Helper() r := NewRunner(&Config{MaxConcurrency: maxConc}) r.executor = fn require.NotNil(t, r.executor) atomic.StoreInt64(&r.concurrency, int64(poolSize)) r.jobQueue = make(chan wire.CheckJob, r.queueCapacity()) r.results = make(chan resultEnvelope, r.queueCapacity()) r.pool = tunny.NewFunc(poolSize, r.executor) for i := 0; i < maxConc; i++ { r.wg.Add(1) go r.dispatcher() } t.Cleanup(func() { r.Stop() r.wg.Wait() r.pool.Close() close(r.results) }) return r } func TestQueueCapacityScalesWithMaxConcurrency(t *testing.T) { cases := []struct { maxConc int wantCapacity int }{ // 2*maxConc, floored at minQueueCapacity so a small pool still // has backpressure headroom. {maxConc: 1, wantCapacity: minQueueCapacity}, {maxConc: 4, wantCapacity: minQueueCapacity}, {maxConc: 16, wantCapacity: 2 * 16}, {maxConc: 64, wantCapacity: 2 * 64}, } for _, tc := range cases { t.Run("max="+strconv.Itoa(tc.maxConc), func(t *testing.T) { r := NewRunner(&Config{MaxConcurrency: tc.maxConc}) assert.Equal(t, tc.wantCapacity, r.QueueCapacity(), "queue capacity should scale with max concurrency") }) } } func TestDispatcherRunsJobsConcurrently(t *testing.T) { const ( maxConc = 8 poolSize = 4 jobCount = 12 hold = 80 * time.Millisecond ) var ( inFlight atomic.Int64 peak atomic.Int64 ) executor := func(payload interface{}) interface{} { cur := inFlight.Add(1) for { p := peak.Load() if cur <= p || peak.CompareAndSwap(p, cur) { break } } time.Sleep(hold) inFlight.Add(-1) job := payload.(wire.CheckJob) return []wire.CheckResultReport{{ JobID: job.JobID, CheckID: job.CheckID, State: "OK", }} } r := newTestRunner(t, maxConc, poolSize, executor) // Drain the results channel so dispatchers do not block. var drainWG sync.WaitGroup drainWG.Add(1) go func() { defer drainWG.Done() for i := 0; i < jobCount; i++ { select { case <-r.results: case <-r.stopCh: return } } }() for i := 0; i < jobCount; i++ { job := wire.CheckJob{ JobID: "job-" + strconv.Itoa(i), CheckID: int64(i + 1), Kind: "http", Host: "example.com", } require.True(t, r.Enqueue(job)) } drainWG.Wait() // The tunny.Pool size (poolSize) limits how many jobs run in // parallel, so the peak should be at most poolSize and at least 2 // (otherwise the test would pass on a serial pool). observed := peak.Load() assert.GreaterOrEqual(t, observed, int64(2), "expected concurrent execution, observed peak=%d", observed) assert.LessOrEqual(t, observed, int64(poolSize), "peak should be bounded by pool size, observed peak=%d", observed) } func TestEnqueueRespectsBackpressure(t *testing.T) { // Use a small queue with no dispatchers consuming it, so the bounded // channel is the only source of backpressure. This is the cleanest // way to assert that Enqueue parks when the buffer is full. const queueCap = 4 r := NewRunner(&Config{MaxConcurrency: 2}) r.jobQueue = make(chan wire.CheckJob, queueCap) r.results = make(chan resultEnvelope, queueCap) t.Cleanup(func() { r.Stop() close(r.results) }) // Fill the bounded queue. for i := 0; i < queueCap; i++ { require.True(t, r.Enqueue(wire.CheckJob{JobID: "prefill-" + strconv.Itoa(i)})) } assert.Equal(t, queueCap, r.QueueDepth(), "queue should be full after %d enqueues", queueCap) // The next Enqueue must block because the queue is full. enqueueDone := make(chan bool, 1) go func() { enqueueDone <- r.Enqueue(wire.CheckJob{JobID: "blocking"}) }() select { case got := <-enqueueDone: t.Fatalf("Enqueue returned %v while the queue was full; expected backpressure", got) case <-time.After(50 * time.Millisecond): // expected: still parked } // Free a slot and confirm the parked Enqueue unblocks. select { case <-r.jobQueue: case <-r.stopCh: t.Fatal("runner stopped unexpectedly") } select { case ok := <-enqueueDone: assert.True(t, ok, "Enqueue should succeed once a slot is free") case <-time.After(time.Second): t.Fatal("Enqueue did not unblock after slot was freed") } } func TestStopRejectsTerminalResultsWithoutClosingResultChannel(t *testing.T) { r := NewRunner(&Config{MaxConcurrency: 1}) r.results = make(chan resultEnvelope, 1) r.Stop() assert.False(t, r.enqueueFailedCheck(wire.CheckJob{JobID: "job-1", LeaseToken: "lease-1"}, "unsupported_kind: rkn")) assert.Empty(t, r.results) select { case r.results <- resultEnvelope{}: default: t.Fatal("results channel should remain open after Stop") } } func TestRotateTokenReconnectsWithoutStoppingRunner(t *testing.T) { const ( oldToken = "old-token" newToken = "new-token" ) connections := make(chan string, 2) results := make(chan wire.WorkerMessage, 1) upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }} server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { switch req.URL.Path { case "/api/internal/workers/rotate-token": if req.Header.Get("Authorization") != "Bearer "+oldToken { http.Error(w, "unexpected rotation token", http.StatusUnauthorized) return } _ = json.NewEncoder(w).Encode(struct { AuthToken string `json:"auth_token"` }{AuthToken: newToken}) case "/worker": conn, err := upgrader.Upgrade(w, req, nil) if err != nil { return } defer conn.Close() token := req.URL.Query().Get("token") connections <- token if token == oldToken { _, _, _ = conn.ReadMessage() // Rotation must close this connection. return } if token != newToken { return } if conn.WriteJSON(wire.WorkerMessage{Kind: "task", Task: &wire.CheckJob{ JobID: "after-rotation", LeaseToken: "lease", CheckID: 1, Kind: "http", }}) != nil { return } for { var message wire.WorkerMessage if err := conn.ReadJSON(&message); err != nil { return } if message.Kind == "result" && message.Result != nil && message.Result.JobID == "after-rotation" { results <- message return } } default: http.NotFound(w, req) } })) defer server.Close() r := NewRunner(&Config{URL: server.URL, Token: oldToken, MaxConcurrency: 1}) r.executor = func(payload interface{}) interface{} { job := payload.(wire.CheckJob) return []wire.CheckResultReport{{JobID: job.JobID, CheckID: job.CheckID, State: "OK"}} } startDone := make(chan error, 1) go func() { startDone <- r.Start() }() t.Cleanup(func() { r.Stop() select { case err := <-startDone: require.NoError(t, err) case <-time.After(time.Second): t.Fatal("runner did not stop") } }) select { case token := <-connections: require.Equal(t, oldToken, token) case <-time.After(time.Second): t.Fatal("worker did not establish its initial control connection") } gotToken, err := r.RotateToken(t.Context()) require.NoError(t, err) require.Equal(t, newToken, gotToken) require.Equal(t, newToken, r.Token()) select { case token := <-connections: require.Equal(t, newToken, token) case <-time.After(time.Second): t.Fatal("worker did not reconnect with the replacement token") } select { case result := <-results: require.NotNil(t, result.Result) assert.Equal(t, "OK", result.Result.State) case <-time.After(time.Second): t.Fatal("runner did not execute work after token rotation") } assert.False(t, r.stopped(), "rotation must not stop runner-owned subsystems") } func TestStopClosesAndJoinsIdleControlConnection(t *testing.T) { closed := make(chan struct{}, 1) connected := make(chan struct{}, 1) upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }} server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { if req.URL.Path != "/worker" { http.NotFound(w, req) return } conn, err := upgrader.Upgrade(w, req, nil) if err != nil { return } defer conn.Close() connected <- struct{}{} _, _, _ = conn.ReadMessage() closed <- struct{}{} })) defer server.Close() r := NewRunner(&Config{URL: server.URL, Token: "token", MaxConcurrency: 1}) startDone := make(chan error, 1) go func() { startDone <- r.Start() }() t.Cleanup(func() { r.Stop() }) select { case <-time.After(time.Second): t.Fatal("worker did not establish idle control connection") case <-connected: } r.Stop() select { case <-closed: case <-time.After(time.Second): t.Fatal("Stop did not close the idle control connection") } select { case err := <-startDone: require.NoError(t, err) case <-time.After(time.Second): t.Fatal("Stop did not join the control loop") } } func TestStopCancelsDialInProgress(t *testing.T) { dialStarted := make(chan struct{}) allowUpgrade := make(chan struct{}) upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }} server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { if req.URL.Path != "/worker" { http.NotFound(w, req) return } close(dialStarted) <-allowUpgrade _, _ = upgrader.Upgrade(w, req, nil) })) defer func() { close(allowUpgrade) server.Close() }() r := NewRunner(&Config{URL: server.URL, Token: "token", MaxConcurrency: 1}) startDone := make(chan error, 1) go func() { startDone <- r.Start() }() select { case <-dialStarted: case <-time.After(time.Second): t.Fatal("worker did not begin websocket dial") } r.Stop() select { case err := <-startDone: require.NoError(t, err) case <-time.After(time.Second): t.Fatal("Stop did not join a canceled websocket dial") } } func TestStartAndStopRegisterControlLoopSafely(t *testing.T) { server := httptest.NewServer(http.NotFoundHandler()) defer server.Close() for i := 0; i < 25; i++ { r := NewRunner(&Config{URL: server.URL, Token: "token", MaxConcurrency: 1}) startDone := make(chan error, 1) stopDone := make(chan struct{}) go func() { startDone <- r.Start() }() go func() { r.Stop() close(stopDone) }() select { case <-stopDone: case <-time.After(time.Second): t.Fatal("Stop did not complete") } select { case <-startDone: case <-time.After(time.Second): t.Fatal("Start did not return after concurrent Stop") } } } func TestRotateTokenPreservesDequeuedResult(t *testing.T) { const ( oldToken = "old-token" newToken = "new-token" ) connected := make(chan string, 2) delivered := make(chan wire.WorkerMessage, 1) upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }} server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { switch req.URL.Path { case "/api/internal/workers/rotate-token": _ = json.NewEncoder(w).Encode(struct { AuthToken string `json:"auth_token"` }{AuthToken: newToken}) case "/worker": conn, err := upgrader.Upgrade(w, req, nil) if err != nil { return } defer conn.Close() token := req.URL.Query().Get("token") connected <- token if token == oldToken { var message wire.WorkerMessage if conn.ReadJSON(&message) == nil { delivered <- message } return } if token != newToken { return } var message wire.WorkerMessage if conn.ReadJSON(&message) == nil { delivered <- message } } })) defer server.Close() r := NewRunner(&Config{URL: server.URL, Token: oldToken, MaxConcurrency: 1}) enteredWrite := make(chan struct{}) releaseWrite := make(chan struct{}) var once sync.Once r.beforeControlWrite = func() { once.Do(func() { close(enteredWrite) <-releaseWrite }) } startDone := make(chan error, 1) go func() { startDone <- r.Start() }() t.Cleanup(func() { r.Stop() select { case <-startDone: case <-time.After(time.Second): t.Fatal("runner did not stop") } }) select { case token := <-connected: require.Equal(t, oldToken, token) case <-time.After(time.Second): t.Fatal("worker did not establish its initial control connection") } r.results <- resultEnvelope{job: wire.CheckJob{JobID: "result", LeaseToken: "lease"}, reports: []wire.CheckResultReport{{JobID: "result", State: "OK"}}} select { case <-enteredWrite: case <-time.After(time.Second): t.Fatal("writer did not dequeue result") } rotated := make(chan error, 1) go func() { _, err := r.RotateToken(t.Context()) rotated <- err }() close(releaseWrite) require.NoError(t, <-rotated) select { case message := <-delivered: require.NotNil(t, message.Result) assert.Equal(t, "result", message.Result.JobID) assert.Equal(t, "lease", message.Result.LeaseToken) case <-time.After(time.Second): t.Fatal("dequeued result was lost during rotation") } } func TestRotateTokenSerializesWithStop(t *testing.T) { rotationStarted := make(chan struct{}) releaseHandler := make(chan struct{}) upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }} server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { switch req.URL.Path { case "/api/internal/workers/rotate-token": close(rotationStarted) <-releaseHandler case "/worker": conn, err := upgrader.Upgrade(w, req, nil) if err == nil { defer conn.Close() _, _, _ = conn.ReadMessage() } } })) defer func() { close(releaseHandler) server.Close() }() r := NewRunner(&Config{URL: server.URL, Token: "old-token", MaxConcurrency: 1}) startDone := make(chan error, 1) go func() { startDone <- r.Start() }() t.Cleanup(func() { r.Stop() }) // Wait for Start to install its client before beginning rotation. deadline := time.After(time.Second) for { r.clientMu.Lock() started := r.client != nil r.clientMu.Unlock() if started { break } select { case <-deadline: t.Fatal("runner did not start") default: time.Sleep(time.Millisecond) } } rotated := make(chan error, 1) go func() { _, err := r.RotateToken(t.Context()) rotated <- err }() select { case <-rotationStarted: case <-time.After(time.Second): t.Fatal("rotation request did not start") } stopped := make(chan struct{}) go func() { r.Stop() close(stopped) }() require.Error(t, <-rotated, "Stop must cancel an in-flight rotation request") select { case <-stopped: case <-time.After(time.Second): t.Fatal("Stop did not complete after canceling rotation") } select { case err := <-startDone: require.NoError(t, err) case <-time.After(time.Second): t.Fatal("runner did not stop") } _, err := r.RotateToken(t.Context()) require.Error(t, err, "rotation cannot succeed after shutdown") } func TestStopUnblocksRotationWaitingForWriter(t *testing.T) { upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }} connected := make(chan struct{}, 1) server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { switch req.URL.Path { case "/api/internal/workers/rotate-token": _ = json.NewEncoder(w).Encode(struct { AuthToken string `json:"auth_token"` }{AuthToken: "new-token"}) case "/worker": conn, err := upgrader.Upgrade(w, req, nil) if err != nil { return } defer conn.Close() connected <- struct{}{} _, _, _ = conn.ReadMessage() } })) defer server.Close() r := NewRunner(&Config{URL: server.URL, Token: "old-token", MaxConcurrency: 1}) writeBlocked := make(chan struct{}) var once sync.Once r.beforeControlWrite = func() { once.Do(func() { close(writeBlocked) <-r.controlCtx.Done() }) } rotationReady := make(chan struct{}) r.beforeTokenCommit = func() { close(rotationReady) } startDone := make(chan error, 1) go func() { startDone <- r.Start() }() t.Cleanup(func() { r.Stop() }) select { case <-connected: case <-time.After(time.Second): t.Fatal("worker did not connect") } r.results <- resultEnvelope{job: wire.CheckJob{JobID: "blocked", LeaseToken: "lease"}, reports: []wire.CheckResultReport{{JobID: "blocked", State: "OK"}}} select { case <-writeBlocked: case <-time.After(time.Second): t.Fatal("writer did not block") } rotated := make(chan error, 1) go func() { _, err := r.RotateToken(t.Context()) rotated <- err }() select { case <-rotationReady: case <-time.After(time.Second): t.Fatal("rotation did not reach writer serialization") } stopped := make(chan struct{}) go func() { r.Stop() close(stopped) }() select { case <-stopped: case <-time.After(time.Second): t.Fatal("Stop deadlocked behind rotation waiting for writer") } select { case <-rotated: case <-time.After(time.Second): t.Fatal("rotation did not unblock after Stop closed the connection") } select { case err := <-startDone: require.NoError(t, err) case <-time.After(time.Second): t.Fatal("runner did not stop") } } func TestStopWinsBeforePostHTTPRotationCommit(t *testing.T) { commitReady := make(chan struct{}) releaseCommit := make(chan struct{}) server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { if req.URL.Path != "/api/internal/workers/rotate-token" { http.NotFound(w, req) return } _ = json.NewEncoder(w).Encode(struct { AuthToken string `json:"auth_token"` }{AuthToken: "new-token"}) })) defer server.Close() r := NewRunner(&Config{URL: server.URL, Token: "old-token", MaxConcurrency: 1}) r.client = NewClient(server.URL, "old-token") r.beforeTokenCommit = func() { close(commitReady) <-releaseCommit } rotated := make(chan error, 1) go func() { _, err := r.RotateToken(t.Context()) rotated <- err }() select { case <-commitReady: case <-time.After(time.Second): t.Fatal("rotation did not reach post-HTTP commit") } stopped := make(chan struct{}) go func() { r.Stop() close(stopped) }() select { case <-r.stopCh: case <-time.After(time.Second): t.Fatal("Stop did not win lifecycle ownership") } close(releaseCommit) require.Error(t, <-rotated, "rotation cannot succeed after Stop wins") select { case <-stopped: case <-time.After(time.Second): t.Fatal("Stop did not finish") } assert.Equal(t, "old-token", r.Token()) } func TestRotateTokenUnchangedResponseReleasesLifecycle(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { _ = json.NewEncoder(w).Encode(struct { AuthToken string `json:"auth_token"` }{AuthToken: "old-token"}) })) defer server.Close() r := NewRunner(&Config{URL: server.URL, Token: "old-token", MaxConcurrency: 1}) r.client = NewClient(server.URL, "old-token") _, err := r.RotateToken(t.Context()) require.Error(t, err) stopped := make(chan struct{}) go func() { r.Stop() close(stopped) }() select { case <-stopped: case <-time.After(time.Second): t.Fatal("Stop deadlocked after unchanged rotation response") } _, err = r.RotateToken(t.Context()) require.Error(t, err, "later rotation must observe shutdown") } func TestRotateTokenClosesStalledWriterBeforeWaiting(t *testing.T) { upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }} connected := make(chan struct{}, 1) connectionClosed := make(chan struct{}) resent := make(chan wire.WorkerMessage, 1) server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { switch req.URL.Path { case "/api/internal/workers/rotate-token": _ = json.NewEncoder(w).Encode(struct { AuthToken string `json:"auth_token"` }{AuthToken: "new-token"}) case "/worker": conn, err := upgrader.Upgrade(w, req, nil) if err != nil { return } defer conn.Close() if req.URL.Query().Get("token") == "old-token" { connected <- struct{}{} _, _, _ = conn.ReadMessage() connectionClosed <- struct{}{} return } var message wire.WorkerMessage if conn.ReadJSON(&message) == nil { resent <- message } } })) defer server.Close() r := NewRunner(&Config{URL: server.URL, Token: "old-token", MaxConcurrency: 1}) writerBlocked := make(chan struct{}) var once sync.Once r.beforeControlWrite = func() { once.Do(func() { close(writerBlocked) <-connectionClosed }) } startDone := make(chan error, 1) go func() { startDone <- r.Start() }() t.Cleanup(func() { r.Stop() }) select { case <-connected: case <-time.After(time.Second): t.Fatal("worker did not connect") } r.results <- resultEnvelope{job: wire.CheckJob{JobID: "stalled", LeaseToken: "lease"}, reports: []wire.CheckResultReport{{JobID: "stalled", State: "OK"}}} select { case <-writerBlocked: case <-time.After(time.Second): t.Fatal("writer did not stall") } rotated := make(chan error, 1) go func() { _, err := r.RotateToken(t.Context()) rotated <- err }() select { case err := <-rotated: require.NoError(t, err) case <-time.After(time.Second): t.Fatal("rotation waited for stalled writer before closing its connection") } select { case message := <-resent: require.NotNil(t, message.Result) assert.Equal(t, "stalled", message.Result.JobID) case <-time.After(time.Second): t.Fatal("failed check result was not requeued after rotation") } r.Stop() select { case err := <-startDone: require.NoError(t, err) case <-time.After(time.Second): t.Fatal("runner did not stop") } } func TestWriterDropsFailedServerMetricSnapshot(t *testing.T) { upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }} upgraded := make(chan struct{}) closeServer := make(chan struct{}) serverClosed := make(chan struct{}) server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { conn, err := upgrader.Upgrade(w, req, nil) if err != nil { return } defer conn.Close() close(upgraded) <-closeServer close(serverClosed) })) defer server.Close() conn, err := NewClient(server.URL, "token").WorkerSocket() require.NoError(t, err) defer conn.Close() select { case <-upgraded: case <-time.After(time.Second): t.Fatal("websocket did not connect") } r := NewRunner(&Config{}) r.metricResults = make(chan metricEnvelope, 1) enteredWrite := make(chan struct{}) releaseWrite := make(chan struct{}) r.beforeControlWrite = func() { close(enteredWrite) <-releaseWrite } done := make(chan struct{}) writerDone := make(chan struct{}) go func() { var writeMu sync.Mutex r.writer(conn, &writeMu, done, 1) close(writerDone) }() r.metricResults <- metricEnvelope{generation: 1, report: wire.ServerMetricReport{ServerID: 1}} select { case <-enteredWrite: case <-time.After(time.Second): t.Fatal("writer did not dequeue metric snapshot") } close(closeServer) select { case <-serverClosed: case <-time.After(time.Second): t.Fatal("server did not close websocket") } _ = conn.Close() close(releaseWrite) select { case <-writerDone: case <-time.After(time.Second): t.Fatal("writer did not return after failed metric write") } _, replayable := r.takeOutbox() assert.False(t, replayable, "failed metric snapshot must not enter result outbox") } func TestWriterRequeuesFailedNotificationResult(t *testing.T) { upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }} upgraded := make(chan struct{}) closeServer := make(chan struct{}) serverClosed := make(chan struct{}) server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { conn, err := upgrader.Upgrade(w, req, nil) if err != nil { return } defer conn.Close() close(upgraded) <-closeServer close(serverClosed) })) defer server.Close() conn, err := NewClient(server.URL, "token").WorkerSocket() require.NoError(t, err) defer conn.Close() select { case <-upgraded: case <-time.After(time.Second): t.Fatal("websocket did not connect") } r := NewRunner(&Config{}) r.notifyResults = make(chan notifyResultEnvelope, 1) enteredWrite := make(chan struct{}) releaseWrite := make(chan struct{}) r.beforeControlWrite = func() { close(enteredWrite) <-releaseWrite } done := make(chan struct{}) writerDone := make(chan struct{}) go func() { var writeMu sync.Mutex r.writer(conn, &writeMu, done, 0) close(writerDone) }() r.notifyResults <- notifyResultEnvelope{report: wire.NotificationResultReport{JobID: "notification", LeaseToken: "lease"}} select { case <-enteredWrite: case <-time.After(time.Second): t.Fatal("writer did not dequeue notification result") } close(closeServer) select { case <-serverClosed: case <-time.After(time.Second): t.Fatal("server did not close websocket") } _ = conn.Close() close(releaseWrite) select { case <-writerDone: case <-time.After(time.Second): t.Fatal("writer did not return after failed notification write") } message, replayable := r.takeOutbox() require.True(t, replayable, "failed notification result must enter result outbox") require.NotNil(t, message.NotificationResult) assert.Equal(t, "notification", message.NotificationResult.JobID) } func TestMetricGenerationRejectsDisconnectedAndSendsFreshMetric(t *testing.T) { upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }} firstClosed := make(chan struct{}) secondConnected := make(chan struct{}) received := make(chan wire.WorkerMessage, 1) var connections atomic.Int32 server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { conn, err := upgrader.Upgrade(w, req, nil) if err != nil { return } defer conn.Close() if connections.Add(1) == 1 { close(firstClosed) return } close(secondConnected) var message wire.WorkerMessage if conn.ReadJSON(&message) == nil { received <- message } })) defer server.Close() r := NewRunner(&Config{URL: server.URL, Token: "token", MaxConcurrency: 1}) startDone := make(chan error, 1) go func() { startDone <- r.Start() }() t.Cleanup(func() { r.Stop() }) select { case <-firstClosed: case <-time.After(time.Second): t.Fatal("worker did not establish initial websocket") } // Wait until the old connection has fully torn down, then try to enqueue a // metric in the old-drain/disconnected interleaving. deadline := time.After(time.Second) for r.metricGeneration.Load() != 0 { select { case <-deadline: t.Fatal("old control generation did not clear") default: time.Sleep(time.Millisecond) } } staleCount := 1 r.enqueueMetric(wire.ServerMetricReport{ServerID: 1, ProcessCount: &staleCount}) assert.Empty(t, r.metricResults, "disconnected metric must not enter bounded channel") select { case r.reconnectCh <- struct{}{}: default: } select { case <-secondConnected: case <-time.After(time.Second): t.Fatal("worker did not reconnect") } deadline = time.After(time.Second) for r.metricGeneration.Load() == 0 { select { case <-deadline: t.Fatal("new control generation did not install") default: time.Sleep(time.Millisecond) } } // This enqueue occurs immediately after the new connection installation. freshCount := 2 r.enqueueMetric(wire.ServerMetricReport{ServerID: 1, ProcessCount: &freshCount}) select { case message := <-received: require.NotNil(t, message.ServerMetric) assert.Equal(t, 2, *message.ServerMetric.ProcessCount) case <-time.After(time.Second): t.Fatal("fresh metric was not sent after reconnect") } } func TestApplyInitResizesPool(t *testing.T) { executor := func(payload interface{}) interface{} { return []wire.CheckResultReport{} } r := newTestRunner(t, 8, 1, executor) // Replace the pool with a proxy that records SetSize calls. The // proxy delegates Close back to the underlying pool so the test // runner cleanup only closes the pool once. original := r.pool proxy := newSizeProbe(original) r.pool = proxy r.applyInit(&wire.WorkerInit{Concurrency: 5, WorkerID: "w-1"}) assert.Equal(t, 5, r.Concurrency(), "applyInit should update pool size") require.NotEmpty(t, proxy.sizes, "expected pool.SetSize to be called") assert.Equal(t, 5, proxy.sizes[len(proxy.sizes)-1]) // Concurrency above maxConcurrency should be clamped. r.applyInit(&wire.WorkerInit{Concurrency: 999, WorkerID: "w-1"}) assert.Equal(t, r.MaxConcurrency(), r.Concurrency(), "applyInit should clamp concurrency to maxConcurrency") } // sizeProbePool wraps a jobPool and records SetSize calls. Close is // forwarded to the wrapped pool so cleanup happens exactly once. type sizeProbePool struct { jobPool sizes []int } func newSizeProbe(p jobPool) *sizeProbePool { return &sizeProbePool{jobPool: p} } func (s *sizeProbePool) SetSize(n int) { s.sizes = append(s.sizes, n) s.jobPool.SetSize(n) } func TestApplyInitIgnoresNonPositive(t *testing.T) { executor := func(payload interface{}) interface{} { return []wire.CheckResultReport{} } r := newTestRunner(t, 4, 1, executor) // newTestRunner mirrors Start()'s initial concurrency of 1. require.Equal(t, 1, r.Concurrency()) // Concurrency = 0 must not break the runner or change its size. r.applyInit(&wire.WorkerInit{Concurrency: 0}) assert.Equal(t, 1, r.Concurrency(), "non-positive concurrency should be ignored") // Concurrency = -5 should also be ignored. r.applyInit(&wire.WorkerInit{Concurrency: -5}) assert.Equal(t, 1, r.Concurrency()) } // TestApplyInitStoresCredentialsInMemory verifies that WorkerInit.Credentials // is stored on the Runner via applyInit and is retrievable through // Credentials(). It also confirms that a subsequent applyInit with nil // credentials replaces the previous value. func TestApplyInitStoresCredentialsInMemory(t *testing.T) { executor := func(payload interface{}) interface{} { return []wire.CheckResultReport{} } r := newTestRunner(t, 4, 1, executor) // Before any applyInit, Credentials() returns nil. assert.Nil(t, r.Credentials(), "Credentials() must return nil before any init") want := &wire.NotificationCredentials{ SMTP: []wire.SMTPCredential{{ ID: 11, Name: "primary", Server: "smtp.example.com", Port: 587, Login: "alerts@example.com", Password: "smtp-password-xyz", }}, Telegram: []wire.TelegramCredential{{ ID: 22, Name: "main-bot", Token: "bot-token-9876543210:ABCDEFG", }}, } r.applyInit(&wire.WorkerInit{ WorkerID: "w-1", Concurrency: 2, Credentials: want, }) got := r.Credentials() require.NotNil(t, got, "Credentials() must return non-nil after applyInit") require.Len(t, got.SMTP, 1) require.Len(t, got.Telegram, 1) assert.Equal(t, "primary", got.SMTP[0].Name) assert.Equal(t, "smtp-password-xyz", got.SMTP[0].Password) assert.Equal(t, "main-bot", got.Telegram[0].Name) assert.Equal(t, "bot-token-9876543210:ABCDEFG", got.Telegram[0].Token) // A subsequent applyInit with nil Credentials replaces the value. r.applyInit(&wire.WorkerInit{WorkerID: "w-1", Concurrency: 2}) assert.Nil(t, r.Credentials(), "Credentials() must return nil after applyInit with nil Credentials") } func TestNotificationHeartbeatCounters(t *testing.T) { r := NewRunner(&Config{MaxConcurrency: 1}) atomic.StoreInt64(&r.notifyDepth, 2) atomic.StoreInt64(&r.notifyActive, 3) assert.Equal(t, 5, r.ActiveNotifications()) assert.Equal(t, 2, r.NotificationQueueDepth()) } // TestApplyInitStoresURLInMemory verifies that the URL field added in // Task 2 is stored on the Runner via applyInit and exposed through // URL(). A subsequent applyInit with empty URL replaces the previous // value (consistent with the Credentials contract). func TestApplyInitStoresURLInMemory(t *testing.T) { executor := func(payload interface{}) interface{} { return []wire.CheckResultReport{} } r := newTestRunner(t, 4, 1, executor) // Before any applyInit, URL() returns empty. assert.Equal(t, "", r.URL(), "URL() must return empty before any init") r.applyInit(&wire.WorkerInit{ WorkerID: "w-1", Concurrency: 2, URL: "https://worker-eu.example.com", }) assert.Equal(t, "https://worker-eu.example.com", r.URL(), "URL() must return the value pushed by applyInit") // Empty URL in a subsequent applyInit must clear the stored value. r.applyInit(&wire.WorkerInit{WorkerID: "w-1", Concurrency: 2}) assert.Equal(t, "", r.URL(), "URL() must return empty after applyInit with empty URL") } // TestApplyInitPrefersPublicURLOverLegacyURL verifies the bounded-migration // precedence on the init/config frame: PublicURL (canonical) wins whenever // it is non-empty, and the legacy URL field remains the fallback for old // control planes. func TestApplyInitPrefersPublicURLOverLegacyURL(t *testing.T) { executor := func(payload interface{}) interface{} { return []wire.CheckResultReport{} } r := newTestRunner(t, 4, 1, executor) r.applyInit(&wire.WorkerInit{ WorkerID: "w-1", Concurrency: 2, PublicURL: "https://canonical.example.com", URL: "https://legacy.example.com", }) assert.Equal(t, "https://canonical.example.com", r.URL(), "PublicURL must win over the legacy URL field") r.applyInit(&wire.WorkerInit{ WorkerID: "w-1", Concurrency: 2, URL: "https://legacy.example.com", }) assert.Equal(t, "https://legacy.example.com", r.URL(), "legacy URL field must be used when PublicURL is empty") } // TestApplyInitInvalidAcceptedURLKeepsPrior verifies the safe-fallback // behavior when the control plane supplies an unusable advertised URL: // the previous accepted value is kept (never regressed to a garbage // endpoint), while a valid empty init still clears it. func TestApplyInitInvalidAcceptedURLKeepsPrior(t *testing.T) { executor := func(payload interface{}) interface{} { return []wire.CheckResultReport{} } r := newTestRunner(t, 4, 1, executor) r.applyInit(&wire.WorkerInit{WorkerID: "w-1", Concurrency: 2, PublicURL: "https://worker.example.com"}) assert.Equal(t, "https://worker.example.com", r.URL()) // Invalid value: keep the previous accepted URL. r.applyInit(&wire.WorkerInit{WorkerID: "w-1", Concurrency: 2, PublicURL: "https://:27401"}) assert.Equal(t, "https://worker.example.com", r.URL(), "invalid accepted URL must not replace the stored value") r.applyInit(&wire.WorkerInit{WorkerID: "w-1", Concurrency: 2, URL: "not a url"}) assert.Equal(t, "https://worker.example.com", r.URL(), "invalid legacy url field must not replace the stored value") // Valid empty init clears, as before. r.applyInit(&wire.WorkerInit{WorkerID: "w-1", Concurrency: 2}) assert.Equal(t, "", r.URL(), "valid empty init must clear the stored URL") }