195 lines
6.5 KiB
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 = ©
|
|
}
|
|
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, ¶ms)
|
|
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": "清理数据库过期日志(操作记录、登录日志、定时任务日志)"}}
|
|
}
|