package worker import ( "bytes" "context" "encoding/json" "errors" "fmt" "io" "net" "net/http" "net/url" "strings" "syscall" "time" "kra/internal/conf" "kra/internal/modules/system/biz" ) type TaskExecutor struct { tasks *biz.TaskUsecase media *biz.MediaUsecase runtime *conf.Runtime } func NewTaskExecutor(tasks *biz.TaskUsecase, media *biz.MediaUsecase, runtime *conf.Runtime) *TaskExecutor { executor := &TaskExecutor{tasks: tasks, media: media, runtime: runtime} registerTaskMethods(executor) return executor } func privateIP(ip net.IP) bool { return ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalMulticast() || ip.IsLinkLocalUnicast() || ip.IsUnspecified() } func taskHTTPClient(allowPrivate bool) *http.Client { dialer := &net.Dialer{Timeout: 10 * time.Second, Control: func(_ string, address string, _ syscall.RawConn) error { if allowPrivate { return nil } host, _, err := net.SplitHostPort(address) if err != nil { return fmt.Errorf("解析拨号地址失败: %w", err) } ip := net.ParseIP(host) if ip == nil { return fmt.Errorf("非法拨号 IP: %s", host) } if privateIP(ip) { return fmt.Errorf("目标解析为内网/环回/链路本地地址, 已被 SSRF 防护拒绝(可在任务上开启\"允许内网\"豁免): %s", ip) } return nil }} return &http.Client{Timeout: 30 * time.Second, Transport: &http.Transport{Proxy: nil, DialContext: dialer.DialContext}} } func (e *TaskExecutor) runHTTP(ctx context.Context, task *biz.TimedTask) (string, error) { parsed, err := url.Parse(task.HTTPURL) if err != nil { return "", fmt.Errorf("URL 非法: %w", err) } if parsed.Scheme != "http" && parsed.Scheme != "https" { return "", fmt.Errorf("仅允许 http/https, 实际为 %q", parsed.Scheme) } method := strings.ToUpper(strings.TrimSpace(task.HTTPMethod)) if method == "" { method = http.MethodGet } request, err := http.NewRequestWithContext(ctx, method, task.HTTPURL, bytes.NewBufferString(task.HTTPBody)) if err != nil { return "", err } headers := map[string]string{} if len(task.HTTPHeader) > 0 { if err = json.Unmarshal(task.HTTPHeader, &headers); err != nil { return "", fmt.Errorf("http_header 必须是 JSON 对象: %w", err) } } for key, value := range headers { request.Header.Set(key, value) } response, err := taskHTTPClient(task.HTTPAllowPrivate).Do(request) if err != nil { return "", err } defer response.Body.Close() body, _ := io.ReadAll(io.LimitReader(response.Body, 1<<20)) output := fmt.Sprintf("HTTP %d: %s", response.StatusCode, string(body)) if response.StatusCode < 200 || response.StatusCode >= 300 { return output, fmt.Errorf("非 2xx 响应: %d", response.StatusCode) } return output, nil } var errTaskTimeout = errors.New("任务执行超时") func truncateTaskText(value string) string { const limit = 4000 if len(value) <= limit { return value } return value[:limit] + "...(截断)" } func (e *TaskExecutor) runMethod(ctx context.Context, task *biz.TimedTask) error { method, ok := biz.TaskMethodByName(task.MethodName) if !ok { return fmt.Errorf("方法 %s 未注册(需在 internal/worker/task_registry.go 经 biz.RegisterTaskMethod 注册)", task.MethodName) } runCtx, cancel := context.WithTimeout(ctx, 5*time.Minute) defer cancel() done := make(chan error, 1) go func() { defer func() { if recovered := recover(); recovered != nil { done <- fmt.Errorf("panic: %v", recovered) } }() done <- method(runCtx, json.RawMessage(task.Params)) }() select { case err := <-done: if errors.Is(err, context.DeadlineExceeded) { return errTaskTimeout } return err case <-runCtx.Done(): if errors.Is(runCtx.Err(), context.DeadlineExceeded) { return errTaskTimeout } return runCtx.Err() } } func (e *TaskExecutor) Run(ctx context.Context, task *biz.TimedTask, trigger string) (log *biz.TimedTaskLog) { if ctx == nil { ctx = context.Background() } started := time.Now() log = &biz.TimedTaskLog{TaskID: task.ID, TaskName: task.Name, TriggerType: trigger, StartedAt: started, Status: "success"} defer func() { if recovered := recover(); recovered != nil { log.Status = "fail" log.ErrorMsg = truncateTaskText(fmt.Sprint(recovered)) } log.Output = truncateTaskText(log.Output) log.FinishedAt = time.Now() log.DurationMS = log.FinishedAt.Sub(started).Milliseconds() logCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 3*time.Second) defer cancel() _ = e.tasks.RecordTaskLog(logCtx, log) }() var err error switch task.ExecutorType { case biz.TaskExecutorMethod: err = e.runMethod(ctx, task) case biz.TaskExecutorHTTP: runCtx, cancel := context.WithTimeout(ctx, 30*time.Second) log.Output, err = e.runHTTP(runCtx, task) cancel() default: err = fmt.Errorf("未知执行器类型: %s", task.ExecutorType) } if err != nil { if errors.Is(err, errTaskTimeout) || errors.Is(err, context.DeadlineExceeded) { log.Status = "timeout" } else { log.Status = "fail" } log.ErrorMsg = truncateTaskText(err.Error()) } return log }