156 lines
4.6 KiB
Go
156 lines
4.6 KiB
Go
// 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"
|
|
)
|
|
|
|
// 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,
|
|
})
|
|
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
|
|
}
|
|
|
|
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,
|
|
)
|
|
}
|
|
|
|
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
|
|
}
|