9501751792
三轮审查修复(60+ 项),相对远端 main(b536672)的关键变更:
- 安全: 数据面 SSRF 拨号防护(防 DNS rebinding)/上游凭据剥离/登录防枚举
与锁定态统一/可信代理(X-Forwarded-For)限流加固/会话版本失效机制/
撤销即时传播/弱密钥拒绝启动/脱敏字节级重写(保签名契约)
- 业务逻辑: 裸 body 上传 panic/bootstrap 审计管线卡死/定价通配符优先级/
全局工具可见性/调度器停机补跑/TOTP 挑战令牌消费顺序/熔断探针语义/
>4MB 响应 token 计量/管理员重置密码作废会话 等
- 前端: 新 logo(语枢 AI 网关主题)/Provider 凭据异常警示/删除入口/
后端错误消息透传/localStorage 敏感数据收敛
- 部署: CREDENTIAL_MASTER_KEY 持久化与弱值拒绝/Provider DELETE 接口/
nginx 安全头/worker 内存限制
- 新增迁移 000029(key_hash 索引)/000030(usage_daily 币种维度)
442 lines
17 KiB
Go
442 lines
17 KiB
Go
package scheduler
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"regexp"
|
|
"strings"
|
|
"time"
|
|
|
|
"aigateway.local/core/internal/apikey"
|
|
"aigateway.local/core/internal/platform/cryptox"
|
|
platformid "aigateway.local/core/internal/platform/id"
|
|
"github.com/jackc/pgx/v5"
|
|
"github.com/jackc/pgx/v5/pgxpool"
|
|
)
|
|
|
|
var (
|
|
ErrNotFound = errors.New("scheduled task not found")
|
|
codePattern = regexp.MustCompile(`^[a-z][a-z0-9_-]{1,63}$`)
|
|
uuidPattern = regexp.MustCompile(`^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[1-5][0-9a-fA-F]{3}-[89abAB][0-9a-fA-F]{3}-[0-9a-fA-F]{12}$`)
|
|
)
|
|
|
|
type Task struct {
|
|
ID string `json:"id"`
|
|
Code string `json:"code"`
|
|
Name string `json:"name"`
|
|
Description string `json:"description"`
|
|
CronExpression string `json:"cron_expression"`
|
|
Timezone string `json:"timezone"`
|
|
TargetType string `json:"target_type"`
|
|
TargetCode string `json:"target_code"`
|
|
Prompt string `json:"prompt"`
|
|
Variables json.RawMessage `json:"variables"`
|
|
SkillIDs []string `json:"skill_ids"`
|
|
MCPServerIDs []string `json:"mcp_server_ids"`
|
|
ConversationID string `json:"conversation_id"`
|
|
NotificationChannelID *string `json:"notification_channel_id,omitempty"`
|
|
HasAPIKey bool `json:"has_api_key"`
|
|
Enabled bool `json:"enabled"`
|
|
NextRunAt *time.Time `json:"next_run_at,omitempty"`
|
|
LastRunAt *time.Time `json:"last_run_at,omitempty"`
|
|
LastStatus string `json:"last_status"`
|
|
LastError string `json:"last_error"`
|
|
CreatedBy string `json:"created_by"`
|
|
Revision int64 `json:"revision"`
|
|
CreatedAt time.Time `json:"created_at"`
|
|
UpdatedAt time.Time `json:"updated_at"`
|
|
encryptedAPIKey []byte
|
|
apiKeyKEKVersion int
|
|
}
|
|
|
|
type TaskInput struct {
|
|
Code string `json:"code"`
|
|
Name string `json:"name"`
|
|
Description string `json:"description"`
|
|
CronExpression string `json:"cron_expression"`
|
|
Timezone string `json:"timezone"`
|
|
TargetType string `json:"target_type"`
|
|
TargetCode string `json:"target_code"`
|
|
Prompt string `json:"prompt"`
|
|
Variables json.RawMessage `json:"variables"`
|
|
SkillIDs []string `json:"skill_ids"`
|
|
MCPServerIDs []string `json:"mcp_server_ids"`
|
|
ConversationID string `json:"conversation_id"`
|
|
NotificationChannelID *string `json:"notification_channel_id"`
|
|
APIKey string `json:"api_key"`
|
|
Enabled bool `json:"enabled"`
|
|
}
|
|
|
|
type Run struct {
|
|
ID string `json:"id"`
|
|
TaskID string `json:"task_id"`
|
|
TaskCode string `json:"task_code"`
|
|
TriggerType string `json:"trigger_type"`
|
|
ScheduledFor time.Time `json:"scheduled_for"`
|
|
Status string `json:"status"`
|
|
Attempts int `json:"attempts"`
|
|
WorkerID string `json:"worker_id"`
|
|
StartedAt *time.Time `json:"started_at,omitempty"`
|
|
FinishedAt *time.Time `json:"finished_at,omitempty"`
|
|
Response json.RawMessage `json:"response,omitempty"`
|
|
Error string `json:"error"`
|
|
CreatedAt time.Time `json:"created_at"`
|
|
}
|
|
|
|
type Service struct {
|
|
pool *pgxpool.Pool
|
|
cipher cryptox.Cipher
|
|
now func() time.Time
|
|
}
|
|
|
|
func NewService(pool *pgxpool.Pool, cipher cryptox.Cipher) *Service {
|
|
return &Service{pool: pool, cipher: cipher, now: time.Now}
|
|
}
|
|
|
|
const taskSelect = `SELECT id::text,code,name,description,cron_expression,timezone,target_type,target_code,prompt,variables,skill_ids::text[],mcp_server_ids::text[],conversation_id,notification_channel_id::text,encrypted_api_key,api_key_kek_version,enabled,next_run_at,last_run_at,last_status,last_error,created_by::text,revision,created_at,updated_at FROM gateway.scheduled_tasks`
|
|
|
|
func scanTask(row pgx.Row) (Task, error) {
|
|
var task Task
|
|
err := row.Scan(&task.ID, &task.Code, &task.Name, &task.Description, &task.CronExpression, &task.Timezone, &task.TargetType, &task.TargetCode, &task.Prompt, &task.Variables, &task.SkillIDs, &task.MCPServerIDs, &task.ConversationID, &task.NotificationChannelID, &task.encryptedAPIKey, &task.apiKeyKEKVersion, &task.Enabled, &task.NextRunAt, &task.LastRunAt, &task.LastStatus, &task.LastError, &task.CreatedBy, &task.Revision, &task.CreatedAt, &task.UpdatedAt)
|
|
if errors.Is(err, pgx.ErrNoRows) {
|
|
return Task{}, ErrNotFound
|
|
}
|
|
task.HasAPIKey = len(task.encryptedAPIKey) > 0
|
|
if task.SkillIDs == nil {
|
|
task.SkillIDs = []string{}
|
|
}
|
|
if task.MCPServerIDs == nil {
|
|
task.MCPServerIDs = []string{}
|
|
}
|
|
return task, err
|
|
}
|
|
|
|
func (s *Service) List(ctx context.Context) ([]Task, error) {
|
|
rows, err := s.pool.Query(ctx, taskSelect+` ORDER BY updated_at DESC`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
items := []Task{}
|
|
for rows.Next() {
|
|
item, scanErr := scanTask(rows)
|
|
if scanErr != nil {
|
|
return nil, scanErr
|
|
}
|
|
items = append(items, item)
|
|
}
|
|
return items, rows.Err()
|
|
}
|
|
|
|
func (s *Service) Get(ctx context.Context, id string) (Task, error) {
|
|
return scanTask(s.pool.QueryRow(ctx, taskSelect+` WHERE id=$1`, id))
|
|
}
|
|
|
|
func normalizeIDs(values []string, maximum int) ([]string, error) {
|
|
seen := map[string]bool{}
|
|
result := []string{}
|
|
for _, value := range values {
|
|
value = strings.TrimSpace(value)
|
|
if value == "" || seen[value] {
|
|
continue
|
|
}
|
|
if !uuidPattern.MatchString(value) {
|
|
return nil, errors.New("资源 ID 格式无效")
|
|
}
|
|
seen[value] = true
|
|
result = append(result, value)
|
|
}
|
|
if len(result) > maximum {
|
|
return nil, fmt.Errorf("资源绑定最多允许 %d 项", maximum)
|
|
}
|
|
return result, nil
|
|
}
|
|
|
|
func (s *Service) validate(ctx context.Context, input *TaskInput, current *Task) (time.Time, []byte, int, error) {
|
|
input.Code = strings.ToLower(strings.TrimSpace(input.Code))
|
|
input.Name = strings.TrimSpace(input.Name)
|
|
input.Description = strings.TrimSpace(input.Description)
|
|
input.CronExpression = strings.TrimSpace(input.CronExpression)
|
|
input.Timezone = strings.TrimSpace(input.Timezone)
|
|
input.TargetType = strings.TrimSpace(input.TargetType)
|
|
input.TargetCode = strings.ToLower(strings.TrimSpace(input.TargetCode))
|
|
input.Prompt = strings.TrimSpace(input.Prompt)
|
|
input.ConversationID = strings.TrimSpace(input.ConversationID)
|
|
if !codePattern.MatchString(input.Code) || input.Name == "" || len(input.Name) > 128 || len(input.Description) > 4000 {
|
|
return time.Time{}, nil, 0, errors.New("任务编码、名称或描述格式无效")
|
|
}
|
|
if len(input.Prompt) < 1 || len(input.Prompt) > 100000 || len(input.ConversationID) > 128 {
|
|
return time.Time{}, nil, 0, errors.New("提示词或会话 ID 格式无效")
|
|
}
|
|
if input.Timezone == "" {
|
|
input.Timezone = "UTC"
|
|
}
|
|
location, err := time.LoadLocation(input.Timezone)
|
|
if err != nil {
|
|
return time.Time{}, nil, 0, errors.New("时区名称无效")
|
|
}
|
|
schedule, err := ParseCron(input.CronExpression)
|
|
if err != nil {
|
|
return time.Time{}, nil, 0, err
|
|
}
|
|
next, err := schedule.Next(s.now(), location)
|
|
if err != nil {
|
|
return time.Time{}, nil, 0, err
|
|
}
|
|
if len(input.Variables) == 0 {
|
|
input.Variables = json.RawMessage(`{}`)
|
|
}
|
|
var variables map[string]any
|
|
if json.Unmarshal(input.Variables, &variables) != nil {
|
|
return time.Time{}, nil, 0, errors.New("变量必须是 JSON 对象")
|
|
}
|
|
input.Variables, _ = json.Marshal(variables)
|
|
if input.SkillIDs, err = normalizeIDs(input.SkillIDs, 100); err != nil {
|
|
return time.Time{}, nil, 0, err
|
|
}
|
|
if input.MCPServerIDs, err = normalizeIDs(input.MCPServerIDs, 100); err != nil {
|
|
return time.Time{}, nil, 0, err
|
|
}
|
|
if err = s.validateTarget(ctx, input); err != nil {
|
|
return time.Time{}, nil, 0, err
|
|
}
|
|
if input.NotificationChannelID != nil {
|
|
trimmed := strings.TrimSpace(*input.NotificationChannelID)
|
|
if trimmed == "" {
|
|
input.NotificationChannelID = nil
|
|
} else {
|
|
var exists bool
|
|
if err = s.pool.QueryRow(ctx, `SELECT EXISTS(SELECT 1 FROM gateway.notification_channels WHERE id=$1 AND enabled)`, trimmed).Scan(&exists); err != nil || !exists {
|
|
return time.Time{}, nil, 0, errors.New("通知渠道不存在或未启用")
|
|
}
|
|
input.NotificationChannelID = &trimmed
|
|
}
|
|
}
|
|
secret := strings.TrimSpace(input.APIKey)
|
|
if secret == "" && current == nil {
|
|
return time.Time{}, nil, 0, errors.New("首次创建必须填写执行 API Key")
|
|
}
|
|
if secret == "" {
|
|
return next, current.encryptedAPIKey, current.apiKeyKEKVersion, nil
|
|
}
|
|
if len(secret) > 512 {
|
|
return time.Time{}, nil, 0, errors.New("执行 API Key 过长")
|
|
}
|
|
encrypted, version, err := s.cipher.Encrypt([]byte(secret))
|
|
if err != nil {
|
|
return time.Time{}, nil, 0, fmt.Errorf("加密执行 API Key: %w", err)
|
|
}
|
|
return next, encrypted, version, nil
|
|
}
|
|
|
|
func subset(selected, allowed []string) bool {
|
|
set := map[string]bool{}
|
|
for _, id := range allowed {
|
|
set[id] = true
|
|
}
|
|
for _, id := range selected {
|
|
if !set[id] {
|
|
return false
|
|
}
|
|
}
|
|
return true
|
|
}
|
|
|
|
func (s *Service) validateTarget(ctx context.Context, input *TaskInput) error {
|
|
switch input.TargetType {
|
|
case "application":
|
|
if len(input.SkillIDs) > 0 || len(input.MCPServerIDs) > 0 {
|
|
return errors.New("应用任务不支持额外绑定 Skill 或 MCP")
|
|
}
|
|
var exists bool
|
|
err := s.pool.QueryRow(ctx, `SELECT EXISTS(SELECT 1 FROM gateway.applications WHERE code=$1 AND status='active' AND published_version IS NOT NULL)`, input.TargetCode).Scan(&exists)
|
|
if err != nil || !exists {
|
|
return errors.New("目标应用不存在或未发布")
|
|
}
|
|
var departmentIDs []string
|
|
_ = s.pool.QueryRow(ctx, `SELECT COALESCE(department_ids,'{}'::text[]) FROM gateway.applications WHERE code=$1`, input.TargetCode).Scan(&departmentIDs)
|
|
if len(departmentIDs) > 0 {
|
|
if err := s.requireKeyTenant(ctx, input.APIKey, departmentIDs); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
case "digital_employee":
|
|
var skillIDs, mcpIDs, departmentIDs []string
|
|
var enabled bool
|
|
var status string
|
|
err := s.pool.QueryRow(ctx, `SELECT skill_ids::text[],mcp_server_ids::text[],enabled,status FROM gateway.digital_employees WHERE code=$1`, input.TargetCode).Scan(&skillIDs, &mcpIDs, &enabled, &status)
|
|
if err != nil || !enabled || status != "published" {
|
|
return errors.New("目标数字员工不存在或未发布")
|
|
}
|
|
if !subset(input.SkillIDs, skillIDs) || !subset(input.MCPServerIDs, mcpIDs) {
|
|
return errors.New("任务选择的 Skill/MCP 必须已绑定到目标数字员工")
|
|
}
|
|
_ = s.pool.QueryRow(ctx, `SELECT COALESCE(department_ids,'{}'::text[]) FROM gateway.digital_employees WHERE code=$1`, input.TargetCode).Scan(&departmentIDs)
|
|
if len(departmentIDs) > 0 {
|
|
if err := s.requireKeyTenant(ctx, input.APIKey, departmentIDs); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
default:
|
|
return errors.New("目标类型必须是 application 或 digital_employee")
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// requireKeyTenant 校验任务 API Key 的部门归属能访问部门限定目标,在保存
|
|
// 阶段就失败,而不是让任务创建成功后永远执行失败(执行 key 无 tenant 时
|
|
// 运行时对部门限定资源一律不可见)。secret 为空(更新时沿用旧 key)跳过。
|
|
func (s *Service) requireKeyTenant(ctx context.Context, secret string, departmentIDs []string) error {
|
|
secret = strings.TrimSpace(secret)
|
|
if secret == "" {
|
|
return nil
|
|
}
|
|
hash, _ := apikey.Digest(secret)
|
|
var tenant *string
|
|
err := s.pool.QueryRow(ctx, `SELECT tenant_id::text FROM gateway.api_keys WHERE key_hash=$1 AND enabled`, hash).Scan(&tenant)
|
|
if errors.Is(err, pgx.ErrNoRows) {
|
|
return errors.New("执行 API Key 不存在或已停用")
|
|
}
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if tenant == nil {
|
|
return errors.New("目标资源按部门限定,但执行 API Key 未绑定部门;请使用该部门下的 API Key")
|
|
}
|
|
for _, id := range departmentIDs {
|
|
if id == *tenant {
|
|
return nil
|
|
}
|
|
}
|
|
return errors.New("执行 API Key 所属部门与目标资源部门不匹配")
|
|
}
|
|
|
|
func (s *Service) Save(ctx context.Context, id string, input TaskInput, actorID string) (Task, error) {
|
|
var current *Task
|
|
if id != "" {
|
|
item, err := s.Get(ctx, id)
|
|
if err != nil {
|
|
return Task{}, err
|
|
}
|
|
current = &item
|
|
}
|
|
next, encrypted, version, err := s.validate(ctx, &input, current)
|
|
if err != nil {
|
|
return Task{}, err
|
|
}
|
|
if id == "" {
|
|
id, err = platformid.NewUUID()
|
|
if err != nil {
|
|
return Task{}, err
|
|
}
|
|
_, err = s.pool.Exec(ctx, `INSERT INTO gateway.scheduled_tasks(id,code,name,description,cron_expression,timezone,target_type,target_code,prompt,variables,skill_ids,mcp_server_ids,conversation_id,notification_channel_id,encrypted_api_key,api_key_kek_version,enabled,next_run_at,created_by) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19)`, id, input.Code, input.Name, input.Description, input.CronExpression, input.Timezone, input.TargetType, input.TargetCode, input.Prompt, input.Variables, input.SkillIDs, input.MCPServerIDs, input.ConversationID, input.NotificationChannelID, encrypted, version, input.Enabled, nullableNext(input.Enabled, next), actorID)
|
|
} else {
|
|
_, err = s.pool.Exec(ctx, `UPDATE gateway.scheduled_tasks SET code=$2,name=$3,description=$4,cron_expression=$5,timezone=$6,target_type=$7,target_code=$8,prompt=$9,variables=$10,skill_ids=$11,mcp_server_ids=$12,conversation_id=$13,notification_channel_id=$14,encrypted_api_key=$15,api_key_kek_version=$16,enabled=$17,next_run_at=$18,revision=revision+1,updated_at=clock_timestamp() WHERE id=$1`, id, input.Code, input.Name, input.Description, input.CronExpression, input.Timezone, input.TargetType, input.TargetCode, input.Prompt, input.Variables, input.SkillIDs, input.MCPServerIDs, input.ConversationID, input.NotificationChannelID, encrypted, version, input.Enabled, nullableNext(input.Enabled, next))
|
|
}
|
|
if err != nil {
|
|
return Task{}, err
|
|
}
|
|
return s.Get(ctx, id)
|
|
}
|
|
|
|
func nullableNext(enabled bool, next time.Time) any {
|
|
if !enabled {
|
|
return nil
|
|
}
|
|
return next
|
|
}
|
|
|
|
func (s *Service) SetEnabled(ctx context.Context, id string, enabled bool) (Task, error) {
|
|
task, err := s.Get(ctx, id)
|
|
if err != nil {
|
|
return Task{}, err
|
|
}
|
|
var next any
|
|
if enabled {
|
|
location, _ := time.LoadLocation(task.Timezone)
|
|
schedule, _ := ParseCron(task.CronExpression)
|
|
value, nextErr := schedule.Next(s.now(), location)
|
|
if nextErr != nil {
|
|
return Task{}, nextErr
|
|
}
|
|
next = value
|
|
}
|
|
tag, err := s.pool.Exec(ctx, `UPDATE gateway.scheduled_tasks SET enabled=$2,next_run_at=$3,revision=revision+1,updated_at=clock_timestamp() WHERE id=$1`, id, enabled, next)
|
|
if err != nil || tag.RowsAffected() == 0 {
|
|
return Task{}, ErrNotFound
|
|
}
|
|
return s.Get(ctx, id)
|
|
}
|
|
|
|
func (s *Service) Delete(ctx context.Context, id string) error {
|
|
tag, err := s.pool.Exec(ctx, `DELETE FROM gateway.scheduled_tasks WHERE id=$1`, id)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if tag.RowsAffected() == 0 {
|
|
return ErrNotFound
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (s *Service) QueueManual(ctx context.Context, id string) (Run, error) {
|
|
if _, err := s.Get(ctx, id); err != nil {
|
|
return Run{}, err
|
|
}
|
|
runID, err := platformid.NewUUID()
|
|
if err != nil {
|
|
return Run{}, err
|
|
}
|
|
now := s.now().UTC()
|
|
_, err = s.pool.Exec(ctx, `INSERT INTO gateway.scheduled_task_runs(id,task_id,trigger_type,scheduled_for) VALUES($1,$2,'manual',$3)`, runID, id, now)
|
|
if err != nil {
|
|
return Run{}, err
|
|
}
|
|
return s.getRun(ctx, runID)
|
|
}
|
|
|
|
const runSelect = `SELECT r.id::text,r.task_id::text,t.code,r.trigger_type,r.scheduled_for,r.status,r.attempts,r.worker_id,r.started_at,r.finished_at,r.response,r.error,r.created_at FROM gateway.scheduled_task_runs r JOIN gateway.scheduled_tasks t ON t.id=r.task_id`
|
|
|
|
func scanRun(row pgx.Row) (Run, error) {
|
|
var run Run
|
|
err := row.Scan(&run.ID, &run.TaskID, &run.TaskCode, &run.TriggerType, &run.ScheduledFor, &run.Status, &run.Attempts, &run.WorkerID, &run.StartedAt, &run.FinishedAt, &run.Response, &run.Error, &run.CreatedAt)
|
|
if errors.Is(err, pgx.ErrNoRows) {
|
|
return Run{}, ErrNotFound
|
|
}
|
|
return run, err
|
|
}
|
|
|
|
func (s *Service) getRun(ctx context.Context, id string) (Run, error) {
|
|
return scanRun(s.pool.QueryRow(ctx, runSelect+` WHERE r.id=$1`, id))
|
|
}
|
|
|
|
func (s *Service) Runs(ctx context.Context, taskID string, limit int) ([]Run, error) {
|
|
if limit < 1 || limit > 500 {
|
|
limit = 100
|
|
}
|
|
rows, err := s.pool.Query(ctx, runSelect+` WHERE r.task_id=$1 ORDER BY r.created_at DESC LIMIT $2`, taskID, limit)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
items := []Run{}
|
|
for rows.Next() {
|
|
item, scanErr := scanRun(rows)
|
|
if scanErr != nil {
|
|
return nil, scanErr
|
|
}
|
|
items = append(items, item)
|
|
}
|
|
return items, rows.Err()
|
|
}
|
|
|
|
func (s *Service) decryptAPIKey(task Task) (string, error) {
|
|
plain, err := s.cipher.Decrypt(task.encryptedAPIKey, task.apiKeyKEKVersion)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
return string(plain), nil
|
|
}
|