d534865b33
- 节点任务:管理端向节点池下发(Prompt/HTTP/MCP/Skill/数字员工/自定义), 指定节点或池路由,认领 SKIP LOCKED + 15 分钟租约,认领令牌防重放上报, 失败 30s×次数退避重入队,达上限 failed,支持取消/重试,完成与失败站内信。 - 个人智能体安全策略:auto_approve_tools 跳过个人调用审批门; rate_limit_multiplier 按 (tool,user) 独立窗口放宽个人限流(全局额度不受影响)。 - 修复存量缺陷:/v1/agent/nodes/ 未挂 publicMux,节点心跳/认领端点在部署 拓扑下不可达。 - 迁移 000046;任务全链路集成测试连真实库通过,HTTP 端到端验证 (下发→认领→伪造令牌拒绝→上报→succeeded,列表不泄露认领令牌); 25 包测试通过,前后端构建通过。
337 lines
13 KiB
Go
337 lines
13 KiB
Go
package agentnode
|
|
|
|
import (
|
|
"context"
|
|
"crypto/sha256"
|
|
"crypto/subtle"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"strings"
|
|
"time"
|
|
|
|
platformid "aigateway.local/core/internal/platform/id"
|
|
"github.com/jackc/pgx/v5"
|
|
)
|
|
|
|
// Task 是一条下发给节点池的任务。
|
|
type Task struct {
|
|
ID string `json:"id"`
|
|
TaskType string `json:"task_type"`
|
|
Payload json.RawMessage `json:"payload"`
|
|
PoolType string `json:"pool_type"`
|
|
PoolCode string `json:"pool_code"`
|
|
NodeID *string `json:"node_id,omitempty"`
|
|
NodeCode string `json:"node_code,omitempty"`
|
|
Status string `json:"status"`
|
|
Attempts int `json:"attempts"`
|
|
MaxAttempts int `json:"max_attempts"`
|
|
Result json.RawMessage `json:"result,omitempty"`
|
|
Error string `json:"error,omitempty"`
|
|
// ClaimToken 是认领时签发的上报凭证,仅认领响应返回,列表/详情不回传。
|
|
ClaimToken string `json:"claim_token,omitempty"`
|
|
CreatedAt time.Time `json:"created_at"`
|
|
ClaimedAt *time.Time `json:"claimed_at,omitempty"`
|
|
FinishedAt *time.Time `json:"finished_at,omitempty"`
|
|
AvailableAt time.Time `json:"available_at"`
|
|
}
|
|
|
|
// TaskInput 是创建任务的输入。
|
|
type TaskInput struct {
|
|
TaskType string
|
|
Payload json.RawMessage
|
|
PoolType string
|
|
PoolCode string
|
|
NodeID *string
|
|
MaxAttempts int
|
|
}
|
|
|
|
var (
|
|
ErrTaskNotFound = errors.New("agent task not found")
|
|
ErrTaskConflict = errors.New("agent task state conflict")
|
|
ErrTaskInvalid = errors.New("agent task input invalid")
|
|
ErrTaskUnauthorized = errors.New("agent node token invalid")
|
|
)
|
|
|
|
var taskTypes = map[string]bool{"prompt": true, "http": true, "mcp_invoke": true, "skill_run": true, "digital_employee": true, "custom": true}
|
|
|
|
const taskSelect = `SELECT t.id::text,t.task_type,t.payload,t.pool_type,t.pool_code,t.node_id::text,coalesce(n.code,''),t.status,t.attempts,t.max_attempts,t.result,t.error,t.created_at,t.claimed_at,t.finished_at,t.available_at
|
|
FROM gateway.agent_tasks t LEFT JOIN gateway.agent_nodes n ON n.id=t.node_id`
|
|
|
|
func scanTask(row pgx.Row) (Task, error) {
|
|
var item Task
|
|
var nodeID *string
|
|
err := row.Scan(&item.ID, &item.TaskType, &item.Payload, &item.PoolType, &item.PoolCode, &nodeID, &item.NodeCode, &item.Status, &item.Attempts, &item.MaxAttempts, &item.Result, &item.Error, &item.CreatedAt, &item.ClaimedAt, &item.FinishedAt, &item.AvailableAt)
|
|
if errors.Is(err, pgx.ErrNoRows) {
|
|
return Task{}, ErrTaskNotFound
|
|
}
|
|
if err != nil {
|
|
return Task{}, err
|
|
}
|
|
item.NodeID = nodeID
|
|
item.Payload = normalizeJSON(item.Payload)
|
|
return item, nil
|
|
}
|
|
|
|
// CreateTask 下发任务到节点池(queued)。
|
|
func (s *Store) CreateTask(ctx context.Context, input TaskInput, actorID string) (Task, error) {
|
|
if s == nil || s.pool == nil {
|
|
return Task{}, ErrStore
|
|
}
|
|
input.TaskType = strings.TrimSpace(input.TaskType)
|
|
input.PoolType = strings.ToLower(strings.TrimSpace(input.PoolType))
|
|
input.PoolCode = strings.TrimSpace(input.PoolCode)
|
|
if !taskTypes[input.TaskType] {
|
|
return Task{}, fmt.Errorf("%w: task type must be one of prompt/http/mcp_invoke/skill_run/digital_employee/custom", ErrTaskInvalid)
|
|
}
|
|
if len(input.Payload) == 0 || !json.Valid(input.Payload) {
|
|
return Task{}, fmt.Errorf("%w: payload must be a JSON object", ErrTaskInvalid)
|
|
}
|
|
var object map[string]any
|
|
if json.Unmarshal(input.Payload, &object) != nil || object == nil {
|
|
return Task{}, fmt.Errorf("%w: payload must be a JSON object", ErrTaskInvalid)
|
|
}
|
|
if input.PoolType == "" {
|
|
input.PoolType = "private"
|
|
}
|
|
if input.PoolType != "public" && input.PoolType != "private" {
|
|
return Task{}, fmt.Errorf("%w: pool type is invalid", ErrTaskInvalid)
|
|
}
|
|
if input.PoolCode == "" {
|
|
input.PoolCode = "default"
|
|
}
|
|
if len(input.PoolCode) > 64 {
|
|
return Task{}, fmt.Errorf("%w: pool code is invalid", ErrTaskInvalid)
|
|
}
|
|
if input.MaxAttempts < 1 || input.MaxAttempts > 10 {
|
|
input.MaxAttempts = 3
|
|
}
|
|
if input.NodeID != nil && strings.TrimSpace(*input.NodeID) != "" {
|
|
var exists bool
|
|
if err := s.pool.QueryRow(ctx, `SELECT EXISTS(SELECT 1 FROM gateway.agent_nodes WHERE id=$1)`, *input.NodeID).Scan(&exists); err != nil {
|
|
return Task{}, fmt.Errorf("%w: %v", ErrStore, err)
|
|
}
|
|
if !exists {
|
|
return Task{}, fmt.Errorf("%w: target node does not exist", ErrTaskInvalid)
|
|
}
|
|
} else {
|
|
input.NodeID = nil
|
|
}
|
|
id, err := platformid.NewUUID()
|
|
if err != nil {
|
|
return Task{}, err
|
|
}
|
|
eventID, err := platformid.NewUUID()
|
|
if err != nil {
|
|
return Task{}, err
|
|
}
|
|
tx, err := s.pool.Begin(ctx)
|
|
if err != nil {
|
|
return Task{}, fmt.Errorf("%w: %v", ErrStore, err)
|
|
}
|
|
defer func() { _ = tx.Rollback(ctx) }()
|
|
_, err = tx.Exec(ctx, `INSERT INTO gateway.agent_tasks(id,task_type,payload,pool_type,pool_code,node_id,max_attempts,created_by) VALUES($1,$2,$3,$4,$5,nullif($6,'')::uuid,$7,nullif($8,'')::uuid)`, id, input.TaskType, input.Payload, input.PoolType, input.PoolCode, input.NodeID, input.MaxAttempts, actorID)
|
|
if err != nil {
|
|
return Task{}, fmt.Errorf("%w: %v", ErrStore, err)
|
|
}
|
|
payload, _ := json.Marshal(map[string]any{"task_id": id, "task_type": input.TaskType, "pool_type": input.PoolType, "pool_code": input.PoolCode})
|
|
if _, err = tx.Exec(ctx, `INSERT INTO gateway.outbox_events(event_id,event_type,event_version,aggregate_type,aggregate_id,payload) VALUES($1,'agent_task.created',1,'agent_task',$2,$3)`, eventID, id, payload); err != nil {
|
|
return Task{}, fmt.Errorf("%w: %v", ErrStore, err)
|
|
}
|
|
if err = tx.Commit(ctx); err != nil {
|
|
return Task{}, fmt.Errorf("%w: %v", ErrStore, err)
|
|
}
|
|
return s.GetTask(ctx, id)
|
|
}
|
|
|
|
// ListTasks 返回任务列表,可按状态过滤。
|
|
func (s *Store) ListTasks(ctx context.Context, status string) ([]Task, error) {
|
|
if s == nil || s.pool == nil {
|
|
return nil, ErrStore
|
|
}
|
|
where, args := " WHERE true", []any{}
|
|
if status != "" {
|
|
args = append(args, status)
|
|
where = " WHERE t.status=$1"
|
|
}
|
|
rows, err := s.pool.Query(ctx, taskSelect+where+` ORDER BY t.created_at DESC LIMIT 200`, args...)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("%w: %v", ErrStore, err)
|
|
}
|
|
defer rows.Close()
|
|
items := []Task{}
|
|
for rows.Next() {
|
|
item, err := scanTask(rows)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
items = append(items, item)
|
|
}
|
|
return items, rows.Err()
|
|
}
|
|
|
|
// GetTask 按 ID 返回任务。
|
|
func (s *Store) GetTask(ctx context.Context, id string) (Task, error) {
|
|
if s == nil || s.pool == nil {
|
|
return Task{}, ErrStore
|
|
}
|
|
return scanTask(s.pool.QueryRow(ctx, taskSelect+` WHERE t.id=$1`, strings.TrimSpace(id)))
|
|
}
|
|
|
|
// CancelTask 取消排队中的任务(已认领/运行中不可取消)。
|
|
func (s *Store) CancelTask(ctx context.Context, id string) error {
|
|
if s == nil || s.pool == nil {
|
|
return ErrStore
|
|
}
|
|
tag, err := s.pool.Exec(ctx, `UPDATE gateway.agent_tasks SET status='cancelled',finished_at=clock_timestamp() WHERE id=$1 AND status='queued'`, strings.TrimSpace(id))
|
|
if err != nil {
|
|
return fmt.Errorf("%w: %v", ErrStore, err)
|
|
}
|
|
if tag.RowsAffected() == 0 {
|
|
return ErrTaskConflict
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// RetryTask 把失败任务重新入队(重置失败状态与退避)。
|
|
func (s *Store) RetryTask(ctx context.Context, id string) error {
|
|
if s == nil || s.pool == nil {
|
|
return ErrStore
|
|
}
|
|
tag, err := s.pool.Exec(ctx, `UPDATE gateway.agent_tasks SET status='queued',error='',available_at=clock_timestamp(),finished_at=NULL,claimed_at=NULL WHERE id=$1 AND status='failed'`, strings.TrimSpace(id))
|
|
if err != nil {
|
|
return fmt.Errorf("%w: %v", ErrStore, err)
|
|
}
|
|
if tag.RowsAffected() == 0 {
|
|
return ErrTaskConflict
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// authenticateNode 校验节点令牌(与心跳同一套恒定时间比较),返回节点行。
|
|
func (s *Store) authenticateNode(ctx context.Context, code, token string) (Node, error) {
|
|
code = strings.ToLower(strings.TrimSpace(code))
|
|
token = strings.TrimSpace(token)
|
|
if !nodeCodePattern.MatchString(code) || token == "" || len(token) > 512 {
|
|
return Node{}, ErrTaskUnauthorized
|
|
}
|
|
var storedHash []byte
|
|
err := s.pool.QueryRow(ctx, `SELECT token_hash FROM gateway.agent_nodes WHERE code=$1 AND enabled`, code).Scan(&storedHash)
|
|
if errors.Is(err, pgx.ErrNoRows) {
|
|
return Node{}, ErrTaskUnauthorized
|
|
}
|
|
if err != nil {
|
|
return Node{}, fmt.Errorf("%w: %v", ErrStore, err)
|
|
}
|
|
digest := sha256.Sum256([]byte(token))
|
|
if subtle.ConstantTimeCompare(digest[:], storedHash) != 1 {
|
|
return Node{}, ErrTaskUnauthorized
|
|
}
|
|
var id string
|
|
if err := s.pool.QueryRow(ctx, `SELECT id::text FROM gateway.agent_nodes WHERE code=$1 AND enabled`, code).Scan(&id); err != nil {
|
|
return Node{}, fmt.Errorf("%w: %v", ErrStore, err)
|
|
}
|
|
return s.Get(ctx, id)
|
|
}
|
|
|
|
// ClaimTask 节点认领下一个可用任务:目标节点(或所在池)中最早排队且到期的任务,
|
|
// SKIP LOCKED 保证多节点并发不重复认领。返回空 Task 表示当前无任务。
|
|
func (s *Store) ClaimTask(ctx context.Context, code, token string) (Task, error) {
|
|
if s == nil || s.pool == nil {
|
|
return Task{}, ErrStore
|
|
}
|
|
node, err := s.authenticateNode(ctx, code, token)
|
|
if err != nil {
|
|
return Task{}, err
|
|
}
|
|
claim, err := platformid.NewUUID()
|
|
if err != nil {
|
|
return Task{}, err
|
|
}
|
|
var item Task
|
|
err = s.pool.QueryRow(ctx, `UPDATE gateway.agent_tasks SET status='claimed',claim_token=$1,claimed_at=clock_timestamp(),available_at=clock_timestamp()+interval '15 minutes'
|
|
WHERE id = (
|
|
SELECT t.id FROM gateway.agent_tasks t
|
|
WHERE t.status='queued' AND t.available_at<=clock_timestamp()
|
|
AND (t.node_id IS NULL OR t.node_id=$2)
|
|
AND (t.node_id IS NOT NULL OR (t.pool_type=$3 AND t.pool_code=$4))
|
|
ORDER BY t.created_at LIMIT 1 FOR UPDATE SKIP LOCKED
|
|
)
|
|
RETURNING id::text,task_type,payload,pool_type,pool_code,node_id::text,status,attempts,max_attempts,result,error,created_at,claimed_at,finished_at,available_at,claim_token`, claim, node.ID, node.PoolType, node.PoolCode).Scan(
|
|
&item.ID, &item.TaskType, &item.Payload, &item.PoolType, &item.PoolCode, &item.NodeID, &item.Status, &item.Attempts, &item.MaxAttempts, &item.Result, &item.Error, &item.CreatedAt, &item.ClaimedAt, &item.FinishedAt, &item.AvailableAt, &item.ClaimToken)
|
|
if errors.Is(err, pgx.ErrNoRows) {
|
|
return Task{}, nil
|
|
}
|
|
if err != nil {
|
|
return Task{}, fmt.Errorf("%w: %v", ErrStore, err)
|
|
}
|
|
item.Payload = normalizeJSON(item.Payload)
|
|
return item, nil
|
|
}
|
|
|
|
// CompleteTask 节点上报任务结果;claim_token 匹配才生效(防重放/防串扰)。
|
|
// 失败且未达最大重试次数时按退避重新入队。
|
|
func (s *Store) CompleteTask(ctx context.Context, code, token, taskID, claimToken string, result json.RawMessage, taskError string) (Task, error) {
|
|
if s == nil || s.pool == nil {
|
|
return Task{}, ErrStore
|
|
}
|
|
if _, err := s.authenticateNode(ctx, code, token); err != nil {
|
|
return Task{}, err
|
|
}
|
|
taskID = strings.TrimSpace(taskID)
|
|
claimToken = strings.TrimSpace(claimToken)
|
|
if taskID == "" || claimToken == "" || len(taskError) > 4000 {
|
|
return Task{}, ErrTaskInvalid
|
|
}
|
|
if len(result) > 0 && !json.Valid(result) {
|
|
return Task{}, fmt.Errorf("%w: result must be valid JSON", ErrTaskInvalid)
|
|
}
|
|
tx, err := s.pool.Begin(ctx)
|
|
if err != nil {
|
|
return Task{}, fmt.Errorf("%w: %v", ErrStore, err)
|
|
}
|
|
defer func() { _ = tx.Rollback(ctx) }()
|
|
var status, currentError string
|
|
var attempts, maxAttempts int
|
|
err = tx.QueryRow(ctx, `SELECT status,error,attempts,max_attempts FROM gateway.agent_tasks WHERE id=$1 AND claim_token=$2`, taskID, claimToken).Scan(&status, ¤tError, &attempts, &maxAttempts)
|
|
if errors.Is(err, pgx.ErrNoRows) {
|
|
return Task{}, ErrTaskConflict
|
|
}
|
|
if err != nil {
|
|
return Task{}, fmt.Errorf("%w: %v", ErrStore, err)
|
|
}
|
|
if status != "claimed" && status != "running" {
|
|
return Task{}, ErrTaskConflict
|
|
}
|
|
attempts++
|
|
nextStatus := "succeeded"
|
|
if taskError != "" {
|
|
nextStatus = "failed"
|
|
if attempts < maxAttempts {
|
|
nextStatus = "queued"
|
|
}
|
|
}
|
|
// 失败退避:30s * 已尝试次数。
|
|
backoff := 30 * time.Second * time.Duration(attempts)
|
|
_, err = tx.Exec(ctx, `UPDATE gateway.agent_tasks SET status=$2::text,attempts=$3::int,result=$4::jsonb,error=$5::text,claim_token=NULL,finished_at=CASE WHEN $2::text IN ('succeeded','failed','cancelled') THEN clock_timestamp() ELSE NULL END,available_at=CASE WHEN $2::text='queued' THEN clock_timestamp()+$6::interval ELSE available_at END WHERE id=$1::uuid`,
|
|
taskID, nextStatus, attempts, normalizeJSON(result), taskError, backoff)
|
|
if err != nil {
|
|
return Task{}, fmt.Errorf("%w: %v", ErrStore, err)
|
|
}
|
|
eventID, _ := platformid.NewUUID()
|
|
eventType := "agent_task.completed"
|
|
if nextStatus == "failed" {
|
|
eventType = "agent_task.failed"
|
|
} else if nextStatus == "queued" {
|
|
eventType = "agent_task.retry"
|
|
}
|
|
payload, _ := json.Marshal(map[string]any{"task_id": taskID, "status": nextStatus, "attempts": attempts, "error": taskError})
|
|
if _, err = tx.Exec(ctx, `INSERT INTO gateway.outbox_events(event_id,event_type,event_version,aggregate_type,aggregate_id,payload) VALUES($1,$2,1,'agent_task',$3,$4)`, eventID, eventType, taskID, payload); err != nil {
|
|
return Task{}, fmt.Errorf("%w: %v", ErrStore, err)
|
|
}
|
|
if err = tx.Commit(ctx); err != nil {
|
|
return Task{}, fmt.Errorf("%w: %v", ErrStore, err)
|
|
}
|
|
return s.GetTask(ctx, taskID)
|
|
}
|