package service import ( "bytes" "context" "encoding/json" "errors" "fmt" "github.com/robfig/cron/v3" "io" "kra/internal/biz" "kra/internal/conf" "kra/internal/service/dto" "net" "net/http" "net/url" "strings" "time" ) type TaskService struct { uc *biz.TaskUsecase media *biz.MediaUsecase config *conf.AdminBackend } func NewTaskService(uc *biz.TaskUsecase, media *biz.MediaUsecase, config *conf.AdminBackend) *TaskService { return &TaskService{uc: uc, media: media, config: config} } func taskDomain(v *dto.TaskRequest) *biz.TimedTask { return &biz.TimedTask{ID: v.ID, Name: v.Name, Description: v.Description, Spec: v.Spec, WithSeconds: v.WithSeconds, ExecutorType: v.ExecutorType, MethodName: v.MethodName, Params: v.Params, HTTPURL: v.HTTPURL, HTTPMethod: v.HTTPMethod, HTTPHeader: v.HTTPHeader, HTTPBody: v.HTTPBody, HTTPAllowPrivate: v.HTTPAllowPrivate, Enabled: v.Enabled} } func (s *TaskService) CreateRequest(ctx context.Context, req *dto.TaskRequest) (uint, error) { value := taskDomain(req) if err := s.Create(ctx, value); err != nil { return 0, err } return value.ID, nil } func (s *TaskService) UpdateRequest(ctx context.Context, req *dto.TaskRequest) error { return s.Update(ctx, taskDomain(req)) } func (s *TaskService) ListRequest(ctx context.Context, page, size int, name, executorType string, enabled *bool, next map[uint]time.Time) ([]map[string]any, int64, error) { return s.Tasks(ctx, page, size, &biz.TimedTask{Name: name, ExecutorType: executorType, EnabledFilter: enabled}, next) } func (s *TaskService) Validate(v *biz.TimedTask) error { if v.Name == "" { return errors.New("任务名不能为空") } var err error if v.WithSeconds { _, err = cron.NewParser(cron.Second | cron.Minute | cron.Hour | cron.Dom | cron.Month | cron.Dow | cron.Descriptor).Parse(v.Spec) } else { _, err = cron.ParseStandard(v.Spec) } if err != nil { return fmt.Errorf("cron 表达式非法: %w", err) } switch v.ExecutorType { case "method": if v.MethodName != "ClearDB" && v.MethodName != "CleanStaleUploads" { return errors.New("方法未注册") } if len(v.Params) > 0 && !json.Valid(v.Params) { return errors.New("params 必须是合法 JSON") } case "http": parsed, parseErr := url.Parse(v.HTTPURL) if parseErr != nil || (parsed.Scheme != "http" && parsed.Scheme != "https") || parsed.Hostname() == "" || parsed.User != nil { return errors.New("httpUrl 必须是合法的 http/https 地址") } if len(v.HTTPHeader) > 0 { headers := map[string]string{} if json.Unmarshal(v.HTTPHeader, &headers) != nil { return errors.New("httpHeader 必须是 JSON 对象") } } default: return errors.New("executorType 必须为 method 或 http") } return nil } func (s *TaskService) Create(ctx context.Context, v *biz.TimedTask) error { if err := s.Validate(v); err != nil { return err } exists, err := s.uc.Repo().TaskNameExists(ctx, v.Name, 0) if err != nil { return err } if exists { return fmt.Errorf("任务名 %s 已存在", v.Name) } return s.uc.Repo().CreateTask(ctx, v) } func (s *TaskService) Update(ctx context.Context, v *biz.TimedTask) error { if v.ID == 0 { return errors.New("缺少任务 ID") } if err := s.Validate(v); err != nil { return err } exists, err := s.uc.Repo().TaskNameExists(ctx, v.Name, v.ID) if err != nil { return err } if exists { return fmt.Errorf("任务名 %s 已存在", v.Name) } return s.uc.Repo().UpdateTask(ctx, v) } func (s *TaskService) Delete(ctx context.Context, id uint) error { return s.uc.Repo().DeleteTask(ctx, id) } func (s *TaskService) Toggle(ctx context.Context, id uint, enabled bool) error { return s.uc.Repo().ToggleTask(ctx, id, enabled) } func (s *TaskService) Task(ctx context.Context, id uint) (*biz.TimedTask, error) { return s.uc.Repo().FindTask(ctx, id) } func taskDTO(v *biz.TimedTask, next *time.Time) map[string]any { return map[string]any{"ID": v.ID, "CreatedAt": v.CreatedAt, "UpdatedAt": v.UpdatedAt, "DeletedAt": nil, "name": v.Name, "description": v.Description, "spec": v.Spec, "withSeconds": v.WithSeconds, "executorType": v.ExecutorType, "methodName": v.MethodName, "params": json.RawMessage(v.Params), "httpUrl": v.HTTPURL, "httpMethod": v.HTTPMethod, "httpHeader": json.RawMessage(v.HTTPHeader), "httpBody": v.HTTPBody, "httpAllowPrivate": v.HTTPAllowPrivate, "enabled": v.Enabled, "nextRunAt": next} } func (s *TaskService) Tasks(ctx context.Context, page, size int, q *biz.TimedTask, next map[uint]time.Time) ([]map[string]any, int64, error) { items, total, err := s.uc.Repo().ListTasks(ctx, page, size, q) if err != nil { return nil, 0, err } out := make([]map[string]any, 0, len(items)) for _, v := range items { var ptr *time.Time if value, ok := next[v.ID]; ok { copy := value ptr = © } out = append(out, taskDTO(v, ptr)) } return out, total, nil } func (s *TaskService) Logs(ctx context.Context, page, size int, taskID uint, status string) ([]map[string]any, int64, error) { items, total, err := s.uc.Repo().ListTaskLogs(ctx, page, size, taskID, status) if err != nil { return nil, 0, err } out := make([]map[string]any, 0, len(items)) for _, v := range items { out = append(out, map[string]any{"ID": v.ID, "CreatedAt": v.CreatedAt, "DeletedAt": nil, "taskId": v.TaskID, "taskName": v.TaskName, "triggerType": v.TriggerType, "startedAt": v.StartedAt, "finishedAt": v.FinishedAt, "durationMs": v.DurationMS, "status": v.Status, "errorMsg": v.ErrorMsg, "output": v.Output}) } return out, total, nil } func privateIP(ip net.IP) bool { return ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalMulticast() || ip.IsLinkLocalUnicast() || ip.IsUnspecified() } func (s *TaskService) runHTTP(ctx context.Context, v *biz.TimedTask) (string, error) { parsed, _ := url.Parse(v.HTTPURL) if !v.HTTPAllowPrivate { ips, err := net.DefaultResolver.LookupIP(ctx, "ip", parsed.Hostname()) if err != nil { return "", err } for _, ip := range ips { if privateIP(ip) { return "", errors.New("禁止访问内网或环回地址") } } } method := strings.ToUpper(v.HTTPMethod) if method == "" { method = http.MethodGet } request, err := http.NewRequestWithContext(ctx, method, v.HTTPURL, bytes.NewBufferString(v.HTTPBody)) if err != nil { return "", err } headers := map[string]string{} _ = json.Unmarshal(v.HTTPHeader, &headers) for key, value := range headers { request.Header.Set(key, value) } response, err := (&http.Client{Timeout: 30 * time.Second}).Do(request) if err != nil { return "", err } defer response.Body.Close() body, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) output := fmt.Sprintf("HTTP %d: %s", response.StatusCode, string(body)) if response.StatusCode < 200 || response.StatusCode >= 300 { return output, fmt.Errorf("HTTP 状态码 %d", response.StatusCode) } return output, nil } func (s *TaskService) Run(ctx context.Context, v *biz.TimedTask, trigger string) (log *biz.TimedTaskLog) { started := time.Now() log = &biz.TimedTaskLog{TaskID: v.ID, TaskName: v.Name, TriggerType: trigger, StartedAt: started, Status: "success"} defer func() { if recovered := recover(); recovered != nil { log.Status = "fail" log.ErrorMsg = fmt.Sprint(recovered) } log.FinishedAt = time.Now() log.DurationMS = log.FinishedAt.Sub(started).Milliseconds() _ = s.uc.Repo().RecordTaskLog(context.Background(), log) }() runCtx, cancel := context.WithTimeout(ctx, 35*time.Second) defer cancel() var err error switch v.ExecutorType { case "method": switch v.MethodName { case "ClearDB": err = s.uc.Repo().CleanupLogs(runCtx) log.Output = "过期日志清理完成" case "CleanStaleUploads": ttl := 24 if s.config != nil && s.config.Media != nil && s.config.Media.SessionTtl > 0 { ttl = int(s.config.Media.SessionTtl) } err = s.media.CleanupStale(runCtx, ttl) log.Output = "过期大文件上传会话清理完成" } case "http": log.Output, err = s.runHTTP(runCtx, v) } if err != nil { log.Status = "fail" log.ErrorMsg = err.Error() } return log } func (s *TaskService) RegisteredMethods() []map[string]any { return []map[string]any{{"name": "ClearDB", "description": "清理数据库过期日志(操作记录、JWT黑名单、定时任务日志)"}, {"name": "CleanStaleUploads", "description": "清理过期大文件上传会话"}} }