Merge pull request #125 from 2311719626/fix/logs

fix(server): 修复 task log 模块多个问题并添加单元测试
This commit is contained in:
Wonder
2026-07-29 14:38:55 +08:00
committed by GitHub
6 changed files with 1308 additions and 52 deletions
+98 -6
View File
@@ -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) {
@@ -992,18 +1007,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 {
@@ -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()
}