Files
ai-gateway-go/internal/contentpolicy/store.go
T
superidou 5759c1862e AI Gateway Go 0.10.0 源码快照 + 旗舰版需求规划报告
M0-M7 已完成:核心网关(身份/RBAC/TOTP/OIDC/SAML/Provider/配额/路由/内容策略/审计/定价)+ 资源市场(MCP/Skills/数字员工)。
含 22 个 PostgreSQL 迁移、管理端/门户端前端源码、OpenAPI 契约、部署 compose。

Co-Authored-By: Claude <noreply@anthropic.com>
2026-08-12 11:45:54 +08:00

157 lines
5.4 KiB
Go

package contentpolicy
import (
"context"
"encoding/json"
"errors"
"fmt"
"strings"
platformid "aigateway.local/core/internal/platform/id"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgtype"
"github.com/jackc/pgx/v5/pgxpool"
)
var ErrNotFound = errors.New("content policy not found")
type Store struct{ pool *pgxpool.Pool }
func NewStore(pool *pgxpool.Pool) *Store { return &Store{pool: pool} }
func (s *Store) List(ctx context.Context) ([]Policy, error) {
rows, err := s.pool.Query(ctx, `SELECT id::text,name,description,action,priority,paths,models,api_key_ids::text[],rules,enabled,revision,created_at,updated_at FROM gateway.content_policies ORDER BY priority DESC,name`)
if err != nil {
return nil, err
}
defer rows.Close()
result := []Policy{}
for rows.Next() {
var p Policy
var raw []byte
if err := rows.Scan(&p.ID, &p.Name, &p.Description, &p.Action, &p.Priority, &p.Paths, &p.Models, &p.APIKeyIDs, &raw, &p.Enabled, &p.Revision, &p.CreatedAt, &p.UpdatedAt); err != nil {
return nil, err
}
if err := json.Unmarshal(raw, &p.Rules); err != nil {
return nil, err
}
result = append(result, p)
}
return result, rows.Err()
}
func (s *Store) Save(ctx context.Context, p Policy, actorID string, create bool) (Policy, error) {
if err := Validate(p); err != nil {
return Policy{}, err
}
rules, _ := json.Marshal(p.Rules)
eventID, _ := platformid.NewUUID()
tx, err := s.pool.Begin(ctx)
if err != nil {
return Policy{}, err
}
defer func() { _ = tx.Rollback(ctx) }()
if create {
p.ID, _ = platformid.NewUUID()
err = tx.QueryRow(ctx, `INSERT INTO gateway.content_policies(id,name,description,action,priority,paths,models,api_key_ids,rules,enabled,created_by) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11) RETURNING revision,created_at,updated_at`, p.ID, p.Name, p.Description, p.Action, p.Priority, p.Paths, p.Models, p.APIKeyIDs, rules, p.Enabled, actorID).Scan(&p.Revision, &p.CreatedAt, &p.UpdatedAt)
} else {
err = tx.QueryRow(ctx, `UPDATE gateway.content_policies SET name=$2,description=$3,action=$4,priority=$5,paths=$6,models=$7,api_key_ids=$8,rules=$9,enabled=$10,revision=revision+1,updated_at=clock_timestamp() WHERE id=$1 RETURNING revision,created_at,updated_at`, p.ID, p.Name, p.Description, p.Action, p.Priority, p.Paths, p.Models, p.APIKeyIDs, rules, p.Enabled).Scan(&p.Revision, &p.CreatedAt, &p.UpdatedAt)
}
if errors.Is(err, pgx.ErrNoRows) {
return Policy{}, ErrNotFound
}
if err != nil {
return Policy{}, err
}
eventType := "content_policy.updated"
if create {
eventType = "content_policy.created"
}
payload, _ := json.Marshal(map[string]any{"content_policy_id": p.ID, "actor_id": actorID, "revision": p.Revision})
_, err = tx.Exec(ctx, `INSERT INTO gateway.outbox_events(event_id,event_type,event_version,aggregate_type,aggregate_id,payload) VALUES($1,$2,1,'content_policy',$3,$4)`, eventID, eventType, p.ID, payload)
if err != nil {
return Policy{}, err
}
return p, tx.Commit(ctx)
}
func (s *Store) Delete(ctx context.Context, id, actorID string) error {
eventID, _ := platformid.NewUUID()
tx, err := s.pool.Begin(ctx)
if err != nil {
return err
}
defer func() { _ = tx.Rollback(ctx) }()
tag, err := tx.Exec(ctx, `DELETE FROM gateway.content_policies WHERE id=$1`, id)
if err != nil {
return err
}
if tag.RowsAffected() == 0 {
return ErrNotFound
}
payload, _ := json.Marshal(map[string]string{"content_policy_id": id, "actor_id": actorID})
if _, err = tx.Exec(ctx, `INSERT INTO gateway.outbox_events(event_id,event_type,event_version,aggregate_type,aggregate_id,payload) VALUES($1,'content_policy.deleted',1,'content_policy',$2,$3)`, eventID, id, payload); err != nil {
return err
}
return tx.Commit(ctx)
}
func Normalize(p *Policy) {
p.Name = strings.TrimSpace(p.Name)
p.Description = strings.TrimSpace(p.Description)
p.Action = strings.ToLower(strings.TrimSpace(p.Action))
p.Paths = normalizeStrings(p.Paths)
p.Models = normalizeStrings(p.Models)
p.APIKeyIDs = normalizeStrings(p.APIKeyIDs)
}
func ValidateInput(p Policy) error {
if p.Name == "" || len(p.Name) > 128 {
return fmt.Errorf("name is required and must not exceed 128 characters")
}
if len(p.Description) > 1000 {
return fmt.Errorf("description is too long")
}
if p.Priority < -100000 || p.Priority > 100000 {
return fmt.Errorf("priority 必须在 -100000 到 100000 之间")
}
if len(p.Paths) > 20 || len(p.Models) > 100 || len(p.APIKeyIDs) > 100 {
return fmt.Errorf("策略范围条目过多")
}
allowedPaths := map[string]bool{"*": true, "/v1/chat/completions": true, "/v1/responses": true, "/v1/embeddings": true, "/v1/messages": true}
for _, value := range p.Paths {
if !allowedPaths[value] {
return fmt.Errorf("不支持的端点范围 %s", value)
}
}
for _, value := range p.Models {
if len(value) > 512 {
return fmt.Errorf("模型名过长")
}
}
for _, value := range p.APIKeyIDs {
var id pgtype.UUID
if id.Scan(value) != nil || !id.Valid {
return fmt.Errorf("API Key ID 格式无效")
}
}
for _, rule := range p.Rules {
if len(rule.Name) > 128 || len(rule.Replacement) > 1024 {
return fmt.Errorf("规则名称或替换文本过长")
}
}
return Validate(p)
}
func normalizeStrings(values []string) []string {
result := make([]string, 0, len(values))
seen := map[string]bool{}
for _, value := range values {
value = strings.TrimSpace(value)
if value != "" && !seen[value] {
seen[value] = true
result = append(result, value)
}
}
return result
}