185 lines
5.3 KiB
Go
185 lines
5.3 KiB
Go
// Package backend owns validated, immutable OpenAI-compatible backend
|
|
// definitions.
|
|
package backend
|
|
|
|
import (
|
|
"errors"
|
|
"fmt"
|
|
"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) {
|
|
endpoint, err := domain.NormalizeOpenAICompatibleBaseEndpoint(definition.Endpoint)
|
|
if err != nil {
|
|
return domain.Backend{}, fmt.Errorf("backend %q endpoint: %w", definition.ID, err)
|
|
}
|
|
definition.Endpoint = endpoint
|
|
|
|
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
|
|
}
|