233 lines
8.2 KiB
Go
233 lines
8.2 KiB
Go
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": "清理过期大文件上传会话"}}
|
||
}
|