Files
xinfra/server/internal/service/delivery.go
T

1126 lines
40 KiB
Go

package service
import (
"bytes"
"context"
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"log"
"net"
"net/http"
"regexp"
"strings"
"sync"
"time"
"github.com/1024XEngineer/xinfra/server/internal/config"
"github.com/1024XEngineer/xinfra/server/internal/model"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
var dnsLabelPattern = regexp.MustCompile(`^[a-z0-9](?:[-a-z0-9]*[a-z0-9])?$`)
type MySQLDeliveryInput struct {
BusinessLineID uint64 `json:"business_line_id" binding:"required"`
TargetID uint64 `json:"target_id" binding:"required"`
Namespace string `json:"namespace" binding:"required"`
InstanceName string `json:"instance_name" binding:"required"`
MySQLVersion string `json:"mysql_version"`
Topology string `json:"topology"`
MySQLPort int `json:"mysql_port"`
DataDisk string `json:"data_disk"`
CPUCores int64 `json:"cpu_cores" binding:"required"`
MemoryGB int64 `json:"memory_gb" binding:"required"`
StorageGB int64 `json:"storage_gb" binding:"required"`
ParamTemplate string `json:"param_template"`
TimeZone string `json:"timezone"`
LowerCaseTableNames int `json:"lower_case_table_names"`
CharacterSet string `json:"character_set"`
Collation string `json:"collation"`
MaxConnections string `json:"max_connections"`
InnoDBRedoLogCapacity string `json:"innodb_redo_log_capacity"`
InnoDBFlushLogAtTrxCommit int `json:"innodb_flush_log_at_trx_commit"`
SyncBinlog int `json:"sync_binlog"`
InnoDBIOCapacity int `json:"innodb_io_capacity"`
LongQueryTime float64 `json:"long_query_time"`
BinlogExpireLogsSeconds int64 `json:"binlog_expire_logs_seconds"`
MaxBinlogSize string `json:"max_binlog_size"`
CPUMilli int64 `json:"-"`
MemoryMi int64 `json:"-"`
StorageGi int64 `json:"-"`
}
type deliveryPayload struct {
MySQLDeliveryInput
TargetType string `json:"target_type"`
}
type DeliveryTarget struct {
ID uint64 `json:"id"`
Name string `json:"name"`
TargetType string `json:"target_type"`
AWXInventoryID uint64 `json:"awx_inventory_id"`
AWXTemplateID uint64 `json:"awx_template_id"`
Enabled bool `json:"enabled"`
Metadata string `json:"metadata"`
}
type DeliveryStageEventInput struct {
Stage string `json:"stage" binding:"required"`
Status string `json:"status" binding:"required"`
Message string `json:"message"`
AWXJobID string `json:"awx_job_id"`
}
type AWXJobNotificationInput struct {
ID uint64 `json:"id"`
Status string `json:"status"`
Name string `json:"name"`
URL string `json:"url"`
Traceback string `json:"traceback"`
ExtraVars any `json:"extra_vars"`
}
type DeliveryTaskListFilter struct {
BusinessLineID uint64
Component string
ActiveOnly bool
}
// targetMetadata describes the native VM候选节点池以及部署形态,由 AWX inventory hosts 动态组装。
type targetMetadata struct {
Topology string `json:"topology"`
MySQLPort int `json:"mysql_port"`
Hosts []targetHost `json:"hosts"`
}
type targetHost struct {
Name string `json:"name"`
IP string `json:"ip"`
}
func parseTargetMetadata(raw string) targetMetadata {
meta := targetMetadata{}
if strings.TrimSpace(raw) != "" {
_ = json.Unmarshal([]byte(raw), &meta)
}
if meta.Topology == "" {
meta.Topology = "standalone"
}
if meta.MySQLPort == 0 {
meta.MySQLPort = 13306
}
return meta
}
// firstFreeHost 返回候选池中第一个未被占用的节点。
func firstFreeHost(hosts []targetHost, occupied []string) *targetHost {
taken := make(map[string]bool, len(occupied))
for _, h := range occupied {
taken[h] = true
}
for i := range hosts {
if !taken[hosts[i].Name] {
return &hosts[i]
}
}
return nil
}
func firstFreePort(occupied []int) int {
taken := make(map[int]bool, len(occupied))
for _, port := range occupied {
taken[port] = true
}
for port := 13306; port <= 13999; port++ {
if !taken[port] {
return port
}
}
return 0
}
type DeliveryService struct {
db *gorm.DB
cfg config.Config
awx *AWXClient
audit *AuditService
executionMu sync.Mutex
}
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}
}
func (s *DeliveryService) ListTargets(ctx context.Context, component string) ([]DeliveryTarget, error) {
templates, err := s.awx.ListJobTemplates(ctx)
if err != nil {
return nil, err
}
component = strings.ToLower(strings.TrimSpace(component))
var targets []DeliveryTarget
for _, template := range templates {
if template.Inventory == 0 {
continue
}
if component != "" && component != "all" {
text := strings.ToLower(template.Name + " " + template.Description)
if !strings.Contains(text, component) {
continue
}
}
target, err := s.awxDeliveryTarget(ctx, template)
if err != nil {
continue
}
targets = append(targets, target)
}
return targets, nil
}
func (s *DeliveryService) getTarget(ctx context.Context, templateID uint64) (DeliveryTarget, error) {
template, err := s.awx.GetJobTemplate(ctx, templateID)
if err != nil {
return DeliveryTarget{}, fmt.Errorf("deployment target is unavailable: %w", err)
}
return s.awxDeliveryTarget(ctx, *template)
}
func (s *DeliveryService) awxDeliveryTarget(ctx context.Context, template AWXJobTemplate) (DeliveryTarget, error) {
if !template.AskVariablesOnLaunch || !template.AskLimitOnLaunch {
return DeliveryTarget{}, fmt.Errorf("AWX job template %d must enable Prompt on launch for Variables and Limit", template.ID)
}
hosts, err := s.awx.ListInventoryHosts(ctx, template.Inventory)
if err != nil {
return DeliveryTarget{}, err
}
meta := targetMetadata{Topology: "standalone", MySQLPort: 13306}
for _, host := range hosts {
if !host.Enabled {
continue
}
meta.Hosts = append(meta.Hosts, targetHost{Name: host.Name, IP: AWXHostIP(host)})
}
raw, err := json.Marshal(meta)
if err != nil {
return DeliveryTarget{}, err
}
return DeliveryTarget{
ID: template.ID,
Name: template.Name,
TargetType: "k8s",
AWXInventoryID: template.Inventory,
AWXTemplateID: template.ID,
Enabled: true,
Metadata: string(raw),
}, nil
}
func (s *DeliveryService) CreateTask(ctx context.Context, userID uint64, isAdmin bool, idempotencyKey string, input MySQLDeliveryInput) (*model.DeliveryTask, bool, error) {
idempotencyKey = strings.TrimSpace(idempotencyKey)
if idempotencyKey == "" || len(idempotencyKey) > 128 {
return nil, false, fmt.Errorf("Idempotency-Key header is required and must not exceed 128 characters")
}
normalizeMySQLDeliveryInput(&input)
if err := validateDeliveryInput(input); err != nil {
return nil, false, err
}
var existing model.DeliveryTask
if err := s.db.WithContext(ctx).Where("idempotency_key = ?", idempotencyKey).First(&existing).Error; err == nil {
if existing.RequestedBy != userID {
return nil, false, fmt.Errorf("idempotency key is already in use by another user")
}
return &existing, true, nil
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, false, err
}
target, err := s.getTarget(ctx, input.TargetID)
if err != nil {
return nil, false, err
}
if target.TargetType != "k8s" {
return nil, false, fmt.Errorf("target type %q is not supported in the MVP", target.TargetType)
}
if !isAdmin {
var count int64
if err := s.db.WithContext(ctx).Model(&model.BusinessLineUser{}).
Where("business_line_id = ? AND user_id = ?", input.BusinessLineID, userID).Count(&count).Error; err != nil {
return nil, false, err
}
if count == 0 {
return nil, false, fmt.Errorf("user is not authorized for this business line")
}
}
payload := deliveryPayload{MySQLDeliveryInput: input, TargetType: target.TargetType}
raw, err := json.Marshal(payload)
if err != nil {
return nil, false, err
}
digest := sha256.Sum256(raw)
task := model.DeliveryTask{
ID: randomUUID(),
BusinessLineID: input.BusinessLineID,
RequestedBy: userID,
Component: "mysql",
TargetType: target.TargetType,
TargetID: target.ID,
Namespace: input.Namespace,
InstanceName: input.InstanceName,
Status: model.TaskPending,
ImmutablePayload: string(raw),
PayloadHash: hex.EncodeToString(digest[:]),
IdempotencyKey: idempotencyKey,
}
if err := s.db.WithContext(ctx).Create(&task).Error; err != nil {
if lookupErr := s.db.WithContext(ctx).Where("idempotency_key = ?", idempotencyKey).First(&existing).Error; lookupErr == nil {
return &existing, true, nil
}
return nil, false, err
}
_ = s.db.WithContext(ctx).Create(&model.TaskEvent{TaskID: task.ID, ToState: model.TaskPending, Message: "delivery task created"}).Error
return &task, false, nil
}
var supportedMySQLVersions = map[string]bool{"8.0": true}
var supportedMySQLTopologies = map[string]bool{"standalone": true}
var supportedDataDisks = map[string]bool{"/data": true, "/disk1": true, "/mnt/vol-1": true}
var supportedParamTemplates = map[string]bool{"default": true, "high_performance": true, "high_safety": true}
var supportedCharacterSets = map[string]bool{"utf8mb4": true, "utf8": true, "gbk": true, "latin1": true}
var supportedCollations = map[string]bool{
"utf8mb4_general_ci": true, "utf8mb4_unicode_ci": true, "utf8mb4_0900_ai_ci": true,
"utf8_general_ci": true, "gbk_chinese_ci": true, "latin1_swedish_ci": true,
}
var supportedRedoLogCapacity = map[string]bool{"auto": true, "128M": true, "256M": true, "512M": true, "1G": true}
var supportedMaxBinlogSize = map[string]bool{"128M": true, "256M": true, "512M": true, "1G": true}
func normalizeMySQLDeliveryInput(input *MySQLDeliveryInput) {
if input.MySQLVersion == "" {
input.MySQLVersion = "8.0"
}
if input.Topology == "" {
input.Topology = "standalone"
}
if input.DataDisk == "" {
input.DataDisk = "/data"
}
if input.ParamTemplate == "" {
input.ParamTemplate = "default"
}
if input.TimeZone == "" {
input.TimeZone = "+08:00"
}
if input.CharacterSet == "" {
input.CharacterSet = "utf8mb4"
}
if input.Collation == "" {
input.Collation = defaultCollation(input.CharacterSet)
}
if input.MaxConnections == "" {
input.MaxConnections = "auto"
}
if input.InnoDBRedoLogCapacity == "" {
input.InnoDBRedoLogCapacity = "auto"
}
if input.InnoDBFlushLogAtTrxCommit == 0 {
input.InnoDBFlushLogAtTrxCommit = 1
}
if input.SyncBinlog == 0 {
input.SyncBinlog = 1
}
if input.InnoDBIOCapacity == 0 {
input.InnoDBIOCapacity = 2000
}
if input.LongQueryTime == 0 {
input.LongQueryTime = 1
}
if input.BinlogExpireLogsSeconds == 0 {
input.BinlogExpireLogsSeconds = 604800
}
if input.MaxBinlogSize == "" {
input.MaxBinlogSize = "256M"
}
input.CPUMilli = input.CPUCores * 1000
input.MemoryMi = input.MemoryGB * 1024
input.StorageGi = input.StorageGB
}
func defaultCollation(characterSet string) string {
switch characterSet {
case "utf8":
return "utf8_general_ci"
case "gbk":
return "gbk_chinese_ci"
case "latin1":
return "latin1_swedish_ci"
default:
return "utf8mb4_general_ci"
}
}
func validateDeliveryInput(input MySQLDeliveryInput) error {
if len(input.Namespace) > 63 || !dnsLabelPattern.MatchString(input.Namespace) {
return fmt.Errorf("namespace must be a valid Kubernetes DNS label")
}
if len(input.InstanceName) > 63 || !dnsLabelPattern.MatchString(input.InstanceName) {
return fmt.Errorf("instance_name must be a valid Kubernetes DNS label")
}
if !oneOfInt64(input.CPUCores, []int64{1, 2, 4, 8, 16}) {
return fmt.Errorf("cpu_cores must be one of 1, 2, 4, 8, 16")
}
if !oneOfInt64(input.MemoryGB, []int64{2, 4, 8, 16, 32, 64}) {
return fmt.Errorf("memory_gb must be one of 2, 4, 8, 16, 32, 64")
}
if input.StorageGB < 20 || input.StorageGB > 2000 {
return fmt.Errorf("storage_gb must be between 20 and 2000")
}
if input.MySQLVersion != "" && !supportedMySQLVersions[input.MySQLVersion] {
return fmt.Errorf("unsupported mysql_version %q, supported: 8.0", input.MySQLVersion)
}
if input.Topology != "" && !supportedMySQLTopologies[input.Topology] {
return fmt.Errorf("unsupported topology %q, supported: standalone", input.Topology)
}
if input.MySQLPort != 0 && (input.MySQLPort < 13306 || input.MySQLPort > 13999) {
return fmt.Errorf("mysql_port must be empty for auto assignment or between 13306 and 13999")
}
if !supportedDataDisks[input.DataDisk] {
return fmt.Errorf("unsupported data_disk %q", input.DataDisk)
}
if !supportedParamTemplates[input.ParamTemplate] {
return fmt.Errorf("unsupported param_template %q", input.ParamTemplate)
}
if !validTimeZone(input.TimeZone) {
return fmt.Errorf("unsupported timezone %q", input.TimeZone)
}
if input.LowerCaseTableNames != 0 && input.LowerCaseTableNames != 1 {
return fmt.Errorf("lower_case_table_names must be 0 or 1")
}
if !supportedCharacterSets[input.CharacterSet] {
return fmt.Errorf("unsupported character_set %q", input.CharacterSet)
}
if !supportedCollations[input.Collation] || !strings.HasPrefix(input.Collation, input.CharacterSet+"_") {
return fmt.Errorf("collation %q is not valid for character_set %q", input.Collation, input.CharacterSet)
}
if !validMaxConnections(input.MaxConnections) {
return fmt.Errorf("max_connections must be auto or one of 200, 500, 1000, 2000, 4000, 8000, 16000")
}
if !supportedRedoLogCapacity[input.InnoDBRedoLogCapacity] {
return fmt.Errorf("unsupported innodb_redo_log_capacity %q", input.InnoDBRedoLogCapacity)
}
if !oneOfInt(input.InnoDBFlushLogAtTrxCommit, []int{0, 1, 2}) {
return fmt.Errorf("innodb_flush_log_at_trx_commit must be one of 0, 1, 2")
}
if input.SyncBinlog != 0 && input.SyncBinlog != 1 {
return fmt.Errorf("sync_binlog must be 0 or 1")
}
if !oneOfInt(input.InnoDBIOCapacity, []int{200, 2000, 5000}) {
return fmt.Errorf("innodb_io_capacity must be one of 200, 2000, 5000")
}
if !oneOfFloat(input.LongQueryTime, []float64{0.5, 1, 2, 5, 10}) {
return fmt.Errorf("long_query_time must be one of 0.5, 1, 2, 5, 10")
}
if !oneOfInt64(input.BinlogExpireLogsSeconds, []int64{86400, 259200, 604800, 1209600}) {
return fmt.Errorf("binlog_expire_logs_seconds must be one of 86400, 259200, 604800, 1209600")
}
if !supportedMaxBinlogSize[input.MaxBinlogSize] {
return fmt.Errorf("unsupported max_binlog_size %q", input.MaxBinlogSize)
}
return nil
}
func oneOfInt(value int, allowed []int) bool {
for _, item := range allowed {
if value == item {
return true
}
}
return false
}
func oneOfInt64(value int64, allowed []int64) bool {
for _, item := range allowed {
if value == item {
return true
}
}
return false
}
func oneOfFloat(value float64, allowed []float64) bool {
for _, item := range allowed {
if value == item {
return true
}
}
return false
}
func validTimeZone(value string) bool {
return value == "SYSTEM" || value == "+08:00" || value == "+00:00" || value == "Asia/Shanghai"
}
func validMaxConnections(value string) bool {
if value == "auto" {
return true
}
switch value {
case "200", "500", "1000", "2000", "4000", "8000", "16000":
return true
default:
return false
}
}
func (s *DeliveryService) ListTasks(ctx context.Context, userID uint64, isAdmin bool, filter DeliveryTaskListFilter) ([]model.DeliveryTask, error) {
query := s.db.WithContext(ctx).Order("created_at DESC")
if filter.BusinessLineID != 0 {
query = query.Where("business_line_id = ?", filter.BusinessLineID)
}
if filter.Component != "" {
query = query.Where("component = ?", strings.ToLower(strings.TrimSpace(filter.Component)))
}
if filter.ActiveOnly {
query = query.Where("status NOT IN ?", terminalTaskStatuses())
}
if !isAdmin {
query = query.Where("business_line_id IN (?)", s.db.Model(&model.BusinessLineUser{}).Select("business_line_id").Where("user_id = ?", userID))
}
var tasks []model.DeliveryTask
return tasks, query.Find(&tasks).Error
}
func (s *DeliveryService) GetTask(ctx context.Context, taskID string, userID uint64, isAdmin bool) (*model.DeliveryTask, []model.TaskEvent, error) {
query := s.db.WithContext(ctx).Where("id = ?", taskID)
if !isAdmin {
query = query.Where("business_line_id IN (?)", s.db.Model(&model.BusinessLineUser{}).Select("business_line_id").Where("user_id = ?", userID))
}
var task model.DeliveryTask
if err := query.First(&task).Error; err != nil {
return nil, nil, err
}
var events []model.TaskEvent
if err := s.db.WithContext(ctx).Where("task_id = ?", taskID).Order("id ASC").Find(&events).Error; err != nil {
return nil, nil, err
}
return &task, events, nil
}
func (s *DeliveryService) AWXJobStdout(ctx context.Context, jobID string) (string, error) {
return s.awx.JobStdout(ctx, jobID)
}
func (s *DeliveryService) Cancel(ctx context.Context, taskID string, userID uint64, isAdmin bool) error {
task, _, err := s.GetTask(ctx, taskID, userID, isAdmin)
if err != nil {
return err
}
if task.Status == model.TaskPending {
return s.transition(ctx, task, model.TaskCanceled, "canceled before dispatch", "")
}
if task.Status != model.TaskRunning && task.Status != model.TaskDispatching {
return fmt.Errorf("task in state %q cannot be canceled", task.Status)
}
var job model.ExecutionJob
if err := s.db.WithContext(ctx).Where("task_id = ?", task.ID).First(&job).Error; err != nil {
return err
}
if err := s.awx.Cancel(ctx, job.ExecutorJobID); err != nil {
return err
}
return s.transition(ctx, task, model.TaskCanceling, "cancel requested in AWX", "")
}
func (s *DeliveryService) claimAndReserve(ctx context.Context) (*model.DeliveryTask, error) {
var task model.DeliveryTask
var payload deliveryPayload
var target DeliveryTarget
dispatchable := false
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Clauses(clause.Locking{Strength: "UPDATE", Options: "SKIP LOCKED"}).Where("status = ?", model.TaskPending).Order("created_at DESC").First(&task).Error; err != nil {
return err
}
var targetErr error
target, targetErr = s.getTarget(ctx, task.TargetID)
if targetErr != nil {
return s.failInTransaction(tx, &task, model.TaskValidationFailed, targetErr.Error())
}
if err := json.Unmarshal([]byte(task.ImmutablePayload), &payload); err != nil {
return s.failInTransaction(tx, &task, model.TaskValidationFailed, "stored deployment payload is invalid")
}
normalizeMySQLDeliveryInput(&payload.MySQLDeliveryInput)
activeStates := []string{model.TaskValidating, model.TaskDispatching, model.TaskRunning, model.TaskCanceling}
checks := []struct {
query string
args []any
limit int
message string
}{
{"status IN ?", []any{activeStates}, s.cfg.DeliveryGlobalLimit, "global concurrency limit reached"},
{"status IN ? AND target_id = ?", []any{activeStates, task.TargetID}, s.cfg.DeliveryTargetLimit, "target concurrency limit reached"},
}
for _, check := range checks {
var count int64
if check.limit > 0 {
if err := tx.Model(&model.DeliveryTask{}).Where(check.query, check.args...).Count(&count).Error; err != nil {
return err
}
if count >= int64(check.limit) {
return fmt.Errorf("defer: %s", check.message)
}
}
}
quotaOK, err := checkResourceQuota(tx, task.BusinessLineID, task.TargetID, payload)
if err != nil {
return err
}
if !quotaOK {
return s.failInTransaction(tx, &task, model.TaskValidationFailed, "resource quota is insufficient")
}
meta := parseTargetMetadata(target.Metadata)
if len(meta.Hosts) == 0 {
return s.failInTransaction(tx, &task, model.TaskValidationFailed, "deployment target has no candidate hosts")
}
var occupied []string
occupiedExclude := []string{model.TaskExecutionFailed, model.TaskValidationFailed, model.TaskCanceled}
if err := tx.Model(&model.DeliveryTask{}).Where("target_id = ? AND target_host <> ? AND status NOT IN ?", task.TargetID, "", occupiedExclude).Pluck("target_host", &occupied).Error; err != nil {
return err
}
host := firstFreeHost(meta.Hosts, occupied)
if host == nil {
return fmt.Errorf("defer: no free host available on target")
}
mysqlPort := payload.MySQLPort
if mysqlPort == 0 {
var occupiedPorts []int
if err := tx.Model(&model.DeliveryTask{}).Where("target_id = ? AND target_host = ? AND mysql_port <> ? AND status NOT IN ?", task.TargetID, host.Name, 0, occupiedExclude).Pluck("mysql_port", &occupiedPorts).Error; err != nil {
return err
}
mysqlPort = firstFreePort(occupiedPorts)
if mysqlPort == 0 {
return fmt.Errorf("defer: no free MySQL port available on target host")
}
}
reservation := model.ResourceReservation{TaskID: task.ID, BusinessLineID: task.BusinessLineID, TargetID: task.TargetID, CPUMilli: payload.CPUMilli, MemoryMi: payload.MemoryMi, StorageGi: payload.StorageGi, InstanceCount: 1, Status: "reserved", ExpiresAt: time.Now().Add(time.Duration(s.cfg.ReservationTTLMinutes) * time.Minute)}
if err := tx.Create(&reservation).Error; err != nil {
return err
}
if err := tx.Model(&model.DeliveryTask{}).Where("id = ?", task.ID).Updates(map[string]any{"target_host": host.Name, "target_host_ip": host.IP, "mysql_port": mysqlPort}).Error; err != nil {
return err
}
task.TargetHost = host.Name
task.TargetHostIP = host.IP
task.MySQLPort = mysqlPort
if err := s.transitionTx(tx, &task, model.TaskDispatching, "resources reserved", ""); err != nil {
return err
}
dispatchable = true
return nil
})
if err == nil && !dispatchable {
err = gorm.ErrRecordNotFound
}
return &task, err
}
func checkResourceQuota(tx *gorm.DB, businessLineID, targetID uint64, payload deliveryPayload) (bool, error) {
var quota model.ResourceQuota
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Where("business_line_id = ? AND target_id = ?", businessLineID, targetID).First(&quota).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return true, nil
}
return false, err
}
type totals struct{ CPU, Memory, Storage, Instances int64 }
var used, reserved totals
if err := tx.Model(&model.ResourceUsage{}).Select("COALESCE(SUM(cpu_milli),0) cpu, COALESCE(SUM(memory_mi),0) memory, COALESCE(SUM(storage_gi),0) storage, COALESCE(SUM(instance_count),0) instances").Where("business_line_id = ? AND target_id = ? AND status = ?", businessLineID, targetID, "active").Scan(&used).Error; err != nil {
return false, err
}
if err := tx.Model(&model.ResourceReservation{}).Select("COALESCE(SUM(cpu_milli),0) cpu, COALESCE(SUM(memory_mi),0) memory, COALESCE(SUM(storage_gi),0) storage, COALESCE(SUM(instance_count),0) instances").Where("business_line_id = ? AND target_id = ? AND status = ? AND expires_at > ?", businessLineID, targetID, "reserved", time.Now()).Scan(&reserved).Error; err != nil {
return false, err
}
return used.CPU+reserved.CPU+payload.CPUMilli <= quota.CPUMilli &&
used.Memory+reserved.Memory+payload.MemoryMi <= quota.MemoryMi &&
used.Storage+reserved.Storage+payload.StorageGi <= quota.StorageGi &&
used.Instances+reserved.Instances+1 <= quota.InstanceLimit, nil
}
func (s *DeliveryService) failInTransaction(tx *gorm.DB, task *model.DeliveryTask, status, message string) error {
if err := s.transitionTx(tx, task, status, message, message); err != nil {
return err
}
return nil
}
func (s *DeliveryService) transition(ctx context.Context, task *model.DeliveryTask, status, message, errorMessage string) error {
return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { return s.transitionTx(tx, task, status, message, errorMessage) })
}
func (s *DeliveryService) transitionTx(tx *gorm.DB, task *model.DeliveryTask, status, message, errorMessage string) error {
from := task.Status
updates := map[string]any{"status": status, "error_message": errorMessage}
now := time.Now()
if status == model.TaskRunning {
updates["started_at"] = now
}
if status == model.TaskFinished || status == model.TaskExecutionFailed || status == model.TaskValidationFailed || status == model.TaskCanceled || status == model.TaskRegisterFailed {
updates["finished_at"] = now
}
result := tx.Model(&model.DeliveryTask{}).Where("id = ? AND status = ?", task.ID, from).Updates(updates)
if result.Error != nil {
return result.Error
}
if result.RowsAffected != 1 {
return fmt.Errorf("task %s changed concurrently", task.ID)
}
task.Status = status
task.ErrorMessage = errorMessage
return tx.Create(&model.TaskEvent{TaskID: task.ID, FromState: from, ToState: status, Message: message}).Error
}
func (s *DeliveryService) releaseReservation(tx *gorm.DB, taskID string) error {
return tx.Model(&model.ResourceReservation{}).Where("task_id = ? AND status = ?", taskID, "reserved").Update("status", "released").Error
}
func randomUUID() string {
b := make([]byte, 16)
if _, err := rand.Read(b); err != nil {
panic(err)
}
b[6] = (b[6] & 0x0f) | 0x40
b[8] = (b[8] & 0x3f) | 0x80
return fmt.Sprintf("%08x-%04x-%04x-%04x-%012x", b[0:4], b[4:6], b[6:8], b[8:10], b[10:16])
}
func mysqlReady(ctx context.Context, address string) error {
dialer := net.Dialer{Timeout: 5 * time.Second}
conn, err := dialer.DialContext(ctx, "tcp", address)
if err != nil {
return err
}
return conn.Close()
}
func (s *DeliveryService) DispatchOnce(ctx context.Context) error {
task, err := s.claimAndReserve(ctx)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) || strings.HasPrefix(err.Error(), "defer:") {
return nil
}
return err
}
_, _, err = s.CreateExecution(ctx, task.ID, task.PayloadHash, task.IdempotencyKey)
if err != nil {
return s.failTask(ctx, task, model.TaskExecutionFailed, err.Error())
}
return nil
}
func (s *DeliveryService) CreateExecution(ctx context.Context, taskID, payloadHash, idempotencyKey string) (*model.ExecutionJob, bool, error) {
s.executionMu.Lock()
defer s.executionMu.Unlock()
var existing model.ExecutionJob
if err := s.db.WithContext(ctx).Where("task_id = ? OR idempotency_key = ?", taskID, idempotencyKey).First(&existing).Error; err == nil {
if existing.TaskID != taskID || existing.IdempotencyKey != idempotencyKey {
return nil, false, fmt.Errorf("idempotency key is already bound to another task")
}
return &existing, true, nil
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, false, err
}
var task model.DeliveryTask
if err := s.db.WithContext(ctx).First(&task, "id = ?", taskID).Error; err != nil {
return nil, false, err
}
if task.PayloadHash != payloadHash || task.IdempotencyKey != idempotencyKey {
return nil, false, fmt.Errorf("execution request does not match the immutable task payload")
}
if task.Status != model.TaskDispatching {
return nil, false, fmt.Errorf("task in state %q is not ready for execution", task.Status)
}
target, err := s.getTarget(ctx, task.TargetID)
if err != nil {
return nil, false, err
}
var payload deliveryPayload
if err := json.Unmarshal([]byte(task.ImmutablePayload), &payload); err != nil {
return nil, false, err
}
normalizeMySQLDeliveryInput(&payload.MySQLDeliveryInput)
meta := parseTargetMetadata(target.Metadata)
topology := payload.Topology
if topology == "" {
topology = meta.Topology
}
// Persist execution record BEFORE launching AWX to ensure crash recovery.
now := time.Now()
execution := model.ExecutionJob{TaskID: task.ID, IdempotencyKey: task.IdempotencyKey, ExecutorJobID: pendingExecutorJobID(task.ID), Status: "launching", StartedAt: &now}
if err := s.db.WithContext(ctx).Create(&execution).Error; err != nil {
return nil, false, err
}
job, err := s.awx.Launch(ctx, target.AWXTemplateID, AWXLaunchRequest{InventoryID: target.AWXInventoryID, Limit: task.TargetHost, ExtraVars: map[string]any{
"task_id": task.ID, "payload_hash": task.PayloadHash,
"target_hosts": task.TargetHost, "topology": topology,
"instance_name": payload.InstanceName, "mysql_port": task.MySQLPort,
"data_disk": payload.DataDisk, "cpu_cores": payload.CPUCores,
"memory_gb": payload.MemoryGB, "storage_gb": payload.StorageGB,
"mysql_version": payload.MySQLVersion, "param_template": payload.ParamTemplate,
"timezone": payload.TimeZone, "lower_case_table_names": payload.LowerCaseTableNames,
"character_set": payload.CharacterSet, "collation": payload.Collation,
"max_connections": payload.MaxConnections, "innodb_redo_log_capacity": payload.InnoDBRedoLogCapacity,
"innodb_flush_log_at_trx_commit": payload.InnoDBFlushLogAtTrxCommit, "sync_binlog": payload.SyncBinlog,
"innodb_io_capacity": payload.InnoDBIOCapacity, "long_query_time": payload.LongQueryTime,
"binlog_expire_logs_seconds": payload.BinlogExpireLogsSeconds, "max_binlog_size": payload.MaxBinlogSize,
"delivery_callback_url": s.deliveryCallbackURL(task.ID), "delivery_callback_token": s.cfg.AWXWebhookToken,
}})
if err != nil {
_ = s.db.WithContext(ctx).Model(&execution).Updates(map[string]any{"status": "failed", "finished_at": time.Now()})
return nil, false, err
}
if err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Model(&execution).Updates(map[string]any{"executor_job_id": fmt.Sprint(job.ID), "status": "running"}).Error; err != nil {
return err
}
return s.transitionTx(tx, &task, model.TaskRunning, "AWX job started", "")
}); err != nil {
return nil, false, err
}
return &execution, false, nil
}
func (s *DeliveryService) deliveryCallbackURL(taskID string) string {
base := strings.TrimRight(strings.TrimSpace(s.cfg.DeliveryCallbackBaseURL), "/")
if base == "" {
return ""
}
return base + "/auth/internal/delivery/tasks/" + taskID + "/events"
}
func (s *DeliveryService) HandleStageEvent(ctx context.Context, taskID string, input DeliveryStageEventInput) error {
stage := strings.ToLower(strings.TrimSpace(input.Stage))
status := strings.ToLower(strings.TrimSpace(input.Status))
if !validDeliveryStage(stage) {
return fmt.Errorf("invalid delivery stage %q", input.Stage)
}
if !validDeliveryStageStatus(status) {
return fmt.Errorf("invalid delivery stage status %q", input.Status)
}
message := strings.TrimSpace(input.Message)
if message == "" {
message = stage + " " + status
}
eventState := "stage_" + stage + "_" + status
return 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
}
if input.AWXJobID != "" {
var execution model.ExecutionJob
if err := tx.Where("task_id = ?", task.ID).First(&execution).Error; err != nil {
return err
}
if execution.ExecutorJobID != strings.TrimSpace(input.AWXJobID) {
return fmt.Errorf("AWX job %s does not match task %s", input.AWXJobID, task.ID)
}
}
if isTerminalTaskStatus(task.Status) {
return nil
}
if err := tx.Create(&model.TaskEvent{
TaskID: task.ID,
FromState: task.Status,
ToState: eventState,
Stage: stage,
EventStatus: status,
Message: message,
}).Error; err != nil {
return err
}
if status == "failed" {
now := time.Now()
if err := tx.Model(&model.ExecutionJob{}).Where("task_id = ?", task.ID).Updates(map[string]any{"status": "failed", "finished_at": now}).Error; err != nil {
return err
}
if err := s.transitionTx(tx, &task, model.TaskExecutionFailed, message, message); err != nil {
return err
}
return s.releaseReservation(tx, task.ID)
}
return nil
})
}
func (s *DeliveryService) HandleAWXJobNotification(ctx context.Context, input AWXJobNotificationInput) error {
status := strings.ToLower(strings.TrimSpace(input.Status))
if status == "" && input.ID == 0 && awxNotificationTaskID(input.ExtraVars) == "" {
return nil
}
if status == "" {
return fmt.Errorf("missing AWX job status")
}
execution, err := s.findExecutionForAWXNotification(ctx, input)
if err != nil {
return err
}
message := awxNotificationMessage(input)
switch status {
case "pending", "waiting", "running", "new":
return s.recordAWXEvent(ctx, execution.TaskID, status, message)
case "successful":
s.finishExecution(ctx, execution, "successful")
if err := s.completeTask(ctx, execution.TaskID); err != nil {
_ = s.failTask(ctx, &model.DeliveryTask{ID: execution.TaskID, Status: model.TaskRunning}, model.TaskValidationFailed, err.Error())
return err
}
return nil
case "canceled", "cancelled":
s.finishExecution(ctx, execution, "canceled")
return s.failTask(ctx, &model.DeliveryTask{ID: execution.TaskID, Status: model.TaskRunning}, model.TaskCanceled, message)
case "failed", "error":
s.finishExecution(ctx, execution, "failed")
return s.failTask(ctx, &model.DeliveryTask{ID: execution.TaskID, Status: model.TaskRunning}, model.TaskExecutionFailed, message)
default:
s.finishExecution(ctx, execution, "failed")
return s.failTask(ctx, &model.DeliveryTask{ID: execution.TaskID, Status: model.TaskRunning}, model.TaskExecutionFailed, "AWX job finished with status "+status)
}
}
func (s *DeliveryService) findExecutionForAWXNotification(ctx context.Context, input AWXJobNotificationInput) (*model.ExecutionJob, error) {
var execution model.ExecutionJob
if input.ID != 0 {
if err := s.db.WithContext(ctx).Where("executor_job_id = ?", fmt.Sprint(input.ID)).First(&execution).Error; err == nil {
return &execution, nil
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err
}
}
taskID := awxNotificationTaskID(input.ExtraVars)
if taskID == "" {
return nil, fmt.Errorf("AWX notification did not include a known job id or task_id")
}
if err := s.db.WithContext(ctx).Where("task_id = ?", taskID).First(&execution).Error; err != nil {
return nil, err
}
return &execution, nil
}
func (s *DeliveryService) recordAWXEvent(ctx context.Context, taskID, status, message string) error {
var task model.DeliveryTask
if err := s.db.WithContext(ctx).First(&task, "id = ?", taskID).Error; err != nil {
return err
}
if isTerminalTaskStatus(task.Status) {
return nil
}
return s.db.WithContext(ctx).Create(&model.TaskEvent{
TaskID: task.ID,
FromState: task.Status,
ToState: "awx_" + status,
Stage: "awx",
EventStatus: status,
Message: message,
}).Error
}
func validDeliveryStage(stage string) bool {
switch stage {
case "precheck", "install", "configure", "healthcheck", "register":
return true
default:
return false
}
}
func validDeliveryStageStatus(status string) bool {
switch status {
case "running", "success", "failed":
return true
default:
return false
}
}
func isTerminalTaskStatus(status string) bool {
for _, item := range terminalTaskStatuses() {
if status == item {
return true
}
}
return false
}
func terminalTaskStatuses() []string {
return []string{model.TaskFinished, model.TaskExecutionFailed, model.TaskValidationFailed, model.TaskRegisterFailed, model.TaskCanceled}
}
func awxNotificationMessage(input AWXJobNotificationInput) string {
status := strings.TrimSpace(input.Status)
name := strings.TrimSpace(input.Name)
if strings.TrimSpace(input.Traceback) != "" {
return strings.TrimSpace(input.Traceback)
}
if name == "" {
return "AWX job " + status
}
return "AWX job " + name + " " + status
}
func awxNotificationTaskID(extraVars any) string {
var values map[string]any
switch v := extraVars.(type) {
case map[string]any:
values = v
case string:
if strings.TrimSpace(v) == "" {
return ""
}
_ = json.Unmarshal([]byte(v), &values)
default:
raw, err := json.Marshal(v)
if err != nil {
return ""
}
_ = json.Unmarshal(raw, &values)
}
if values == nil {
return ""
}
if taskID, ok := values["task_id"].(string); ok {
return strings.TrimSpace(taskID)
}
return ""
}
func (s *DeliveryService) finishExecution(ctx context.Context, execution *model.ExecutionJob, status string) {
now := time.Now()
_ = s.db.WithContext(ctx).Model(execution).Updates(map[string]any{"status": status, "finished_at": now})
}
func (s *DeliveryService) completeTask(ctx context.Context, taskID string) error {
var task model.DeliveryTask
if err := s.db.WithContext(ctx).First(&task, "id = ?", taskID).Error; err != nil {
return err
}
if isTerminalTaskStatus(task.Status) {
return nil
}
var payload deliveryPayload
if err := json.Unmarshal([]byte(task.ImmutablePayload), &payload); err != nil {
return err
}
normalizeMySQLDeliveryInput(&payload.MySQLDeliveryInput)
addr := fmt.Sprintf("%s:%d", task.TargetHostIP, task.MySQLPort)
if err := mysqlReady(ctx, addr); err != nil {
return fmt.Errorf("MySQL health check failed: %w", err)
}
now := time.Now()
if err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if err := s.transitionTx(tx, &task, model.TaskRegistering, "AWX succeeded and MySQL health check passed", ""); err != nil {
return err
}
instance := model.MySQLInstance{TaskID: task.ID, BusinessLineID: task.BusinessLineID, TargetID: task.TargetID, Namespace: payload.Namespace, Name: payload.InstanceName, NodeName: task.TargetHost, Host: task.TargetHostIP, Port: task.MySQLPort, Version: payload.MySQLVersion, Status: "active"}
if err := tx.Create(&instance).Error; err != nil {
return err
}
if err := tx.Create(&model.ResourceUsage{TaskID: task.ID, InstanceID: instance.ID, BusinessLineID: task.BusinessLineID, TargetID: task.TargetID, CPUMilli: payload.CPUMilli, MemoryMi: payload.MemoryMi, StorageGi: payload.StorageGi, InstanceCount: 1, Status: "active"}).Error; err != nil {
return err
}
if err := tx.Model(&model.ResourceReservation{}).Where("task_id = ? AND status = ?", task.ID, "reserved").Update("status", "consumed").Error; err != nil {
return err
}
if err := tx.Model(&model.ExecutionJob{}).Where("task_id = ?", task.ID).Updates(map[string]any{"status": "successful", "finished_at": now}).Error; err != nil {
return err
}
return nil
}); err != nil {
return err
}
if err := s.RegisterCloudDM(ctx, task.ID); err != nil {
return s.transition(ctx, &task, model.TaskRegisterFailed, "CloudDM registration failed", err.Error())
}
return s.transition(ctx, &task, model.TaskFinished, "MySQL delivery completed", "")
}
func (s *DeliveryService) RegisterCloudDM(ctx context.Context, taskID string) error {
if s.cfg.CloudDMRegisterURL == "" {
return nil
}
var instance model.MySQLInstance
if err := s.db.WithContext(ctx).Where("task_id = ?", taskID).First(&instance).Error; err != nil {
return err
}
body := map[string]any{"name": instance.Name, "host": instance.Host, "port": instance.Port, "username": "root", "database_type": "mysql"}
raw, _ := json.Marshal(body)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, s.cfg.CloudDMRegisterURL, bytes.NewReader(raw))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
if s.cfg.CloudDMAPIToken != "" {
req.Header.Set("Authorization", "Bearer "+s.cfg.CloudDMAPIToken)
}
resp, err := http.DefaultClient.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return fmt.Errorf("CloudDM returned %s", resp.Status)
}
return nil
}
func (s *DeliveryService) failTask(ctx context.Context, task *model.DeliveryTask, status, message string) error {
return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
var current model.DeliveryTask
if err := tx.First(&current, "id = ?", task.ID).Error; err != nil {
return err
}
if current.Status != model.TaskPending && current.Status != model.TaskDispatching && current.Status != model.TaskRunning && current.Status != model.TaskRegistering && current.Status != model.TaskCanceling {
return nil
}
if err := s.transitionTx(tx, &current, status, message, message); err != nil {
return err
}
if status == model.TaskExecutionFailed || status == model.TaskValidationFailed || status == model.TaskCanceled {
return s.releaseReservation(tx, current.ID)
}
return nil
})
}
func (s *DeliveryService) Run(ctx context.Context) {
interval := time.Duration(s.cfg.DeliveryDispatchSeconds) * time.Second
if interval < time.Second {
interval = time.Second
}
ticker := time.NewTicker(interval)
defer ticker.Stop()
if err := s.DispatchOnce(ctx); err != nil {
log.Printf("[delivery] dispatch pending task failed: %v", err)
}
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
if err := s.DispatchOnce(ctx); err != nil {
log.Printf("[delivery] dispatch pending task failed: %v", err)
}
}
}
}
func pendingExecutorJobID(taskID string) string {
return "pending:" + taskID
}