package worker import ( "bytes" "context" "encoding/json" "errors" "fmt" "io" "net" "net/http" "net/url" "strings" "syscall" "time" taskbiz "kra/internal/biz/task" ) type TaskExecutor struct { tasks *taskbiz.TaskUsecase methods taskbiz.TaskMethodRegistry } func NewTaskExecutorWithRegistry(tasks *taskbiz.TaskUsecase, methods taskbiz.TaskMethodRegistry) *TaskExecutor { return &TaskExecutor{tasks: tasks, methods: methods} } 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 *taskbiz.TimedTask) (string, error) { if task == nil { return "", errors.New("任务不能为空") } if ctx == nil { ctx = context.Background() } 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) } if parsed.Host == "" { return "", errors.New("URL 缺少主机名") } 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) } if headers == nil { return "", errors.New("http_header 必须是 JSON 对象") } } 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, readErr := io.ReadAll(io.LimitReader(response.Body, 1<<20)) output := fmt.Sprintf("HTTP %d: %s", response.StatusCode, string(body)) if readErr != nil { return output, fmt.Errorf("读取 HTTP 响应失败: %w", readErr) } 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) recordTaskLog(ctx context.Context, log *taskbiz.TimedTaskLog) { defer func() { _ = recover() }() if e == nil || e.tasks == nil || log == nil { return } _ = e.tasks.RecordTaskLog(ctx, log) } func (e *TaskExecutor) runMethod(ctx context.Context, task *taskbiz.TimedTask) error { if task == nil { return errors.New("任务不能为空") } if e == nil { return errors.New("任务执行器未初始化") } if ctx == nil { ctx = context.Background() } var method taskbiz.TaskMethodFunc var ok bool if e.methods != nil { method, ok = e.methods.Lookup(task.MethodName) } if !ok { return fmt.Errorf("方法 %s 未注册(需通过 platform/task.Registry 注册)", 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 *taskbiz.TimedTask, trigger string) (log *taskbiz.TimedTaskLog) { if ctx == nil { ctx = context.Background() } started := time.Now() log = &taskbiz.TimedTaskLog{TriggerType: trigger, StartedAt: started, Status: "success"} if task != nil { log.TaskID = task.ID log.TaskName = task.Name } 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() if task != nil { e.recordTaskLog(logCtx, log) } }() if task == nil { log.Status = "fail" log.ErrorMsg = "任务不能为空" return log } if e == nil { log.Status = "fail" log.ErrorMsg = "任务执行器未初始化" return log } var err error switch task.ExecutorType { case taskbiz.TaskExecutorMethod: err = e.runMethod(ctx, task) case taskbiz.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 }