From 61eeb0c8acab3775121ec7a38725996713b88ad3 Mon Sep 17 00:00:00 2001 From: hezhaohui Date: Wed, 29 Jul 2026 14:00:40 +0800 Subject: [PATCH 1/3] =?UTF-8?q?fix(server):=20=E4=BF=AE=E5=A4=8D=20task=20?= =?UTF-8?q?log=20=E6=8E=A5=E5=8F=A3=E5=A4=9A=E4=B8=AA=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 修复 Wayne 任务状态映射,将数字状态统一为字符串状态 - 添加分页支持,支持 page/page_size 参数 - 添加 AWXJobStdout 内存缓存,减少重复 API 调用 - 修复 SubscribeTask 内存泄漏,自动清理 closed channel - 统一 Stream 端点,AWX 使用 pub/sub,Wayne 使用轮询 背景:task log 接口存在多个潜在问题,包括状态判断错误导致无限轮询、 缺少分页、重复 API 调用性能问题、内存泄漏风险等 关联 commit:fix/logs 分支 --- server/internal/handler/task_log.go | 169 +++++++++++++++++++++++++--- server/internal/service/delivery.go | 104 ++++++++++++++++- 2 files changed, 250 insertions(+), 23 deletions(-) diff --git a/server/internal/handler/task_log.go b/server/internal/handler/task_log.go index 81bf4bc..5857053 100644 --- a/server/internal/handler/task_log.go +++ b/server/internal/handler/task_log.go @@ -8,6 +8,7 @@ import ( "strings" "time" + "github.com/1024XEngineer/xinfra/server/internal/auth" "github.com/1024XEngineer/xinfra/server/internal/model" "github.com/1024XEngineer/xinfra/server/internal/service" "github.com/gin-gonic/gin" @@ -55,6 +56,16 @@ func (h *TaskLogHandler) List(c *gin.Context) { source := strings.ToLower(strings.TrimSpace(c.DefaultQuery("source", "all"))) businessLineID, _ := strconv.ParseUint(c.Query("business_line_id"), 10, 64) + // 分页参数 + page, _ := strconv.Atoi(c.DefaultQuery("page", "1")) + pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20")) + if page < 1 { + page = 1 + } + if pageSize < 1 || pageSize > 100 { + pageSize = 20 + } + items := make([]taskLogSummary, 0) if source == "" || source == "all" || source == "awx" { awxItems, err := h.listAWXTasks(c, claims.UserID, claims.IsAdmin, businessLineID) @@ -73,10 +84,27 @@ func (h *TaskLogHandler) List(c *gin.Context) { items = append(items, wayneItems...) } sortTaskLogSummaries(items) - if len(items) > 100 { - items = items[:100] + + // 计算分页 + total := len(items) + start := (page - 1) * pageSize + if start >= total { + items = make([]taskLogSummary, 0) + } else { + end := start + pageSize + if end > total { + end = total + } + items = items[start:end] } - c.JSON(http.StatusOK, gin.H{"items": items}) + + c.JSON(http.StatusOK, gin.H{ + "items": items, + "total": total, + "page": page, + "page_size": pageSize, + "total_pages": (total + pageSize - 1) / pageSize, + }) } func (h *TaskLogHandler) Get(c *gin.Context) { @@ -211,7 +239,7 @@ func wayneTaskSummary(history service.WayneDeploymentHistory) taskLogSummary { Service: "wayne-deployment", Name: wayneDeploymentTaskName(history), Runner: "Wayne Native API", - Status: strconv.Itoa(history.Status), + Status: wayneStatusToString(history.Status), StatusText: textForWaynePublishStatus(history.Status), StatusClass: classForWaynePublishStatus(history.Status), BusinessLineID: history.BusinessLineID, @@ -325,6 +353,18 @@ func classForWaynePublishStatus(status int) string { } } +// wayneStatusToString 将 Wayne 数字状态映射为与 AWX 一致的字符串状态 +func wayneStatusToString(status int) string { + switch status { + case 1: + return model.TaskFinished + case 0: + return model.TaskExecutionFailed + default: + return "unknown" + } +} + func splitStdoutLines(stdout string) []taskLogLine { lines := make([]taskLogLine, 0) for _, line := range strings.Split(stdout, "\n") { @@ -352,6 +392,7 @@ func classForOutputLine(line string) string { } // Stream SSE 端点,用于实时推送任务日志。 +// AWX 任务使用 pub/sub 模式,Wayne 任务使用轮询模式。 func (h *TaskLogHandler) Stream(c *gin.Context) { claims, ok := CurrentClaims(c) if !ok { @@ -373,10 +414,94 @@ func (h *TaskLogHandler) Stream(c *gin.Context) { ctx := c.Request.Context() + // AWX 任务使用 pub/sub 模式 + if strings.HasPrefix(taskID, "awx:") { + h.streamAWXTask(c, taskID, ctx, claims) + return + } + + // Wayne 任务使用轮询模式(因为 Wayne API 不支持推送) + if strings.HasPrefix(taskID, "wayne:") { + h.streamWayneTask(c, taskID, ctx, claims) + return + } + + c.SSEvent("message", gin.H{"type": "error", "error": "invalid task_id format"}) + c.Writer.Flush() +} + +// streamAWXTask AWX 任务使用 pub/sub 模式 +func (h *TaskLogHandler) streamAWXTask(c *gin.Context, taskID string, ctx context.Context, claims *auth.Claims) { + actualTaskID := strings.TrimPrefix(taskID, "awx:") + // 发送初始日志 initialLines, taskStatus, err := h.fetchTaskLogLines(ctx, taskID, claims.UserID, claims.IsAdmin) if err != nil { c.SSEvent("message", gin.H{"type": "error", "error": err.Error()}) + c.Writer.Flush() + return + } + + // 发送初始数据 + c.SSEvent("message", gin.H{"type": "init", "lines": initialLines}) + c.Writer.Flush() + + // 如果任务已完成,直接发送结束事件 + if isTerminalStatus(taskStatus) { + c.SSEvent("message", gin.H{"type": "finished"}) + c.Writer.Flush() + return + } + + // 订阅任务更新(使用 pub/sub 模式) + updates, cancel := h.delivery.SubscribeTask(actualTaskID) + defer cancel() + + heartbeatTicker := time.NewTicker(15 * time.Second) + defer heartbeatTicker.Stop() + + for { + select { + case <-ctx.Done(): + return + case _, ok := <-updates: + if !ok { + return + } + + // 重新获取日志行 + lines, status, err := h.fetchTaskLogLines(ctx, taskID, claims.UserID, claims.IsAdmin) + if err != nil { + c.SSEvent("message", gin.H{"type": "error", "error": err.Error()}) + c.Writer.Flush() + continue + } + + // 发送完整日志 + c.SSEvent("message", gin.H{"type": "update", "lines": lines}) + c.Writer.Flush() + + // 如果任务完成,发送结束事件 + if isTerminalStatus(status) { + c.SSEvent("message", gin.H{"type": "finished"}) + c.Writer.Flush() + return + } + + case <-heartbeatTicker.C: + c.SSEvent("heartbeat", nil) + c.Writer.Flush() + } + } +} + +// streamWayneTask Wayne 任务使用轮询模式 +func (h *TaskLogHandler) streamWayneTask(c *gin.Context, taskID string, ctx context.Context, claims *auth.Claims) { + // 发送初始日志 + initialLines, taskStatus, err := h.fetchTaskLogLines(ctx, taskID, claims.UserID, claims.IsAdmin) + if err != nil { + c.SSEvent("message", gin.H{"type": "error", "error": err.Error()}) + c.Writer.Flush() return } @@ -392,14 +517,12 @@ func (h *TaskLogHandler) Stream(c *gin.Context) { } // 轮询循环 - ticker := time.NewTicker(2 * time.Second) + ticker := time.NewTicker(5 * time.Second) defer ticker.Stop() heartbeatTicker := time.NewTicker(15 * time.Second) defer heartbeatTicker.Stop() - lastLineCount := len(initialLines) - for { select { case <-ctx.Done(): @@ -407,19 +530,14 @@ func (h *TaskLogHandler) Stream(c *gin.Context) { case <-ticker.C: lines, status, err := h.fetchTaskLogLines(ctx, taskID, claims.UserID, claims.IsAdmin) if err != nil { - // 发送错误但不关闭连接 c.SSEvent("message", gin.H{"type": "error", "error": err.Error()}) c.Writer.Flush() continue } - // 发送增量日志 - if len(lines) > lastLineCount { - newLines := lines[lastLineCount:] - c.SSEvent("message", gin.H{"type": "update", "lines": newLines}) - c.Writer.Flush() - lastLineCount = len(lines) - } + // 发送完整日志 + c.SSEvent("message", gin.H{"type": "update", "lines": lines}) + c.Writer.Flush() // 如果任务完成,发送结束事件 if isTerminalStatus(status) { @@ -429,13 +547,29 @@ func (h *TaskLogHandler) Stream(c *gin.Context) { } case <-heartbeatTicker.C: - // 发送心跳保持连接 c.SSEvent("heartbeat", nil) c.Writer.Flush() } } } +// parseActualTaskID 从 task log ID 解析出实际的任务 ID +func (h *TaskLogHandler) parseActualTaskID(taskID string) string { + switch { + case strings.HasPrefix(taskID, "awx:"): + return strings.TrimPrefix(taskID, "awx:") + case strings.HasPrefix(taskID, "wayne:publish:"): + // Wayne 任务需要使用 resourceID 作为订阅 key + parts := strings.Split(taskID, ":") + if len(parts) == 4 { + return "wayne:" + parts[2] + } + return "" + default: + return "" + } +} + // fetchTaskLogLines 获取任务的日志行和状态。 func (h *TaskLogHandler) fetchTaskLogLines(ctx context.Context, taskID string, userID uint64, isAdmin bool) ([]taskLogLine, string, error) { switch { @@ -513,7 +647,8 @@ func (h *TaskLogHandler) fetchWayneLogLines(ctx context.Context, id string, user lines = append(lines, taskLogLine{Time: formatTaskLogTime(history.CreatedAt), Message: "[message] " + history.Message, Class: classForWaynePublishStatus(history.Status)}) } - statusText := strconv.Itoa(history.Status) + // 将 Wayne 数字状态映射为与 AWX 一致的字符串状态 + statusText := wayneStatusToString(history.Status) return lines, statusText, nil } diff --git a/server/internal/service/delivery.go b/server/internal/service/delivery.go index 50f34af..94e3f9f 100644 --- a/server/internal/service/delivery.go +++ b/server/internal/service/delivery.go @@ -332,6 +332,12 @@ func allocatePort(requested int, used []int) (int, error) { return 0, fmt.Errorf("mysql port pool %d-%d is exhausted on the target host", mysqlPortPoolStart, mysqlPortPoolEnd) } +// stdoutCacheItem 缓存 AWX Job stdout 的结果 +type stdoutCacheItem struct { + stdout string + createdAt time.Time +} + type DeliveryService struct { db *gorm.DB cfg config.Config @@ -340,12 +346,21 @@ type DeliveryService struct { executionMu sync.Mutex streamMu sync.Mutex streams map[string]map[chan DeliveryTaskSnapshot]struct{} + stdoutCache map[string]*stdoutCacheItem + cacheMu sync.RWMutex } func (s *DeliveryService) DB() *gorm.DB { return s.db } func NewDeliveryService(cfg config.Config, db *gorm.DB, audit *AuditService) *DeliveryService { - return &DeliveryService{db: db, cfg: cfg, awx: NewAWXClient(cfg.AWXBaseURL, cfg.AWXToken, cfg.AWXUsername, cfg.AWXPassword), audit: audit, streams: make(map[string]map[chan DeliveryTaskSnapshot]struct{})} + return &DeliveryService{ + db: db, + cfg: cfg, + awx: NewAWXClient(cfg.AWXBaseURL, cfg.AWXToken, cfg.AWXUsername, cfg.AWXPassword), + audit: audit, + streams: make(map[string]map[chan DeliveryTaskSnapshot]struct{}), + stdoutCache: make(map[string]*stdoutCacheItem), + } } func (s *DeliveryService) ListTargets(ctx context.Context, component string) ([]DeliveryTarget, error) { @@ -942,18 +957,95 @@ func (s *DeliveryService) broadcastTask(ctx context.Context, taskID string) { if err != nil { return } + + // 任务状态变化时清除 stdout 缓存 + s.invalidateStdoutCacheForTask(taskID) + s.streamMu.Lock() defer s.streamMu.Unlock() + + // 收集需要清理的 closed channel + var closedChannels []chan DeliveryTaskSnapshot + for ch := range s.streams[taskID] { - select { - case ch <- snapshot: - default: - } + // 使用 recover 捕获 send on closed channel 的错误 + func() { + defer func() { + if r := recover(); r != nil { + // channel 已关闭,标记需要清理 + closedChannels = append(closedChannels, ch) + } + }() + select { + case ch <- snapshot: + default: + } + }() + } + + // 清理 closed channels + for _, ch := range closedChannels { + delete(s.streams[taskID], ch) + } + if len(s.streams[taskID]) == 0 { + delete(s.streams, taskID) } } +// invalidateStdoutCacheForTask 清除与任务相关的 stdout 缓存 +func (s *DeliveryService) invalidateStdoutCacheForTask(taskID string) { + var execution model.ExecutionJob + var rollback model.RollbackJob + + s.cacheMu.Lock() + defer s.cacheMu.Unlock() + + // 清除 execution job 的缓存 + if err := s.db.Where("task_id = ?", taskID).First(&execution).Error; err == nil && execution.ExecutorJobID != "" { + delete(s.stdoutCache, execution.ExecutorJobID) + } + + // 清除 rollback job 的缓存 + if err := s.db.Where("task_id = ?", taskID).First(&rollback).Error; err == nil && rollback.ExecutorJobID != "" { + delete(s.stdoutCache, rollback.ExecutorJobID) + } +} + +const stdoutCacheTTL = 30 * time.Second + func (s *DeliveryService) AWXJobStdout(ctx context.Context, jobID string) (string, error) { - return s.awx.JobStdout(ctx, jobID) + // 检查缓存 + s.cacheMu.RLock() + if item, ok := s.stdoutCache[jobID]; ok { + if time.Since(item.createdAt) < stdoutCacheTTL { + s.cacheMu.RUnlock() + return item.stdout, nil + } + } + s.cacheMu.RUnlock() + + // 缓存未命中或已过期,重新获取 + stdout, err := s.awx.JobStdout(ctx, jobID) + if err != nil { + return "", err + } + + // 更新缓存 + s.cacheMu.Lock() + s.stdoutCache[jobID] = &stdoutCacheItem{ + stdout: stdout, + createdAt: time.Now(), + } + s.cacheMu.Unlock() + + return stdout, nil +} + +// InvalidateStdoutCache 清除指定 jobID 的 stdout 缓存 +func (s *DeliveryService) InvalidateStdoutCache(jobID string) { + s.cacheMu.Lock() + delete(s.stdoutCache, jobID) + s.cacheMu.Unlock() } func (s *DeliveryService) Cancel(ctx context.Context, taskID string, userID uint64, isAdmin bool) error { From dbb930d0cb128d5cdea9ac4bd7582c00945e57fc Mon Sep 17 00:00:00 2001 From: hezhaohui Date: Wed, 29 Jul 2026 14:01:12 +0800 Subject: [PATCH 2/3] =?UTF-8?q?fix(frontend):=20=E4=BC=98=E5=8C=96=20task?= =?UTF-8?q?=20log=20=E5=89=8D=E7=AB=AF=E4=BA=A4=E4=BA=92=E4=BD=93=E9=AA=8C?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 简化虚拟滚动实现,移除 @tanstack/vue-virtual 依赖 - 添加分页组件,支持翻页查看历史任务 - 添加 SSE 指数退避重试机制(最多 3 次) - 更新 API 接口支持分页参数 - 添加重试状态指示器 背景:前端存在虚拟滚动与 SSE 不兼容、缺少分页、连接断开无重试等问题 关联 commit:fix/logs 分支 --- frontend/src/api/taskLog.ts | 28 +++++- frontend/src/views/task/TaskCenter.vue | 113 +++++++++++++++++++------ 2 files changed, 112 insertions(+), 29 deletions(-) diff --git a/frontend/src/api/taskLog.ts b/frontend/src/api/taskLog.ts index 4791b15..c4505d9 100644 --- a/frontend/src/api/taskLog.ts +++ b/frontend/src/api/taskLog.ts @@ -24,10 +24,20 @@ export interface TaskLogLine { export interface TaskLogListParams { source?: string businessLineId?: number + page?: number + pageSize?: number +} + +export interface TaskLogListResult { + items: TaskLogSummary[] + total: number + page: number + page_size: number + total_pages: number } export const taskLogApi = { - async list(params: TaskLogListParams = {}): Promise { + async list(params: TaskLogListParams = {}): Promise { const query = new URLSearchParams() if (params.source && params.source !== 'all') { query.set('source', params.source) @@ -35,9 +45,21 @@ export const taskLogApi = { if (params.businessLineId) { query.set('business_line_id', String(params.businessLineId)) } + if (params.page) { + query.set('page', String(params.page)) + } + if (params.pageSize) { + query.set('page_size', String(params.pageSize)) + } const suffix = query.toString() ? `?${query.toString()}` : '' const data = await authRequest(`/auth/api/v1/task-logs${suffix}`) - return Array.isArray(data.items) ? data.items : [] + return { + items: Array.isArray(data.items) ? data.items : [], + total: data.total || 0, + page: data.page || 1, + page_size: data.page_size || 20, + total_pages: data.total_pages || 0, + } }, async get(id: string): Promise<{ task: TaskLogSummary; lines: TaskLogLine[] }> { @@ -89,6 +111,8 @@ export interface TaskLogStreamCallbacks { export function createTaskLogStream(taskId: string, callbacks: TaskLogStreamCallbacks): EventSource { const token = getToken() + // 注意:EventSource 不支持自定义请求头,只能通过 query parameter 传递 token + // TODO: 在生产环境中,应考虑使用 WebSocket 或 HttpOnly Cookie 方式以提高安全性 const url = `/auth/api/v1/task-logs/stream?task_id=${encodeURIComponent(taskId)}&access_token=${encodeURIComponent(token || '')}` const es = new EventSource(url) diff --git a/frontend/src/views/task/TaskCenter.vue b/frontend/src/views/task/TaskCenter.vue index 402ac67..c68f673 100644 --- a/frontend/src/views/task/TaskCenter.vue +++ b/frontend/src/views/task/TaskCenter.vue @@ -42,7 +42,11 @@ 暂无任务记录 @@ -53,6 +57,7 @@ 任务日志 {{ selectedTaskName }} ● 实时更新中 + 重试中 ({{ streamRetryCount }}/3)
- +
-
- {{ logs[virtualRow.index].time }}{{ logs[virtualRow.index].message }} -
+ {{ log.time }}{{ log.message }}
...加载中 @@ -104,7 +98,6 @@ import { computed, nextTick, onMounted, onUnmounted, ref, watch } from 'vue' import { taskLogApi, createTaskLogStream, type TaskLogLine, type TaskLogSummary } from '@/api/taskLog' import { useBusinessLineStore } from '@/stores/businessLine' import { useBusinessLineMockProfile } from '@/utils/businessLineMock' -import { useVirtualizer } from '@tanstack/vue-virtual' const { currentName } = useBusinessLineMockProfile() const businessLineStore = useBusinessLineStore() @@ -121,13 +114,16 @@ let eventSource: EventSource | null = null const autoScroll = ref(true) const logContainerRef = ref(null) -// 虚拟滚动配置 -const virtualizer = useVirtualizer({ - count: logs.value.length, - getScrollElement: () => logContainerRef.value, - estimateSize: () => 20, - overscan: 5, -}) +// 分页相关 +const currentPage = ref(1) +const pageSize = ref(20) +const totalTasks = ref(0) +const totalPages = computed(() => Math.ceil(totalTasks.value / pageSize.value)) + +// SSE 重试相关 +const streamRetryCount = ref(0) +const maxRetryCount = 3 +let retryTimeout: ReturnType | null = null const selectedTask = computed(() => tasks.value.find((task) => task.id === selectedTaskId.value)) const selectedTaskName = computed(() => selectedTask.value?.name || '未选择') @@ -137,10 +133,14 @@ const lastLoadedText = computed(() => lastLoadedAt.value ? `更新于 ${formatTi async function loadTasks() { loadingTasks.value = true try { - tasks.value = await taskLogApi.list({ + const result = await taskLogApi.list({ source: sourceFilter.value, businessLineId: businessLineStore.current?.id, + page: currentPage.value, + pageSize: pageSize.value, }) + tasks.value = result.items + totalTasks.value = result.total lastLoadedAt.value = new Date() if (!tasks.value.some((task) => task.id === selectedTaskId.value)) { selectedTaskId.value = tasks.value[0]?.id || '' @@ -163,6 +163,7 @@ async function loadLogs(taskId: string) { const data = await taskLogApi.get(taskId) logs.value = data.lines lastLoadedAt.value = new Date() + scrollToBottom() // 启动 SSE 流式更新 startStream(taskId) } catch (error) { @@ -174,6 +175,7 @@ async function loadLogs(taskId: string) { function startStream(taskId: string) { stopStream() // 关闭之前的连接 + streamRetryCount.value = 0 eventSource = createTaskLogStream(taskId, { onInit: (lines) => { @@ -182,23 +184,45 @@ function startStream(taskId: string) { scrollToBottom() }, onUpdate: (newLines) => { - logs.value = [...logs.value, ...newLines] + logs.value = newLines lastLoadedAt.value = new Date() scrollToBottom() }, onFinished: () => { isStreaming.value = false + streamRetryCount.value = 0 }, onError: (error) => { console.error('SSE error:', error) isStreaming.value = false + // 尝试重连 + retryStream(taskId) }, }) isStreaming.value = true } +function retryStream(taskId: string) { + if (streamRetryCount.value >= maxRetryCount) { + console.error('Max retry count reached') + return + } + + streamRetryCount.value++ + // 指数退避:1s, 2s, 4s + const delay = Math.pow(2, streamRetryCount.value - 1) * 1000 + + retryTimeout = setTimeout(() => { + startStream(taskId) + }, delay) +} + function stopStream() { + if (retryTimeout) { + clearTimeout(retryTimeout) + retryTimeout = null + } if (eventSource) { eventSource.close() eventSource = null @@ -244,6 +268,20 @@ function refreshCurrent() { void loadTasks() } +function prevPage() { + if (currentPage.value > 1) { + currentPage.value-- + void loadTasks() + } +} + +function nextPage() { + if (currentPage.value < totalPages.value) { + currentPage.value++ + void loadTasks() + } +} + function formatTime(date: Date) { return date.toLocaleTimeString('zh-CN', { hour12: false }) } @@ -251,6 +289,7 @@ function formatTime(date: Date) { watch( () => businessLineStore.current?.id, () => { + currentPage.value = 1 void loadTasks() }, ) @@ -419,6 +458,26 @@ onUnmounted(() => { font-size: 13px; } +.pagination { + display: flex; + justify-content: space-between; + align-items: center; + padding: 12px 0; + font-size: 12px; + color: var(--text-dim); +} + +.pagination-actions { + display: flex; + gap: 8px; +} + +.retry-indicator { + font-size: 11px; + color: var(--warn); + animation: pulse 1s ease-in-out infinite; +} + @media (max-width: 1180px) { .task-layout { grid-template-columns: minmax(280px, 340px) minmax(0, 1fr); From 11a92c429fcf22100922e29334c5fe6ca1f04dc4 Mon Sep 17 00:00:00 2001 From: hezhaohui Date: Wed, 29 Jul 2026 14:13:53 +0800 Subject: [PATCH 3/3] =?UTF-8?q?test(server):=20=E6=B7=BB=E5=8A=A0=20task?= =?UTF-8?q?=20log=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() +}