Files
chorus/admin/app/chorus/service.go
T

1393 lines
58 KiB
Go

package chorus
import (
"context"
"encoding/json"
"errors"
"fmt"
"math"
"net/url"
"slices"
"strings"
"time"
"git.ilapage.cn/OPC/chorus/internal/core/model"
"git.ilapage.cn/OPC/chorus/internal/core/provider"
"github.com/google/uuid"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
const maxRouteMembers = 32
type ProbeConfiguration struct {
BaseURL string
AuthType provider.AuthType
APIKey string
ModelID string
APIType model.APIType
Kind model.GenerationKind
Capabilities []model.Capability
ExtraBody json.RawMessage
Timeout time.Duration
MaxResponseBytes int64
}
// ProbeRunner is deliberately injected. The production implementation is
// created with the same SSRF-safe HTTP client as workers; tests provide a
// mock and the default server setting refuses to call it.
type ProbeRunner interface {
Probe(context.Context, ProbeConfiguration) error
}
type Config struct {
AllowConnectivityChecks bool
ConnectivityCooldown time.Duration
MaxResponseBytes int64
Now func() time.Time
Probe ProbeRunner
}
type Service struct {
db *gorm.DB
allowConnectivityChecks bool
connectivityCooldown time.Duration
maxResponseBytes int64
now func() time.Time
probe ProbeRunner
}
func NewService(db *gorm.DB, cfg Config) (*Service, error) {
if db == nil {
return nil, errors.New("chorus admin service dependencies are required")
}
if cfg.Now == nil {
cfg.Now = time.Now
}
if cfg.AllowConnectivityChecks && (cfg.ConnectivityCooldown <= 0 || cfg.Probe == nil || cfg.MaxResponseBytes <= 0) {
return nil, errors.New("authorized connectivity checks require cooldown, probe, and response limit")
}
return &Service{
db: db, allowConnectivityChecks: cfg.AllowConnectivityChecks,
connectivityCooldown: cfg.ConnectivityCooldown, maxResponseBytes: cfg.MaxResponseBytes,
now: cfg.Now, probe: cfg.Probe,
}, nil
}
type providerRow struct {
ID uint64 `gorm:"column:id;primaryKey"`
Slug string `gorm:"column:slug"`
Name string `gorm:"column:name"`
BaseURL string `gorm:"column:base_url"`
AuthType string `gorm:"column:auth_type"`
ActiveCredentialID *uint64 `gorm:"column:active_credential_id"`
Enabled bool `gorm:"column:enabled"`
CreatedAt time.Time `gorm:"column:created_at"`
UpdatedAt time.Time `gorm:"column:updated_at"`
}
func (providerRow) TableName() string { return "providers" }
type credentialRow struct {
ID uint64 `gorm:"column:id;primaryKey"`
ProviderID uint64 `gorm:"column:provider_id"`
CredentialVersion uint32 `gorm:"column:credential_version"`
APIKey string `gorm:"column:api_key"`
Status string `gorm:"column:status"`
CreatedBy uint64 `gorm:"column:created_by"`
UpdatedBy uint64 `gorm:"column:updated_by"`
CreatedAt time.Time `gorm:"column:created_at"`
UpdatedAt time.Time `gorm:"column:updated_at"`
RetiredAt *time.Time `gorm:"column:retired_at"`
}
func (credentialRow) TableName() string { return "provider_credentials" }
type providerModelRow struct {
ID uint64 `gorm:"column:id;primaryKey"`
ProviderID uint64 `gorm:"column:provider_id"`
Name string `gorm:"column:name"`
ModelID string `gorm:"column:model_id"`
APIType model.APIType `gorm:"column:api_type"`
Kind model.GenerationKind `gorm:"column:kind"`
ExtraBody json.RawMessage `gorm:"column:extra_body"`
TimeoutMS uint32 `gorm:"column:timeout_ms"`
Weight uint32 `gorm:"column:weight"`
Enabled bool `gorm:"column:enabled"`
CreatedAt time.Time `gorm:"column:created_at"`
UpdatedAt time.Time `gorm:"column:updated_at"`
}
func (providerModelRow) TableName() string { return "provider_models" }
type modelCapabilityRow struct {
ProviderModelID uint64 `gorm:"column:provider_model_id;primaryKey"`
Capability model.Capability `gorm:"column:capability;primaryKey"`
}
func (modelCapabilityRow) TableName() string { return "provider_model_capabilities" }
type promptTemplateRow struct {
ID uint64 `gorm:"column:id;primaryKey"`
TemplateKey string `gorm:"column:template_key"`
Kind model.GenerationKind `gorm:"column:kind"`
APIType model.APIType `gorm:"column:api_type"`
Capability model.Capability `gorm:"column:capability"`
Name string `gorm:"column:name"`
Version uint32 `gorm:"column:version"`
TemplateText string `gorm:"column:template_text"`
DefaultRoleRule string `gorm:"column:default_role_rule"`
Enabled bool `gorm:"column:enabled"`
CreatedAt time.Time `gorm:"column:created_at"`
UpdatedAt time.Time `gorm:"column:updated_at"`
}
func (promptTemplateRow) TableName() string { return "prompt_templates" }
type routePoolRow struct {
ID uint64 `gorm:"column:id;primaryKey"`
Slug string `gorm:"column:slug"`
Name string `gorm:"column:name"`
Capability model.Capability `gorm:"column:capability"`
PromptTemplateID uint64 `gorm:"column:prompt_template_id"`
MaxFailover uint32 `gorm:"column:max_failover"`
Version uint32 `gorm:"column:version"`
Enabled bool `gorm:"column:enabled"`
CreatedAt time.Time `gorm:"column:created_at"`
UpdatedAt time.Time `gorm:"column:updated_at"`
}
func (routePoolRow) TableName() string { return "route_pools" }
type routeMemberRow struct {
ID uint64 `gorm:"column:id;primaryKey"`
RoutePoolID uint64 `gorm:"column:route_pool_id"`
ProviderModelID uint64 `gorm:"column:provider_model_id"`
Weight uint32 `gorm:"column:weight"`
FailureThreshold uint32 `gorm:"column:failure_threshold"`
OpenSeconds uint32 `gorm:"column:open_seconds"`
HalfOpenMax uint32 `gorm:"column:half_open_max"`
Enabled bool `gorm:"column:enabled"`
Position uint32 `gorm:"column:position"`
CreatedAt time.Time `gorm:"column:created_at"`
UpdatedAt time.Time `gorm:"column:updated_at"`
}
func (routeMemberRow) TableName() string { return "route_pool_members" }
type activeRouteRow struct {
Capability model.Capability `gorm:"column:capability;primaryKey"`
RoutePoolID uint64 `gorm:"column:route_pool_id"`
OperatorID uint64 `gorm:"column:operator_id"`
CreatedAt time.Time `gorm:"column:created_at"`
UpdatedAt time.Time `gorm:"column:updated_at"`
}
func (activeRouteRow) TableName() string { return "active_routes" }
type connectivityRow struct {
ID uint64 `gorm:"column:id;primaryKey"`
ProviderModelID uint64 `gorm:"column:provider_model_id"`
OperatorID uint64 `gorm:"column:operator_id"`
Status string `gorm:"column:status"`
ErrorCode *string `gorm:"column:error_code"`
ErrorMessage *string `gorm:"column:error_message"`
LatencyMS *uint32 `gorm:"column:latency_ms"`
StartedAt time.Time `gorm:"column:started_at"`
CompletedAt *time.Time `gorm:"column:completed_at"`
CooldownUntil time.Time `gorm:"column:cooldown_until"`
CreatedAt time.Time `gorm:"column:created_at"`
}
func (connectivityRow) TableName() string { return "provider_connectivity_checks" }
type auditRow struct {
ID uint64 `gorm:"column:id;primaryKey"`
Operator uint64 `gorm:"column:operator_id"`
Action string `gorm:"column:action"`
Target string `gorm:"column:target_type"`
TargetID string `gorm:"column:target_id"`
Result string `gorm:"column:result"`
RequestID string `gorm:"column:request_id"`
Summary json.RawMessage `gorm:"column:summary"`
CreatedAt time.Time `gorm:"column:created_at"`
}
func (auditRow) TableName() string { return "admin_audit_events" }
type portalUserRow struct {
ID uint64 `gorm:"column:id;primaryKey"`
Email string `gorm:"column:email"`
DisplayName string `gorm:"column:display_name"`
Status string `gorm:"column:status"`
CreatedAt time.Time `gorm:"column:created_at"`
UpdatedAt time.Time `gorm:"column:updated_at"`
}
func (portalUserRow) TableName() string { return "users" }
type generationRow struct {
ID uint64 `gorm:"column:id;primaryKey"`
UserID uint64 `gorm:"column:user_id"`
Status model.GenerationStatus `gorm:"column:status"`
Kind model.GenerationKind `gorm:"column:kind"`
RenderedPrompt string `gorm:"column:rendered_prompt"`
Attempts json.RawMessage `gorm:"column:attempts"`
ErrorCode *string `gorm:"column:error_code"`
ErrorMessage *string `gorm:"column:error_message"`
ProviderAttemptCount uint32 `gorm:"column:provider_attempt_count"`
CreatedAt time.Time `gorm:"column:created_at"`
CompletedAt *time.Time `gorm:"column:completed_at"`
}
func (generationRow) TableName() string { return "generations" }
func (s *Service) CreateProvider(ctx context.Context, actor uint64, requestID string, input ProviderInput) (ProviderView, error) {
if err := validateProvider(input, true); err != nil {
return ProviderView{}, err
}
var id uint64
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
// MySQL validates the provider's active-credential foreign-key cycle per
// statement. Start credential-backed providers as auth_type=none, then
// switch both columns together after the credential exists.
initialAuthType := input.AuthType
if input.APIKey != nil {
initialAuthType = string(provider.AuthNone)
}
row := providerRow{Slug: strings.TrimSpace(input.Slug), Name: strings.TrimSpace(input.Name), BaseURL: strings.TrimSpace(input.BaseURL), AuthType: initialAuthType, Enabled: input.Enabled}
if err := tx.Create(&row).Error; err != nil {
return fmt.Errorf("create provider: %w", err)
}
if input.APIKey != nil {
credential, err := s.newCredential(row.ID, 1, actor, *input.APIKey)
if err != nil {
return err
}
if err := tx.Create(&credential).Error; err != nil {
return fmt.Errorf("create provider credential: %w", err)
}
if err := tx.Model(&providerRow{}).Where("id = ?", row.ID).Updates(map[string]any{"auth_type": input.AuthType, "active_credential_id": credential.ID}).Error; err != nil {
return fmt.Errorf("activate provider credential: %w", err)
}
row.AuthType = input.AuthType
}
if row.Enabled && input.AuthType != string(provider.AuthNone) && input.APIKey == nil {
return FieldError{Field: "api_key", Message: "is required before enabling this provider"}
}
if err := s.audit(tx, actor, requestID, "provider.create", "provider", row.ID, "succeeded", map[string]any{"enabled": row.Enabled, "auth_type": row.AuthType}); err != nil {
return err
}
id = row.ID
return nil
})
if err != nil {
return ProviderView{}, err
}
return s.Provider(ctx, id)
}
func (s *Service) UpdateProvider(ctx context.Context, actor uint64, requestID string, id uint64, input ProviderInput) (ProviderView, error) {
if id == 0 {
return ProviderView{}, FieldError{Field: "id", Message: "must be positive"}
}
if input.APIKey != nil {
return ProviderView{}, FieldError{Field: "api_key", Message: "use the credential rotation endpoint"}
}
if err := validateProvider(input, false); err != nil {
return ProviderView{}, err
}
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
var row providerRow
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&row, id).Error; err != nil {
return translateNotFound(err)
}
if input.AuthType == string(provider.AuthNone) && row.ActiveCredentialID != nil {
return FieldError{Field: "auth_type", Message: "cannot be none while a credential is active"}
}
if input.Enabled && input.AuthType != string(provider.AuthNone) && row.ActiveCredentialID == nil {
return FieldError{Field: "enabled", Message: "requires an active credential"}
}
updates := map[string]any{"slug": strings.TrimSpace(input.Slug), "name": strings.TrimSpace(input.Name), "base_url": strings.TrimSpace(input.BaseURL), "auth_type": input.AuthType, "enabled": input.Enabled}
if err := tx.Model(&providerRow{}).Where("id = ?", id).Updates(updates).Error; err != nil {
return fmt.Errorf("update provider: %w", err)
}
return s.audit(tx, actor, requestID, "provider.update", "provider", id, "succeeded", map[string]any{"enabled": input.Enabled, "auth_type": input.AuthType})
})
if err != nil {
return ProviderView{}, err
}
return s.Provider(ctx, id)
}
func (s *Service) RotateCredential(ctx context.Context, actor uint64, requestID string, providerID uint64, apiKey string) (ProviderView, error) {
if providerID == 0 {
return ProviderView{}, FieldError{Field: "provider_id", Message: "must be positive"}
}
if strings.TrimSpace(apiKey) == "" {
return ProviderView{}, FieldError{Field: "api_key", Message: "must not be empty"}
}
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
var row providerRow
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&row, providerID).Error; err != nil {
return translateNotFound(err)
}
if row.AuthType == string(provider.AuthNone) {
return FieldError{Field: "auth_type", Message: "does not accept credentials"}
}
var latest uint32
if err := tx.Model(&credentialRow{}).Where("provider_id = ?", row.ID).Select("COALESCE(MAX(credential_version), 0)").Scan(&latest).Error; err != nil {
return fmt.Errorf("read credential version: %w", err)
}
credential, err := s.newCredential(row.ID, latest+1, actor, apiKey)
if err != nil {
return err
}
now := s.nowUTC()
if err := tx.Model(&credentialRow{}).Where("provider_id = ? AND status = 'active'", row.ID).Updates(map[string]any{"status": "retired", "retired_at": now, "updated_by": actor}).Error; err != nil {
return fmt.Errorf("retire provider credential: %w", err)
}
if err := tx.Create(&credential).Error; err != nil {
return fmt.Errorf("create rotated credential: %w", err)
}
if err := tx.Model(&providerRow{}).Where("id = ?", row.ID).Update("active_credential_id", credential.ID).Error; err != nil {
return fmt.Errorf("activate rotated credential: %w", err)
}
return s.audit(tx, actor, requestID, "provider.credential.rotate", "provider", row.ID, "succeeded", map[string]any{"credential_version": credential.CredentialVersion})
})
if err != nil {
return ProviderView{}, err
}
return s.Provider(ctx, providerID)
}
func (s *Service) ActivateCredential(ctx context.Context, actor uint64, requestID string, providerID uint64, version uint32) (ProviderView, error) {
if providerID == 0 || version == 0 {
return ProviderView{}, FieldError{Field: "credential_version", Message: "must be positive"}
}
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
var providerRecord providerRow
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&providerRecord, providerID).Error; err != nil {
return translateNotFound(err)
}
if providerRecord.AuthType == string(provider.AuthNone) {
return FieldError{Field: "auth_type", Message: "does not accept credentials"}
}
var selected credentialRow
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Where("provider_id = ? AND credential_version = ?", providerID, version).First(&selected).Error; err != nil {
return translateNotFound(err)
}
now := s.nowUTC()
if err := tx.Model(&credentialRow{}).Where("provider_id = ? AND status = 'active'", providerID).Updates(map[string]any{"status": "retired", "retired_at": now, "updated_by": actor}).Error; err != nil {
return fmt.Errorf("retire active credential: %w", err)
}
if err := tx.Model(&credentialRow{}).Where("id = ?", selected.ID).Updates(map[string]any{"status": "active", "retired_at": nil, "updated_by": actor}).Error; err != nil {
return fmt.Errorf("reactivate provider credential: %w", err)
}
if err := tx.Model(&providerRow{}).Where("id = ?", providerID).Update("active_credential_id", selected.ID).Error; err != nil {
return fmt.Errorf("set active credential: %w", err)
}
return s.audit(tx, actor, requestID, "provider.credential.activate", "provider", providerID, "succeeded", map[string]any{"credential_version": version})
})
if err != nil {
return ProviderView{}, err
}
return s.Provider(ctx, providerID)
}
func (s *Service) Providers(ctx context.Context) ([]ProviderView, error) {
var rows []providerRow
if err := s.db.WithContext(ctx).Order("id ASC").Find(&rows).Error; err != nil {
return nil, fmt.Errorf("list providers: %w", err)
}
return s.providerViews(ctx, rows)
}
func (s *Service) Provider(ctx context.Context, id uint64) (ProviderView, error) {
if id == 0 {
return ProviderView{}, ErrNotFound
}
var row providerRow
if err := s.db.WithContext(ctx).First(&row, id).Error; err != nil {
return ProviderView{}, translateNotFound(err)
}
items, err := s.providerViews(ctx, []providerRow{row})
if err != nil {
return ProviderView{}, err
}
return items[0], nil
}
func (s *Service) ProviderCredential(ctx context.Context, actor uint64, requestID string, providerID uint64) (ProviderCredentialView, error) {
if providerID == 0 {
return ProviderCredentialView{}, ErrNotFound
}
var view ProviderCredentialView
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
var providerRecord providerRow
if err := tx.First(&providerRecord, providerID).Error; err != nil {
return translateNotFound(err)
}
if providerRecord.ActiveCredentialID == nil {
return ErrNotFound
}
var credential credentialRow
if err := tx.Where("id = ? AND provider_id = ? AND status = 'active'", *providerRecord.ActiveCredentialID, providerID).First(&credential).Error; err != nil {
return translateNotFound(err)
}
view = ProviderCredentialView{APIKey: credential.APIKey, Version: credential.CredentialVersion}
return s.audit(tx, actor, requestID, "provider.credential.read", "provider", providerID, "succeeded", map[string]any{"credential_version": credential.CredentialVersion})
})
if err != nil {
return ProviderCredentialView{}, err
}
return view, nil
}
func (s *Service) CreateProviderModel(ctx context.Context, actor uint64, requestID string, input ProviderModelInput) (ProviderModelView, error) {
if err := validateProviderModel(input); err != nil {
return ProviderModelView{}, err
}
var id uint64
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if err := ensureProvider(tx, input.ProviderID); err != nil {
return err
}
row := providerModelRow{ProviderID: input.ProviderID, Name: strings.TrimSpace(input.Name), ModelID: strings.TrimSpace(input.ModelID), APIType: input.APIType, Kind: input.Kind, ExtraBody: normalizedExtraBody(input.ExtraBody), TimeoutMS: input.TimeoutMS, Weight: input.Weight, Enabled: input.Enabled}
if err := tx.Create(&row).Error; err != nil {
return fmt.Errorf("create provider model: %w", err)
}
if err := replaceCapabilities(tx, row.ID, input.Capabilities); err != nil {
return err
}
if err := s.audit(tx, actor, requestID, "provider_model.create", "provider_model", row.ID, "succeeded", map[string]any{"enabled": row.Enabled, "api_type": row.APIType}); err != nil {
return err
}
id = row.ID
return nil
})
if err != nil {
return ProviderModelView{}, err
}
return s.ProviderModel(ctx, id)
}
func (s *Service) UpdateProviderModel(ctx context.Context, actor uint64, requestID string, id uint64, input ProviderModelInput) (ProviderModelView, error) {
if id == 0 {
return ProviderModelView{}, FieldError{Field: "id", Message: "must be positive"}
}
if err := validateProviderModel(input); err != nil {
return ProviderModelView{}, err
}
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
var existing providerModelRow
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&existing, id).Error; err != nil {
return translateNotFound(err)
}
if err := ensureProvider(tx, input.ProviderID); err != nil {
return err
}
updates := map[string]any{"provider_id": input.ProviderID, "name": strings.TrimSpace(input.Name), "model_id": strings.TrimSpace(input.ModelID), "api_type": input.APIType, "kind": input.Kind, "extra_body": normalizedExtraBody(input.ExtraBody), "timeout_ms": input.TimeoutMS, "weight": input.Weight, "enabled": input.Enabled}
if err := tx.Model(&providerModelRow{}).Where("id = ?", id).Updates(updates).Error; err != nil {
return fmt.Errorf("update provider model: %w", err)
}
if err := replaceCapabilities(tx, id, input.Capabilities); err != nil {
return err
}
return s.audit(tx, actor, requestID, "provider_model.update", "provider_model", id, "succeeded", map[string]any{"enabled": input.Enabled, "api_type": input.APIType})
})
if err != nil {
return ProviderModelView{}, err
}
return s.ProviderModel(ctx, id)
}
func (s *Service) ProviderModels(ctx context.Context) ([]ProviderModelView, error) {
var rows []providerModelRow
if err := s.db.WithContext(ctx).Order("id ASC").Find(&rows).Error; err != nil {
return nil, fmt.Errorf("list provider models: %w", err)
}
return s.providerModelViews(ctx, rows)
}
func (s *Service) ProviderModel(ctx context.Context, id uint64) (ProviderModelView, error) {
var row providerModelRow
if err := s.db.WithContext(ctx).First(&row, id).Error; err != nil {
return ProviderModelView{}, translateNotFound(err)
}
items, err := s.providerModelViews(ctx, []providerModelRow{row})
if err != nil {
return ProviderModelView{}, err
}
return items[0], nil
}
func (s *Service) CreatePromptTemplate(ctx context.Context, actor uint64, requestID string, input PromptTemplateInput) (PromptTemplateView, error) {
if err := validatePromptTemplate(input); err != nil {
return PromptTemplateView{}, err
}
var id uint64
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
row := promptTemplateRow{TemplateKey: strings.TrimSpace(input.TemplateKey), Kind: input.Kind, APIType: input.APIType, Capability: input.Capability, Name: strings.TrimSpace(input.Name), Version: 1, TemplateText: input.TemplateText, DefaultRoleRule: input.DefaultRoleRule, Enabled: input.Enabled}
if err := tx.Create(&row).Error; err != nil {
return fmt.Errorf("create prompt template: %w", err)
}
if err := s.audit(tx, actor, requestID, "prompt_template.create", "prompt_template", row.ID, "succeeded", map[string]any{"enabled": row.Enabled, "capability": row.Capability, "version": row.Version}); err != nil {
return err
}
id = row.ID
return nil
})
if err != nil {
return PromptTemplateView{}, err
}
return s.PromptTemplate(ctx, id)
}
func (s *Service) UpdatePromptTemplate(ctx context.Context, actor uint64, requestID string, id uint64, input PromptTemplateInput) (PromptTemplateView, error) {
if id == 0 {
return PromptTemplateView{}, FieldError{Field: "id", Message: "must be positive"}
}
if err := validatePromptTemplate(input); err != nil {
return PromptTemplateView{}, err
}
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
var existing promptTemplateRow
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&existing, id).Error; err != nil {
return translateNotFound(err)
}
updates := map[string]any{"template_key": strings.TrimSpace(input.TemplateKey), "kind": input.Kind, "api_type": input.APIType, "capability": input.Capability, "name": strings.TrimSpace(input.Name), "template_text": input.TemplateText, "default_role_rule": input.DefaultRoleRule, "enabled": input.Enabled, "version": existing.Version + 1}
if err := tx.Model(&promptTemplateRow{}).Where("id = ?", id).Updates(updates).Error; err != nil {
return fmt.Errorf("update prompt template: %w", err)
}
return s.audit(tx, actor, requestID, "prompt_template.update", "prompt_template", id, "succeeded", map[string]any{"enabled": input.Enabled, "capability": input.Capability, "version": existing.Version + 1})
})
if err != nil {
return PromptTemplateView{}, err
}
return s.PromptTemplate(ctx, id)
}
func (s *Service) PromptTemplates(ctx context.Context) ([]PromptTemplateView, error) {
var rows []promptTemplateRow
if err := s.db.WithContext(ctx).Order("id ASC").Find(&rows).Error; err != nil {
return nil, fmt.Errorf("list prompt templates: %w", err)
}
return promptViews(rows), nil
}
func (s *Service) PromptTemplate(ctx context.Context, id uint64) (PromptTemplateView, error) {
var row promptTemplateRow
if err := s.db.WithContext(ctx).First(&row, id).Error; err != nil {
return PromptTemplateView{}, translateNotFound(err)
}
return promptViews([]promptTemplateRow{row})[0], nil
}
func (s *Service) CreateRoutePool(ctx context.Context, actor uint64, requestID string, input RoutePoolInput) (RoutePoolView, error) {
if input.Version != 0 {
return RoutePoolView{}, FieldError{Field: "version", Message: "must be omitted when creating a route pool"}
}
if err := validateRoutePoolInput(input, true); err != nil {
return RoutePoolView{}, err
}
var id uint64
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if err := s.validateRoutePoolReferences(tx, input); err != nil {
return err
}
row := routePoolRow{Slug: strings.TrimSpace(input.Slug), Name: strings.TrimSpace(input.Name), Capability: input.Capability, PromptTemplateID: input.PromptTemplateID, MaxFailover: input.MaxFailover, Version: 1, Enabled: input.Enabled}
if err := tx.Create(&row).Error; err != nil {
return fmt.Errorf("create route pool: %w", err)
}
if err := replaceRouteMembers(tx, row.ID, input.Members); err != nil {
return err
}
if input.Publish {
if err := s.publishRoute(tx, actor, row, input); err != nil {
return err
}
}
if err := s.audit(tx, actor, requestID, "route_pool.create", "route_pool", row.ID, "succeeded", map[string]any{"capability": row.Capability, "version": row.Version, "published": input.Publish}); err != nil {
return err
}
id = row.ID
return nil
})
if err != nil {
return RoutePoolView{}, err
}
return s.RoutePool(ctx, id)
}
func (s *Service) UpdateRoutePool(ctx context.Context, actor uint64, requestID string, id uint64, input RoutePoolInput) (RoutePoolView, error) {
if id == 0 || input.Version == 0 {
return RoutePoolView{}, FieldError{Field: "version", Message: "is required for route pool updates"}
}
if err := validateRoutePoolInput(input, true); err != nil {
return RoutePoolView{}, err
}
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
var existing routePoolRow
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&existing, id).Error; err != nil {
return translateNotFound(err)
}
if existing.Version != input.Version {
return ErrConflict
}
if err := s.validateRoutePoolReferences(tx, input); err != nil {
return err
}
updates := map[string]any{"slug": strings.TrimSpace(input.Slug), "name": strings.TrimSpace(input.Name), "capability": input.Capability, "prompt_template_id": input.PromptTemplateID, "max_failover": input.MaxFailover, "enabled": input.Enabled, "version": existing.Version + 1}
if err := tx.Model(&routePoolRow{}).Where("id = ? AND version = ?", id, existing.Version).Updates(updates).Error; err != nil {
return fmt.Errorf("update route pool: %w", err)
}
if err := replaceRouteMembers(tx, id, input.Members); err != nil {
return err
}
updated := existing
updated.Version++
updated.Capability, updated.PromptTemplateID, updated.MaxFailover, updated.Enabled = input.Capability, input.PromptTemplateID, input.MaxFailover, input.Enabled
if input.Publish {
if err := s.publishRoute(tx, actor, updated, input); err != nil {
return err
}
}
return s.audit(tx, actor, requestID, "route_pool.update", "route_pool", id, "succeeded", map[string]any{"capability": input.Capability, "version": updated.Version, "published": input.Publish})
})
if err != nil {
return RoutePoolView{}, err
}
return s.RoutePool(ctx, id)
}
func (s *Service) RoutePools(ctx context.Context) ([]RoutePoolView, error) {
var rows []routePoolRow
if err := s.db.WithContext(ctx).Order("id ASC").Find(&rows).Error; err != nil {
return nil, fmt.Errorf("list route pools: %w", err)
}
items := make([]RoutePoolView, 0, len(rows))
for _, row := range rows {
item, err := s.routePoolView(ctx, row)
if err != nil {
return nil, err
}
items = append(items, item)
}
return items, nil
}
func (s *Service) RoutePool(ctx context.Context, id uint64) (RoutePoolView, error) {
var row routePoolRow
if err := s.db.WithContext(ctx).First(&row, id).Error; err != nil {
return RoutePoolView{}, translateNotFound(err)
}
return s.routePoolView(ctx, row)
}
func (s *Service) Users(ctx context.Context) ([]UserView, error) {
var rows []portalUserRow
if err := s.db.WithContext(ctx).Order("id ASC").Find(&rows).Error; err != nil {
return nil, fmt.Errorf("list users: %w", err)
}
items := make([]UserView, 0, len(rows))
for _, row := range rows {
items = append(items, UserView{ID: row.ID, Email: row.Email, DisplayName: row.DisplayName, Status: row.Status, CreatedAt: row.CreatedAt, UpdatedAt: row.UpdatedAt})
}
return items, nil
}
func (s *Service) UpdateUserStatus(ctx context.Context, actor uint64, requestID string, id uint64, input UserStatusInput) (UserView, error) {
input.Status = strings.TrimSpace(input.Status)
if input.Status != "active" && input.Status != "disabled" {
return UserView{}, FieldError{Field: "status", Message: "must be active or disabled"}
}
var row portalUserRow
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&row, id).Error; err != nil {
return translateNotFound(err)
}
if err := tx.Model(&portalUserRow{}).Where("id = ?", id).Update("status", input.Status).Error; err != nil {
return fmt.Errorf("update user status: %w", err)
}
row.Status = input.Status
return s.audit(tx, actor, requestID, "user.status.update", "user", id, "succeeded", map[string]any{"status": input.Status})
})
if err != nil {
return UserView{}, err
}
if err := s.db.WithContext(ctx).First(&row, id).Error; err != nil {
return UserView{}, translateNotFound(err)
}
return UserView{ID: row.ID, Email: row.Email, DisplayName: row.DisplayName, Status: row.Status, CreatedAt: row.CreatedAt, UpdatedAt: row.UpdatedAt}, nil
}
func (s *Service) Generations(ctx context.Context) ([]GenerationView, error) {
var rows []generationRow
if err := s.db.WithContext(ctx).Order("created_at DESC, id DESC").Limit(100).Find(&rows).Error; err != nil {
return nil, fmt.Errorf("list generations: %w", err)
}
items := make([]GenerationView, 0, len(rows))
for _, row := range rows {
item := GenerationView{ID: row.ID, UserID: row.UserID, Status: row.Status, Kind: row.Kind, RenderedPrompt: truncate(row.RenderedPrompt, 1024), Attempts: append(json.RawMessage(nil), row.Attempts...), ProviderAttemptCnt: row.ProviderAttemptCount, CreatedAt: row.CreatedAt, CompletedAt: row.CompletedAt}
if row.ErrorCode != nil {
item.ErrorCode = *row.ErrorCode
}
if row.ErrorMessage != nil {
item.ErrorMessage = truncate(*row.ErrorMessage, 256)
}
items = append(items, item)
}
return items, nil
}
func (s *Service) StartConnectivityCheck(ctx context.Context, actor uint64, requestID string, providerModelID uint64) (ConnectivityCheckView, error) {
if providerModelID == 0 {
return ConnectivityCheckView{}, FieldError{Field: "provider_model_id", Message: "must be positive"}
}
if !s.allowConnectivityChecks {
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if err := ensureProviderModel(tx, providerModelID); err != nil {
return err
}
return s.audit(tx, actor, requestID, "provider_model.connectivity.blocked", "provider_model", providerModelID, "failed", map[string]any{"reason": "explicit authorization required"})
})
if err != nil {
return ConnectivityCheckView{}, err
}
return ConnectivityCheckView{}, ErrConnectivityDisabled
}
var check connectivityRow
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if err := ensureProviderModelLocked(tx, providerModelID); err != nil {
return err
}
var latest connectivityRow
result := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Where("provider_model_id = ?", providerModelID).Order("started_at DESC, id DESC").Limit(1).Find(&latest)
if result.Error != nil {
return fmt.Errorf("read latest connectivity check: %w", result.Error)
}
now := s.nowUTC()
if result.RowsAffected == 1 && latest.CooldownUntil.After(now) {
return CooldownError{RetryAfterSeconds: secondsUntil(latest.CooldownUntil, now)}
}
check = connectivityRow{ProviderModelID: providerModelID, OperatorID: actor, Status: "running", StartedAt: now, CooldownUntil: now.Add(s.connectivityCooldown)}
if err := tx.Create(&check).Error; err != nil {
return fmt.Errorf("reserve connectivity check: %w", err)
}
return s.audit(tx, actor, requestID, "provider_model.connectivity.start", "provider_model", providerModelID, "succeeded", map[string]any{"check_id": check.ID})
})
if err != nil {
return ConnectivityCheckView{}, err
}
started := s.nowUTC()
configuration, configErr := s.probeConfiguration(ctx, providerModelID)
var probeErr error
if configErr != nil {
probeErr = configErr
} else {
probeErr = s.probe.Probe(ctx, configuration)
}
completed := s.nowUTC()
latencyMillis := completed.Sub(started).Milliseconds()
if latencyMillis < 0 {
latencyMillis = 0
}
if uint64(latencyMillis) > uint64(^uint32(0)) {
latencyMillis = int64(^uint32(0))
}
latency := uint32(latencyMillis)
check.Status, check.CompletedAt, check.LatencyMS = "succeeded", &completed, &latency
if probeErr != nil {
code, message := "connectivity_failed", "connectivity check failed"
check.Status, check.ErrorCode, check.ErrorMessage = "failed", &code, &message
}
if err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
updates := map[string]any{"status": check.Status, "completed_at": completed, "latency_ms": latency, "error_code": check.ErrorCode, "error_message": check.ErrorMessage}
if err := tx.Model(&connectivityRow{}).Where("id = ? AND status = 'running'", check.ID).Updates(updates).Error; err != nil {
return fmt.Errorf("complete connectivity check: %w", err)
}
result := "succeeded"
if probeErr != nil {
result = "failed"
}
return s.audit(tx, actor, requestID+":complete", "provider_model.connectivity.complete", "provider_model", providerModelID, result, map[string]any{"check_id": check.ID, "latency_ms": latency, "error_code": check.ErrorCode})
}); err != nil {
return ConnectivityCheckView{}, err
}
return connectivityView(check), nil
}
func (s *Service) ProviderHealth(ctx context.Context) ([]ProviderHealthView, error) {
type healthRow struct {
ProviderID uint64
ProviderName string
ProviderModelID uint64
ModelName string
Status *string
ErrorCode *string
ErrorMessage *string
LatencyMS *uint32
CompletedAt *time.Time
CooldownUntil *time.Time
}
var rows []healthRow
err := s.db.WithContext(ctx).Raw(`
SELECT p.id AS provider_id, p.name AS provider_name,
pm.id AS provider_model_id, pm.name AS model_name,
checks.status, checks.error_code, checks.error_message,
checks.latency_ms, checks.completed_at, checks.cooldown_until
FROM provider_models pm
JOIN providers p ON p.id = pm.provider_id
LEFT JOIN provider_connectivity_checks checks ON checks.id = (
SELECT latest.id FROM provider_connectivity_checks latest
WHERE latest.provider_model_id = pm.id
ORDER BY latest.started_at DESC, latest.id DESC LIMIT 1
)
ORDER BY p.name ASC, pm.name ASC, pm.id ASC`).Scan(&rows).Error
if err != nil {
return nil, fmt.Errorf("list provider health: %w", err)
}
items := make([]ProviderHealthView, 0, len(rows))
for _, row := range rows {
item := ProviderHealthView{ProviderID: row.ProviderID, ProviderName: row.ProviderName, ProviderModelID: row.ProviderModelID, ModelName: row.ModelName, Status: "never", LatencyMS: row.LatencyMS, CompletedAt: row.CompletedAt, CooldownUntil: row.CooldownUntil}
if row.Status != nil {
item.Status = *row.Status
}
if row.ErrorCode != nil {
item.ErrorCode = *row.ErrorCode
}
if row.ErrorMessage != nil {
item.ErrorMessage = truncate(*row.ErrorMessage, 256)
}
items = append(items, item)
}
return items, nil
}
func (s *Service) newCredential(providerID uint64, version uint32, actor uint64, apiKey string) (credentialRow, error) {
apiKey = strings.TrimSpace(apiKey)
if apiKey == "" || len(apiKey) > 4096 {
return credentialRow{}, FieldError{Field: "api_key", Message: "must be 1 to 4096 characters"}
}
return credentialRow{ProviderID: providerID, CredentialVersion: version, APIKey: apiKey, Status: "active", CreatedBy: actor, UpdatedBy: actor}, nil
}
func (s *Service) providerModelViews(ctx context.Context, rows []providerModelRow) ([]ProviderModelView, error) {
if len(rows) == 0 {
return []ProviderModelView{}, nil
}
ids := make([]uint64, 0, len(rows))
providerIDs := make([]uint64, 0, len(rows))
for _, row := range rows {
ids, providerIDs = append(ids, row.ID), append(providerIDs, row.ProviderID)
}
var capabilities []modelCapabilityRow
if err := s.db.WithContext(ctx).Where("provider_model_id IN ?", ids).Find(&capabilities).Error; err != nil {
return nil, fmt.Errorf("list model capabilities: %w", err)
}
var providers []providerRow
if err := s.db.WithContext(ctx).Where("id IN ?", providerIDs).Find(&providers).Error; err != nil {
return nil, fmt.Errorf("list model providers: %w", err)
}
capabilityByID := make(map[uint64][]model.Capability, len(rows))
for _, capability := range capabilities {
capabilityByID[capability.ProviderModelID] = append(capabilityByID[capability.ProviderModelID], capability.Capability)
}
providerName := make(map[uint64]string, len(providers))
for _, provider := range providers {
providerName[provider.ID] = provider.Name
}
items := make([]ProviderModelView, 0, len(rows))
for _, row := range rows {
caps := capabilityByID[row.ID]
slices.Sort(caps)
items = append(items, ProviderModelView{ID: row.ID, ProviderID: row.ProviderID, ProviderName: providerName[row.ProviderID], Name: row.Name, ModelID: row.ModelID, APIType: row.APIType, Kind: row.Kind, Capabilities: caps, ExtraBody: append(json.RawMessage(nil), row.ExtraBody...), TimeoutMS: row.TimeoutMS, Weight: row.Weight, Enabled: row.Enabled, CreatedAt: row.CreatedAt, UpdatedAt: row.UpdatedAt})
}
return items, nil
}
func (s *Service) providerViews(ctx context.Context, rows []providerRow) ([]ProviderView, error) {
if len(rows) == 0 {
return []ProviderView{}, nil
}
credentialIDs := make([]uint64, 0, len(rows))
for _, row := range rows {
if row.ActiveCredentialID != nil {
credentialIDs = append(credentialIDs, *row.ActiveCredentialID)
}
}
credentialsByID := make(map[uint64]credentialRow, len(credentialIDs))
if len(credentialIDs) > 0 {
var credentials []credentialRow
if err := s.db.WithContext(ctx).Where("id IN ? AND status = 'active'", credentialIDs).Find(&credentials).Error; err != nil {
return nil, fmt.Errorf("list active provider credentials: %w", err)
}
for _, credential := range credentials {
credentialsByID[credential.ID] = credential
}
}
items := make([]ProviderView, 0, len(rows))
for _, row := range rows {
var version *uint32
if row.ActiveCredentialID != nil {
if credential, ok := credentialsByID[*row.ActiveCredentialID]; ok {
version = &credential.CredentialVersion
}
}
items = append(items, providerView(row, version))
}
return items, nil
}
func (s *Service) validateRoutePoolReferences(tx *gorm.DB, input RoutePoolInput) error {
var template promptTemplateRow
if err := tx.First(&template, input.PromptTemplateID).Error; err != nil {
return FieldError{Field: "prompt_template_id", Message: "does not exist"}
}
if template.Capability != input.Capability {
return FieldError{Field: "prompt_template_id", Message: "does not match route capability"}
}
if len(input.Members) == 0 {
return nil
}
modelIDs := make([]uint64, 0, len(input.Members))
for _, member := range input.Members {
modelIDs = append(modelIDs, member.ProviderModelID)
}
var models []providerModelRow
if err := tx.Where("id IN ?", modelIDs).Find(&models).Error; err != nil {
return fmt.Errorf("load route models: %w", err)
}
if len(models) != len(modelIDs) {
return FieldError{Field: "members", Message: "contains an unknown provider model"}
}
var capabilities []modelCapabilityRow
if err := tx.Where("provider_model_id IN ? AND capability = ?", modelIDs, input.Capability).Find(&capabilities).Error; err != nil {
return fmt.Errorf("load route capabilities: %w", err)
}
if len(capabilities) != len(modelIDs) {
return FieldError{Field: "members", Message: "contains a model without the route capability"}
}
if input.Publish {
if !input.Enabled {
return FieldError{Field: "enabled", Message: "must be true when publishing"}
}
providerIDs := make([]uint64, 0, len(models))
for _, row := range models {
providerIDs = append(providerIDs, row.ProviderID)
}
var providers []providerRow
if err := tx.Where("id IN ?", providerIDs).Find(&providers).Error; err != nil {
return fmt.Errorf("load route providers: %w", err)
}
providersByID := make(map[uint64]providerRow, len(providers))
for _, row := range providers {
providersByID[row.ID] = row
}
hasEnabled := false
for _, member := range input.Members {
modelRecord := findModel(models, member.ProviderModelID)
providerRecord := providersByID[modelRecord.ProviderID]
if member.Enabled && modelRecord.Enabled && providerRecord.Enabled && (providerRecord.AuthType == string(provider.AuthNone) || providerRecord.ActiveCredentialID != nil) {
hasEnabled = true
}
}
if !hasEnabled {
return FieldError{Field: "members", Message: "publishing requires an enabled, credential-ready member"}
}
}
return nil
}
func (s *Service) publishRoute(tx *gorm.DB, actor uint64, row routePoolRow, input RoutePoolInput) error {
if !input.Publish {
return nil
}
active := activeRouteRow{Capability: row.Capability, RoutePoolID: row.ID, OperatorID: actor}
if err := tx.Clauses(clause.OnConflict{Columns: []clause.Column{{Name: "capability"}}, DoUpdates: clause.Assignments(map[string]any{"route_pool_id": row.ID, "operator_id": actor, "updated_at": s.nowUTC()})}).Create(&active).Error; err != nil {
return fmt.Errorf("publish active route: %w", err)
}
return nil
}
func (s *Service) routePoolView(ctx context.Context, row routePoolRow) (RoutePoolView, error) {
var members []routeMemberRow
if err := s.db.WithContext(ctx).Where("route_pool_id = ?", row.ID).Order("position ASC, id ASC").Find(&members).Error; err != nil {
return RoutePoolView{}, fmt.Errorf("list route members: %w", err)
}
var activeCount int64
if err := s.db.WithContext(ctx).Model(&activeRouteRow{}).Where("capability = ? AND route_pool_id = ?", row.Capability, row.ID).Count(&activeCount).Error; err != nil {
return RoutePoolView{}, fmt.Errorf("read active route: %w", err)
}
items := make([]RouteMemberInput, 0, len(members))
for _, member := range members {
items = append(items, RouteMemberInput{ProviderModelID: member.ProviderModelID, Weight: member.Weight, FailureThreshold: member.FailureThreshold, OpenSeconds: member.OpenSeconds, HalfOpenMax: member.HalfOpenMax, Enabled: member.Enabled, Position: member.Position})
}
return RoutePoolView{ID: row.ID, Slug: row.Slug, Name: row.Name, Capability: row.Capability, PromptTemplateID: row.PromptTemplateID, MaxFailover: row.MaxFailover, Version: row.Version, Enabled: row.Enabled, Active: activeCount == 1, Members: items, CreatedAt: row.CreatedAt, UpdatedAt: row.UpdatedAt}, nil
}
func (s *Service) probeConfiguration(ctx context.Context, providerModelID uint64) (ProbeConfiguration, error) {
type selected struct {
BaseURL string
AuthType string
APIKey string
ModelID string
APIType model.APIType
Kind model.GenerationKind
ExtraBody json.RawMessage
TimeoutMS uint32
}
var row selected
result := s.db.WithContext(ctx).Raw(`
SELECT p.base_url, p.auth_type, credential.api_key, pm.model_id, pm.api_type, pm.kind, pm.extra_body, pm.timeout_ms
FROM provider_models pm
JOIN providers p ON p.id = pm.provider_id
LEFT JOIN provider_credentials credential ON credential.id = p.active_credential_id AND credential.status = 'active'
WHERE pm.id = ? AND pm.enabled = TRUE AND p.enabled = TRUE
LIMIT 1`, providerModelID).Scan(&row)
if result.Error != nil {
return ProbeConfiguration{}, fmt.Errorf("load connectivity configuration: %w", result.Error)
}
if result.RowsAffected != 1 {
return ProbeConfiguration{}, errors.New("connectivity target is not enabled")
}
authType := provider.AuthType(row.AuthType)
if !authType.Valid() {
return ProbeConfiguration{}, errors.New("connectivity target authentication is invalid")
}
var apiKey string
if authType != provider.AuthNone {
if strings.TrimSpace(row.APIKey) == "" {
return ProbeConfiguration{}, errors.New("connectivity credential is unavailable")
}
apiKey = row.APIKey
}
var capabilityRows []modelCapabilityRow
if err := s.db.WithContext(ctx).Where("provider_model_id = ?", providerModelID).Find(&capabilityRows).Error; err != nil {
return ProbeConfiguration{}, fmt.Errorf("load connectivity capabilities: %w", err)
}
capabilities := make([]model.Capability, 0, len(capabilityRows))
for _, capability := range capabilityRows {
capabilities = append(capabilities, capability.Capability)
}
return ProbeConfiguration{BaseURL: row.BaseURL, AuthType: authType, APIKey: apiKey, ModelID: row.ModelID, APIType: row.APIType, Kind: row.Kind, Capabilities: capabilities, ExtraBody: row.ExtraBody, Timeout: time.Duration(row.TimeoutMS) * time.Millisecond, MaxResponseBytes: s.maxResponseBytes}, nil
}
func (s *Service) audit(tx *gorm.DB, actor uint64, requestID, action, target string, targetID uint64, result string, summary map[string]any) error {
if requestID == "" {
requestID = uuid.NewString()
}
if len(requestID) > 128 {
requestID = requestID[:128]
}
encoded, err := json.Marshal(summary)
if err != nil {
return errors.New("encode admin audit summary")
}
event := auditRow{Operator: actor, Action: action, Target: target, TargetID: fmt.Sprintf("%d", targetID), Result: result, RequestID: requestID, Summary: encoded}
if err := tx.Create(&event).Error; err != nil {
return fmt.Errorf("write admin audit event: %w", err)
}
return nil
}
func (s *Service) nowUTC() time.Time { return s.now().UTC() }
func validateProvider(input ProviderInput, creating bool) error {
if strings.TrimSpace(input.Slug) == "" || len(input.Slug) > 80 {
return FieldError{Field: "slug", Message: "must be 1 to 80 characters"}
}
if strings.TrimSpace(input.Name) == "" || len(input.Name) > 120 {
return FieldError{Field: "name", Message: "must be 1 to 120 characters"}
}
if err := validateBaseURL(input.BaseURL); err != nil {
return err
}
authType := provider.AuthType(input.AuthType)
if !authType.Valid() {
return FieldError{Field: "auth_type", Message: "must be none, bearer, or x-goog-api-key"}
}
if input.APIKey != nil && strings.TrimSpace(*input.APIKey) == "" {
return FieldError{Field: "api_key", Message: "must not be empty when supplied"}
}
if authType == provider.AuthNone && input.APIKey != nil {
return FieldError{Field: "api_key", Message: "is not allowed when auth_type is none"}
}
if creating && input.Enabled && authType != provider.AuthNone && input.APIKey == nil {
return FieldError{Field: "api_key", Message: "is required before enabling this provider"}
}
return nil
}
func validateBaseURL(raw string) error {
parsed, err := url.Parse(strings.TrimSpace(raw))
if err != nil || (parsed.Scheme != "http" && parsed.Scheme != "https") || parsed.Hostname() == "" || parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" {
return FieldError{Field: "base_url", Message: "must be an http or https origin or path without credentials, query, or fragment"}
}
return nil
}
func validateProviderModel(input ProviderModelInput) error {
if input.ProviderID == 0 {
return FieldError{Field: "provider_id", Message: "must be positive"}
}
if strings.TrimSpace(input.Name) == "" || len(input.Name) > 120 || strings.TrimSpace(input.ModelID) == "" || len(input.ModelID) > 191 {
return FieldError{Field: "model_id", Message: "name and model_id are required within their limits"}
}
if input.TimeoutMS == 0 || input.Weight == 0 {
return FieldError{Field: "timeout_ms", Message: "timeout_ms and weight must be positive"}
}
if !input.Kind.Valid() || !validAPIType(input.APIType) {
return FieldError{Field: "api_type", Message: "is invalid"}
}
if err := validateCapabilities(input.APIType, input.Kind, input.Capabilities); err != nil {
return err
}
return validateExtraBody(input.ExtraBody)
}
func validatePromptTemplate(input PromptTemplateInput) error {
if strings.TrimSpace(input.TemplateKey) == "" || len(input.TemplateKey) > 120 || strings.TrimSpace(input.Name) == "" || len(input.Name) > 120 {
return FieldError{Field: "template_key", Message: "template_key and name are required within their limits"}
}
if strings.TrimSpace(input.TemplateText) == "" || len(input.DefaultRoleRule) > 65535 || !input.Capability.Valid() || !input.Kind.Valid() || !validAPIType(input.APIType) {
return FieldError{Field: "template", Message: "template fields are invalid"}
}
if err := validateCapabilities(input.APIType, input.Kind, []model.Capability{input.Capability}); err != nil {
return FieldError{Field: "capability", Message: "does not match api_type and kind"}
}
return nil
}
func validateRoutePoolInput(input RoutePoolInput, requireMembers bool) error {
if strings.TrimSpace(input.Slug) == "" || len(input.Slug) > 80 || strings.TrimSpace(input.Name) == "" || len(input.Name) > 120 {
return FieldError{Field: "slug", Message: "slug and name are required within their limits"}
}
if !input.Capability.Valid() || input.PromptTemplateID == 0 {
return FieldError{Field: "capability", Message: "capability and prompt_template_id are required"}
}
if len(input.Members) > maxRouteMembers || (requireMembers && input.Publish && len(input.Members) == 0) {
return FieldError{Field: "members", Message: "published route pools need 1 to 32 members"}
}
seenModels, seenPositions := map[uint64]bool{}, map[uint32]bool{}
for _, member := range input.Members {
if member.ProviderModelID == 0 || member.Weight == 0 || member.FailureThreshold == 0 || member.OpenSeconds == 0 || member.HalfOpenMax == 0 {
return FieldError{Field: "members", Message: "member model, weight, breaker threshold, open duration, and half-open limit must be positive"}
}
if seenModels[member.ProviderModelID] || seenPositions[member.Position] {
return FieldError{Field: "members", Message: "model and position must be unique"}
}
seenModels[member.ProviderModelID], seenPositions[member.Position] = true, true
}
for position := range input.Members {
if !seenPositions[uint32(position)] {
return FieldError{Field: "members", Message: "positions must be continuous from zero"}
}
}
return nil
}
func validateCapabilities(apiType model.APIType, kind model.GenerationKind, capabilities []model.Capability) error {
if len(capabilities) == 0 || len(capabilities) > 3 {
return FieldError{Field: "capabilities", Message: "must contain one or more supported capabilities"}
}
seen := map[model.Capability]bool{}
for _, capability := range capabilities {
if !capability.Valid() || seen[capability] {
return FieldError{Field: "capabilities", Message: "contains an invalid or duplicate capability"}
}
seen[capability] = true
}
switch apiType {
case model.APIChat:
if kind != model.KindText || !seen[model.CapabilityText] || len(seen) != 1 {
return FieldError{Field: "capabilities", Message: "chat only supports text"}
}
case model.APIImages:
if kind != model.KindImage || !seen[model.CapabilityImageGenerate] || len(seen) != 1 {
return FieldError{Field: "capabilities", Message: "images only supports image_generate"}
}
case model.APIImagesEdits:
if kind != model.KindImage || !seen[model.CapabilityImageEdit] || len(seen) != 1 {
return FieldError{Field: "capabilities", Message: "images_edits only supports image_edit"}
}
case model.APIGemini:
if kind == model.KindText && (!seen[model.CapabilityText] || len(seen) != 1) {
return FieldError{Field: "capabilities", Message: "text Gemini models must declare text only"}
}
if kind == model.KindImage && seen[model.CapabilityText] {
return FieldError{Field: "capabilities", Message: "image Gemini models cannot declare text"}
}
default:
return FieldError{Field: "api_type", Message: "is invalid"}
}
return nil
}
func validateExtraBody(raw json.RawMessage) error {
if len(raw) == 0 || string(raw) == "null" {
return nil
}
var fields map[string]any
if err := json.Unmarshal(raw, &fields); err != nil {
return FieldError{Field: "extra_body", Message: "must be a JSON object"}
}
allowed := map[string]bool{"temperature": true, "max_tokens": true, "size": true, "quality": true}
for key, value := range fields {
if !allowed[key] {
return FieldError{Field: "extra_body", Message: "contains a reserved or unsupported field"}
}
switch value.(type) {
case string, float64, bool:
default:
return FieldError{Field: "extra_body", Message: "values must be strings, numbers, or booleans"}
}
}
return nil
}
func validAPIType(apiType model.APIType) bool {
return apiType == model.APIChat || apiType == model.APIImages || apiType == model.APIImagesEdits || apiType == model.APIGemini
}
func replaceCapabilities(tx *gorm.DB, providerModelID uint64, capabilities []model.Capability) error {
if err := tx.Where("provider_model_id = ?", providerModelID).Delete(&modelCapabilityRow{}).Error; err != nil {
return fmt.Errorf("clear model capabilities: %w", err)
}
rows := make([]modelCapabilityRow, 0, len(capabilities))
for _, capability := range capabilities {
rows = append(rows, modelCapabilityRow{ProviderModelID: providerModelID, Capability: capability})
}
if err := tx.Create(&rows).Error; err != nil {
return fmt.Errorf("create model capabilities: %w", err)
}
return nil
}
func replaceRouteMembers(tx *gorm.DB, routePoolID uint64, members []RouteMemberInput) error {
if err := tx.Where("route_pool_id = ?", routePoolID).Delete(&routeMemberRow{}).Error; err != nil {
return fmt.Errorf("replace route members: %w", err)
}
if len(members) == 0 {
return nil
}
rows := make([]routeMemberRow, 0, len(members))
for _, member := range members {
rows = append(rows, routeMemberRow{RoutePoolID: routePoolID, ProviderModelID: member.ProviderModelID, Weight: member.Weight, FailureThreshold: member.FailureThreshold, OpenSeconds: member.OpenSeconds, HalfOpenMax: member.HalfOpenMax, Enabled: member.Enabled, Position: member.Position})
}
if err := tx.Create(&rows).Error; err != nil {
return fmt.Errorf("create route members: %w", err)
}
return nil
}
func ensureProvider(tx *gorm.DB, id uint64) error {
if id == 0 {
return FieldError{Field: "provider_id", Message: "must be positive"}
}
var count int64
if err := tx.Model(&providerRow{}).Where("id = ?", id).Count(&count).Error; err != nil {
return fmt.Errorf("read provider: %w", err)
}
if count != 1 {
return FieldError{Field: "provider_id", Message: "does not exist"}
}
return nil
}
func ensureProviderModel(tx *gorm.DB, id uint64) error {
var count int64
if err := tx.Model(&providerModelRow{}).Where("id = ?", id).Count(&count).Error; err != nil {
return fmt.Errorf("read provider model: %w", err)
}
if count != 1 {
return ErrNotFound
}
return nil
}
func ensureProviderModelLocked(tx *gorm.DB, id uint64) error {
var row providerModelRow
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&row, id).Error; err != nil {
return translateNotFound(err)
}
return nil
}
func providerView(row providerRow, version *uint32) ProviderView {
view := ProviderView{ID: row.ID, Slug: row.Slug, Name: row.Name, BaseURL: row.BaseURL, AuthType: row.AuthType, Enabled: row.Enabled, CreatedAt: row.CreatedAt, UpdatedAt: row.UpdatedAt}
if row.ActiveCredentialID != nil && version != nil {
view.Credential = CredentialInfo{HasCredential: true, Version: *version}
}
return view
}
func promptViews(rows []promptTemplateRow) []PromptTemplateView {
items := make([]PromptTemplateView, 0, len(rows))
for _, row := range rows {
items = append(items, PromptTemplateView{ID: row.ID, TemplateKey: row.TemplateKey, Kind: row.Kind, APIType: row.APIType, Capability: row.Capability, Name: row.Name, Version: row.Version, TemplateText: row.TemplateText, DefaultRoleRule: row.DefaultRoleRule, Enabled: row.Enabled, CreatedAt: row.CreatedAt, UpdatedAt: row.UpdatedAt})
}
return items
}
func connectivityView(row connectivityRow) ConnectivityCheckView {
view := ConnectivityCheckView{ID: row.ID, ProviderModel: row.ProviderModelID, Status: row.Status, LatencyMS: row.LatencyMS, StartedAt: row.StartedAt, CompletedAt: row.CompletedAt, CooldownUntil: row.CooldownUntil}
if row.ErrorCode != nil {
view.ErrorCode = *row.ErrorCode
}
if row.ErrorMessage != nil {
view.ErrorMessage = *row.ErrorMessage
}
return view
}
func findModel(rows []providerModelRow, id uint64) providerModelRow {
for _, row := range rows {
if row.ID == id {
return row
}
}
return providerModelRow{}
}
func translateNotFound(err error) error {
if errors.Is(err, gorm.ErrRecordNotFound) {
return ErrNotFound
}
return err
}
func normalizedExtraBody(raw json.RawMessage) json.RawMessage {
if len(raw) == 0 || string(raw) == "null" {
return json.RawMessage("{}")
}
return append(json.RawMessage(nil), raw...)
}
func clear(value []byte) {
for index := range value {
value[index] = 0
}
}
func truncate(value string, limit int) string {
if len(value) <= limit {
return value
}
return value[:limit]
}
func secondsUntil(until, now time.Time) uint32 {
seconds := math.Ceil(until.Sub(now).Seconds())
if seconds < 1 {
return 1
}
if seconds > float64(^uint32(0)) {
return ^uint32(0)
}
return uint32(seconds)
}