From 11a92c429fcf22100922e29334c5fe6ca1f04dc4 Mon Sep 17 00:00:00 2001 From: hezhaohui Date: Wed, 29 Jul 2026 14:13:53 +0800 Subject: [PATCH] =?UTF-8?q?test(server):=20=E6=B7=BB=E5=8A=A0=20task=20log?= =?UTF-8?q?=20=E6=A8=A1=E5=9D=97=E5=8D=95=E5=85=83=E6=B5=8B=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 添加 handler 层 16 个纯函数单元测试(状态映射、ID 解析、排序、格式化等) - 添加 service 层 11 个 pub/sub 与缓存单元测试(订阅、取消、并发安全、TTL 过期等) - 共 27 个测试用例,覆盖 task log 核心逻辑 背景:task log 模块经历了大量修改但缺乏测试覆盖,需要确保日志模块可用 关联 commit:61eeb0c, dbb930d --- server/internal/handler/task_log_test.go | 506 ++++++++++++++++++ server/internal/service/delivery_task_test.go | 440 +++++++++++++++ 2 files changed, 946 insertions(+) create mode 100644 server/internal/handler/task_log_test.go create mode 100644 server/internal/service/delivery_task_test.go diff --git a/server/internal/handler/task_log_test.go b/server/internal/handler/task_log_test.go new file mode 100644 index 0000000..2417413 --- /dev/null +++ b/server/internal/handler/task_log_test.go @@ -0,0 +1,506 @@ +package handler + +import ( + "testing" + "time" + + "github.com/1024XEngineer/xinfra/server/internal/model" + "github.com/1024XEngineer/xinfra/server/internal/service" +) + +func TestWayneStatusToString(t *testing.T) { + for name, tc := range map[string]struct { + input int + want string + }{ + "success": {input: 1, want: model.TaskFinished}, + "failed": {input: 0, want: model.TaskExecutionFailed}, + "unknown": {input: -1, want: "unknown"}, + "large number": {input: 99, want: "unknown"}, + } { + t.Run(name, func(t *testing.T) { + got := wayneStatusToString(tc.input) + if got != tc.want { + t.Fatalf("wayneStatusToString(%d) = %q, want %q", tc.input, got, tc.want) + } + }) + } +} + +func TestParseWaynePublishTaskID(t *testing.T) { + for name, tc := range map[string]struct { + input string + wantResID int64 + wantHisID int64 + wantOk bool + }{ + "valid": {input: "wayne:publish:123:456", wantResID: 123, wantHisID: 456, wantOk: true}, + "large ids": {input: "wayne:publish:999999999:888888888", wantResID: 999999999, wantHisID: 888888888, wantOk: true}, + "wrong prefix": {input: "awx:publish:123:456", wantOk: false}, + "missing parts": {input: "wayne:publish:123", wantOk: false}, + "extra parts": {input: "wayne:publish:123:456:789", wantOk: false}, + "non-numeric": {input: "wayne:publish:abc:456", wantOk: false}, + "non-numeric hist": {input: "wayne:publish:123:abc", wantOk: false}, + "empty": {input: "", wantOk: false}, + "random string": {input: "hello", wantOk: false}, + } { + t.Run(name, func(t *testing.T) { + resID, hisID, ok := parseWaynePublishTaskID(tc.input) + if ok != tc.wantOk { + t.Fatalf("parseWaynePublishTaskID(%q) ok=%v, want %v", tc.input, ok, tc.wantOk) + } + if ok && (resID != tc.wantResID || hisID != tc.wantHisID) { + t.Fatalf("parseWaynePublishTaskID(%q) = (%d, %d), want (%d, %d)", tc.input, resID, hisID, tc.wantResID, tc.wantHisID) + } + }) + } +} + +func TestWaynePublishTaskID(t *testing.T) { + history := service.WayneDeploymentHistory{ResourceID: 42, ID: 100} + got := waynePublishTaskID(history) + want := "wayne:publish:42:100" + if got != want { + t.Fatalf("waynePublishTaskID() = %q, want %q", got, want) + } +} + +func TestWayneDeploymentTaskName(t *testing.T) { + for name, tc := range map[string]struct { + history service.WayneDeploymentHistory + want string + }{ + "with name": { + history: service.WayneDeploymentHistory{ResourceName: "my-service", ResourceID: 10}, + want: "Wayne 服务部署 · my-service", + }, + "empty name uses ID": { + history: service.WayneDeploymentHistory{ResourceName: "", ResourceID: 42}, + want: "Wayne 服务部署 · 42", + }, + "whitespace name uses ID": { + history: service.WayneDeploymentHistory{ResourceName: " ", ResourceID: 7}, + want: "Wayne 服务部署 · 7", + }, + } { + t.Run(name, func(t *testing.T) { + got := wayneDeploymentTaskName(tc.history) + if got != tc.want { + t.Fatalf("wayneDeploymentTaskName() = %q, want %q", got, tc.want) + } + }) + } +} + +func TestTextForTaskStatus(t *testing.T) { + for name, tc := range map[string]struct { + status string + want string + }{ + "pending": {status: model.TaskPending, want: "等待"}, + "validating": {status: model.TaskValidating, want: "等待"}, + "dispatching": {status: model.TaskDispatching, want: "等待"}, + "running": {status: model.TaskRunning, want: "执行中"}, + "registering": {status: model.TaskRegistering, want: "执行中"}, + "canceling": {status: model.TaskCanceling, want: "执行中"}, + "rollback_pending": {status: model.TaskRollbackPending, want: "回退中"}, + "rolling_back": {status: model.TaskRollingBack, want: "回退中"}, + "finished": {status: model.TaskFinished, want: "成功"}, + "rolled_back": {status: model.TaskRolledBack, want: "已回退"}, + "rollback_failed": {status: model.TaskRollbackFailed, want: "回退失败"}, + "rollback_ack": {status: model.TaskRollbackAck, want: "已确认释放"}, + "register_failed": {status: model.TaskRegisterFailed, want: "注册失败(实例保留)"}, + "canceled": {status: model.TaskCanceled, want: "已取消"}, + "execution_failed": {status: model.TaskExecutionFailed, want: "失败"}, + "validation_failed": {status: model.TaskValidationFailed, want: "失败"}, + "unknown": {status: "unknown_status", want: "失败"}, + } { + t.Run(name, func(t *testing.T) { + got := textForTaskStatus(tc.status) + if got != tc.want { + t.Fatalf("textForTaskStatus(%q) = %q, want %q", tc.status, got, tc.want) + } + }) + } +} + +func TestClassForTaskStatus(t *testing.T) { + for name, tc := range map[string]struct { + status string + want string + }{ + "finished": {status: model.TaskFinished, want: "ok"}, + "rolled_back": {status: model.TaskRolledBack, want: "ok"}, + "execution_failed": {status: model.TaskExecutionFailed, want: "err"}, + "validation_failed": {status: model.TaskValidationFailed, want: "err"}, + "canceled": {status: model.TaskCanceled, want: "err"}, + "rollback_failed": {status: model.TaskRollbackFailed, want: "err"}, + "rollback_ack": {status: model.TaskRollbackAck, want: "warn"}, + "register_failed": {status: model.TaskRegisterFailed, want: "warn"}, + "running": {status: model.TaskRunning, want: "warn"}, + "dispatching": {status: model.TaskDispatching, want: "warn"}, + "pending": {status: model.TaskPending, want: ""}, + "unknown": {status: "unknown_status", want: ""}, + } { + t.Run(name, func(t *testing.T) { + got := classForTaskStatus(tc.status) + if got != tc.want { + t.Fatalf("classForTaskStatus(%q) = %q, want %q", tc.status, got, tc.want) + } + }) + } +} + +func TestTextForWaynePublishStatus(t *testing.T) { + for name, tc := range map[string]struct { + status int + want string + }{ + "success": {status: 1, want: "成功"}, + "failed": {status: 0, want: "失败"}, + "unknown": {status: -1, want: "未知"}, + "large": {status: 99, want: "未知"}, + } { + t.Run(name, func(t *testing.T) { + got := textForWaynePublishStatus(tc.status) + if got != tc.want { + t.Fatalf("textForWaynePublishStatus(%d) = %q, want %q", tc.status, got, tc.want) + } + }) + } +} + +func TestClassForWaynePublishStatus(t *testing.T) { + for name, tc := range map[string]struct { + status int + want string + }{ + "success": {status: 1, want: "ok"}, + "failed": {status: 0, want: "err"}, + "unknown": {status: -1, want: ""}, + } { + t.Run(name, func(t *testing.T) { + got := classForWaynePublishStatus(tc.status) + if got != tc.want { + t.Fatalf("classForWaynePublishStatus(%d) = %q, want %q", tc.status, got, tc.want) + } + }) + } +} + +func TestSplitStdoutLines(t *testing.T) { + for name, tc := range map[string]struct { + input string + want []taskLogLine + }{ + "single line": { + input: "hello world", + want: []taskLogLine{{Time: "", Message: "hello world", Class: ""}}, + }, + "multiple lines": { + input: "line1\nline2\nline3", + want: []taskLogLine{ + {Time: "", Message: "line1", Class: ""}, + {Time: "", Message: "line2", Class: ""}, + {Time: "", Message: "line3", Class: ""}, + }, + }, + "skip blank lines": { + input: "line1\n\n\nline2", + want: []taskLogLine{ + {Time: "", Message: "line1", Class: ""}, + {Time: "", Message: "line2", Class: ""}, + }, + }, + "trim cr": { + input: "line1\r\nline2\r\n", + want: []taskLogLine{ + {Time: "", Message: "line1", Class: ""}, + {Time: "", Message: "line2", Class: ""}, + }, + }, + "error line": { + input: "TASK FAILED: something went wrong", + want: []taskLogLine{{Time: "", Message: "TASK FAILED: something went wrong", Class: "err"}}, + }, + "ok line": { + input: "ok: [task 1] Apply role", + want: []taskLogLine{{Time: "", Message: "ok: [task 1] Apply role", Class: "ok"}}, + }, + "changed line": { + input: "changed: [host1] Task result changed", + want: []taskLogLine{{Time: "", Message: "changed: [host1] Task result changed", Class: "tag-ok"}}, + }, + "empty input": { + input: "", + want: []taskLogLine{}, + }, + "only whitespace": { + input: " \n \n ", + want: []taskLogLine{}, + }, + } { + t.Run(name, func(t *testing.T) { + got := splitStdoutLines(tc.input) + if len(got) != len(tc.want) { + t.Fatalf("splitStdoutLines(%q) returned %d lines, want %d", tc.input, len(got), len(tc.want)) + } + for i := range got { + if got[i] != tc.want[i] { + t.Fatalf("splitStdoutLines(%q)[%d] = %+v, want %+v", tc.input, i, got[i], tc.want[i]) + } + } + }) + } +} + +func TestClassForOutputLine(t *testing.T) { + for name, tc := range map[string]struct { + input string + want string + }{ + "failed keyword": {input: "TASK FAILED: error occurred", want: "err"}, + "fatal keyword": {input: "fatal: [host] unresolvable", want: "err"}, + "error keyword": {input: "ERROR: something bad", want: "err"}, + "error uppercase": {input: "ConnectionError: timeout", want: "err"}, + "ok keyword": {input: "ok: [host1] Apply task", want: "ok"}, + "successful": {input: "PLAY RECAP: successful", want: "ok"}, + "success keyword": {input: "task completed with success", want: "ok"}, + "changed keyword": {input: "changed: [host1] Executed task", want: "tag-ok"}, + "plain line": {input: "some random output", want: ""}, + "mixed case error": {input: "FAILED: task failed", want: "err"}, + } { + t.Run(name, func(t *testing.T) { + got := classForOutputLine(tc.input) + if got != tc.want { + t.Fatalf("classForOutputLine(%q) = %q, want %q", tc.input, got, tc.want) + } + }) + } +} + +func TestIsTerminalStatus(t *testing.T) { + terminalStatuses := []string{ + model.TaskFinished, + model.TaskCanceled, + model.TaskExecutionFailed, + model.TaskValidationFailed, + model.TaskRegisterFailed, + model.TaskRolledBack, + model.TaskRollbackFailed, + model.TaskRollbackAck, + } + for _, status := range terminalStatuses { + t.Run("terminal_"+status, func(t *testing.T) { + if !isTerminalStatus(status) { + t.Fatalf("isTerminalStatus(%q) = false, want true", status) + } + }) + } + + nonTerminalStatuses := []string{ + model.TaskPending, + model.TaskValidating, + model.TaskDispatching, + model.TaskRunning, + model.TaskRegistering, + model.TaskCanceling, + model.TaskRollbackPending, + model.TaskRollingBack, + "unknown", + "", + } + for _, status := range nonTerminalStatuses { + t.Run("non_terminal_"+status, func(t *testing.T) { + if isTerminalStatus(status) { + t.Fatalf("isTerminalStatus(%q) = true, want false", status) + } + }) + } +} + +func TestFormatTaskLogTime(t *testing.T) { + t.Run("zero time", func(t *testing.T) { + got := formatTaskLogTime(time.Time{}) + if got != "" { + t.Fatalf("formatTaskLogTime(zero) = %q, want empty", got) + } + }) + + t.Run("normal time", func(t *testing.T) { + ts := time.Date(2026, 7, 29, 14, 30, 45, 0, time.UTC) + got := formatTaskLogTime(ts) + if got != "14:30:45" { + t.Fatalf("formatTaskLogTime() = %q, want %q", got, "14:30:45") + } + }) + + t.Run("midnight", func(t *testing.T) { + ts := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) + got := formatTaskLogTime(ts) + if got != "00:00:00" { + t.Fatalf("formatTaskLogTime() = %q, want %q", got, "00:00:00") + } + }) +} + +func TestSortTaskLogSummaries(t *testing.T) { + t.Run("empty", func(t *testing.T) { + items := []taskLogSummary{} + sortTaskLogSummaries(items) + if len(items) != 0 { + t.Fatalf("sortTaskLogSummaries(empty) produced %d items", len(items)) + } + }) + + t.Run("single", func(t *testing.T) { + items := []taskLogSummary{{UpdatedAt: time.Now()}} + sortTaskLogSummaries(items) + if len(items) != 1 { + t.Fatalf("sortTaskLogSummaries(single) produced %d items", len(items)) + } + }) + + t.Run("sorts descending by UpdatedAt", func(t *testing.T) { + now := time.Now() + items := []taskLogSummary{ + {ID: "oldest", UpdatedAt: now.Add(-3 * time.Hour)}, + {ID: "newest", UpdatedAt: now}, + {ID: "middle", UpdatedAt: now.Add(-1 * time.Hour)}, + } + sortTaskLogSummaries(items) + if items[0].ID != "newest" || items[1].ID != "middle" || items[2].ID != "oldest" { + t.Fatalf("sortTaskLogSummaries order wrong: got IDs [%s, %s, %s]", items[0].ID, items[1].ID, items[2].ID) + } + }) + + t.Run("already sorted", func(t *testing.T) { + now := time.Now() + items := []taskLogSummary{ + {ID: "a", UpdatedAt: now.Add(-2 * time.Hour)}, + {ID: "b", UpdatedAt: now.Add(-1 * time.Hour)}, + {ID: "c", UpdatedAt: now}, + } + sortTaskLogSummaries(items) + if items[0].ID != "c" || items[1].ID != "b" || items[2].ID != "a" { + t.Fatalf("sortTaskLogSummaries order wrong: got IDs [%s, %s, %s]", items[0].ID, items[1].ID, items[2].ID) + } + }) +} + +func TestAwxTaskSummary(t *testing.T) { + now := time.Now() + task := model.DeliveryTask{ + ID: "task-123", + BusinessLineID: 5, + Status: model.TaskFinished, + InstanceName: "mysql-01", + TargetID: 10, + CreatedAt: now.Add(-1 * time.Hour), + UpdatedAt: now, + } + summary := awxTaskSummary(task) + + if summary.ID != "awx:task-123" { + t.Fatalf("awxTaskSummary ID = %q, want %q", summary.ID, "awx:task-123") + } + if summary.Source != "awx" { + t.Fatalf("awxTaskSummary Source = %q, want %q", summary.Source, "awx") + } + if summary.Service != "mysql" { + t.Fatalf("awxTaskSummary Service = %q, want %q", summary.Service, "mysql") + } + if summary.Name != "MySQL 标准化交付 · mysql-01" { + t.Fatalf("awxTaskSummary Name = %q, want %q", summary.Name, "MySQL 标准化交付 · mysql-01") + } + if summary.Runner != "AWX Job Template #10" { + t.Fatalf("awxTaskSummary Runner = %q, want %q", summary.Runner, "AWX Job Template #10") + } + if summary.Status != model.TaskFinished { + t.Fatalf("awxTaskSummary Status = %q, want %q", summary.Status, model.TaskFinished) + } + if summary.StatusText != "成功" { + t.Fatalf("awxTaskSummary StatusText = %q, want %q", summary.StatusText, "成功") + } + if summary.StatusClass != "ok" { + t.Fatalf("awxTaskSummary StatusClass = %q, want %q", summary.StatusClass, "ok") + } + if summary.BusinessLineID != 5 { + t.Fatalf("awxTaskSummary BusinessLineID = %d, want 5", summary.BusinessLineID) + } + if summary.ReferenceID != "task-123" { + t.Fatalf("awxTaskSummary ReferenceID = %q, want %q", summary.ReferenceID, "task-123") + } + if !summary.CreatedAt.Equal(now.Add(-1 * time.Hour)) { + t.Fatalf("awxTaskSummary CreatedAt = %v, want %v", summary.CreatedAt, now.Add(-1*time.Hour)) + } + if !summary.UpdatedAt.Equal(now) { + t.Fatalf("awxTaskSummary UpdatedAt = %v, want %v", summary.UpdatedAt, now) + } +} + +func TestWayneTaskSummary(t *testing.T) { + now := time.Now() + history := service.WayneDeploymentHistory{ + ID: 200, + ResourceID: 42, + ResourceName: "my-service", + Status: 1, + BusinessLineID: 3, + CreatedAt: now, + } + summary := wayneTaskSummary(history) + + if summary.ID != "wayne:publish:42:200" { + t.Fatalf("wayneTaskSummary ID = %q, want %q", summary.ID, "wayne:publish:42:200") + } + if summary.Source != "wayne" { + t.Fatalf("wayneTaskSummary Source = %q, want %q", summary.Source, "wayne") + } + if summary.Service != "wayne-deployment" { + t.Fatalf("wayneTaskSummary Service = %q, want %q", summary.Service, "wayne-deployment") + } + if summary.Name != "Wayne 服务部署 · my-service" { + t.Fatalf("wayneTaskSummary Name = %q, want %q", summary.Name, "Wayne 服务部署 · my-service") + } + if summary.Runner != "Wayne Native API" { + t.Fatalf("wayneTaskSummary Runner = %q, want %q", summary.Runner, "Wayne Native API") + } + if summary.Status != model.TaskFinished { + t.Fatalf("wayneTaskSummary Status = %q, want %q", summary.Status, model.TaskFinished) + } + if summary.StatusText != "成功" { + t.Fatalf("wayneTaskSummary StatusText = %q, want %q", summary.StatusText, "成功") + } + if summary.StatusClass != "ok" { + t.Fatalf("wayneTaskSummary StatusClass = %q, want %q", summary.StatusClass, "ok") + } + if summary.BusinessLineID != 3 { + t.Fatalf("wayneTaskSummary BusinessLineID = %d, want 3", summary.BusinessLineID) + } + if summary.ReferenceID != "200" { + t.Fatalf("wayneTaskSummary ReferenceID = %q, want %q", summary.ReferenceID, "200") + } +} + +func TestParseActualTaskID(t *testing.T) { + h := &TaskLogHandler{} + + for name, tc := range map[string]struct { + input string + want string + }{ + "awx task": {input: "awx:task-123", want: "task-123"}, + "wayne task": {input: "wayne:publish:42:200", want: "wayne:42"}, + "empty": {input: "", want: ""}, + "unknown": {input: "unknown:id", want: ""}, + "awx no prefix": {input: "awx:", want: ""}, + } { + t.Run(name, func(t *testing.T) { + got := h.parseActualTaskID(tc.input) + if got != tc.want { + t.Fatalf("parseActualTaskID(%q) = %q, want %q", tc.input, got, tc.want) + } + }) + } +} diff --git a/server/internal/service/delivery_task_test.go b/server/internal/service/delivery_task_test.go new file mode 100644 index 0000000..77c5ea4 --- /dev/null +++ b/server/internal/service/delivery_task_test.go @@ -0,0 +1,440 @@ +package service + +import ( + "context" + "sync" + "testing" + "time" + + "github.com/1024XEngineer/xinfra/server/internal/model" +) + +// newTestDeliveryService 创建一个用于测试的 DeliveryService,不需要真实的 DB 和配置。 +// 仅适用于测试 pub/sub、缓存等内存逻辑。 +func newTestDeliveryService() *DeliveryService { + return &DeliveryService{ + streams: make(map[string]map[chan DeliveryTaskSnapshot]struct{}), + stdoutCache: make(map[string]*stdoutCacheItem), + } +} + +func TestSubscribeTask(t *testing.T) { + svc := newTestDeliveryService() + ctx := context.Background() + _ = ctx + + taskID := "test-task-1" + + // 订阅任务 + ch, cancel := svc.SubscribeTask(taskID) + defer cancel() + + // 验证 channel 已注册 + svc.streamMu.Lock() + if _, ok := svc.streams[taskID]; !ok { + svc.streamMu.Unlock() + t.Fatal("SubscribeTask did not register channel in streams map") + } + svc.streamMu.Unlock() + + // 模拟 broadcastTask 推送 snapshot + snapshot := DeliveryTaskSnapshot{ + Task: &model.DeliveryTask{ID: taskID, Status: model.TaskRunning}, + } + svc.streamMu.Lock() + for ch := range svc.streams[taskID] { + select { + case ch <- snapshot: + default: + } + } + svc.streamMu.Unlock() + + // 接收推送 + select { + case received := <-ch: + if received.Task == nil || received.Task.ID != taskID { + t.Fatalf("received snapshot task ID = %v, want %q", received.Task, taskID) + } + if received.Task.Status != model.TaskRunning { + t.Fatalf("received snapshot status = %q, want %q", received.Task.Status, model.TaskRunning) + } + case <-time.After(time.Second): + t.Fatal("timeout waiting for snapshot from SubscribeTask") + } +} + +func TestSubscribeTask_Cancel(t *testing.T) { + svc := newTestDeliveryService() + taskID := "test-task-cancel" + + ch, cancel := svc.SubscribeTask(taskID) + + // 调用 cancel + cancel() + + // 验证 channel 已关闭 + select { + case _, ok := <-ch: + if ok { + t.Fatal("channel should be closed after cancel, but got a value") + } + case <-time.After(time.Second): + t.Fatal("timeout waiting for channel close") + } + + // 验证已从 streams 中移除 + svc.streamMu.Lock() + if subs := svc.streams[taskID]; subs != nil && len(subs) > 0 { + svc.streamMu.Unlock() + t.Fatal("cancel did not remove channel from streams map") + } + svc.streamMu.Unlock() +} + +func TestSubscribeTask_MultipleSubscribers(t *testing.T) { + svc := newTestDeliveryService() + taskID := "test-task-multi" + + ch1, cancel1 := svc.SubscribeTask(taskID) + defer cancel1() + ch2, cancel2 := svc.SubscribeTask(taskID) + defer cancel2() + ch3, cancel3 := svc.SubscribeTask(taskID) + defer cancel3() + + // 验证三个订阅者都已注册 + svc.streamMu.Lock() + subs := svc.streams[taskID] + if len(subs) != 3 { + svc.streamMu.Unlock() + t.Fatalf("expected 3 subscribers, got %d", len(subs)) + } + svc.streamMu.Unlock() + + // 推送 snapshot + snapshot := DeliveryTaskSnapshot{ + Task: &model.DeliveryTask{ID: taskID, Status: model.TaskFinished}, + } + svc.streamMu.Lock() + for ch := range svc.streams[taskID] { + select { + case ch <- snapshot: + default: + } + } + svc.streamMu.Unlock() + + // 验证三个订阅者都收到 + for i, ch := range []<-chan DeliveryTaskSnapshot{ch1, ch2, ch3} { + select { + case received := <-ch: + if received.Task == nil || received.Task.Status != model.TaskFinished { + t.Fatalf("subscriber %d: expected TaskFinished, got %+v", i+1, received) + } + case <-time.After(time.Second): + t.Fatalf("subscriber %d: timeout waiting for snapshot", i+1) + } + } +} + +func TestSubscribeTask_CancelOneDoesNotAffectOthers(t *testing.T) { + svc := newTestDeliveryService() + taskID := "test-task-cancel-one" + + ch1, cancel1 := svc.SubscribeTask(taskID) + _, cancel2 := svc.SubscribeTask(taskID) + _ = cancel2 + + // 取消第一个订阅者 + cancel1() + + // 验证还剩一个订阅者 + svc.streamMu.Lock() + subs := svc.streams[taskID] + if len(subs) != 1 { + svc.streamMu.Unlock() + t.Fatalf("expected 1 subscriber after cancel, got %d", len(subs)) + } + svc.streamMu.Unlock() + + // 推送 snapshot + snapshot := DeliveryTaskSnapshot{ + Task: &model.DeliveryTask{ID: taskID, Status: model.TaskRunning}, + } + svc.streamMu.Lock() + for ch := range svc.streams[taskID] { + select { + case ch <- snapshot: + default: + } + } + svc.streamMu.Unlock() + + // ch1 已关闭,不应收到消息 + select { + case _, ok := <-ch1: + if ok { + t.Fatal("ch1 should be closed after cancel") + } + case <-time.After(100 * time.Millisecond): + // OK: channel is closed, no value received + } +} + +func TestBroadcastTask_ClosedChannelCleanup(t *testing.T) { + svc := newTestDeliveryService() + taskID := "test-task-cleanup" + + // 创建一个订阅者并立即关闭 + _, cancel := svc.SubscribeTask(taskID) + cancel() + // 等待 cancel 完成 + time.Sleep(10 * time.Millisecond) + + // 创建一个新的正常订阅者 + ch2, cancel2 := svc.SubscribeTask(taskID) + defer cancel2() + + // 模拟 broadcastTask 行为(带 recover) + snapshot := DeliveryTaskSnapshot{ + Task: &model.DeliveryTask{ID: taskID, Status: model.TaskRunning}, + } + + svc.streamMu.Lock() + var closedChannels []chan DeliveryTaskSnapshot + for ch := range svc.streams[taskID] { + func() { + defer func() { + if r := recover(); r != nil { + closedChannels = append(closedChannels, ch) + } + }() + select { + case ch <- snapshot: + default: + } + }() + } + for _, ch := range closedChannels { + delete(svc.streams[taskID], ch) + } + if len(svc.streams[taskID]) == 0 { + delete(svc.streams, taskID) + } + svc.streamMu.Unlock() + + // 验证 closed channel 被清理 + svc.streamMu.Lock() + if subs := svc.streams[taskID]; subs != nil && len(subs) != 1 { + svc.streamMu.Unlock() + t.Fatalf("expected 1 subscriber after cleanup, got %d", len(subs)) + } + svc.streamMu.Unlock() + + // 正常订阅者应该收到消息 + select { + case received := <-ch2: + if received.Task == nil || received.Task.ID != taskID { + t.Fatalf("expected task ID %q, got %+v", taskID, received) + } + case <-time.After(time.Second): + t.Fatal("timeout waiting for snapshot from normal subscriber") + } +} + +func TestInvalidateStdoutCache(t *testing.T) { + svc := newTestDeliveryService() + + // 填充缓存 + svc.cacheMu.Lock() + svc.stdoutCache["job-1"] = &stdoutCacheItem{stdout: "cached output", createdAt: time.Now()} + svc.cacheMu.Unlock() + + // 验证缓存命中 + svc.cacheMu.RLock() + item, ok := svc.stdoutCache["job-1"] + svc.cacheMu.RUnlock() + if !ok { + t.Fatal("cache entry not found before invalidation") + } + if item.stdout != "cached output" { + t.Fatalf("cache stdout = %q, want %q", item.stdout, "cached output") + } + + // 失效缓存 + svc.InvalidateStdoutCache("job-1") + + // 验证缓存已失效 + svc.cacheMu.RLock() + _, ok = svc.stdoutCache["job-1"] + svc.cacheMu.RUnlock() + if ok { + t.Fatal("cache entry still exists after invalidation") + } +} + +func TestInvalidateStdoutCache_NonExistent(t *testing.T) { + svc := newTestDeliveryService() + + // 对不存在的 key 调用 invalidate 不应 panic + svc.InvalidateStdoutCache("non-existent-job") + + // 验证缓存为空 + svc.cacheMu.RLock() + size := len(svc.stdoutCache) + svc.cacheMu.RUnlock() + if size != 0 { + t.Fatalf("cache size = %d, want 0", size) + } +} + +func TestStdoutCache_TTLExpiry(t *testing.T) { + svc := newTestDeliveryService() + + // 填充一个已过期的缓存条目 + svc.cacheMu.Lock() + svc.stdoutCache["job-expired"] = &stdoutCacheItem{ + stdout: "old output", + createdAt: time.Now().Add(-stdoutCacheTTL - time.Second), + } + svc.cacheMu.Unlock() + + // 模拟 AWXJobStdout 的缓存检查逻辑 + svc.cacheMu.RLock() + cacheHit := false + if item, ok := svc.stdoutCache["job-expired"]; ok { + if time.Since(item.createdAt) < stdoutCacheTTL { + cacheHit = true + } + } + svc.cacheMu.RUnlock() + + if cacheHit { + t.Fatal("expired cache entry should not be a hit") + } +} + +func TestStdoutCache_FreshEntry(t *testing.T) { + svc := newTestDeliveryService() + + // 填充一个新鲜的缓存条目 + svc.cacheMu.Lock() + svc.stdoutCache["job-fresh"] = &stdoutCacheItem{ + stdout: "fresh output", + createdAt: time.Now(), + } + svc.cacheMu.Unlock() + + // 模拟 AWXJobStdout 的缓存检查逻辑 + svc.cacheMu.RLock() + cacheHit := false + var cachedStdout string + if item, ok := svc.stdoutCache["job-fresh"]; ok { + if time.Since(item.createdAt) < stdoutCacheTTL { + cacheHit = true + cachedStdout = item.stdout + } + } + svc.cacheMu.RUnlock() + + if !cacheHit { + t.Fatal("fresh cache entry should be a hit") + } + if cachedStdout != "fresh output" { + t.Fatalf("cached stdout = %q, want %q", cachedStdout, "fresh output") + } +} + +func TestSubscribeTask_ConcurrentSafety(t *testing.T) { + svc := newTestDeliveryService() + taskID := "test-concurrent" + + var wg sync.WaitGroup + const goroutines = 50 + + // 并发订阅 + cancels := make([]func(), 0, goroutines) + for i := 0; i < goroutines; i++ { + wg.Add(1) + go func() { + defer wg.Done() + _, cancel := svc.SubscribeTask(taskID) + cancels = append(cancels, cancel) + }() + } + wg.Wait() + + // 验证所有订阅者都已注册 + svc.streamMu.Lock() + subs := svc.streams[taskID] + if len(subs) != goroutines { + svc.streamMu.Unlock() + t.Fatalf("expected %d subscribers, got %d", goroutines, len(subs)) + } + svc.streamMu.Unlock() + + // 并发取消 + for _, cancel := range cancels { + wg.Add(1) + go func(c func()) { + defer wg.Done() + c() + }(cancel) + } + wg.Wait() + + // 验证所有订阅者都已移除 + svc.streamMu.Lock() + if subs := svc.streams[taskID]; subs != nil && len(subs) > 0 { + svc.streamMu.Unlock() + t.Fatalf("expected 0 subscribers after concurrent cancel, got %d", len(subs)) + } + svc.streamMu.Unlock() +} + +func TestStdoutCache_ConcurrentAccess(t *testing.T) { + svc := newTestDeliveryService() + + var wg sync.WaitGroup + const goroutines = 50 + + // 并发写入缓存 + for i := 0; i < goroutines; i++ { + wg.Add(1) + go func(idx int) { + defer wg.Done() + jobID := "job-" + string(rune('A'+idx%26)) + svc.cacheMu.Lock() + svc.stdoutCache[jobID] = &stdoutCacheItem{ + stdout: "output", + createdAt: time.Now(), + } + svc.cacheMu.Unlock() + }(i) + } + + // 并发读取缓存 + for i := 0; i < goroutines; i++ { + wg.Add(1) + go func(idx int) { + defer wg.Done() + jobID := "job-" + string(rune('A'+idx%26)) + svc.cacheMu.RLock() + _ = svc.stdoutCache[jobID] + svc.cacheMu.RUnlock() + }(i) + } + + // 并发失效缓存 + for i := 0; i < goroutines; i++ { + wg.Add(1) + go func(idx int) { + defer wg.Done() + jobID := "job-" + string(rune('A'+idx%26)) + svc.InvalidateStdoutCache(jobID) + }(i) + } + + wg.Wait() +}