Files
ai-gateway-go/internal/provider/repository.go
T
superidou 9501751792 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 币种维度)
2026-08-13 10:50:51 +08:00

391 lines
14 KiB
Go

package provider
import (
"context"
"encoding/json"
"errors"
"fmt"
"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 (
ErrProviderNotFound = errors.New("provider not found")
ErrProviderExists = errors.New("provider already exists")
ErrMultipleDefaults = errors.New("multiple default providers")
ErrProviderStore = errors.New("provider store unavailable")
ErrProviderUpstream = errors.New("provider upstream unavailable")
)
type Record struct {
ID string
TenantID *string
Code string
Adapter string
BaseURL string
EncryptedCredentials []byte
CredentialKEKVersion int
Capabilities []string
Config json.RawMessage
Enabled bool
Revision int64
}
func (r *Repository) ListModels(ctx context.Context, providerID string) ([]Model, error) {
if r.pool == nil {
return nil, ErrProviderStore
}
rows, err := r.pool.Query(ctx, `
SELECT id::text, provider_id::text, provider_model_id, owned_by, metadata,
enabled, discovered_at, last_seen_at
FROM gateway.provider_models
WHERE provider_id = $1
ORDER BY enabled DESC, provider_model_id`, providerID)
if err != nil {
return nil, fmt.Errorf("%w: %v", ErrProviderStore, err)
}
defer rows.Close()
models := make([]Model, 0)
for rows.Next() {
var model Model
if err := rows.Scan(
&model.ID, &model.ProviderID, &model.ProviderModelID, &model.OwnedBy,
&model.Metadata, &model.Enabled, &model.DiscoveredAt, &model.LastSeenAt,
); err != nil {
return nil, fmt.Errorf("%w: %v", ErrProviderStore, err)
}
models = append(models, model)
}
if err := rows.Err(); err != nil {
return nil, fmt.Errorf("%w: %v", ErrProviderStore, err)
}
return models, nil
}
func (r *Repository) SyncModels(ctx context.Context, providerID, actorID string, models []DiscoveredModel) (ModelSyncResult, error) {
if r.pool == nil {
return ModelSyncResult{}, ErrProviderStore
}
eventID, err := platformid.NewUUID()
if err != nil {
return ModelSyncResult{}, err
}
transaction, err := r.pool.Begin(ctx)
if err != nil {
return ModelSyncResult{}, fmt.Errorf("%w: %v", ErrProviderStore, err)
}
defer func() { _ = transaction.Rollback(ctx) }()
var exists bool
if err := transaction.QueryRow(ctx, `
SELECT EXISTS (
SELECT 1 FROM gateway.providers WHERE id = $1 AND tenant_id IS NULL
)`, providerID).Scan(&exists); err != nil {
return ModelSyncResult{}, fmt.Errorf("%w: %v", ErrProviderStore, err)
}
if !exists {
return ModelSyncResult{}, ErrProviderNotFound
}
if _, err := transaction.Exec(ctx, `
UPDATE gateway.provider_models
SET enabled = false, updated_at = clock_timestamp()
WHERE provider_id = $1 AND enabled = true`, providerID); err != nil {
return ModelSyncResult{}, fmt.Errorf("%w: %v", ErrProviderStore, err)
}
for _, model := range models {
id, err := platformid.NewUUID()
if err != nil {
return ModelSyncResult{}, err
}
if _, err := transaction.Exec(ctx, `
INSERT INTO gateway.provider_models
(id, provider_id, provider_model_id, owned_by, metadata, enabled)
VALUES ($1, $2, $3, $4, $5, true)
ON CONFLICT (provider_id, provider_model_id) DO UPDATE
SET owned_by = EXCLUDED.owned_by,
metadata = EXCLUDED.metadata,
enabled = true,
last_seen_at = clock_timestamp(),
updated_at = clock_timestamp()`,
id, providerID, model.ProviderModelID, model.OwnedBy, model.Metadata,
); err != nil {
return ModelSyncResult{}, fmt.Errorf("%w: %v", ErrProviderStore, err)
}
}
result := ModelSyncResult{Discovered: len(models), SyncedAt: time.Now().UTC()}
if err := transaction.QueryRow(ctx, `
SELECT count(*) FILTER (WHERE enabled), count(*) FILTER (WHERE NOT enabled)
FROM gateway.provider_models
WHERE provider_id = $1`, providerID).Scan(&result.Active, &result.Disabled); err != nil {
return ModelSyncResult{}, fmt.Errorf("%w: %v", ErrProviderStore, err)
}
payload, _ := json.Marshal(map[string]any{
"provider_id": providerID, "discovered": result.Discovered,
"active": result.Active, "disabled": result.Disabled, "actor_id": actorID,
})
if _, err := transaction.Exec(ctx, `
INSERT INTO gateway.outbox_events
(event_id, event_type, event_version, aggregate_type, aggregate_id, payload)
VALUES ($1, 'provider.models_synced', 1, 'provider', $2, $3)`, eventID, providerID, payload); err != nil {
return ModelSyncResult{}, fmt.Errorf("%w: %v", ErrProviderStore, err)
}
if err := transaction.Commit(ctx); err != nil {
return ModelSyncResult{}, fmt.Errorf("%w: %v", ErrProviderStore, err)
}
return result, nil
}
func (r *Repository) RotateCredentials(ctx context.Context, actorID string, rotations []CredentialRotation) error {
if r.pool == nil {
return ErrProviderStore
}
if len(rotations) == 0 {
return nil
}
transaction, err := r.pool.Begin(ctx)
if err != nil {
return fmt.Errorf("%w: %v", ErrProviderStore, err)
}
defer func() { _ = transaction.Rollback(ctx) }()
for _, rotation := range rotations {
result, err := transaction.Exec(ctx, `
UPDATE gateway.providers
SET encrypted_credentials = $2,
credential_kek_version = $3,
revision = revision + 1,
updated_at = clock_timestamp()
WHERE id = $1
AND tenant_id IS NULL
AND credential_kek_version = $4`,
rotation.ProviderID, rotation.Ciphertext, rotation.ToVersion, rotation.FromVersion,
)
if err != nil {
return fmt.Errorf("%w: %v", ErrProviderStore, err)
}
if result.RowsAffected() != 1 {
return fmt.Errorf("%w: provider %s changed during credential rotation", ErrProviderStore, rotation.ProviderID)
}
eventID, err := platformid.NewUUID()
if err != nil {
return err
}
payload, _ := json.Marshal(map[string]any{
"provider_id": rotation.ProviderID, "from_version": rotation.FromVersion,
"to_version": rotation.ToVersion, "actor_id": actorID,
})
if _, err := transaction.Exec(ctx, `
INSERT INTO gateway.outbox_events
(event_id, event_type, event_version, aggregate_type, aggregate_id, payload)
VALUES ($1, 'provider.credentials_rotated', 1, 'provider', $2, $3)`,
eventID, rotation.ProviderID, payload,
); err != nil {
return fmt.Errorf("%w: %v", ErrProviderStore, err)
}
}
if err := transaction.Commit(ctx); err != nil {
return fmt.Errorf("%w: %v", ErrProviderStore, err)
}
return nil
}
type Repository struct {
pool *pgxpool.Pool
}
func NewRepository(pool *pgxpool.Pool) *Repository {
return &Repository{pool: pool}
}
func (r *Repository) List(ctx context.Context) ([]Record, error) {
if r.pool == nil {
return nil, ErrProviderStore
}
rows, err := r.pool.Query(ctx, `
SELECT id::text, tenant_id::text, code, adapter, base_url,
encrypted_credentials, credential_kek_version, capabilities,
config, enabled, revision
FROM gateway.providers
WHERE tenant_id IS NULL
ORDER BY code`)
if err != nil {
return nil, fmt.Errorf("%w: %v", ErrProviderStore, err)
}
defer rows.Close()
var records []Record
for rows.Next() {
var record Record
if err := rows.Scan(
&record.ID, &record.TenantID, &record.Code, &record.Adapter, &record.BaseURL,
&record.EncryptedCredentials, &record.CredentialKEKVersion, &record.Capabilities,
&record.Config, &record.Enabled, &record.Revision,
); err != nil {
return nil, fmt.Errorf("%w: %v", ErrProviderStore, err)
}
records = append(records, record)
}
if err := rows.Err(); err != nil {
return nil, fmt.Errorf("%w: %v", ErrProviderStore, err)
}
return records, nil
}
func (r *Repository) Get(ctx context.Context, id string) (Record, error) {
if r.pool == nil {
return Record{}, ErrProviderStore
}
var record Record
err := r.pool.QueryRow(ctx, `
SELECT id::text, tenant_id::text, code, adapter, base_url,
encrypted_credentials, credential_kek_version, capabilities,
config, enabled, revision
FROM gateway.providers
WHERE id = $1 AND tenant_id IS NULL`, id).Scan(
&record.ID, &record.TenantID, &record.Code, &record.Adapter, &record.BaseURL,
&record.EncryptedCredentials, &record.CredentialKEKVersion, &record.Capabilities,
&record.Config, &record.Enabled, &record.Revision,
)
return record, mapProviderError(err)
}
// Delete removes a provider and its cascaded model routes / synced models
// (FKs are ON DELETE CASCADE), emitting a provider.deleted outbox event in the
// same transaction.
func (r *Repository) Delete(ctx context.Context, id, actorID string) error {
if r.pool == nil {
return ErrProviderStore
}
eventID, err := platformid.NewUUID()
if err != nil {
return err
}
tx, err := r.pool.Begin(ctx)
if err != nil {
return fmt.Errorf("%w: %v", ErrProviderStore, err)
}
defer func() { _ = tx.Rollback(ctx) }()
result, err := tx.Exec(ctx, `DELETE FROM gateway.providers WHERE id=$1 AND tenant_id IS NULL`, id)
if err != nil {
return fmt.Errorf("%w: %v", ErrProviderStore, err)
}
if result.RowsAffected() == 0 {
return ErrProviderNotFound
}
payload, _ := json.Marshal(map[string]any{"provider_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, 'provider.deleted', 1, 'provider', $2, $3)`, eventID, id, payload); err != nil {
return fmt.Errorf("%w: %v", ErrProviderStore, err)
}
return tx.Commit(ctx)
}
func (r *Repository) Create(ctx context.Context, record Record, actorID string) (Record, error) {
if r.pool == nil {
return Record{}, ErrProviderStore
}
id, err := platformid.NewUUID()
if err != nil {
return Record{}, err
}
eventID, err := platformid.NewUUID()
if err != nil {
return Record{}, err
}
record.ID = id
transaction, err := r.pool.Begin(ctx)
if err != nil {
return Record{}, fmt.Errorf("%w: %v", ErrProviderStore, err)
}
defer func() { _ = transaction.Rollback(ctx) }()
err = transaction.QueryRow(ctx, `
INSERT INTO gateway.providers
(id, code, adapter, base_url, encrypted_credentials,
credential_kek_version, capabilities, config, enabled, created_by)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10)
RETURNING revision`, record.ID, record.Code, record.Adapter, record.BaseURL,
record.EncryptedCredentials, record.CredentialKEKVersion, record.Capabilities,
record.Config, record.Enabled, actorID).Scan(&record.Revision)
if err != nil {
return Record{}, mapProviderError(err)
}
payload, _ := json.Marshal(map[string]any{"provider_id": record.ID, "code": record.Code, "revision": record.Revision})
if _, err := transaction.Exec(ctx, `
INSERT INTO gateway.outbox_events
(event_id, event_type, event_version, aggregate_type, aggregate_id, payload)
VALUES ($1, 'provider.created', 1, 'provider', $2, $3)`, eventID, record.ID, payload); err != nil {
return Record{}, fmt.Errorf("%w: %v", ErrProviderStore, err)
}
if err := transaction.Commit(ctx); err != nil {
return Record{}, fmt.Errorf("%w: %v", ErrProviderStore, err)
}
return record, nil
}
func (r *Repository) Update(ctx context.Context, record Record, actorID string, replaceCredentials bool) (Record, error) {
if r.pool == nil {
return Record{}, ErrProviderStore
}
eventID, err := platformid.NewUUID()
if err != nil {
return Record{}, err
}
transaction, err := r.pool.Begin(ctx)
if err != nil {
return Record{}, fmt.Errorf("%w: %v", ErrProviderStore, err)
}
defer func() { _ = transaction.Rollback(ctx) }()
err = transaction.QueryRow(ctx, `
UPDATE gateway.providers
SET code = $2,
adapter = $3,
base_url = $4,
encrypted_credentials = CASE WHEN $10 THEN $5 ELSE encrypted_credentials END,
credential_kek_version = CASE WHEN $10 THEN $6 ELSE credential_kek_version END,
capabilities = $7,
config = $8,
enabled = $9,
revision = revision + 1,
updated_at = clock_timestamp()
WHERE id = $1 AND tenant_id IS NULL
RETURNING revision`, record.ID, record.Code, record.Adapter, record.BaseURL,
record.EncryptedCredentials, record.CredentialKEKVersion, record.Capabilities,
record.Config, record.Enabled, replaceCredentials).Scan(&record.Revision)
if err != nil {
return Record{}, mapProviderError(err)
}
payload, _ := json.Marshal(map[string]any{
"provider_id": record.ID, "code": record.Code, "revision": record.Revision, "actor_id": actorID,
})
if _, err := transaction.Exec(ctx, `
INSERT INTO gateway.outbox_events
(event_id, event_type, event_version, aggregate_type, aggregate_id, payload)
VALUES ($1, 'provider.updated', 1, 'provider', $2, $3)`, eventID, record.ID, payload); err != nil {
return Record{}, fmt.Errorf("%w: %v", ErrProviderStore, err)
}
if err := transaction.Commit(ctx); err != nil {
return Record{}, fmt.Errorf("%w: %v", ErrProviderStore, err)
}
return record, nil
}
func mapProviderError(err error) error {
if errors.Is(err, pgx.ErrNoRows) {
return ErrProviderNotFound
}
var pgError *pgconn.PgError
if errors.As(err, &pgError) && pgError.Code == "23505" {
if pgError.ConstraintName == "providers_single_global_default_idx" {
return ErrMultipleDefaults
}
return ErrProviderExists
}
if err != nil {
return fmt.Errorf("%w: %v", ErrProviderStore, err)
}
return nil
}