kra-oa/internal/service/task.go

195 lines
6.5 KiB
Go

package service
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"github.com/robfig/cron/v3"
"io"
"kra/internal/biz"
"net"
"net/http"
"net/url"
"strings"
"time"
)
type TaskService struct{ uc *biz.TaskUsecase }
func NewTaskService(uc *biz.TaskUsecase) *TaskService { return &TaskService{uc: uc} }
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" {
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
}
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
}
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":
if v.MethodName == "ClearDB" {
days := 30
var params struct {
Days int `json:"days"`
}
_ = json.Unmarshal(v.Params, &params)
if params.Days > 0 {
days = params.Days
}
err = s.uc.Repo().CleanupLogs(runCtx, time.Now().AddDate(0, 0, -days))
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": "清理数据库过期日志(操作记录、登录日志、定时任务日志)"}}
}