- go.mod module name
- All 714 import references across 289 Go files
- VERSION: 0.1.0
- CI DB names: switchboard_core_{ci,dev,test}
- Docker image: gobha/switchboard-core
- Test fixtures: JWT issuer, repo names
- .env.example, docker-compose container name
- Compiles clean (go build exit 0)
478 lines
14 KiB
Go
478 lines
14 KiB
Go
package roles
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"log"
|
|
"sync"
|
|
"time"
|
|
|
|
"switchboard-core/crypto"
|
|
"switchboard-core/events"
|
|
"switchboard-core/models"
|
|
"switchboard-core/providers"
|
|
"switchboard-core/store"
|
|
)
|
|
|
|
// ── Known Roles ────────────────────────────
|
|
|
|
const (
|
|
RoleUtility = "utility" // Internal tasks: summarization, title generation
|
|
RoleEmbedding = "embedding" // Vector embedding for knowledge bases
|
|
)
|
|
|
|
// ValidRoles lists all recognized role names.
|
|
// Note: "generation" (image/media) was removed in v0.10.2 — image gen
|
|
// will be extension-managed with its own provider config, not a role slot.
|
|
var ValidRoles = []string{RoleUtility, RoleEmbedding}
|
|
|
|
// IsValidRole returns true if the given role name is recognized.
|
|
func IsValidRole(role string) bool {
|
|
for _, r := range ValidRoles {
|
|
if r == role {
|
|
return true
|
|
}
|
|
}
|
|
return false
|
|
}
|
|
|
|
// ── Types ──────────────────────────────────
|
|
|
|
// RoleBinding pairs a provider config with a specific model.
|
|
type RoleBinding struct {
|
|
ProviderConfigID string `json:"provider_config_id"`
|
|
ModelID string `json:"model_id"`
|
|
}
|
|
|
|
// RoleConfig holds primary and fallback bindings for a role slot.
|
|
type RoleConfig struct {
|
|
Primary *RoleBinding `json:"primary"`
|
|
Fallback *RoleBinding `json:"fallback"`
|
|
}
|
|
|
|
// CompletionResult wraps a completion response with usage metadata.
|
|
type CompletionResult struct {
|
|
Content string
|
|
Model string
|
|
ProviderID string
|
|
ConfigID string
|
|
ProviderScope string
|
|
InputTokens int
|
|
OutputTokens int
|
|
CacheCreationTokens int
|
|
CacheReadTokens int
|
|
Role string
|
|
UsedFallback bool
|
|
}
|
|
|
|
// EmbeddingResult wraps an embedding response with usage metadata.
|
|
type EmbeddingResult struct {
|
|
Embeddings [][]float64
|
|
Model string
|
|
ProviderID string
|
|
ConfigID string
|
|
ProviderScope string
|
|
InputTokens int
|
|
Role string
|
|
UsedFallback bool
|
|
}
|
|
|
|
// ── Errors ─────────────────────────────────
|
|
|
|
var (
|
|
ErrRoleNotConfigured = fmt.Errorf("role not configured")
|
|
ErrNoPrimary = fmt.Errorf("no primary binding for role")
|
|
)
|
|
|
|
// ── Resolver ───────────────────────────────
|
|
|
|
// Resolver resolves named model roles to provider+model pairs and
|
|
// executes completions/embeddings against them. It handles the full
|
|
// pipeline: config lookup → key decryption → provider dispatch → fallback.
|
|
type Resolver struct {
|
|
stores store.Stores
|
|
vault *crypto.KeyResolver
|
|
bus *events.Bus // optional — nil-safe
|
|
|
|
// Fallback cooldown: suppress duplicate alerts per role
|
|
coolMu sync.Mutex
|
|
cooldown map[string]time.Time // role → last alert timestamp
|
|
}
|
|
|
|
const fallbackCooldown = 5 * time.Minute
|
|
|
|
// NewResolver creates a role resolver with access to stores and vault.
|
|
func NewResolver(s store.Stores, vault *crypto.KeyResolver) *Resolver {
|
|
return &Resolver{stores: s, vault: vault, cooldown: make(map[string]time.Time)}
|
|
}
|
|
|
|
// WithBus attaches an event bus for fallback alert publishing.
|
|
func (r *Resolver) WithBus(bus *events.Bus) *Resolver {
|
|
r.bus = bus
|
|
return r
|
|
}
|
|
|
|
// Complete sends a chat completion using the named role.
|
|
// Resolution: personal override → team override → global config → try primary → fallback on error.
|
|
func (r *Resolver) Complete(ctx context.Context, role string, userID string, teamID *string, messages []providers.Message) (*CompletionResult, error) {
|
|
cfg, err := r.GetConfig(ctx, role, userID, teamID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// Try primary
|
|
if cfg.Primary != nil {
|
|
result, err := r.doComplete(ctx, role, cfg.Primary, messages)
|
|
if err == nil {
|
|
return result, nil
|
|
}
|
|
log.Printf("⚠ Role %q primary failed: %v", role, err)
|
|
}
|
|
|
|
// Try fallback
|
|
if cfg.Fallback != nil {
|
|
result, err := r.doComplete(ctx, role, cfg.Fallback, messages)
|
|
if err == nil {
|
|
result.UsedFallback = true
|
|
r.onFallback(ctx, role, "completion", cfg.Primary, cfg.Fallback)
|
|
return result, nil
|
|
}
|
|
return nil, fmt.Errorf("role %q: both primary and fallback failed: %w", role, err)
|
|
}
|
|
|
|
if cfg.Primary == nil {
|
|
return nil, fmt.Errorf("%w: %s", ErrNoPrimary, role)
|
|
}
|
|
return nil, fmt.Errorf("role %q: primary failed with no fallback configured", role)
|
|
}
|
|
|
|
// Embed generates embeddings using the named role.
|
|
func (r *Resolver) Embed(ctx context.Context, role string, userID string, teamID *string, input []string) (*EmbeddingResult, error) {
|
|
cfg, err := r.GetConfig(ctx, role, userID, teamID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// Try primary
|
|
if cfg.Primary != nil {
|
|
result, err := r.doEmbed(ctx, role, cfg.Primary, input)
|
|
if err == nil {
|
|
return result, nil
|
|
}
|
|
log.Printf("⚠ Role %q primary embed failed: %v", role, err)
|
|
}
|
|
|
|
// Try fallback
|
|
if cfg.Fallback != nil {
|
|
result, err := r.doEmbed(ctx, role, cfg.Fallback, input)
|
|
if err == nil {
|
|
result.UsedFallback = true
|
|
r.onFallback(ctx, role, "embedding", cfg.Primary, cfg.Fallback)
|
|
return result, nil
|
|
}
|
|
return nil, fmt.Errorf("role %q: both primary and fallback embed failed: %w", role, err)
|
|
}
|
|
|
|
if cfg.Primary == nil {
|
|
return nil, fmt.Errorf("%w: %s", ErrNoPrimary, role)
|
|
}
|
|
return nil, fmt.Errorf("role %q: primary embed failed with no fallback", role)
|
|
}
|
|
|
|
// GetConfig returns the resolved role config for the given role name.
|
|
// Resolution order: personal override → team override → global config.
|
|
func (r *Resolver) GetConfig(ctx context.Context, role string, userID string, teamID *string) (*RoleConfig, error) {
|
|
if !IsValidRole(role) {
|
|
return nil, fmt.Errorf("unknown role: %q", role)
|
|
}
|
|
|
|
// Check personal override first (BYOK users)
|
|
if userID != "" {
|
|
personalCfg, err := r.getPersonalRoleConfig(ctx, userID, role)
|
|
if err == nil && personalCfg != nil && (personalCfg.Primary != nil || personalCfg.Fallback != nil) {
|
|
return personalCfg, nil
|
|
}
|
|
}
|
|
|
|
// Check team override
|
|
if teamID != nil && *teamID != "" {
|
|
teamCfg, err := r.getTeamRoleConfig(ctx, *teamID, role)
|
|
if err == nil && teamCfg != nil && (teamCfg.Primary != nil || teamCfg.Fallback != nil) {
|
|
return teamCfg, nil
|
|
}
|
|
}
|
|
|
|
// Fall back to global
|
|
return r.getGlobalRoleConfig(ctx, role)
|
|
}
|
|
|
|
// IsConfigured returns true if the named role has at least a primary binding.
|
|
func (r *Resolver) IsConfigured(ctx context.Context, role string) bool {
|
|
cfg, err := r.GetConfig(ctx, role, "", nil)
|
|
return err == nil && cfg != nil && cfg.Primary != nil
|
|
}
|
|
|
|
// IsPersonalOverride returns true if the user has a personal binding for the role.
|
|
func (r *Resolver) IsPersonalOverride(ctx context.Context, userID, role string) bool {
|
|
cfg, err := r.getPersonalRoleConfig(ctx, userID, role)
|
|
return err == nil && cfg != nil && cfg.Primary != nil
|
|
}
|
|
|
|
// ── Internal: Completion ───────────────────
|
|
|
|
func (r *Resolver) doComplete(ctx context.Context, role string, binding *RoleBinding, messages []providers.Message) (*CompletionResult, error) {
|
|
prov, provCfg, scope, err := r.resolveBinding(ctx, binding)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
resp, err := prov.ChatCompletion(ctx, provCfg, providers.CompletionRequest{
|
|
Model: binding.ModelID,
|
|
Messages: messages,
|
|
})
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
return &CompletionResult{
|
|
Content: resp.Content,
|
|
Model: resp.Model,
|
|
ProviderID: prov.ID(),
|
|
ConfigID: binding.ProviderConfigID,
|
|
ProviderScope: scope,
|
|
InputTokens: resp.InputTokens,
|
|
OutputTokens: resp.OutputTokens,
|
|
CacheCreationTokens: resp.CacheCreationTokens,
|
|
CacheReadTokens: resp.CacheReadTokens,
|
|
Role: role,
|
|
}, nil
|
|
}
|
|
|
|
// ── Internal: Embedding ────────────────────
|
|
|
|
func (r *Resolver) doEmbed(ctx context.Context, role string, binding *RoleBinding, input []string) (*EmbeddingResult, error) {
|
|
prov, provCfg, scope, err := r.resolveBinding(ctx, binding)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
resp, err := prov.Embed(ctx, provCfg, providers.EmbeddingRequest{
|
|
Model: binding.ModelID,
|
|
Input: input,
|
|
})
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
return &EmbeddingResult{
|
|
Embeddings: resp.Embeddings,
|
|
Model: resp.Model,
|
|
ProviderID: prov.ID(),
|
|
ConfigID: binding.ProviderConfigID,
|
|
ProviderScope: scope,
|
|
InputTokens: resp.InputTokens,
|
|
Role: role,
|
|
}, nil
|
|
}
|
|
|
|
// ── Internal: Provider Resolution ──────────
|
|
|
|
// resolveBinding loads a provider config, decrypts the API key, and returns
|
|
// a ready-to-use Provider + ProviderConfig pair + scope.
|
|
func (r *Resolver) resolveBinding(ctx context.Context, binding *RoleBinding) (providers.Provider, providers.ProviderConfig, string, error) {
|
|
cfg, err := r.stores.Providers.GetByID(ctx, binding.ProviderConfigID)
|
|
if err != nil {
|
|
return nil, providers.ProviderConfig{}, "", fmt.Errorf("load provider config %s: %w", binding.ProviderConfigID, err)
|
|
}
|
|
|
|
prov, err := providers.Get(cfg.Provider)
|
|
if err != nil {
|
|
return nil, providers.ProviderConfig{}, "", fmt.Errorf("unknown provider %s: %w", cfg.Provider, err)
|
|
}
|
|
|
|
// Decrypt API key
|
|
apiKey := ""
|
|
if len(cfg.APIKeyEnc) > 0 && r.vault != nil {
|
|
apiKey, err = r.vault.Decrypt(cfg.APIKeyEnc, cfg.KeyNonce, cfg.KeyScope, "")
|
|
if err != nil {
|
|
return nil, providers.ProviderConfig{}, "", fmt.Errorf("decrypt API key: %w", err)
|
|
}
|
|
} else if len(cfg.APIKeyEnc) > 0 {
|
|
// No vault — key stored as raw bytes (unencrypted fallback)
|
|
apiKey = string(cfg.APIKeyEnc)
|
|
}
|
|
|
|
// Parse headers
|
|
headers := make(map[string]string)
|
|
if cfg.Headers != nil {
|
|
for k, v := range cfg.Headers {
|
|
if s, ok := v.(string); ok {
|
|
headers[k] = s
|
|
}
|
|
}
|
|
}
|
|
|
|
proxyURL := ""
|
|
if cfg.ProxyURL != nil {
|
|
proxyURL = *cfg.ProxyURL
|
|
}
|
|
return prov, providers.ProviderConfig{
|
|
Endpoint: cfg.Endpoint,
|
|
APIKey: apiKey,
|
|
CustomHeaders: headers,
|
|
ProxyMode: cfg.ProxyMode,
|
|
ProxyURL: proxyURL,
|
|
}, cfg.Scope, nil
|
|
}
|
|
|
|
// ── Internal: Config Loading ───────────────
|
|
|
|
func (r *Resolver) getPersonalRoleConfig(ctx context.Context, userID, role string) (*RoleConfig, error) {
|
|
user, err := r.stores.Users.GetByID(ctx, userID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
settings := user.Settings
|
|
if settings == nil {
|
|
return nil, nil
|
|
}
|
|
|
|
rolesRaw, ok := settings["model_roles"]
|
|
if !ok {
|
|
return nil, nil
|
|
}
|
|
|
|
rolesMap, ok := rolesRaw.(map[string]interface{})
|
|
if !ok {
|
|
return nil, nil
|
|
}
|
|
|
|
roleData, ok := rolesMap[role]
|
|
if !ok {
|
|
return nil, nil
|
|
}
|
|
|
|
return parseRoleConfig(roleData)
|
|
}
|
|
|
|
func (r *Resolver) getGlobalRoleConfig(ctx context.Context, role string) (*RoleConfig, error) {
|
|
allRoles, err := r.stores.GlobalConfig.Get(ctx, "model_roles")
|
|
if err != nil {
|
|
return nil, fmt.Errorf("load global model_roles: %w", err)
|
|
}
|
|
|
|
roleData, ok := allRoles[role]
|
|
if !ok {
|
|
return nil, fmt.Errorf("%w: %s (not in global settings)", ErrRoleNotConfigured, role)
|
|
}
|
|
|
|
return parseRoleConfig(roleData)
|
|
}
|
|
|
|
func (r *Resolver) getTeamRoleConfig(ctx context.Context, teamID, role string) (*RoleConfig, error) {
|
|
team, err := r.stores.Teams.GetByID(ctx, teamID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
settings := team.Settings
|
|
if settings == nil {
|
|
return nil, nil
|
|
}
|
|
|
|
rolesRaw, ok := settings["model_roles"]
|
|
if !ok {
|
|
return nil, nil
|
|
}
|
|
|
|
rolesMap, ok := rolesRaw.(map[string]interface{})
|
|
if !ok {
|
|
return nil, nil
|
|
}
|
|
|
|
roleData, ok := rolesMap[role]
|
|
if !ok {
|
|
return nil, nil
|
|
}
|
|
|
|
return parseRoleConfig(roleData)
|
|
}
|
|
|
|
// parseRoleConfig converts an interface{} (from JSONB) into a typed RoleConfig.
|
|
func parseRoleConfig(data interface{}) (*RoleConfig, error) {
|
|
b, err := json.Marshal(data)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
var cfg RoleConfig
|
|
if err := json.Unmarshal(b, &cfg); err != nil {
|
|
return nil, err
|
|
}
|
|
return &cfg, nil
|
|
}
|
|
|
|
// ── Fallback Alerting ──────────────────────
|
|
|
|
// onFallback fires on successful fallback activation. It emits:
|
|
// - log line (always)
|
|
// - audit_log entry (always)
|
|
// - event bus "role.fallback" (once per cooldown window)
|
|
//
|
|
// The cooldown prevents flooding the admin UI with identical alerts
|
|
// when a primary is down and every request triggers the fallback.
|
|
func (r *Resolver) onFallback(ctx context.Context, role, opType string, primary, fallback *RoleBinding) {
|
|
primaryModel := ""
|
|
fallbackModel := ""
|
|
if primary != nil {
|
|
primaryModel = primary.ModelID
|
|
}
|
|
if fallback != nil {
|
|
fallbackModel = fallback.ModelID
|
|
}
|
|
|
|
log.Printf("⚠ Role %q fallback activated: %s → %s (%s)", role, primaryModel, fallbackModel, opType)
|
|
|
|
// Audit log — always written (one row per fallback fire)
|
|
_ = r.stores.Audit.Log(ctx, &models.AuditEntry{
|
|
Action: "role.fallback",
|
|
ResourceType: "role",
|
|
ResourceID: role,
|
|
Metadata: models.JSONMap{
|
|
"operation": opType,
|
|
"primary_model": primaryModel,
|
|
"fallback_model": fallbackModel,
|
|
},
|
|
})
|
|
|
|
// Bus event — cooldown-gated
|
|
if r.bus == nil {
|
|
return
|
|
}
|
|
|
|
r.coolMu.Lock()
|
|
last, exists := r.cooldown[role]
|
|
now := time.Now()
|
|
if exists && now.Sub(last) < fallbackCooldown {
|
|
r.coolMu.Unlock()
|
|
return
|
|
}
|
|
r.cooldown[role] = now
|
|
r.coolMu.Unlock()
|
|
|
|
payload, _ := json.Marshal(map[string]string{
|
|
"role": role,
|
|
"operation": opType,
|
|
"primary_model": primaryModel,
|
|
"fallback_model": fallbackModel,
|
|
"message": fmt.Sprintf("Role %q primary (%s) failed — using fallback (%s)", role, primaryModel, fallbackModel),
|
|
})
|
|
r.bus.PublishAsync(events.Event{
|
|
Label: "role.fallback",
|
|
Room: "admin", // admin-targeted
|
|
Payload: payload,
|
|
Ts: now.UnixMilli(),
|
|
})
|
|
}
|