Files
2026-09-11 06:14:38 +02:00

180 lines
5.5 KiB
Go

package server
import (
"encoding/json"
"errors"
"fmt"
"net/http"
"net/url"
"os"
"sort"
"strings"
"github.com/example/ollama-fair-gateway/internal/auth"
"github.com/example/ollama-fair-gateway/internal/config"
"github.com/example/ollama-fair-gateway/internal/proxy"
)
func cloneModelAccess(in config.ModelAccessConfig) config.ModelAccessConfig {
cloneRule := func(r config.ModelAccessRule) config.ModelAccessRule {
r.AllowedModels = append([]string(nil), r.AllowedModels...)
r.DeniedModels = append([]string(nil), r.DeniedModels...)
return r
}
out := config.ModelAccessConfig{Default: cloneRule(in.Default), Tenants: make(map[string]config.ModelAccessRule, len(in.Tenants))}
for name, r := range in.Tenants {
out.Tenants[name] = cloneRule(r)
}
return out
}
func (s *Server) modelAccessSnapshot() config.ModelAccessConfig {
if s == nil {
return config.ModelAccessConfig{Tenants: map[string]config.ModelAccessRule{}}
}
if v := s.modelAccess.Load(); v != nil {
return cloneModelAccess(v.(config.ModelAccessConfig))
}
return cloneModelAccess(s.cfg.ModelAccess)
}
func (s *Server) runtimeTenantModelAccessRule(tenant string) config.ModelAccessRule {
m := s.modelAccessSnapshot()
if r, ok := m.Tenants[tenant]; ok {
if r.Mode == "" {
r.Mode = "allow_all"
}
return r
}
r := m.Default
if r.Mode == "" {
r.Mode = "allow_all"
}
return r
}
func normalizeAccessRule(r config.ModelAccessRule) config.ModelAccessRule {
r.Mode = strings.TrimSpace(r.Mode)
if r.Mode == "" {
r.Mode = "allow_all"
}
r.AllowedModels = cleanStrings(r.AllowedModels)
r.DeniedModels = cleanStrings(r.DeniedModels)
return r
}
func validateModelAccessConfig(m config.ModelAccessConfig) error {
if err := config.ValidateModelAccessRule(normalizeAccessRule(m.Default)); err != nil {
return fmt.Errorf("default model access: %w", err)
}
for tenant, r := range m.Tenants {
if strings.TrimSpace(tenant) == "" {
return errors.New("tenant name is required")
}
if err := config.ValidateModelAccessRule(normalizeAccessRule(r)); err != nil {
return fmt.Errorf("tenant %q model access: %w", tenant, err)
}
}
return nil
}
func (s *Server) storeModelAccess(next config.ModelAccessConfig) error {
if s.configStore == nil {
return errors.New("persistent configuration store unavailable")
}
base := s.cfg
if loader, ok := s.configStore.(configOverrideLoader); ok {
if loaded, err := loader.LoadWithBootstrap(s.cfg); err == nil {
base = loaded
} else if !errors.Is(err, os.ErrNotExist) {
return err
}
}
b, err := json.Marshal(base)
if err != nil {
return err
}
var candidate config.Config
if err := json.Unmarshal(b, &candidate); err != nil {
return err
}
next = cloneModelAccess(next)
if next.Tenants == nil {
next.Tenants = map[string]config.ModelAccessRule{}
}
candidate.ModelAccess = next
if err := validateModelAccessConfig(next); err != nil {
return err
}
if err := s.configStore.Save(&candidate); err != nil {
return err
}
s.modelAccess.Store(cloneModelAccess(next))
return nil
}
func (s *Server) uiModelAccess(w http.ResponseWriter, r *http.Request, actor auth.Identity) {
if r.URL.Path == "/gateway/ui-api/model-access" {
if r.Method != http.MethodGet {
proxy.WriteJSONError(w, http.StatusMethodNotAllowed, "method_not_allowed", "GET required")
return
}
m := s.modelAccessSnapshot()
names := make([]string, 0, len(m.Tenants))
for n := range m.Tenants {
names = append(names, n)
}
sort.Strings(names)
writeJSON(w, http.StatusOK, map[string]any{"default": m.Default, "tenants": m.Tenants, "tenant_names": names, "runtime": true, "persistent": s.configStore != nil})
return
}
raw := strings.TrimPrefix(r.URL.Path, "/gateway/ui-api/model-access/")
tenant, err := url.PathUnescape(raw)
tenant = strings.TrimSpace(tenant)
if err != nil || tenant == "" || strings.Contains(tenant, "/") || len(tenant) > 256 {
proxy.WriteJSONError(w, 400, "bad_tenant", "invalid tenant")
return
}
if s.configStore == nil {
proxy.WriteJSONError(w, 503, "config_store", "persistent configuration store unavailable")
return
}
s.modelAccessMu.Lock()
defer s.modelAccessMu.Unlock()
next := s.modelAccessSnapshot()
switch r.Method {
case http.MethodPut:
var in config.ModelAccessRule
if err := decodeJSON(r, &in, 128<<10); err != nil {
proxy.WriteJSONError(w, 400, "bad_model_access", err.Error())
return
}
in = normalizeAccessRule(in)
if err := config.ValidateModelAccessRule(in); err != nil {
proxy.WriteJSONError(w, 400, "bad_model_access", err.Error())
return
}
next.Tenants[tenant] = in
if err := s.storeModelAccess(next); err != nil {
proxy.WriteJSONError(w, 400, "model_access_store", err.Error())
return
}
s.log.Info("tenant model access saved", "tenant", tenant, "mode", in.Mode, "admin_subject", actor.Subject)
writeJSON(w, 200, map[string]any{"tenant": tenant, "rule": in, "restart_required": false})
case http.MethodDelete:
if _, ok := next.Tenants[tenant]; !ok {
proxy.WriteJSONError(w, 404, "model_access_not_found", fmt.Sprintf("tenant model access %q not found", tenant))
return
}
delete(next.Tenants, tenant)
if err := s.storeModelAccess(next); err != nil {
proxy.WriteJSONError(w, 400, "model_access_store", err.Error())
return
}
s.log.Info("tenant model access reset", "tenant", tenant, "admin_subject", actor.Subject)
writeJSON(w, 200, map[string]any{"deleted": true, "tenant": tenant, "restart_required": false})
default:
proxy.WriteJSONError(w, http.StatusMethodNotAllowed, "method_not_allowed", "PUT or DELETE required")
}
}