feat(delivery): stream task updates with SSE

This commit is contained in:
mac
2026-07-28 09:55:04 +08:00
parent 7b4b9d2ee7
commit 5a260c443f
5 changed files with 168 additions and 4 deletions
+51
View File
@@ -2,7 +2,9 @@ package handler
import (
"crypto/subtle"
"encoding/json"
"errors"
"fmt"
"net/http"
"strconv"
"strings"
@@ -167,6 +169,55 @@ func (h *DeliveryHandler) Get(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"task": task, "events": events})
}
func (h *DeliveryHandler) Stream(c *gin.Context) {
claims, ok := CurrentClaims(c)
if !ok {
c.JSON(http.StatusUnauthorized, gin.H{"error": "missing current user"})
return
}
taskID := c.Param("id")
task, events, err := h.service.GetTask(c.Request.Context(), taskID, claims.UserID, claims.IsAdmin)
if errors.Is(err, gorm.ErrRecordNotFound) {
c.JSON(http.StatusNotFound, gin.H{"error": "task not found"})
return
}
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
updates, cancel := h.service.SubscribeTask(taskID)
defer cancel()
c.Header("Content-Type", "text/event-stream")
c.Header("Cache-Control", "no-cache")
c.Header("Connection", "keep-alive")
c.Header("X-Accel-Buffering", "no")
c.Status(http.StatusOK)
writeSSE(c.Writer, "snapshot", service.DeliveryTaskSnapshot{Task: task, Events: events})
c.Writer.Flush()
for {
select {
case <-c.Request.Context().Done():
return
case snapshot, ok := <-updates:
if !ok {
return
}
writeSSE(c.Writer, "snapshot", snapshot)
c.Writer.Flush()
}
}
}
func writeSSE(w http.ResponseWriter, event string, payload any) {
raw, err := json.Marshal(payload)
if err != nil {
return
}
_, _ = fmt.Fprintf(w, "event: %s\ndata: %s\n\n", event, raw)
}
// Cancel 取消交付任务
// @Summary 取消交付任务
// @Description 取消一个正在执行或等待中的交付任务
+1
View File
@@ -149,6 +149,7 @@ func registerAuthServerRoutes(r *gin.Engine, deps Dependencies) {
protected.POST("/delivery/mysql", deliveryHandler.CreateMySQL)
protected.GET("/delivery/tasks", deliveryHandler.List)
protected.GET("/delivery/tasks/:id", deliveryHandler.Get)
protected.GET("/delivery/tasks/:id/stream", deliveryHandler.Stream)
protected.POST("/delivery/tasks/:id/cancel", deliveryHandler.Cancel)
protected.GET("/task-logs", taskLogHandler.List)
protected.GET("/task-logs/:id", taskLogHandler.Get)
+64 -3
View File
@@ -94,6 +94,11 @@ type DeliveryTaskListFilter struct {
ActiveOnly bool
}
type DeliveryTaskSnapshot struct {
Task *model.DeliveryTask `json:"task"`
Events []model.TaskEvent `json:"events"`
}
// targetMetadata describes the native VM候选节点池以及部署形态,由 AWX inventory hosts 动态组装。
type targetMetadata struct {
Topology string `json:"topology"`
@@ -173,12 +178,14 @@ type DeliveryService struct {
awx *AWXClient
audit *AuditService
executionMu sync.Mutex
streamMu sync.Mutex
streams map[string]map[chan DeliveryTaskSnapshot]struct{}
}
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}
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{})}
}
func (s *DeliveryService) ListTargets(ctx context.Context, component string) ([]DeliveryTarget, error) {
@@ -312,6 +319,7 @@ func (s *DeliveryService) CreateTask(ctx context.Context, userID uint64, isAdmin
return nil, false, err
}
_ = s.db.WithContext(ctx).Create(&model.TaskEvent{TaskID: task.ID, ToState: model.TaskPending, Message: "delivery task created"}).Error
s.broadcastTask(ctx, task.ID)
return &task, false, nil
}
@@ -554,6 +562,51 @@ func (s *DeliveryService) GetTask(ctx context.Context, taskID string, userID uin
return &task, events, nil
}
func (s *DeliveryService) SubscribeTask(taskID string) (<-chan DeliveryTaskSnapshot, func()) {
ch := make(chan DeliveryTaskSnapshot, 8)
s.streamMu.Lock()
if s.streams[taskID] == nil {
s.streams[taskID] = make(map[chan DeliveryTaskSnapshot]struct{})
}
s.streams[taskID][ch] = struct{}{}
s.streamMu.Unlock()
cancel := func() {
s.streamMu.Lock()
if subscribers := s.streams[taskID]; subscribers != nil {
delete(subscribers, ch)
if len(subscribers) == 0 {
delete(s.streams, taskID)
}
}
s.streamMu.Unlock()
close(ch)
}
return ch, cancel
}
func (s *DeliveryService) taskSnapshot(ctx context.Context, taskID string) (DeliveryTaskSnapshot, error) {
task, events, err := s.GetTask(ctx, taskID, 0, true)
if err != nil {
return DeliveryTaskSnapshot{}, err
}
return DeliveryTaskSnapshot{Task: task, Events: events}, nil
}
func (s *DeliveryService) broadcastTask(ctx context.Context, taskID string) {
snapshot, err := s.taskSnapshot(ctx, taskID)
if err != nil {
return
}
s.streamMu.Lock()
defer s.streamMu.Unlock()
for ch := range s.streams[taskID] {
select {
case ch <- snapshot:
default:
}
}
}
func (s *DeliveryService) AWXJobStdout(ctx context.Context, jobID string) (string, error) {
return s.awx.JobStdout(ctx, jobID)
}
@@ -781,8 +834,11 @@ func (s *DeliveryService) DispatchOnce(ctx context.Context) error {
}
_, _, err = s.CreateExecution(ctx, task.ID, task.PayloadHash, task.IdempotencyKey)
if err != nil {
return s.failTask(ctx, task, model.TaskExecutionFailed, err.Error())
failErr := s.failTask(ctx, task, model.TaskExecutionFailed, err.Error())
s.broadcastTask(ctx, task.ID)
return failErr
}
s.broadcastTask(ctx, task.ID)
return nil
}
@@ -883,7 +939,7 @@ func (s *DeliveryService) HandleStageEvent(ctx context.Context, taskID string, i
}
eventState := "stage_" + stage + "_" + status
return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
var task model.DeliveryTask
if err := tx.First(&task, "id = ?", taskID).Error; err != nil {
return err
@@ -922,6 +978,10 @@ func (s *DeliveryService) HandleStageEvent(ctx context.Context, taskID string, i
}
return nil
})
if err == nil {
s.broadcastTask(ctx, taskID)
}
return err
}
func (s *DeliveryService) HandleAWXJobNotification(ctx context.Context, input AWXJobNotificationInput) error {
@@ -937,6 +997,7 @@ func (s *DeliveryService) HandleAWXJobNotification(ctx context.Context, input AW
return err
}
message := awxNotificationMessage(input)
defer s.broadcastTask(ctx, execution.TaskID)
switch status {
case "pending", "waiting", "running", "new":
return s.recordAWXEvent(ctx, execution.TaskID, status, message)