kra-oa/internal/worker/task_executor.go

171 lines
4.7 KiB
Go

package worker
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"net/http"
"net/url"
"strings"
"syscall"
"time"
"kra/internal/biz"
"kra/internal/conf"
)
type TaskExecutor struct {
tasks *biz.TaskUsecase
media *biz.MediaUsecase
runtime *conf.Runtime
}
func NewTaskExecutor(tasks *biz.TaskUsecase, media *biz.MediaUsecase, runtime *conf.Runtime) *TaskExecutor {
return &TaskExecutor{tasks: tasks, media: media, runtime: runtime}
}
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 "", err
}
if parsed.Scheme != "http" && parsed.Scheme != "https" {
return "", fmt.Errorf("仅允许 http/https, 实际为 %q", parsed.Scheme)
}
method := strings.ToUpper(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{}
_ = json.Unmarshal(task.HTTPHeader, &headers)
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(task *biz.TimedTask) error {
ctx, cancel := context.WithTimeout(context.Background(), 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)
}
}()
var err error
switch task.MethodName {
case biz.TaskMethodClearDB:
err = e.tasks.CleanupLogs(ctx)
case biz.TaskMethodUploads:
ttl := 24
config := e.runtime.Admin()
if config != nil && config.Media != nil && config.Media.SessionTtl > 0 {
ttl = int(config.Media.SessionTtl)
}
err = e.media.CleanupStale(ctx, ttl)
}
done <- err
}()
select {
case err := <-done:
if errors.Is(err, context.DeadlineExceeded) {
return errTaskTimeout
}
return err
case <-ctx.Done():
return errTaskTimeout
}
}
func (e *TaskExecutor) Run(ctx context.Context, task *biz.TimedTask, trigger string) (log *biz.TimedTaskLog) {
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(task)
case biz.TaskExecutorHTTP:
runCtx, cancel := context.WithTimeout(context.Background(), 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
}