This repository has been archived on 2026-04-03. You can view files and clone it. You cannot open issues or pull requests or push a commit.
Files
core/server/roles/roles.go
2026-02-27 16:25:39 +00:00

472 lines
14 KiB
Go

package roles
import (
"context"
"encoding/json"
"fmt"
"log"
"sync"
"time"
"git.gobha.me/xcaliber/chat-switchboard/crypto"
"git.gobha.me/xcaliber/chat-switchboard/events"
"git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/providers"
"git.gobha.me/xcaliber/chat-switchboard/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
}
}
}
return prov, providers.ProviderConfig{
Endpoint: cfg.Endpoint,
APIKey: apiKey,
CustomHeaders: headers,
}, 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(),
})
}