kra-new/internal/service/task.go

233 lines
8.2 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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 = &copy
}
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": "清理过期大文件上传会话"}}
}