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:
@@ -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, "节点心跳处理失败")
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user