5759c1862e
M0-M7 已完成:核心网关(身份/RBAC/TOTP/OIDC/SAML/Provider/配额/路由/内容策略/审计/定价)+ 资源市场(MCP/Skills/数字员工)。 含 22 个 PostgreSQL 迁移、管理端/门户端前端源码、OpenAPI 契约、部署 compose。 Co-Authored-By: Claude <noreply@anthropic.com>
178 lines
6.5 KiB
Go
178 lines
6.5 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"
|
|
)
|
|
|
|
var (
|
|
ErrModelRouteNotFound = errors.New("model route not found")
|
|
ErrModelRouteExists = errors.New("model route already exists")
|
|
)
|
|
|
|
type ModelRoute struct {
|
|
ID string `json:"id"`
|
|
Name string `json:"name"`
|
|
SourceModel string `json:"source_model"`
|
|
TargetModel string `json:"target_model"`
|
|
ProviderID string `json:"provider_id"`
|
|
ProviderCode string `json:"provider_code"`
|
|
Weight int `json:"weight"`
|
|
Priority int `json:"priority"`
|
|
Conditions json.RawMessage `json:"conditions"`
|
|
Enabled bool `json:"enabled"`
|
|
CreatedAt time.Time `json:"created_at"`
|
|
UpdatedAt time.Time `json:"updated_at"`
|
|
}
|
|
|
|
func (r *Repository) ListModelRoutes(ctx context.Context) ([]ModelRoute, error) {
|
|
if r.pool == nil {
|
|
return nil, ErrProviderStore
|
|
}
|
|
rows, err := r.pool.Query(ctx, `
|
|
SELECT mr.id::text, mr.name, mr.source_model, mr.target_model, mr.provider_id::text,
|
|
p.code, mr.weight, mr.priority, mr.conditions, mr.enabled, mr.created_at, mr.updated_at
|
|
FROM gateway.model_routes mr JOIN gateway.providers p ON p.id = mr.provider_id
|
|
ORDER BY mr.source_model, mr.priority DESC, mr.name, p.code`)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("%w: %v", ErrProviderStore, err)
|
|
}
|
|
defer rows.Close()
|
|
routes := make([]ModelRoute, 0)
|
|
for rows.Next() {
|
|
var route ModelRoute
|
|
if err := rows.Scan(&route.ID, &route.Name, &route.SourceModel, &route.TargetModel, &route.ProviderID,
|
|
&route.ProviderCode, &route.Weight, &route.Priority, &route.Conditions, &route.Enabled, &route.CreatedAt, &route.UpdatedAt); err != nil {
|
|
return nil, fmt.Errorf("%w: %v", ErrProviderStore, err)
|
|
}
|
|
routes = append(routes, route)
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return nil, fmt.Errorf("%w: %v", ErrProviderStore, err)
|
|
}
|
|
return routes, nil
|
|
}
|
|
|
|
func (r *Repository) CreateModelRoute(ctx context.Context, route ModelRoute, actorID string) (ModelRoute, error) {
|
|
id, err := platformid.NewUUID()
|
|
if err != nil {
|
|
return ModelRoute{}, err
|
|
}
|
|
route.ID = id
|
|
return r.saveModelRoute(ctx, route, actorID, true)
|
|
}
|
|
|
|
func (r *Repository) UpdateModelRoute(ctx context.Context, route ModelRoute, actorID string) (ModelRoute, error) {
|
|
return r.saveModelRoute(ctx, route, actorID, false)
|
|
}
|
|
|
|
func (r *Repository) saveModelRoute(ctx context.Context, route ModelRoute, actorID string, create bool) (ModelRoute, error) {
|
|
if r.pool == nil {
|
|
return ModelRoute{}, ErrProviderStore
|
|
}
|
|
eventID, err := platformid.NewUUID()
|
|
if err != nil {
|
|
return ModelRoute{}, err
|
|
}
|
|
tx, err := r.pool.Begin(ctx)
|
|
if err != nil {
|
|
return ModelRoute{}, fmt.Errorf("%w: %v", ErrProviderStore, err)
|
|
}
|
|
defer func() { _ = tx.Rollback(ctx) }()
|
|
var providerExists bool
|
|
if err := tx.QueryRow(ctx, `SELECT EXISTS(SELECT 1 FROM gateway.providers WHERE id=$1 AND tenant_id IS NULL)`, route.ProviderID).Scan(&providerExists); err != nil {
|
|
return ModelRoute{}, fmt.Errorf("%w: %v", ErrProviderStore, err)
|
|
}
|
|
if !providerExists {
|
|
return ModelRoute{}, ErrProviderNotFound
|
|
}
|
|
if create {
|
|
err = tx.QueryRow(ctx, `
|
|
INSERT INTO gateway.model_routes
|
|
(id,name,source_model,target_model,provider_id,weight,priority,conditions,enabled,created_by)
|
|
SELECT $1,$2,$3,$4,p.id,$6,$7,$8,$9,$10 FROM gateway.providers p WHERE p.id=$5 AND p.tenant_id IS NULL
|
|
RETURNING created_at,updated_at`, route.ID, route.Name, route.SourceModel, route.TargetModel,
|
|
route.ProviderID, route.Weight, route.Priority, route.Conditions, route.Enabled, actorID,
|
|
).Scan(&route.CreatedAt, &route.UpdatedAt)
|
|
} else {
|
|
err = tx.QueryRow(ctx, `
|
|
UPDATE gateway.model_routes
|
|
SET name=$2,source_model=$3,target_model=$4,provider_id=$5,weight=$6,priority=$7,
|
|
conditions=$8,enabled=$9,updated_at=clock_timestamp()
|
|
WHERE id=$1
|
|
RETURNING created_at,updated_at`, route.ID, route.Name, route.SourceModel, route.TargetModel,
|
|
route.ProviderID, route.Weight, route.Priority, route.Conditions, route.Enabled,
|
|
).Scan(&route.CreatedAt, &route.UpdatedAt)
|
|
}
|
|
if errors.Is(err, pgx.ErrNoRows) {
|
|
if create {
|
|
return ModelRoute{}, ErrProviderNotFound
|
|
}
|
|
return ModelRoute{}, ErrModelRouteNotFound
|
|
}
|
|
if err != nil {
|
|
var pgErr *pgconn.PgError
|
|
if errors.As(err, &pgErr) {
|
|
if pgErr.Code == "23505" {
|
|
return ModelRoute{}, ErrModelRouteExists
|
|
}
|
|
if pgErr.Code == "23503" {
|
|
return ModelRoute{}, ErrProviderNotFound
|
|
}
|
|
}
|
|
return ModelRoute{}, fmt.Errorf("%w: %v", ErrProviderStore, err)
|
|
}
|
|
if err := tx.QueryRow(ctx, `SELECT code FROM gateway.providers WHERE id=$1`, route.ProviderID).Scan(&route.ProviderCode); err != nil {
|
|
return ModelRoute{}, ErrProviderNotFound
|
|
}
|
|
eventType := "model_route.updated"
|
|
if create {
|
|
eventType = "model_route.created"
|
|
}
|
|
payload, _ := json.Marshal(map[string]any{"model_route_id": route.ID, "source_model": route.SourceModel, "target_model": route.TargetModel, "provider_id": route.ProviderID, "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,$2,1,'model_route',$3,$4)`, eventID, eventType, route.ID, payload); err != nil {
|
|
return ModelRoute{}, fmt.Errorf("%w: %v", ErrProviderStore, err)
|
|
}
|
|
if err := tx.Commit(ctx); err != nil {
|
|
return ModelRoute{}, fmt.Errorf("%w: %v", ErrProviderStore, err)
|
|
}
|
|
return route, nil
|
|
}
|
|
|
|
func (r *Repository) DeleteModelRoute(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.model_routes WHERE id=$1`, id)
|
|
if err != nil {
|
|
return fmt.Errorf("%w: %v", ErrProviderStore, err)
|
|
}
|
|
if result.RowsAffected() == 0 {
|
|
return ErrModelRouteNotFound
|
|
}
|
|
payload, _ := json.Marshal(map[string]any{"model_route_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,'model_route.deleted',1,'model_route',$2,$3)`, eventID, id, payload); err != nil {
|
|
return fmt.Errorf("%w: %v", ErrProviderStore, err)
|
|
}
|
|
if err := tx.Commit(ctx); err != nil {
|
|
return fmt.Errorf("%w: %v", ErrProviderStore, err)
|
|
}
|
|
return nil
|
|
}
|