// Package backend owns validated, immutable OpenAI-compatible backend // definitions. package backend import ( "errors" "fmt" "net/url" "regexp" "sort" "strings" "gitea.maximumdirect.net/eric/promptkit/internal/domain" "gitea.maximumdirect.net/eric/promptkit/internal/jsonvalue" "gitea.maximumdirect.net/eric/promptkit/internal/llm" ) const ( // OpenRouterID is the reserved ID of Promptkit's built-in OpenRouter // backend. OpenRouterID = "openrouter" openRouterEndpoint = "https://openrouter.ai/api/v1" openRouterAPIKeyEnv = "OPENROUTER_API_KEY" openRouterConcurrencyLimit = 16 defaultQueueCapacity = 1024 ) // ErrBackendNotFound identifies a registry lookup for an unknown backend ID. var ErrBackendNotFound = errors.New("backend not found") var environmentVariableName = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`) // Registry is an immutable collection of validated backend definitions. type Registry struct { backends map[string]domain.Backend } // NewRegistry constructs a registry containing the built-in OpenRouter // definition followed by the supplied additions. Every ID must be unique. func NewRegistry(additions []domain.Backend) (*Registry, error) { registry := &Registry{ backends: make(map[string]domain.Backend, len(additions)+1), } definitions := make([]domain.Backend, 0, len(additions)+1) definitions = append(definitions, domain.Backend{ ID: OpenRouterID, Endpoint: openRouterEndpoint, APIKeyEnv: openRouterAPIKeyEnv, ConcurrencyLimit: openRouterConcurrencyLimit, }) definitions = append(definitions, additions...) for _, definition := range definitions { definition.ID = strings.TrimSpace(definition.ID) if definition.ID == "" { return nil, errors.New("backend ID must not be blank") } if _, exists := registry.backends[definition.ID]; exists { return nil, fmt.Errorf("backend ID %q is already registered", definition.ID) } normalized, err := normalizeBackend(definition) if err != nil { return nil, err } registry.backends[normalized.ID] = normalized } return registry, nil } // GetBackend returns a defensive copy of the backend registered with id. func (r *Registry) GetBackend(id string) (domain.Backend, error) { if r == nil { return domain.Backend{}, fmt.Errorf("%w: %q", ErrBackendNotFound, id) } definition, ok := r.backends[id] if !ok { return domain.Backend{}, fmt.Errorf("%w: %q", ErrBackendNotFound, id) } extraParams, err := jsonvalue.CopyMap(definition.ExtraParams) if err != nil { return domain.Backend{}, fmt.Errorf("copy backend %q: %w", id, err) } definition.ExtraParams = extraParams return definition, nil } // CapacityPolicies returns a copy of the normalized policies for limited // backends. func (r *Registry) CapacityPolicies() map[string]domain.BackendCapacityPolicy { policies := make(map[string]domain.BackendCapacityPolicy) if r == nil { return policies } for id, definition := range r.backends { if definition.ConcurrencyLimit == 0 { continue } policies[id] = domain.BackendCapacityPolicy{ ConcurrencyLimit: definition.ConcurrencyLimit, QueueCapacity: definition.QueueCapacity, } } return policies } func normalizeBackend(definition domain.Backend) (domain.Backend, error) { definition.Endpoint = strings.TrimSpace(definition.Endpoint) if err := validateEndpoint(definition.Endpoint); err != nil { return domain.Backend{}, fmt.Errorf("backend %q endpoint: %w", definition.ID, err) } definition.APIKeyEnv = strings.TrimSpace(definition.APIKeyEnv) if definition.APIKeyEnv != "" && !environmentVariableName.MatchString(definition.APIKeyEnv) { return domain.Backend{}, fmt.Errorf( "backend %q api key environment variable %q is invalid", definition.ID, definition.APIKeyEnv, ) } if definition.ConcurrencyLimit < 0 { return domain.Backend{}, fmt.Errorf( "backend %q concurrency limit must not be negative", definition.ID, ) } if definition.QueueCapacity < 0 { return domain.Backend{}, fmt.Errorf( "backend %q queue capacity must not be negative", definition.ID, ) } if definition.ConcurrencyLimit == 0 { if definition.QueueCapacitySet { return domain.Backend{}, fmt.Errorf( "backend %q queue capacity requires a positive concurrency limit", definition.ID, ) } definition.QueueCapacity = 0 } else { if !definition.QueueCapacitySet { definition.QueueCapacity = defaultQueueCapacity definition.QueueCapacitySet = true } maxInt := int(^uint(0) >> 1) if definition.QueueCapacity > maxInt-definition.ConcurrencyLimit { return domain.Backend{}, fmt.Errorf( "backend %q total capacity overflows int", definition.ID, ) } } keys := make([]string, 0, len(definition.ExtraParams)) for key := range definition.ExtraParams { keys = append(keys, key) } sort.Strings(keys) for _, key := range keys { if key == "" { return domain.Backend{}, fmt.Errorf("backend %q extra parameter key must not be empty", definition.ID) } if llm.IsReservedOpenAIChatRequestField(key) { return domain.Backend{}, fmt.Errorf( "backend %q extra parameter %q collides with a reserved request field", definition.ID, key, ) } } extraParams, err := jsonvalue.CopyMap(definition.ExtraParams) if err != nil { return domain.Backend{}, fmt.Errorf("backend %q extra parameters: %w", definition.ID, err) } definition.ExtraParams = extraParams return definition, nil } func validateEndpoint(endpoint string) error { if endpoint == "" { return errors.New("must not be blank") } if strings.Contains(endpoint, "#") { return errors.New("must not contain a fragment") } parsed, err := url.Parse(endpoint) if err != nil { return fmt.Errorf("must be a valid URL: %w", err) } scheme := strings.ToLower(parsed.Scheme) if scheme != "http" && scheme != "https" { return errors.New("must use http or https") } if !parsed.IsAbs() || parsed.Hostname() == "" { return errors.New("must be absolute and include a host") } if parsed.User != nil { return errors.New("must not contain user information") } if parsed.RawQuery != "" || parsed.ForceQuery { return errors.New("must not contain a query string") } return nil }