0.10.1: 安全与业务逻辑加固、新品牌与部署加固

三轮审查修复(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 币种维度)
This commit is contained in:
2026-08-13 10:50:51 +08:00
parent b536672000
commit 9501751792
136 changed files with 8024 additions and 1476 deletions
+218
View File
@@ -0,0 +1,218 @@
package agentnode
import (
"encoding/json"
"errors"
"net"
"net/http"
"strings"
"aigateway.local/core/internal/identity"
"aigateway.local/core/internal/platform/apiresponse"
)
type HTTPHandler struct {
store *Store
identity *identity.Service
mux *http.ServeMux
}
type nodeRequest struct {
Code string `json:"code"`
Name string `json:"name"`
Description string `json:"description"`
Endpoint string `json:"endpoint"`
NodeType string `json:"node_type"`
PoolType string `json:"pool_type"`
PoolCode string `json:"pool_code"`
Enabled *bool `json:"enabled"`
}
type routePreviewRequest struct {
PoolType string `json:"pool_type"`
PoolCode string `json:"pool_code"`
RequiredCapabilities []string `json:"required_capabilities"`
RequestKey string `json:"request_key"`
}
func NewHTTPHandler(store *Store, identityService *identity.Service) *HTTPHandler {
h := &HTTPHandler{store: store, identity: identityService, mux: http.NewServeMux()}
h.mux.HandleFunc("GET /api/v1/admin/agent-nodes", h.list)
h.mux.HandleFunc("POST /api/v1/admin/agent-nodes", h.create)
h.mux.HandleFunc("POST /api/v1/admin/agent-nodes/route-preview", h.routePreview)
h.mux.HandleFunc("PUT /api/v1/admin/agent-nodes/{id}", h.update)
h.mux.HandleFunc("DELETE /api/v1/admin/agent-nodes/{id}", h.delete)
h.mux.HandleFunc("POST /api/v1/admin/agent-nodes/{id}/rotate-token", h.rotateToken)
h.mux.HandleFunc("POST /api/v1/agent/nodes/{code}/heartbeat", h.heartbeat)
return h
}
func (h *HTTPHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { h.mux.ServeHTTP(w, r) }
func (h *HTTPHandler) require(w http.ResponseWriter, r *http.Request, permission string) (identity.Account, bool) {
account, err := h.identity.Authenticate(r.Context(), identity.KindAdmin, r.Header.Get("Authorization"))
if err != nil {
apiresponse.Error(w, http.StatusUnauthorized, "登录状态无效或已过期")
return identity.Account{}, false
}
if !identity.HasPermission(account, permission) {
apiresponse.Error(w, http.StatusForbidden, "缺少智能体节点操作权限")
return identity.Account{}, false
}
return account, true
}
func decodeJSON(w http.ResponseWriter, r *http.Request, target any) bool {
decoder := json.NewDecoder(http.MaxBytesReader(w, r.Body, 1<<20))
decoder.DisallowUnknownFields()
if err := decoder.Decode(target); err != nil {
apiresponse.Error(w, http.StatusBadRequest, "请求格式无效")
return false
}
return true
}
func (h *HTTPHandler) list(w http.ResponseWriter, r *http.Request) {
if _, ok := h.require(w, r, identity.PermissionAgentNodeRead); !ok {
return
}
items, err := h.store.List(r.Context())
if err != nil {
writeError(w, err)
return
}
apiresponse.OK(w, items)
}
func (h *HTTPHandler) routePreview(w http.ResponseWriter, r *http.Request) {
if _, ok := h.require(w, r, identity.PermissionAgentNodeRead); !ok {
return
}
var input routePreviewRequest
if !decodeJSON(w, r, &input) {
return
}
preview, err := h.store.PreviewRoute(r.Context(), RoutePreviewInput{
PoolType: input.PoolType, PoolCode: input.PoolCode,
RequiredCapabilities: input.RequiredCapabilities, RequestKey: input.RequestKey,
})
if err != nil {
writeError(w, err)
return
}
apiresponse.OK(w, preview)
}
func (h *HTTPHandler) create(w http.ResponseWriter, r *http.Request) {
actor, ok := h.require(w, r, identity.PermissionAgentNodeManage)
if !ok {
return
}
var input nodeRequest
if !decodeJSON(w, r, &input) {
return
}
enabled := true
if input.Enabled != nil {
enabled = *input.Enabled
}
node, token, err := h.store.Create(r.Context(), CreateInput{Code: input.Code, Name: input.Name, Description: input.Description, Endpoint: input.Endpoint, NodeType: input.NodeType, PoolType: input.PoolType, PoolCode: input.PoolCode, Enabled: enabled}, actor.ID)
if err != nil {
writeError(w, err)
return
}
apiresponse.OK(w, map[string]any{"node": node, "token": token, "warning": "令牌只显示一次,请立即安全保存并配置到节点"})
}
func (h *HTTPHandler) update(w http.ResponseWriter, r *http.Request) {
if _, ok := h.require(w, r, identity.PermissionAgentNodeManage); !ok {
return
}
var input nodeRequest
if !decodeJSON(w, r, &input) {
return
}
if input.Enabled == nil {
apiresponse.Error(w, http.StatusBadRequest, "enabled 字段不能为空")
return
}
node, err := h.store.Update(r.Context(), UpdateInput{ID: r.PathValue("id"), Name: input.Name, Description: input.Description, Endpoint: input.Endpoint, NodeType: input.NodeType, PoolType: input.PoolType, PoolCode: input.PoolCode, Enabled: *input.Enabled})
if err != nil {
writeError(w, err)
return
}
apiresponse.OK(w, node)
}
func (h *HTTPHandler) delete(w http.ResponseWriter, r *http.Request) {
if _, ok := h.require(w, r, identity.PermissionAgentNodeManage); !ok {
return
}
if err := h.store.Delete(r.Context(), r.PathValue("id")); err != nil {
writeError(w, err)
return
}
apiresponse.OK(w, map[string]bool{"deleted": true})
}
func (h *HTTPHandler) rotateToken(w http.ResponseWriter, r *http.Request) {
if _, ok := h.require(w, r, identity.PermissionAgentNodeManage); !ok {
return
}
node, token, err := h.store.RotateToken(r.Context(), r.PathValue("id"))
if err != nil {
writeError(w, err)
return
}
apiresponse.OK(w, map[string]any{"node": node, "token": token, "warning": "旧令牌已立即失效,新令牌只显示一次"})
}
func (h *HTTPHandler) heartbeat(w http.ResponseWriter, r *http.Request) {
var input HeartbeatInput
if !decodeJSON(w, r, &input) {
return
}
remoteIP := parseRemoteIP(r.RemoteAddr)
node, err := h.store.Heartbeat(r.Context(), r.PathValue("code"), r.Header.Get("X-Agent-Token"), remoteIP, input)
if err != nil {
writeHeartbeatError(w, err)
return
}
apiresponse.OK(w, map[string]any{"accepted": true, "node": node})
}
func parseRemoteIP(remoteAddr string) net.IP {
host, _, err := net.SplitHostPort(strings.TrimSpace(remoteAddr))
if err != nil {
host = strings.TrimSpace(remoteAddr)
}
return net.ParseIP(host)
}
func writeError(w http.ResponseWriter, err error) {
switch {
case errors.Is(err, ErrNotFound):
apiresponse.Error(w, http.StatusNotFound, "智能体节点不存在")
case errors.Is(err, ErrConflict):
apiresponse.Error(w, http.StatusConflict, "节点编码已存在")
case errors.Is(err, ErrInvalidInput):
apiresponse.Error(w, http.StatusBadRequest, err.Error())
case errors.Is(err, ErrStore):
apiresponse.Error(w, http.StatusServiceUnavailable, "智能体节点服务暂不可用")
default:
apiresponse.Error(w, http.StatusInternalServerError, "智能体节点处理失败")
}
}
func writeHeartbeatError(w http.ResponseWriter, err error) {
switch {
case errors.Is(err, ErrInvalidToken):
apiresponse.Error(w, http.StatusUnauthorized, "节点令牌无效或节点已停用")
case errors.Is(err, ErrInvalidInput):
apiresponse.Error(w, http.StatusBadRequest, err.Error())
case errors.Is(err, ErrStore):
apiresponse.Error(w, http.StatusServiceUnavailable, "智能体节点服务暂不可用")
default:
apiresponse.Error(w, http.StatusInternalServerError, "节点心跳处理失败")
}
}
+60
View File
@@ -0,0 +1,60 @@
package agentnode
import (
"encoding/json"
"testing"
"time"
)
func routeTestNode(id, status string, enabled bool, capabilities map[string]any) Node {
raw, _ := json.Marshal(capabilities)
return Node{ID: id, Code: id, Status: status, Enabled: enabled, Capabilities: raw, LastHeartbeatAt: func() *time.Time { now := time.Now(); return &now }()}
}
func TestSelectRouteCandidatesFiltersCapabilitiesAndStatus(t *testing.T) {
nodes := []Node{
routeTestNode("online-capable", "online", true, map[string]any{"tool_exec": true, "region": "cn"}),
routeTestNode("online-disabled-capability", "online", true, map[string]any{"tool_exec": false}),
routeTestNode("online-missing-capability", "online", true, map[string]any{"region": "cn"}),
routeTestNode("offline-capable", "offline", true, map[string]any{"tool_exec": true}),
routeTestNode("disabled-capable", "online", false, map[string]any{"tool_exec": true}),
}
selected := selectRouteCandidates(nodes, "request-1", []string{"tool_exec"})
if len(selected) != 1 || selected[0].ID != "online-capable" {
t.Fatalf("unexpected candidates: %#v", selected)
}
}
func TestSelectRouteCandidatesIsDeterministic(t *testing.T) {
nodes := []Node{
routeTestNode("node-a", "online", true, nil),
routeTestNode("node-b", "online", true, nil),
routeTestNode("node-c", "online", true, nil),
}
first := selectRouteCandidates(nodes, "request-42", nil)
second := selectRouteCandidates([]Node{nodes[2], nodes[0], nodes[1]}, "request-42", nil)
if len(first) != len(second) || len(first) == 0 {
t.Fatalf("candidate lengths differ: %d %d", len(first), len(second))
}
for index := range first {
if first[index].ID != second[index].ID {
t.Fatalf("selection order is not stable: first=%v second=%v", first, second)
}
}
if first[0].ID == first[1].ID {
t.Fatal("candidate order contains duplicates")
}
}
func TestNormalizeRoutePreviewInput(t *testing.T) {
input, err := normalizeRoutePreviewInput(RoutePreviewInput{
PoolType: " PUBLIC ", PoolCode: "shared", RequestKey: " request-1 ",
RequiredCapabilities: []string{"tool_exec", " tool_exec ", ""},
})
if err != nil || input.PoolType != "public" || input.RequestKey != "request-1" || len(input.RequiredCapabilities) != 1 {
t.Fatalf("normalized input=%+v err=%v", input, err)
}
if _, err := normalizeRoutePreviewInput(RoutePreviewInput{PoolType: "public", PoolCode: "shared"}); err == nil {
t.Fatal("empty request key must be rejected")
}
}
+502
View File
@@ -0,0 +1,502 @@
package agentnode
import (
"context"
"crypto/rand"
"crypto/sha256"
"crypto/subtle"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"net"
"net/url"
"regexp"
"sort"
"strings"
"time"
platformid "aigateway.local/core/internal/platform/id"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgconn"
"github.com/jackc/pgx/v5/pgxpool"
)
var (
ErrNotFound = errors.New("agent node not found")
ErrConflict = errors.New("agent node already exists")
ErrInvalidToken = errors.New("agent node token invalid")
ErrInvalidInput = errors.New("agent node input invalid")
ErrStore = errors.New("agent node store unavailable")
)
var nodeCodePattern = regexp.MustCompile(`^[a-z0-9][a-z0-9._-]{0,127}$`)
type Store struct{ pool *pgxpool.Pool }
func NewStore(pool *pgxpool.Pool) *Store { return &Store{pool: pool} }
type Node struct {
ID string `json:"id"`
Code string `json:"code"`
Name string `json:"name"`
Description string `json:"description"`
Endpoint string `json:"endpoint"`
NodeType string `json:"node_type"`
PoolType string `json:"pool_type"`
PoolCode string `json:"pool_code"`
Enabled bool `json:"enabled"`
Status string `json:"status"`
TokenPrefix string `json:"token_prefix"`
Version string `json:"version"`
Capabilities json.RawMessage `json:"capabilities"`
Metadata json.RawMessage `json:"metadata"`
LastHeartbeatAt *time.Time `json:"last_heartbeat_at,omitempty"`
LastHeartbeatIP string `json:"last_heartbeat_ip,omitempty"`
LastError string `json:"last_error"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
type CreateInput struct {
Code string
Name string
Description string
Endpoint string
NodeType string
PoolType string
PoolCode string
Enabled bool
}
type UpdateInput struct {
ID string
Name string
Description string
Endpoint string
NodeType string
PoolType string
PoolCode string
Enabled bool
}
type HeartbeatInput struct {
Version string `json:"version"`
Capabilities map[string]any `json:"capabilities"`
Metadata map[string]any `json:"metadata"`
Error string `json:"error"`
}
// RoutePreviewInput describes the node-pool constraints used by the read-only
// routing preview. It deliberately contains no task or request payload: the
// preview only validates candidate selection before a remote executor exists.
type RoutePreviewInput struct {
PoolType string `json:"pool_type"`
PoolCode string `json:"pool_code"`
RequiredCapabilities []string `json:"required_capabilities"`
RequestKey string `json:"request_key"`
}
type RoutePreview struct {
PoolType string `json:"pool_type"`
PoolCode string `json:"pool_code"`
RequiredCapabilities []string `json:"required_capabilities"`
RequestKey string `json:"request_key"`
SelectionPolicy string `json:"selection_policy"`
Reason string `json:"reason"`
Selected *Node `json:"selected"`
Candidates []Node `json:"candidates"`
}
const nodeSelect = `SELECT n.id::text,n.code,n.name,n.description,n.endpoint,n.node_type,n.pool_type,n.pool_code,n.enabled,
CASE WHEN NOT n.enabled THEN 'disabled' WHEN n.last_heartbeat_at IS NULL THEN 'pending' WHEN n.last_heartbeat_at < clock_timestamp()-interval '90 seconds' THEN 'offline' ELSE 'online' END,
n.token_prefix,n.version,n.capabilities,n.metadata,n.last_heartbeat_at,coalesce(host(n.last_heartbeat_ip),''),n.last_error,n.created_at,n.updated_at
FROM gateway.agent_nodes n`
func normalizeJSON(raw []byte) json.RawMessage {
if len(raw) == 0 || !json.Valid(raw) {
return json.RawMessage(`{}`)
}
return raw
}
func objectJSON(value map[string]any) ([]byte, error) {
if value == nil {
return nil, nil
}
raw, err := json.Marshal(value)
if err != nil {
return nil, fmt.Errorf("%w: metadata cannot be encoded", ErrInvalidInput)
}
return raw, nil
}
func validateCommon(code, name, description, endpoint, nodeType, poolType, poolCode string) error {
if !nodeCodePattern.MatchString(code) || strings.ToLower(code) != code {
return fmt.Errorf("%w: code must use lowercase letters, numbers, dot, underscore or hyphen", ErrInvalidInput)
}
if strings.TrimSpace(name) == "" || len(name) > 128 || len(description) > 4000 || len(endpoint) > 512 {
return fmt.Errorf("%w: node fields exceed their limits", ErrInvalidInput)
}
if nodeType != "worker" && nodeType != "gateway" && nodeType != "executor" {
return fmt.Errorf("%w: node type is invalid", ErrInvalidInput)
}
if poolType != "public" && poolType != "private" {
return fmt.Errorf("%w: pool type is invalid", ErrInvalidInput)
}
if strings.TrimSpace(poolCode) == "" || len(poolCode) > 64 {
return fmt.Errorf("%w: pool code is invalid", ErrInvalidInput)
}
// Endpoint 将来可能被节点池路由直接拨号,必须保证是干净的 http(s)
// 绝对地址(无 userinfo/query/fragment)。不做 DNS 解析:节点本身常部署
// 在内网,不能按公网规则校验。
if strings.TrimSpace(endpoint) != "" {
parsed, parseErr := url.Parse(strings.TrimSpace(endpoint))
if parseErr != nil || (parsed.Scheme != "http" && parsed.Scheme != "https") || parsed.Hostname() == "" ||
parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" {
return fmt.Errorf("%w: endpoint must be an absolute http(s) URL without user info, query or fragment", ErrInvalidInput)
}
}
return nil
}
func generateToken() (string, string, []byte, error) {
raw := make([]byte, 32)
if _, err := rand.Read(raw); err != nil {
return "", "", nil, err
}
secret := "agn_" + base64.RawURLEncoding.EncodeToString(raw)
prefix := secret[:12]
digest := sha256.Sum256([]byte(secret))
return secret, prefix, digest[:], nil
}
func scanNode(row pgx.Row) (Node, error) {
var item Node
err := row.Scan(&item.ID, &item.Code, &item.Name, &item.Description, &item.Endpoint, &item.NodeType, &item.PoolType, &item.PoolCode, &item.Enabled, &item.Status, &item.TokenPrefix, &item.Version, &item.Capabilities, &item.Metadata, &item.LastHeartbeatAt, &item.LastHeartbeatIP, &item.LastError, &item.CreatedAt, &item.UpdatedAt)
if errors.Is(err, pgx.ErrNoRows) {
return Node{}, ErrNotFound
}
item.Capabilities = normalizeJSON(item.Capabilities)
item.Metadata = normalizeJSON(item.Metadata)
return item, err
}
func (s *Store) Create(ctx context.Context, input CreateInput, actorID string) (Node, string, error) {
if s == nil || s.pool == nil {
return Node{}, "", ErrStore
}
input.Code = strings.ToLower(strings.TrimSpace(input.Code))
input.Name = strings.TrimSpace(input.Name)
input.Description = strings.TrimSpace(input.Description)
input.Endpoint = strings.TrimSpace(input.Endpoint)
input.NodeType = strings.TrimSpace(input.NodeType)
input.PoolType = strings.TrimSpace(input.PoolType)
input.PoolCode = strings.TrimSpace(input.PoolCode)
if input.NodeType == "" {
input.NodeType = "worker"
}
if input.PoolType == "" {
input.PoolType = "private"
}
if input.PoolCode == "" {
input.PoolCode = "default"
}
if err := validateCommon(input.Code, input.Name, input.Description, input.Endpoint, input.NodeType, input.PoolType, input.PoolCode); err != nil {
return Node{}, "", err
}
id, err := platformid.NewUUID()
if err != nil {
return Node{}, "", err
}
secret, prefix, digest, err := generateToken()
if err != nil {
return Node{}, "", err
}
_, err = s.pool.Exec(ctx, `INSERT INTO gateway.agent_nodes(id,code,name,description,endpoint,node_type,pool_type,pool_code,enabled,token_prefix,token_hash,created_by) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,nullif($12,'')::uuid)`, id, input.Code, input.Name, input.Description, input.Endpoint, input.NodeType, input.PoolType, input.PoolCode, input.Enabled, prefix, digest, actorID)
if err != nil {
var pgErr *pgconn.PgError
if errors.As(err, &pgErr) && pgErr.Code == "23505" {
return Node{}, "", ErrConflict
}
return Node{}, "", fmt.Errorf("%w: %v", ErrStore, err)
}
item, err := s.Get(ctx, id)
return item, secret, err
}
func (s *Store) Get(ctx context.Context, id string) (Node, error) {
if s == nil || s.pool == nil {
return Node{}, ErrStore
}
return scanNode(s.pool.QueryRow(ctx, nodeSelect+` WHERE n.id=$1`, id))
}
func (s *Store) List(ctx context.Context) ([]Node, error) {
if s == nil || s.pool == nil {
return nil, ErrStore
}
rows, err := s.pool.Query(ctx, nodeSelect+` ORDER BY n.updated_at DESC,n.code`)
if err != nil {
return nil, fmt.Errorf("%w: %v", ErrStore, err)
}
defer rows.Close()
items := make([]Node, 0)
for rows.Next() {
item, scanErr := scanNode(rows)
if scanErr != nil {
return nil, fmt.Errorf("%w: %v", ErrStore, scanErr)
}
items = append(items, item)
}
return items, rows.Err()
}
func normalizeRoutePreviewInput(input RoutePreviewInput) (RoutePreviewInput, error) {
input.PoolType = strings.TrimSpace(strings.ToLower(input.PoolType))
input.PoolCode = strings.TrimSpace(input.PoolCode)
input.RequestKey = strings.TrimSpace(input.RequestKey)
if input.PoolType != "public" && input.PoolType != "private" {
return RoutePreviewInput{}, fmt.Errorf("%w: pool type is invalid", ErrInvalidInput)
}
if input.PoolCode == "" || len(input.PoolCode) > 64 {
return RoutePreviewInput{}, fmt.Errorf("%w: pool code is invalid", ErrInvalidInput)
}
if input.RequestKey == "" || len(input.RequestKey) > 512 {
return RoutePreviewInput{}, fmt.Errorf("%w: request key is invalid", ErrInvalidInput)
}
capabilities := make([]string, 0, len(input.RequiredCapabilities))
seen := make(map[string]struct{}, len(input.RequiredCapabilities))
for _, capability := range input.RequiredCapabilities {
capability = strings.TrimSpace(capability)
if capability == "" {
continue
}
if len(capability) > 128 {
return RoutePreviewInput{}, fmt.Errorf("%w: capability is too long", ErrInvalidInput)
}
if _, ok := seen[capability]; ok {
continue
}
seen[capability] = struct{}{}
capabilities = append(capabilities, capability)
}
if len(capabilities) > 32 {
return RoutePreviewInput{}, fmt.Errorf("%w: too many required capabilities", ErrInvalidInput)
}
input.RequiredCapabilities = capabilities
return input, nil
}
func capabilityEnabled(value any) bool {
switch typed := value.(type) {
case nil:
return false
case bool:
return typed
case string:
value := strings.TrimSpace(strings.ToLower(typed))
return value != "" && value != "false" && value != "0" && value != "no"
case float64:
return typed != 0
default:
return true
}
}
func nodeHasCapabilities(node Node, required []string) bool {
if len(required) == 0 {
return true
}
var capabilities map[string]any
if err := json.Unmarshal(node.Capabilities, &capabilities); err != nil {
return false
}
for _, capability := range required {
value, ok := capabilities[capability]
if !ok || !capabilityEnabled(value) {
return false
}
}
return true
}
func orderRouteCandidates(nodes []Node, requestKey string) []Node {
ordered := append([]Node(nil), nodes...)
type candidateHash struct {
digest [32]byte
id string
}
hashes := make(map[string]candidateHash, len(ordered))
for _, node := range ordered {
hashes[node.ID] = candidateHash{digest: sha256.Sum256([]byte(requestKey + "\x00" + node.ID)), id: node.ID}
}
sort.SliceStable(ordered, func(i, j int) bool {
left, right := hashes[ordered[i].ID], hashes[ordered[j].ID]
if string(left.digest[:]) == string(right.digest[:]) {
return left.id < right.id
}
return string(left.digest[:]) < string(right.digest[:])
})
return ordered
}
func selectRouteCandidates(nodes []Node, requestKey string, required []string) []Node {
filtered := make([]Node, 0, len(nodes))
for _, node := range nodes {
if node.Status != "online" || !node.Enabled || !nodeHasCapabilities(node, required) {
continue
}
filtered = append(filtered, node)
}
return orderRouteCandidates(filtered, requestKey)
}
// PreviewRoute returns the online, capability-compatible nodes in stable
// request-key order. It is intentionally read-only and does not invoke an
// endpoint or enqueue a task.
func (s *Store) PreviewRoute(ctx context.Context, input RoutePreviewInput) (RoutePreview, error) {
if s == nil || s.pool == nil {
return RoutePreview{}, ErrStore
}
normalized, err := normalizeRoutePreviewInput(input)
if err != nil {
return RoutePreview{}, err
}
rows, err := s.pool.Query(ctx, nodeSelect+` WHERE n.pool_type=$1 AND n.pool_code=$2 AND n.enabled AND n.last_heartbeat_at IS NOT NULL AND n.last_heartbeat_at >= clock_timestamp()-interval '90 seconds' ORDER BY n.code`, normalized.PoolType, normalized.PoolCode)
if err != nil {
return RoutePreview{}, fmt.Errorf("%w: %v", ErrStore, err)
}
defer rows.Close()
online := make([]Node, 0)
for rows.Next() {
item, scanErr := scanNode(rows)
if scanErr != nil {
return RoutePreview{}, fmt.Errorf("%w: %v", ErrStore, scanErr)
}
online = append(online, item)
}
if err := rows.Err(); err != nil {
return RoutePreview{}, fmt.Errorf("%w: %v", ErrStore, err)
}
candidates := selectRouteCandidates(online, normalized.RequestKey, normalized.RequiredCapabilities)
preview := RoutePreview{
PoolType: normalized.PoolType, PoolCode: normalized.PoolCode,
RequiredCapabilities: normalized.RequiredCapabilities, RequestKey: normalized.RequestKey,
SelectionPolicy: "stable-hash(request_key,node_id)", Candidates: candidates,
}
switch {
case len(online) == 0:
preview.Reason = "no_online_node"
case len(candidates) == 0:
preview.Reason = "no_capable_node"
default:
preview.Reason = "selected_online_node"
preview.Selected = &preview.Candidates[0]
}
return preview, nil
}
func (s *Store) Update(ctx context.Context, input UpdateInput) (Node, error) {
if s == nil || s.pool == nil {
return Node{}, ErrStore
}
input.ID = strings.TrimSpace(input.ID)
input.Name = strings.TrimSpace(input.Name)
input.Description = strings.TrimSpace(input.Description)
input.Endpoint = strings.TrimSpace(input.Endpoint)
input.NodeType = strings.TrimSpace(input.NodeType)
input.PoolType = strings.TrimSpace(input.PoolType)
input.PoolCode = strings.TrimSpace(input.PoolCode)
if err := validateCommon("valid-node", input.Name, input.Description, input.Endpoint, input.NodeType, input.PoolType, input.PoolCode); err != nil {
return Node{}, err
}
tag, err := s.pool.Exec(ctx, `UPDATE gateway.agent_nodes SET name=$2,description=$3,endpoint=$4,node_type=$5,pool_type=$6,pool_code=$7,enabled=$8,updated_at=clock_timestamp() WHERE id=$1`, input.ID, input.Name, input.Description, input.Endpoint, input.NodeType, input.PoolType, input.PoolCode, input.Enabled)
if err != nil {
return Node{}, fmt.Errorf("%w: %v", ErrStore, err)
}
if tag.RowsAffected() == 0 {
return Node{}, ErrNotFound
}
return s.Get(ctx, input.ID)
}
func (s *Store) Delete(ctx context.Context, id string) error {
if s == nil || s.pool == nil {
return ErrStore
}
tag, err := s.pool.Exec(ctx, `DELETE FROM gateway.agent_nodes WHERE id=$1`, strings.TrimSpace(id))
if err != nil {
return fmt.Errorf("%w: %v", ErrStore, err)
}
if tag.RowsAffected() == 0 {
return ErrNotFound
}
return nil
}
func (s *Store) RotateToken(ctx context.Context, id string) (Node, string, error) {
if s == nil || s.pool == nil {
return Node{}, "", ErrStore
}
secret, prefix, digest, err := generateToken()
if err != nil {
return Node{}, "", err
}
tag, err := s.pool.Exec(ctx, `UPDATE gateway.agent_nodes SET token_prefix=$2,token_hash=$3,updated_at=clock_timestamp() WHERE id=$1`, strings.TrimSpace(id), prefix, digest)
if err != nil {
return Node{}, "", fmt.Errorf("%w: %v", ErrStore, err)
}
if tag.RowsAffected() == 0 {
return Node{}, "", ErrNotFound
}
item, err := s.Get(ctx, id)
return item, secret, err
}
func (s *Store) Heartbeat(ctx context.Context, code, token string, remoteIP net.IP, input HeartbeatInput) (Node, error) {
if s == nil || s.pool == nil {
return Node{}, ErrStore
}
code = strings.ToLower(strings.TrimSpace(code))
token = strings.TrimSpace(token)
if !nodeCodePattern.MatchString(code) || token == "" || len(token) > 512 || len(input.Version) > 128 || len(input.Error) > 4000 {
return Node{}, ErrInvalidInput
}
capabilities, err := objectJSON(input.Capabilities)
if err != nil {
return Node{}, err
}
metadata, err := objectJSON(input.Metadata)
if err != nil {
return Node{}, err
}
digest := sha256.Sum256([]byte(token))
ip := ""
if remoteIP != nil {
ip = remoteIP.String()
}
// 先按 code 取出令牌哈希,在 Go 侧做恒定时间比较:未知 code 与错误
// 令牌返回同一个错误,避免通过 404/401 差异枚举有效节点;数据库端
// bytea 比较可能提前短路,不做恒定时间保证。
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{}, ErrInvalidToken
}
if err != nil {
return Node{}, fmt.Errorf("%w: %v", ErrStore, err)
}
if subtle.ConstantTimeCompare(digest[:], storedHash) != 1 {
return Node{}, ErrInvalidToken
}
var id string
err = s.pool.QueryRow(ctx, `UPDATE gateway.agent_nodes SET version=$3,capabilities=coalesce($4::jsonb,capabilities),metadata=coalesce($5::jsonb,metadata),last_error=$6,last_heartbeat_at=clock_timestamp(),last_heartbeat_ip=nullif($7,'')::inet,updated_at=clock_timestamp() WHERE code=$1 AND token_hash=$2 AND enabled RETURNING id::text`, code, digest[:], input.Version, capabilities, metadata, strings.TrimSpace(input.Error), ip).Scan(&id)
if errors.Is(err, pgx.ErrNoRows) {
return Node{}, ErrInvalidToken
}
if err != nil {
return Node{}, fmt.Errorf("%w: %v", ErrStore, err)
}
return s.Get(ctx, id)
}
@@ -0,0 +1,61 @@
package agentnode
import (
"context"
"net"
"os"
"testing"
"time"
"aigateway.local/core/internal/platform/config"
"aigateway.local/core/internal/platform/database"
)
func TestAgentNodePostgreSQLLifecycle(t *testing.T) {
databaseURL := os.Getenv("AGENT_NODE_TEST_DATABASE_URL")
if databaseURL == "" {
t.Skip("AGENT_NODE_TEST_DATABASE_URL is not set")
}
ctx := context.Background()
pool, err := database.Open(ctx, config.Database{URL: databaseURL, MaxConns: 4, MinConns: 0})
if err != nil {
t.Fatal(err)
}
defer pool.Close()
_, _ = pool.Exec(ctx, `DELETE FROM gateway.agent_nodes WHERE code='node-integration'`)
defer pool.Exec(ctx, `DELETE FROM gateway.agent_nodes WHERE code='node-integration'`)
store := NewStore(pool)
node, token, err := store.Create(ctx, CreateInput{Code: "node-integration", Name: "Integration Node", PoolType: "private", PoolCode: "test", Enabled: true}, "")
if err != nil {
t.Fatal(err)
}
if token == "" || node.Status != "pending" || node.TokenPrefix == "" {
t.Fatalf("created node=%+v token=%q", node, token)
}
updated, err := store.Update(ctx, UpdateInput{ID: node.ID, Name: "Integration Node v2", Description: "test", Endpoint: "https://node.invalid", NodeType: "executor", PoolType: "public", PoolCode: "shared", Enabled: true})
if err != nil || updated.Name != "Integration Node v2" || updated.PoolType != "public" {
t.Fatalf("updated node=%+v err=%v", updated, err)
}
online, err := store.Heartbeat(ctx, node.Code, token, net.ParseIP("192.0.2.10"), HeartbeatInput{Version: "0.10.0-node", Capabilities: map[string]any{"tool_exec": true}, Metadata: map[string]any{"region": "test"}})
if err != nil || online.Status != "online" || online.Version != "0.10.0-node" || online.LastHeartbeatIP != "192.0.2.10" {
t.Fatalf("heartbeat node=%+v err=%v", online, err)
}
preview, err := store.PreviewRoute(ctx, RoutePreviewInput{PoolType: "public", PoolCode: "shared", RequestKey: "integration-request", RequiredCapabilities: []string{"tool_exec"}})
if err != nil || preview.Reason != "selected_online_node" || preview.Selected == nil || preview.Selected.ID != online.ID || len(preview.Candidates) != 1 {
t.Fatalf("route preview=%+v err=%v", preview, err)
}
rotated, newToken, err := store.RotateToken(ctx, node.ID)
if err != nil || newToken == token || rotated.TokenPrefix == node.TokenPrefix {
t.Fatalf("rotated node=%+v token=%q err=%v", rotated, newToken, err)
}
if _, err = store.Heartbeat(ctx, node.Code, token, nil, HeartbeatInput{}); err != ErrInvalidToken {
t.Fatalf("old token err=%v", err)
}
if _, err = store.Heartbeat(ctx, node.Code, newToken, nil, HeartbeatInput{}); err != nil {
t.Fatal(err)
}
items, err := store.List(ctx)
if err != nil || len(items) == 0 || items[0].UpdatedAt.Before(time.Now().UTC().Add(-time.Minute)) {
t.Fatalf("items=%+v err=%v", items, err)
}
}