This commit is contained in:
2026-09-11 06:14:38 +02:00
parent bf64652300
commit e581949946
161 changed files with 31126 additions and 1 deletions
+183
View File
@@ -0,0 +1,183 @@
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"
)
type configOverrideLoader interface {
LoadWithBootstrap(*config.Config) (*config.Config, error)
}
func cloneModelAliases(in map[string]config.ModelAliasConfig) map[string]config.ModelAliasConfig {
out := make(map[string]config.ModelAliasConfig, len(in))
for name, a := range in {
a.Models = append([]string(nil), a.Models...)
a.RequiredCapabilities = append([]string(nil), a.RequiredCapabilities...)
if a.Visible != nil {
v := *a.Visible
a.Visible = &v
}
out[name] = a
}
return out
}
func (s *Server) aliasSnapshot() map[string]config.ModelAliasConfig {
if s == nil {
return map[string]config.ModelAliasConfig{}
}
v := s.aliases.Load()
if v == nil {
return cloneModelAliases(s.cfg.ModelAliases)
}
return cloneModelAliases(v.(map[string]config.ModelAliasConfig))
}
func (s *Server) aliasConfig(name string) (config.ModelAliasConfig, bool) {
if s == nil {
return config.ModelAliasConfig{}, false
}
v := s.aliases.Load()
if v == nil {
a, ok := s.cfg.ModelAliases[name]
return a, ok
}
a, ok := v.(map[string]config.ModelAliasConfig)[name]
return a, ok
}
func (s *Server) storeAliases(next map[string]config.ModelAliasConfig) 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
}
candidate.ModelAliases = cloneModelAliases(next)
if err := validateAliasSet(candidate.ModelAliases); err != nil {
return err
}
if err := s.configStore.Save(&candidate); err != nil {
return err
}
// Publish only after persistence succeeds. Readers never observe a map that
// would be lost on restart, and the stored map is never mutated in place.
s.aliases.Store(cloneModelAliases(next))
return nil
}
func validateAliasSet(aliases map[string]config.ModelAliasConfig) error {
for name, a := range aliases {
if strings.TrimSpace(name) == "" {
return errors.New("model alias name is required")
}
if len(a.Models) == 0 {
return fmt.Errorf("model alias %q requires at least one model", name)
}
for _, model := range a.Models {
if strings.TrimSpace(model) == "" {
return fmt.Errorf("model alias %q contains an empty model", name)
}
}
}
return nil
}
func (s *Server) uiModelAliases(w http.ResponseWriter, r *http.Request, actor auth.Identity) {
if r.URL.Path == "/gateway/ui-api/model-aliases" {
if r.Method != http.MethodGet {
proxy.WriteJSONError(w, http.StatusMethodNotAllowed, "method_not_allowed", "GET required")
return
}
aliases := s.aliasSnapshot()
names := make([]string, 0, len(aliases))
for name := range aliases {
names = append(names, name)
}
sort.Strings(names)
writeJSON(w, http.StatusOK, map[string]any{"aliases": aliases, "names": names, "runtime": true, "persistent": s.configStore != nil})
return
}
raw := strings.TrimPrefix(r.URL.Path, "/gateway/ui-api/model-aliases/")
name, err := url.PathUnescape(raw)
name = strings.TrimSpace(name)
if err != nil || name == "" || strings.Contains(name, "/") || len(name) > 256 {
proxy.WriteJSONError(w, http.StatusBadRequest, "bad_alias", "invalid model alias name")
return
}
if s.configStore == nil {
proxy.WriteJSONError(w, http.StatusServiceUnavailable, "config_store", "persistent configuration store unavailable")
return
}
s.aliasMu.Lock()
defer s.aliasMu.Unlock()
next := s.aliasSnapshot()
switch r.Method {
case http.MethodPut:
var in config.ModelAliasConfig
if err := decodeJSON(r, &in, 128<<10); err != nil {
proxy.WriteJSONError(w, http.StatusBadRequest, "bad_alias", err.Error())
return
}
in.Models = cleanStrings(in.Models)
in.RequiredCapabilities = cleanStrings(in.RequiredCapabilities)
next[name] = in
if err := s.storeAliases(next); err != nil {
proxy.WriteJSONError(w, http.StatusBadRequest, "alias_store", err.Error())
return
}
s.log.Info("model alias saved", "alias", name, "models", len(in.Models), "admin_subject", actor.Subject, "admin_auth_type", actor.AuthType)
writeJSON(w, http.StatusOK, map[string]any{"name": name, "alias": in, "restart_required": false})
case http.MethodDelete:
if _, ok := next[name]; !ok {
proxy.WriteJSONError(w, http.StatusNotFound, "alias_not_found", fmt.Sprintf("model alias %q not found", name))
return
}
delete(next, name)
if err := s.storeAliases(next); err != nil {
proxy.WriteJSONError(w, http.StatusBadRequest, "alias_store", err.Error())
return
}
s.log.Info("model alias deleted", "alias", name, "admin_subject", actor.Subject, "admin_auth_type", actor.AuthType)
writeJSON(w, http.StatusOK, map[string]any{"deleted": true, "name": name, "restart_required": false})
default:
proxy.WriteJSONError(w, http.StatusMethodNotAllowed, "method_not_allowed", "PUT or DELETE required")
}
}
func cleanStrings(in []string) []string {
out := make([]string, 0, len(in))
for _, v := range in {
v = strings.TrimSpace(v)
if v != "" {
out = append(out, v)
}
}
return out
}
+298
View File
@@ -0,0 +1,298 @@
package server
import (
"context"
"encoding/json"
"errors"
"io"
"net/http"
"os"
"sort"
"strconv"
"strings"
"github.com/example/ollama-fair-gateway/internal/auth"
"github.com/example/ollama-fair-gateway/internal/batch"
)
type batchCreateRequest struct {
Path string `json:"path"`
Body json.RawMessage `json:"body"`
}
func (s *Server) batchAPI(w http.ResponseWriter, r *http.Request, id auth.Identity) {
if s.batchJobs == nil || !s.batchJobs.Enabled() {
writeProtocolError(w, r, http.StatusNotFound, "batch_disabled", "durable batch jobs are disabled")
return
}
base := "/gateway/v1/batches"
rest := strings.TrimPrefix(r.URL.Path, base)
if rest == "" || rest == "/" {
switch r.Method {
case http.MethodGet:
writeJSON(w, http.StatusOK, map[string]any{"jobs": s.batchJobs.List(id.Tenant, id.Actor(), false)})
case http.MethodPost:
s.batchCreate(w, r, id)
default:
writeProtocolError(w, r, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
}
return
}
parts := strings.Split(strings.Trim(rest, "/"), "/")
if len(parts) == 0 || parts[0] == "" {
writeProtocolError(w, r, http.StatusNotFound, "not_found", "batch job not found")
return
}
jobID := parts[0]
if len(parts) == 1 && r.Method == http.MethodGet {
j, ok := s.batchJobs.Get(jobID, id.Tenant, id.Actor(), false)
if !ok {
writeProtocolError(w, r, http.StatusNotFound, "batch_not_found", "batch job not found")
return
}
writeJSON(w, http.StatusOK, j)
return
}
if len(parts) == 2 && parts[1] == "output" && r.Method == http.MethodGet {
s.batchOutput(w, r, id, jobID, false)
return
}
if len(parts) == 2 && r.Method == http.MethodPost {
var (
j batch.Job
err error
)
switch parts[1] {
case "pause":
j, err = s.batchJobs.Pause(jobID, id.Tenant, id.Actor(), false)
case "resume":
j, err = s.batchJobs.Resume(jobID, id.Tenant, id.Actor(), false)
case "cancel":
j, err = s.batchJobs.Cancel(jobID, id.Tenant, id.Actor(), false)
default:
writeProtocolError(w, r, http.StatusNotFound, "not_found", "unknown batch operation")
return
}
if err != nil {
s.writeBatchError(w, r, err)
return
}
writeJSON(w, http.StatusOK, j)
return
}
writeProtocolError(w, r, http.StatusNotFound, "not_found", "unknown batch endpoint")
}
func (s *Server) batchCreate(w http.ResponseWriter, r *http.Request, id auth.Identity) {
limit := s.cfg.BatchJobs.MaxInputBytes + 64<<10
if limit <= 0 {
limit = 16 << 20
}
var req batchCreateRequest
if err := decodeJSON(r, &req, limit); err != nil {
writeProtocolError(w, r, http.StatusBadRequest, "bad_batch", err.Error())
return
}
if !s.isCompute(http.MethodPost, req.Path) {
writeProtocolError(w, r, http.StatusBadRequest, "bad_batch_path", "batch path must be a configured compute POST endpoint")
return
}
if len(req.Body) == 0 || string(req.Body) == "null" || !json.Valid(req.Body) {
writeProtocolError(w, r, http.StatusBadRequest, "bad_batch_body", "batch body must contain a valid JSON request body")
return
}
model := modelFromBody(req.Body)
if _, _, err := s.resolveModel(r.Context(), id, model); err != nil {
if errors.Is(err, ErrModelAccessDenied) {
writeProtocolError(w, r, http.StatusForbidden, "model_access_denied", err.Error())
} else {
writeProtocolError(w, r, http.StatusNotFound, "model_alias_unavailable", err.Error())
}
return
}
j, err := s.batchJobs.Create(batchIdentitySnapshot(id), req.Path, model, req.Body)
if err != nil {
s.writeBatchError(w, r, err)
return
}
w.Header().Set("Location", "/gateway/v1/batches/"+j.ID)
writeJSON(w, http.StatusAccepted, j)
}
func (s *Server) batchOutput(w http.ResponseWriter, r *http.Request, id auth.Identity, jobID string, all bool) {
f, j, err := s.batchJobs.OpenOutput(jobID, id.Tenant, id.Actor(), all)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
writeProtocolError(w, r, http.StatusConflict, "batch_output_unavailable", "batch output is not available yet")
return
}
s.writeBatchError(w, r, err)
return
}
defer f.Close()
ct := strings.TrimSpace(j.ResponseContentType)
if ct == "" {
ct = "application/octet-stream"
}
w.Header().Set("Content-Type", ct)
w.Header().Set("Cache-Control", "no-store")
if st, err := f.Stat(); err == nil {
w.Header().Set("Content-Length", strconv.FormatInt(st.Size(), 10))
}
w.WriteHeader(http.StatusOK)
_, _ = io.Copy(w, f)
}
func (s *Server) writeBatchError(w http.ResponseWriter, r *http.Request, err error) {
switch {
case errors.Is(err, batch.ErrDisabled):
writeProtocolError(w, r, http.StatusNotFound, "batch_disabled", err.Error())
case errors.Is(err, batch.ErrNotFound):
writeProtocolError(w, r, http.StatusNotFound, "batch_not_found", "batch job not found")
case errors.Is(err, batch.ErrInvalidState):
writeProtocolError(w, r, http.StatusConflict, "batch_state", err.Error())
case errors.Is(err, batch.ErrFull):
writeProtocolError(w, r, http.StatusTooManyRequests, "batch_full", err.Error())
case errors.Is(err, batch.ErrInputTooLarge):
writeProtocolError(w, r, http.StatusRequestEntityTooLarge, "batch_input_too_large", err.Error())
default:
writeProtocolError(w, r, http.StatusInternalServerError, "batch_error", err.Error())
}
}
// ExecuteBatch is the runner used by the durable batch manager. The request is
// replayed directly into the normal compute path with the original identity
// metadata but the batch service class, so quotas, ACLs, scheduling, routing,
// metering, alerts and OpenTelemetry stay consistent with interactive traffic.
func (s *Server) ExecuteBatch(ctx context.Context, j batch.Job, input io.Reader, output io.Writer) batch.RunResult {
id := auth.Identity{
Tenant: j.Identity.Tenant,
Subject: j.Identity.Subject,
Application: j.Identity.Application,
AuthType: j.Identity.AuthType,
ClientIP: j.Identity.ClientIP,
Scopes: make(map[string]bool, len(j.Identity.Scopes)),
ModelACLSet: j.Identity.ModelACLSet,
ModelAccess: j.Identity.ModelAccess,
ServiceClass: "batch",
}
for _, scope := range j.Identity.Scopes {
id.Scopes[scope] = true
}
r, err := http.NewRequestWithContext(ctx, http.MethodPost, "http://gateway.local"+j.Path, input)
if err != nil {
return batch.RunResult{Error: err.Error()}
}
r.Header.Set("Content-Type", "application/json")
r.RemoteAddr = "127.0.0.1:0"
rw := &batchResponseWriter{header: make(http.Header), out: output}
s.forward(rw, r, id)
status := rw.status
if status == 0 {
status = http.StatusOK
}
return batch.RunResult{HTTPStatus: status, ResponseContentType: rw.header.Get("Content-Type"), RequestID: rw.header.Get("X-Request-ID")}
}
type batchResponseWriter struct {
header http.Header
out io.Writer
status int
}
func (w *batchResponseWriter) Header() http.Header { return w.header }
func (w *batchResponseWriter) WriteHeader(status int) {
if w.status == 0 {
w.status = status
}
}
func (w *batchResponseWriter) Write(p []byte) (int, error) {
if w.status == 0 {
w.status = http.StatusOK
}
return w.out.Write(p)
}
func (w *batchResponseWriter) Flush() {}
func batchIdentitySnapshot(id auth.Identity) batch.IdentitySnapshot {
scopes := make([]string, 0, len(id.Scopes))
for scope, ok := range id.Scopes {
if ok {
scopes = append(scopes, scope)
}
}
sort.Strings(scopes)
return batch.IdentitySnapshot{Tenant: id.Tenant, Subject: id.Subject, Actor: id.Actor(), Application: id.Application, AuthType: id.AuthType, ClientIP: id.ClientIP, Scopes: scopes, ModelACLSet: id.ModelACLSet, ModelAccess: id.ModelAccess}
}
func (s *Server) uiBatchJobs(w http.ResponseWriter, r *http.Request, id auth.Identity) {
base := "/gateway/ui-api/batches"
if s.batchJobs == nil || !s.batchJobs.Enabled() {
if r.URL.Path == base && r.Method == http.MethodGet {
writeJSON(w, http.StatusOK, map[string]any{"enabled": false, "jobs": []batch.Job{}})
return
}
writeProtocolError(w, r, http.StatusNotFound, "batch_disabled", "durable batch jobs are disabled")
return
}
rest := strings.TrimPrefix(r.URL.Path, base)
if rest == "" || rest == "/" {
if r.Method != http.MethodGet {
writeProtocolError(w, r, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
return
}
writeJSON(w, http.StatusOK, map[string]any{
"enabled": true,
"jobs": s.batchJobs.List("", "", true),
"retention": s.cfg.BatchJobs.Retention.Value().String(),
"max_jobs": s.cfg.BatchJobs.MaxJobs,
"max_concurrent": s.cfg.BatchJobs.MaxConcurrent,
"max_input_bytes": s.cfg.BatchJobs.MaxInputBytes,
})
return
}
parts := strings.Split(strings.Trim(rest, "/"), "/")
if len(parts) == 0 || parts[0] == "" {
writeProtocolError(w, r, http.StatusNotFound, "batch_not_found", "batch job not found")
return
}
jobID := parts[0]
if len(parts) == 1 && r.Method == http.MethodGet {
j, ok := s.batchJobs.Get(jobID, "", "", true)
if !ok {
writeProtocolError(w, r, http.StatusNotFound, "batch_not_found", "batch job not found")
return
}
writeJSON(w, http.StatusOK, j)
return
}
if len(parts) == 2 && parts[1] == "output" && r.Method == http.MethodGet {
s.batchOutput(w, r, id, jobID, true)
return
}
if len(parts) == 2 && r.Method == http.MethodPost {
var (
j batch.Job
err error
)
switch parts[1] {
case "pause":
j, err = s.batchJobs.Pause(jobID, "", "", true)
case "resume":
j, err = s.batchJobs.Resume(jobID, "", "", true)
case "cancel":
j, err = s.batchJobs.Cancel(jobID, "", "", true)
default:
writeProtocolError(w, r, http.StatusNotFound, "not_found", "unknown batch operation")
return
}
if err != nil {
s.writeBatchError(w, r, err)
return
}
s.log.Info("durable batch control", "batch_id", jobID, "action", parts[1], "admin_subject", id.Subject, "admin_auth_type", id.AuthType)
writeJSON(w, http.StatusOK, j)
return
}
writeProtocolError(w, r, http.StatusNotFound, "not_found", "unknown batch UI endpoint")
}
+165
View File
@@ -0,0 +1,165 @@
package server
import (
"bytes"
"context"
"encoding/json"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"path/filepath"
"strings"
"testing"
"time"
"github.com/example/ollama-fair-gateway/internal/auth"
"github.com/example/ollama-fair-gateway/internal/batch"
"github.com/example/ollama-fair-gateway/internal/config"
"github.com/example/ollama-fair-gateway/internal/cost"
"github.com/example/ollama-fair-gateway/internal/metrics"
px "github.com/example/ollama-fair-gateway/internal/proxy"
"github.com/example/ollama-fair-gateway/internal/quota"
"github.com/example/ollama-fair-gateway/internal/scheduler"
"github.com/example/ollama-fair-gateway/internal/usage"
"github.com/example/ollama-fair-gateway/internal/worker"
)
func TestDurableBatchEndToEndThroughGatewayPipeline(t *testing.T) {
var upstreamBody string
backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/ps":
w.Header().Set("Content-Type", "application/json")
_, _ = io.WriteString(w, `{"models":[{"name":"qwen3:8b","model":"qwen3:8b"}]}`)
case "/api/tags":
w.Header().Set("Content-Type", "application/json")
_, _ = io.WriteString(w, `{"models":[{"name":"qwen3:8b"}]}`)
case "/api/chat":
b, _ := io.ReadAll(r.Body)
upstreamBody = string(b)
w.Header().Set("Content-Type", "application/x-ndjson")
_, _ = io.WriteString(w, "{\"message\":{\"content\":\"batch-ok\"},\"done\":true,\"prompt_eval_count\":3,\"eval_count\":2}\n")
default:
w.WriteHeader(http.StatusNotFound)
}
}))
defer backend.Close()
cfg := &config.Config{
Server: config.ServerConfig{MaxBodyBytes: 1 << 20, MaxRequestDuration: config.Duration(time.Minute)},
Auth: config.AuthConfig{APIKeys: []config.APIKeyConfig{{Name: "client", Key: "batch-key", Tenant: "team-a", Subject: "alice"}}},
Scheduler: config.SchedulerConfig{GlobalConcurrency: 2, MaxQueue: 32, MaxQueuePerActor: 8, QueueTimeout: config.Duration(time.Second), DefaultTenantWeight: 1, DefaultActorWeight: 1, ComputePaths: []string{"/api/chat"}},
Cost: config.CostConfig{Default: config.ModelRate{InputCreditsPer1K: 1, OutputCreditsPer1K: 3}, DefaultMaxOutputTokens: 16},
Workers: []config.WorkerConfig{{Name: "w", URL: backend.URL, MaxConcurrent: 2, HealthInterval: config.Duration(time.Hour)}},
ServiceClasses: config.ServiceClassesConfig{Default: "interactive", Classes: map[string]config.ServiceClassConfig{
"interactive": {Weight: 1, MaxQueueWait: config.Duration(time.Second), MaxConcurrent: 2},
"batch": {Weight: 0.25, MaxQueueWait: config.Duration(time.Second), MaxConcurrent: 1},
}},
BatchJobs: config.BatchJobsConfig{Enabled: true, Retention: config.Duration(time.Hour), MaxJobs: 100, MaxConcurrent: 1, MaxInputBytes: 1 << 20},
}
a, err := auth.New(context.Background(), cfg.Auth)
if err != nil {
t.Fatal(err)
}
wp := worker.New(cfg.Workers, "w")
root, cancel := context.WithCancel(context.Background())
defer cancel()
wp.Start(root)
rec, err := usage.New("", 100, time.Second, nil)
if err != nil {
t.Fatal(err)
}
defer rec.Close()
dir := t.TempDir()
bm, err := batch.New(cfg.BatchJobs, filepath.Join(dir, "batch-jobs.json"), filepath.Join(dir, "batch"))
if err != nil {
t.Fatal(err)
}
sv := New(cfg, Dependencies{Auth: a, Scheduler: scheduler.NewLocal(2, 32, 8), Quota: quota.Disabled{}, Estimator: cost.New(cfg.Cost), Workers: wp, Proxy: px.New(), Usage: rec, Metrics: metrics.New(), Logger: slog.Default(), BatchJobs: bm})
bm.Start(root, sv.ExecuteBatch)
front := httptest.NewServer(sv.Handler())
defer front.Close()
payload := []byte(`{"path":"/api/chat","body":{"model":"qwen3:8b","messages":[{"role":"user","content":"run batch"}]}}`)
req, _ := http.NewRequest(http.MethodPost, front.URL+"/gateway/v1/batches", bytes.NewReader(payload))
req.Header.Set("Authorization", "Bearer batch-key")
req.Header.Set("Content-Type", "application/json")
resp, err := http.DefaultClient.Do(req)
if err != nil {
t.Fatal(err)
}
body, _ := io.ReadAll(resp.Body)
resp.Body.Close()
if resp.StatusCode != http.StatusAccepted {
t.Fatalf("create status=%d body=%s", resp.StatusCode, body)
}
var created batch.Job
if err := json.Unmarshal(body, &created); err != nil {
t.Fatal(err)
}
if created.ID == "" || created.ServiceClass != "batch" || created.Identity.Tenant != "team-a" || created.Identity.Actor != "alice" {
t.Fatalf("created=%#v", created)
}
if !strings.HasPrefix(resp.Header.Get("Location"), "/gateway/v1/batches/") {
t.Fatalf("location=%q", resp.Header.Get("Location"))
}
var completed batch.Job
deadline := time.Now().Add(3 * time.Second)
for time.Now().Before(deadline) {
req, _ = http.NewRequest(http.MethodGet, front.URL+"/gateway/v1/batches/"+created.ID, nil)
req.Header.Set("Authorization", "Bearer batch-key")
resp, err = http.DefaultClient.Do(req)
if err != nil {
t.Fatal(err)
}
body, _ = io.ReadAll(resp.Body)
resp.Body.Close()
if resp.StatusCode != http.StatusOK {
t.Fatalf("get status=%d body=%s", resp.StatusCode, body)
}
if err := json.Unmarshal(body, &completed); err != nil {
t.Fatal(err)
}
if completed.State == batch.StateCompleted {
break
}
if completed.State == batch.StateFailed || completed.State == batch.StateCancelled {
t.Fatalf("unexpected terminal state: %#v", completed)
}
time.Sleep(10 * time.Millisecond)
}
if completed.State != batch.StateCompleted || completed.HTTPStatus != http.StatusOK || completed.ExecutionRequestID == "" || completed.OutputRef == "" {
t.Fatalf("completed=%#v", completed)
}
if !strings.Contains(upstreamBody, `"model":"qwen3:8b"`) || !strings.Contains(upstreamBody, `"run batch"`) {
t.Fatalf("upstream body=%s", upstreamBody)
}
req, _ = http.NewRequest(http.MethodGet, front.URL+"/gateway/v1/batches/"+created.ID+"/output", nil)
req.Header.Set("Authorization", "Bearer batch-key")
resp, err = http.DefaultClient.Do(req)
if err != nil {
t.Fatal(err)
}
body, _ = io.ReadAll(resp.Body)
resp.Body.Close()
if resp.StatusCode != http.StatusOK || !strings.Contains(string(body), `"batch-ok"`) {
t.Fatalf("output status=%d body=%s", resp.StatusCode, body)
}
events := rec.Recent(10)
found := false
for _, e := range events {
if e.ID == completed.ExecutionRequestID {
found = true
if e.ServiceClass != "batch" || e.Tenant != "team-a" || e.Actor != "alice" || e.Usage.PromptTokens != 3 || e.Usage.CompletionTokens != 2 {
t.Fatalf("usage event=%#v", e)
}
}
}
if !found {
t.Fatalf("execution request %s missing from usage: %#v", completed.ExecutionRequestID, events)
}
}
+185
View File
@@ -0,0 +1,185 @@
package server
import (
"bytes"
"encoding/json"
"errors"
"fmt"
"strings"
"time"
"github.com/example/ollama-fair-gateway/internal/auth"
"github.com/example/ollama-fair-gateway/internal/conversation"
)
type responseConversationPlan struct {
Store bool
RequestItems []any
}
// prepareResponseConversation expands previous_response_id into a flattened
// Responses API input history when the optional encrypted conversation store
// is enabled. When disabled the request is left untouched, preserving the
// gateway's content-free default behavior and any upstream-native semantics.
func (s *Server) prepareResponseConversation(body []byte, id auth.Identity) ([]byte, *responseConversationPlan, error) {
if s.conversations == nil || !s.conversations.Enabled() || len(body) == 0 {
return body, nil, nil
}
var req map[string]any
dec := json.NewDecoder(bytes.NewReader(body))
dec.UseNumber()
if err := dec.Decode(&req); err != nil {
return nil, nil, fmt.Errorf("decode Responses request: %w", err)
}
store := true
if v, ok := req["store"].(bool); ok && !v {
store = false
}
current, err := responseInputItems(req["input"])
if err != nil {
return nil, nil, err
}
items := current
if rawPrev, ok := req["previous_response_id"]; ok {
prev, ok := rawPrev.(string)
prev = strings.TrimSpace(prev)
if !ok || prev == "" {
return nil, nil, errors.New("previous_response_id must be a non-empty string")
}
parent, found, err := s.conversations.Get(prev, id.Tenant, id.Actor())
if err != nil {
return nil, nil, fmt.Errorf("load previous response: %w", err)
}
if !found {
return nil, nil, fmt.Errorf("previous_response_id %q was not found for this identity", prev)
}
var prior []any
if err := json.Unmarshal(parent.Context, &prior); err != nil {
return nil, nil, fmt.Errorf("decode stored conversation context: %w", err)
}
items = make([]any, 0, len(prior)+len(current))
items = append(items, prior...)
items = append(items, current...)
req["input"] = items
delete(req, "previous_response_id")
}
out, err := json.Marshal(req)
if err != nil {
return nil, nil, fmt.Errorf("encode expanded Responses request: %w", err)
}
return out, &responseConversationPlan{Store: store, RequestItems: append([]any(nil), items...)}, nil
}
func responseInputItems(v any) ([]any, error) {
if v == nil {
return nil, nil
}
switch x := v.(type) {
case string:
return []any{map[string]any{"role": "user", "content": x}}, nil
case []any:
return append([]any(nil), x...), nil
case map[string]any:
return []any{x}, nil
default:
return nil, errors.New("Responses input must be a string, object, or array")
}
}
func (s *Server) persistResponseConversation(plan *responseConversationPlan, captured []byte, truncated bool, id auth.Identity, model string) {
if plan == nil || !plan.Store || s.conversations == nil || !s.conversations.Enabled() {
return
}
if truncated {
s.log.Warn("conversation response not stored because capture exceeded limit", "tenant", id.Tenant, "actor", id.Actor(), "model", model)
return
}
responseID, output, err := parseResponsesOutput(captured)
if err != nil {
s.log.Warn("conversation response not stored", "tenant", id.Tenant, "actor", id.Actor(), "model", model, "error", err)
return
}
if responseID == "" {
s.log.Warn("conversation response not stored because response id was absent", "tenant", id.Tenant, "actor", id.Actor(), "model", model)
return
}
ctx := make([]any, 0, len(plan.RequestItems)+len(output))
ctx = append(ctx, plan.RequestItems...)
ctx = append(ctx, output...)
raw, err := json.Marshal(ctx)
if err != nil {
s.log.Warn("conversation context encode failed", "response_id", responseID, "error", err)
return
}
if err := s.conversations.Put(conversation.Entry{ID: responseID, Tenant: id.Tenant, Actor: id.Actor(), Model: model, CreatedAt: time.Now().UTC(), Context: raw}); err != nil {
s.log.Warn("conversation persistence failed", "response_id", responseID, "tenant", id.Tenant, "actor", id.Actor(), "error", err)
}
}
func parseResponsesOutput(body []byte) (string, []any, error) {
body = bytes.TrimSpace(body)
if len(body) == 0 {
return "", nil, errors.New("empty Responses response")
}
if body[0] == '{' {
var v map[string]any
if err := json.Unmarshal(body, &v); err != nil {
return "", nil, err
}
return responseIDAndOutput(v)
}
var responseID string
var completedOutput []any
var doneOutput []any
for _, rawLine := range bytes.Split(body, []byte{'\n'}) {
line := bytes.TrimSpace(rawLine)
if !bytes.HasPrefix(line, []byte("data:")) {
continue
}
data := bytes.TrimSpace(bytes.TrimPrefix(line, []byte("data:")))
if len(data) == 0 || bytes.Equal(data, []byte("[DONE]")) {
continue
}
var ev map[string]any
if json.Unmarshal(data, &ev) != nil {
continue
}
typ, _ := ev["type"].(string)
if resp, ok := ev["response"].(map[string]any); ok {
id, out, _ := responseIDAndOutput(resp)
if id != "" {
responseID = id
}
if typ == "response.completed" && out != nil {
completedOutput = out
}
}
if typ == "response.output_item.done" {
if item, ok := ev["item"].(map[string]any); ok {
doneOutput = append(doneOutput, item)
}
}
}
if responseID == "" {
return "", nil, errors.New("stream did not contain a response id")
}
if completedOutput != nil {
return responseID, completedOutput, nil
}
if len(doneOutput) > 0 {
return responseID, doneOutput, nil
}
return "", nil, errors.New("stream did not contain completed output items")
}
func responseIDAndOutput(v map[string]any) (string, []any, error) {
id, _ := v["id"].(string)
out, _ := v["output"].([]any)
if id == "" {
return "", nil, errors.New("response id is missing")
}
if out == nil {
return id, nil, errors.New("response output is missing")
}
return id, out, nil
}
+222
View File
@@ -0,0 +1,222 @@
package server
import (
"context"
"encoding/json"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"path/filepath"
"strings"
"testing"
"time"
"github.com/example/ollama-fair-gateway/internal/auth"
"github.com/example/ollama-fair-gateway/internal/config"
"github.com/example/ollama-fair-gateway/internal/conversation"
"github.com/example/ollama-fair-gateway/internal/cost"
"github.com/example/ollama-fair-gateway/internal/metrics"
"github.com/example/ollama-fair-gateway/internal/proxy"
"github.com/example/ollama-fair-gateway/internal/quota"
"github.com/example/ollama-fair-gateway/internal/scheduler"
"github.com/example/ollama-fair-gateway/internal/usage"
"github.com/example/ollama-fair-gateway/internal/worker"
)
func newConversationTestServer(t *testing.T) (*Server, *conversation.Store) {
t.Helper()
cc := config.ConversationsConfig{Enabled: true, EncryptionKey: strings.Repeat("k", 32), Retention: config.Duration(time.Hour), MaxEntries: 100, MaxContentBytes: 1 << 20}
store, err := conversation.New(cc, filepath.Join(t.TempDir(), "conversations.enc.json"))
if err != nil {
t.Fatal(err)
}
return &Server{cfg: &config.Config{Conversations: cc}, conversations: store, log: slog.New(slog.NewTextHandler(io.Discard, nil))}, store
}
func TestPrepareResponseConversationExpandsPreviousResponse(t *testing.T) {
s, store := newConversationTestServer(t)
id := auth.Identity{Tenant: "tenant-a", Subject: "user-a", AuthType: "oidc"}
prior := json.RawMessage(`[{"role":"user","content":"first"},{"type":"message","role":"assistant","content":[{"type":"output_text","text":"answer"}]}]`)
if err := store.Put(conversation.Entry{ID: "resp_prev", Tenant: id.Tenant, Actor: id.Actor(), Context: prior}); err != nil {
t.Fatal(err)
}
body, plan, err := s.prepareResponseConversation([]byte(`{"model":"qwen3:8b","previous_response_id":"resp_prev","input":"follow up"}`), id)
if err != nil {
t.Fatal(err)
}
if plan == nil || !plan.Store || len(plan.RequestItems) != 3 {
t.Fatalf("plan=%#v", plan)
}
var req map[string]any
if err := json.Unmarshal(body, &req); err != nil {
t.Fatal(err)
}
if _, exists := req["previous_response_id"]; exists {
t.Fatal("previous_response_id must not be forwarded after expansion")
}
items, ok := req["input"].([]any)
if !ok || len(items) != 3 {
t.Fatalf("expanded input=%#v", req["input"])
}
last := items[2].(map[string]any)
if last["role"] != "user" || last["content"] != "follow up" {
t.Fatalf("last item=%#v", last)
}
}
func TestPrepareResponseConversationDoesNotCrossIdentityBoundary(t *testing.T) {
s, store := newConversationTestServer(t)
if err := store.Put(conversation.Entry{ID: "resp_prev", Tenant: "tenant-a", Actor: "user-a", Context: json.RawMessage(`[]`)}); err != nil {
t.Fatal(err)
}
_, _, err := s.prepareResponseConversation([]byte(`{"previous_response_id":"resp_prev","input":"x"}`), auth.Identity{Tenant: "tenant-a", Subject: "user-b", AuthType: "oidc"})
if err == nil || !strings.Contains(err.Error(), "not found") {
t.Fatalf("err=%v", err)
}
}
func TestPrepareResponseConversationRespectsStoreFalse(t *testing.T) {
s, _ := newConversationTestServer(t)
_, plan, err := s.prepareResponseConversation([]byte(`{"input":"x","store":false}`), auth.Identity{Tenant: "t", Subject: "u", AuthType: "oidc"})
if err != nil {
t.Fatal(err)
}
if plan == nil || plan.Store {
t.Fatalf("plan=%#v", plan)
}
}
func TestParseResponsesOutputJSONAndSSE(t *testing.T) {
id, out, err := parseResponsesOutput([]byte(`{"id":"resp_1","output":[{"type":"message","role":"assistant"}]}`))
if err != nil || id != "resp_1" || len(out) != 1 {
t.Fatalf("json: id=%q out=%#v err=%v", id, out, err)
}
sse := "event: response.created\ndata: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_2\",\"output\":[]}}\n\n" +
"event: response.output_item.done\ndata: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"message\",\"role\":\"assistant\"}}\n\n" +
"event: response.completed\ndata: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_2\",\"output\":[{\"type\":\"message\",\"role\":\"assistant\"}]}}\n\n"
id, out, err = parseResponsesOutput([]byte(sse))
if err != nil || id != "resp_2" || len(out) != 1 {
t.Fatalf("sse: id=%q out=%#v err=%v", id, out, err)
}
}
func TestResponsesPreviousResponseEndToEnd(t *testing.T) {
var responseCalls int
backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/ps":
w.Header().Set("Content-Type", "application/json")
io.WriteString(w, `{"models":[{"name":"qwen3:8b","model":"qwen3:8b"}]}`)
case "/api/tags":
w.Header().Set("Content-Type", "application/json")
io.WriteString(w, `{"models":[{"name":"qwen3:8b","model":"qwen3:8b"}]}`)
case "/v1/responses":
responseCalls++
var req map[string]any
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
t.Errorf("decode upstream request: %v", err)
w.WriteHeader(400)
return
}
if _, exists := req["previous_response_id"]; exists {
t.Errorf("previous_response_id leaked upstream on call %d", responseCalls)
}
if responseCalls == 1 {
if req["input"] != "first" {
t.Errorf("first input=%#v", req["input"])
}
w.Header().Set("Content-Type", "application/json")
io.WriteString(w, `{"id":"resp_1","output":[{"type":"message","role":"assistant","content":[{"type":"output_text","text":"one"}]}],"usage":{"input_tokens":1,"output_tokens":1}}`)
return
}
items, ok := req["input"].([]any)
if !ok || len(items) != 3 {
t.Errorf("expanded second input=%#v", req["input"])
} else {
first := items[0].(map[string]any)
assistant := items[1].(map[string]any)
last := items[2].(map[string]any)
if first["role"] != "user" || first["content"] != "first" {
t.Errorf("first history item=%#v", first)
}
if assistant["role"] != "assistant" {
t.Errorf("assistant history item=%#v", assistant)
}
if last["role"] != "user" || last["content"] != "second" {
t.Errorf("last history item=%#v", last)
}
}
w.Header().Set("Content-Type", "application/json")
io.WriteString(w, `{"id":"resp_2","output":[{"type":"message","role":"assistant","content":[{"type":"output_text","text":"two"}]}],"usage":{"input_tokens":3,"output_tokens":1}}`)
default:
w.WriteHeader(http.StatusNotFound)
}
}))
defer backend.Close()
cc := config.ConversationsConfig{Enabled: true, EncryptionKey: strings.Repeat("k", 32), Retention: config.Duration(time.Hour), MaxEntries: 100, MaxContentBytes: 1 << 20}
conv, err := conversation.New(cc, filepath.Join(t.TempDir(), "conversations.enc.json"))
if err != nil {
t.Fatal(err)
}
cfg := &config.Config{
Server: config.ServerConfig{MaxBodyBytes: 1 << 20, MaxRequestDuration: config.Duration(time.Minute)},
Auth: config.AuthConfig{IPBypass: []config.IPBypassConfig{{CIDRs: []string{"127.0.0.1/32"}, Tenant: "test", Subject: "u"}}},
Scheduler: config.SchedulerConfig{GlobalConcurrency: 1, MaxQueue: 8, MaxQueuePerActor: 8, QueueTimeout: config.Duration(time.Second), DefaultTenantWeight: 1, DefaultActorWeight: 1, ComputePaths: []string{"/v1/responses"}},
Cost: config.CostConfig{Default: config.ModelRate{InputCreditsPer1K: 1, OutputCreditsPer1K: 3}, DefaultMaxOutputTokens: 16},
Workers: []config.WorkerConfig{{Name: "w", URL: backend.URL, MaxConcurrent: 1, HealthInterval: config.Duration(time.Hour)}},
ModelCapabilities: config.ModelCapabilitiesConfig{Mode: "off", ContextGuard: "off"},
Conversations: cc,
}
a, err := auth.New(context.Background(), cfg.Auth)
if err != nil {
t.Fatal(err)
}
wp := worker.New(cfg.Workers, "")
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
wp.SetModelCapabilitiesConfig(cfg.ModelCapabilities)
wp.Start(ctx)
rec, _ := usage.New("", 100, time.Second, nil)
defer rec.Close()
sv := New(cfg, Dependencies{Auth: a, Scheduler: scheduler.NewLocal(1, 8, 8), Quota: quota.Disabled{}, Estimator: cost.New(cfg.Cost), Workers: wp, Proxy: proxy.New(), Usage: rec, Metrics: metrics.New(), Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), Conversations: conv})
front := httptest.NewServer(sv.Handler())
defer front.Close()
resp, err := http.Post(front.URL+"/v1/responses", "application/json", strings.NewReader(`{"model":"qwen3:8b","input":"first"}`))
if err != nil {
t.Fatal(err)
}
firstBody, _ := io.ReadAll(resp.Body)
resp.Body.Close()
if resp.StatusCode != 200 || !strings.Contains(string(firstBody), `"id":"resp_1"`) {
t.Fatalf("first status=%d body=%s", resp.StatusCode, firstBody)
}
if _, ok, err := conv.Get("resp_1", "test", "u"); err != nil || !ok {
t.Fatalf("first response not stored: ok=%v err=%v", ok, err)
}
resp, err = http.Post(front.URL+"/v1/responses", "application/json", strings.NewReader(`{"model":"qwen3:8b","previous_response_id":"resp_1","input":"second"}`))
if err != nil {
t.Fatal(err)
}
secondBody, _ := io.ReadAll(resp.Body)
resp.Body.Close()
if resp.StatusCode != 200 || !strings.Contains(string(secondBody), `"id":"resp_2"`) {
t.Fatalf("second status=%d body=%s", resp.StatusCode, secondBody)
}
stored, ok, err := conv.Get("resp_2", "test", "u")
if err != nil || !ok {
t.Fatalf("second response not stored: ok=%v err=%v", ok, err)
}
var final []any
if err := json.Unmarshal(stored.Context, &final); err != nil {
t.Fatal(err)
}
if len(final) != 4 {
t.Fatalf("stored final context has %d items: %s", len(final), stored.Context)
}
}
+103
View File
@@ -0,0 +1,103 @@
package server
import (
"context"
"errors"
"sort"
"sync"
"time"
"github.com/example/ollama-fair-gateway/internal/liveflow"
)
var errJobCancelled = errors.New("job cancelled by administrator")
type jobEntry struct {
ID string `json:"id"`
Tenant string `json:"tenant"`
Actor string `json:"actor"`
Application string `json:"application,omitempty"`
ServiceClass string `json:"service_class,omitempty"`
Model string `json:"model,omitempty"`
Path string `json:"path"`
API string `json:"api"`
Worker string `json:"worker,omitempty"`
CreatedAt time.Time `json:"created_at"`
CancelledAt *time.Time `json:"cancelled_at,omitempty"`
Cancelling bool `json:"cancelling,omitempty"`
}
type jobView struct {
liveflow.Request
Cancellable bool `json:"cancellable"`
Cancelling bool `json:"cancelling,omitempty"`
}
type jobManager struct {
mu sync.RWMutex
jobs map[string]jobEntry
cancel map[string]context.CancelCauseFunc
}
func newJobManager() *jobManager {
return &jobManager{jobs: make(map[string]jobEntry), cancel: make(map[string]context.CancelCauseFunc)}
}
func (m *jobManager) register(j jobEntry, cancel context.CancelCauseFunc) {
m.mu.Lock()
defer m.mu.Unlock()
m.jobs[j.ID] = j
m.cancel[j.ID] = cancel
}
func (m *jobManager) setWorker(id, worker string) {
m.mu.Lock()
defer m.mu.Unlock()
j, ok := m.jobs[id]
if !ok {
return
}
j.Worker = worker
m.jobs[id] = j
}
func (m *jobManager) finish(id string) {
m.mu.Lock()
delete(m.jobs, id)
delete(m.cancel, id)
m.mu.Unlock()
}
func (m *jobManager) cancelJob(id string) bool {
m.mu.Lock()
cancel := m.cancel[id]
j, ok := m.jobs[id]
if cancel == nil || !ok {
m.mu.Unlock()
return false
}
if !j.Cancelling {
now := time.Now().UTC()
j.Cancelling = true
j.CancelledAt = &now
m.jobs[id] = j
}
m.mu.Unlock()
cancel(errJobCancelled)
return true
}
func (m *jobManager) list() []jobEntry {
m.mu.RLock()
out := make([]jobEntry, 0, len(m.jobs))
for _, j := range m.jobs {
out = append(out, j)
}
m.mu.RUnlock()
sort.Slice(out, func(i, j int) bool { return out[i].CreatedAt.Before(out[j].CreatedAt) })
return out
}
func isAdminJobCancel(ctx context.Context) bool {
return errors.Is(context.Cause(ctx), errJobCancelled)
}
+179
View File
@@ -0,0 +1,179 @@
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")
}
}
+107
View File
@@ -0,0 +1,107 @@
package server
import (
"context"
"encoding/json"
"errors"
"fmt"
"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/worker"
)
var ErrModelAccessDenied = errors.New("model access denied")
var ErrAliasUnavailable = errors.New("model alias has no routable target")
func (s *Server) tenantModelAccessRule(id auth.Identity) config.ModelAccessRule {
return s.runtimeTenantModelAccessRule(id.Tenant)
}
func (s *Server) modelAllowed(id auth.Identity, model string) bool {
if strings.TrimSpace(model) == "" {
return true
}
// Tenant policy is the outer security boundary. An API-key ACL may narrow
// that boundary, but must never widen it.
if !config.ModelAccessAllowed(s.tenantModelAccessRule(id), model) {
return false
}
if id.ModelACLSet && !config.ModelAccessAllowed(id.ModelAccess, model) {
return false
}
return true
}
func (s *Server) resolveModel(ctx context.Context, id auth.Identity, requested string) (string, string, error) {
requested = strings.TrimSpace(requested)
if requested == "" {
return "", "", nil
}
if !s.modelAllowed(id, requested) {
return "", "", fmt.Errorf("%w: model %q is not allowed for this identity", ErrModelAccessDenied, requested)
}
a, ok := s.aliasConfig(requested)
if !ok {
return requested, "", nil
}
for _, m := range a.Models {
m = strings.TrimSpace(m)
if m == "" || !s.workers.CanRoute(m) {
continue
}
if len(a.RequiredCapabilities) > 0 {
meta, _, err := s.workers.Metadata(ctx, m)
if err != nil {
continue
}
okCaps := true
for _, capName := range a.RequiredCapabilities {
if !worker.HasCapability(meta, capName) {
okCaps = false
break
}
}
if !okCaps {
continue
}
}
return m, requested, nil
}
return "", requested, fmt.Errorf("%w: alias %q has no healthy/eligible installed target", ErrAliasUnavailable, requested)
}
func rewriteModelBody(body []byte, model string) ([]byte, error) {
if len(body) == 0 || strings.TrimSpace(model) == "" {
return body, nil
}
var v map[string]any
if err := json.Unmarshal(body, &v); err != nil {
return nil, err
}
v["model"] = model
if _, ok := v["name"]; ok {
v["name"] = model
}
return json.Marshal(v)
}
func (s *Server) visibleAliases(ctx context.Context, id auth.Identity) []string {
aliases := s.aliasSnapshot()
out := make([]string, 0, len(aliases))
for name, a := range aliases {
if a.Visible != nil && !*a.Visible {
continue
}
if !s.modelAllowed(id, name) {
continue
}
if _, _, err := s.resolveModel(ctx, id, name); err == nil {
out = append(out, name)
}
}
sort.Strings(out)
return out
}
+164
View File
@@ -0,0 +1,164 @@
package server
import (
"context"
"encoding/json"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"time"
"github.com/example/ollama-fair-gateway/internal/auth"
"github.com/example/ollama-fair-gateway/internal/config"
"github.com/example/ollama-fair-gateway/internal/cost"
"github.com/example/ollama-fair-gateway/internal/metrics"
px "github.com/example/ollama-fair-gateway/internal/proxy"
"github.com/example/ollama-fair-gateway/internal/quota"
"github.com/example/ollama-fair-gateway/internal/scheduler"
"github.com/example/ollama-fair-gateway/internal/usage"
"github.com/example/ollama-fair-gateway/internal/worker"
)
func TestOpenWebUIOllamaCompatibility(t *testing.T) {
var betaShowOnB atomic.Bool
var betaChatOnB atomic.Bool
backendA := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/ps":
writeJSON(w, 200, map[string]any{"models": []any{}})
case "/api/tags":
// Deliberately omit "model" to verify gateway normalization for
// clients such as OpenWebUI that key discovery by that field.
io.WriteString(w, `{"models":[{"name":"alpha:latest","size":100,"details":{"family":"alpha"}}]}`)
case "/api/version":
io.WriteString(w, `{"version":"0.99.0"}`)
case "/api/show":
io.WriteString(w, `{"error":"model not found"}`)
case "/api/chat":
w.WriteHeader(http.StatusNotFound)
default:
w.WriteHeader(http.StatusNotFound)
}
}))
defer backendA.Close()
backendB := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/ps":
writeJSON(w, 200, map[string]any{"models": []any{}})
case "/api/tags":
io.WriteString(w, `{"models":[{"name":"beta:latest","model":"beta:latest","size":200,"details":{"family":"beta"}}]}`)
case "/api/version":
io.WriteString(w, `{"version":"0.99.0"}`)
case "/api/show":
b, _ := io.ReadAll(r.Body)
var v map[string]any
_ = json.Unmarshal(b, &v)
if v["model"] == "beta:latest" {
betaShowOnB.Store(true)
io.WriteString(w, `{"modelfile":"FROM beta"}`)
return
}
w.WriteHeader(http.StatusNotFound)
case "/api/chat":
betaChatOnB.Store(true)
w.Header().Set("Content-Type", "application/x-ndjson")
io.WriteString(w, "{\"message\":{\"content\":\"ok\"},\"done\":true,\"prompt_eval_count\":2,\"eval_count\":1}\n")
default:
w.WriteHeader(http.StatusNotFound)
}
}))
defer backendB.Close()
cfg := &config.Config{
Server: config.ServerConfig{MaxBodyBytes: 1 << 20, MaxRequestDuration: config.Duration(time.Minute)},
Auth: config.AuthConfig{APIKeys: []config.APIKeyConfig{{
Name: "openwebui", Key: "owui-secret", Tenant: "apps", Subject: "openwebui", Application: "openwebui",
}}},
Scheduler: config.SchedulerConfig{
GlobalConcurrency: 2, MaxQueue: 16, MaxQueuePerActor: 8,
QueueTimeout: config.Duration(time.Second), DefaultTenantWeight: 1, DefaultActorWeight: 1,
ComputePaths: []string{"/api/chat", "/api/generate", "/api/embed", "/api/embeddings", "/v1/chat/completions"},
},
Cost: config.CostConfig{Default: config.ModelRate{InputCreditsPer1K: 1, OutputCreditsPer1K: 1}, DefaultMaxOutputTokens: 16},
Workers: []config.WorkerConfig{
{Name: "a", URL: backendA.URL, MaxConcurrent: 1, HealthInterval: config.Duration(time.Hour)},
{Name: "b", URL: backendB.URL, MaxConcurrent: 1, HealthInterval: config.Duration(time.Hour)},
},
Native: config.NativeConfig{ControlWorker: "a"},
}
a, err := auth.New(context.Background(), cfg.Auth)
if err != nil {
t.Fatal(err)
}
wp := worker.New(cfg.Workers, "a")
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
wp.Start(ctx)
rec, _ := usage.New("", 128, time.Second, nil)
sv := New(cfg, Dependencies{
Auth: a, Scheduler: scheduler.NewLocal(2, 16, 8), Quota: quota.Disabled{}, Estimator: cost.New(cfg.Cost),
Workers: wp, Proxy: px.New(), Usage: rec, Metrics: metrics.New(), Logger: slog.Default(),
})
front := httptest.NewServer(sv.Handler())
defer front.Close()
do := func(method, path, body string, withKey bool) (int, string) {
req, _ := http.NewRequest(method, front.URL+path, strings.NewReader(body))
if body != "" {
req.Header.Set("Content-Type", "application/json")
}
if withKey {
req.Header.Set("Authorization", "Bearer owui-secret")
}
resp, err := http.DefaultClient.Do(req)
if err != nil {
t.Fatal(err)
}
b, _ := io.ReadAll(resp.Body)
resp.Body.Close()
return resp.StatusCode, string(b)
}
if status, body := do(http.MethodGet, "/api/tags", "", false); status != http.StatusUnauthorized {
t.Fatalf("unauthenticated /api/tags status=%d, want 401", status)
} else if !strings.Contains(body, `"error":"authentication required"`) {
t.Fatalf("native Ollama error shape is not compatible: %s", body)
}
if status, body := do(http.MethodGet, "/v1/models", "", false); status != http.StatusUnauthorized {
t.Fatalf("unauthenticated /v1/models status=%d, want 401", status)
} else if !strings.Contains(body, `"message":"authentication required"`) {
t.Fatalf("OpenAI error shape changed unexpectedly: %s", body)
}
status, body := do(http.MethodGet, "/api/version", "", true)
if status != 200 || !strings.Contains(body, `"version":"0.99.0"`) {
t.Fatalf("version status=%d body=%s", status, body)
}
status, body = do(http.MethodGet, "/api/tags", "", true)
if status != 200 || !strings.Contains(body, `"model":"alpha:latest"`) || !strings.Contains(body, `"model":"beta:latest"`) {
t.Fatalf("tags status=%d body=%s", status, body)
}
status, body = do(http.MethodGet, "/v1/models", "", true)
if status != 200 || !strings.Contains(body, `"id":"alpha:latest"`) || !strings.Contains(body, `"id":"beta:latest"`) || !strings.Contains(body, `"object":"list"`) {
t.Fatalf("v1 models status=%d body=%s", status, body)
}
status, body = do(http.MethodPost, "/api/show", `{"model":"beta:latest"}`, true)
if status != 200 || !betaShowOnB.Load() {
t.Fatalf("show was not model-routed to backend B: status=%d body=%s", status, body)
}
status, body = do(http.MethodPost, "/api/chat", `{"model":"beta:latest","messages":[{"role":"user","content":"hi"}]}`, true)
if status != 200 || !betaChatOnB.Load() || !strings.Contains(body, `"content":"ok"`) {
t.Fatalf("chat was not model-routed to backend B: status=%d body=%s", status, body)
}
}
+229
View File
@@ -0,0 +1,229 @@
package server
import (
"bufio"
"bytes"
"context"
"crypto/rand"
"encoding/hex"
"encoding/json"
"fmt"
"io"
"net/http"
"net/url"
"sort"
"sync"
"time"
)
type adminOperation struct {
ID string `json:"id"`
Type string `json:"type"`
Worker string `json:"worker"`
Model string `json:"model"`
Status string `json:"status"`
Message string `json:"message,omitempty"`
Completed int64 `json:"completed,omitempty"`
Total int64 `json:"total,omitempty"`
Progress float64 `json:"progress"`
StartedAt time.Time `json:"started_at"`
UpdatedAt time.Time `json:"updated_at"`
Error string `json:"error,omitempty"`
}
type operationManager struct {
mu sync.RWMutex
ops map[string]adminOperation
cancel map[string]context.CancelFunc
maxKeep int
client *http.Client
}
func newOperationManager() *operationManager {
return &operationManager{ops: map[string]adminOperation{}, cancel: map[string]context.CancelFunc{}, maxKeep: 100, client: &http.Client{Timeout: 0}}
}
func (m *operationManager) list() []adminOperation {
m.mu.RLock()
defer m.mu.RUnlock()
out := make([]adminOperation, 0, len(m.ops))
for _, op := range m.ops {
out = append(out, op)
}
sort.Slice(out, func(i, j int) bool { return out[i].StartedAt.After(out[j].StartedAt) })
return out
}
func (m *operationManager) update(id string, fn func(*adminOperation)) {
m.mu.Lock()
defer m.mu.Unlock()
op, ok := m.ops[id]
if !ok {
return
}
fn(&op)
op.UpdatedAt = time.Now().UTC()
m.ops[id] = op
}
func (m *operationManager) add(op adminOperation) {
m.mu.Lock()
defer m.mu.Unlock()
m.ops[op.ID] = op
if len(m.ops) <= m.maxKeep {
return
}
var oldest string
var oldestTime time.Time
for id, x := range m.ops {
if x.Status == "running" || x.Status == "queued" {
continue
}
if oldest == "" || x.StartedAt.Before(oldestTime) {
oldest, oldestTime = id, x.StartedAt
}
}
if oldest != "" {
delete(m.ops, oldest)
}
}
func (m *operationManager) startPull(base *url.URL, worker, model string) adminOperation {
now := time.Now().UTC()
op := adminOperation{ID: operationID(), Type: "pull", Worker: worker, Model: model, Status: "queued", StartedAt: now, UpdatedAt: now}
m.add(op)
ctx, cancel := context.WithCancel(context.Background())
m.mu.Lock()
m.cancel[op.ID] = cancel
m.mu.Unlock()
go m.runPull(ctx, op.ID, base, model)
return op
}
func (m *operationManager) runPull(ctx context.Context, id string, base *url.URL, model string) {
defer func() {
m.mu.Lock()
delete(m.cancel, id)
m.mu.Unlock()
}()
m.update(id, func(op *adminOperation) { op.Status = "running"; op.Message = "starting pull" })
body, _ := json.Marshal(map[string]any{"model": model, "stream": true})
req, err := http.NewRequestWithContext(ctx, http.MethodPost, base.String()+"/api/pull", bytes.NewReader(body))
if err != nil {
m.fail(id, err)
return
}
req.Header.Set("Content-Type", "application/json")
resp, err := m.client.Do(req)
if err != nil {
if ctx.Err() != nil {
m.update(id, func(op *adminOperation) { op.Status = "cancelled"; op.Message = "cancelled" })
return
}
m.fail(id, err)
return
}
defer resp.Body.Close()
if resp.StatusCode/100 != 2 {
b, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
m.fail(id, fmt.Errorf("Ollama HTTP %d: %s", resp.StatusCode, string(b)))
return
}
sc := bufio.NewScanner(resp.Body)
buf := make([]byte, 0, 64<<10)
sc.Buffer(buf, 2<<20)
for sc.Scan() {
var x struct {
Status string `json:"status"`
Digest string `json:"digest"`
Total int64 `json:"total"`
Completed int64 `json:"completed"`
Error string `json:"error"`
}
if json.Unmarshal(sc.Bytes(), &x) != nil {
continue
}
if x.Error != "" {
m.fail(id, fmt.Errorf("%s", x.Error))
return
}
m.update(id, func(op *adminOperation) {
op.Message = x.Status
if x.Total > 0 {
op.Total = x.Total
}
if x.Completed > 0 {
op.Completed = x.Completed
}
if op.Total > 0 {
op.Progress = float64(op.Completed) / float64(op.Total)
if op.Progress > 1 {
op.Progress = 1
}
}
})
}
if err := sc.Err(); err != nil {
if ctx.Err() != nil {
m.update(id, func(op *adminOperation) { op.Status = "cancelled"; op.Message = "cancelled" })
return
}
m.fail(id, err)
return
}
m.update(id, func(op *adminOperation) { op.Status = "completed"; op.Message = "success"; op.Progress = 1 })
}
func (m *operationManager) fail(id string, err error) {
m.update(id, func(op *adminOperation) { op.Status = "failed"; op.Error = err.Error(); op.Message = "failed" })
}
func (m *operationManager) cancelOperation(id string) bool {
m.mu.RLock()
cancel := m.cancel[id]
m.mu.RUnlock()
if cancel == nil {
return false
}
cancel()
return true
}
func (m *operationManager) modelAction(ctx context.Context, base *url.URL, action, model string) error {
var method, path string
var payload any
switch action {
case "delete":
method, path = http.MethodDelete, "/api/delete"
payload = map[string]string{"model": model}
case "stop":
// Ollama unloads a resident model by issuing a generate request with
// keep_alive=0. There is no native /api/stop endpoint.
method, path = http.MethodPost, "/api/generate"
payload = map[string]any{"model": model, "keep_alive": 0, "stream": false}
default:
return fmt.Errorf("unsupported model action %q", action)
}
body, _ := json.Marshal(payload)
req, err := http.NewRequestWithContext(ctx, method, base.String()+path, bytes.NewReader(body))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
resp, err := m.client.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
if resp.StatusCode/100 != 2 {
b, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
return fmt.Errorf("Ollama HTTP %d: %s", resp.StatusCode, string(b))
}
return nil
}
func operationID() string {
b := make([]byte, 12)
_, _ = rand.Read(b)
return hex.EncodeToString(b)
}
+43
View File
@@ -0,0 +1,43 @@
package server
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"net/url"
"testing"
)
func TestModelStopUsesGenerateKeepAliveZero(t *testing.T) {
var gotMethod, gotPath string
var got map[string]any
backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
gotMethod, gotPath = r.Method, r.URL.Path
if err := json.NewDecoder(r.Body).Decode(&got); err != nil {
t.Errorf("decode body: %v", err)
}
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte(`{"done":true}`))
}))
defer backend.Close()
u, _ := url.Parse(backend.URL)
m := newOperationManager()
if err := m.modelAction(context.Background(), u, "stop", "gemma4:latest"); err != nil {
t.Fatal(err)
}
if gotMethod != http.MethodPost || gotPath != "/api/generate" {
t.Fatalf("method/path=%s %s", gotMethod, gotPath)
}
if got["model"] != "gemma4:latest" {
t.Fatalf("model=%v", got["model"])
}
if v, ok := got["keep_alive"].(float64); !ok || v != 0 {
t.Fatalf("keep_alive=%#v", got["keep_alive"])
}
if v, ok := got["stream"].(bool); !ok || v {
t.Fatalf("stream=%#v", got["stream"])
}
}
+284
View File
@@ -0,0 +1,284 @@
package server
import (
"context"
"fmt"
"math"
"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/worker"
)
type policySimulationRequest struct {
Tenant string `json:"tenant"`
APIKeyID string `json:"api_key_id,omitempty"`
APIKeyName string `json:"api_key_name,omitempty"`
Model string `json:"model"`
RequiredCapabilities []string `json:"required_capabilities,omitempty"`
InputTokens int64 `json:"input_tokens,omitempty"`
OutputTokens int64 `json:"output_tokens,omitempty"`
ServiceClass string `json:"service_class,omitempty"`
}
type accessExplanation struct {
TenantAllowed bool `json:"tenant_allowed"`
KeyApplied bool `json:"key_applied"`
KeyAllowed bool `json:"key_allowed"`
Allowed bool `json:"allowed"`
TenantRule config.ModelAccessRule `json:"tenant_rule"`
KeyRule config.ModelAccessRule `json:"key_rule,omitempty"`
}
type aliasCandidateExplanation struct {
Model string `json:"model"`
Routable bool `json:"routable"`
Capabilities []string `json:"capabilities,omitempty"`
Reason string `json:"reason,omitempty"`
}
type policySimulationResult struct {
RequestedModel string `json:"requested_model"`
ResolvedModel string `json:"resolved_model,omitempty"`
Alias string `json:"alias,omitempty"`
AliasCandidates []aliasCandidateExplanation `json:"alias_candidates,omitempty"`
Access accessExplanation `json:"access"`
CapabilitiesRequired []string `json:"capabilities_required,omitempty"`
CapabilitiesKnown []string `json:"capabilities_known,omitempty"`
CapabilitiesOK bool `json:"capabilities_ok"`
ContextLength int64 `json:"context_length,omitempty"`
ContextEffectiveMax int64 `json:"context_effective_max,omitempty"`
ContextRequested int64 `json:"context_requested,omitempty"`
ContextOK bool `json:"context_ok"`
EstimatedCredits float64 `json:"estimated_credits"`
CostRate config.ModelRate `json:"cost_rate"`
TenantPolicy config.TenantPolicy `json:"tenant_policy"`
ServiceClass string `json:"service_class"`
ServiceClassConfig config.ServiceClassConfig `json:"service_class_config"`
Workers []worker.RoutingExplanation `json:"workers"`
SelectedWorker string `json:"selected_worker,omitempty"`
Decision string `json:"decision"`
Errors []string `json:"errors,omitempty"`
}
func (s *Server) simulatedIdentity(in policySimulationRequest) (auth.Identity, *auth.APIKeyInfo, error) {
tenant := strings.TrimSpace(in.Tenant)
var selected *auth.APIKeyInfo
if strings.TrimSpace(in.APIKeyID) != "" || strings.TrimSpace(in.APIKeyName) != "" {
for _, k := range s.auth.APIKeys() {
idMatch := in.APIKeyID != "" && k.ID == in.APIKeyID
nameMatch := in.APIKeyID == "" && in.APIKeyName != "" && k.Name == in.APIKeyName && (tenant == "" || k.Tenant == tenant)
if idMatch || nameMatch {
kk := k
selected = &kk
break
}
}
if selected == nil {
return auth.Identity{}, nil, fmt.Errorf("API key not found")
}
if tenant == "" {
tenant = selected.Tenant
} else if selected.Tenant != tenant {
return auth.Identity{}, nil, fmt.Errorf("API key belongs to tenant %q, not %q", selected.Tenant, tenant)
}
}
if tenant == "" {
tenant = "default"
}
id := auth.Identity{Tenant: tenant, Subject: "policy-simulator", Application: "admin-ui", AuthType: "simulation", Scopes: map[string]bool{}}
if selected != nil {
id.Subject = selected.Subject
id.Application = selected.Application
id.ServiceClass = selected.ServiceClass
if len(selected.AllowedModels) > 0 || len(selected.DeniedModels) > 0 {
id.ModelACLSet = true
id.ModelAccess = config.ModelAccessRule{Mode: "allow_all", AllowedModels: append([]string(nil), selected.AllowedModels...), DeniedModels: append([]string(nil), selected.DeniedModels...)}
}
for _, scope := range selected.Scopes {
id.Scopes[scope] = true
}
}
return id, selected, nil
}
func (s *Server) simulatePolicy(ctx context.Context, in policySimulationRequest) (policySimulationResult, error) {
in.Model = strings.TrimSpace(in.Model)
if in.Model == "" {
return policySimulationResult{}, fmt.Errorf("model is required")
}
if in.InputTokens < 0 || in.OutputTokens < 0 {
return policySimulationResult{}, fmt.Errorf("token estimates must be >= 0")
}
id, _, err := s.simulatedIdentity(in)
if err != nil {
return policySimulationResult{}, err
}
result := policySimulationResult{RequestedModel: in.Model, CapabilitiesOK: true, ContextOK: true}
tenantRule := s.tenantModelAccessRule(id)
result.Access = accessExplanation{TenantAllowed: config.ModelAccessAllowed(tenantRule, in.Model), TenantRule: tenantRule, KeyAllowed: true}
if id.ModelACLSet {
result.Access.KeyApplied = true
result.Access.KeyRule = id.ModelAccess
result.Access.KeyAllowed = config.ModelAccessAllowed(id.ModelAccess, in.Model)
}
result.Access.Allowed = result.Access.TenantAllowed && result.Access.KeyAllowed
if !result.Access.Allowed {
result.Decision = "denied_model_access"
return result, nil
}
resolved := in.Model
if alias, ok := s.aliasConfig(in.Model); ok {
result.Alias = in.Model
for _, candidate := range alias.Models {
candidate = strings.TrimSpace(candidate)
if candidate == "" {
continue
}
x := aliasCandidateExplanation{Model: candidate, Routable: s.workers.CanRoute(candidate)}
if !x.Routable {
x.Reason = "no_eligible_worker"
}
if len(alias.RequiredCapabilities) > 0 {
meta, _, metaErr := s.workers.Metadata(ctx, candidate)
if metaErr != nil {
x.Reason = "metadata_unavailable"
} else {
x.Capabilities = append([]string(nil), meta.Capabilities...)
for _, c := range alias.RequiredCapabilities {
if !worker.HasCapability(meta, c) {
x.Routable = false
x.Reason = "missing_alias_capability:" + c
break
}
}
}
}
result.AliasCandidates = append(result.AliasCandidates, x)
if resolved == in.Model && x.Routable {
resolved = candidate
}
}
if resolved == in.Model {
result.Decision = "alias_unavailable"
return result, nil
}
}
result.ResolvedModel = resolved
requiredMap := map[string]bool{}
for _, c := range in.RequiredCapabilities {
if c = strings.TrimSpace(c); c != "" {
requiredMap[c] = true
}
}
if alias, ok := s.aliasConfig(in.Model); ok {
for _, c := range alias.RequiredCapabilities {
if c = strings.TrimSpace(c); c != "" {
requiredMap[c] = true
}
}
}
for c := range requiredMap {
result.CapabilitiesRequired = append(result.CapabilitiesRequired, c)
}
sort.Strings(result.CapabilitiesRequired)
if len(result.CapabilitiesRequired) > 0 || in.InputTokens+in.OutputTokens > 0 {
meta, _, metaErr := s.workers.Metadata(ctx, resolved)
if metaErr != nil {
result.Errors = append(result.Errors, "model metadata unavailable: "+metaErr.Error())
} else {
result.CapabilitiesKnown = append([]string(nil), meta.Capabilities...)
result.ContextLength = meta.ContextLength
for _, c := range result.CapabilitiesRequired {
if !worker.HasCapability(meta, c) {
result.CapabilitiesOK = false
result.Errors = append(result.Errors, "missing capability: "+c)
}
}
}
}
margin := int64(0)
if pct := s.cfg.ModelCapabilities.Context.EstimationMarginPercent; pct > 0 && in.InputTokens > 0 {
margin = int64(math.Ceil(float64(in.InputTokens) * pct / 100))
}
result.ContextRequested = in.InputTokens + margin + in.OutputTokens
contextAllowed := map[string]bool{}
for _, cw := range s.workers.ContextWindows(ctx, resolved) {
effective := cw.EffectiveTokens
if cap := s.cfg.ModelCapabilities.Context.MaxRequestedTokens; cap > 0 {
effective = minPositive64(effective, cap)
}
if effective > result.ContextEffectiveMax {
result.ContextEffectiveMax = effective
}
if effective > 0 && result.ContextRequested <= effective {
contextAllowed[cw.Worker] = true
}
}
if cap := s.cfg.ModelCapabilities.Context.MaxRequestedTokens; cap > 0 && result.ContextRequested > cap {
result.ContextOK = false
result.Errors = append(result.Errors, fmt.Sprintf("context %d exceeds gateway cap %d", result.ContextRequested, cap))
} else if result.ContextEffectiveMax > 0 && result.ContextRequested > result.ContextEffectiveMax {
result.ContextOK = false
result.Errors = append(result.Errors, fmt.Sprintf("context %d exceeds effective worker context %d", result.ContextRequested, result.ContextEffectiveMax))
} else if result.ContextLength > 0 && result.ContextRequested > result.ContextLength {
result.ContextOK = false
result.Errors = append(result.Errors, fmt.Sprintf("context %d exceeds model limit %d", result.ContextRequested, result.ContextLength))
}
rate := s.estimator.Rate(resolved)
result.CostRate = rate
result.EstimatedCredits = float64(in.InputTokens)/1000*rate.InputCreditsPer1K + float64(in.OutputTokens)/1000*rate.OutputCreditsPer1K
if rate.ComputeCreditsPerSecond > 0 {
sec := 0.0
if rate.ExpectedPromptTokensPerSecond > 0 {
sec += float64(in.InputTokens) / rate.ExpectedPromptTokensPerSecond
}
if rate.ExpectedOutputTokensPerSecond > 0 {
sec += float64(in.OutputTokens) / rate.ExpectedOutputTokensPerSecond
}
result.EstimatedCredits += sec * rate.ComputeCreditsPerSecond
}
result.TenantPolicy = s.policyFor(ctx, id.Tenant)
className := strings.TrimSpace(in.ServiceClass)
if className == "" {
className = strings.TrimSpace(id.ServiceClass)
}
if className == "" {
className = s.cfg.ServiceClasses.Default
}
if className == "" {
className = "interactive"
}
result.ServiceClass = className
if sc, ok := s.cfg.ServiceClasses.Classes[className]; ok {
result.ServiceClassConfig = sc
} else {
result.Errors = append(result.Errors, "unknown service class: "+className)
}
if len(contextAllowed) == 0 {
contextAllowed = nil
}
result.Workers = s.workers.ExplainRoutingAllowed(resolved, contextAllowed, 0)
for _, w := range result.Workers {
if w.Eligible {
result.SelectedWorker = w.Worker
break
}
}
if !result.CapabilitiesOK {
result.Decision = "denied_capability"
} else if !result.ContextOK {
result.Decision = "denied_context"
} else if result.SelectedWorker == "" {
result.Decision = "no_eligible_worker"
} else {
result.Decision = "would_route"
}
return result, nil
}
+343
View File
@@ -0,0 +1,343 @@
package server
import (
"bytes"
"encoding/json"
"fmt"
"math"
"net/http"
"sort"
"strconv"
"strings"
"github.com/example/ollama-fair-gateway/internal/cost"
"github.com/example/ollama-fair-gateway/internal/worker"
)
type requestRequirements struct {
Model string
Capabilities []string
RequestedContext int64
}
type modelPreflight struct {
AllowedWorkers map[string]bool
RequestedContext int64
RequiredContext int64
}
func requirementsFor(path string, body []byte) (requestRequirements, error) {
var doc map[string]json.RawMessage
_ = json.Unmarshal(body, &doc)
var model string
_ = json.Unmarshal(doc["model"], &model)
req := requestRequirements{Model: model}
caps := map[string]bool{}
if strings.Contains(path, "embed") {
caps["embedding"] = true
}
if strings.Contains(path, "chat") || strings.Contains(path, "completion") || strings.Contains(path, "responses") || strings.Contains(path, "messages") || strings.Contains(path, "generate") {
caps["completion"] = true
}
if nonEmptyJSONList(doc["tools"]) {
caps["tools"] = true
}
if thinkingRequested(doc) {
caps["thinking"] = true
}
if containsVision(body) {
caps["vision"] = true
}
if raw := bytes.TrimSpace(doc["options"]); len(raw) > 0 && !bytes.Equal(raw, []byte("null")) {
var opts map[string]json.RawMessage
if err := json.Unmarshal(raw, &opts); err != nil {
return req, fmt.Errorf("options must be a JSON object: %w", err)
}
if rawNum, ok := opts["num_ctx"]; ok {
if !strings.HasPrefix(path, "/api/") {
return req, fmt.Errorf("options.num_ctx is only supported on native Ollama /api endpoints")
}
n, err := strictPositiveInt64(rawNum)
if err != nil {
return req, fmt.Errorf("options.num_ctx must be a positive integer: %w", err)
}
req.RequestedContext = n
}
}
for c := range caps {
req.Capabilities = append(req.Capabilities, c)
}
sort.Strings(req.Capabilities)
return req, nil
}
func strictPositiveInt64(raw json.RawMessage) (int64, error) {
raw = bytes.TrimSpace(raw)
if len(raw) == 0 || bytes.Equal(raw, []byte("null")) || raw[0] == '"' {
return 0, fmt.Errorf("not an integer")
}
n, err := strconv.ParseInt(string(raw), 10, 64)
if err != nil || n <= 0 {
if err == nil {
err = fmt.Errorf("must be greater than zero")
}
return 0, err
}
return n, nil
}
func nonEmptyJSONList(raw json.RawMessage) bool {
raw = bytes.TrimSpace(raw)
return len(raw) > 2 && !bytes.Equal(raw, []byte("null")) && !bytes.Equal(raw, []byte("[]"))
}
func thinkingRequested(doc map[string]json.RawMessage) bool {
if raw := bytes.TrimSpace(doc["think"]); len(raw) > 0 && !bytes.Equal(raw, []byte("false")) && !bytes.Equal(raw, []byte("null")) && !bytes.Equal(raw, []byte(`"none"`)) {
return true
}
var effort string
_ = json.Unmarshal(doc["reasoning_effort"], &effort)
if effort != "" && !strings.EqualFold(effort, "none") {
return true
}
if raw := doc["thinking"]; len(raw) > 0 {
var x struct {
Type string `json:"type"`
}
if json.Unmarshal(raw, &x) == nil && strings.EqualFold(x.Type, "enabled") {
return true
}
}
if raw := doc["reasoning"]; len(raw) > 0 {
var x struct {
Effort string `json:"effort"`
}
if json.Unmarshal(raw, &x) == nil && x.Effort != "" && !strings.EqualFold(x.Effort, "none") {
return true
}
}
return false
}
func countVisionInputs(body []byte) int64 {
var v any
if json.Unmarshal(body, &v) != nil {
return 0
}
var walk func(any) int64
walk = func(x any) int64 {
switch z := x.(type) {
case []any:
var n int64
for _, item := range z {
n += walk(item)
}
return n
case map[string]any:
var n int64
if imgs, ok := z["images"].([]any); ok {
n += int64(len(imgs))
}
typ, _ := z["type"].(string)
isImageObject := typ == "image_url" || typ == "input_image" || typ == "image"
if isImageObject {
n++
}
for k, item := range z {
// Payloads can be very large. The image object/array itself already
// reserves context tokens, so never recurse into base64/URLs.
if k == "image_url" || k == "images" || (isImageObject && (k == "url" || k == "data" || k == "image")) {
continue
}
n += walk(item)
}
return n
}
return 0
}
return walk(v)
}
func containsVision(body []byte) bool { return countVisionInputs(body) > 0 }
func requiredContextTokens(est cost.Estimate, body []byte, visionReserve int64, marginPercent float64) (required, vision, margin int64) {
images := countVisionInputs(body)
if images > 0 && visionReserve > 0 {
vision = images * visionReserve
}
base := est.InputTokens + vision
if marginPercent > 0 && base > 0 {
margin = int64(math.Ceil(float64(base) * marginPercent / 100))
}
return base + margin + est.OutputTokens, vision, margin
}
func minPositive64(values ...int64) int64 {
var out int64
for _, v := range values {
if v <= 0 {
continue
}
if out == 0 || v < out {
out = v
}
}
return out
}
func contextWindowsSummary(windows []worker.ContextWindow) string {
parts := make([]string, 0, len(windows))
for _, x := range windows {
if x.EffectiveTokens > 0 {
parts = append(parts, fmt.Sprintf("%s=%d(%s)", x.Worker, x.EffectiveTokens, x.EffectiveSource))
} else {
parts = append(parts, x.Worker+"=unknown")
}
}
return strings.Join(parts, ", ")
}
// preflightModel validates capability-sensitive requests and derives the set of
// workers that can actually serve the request's context budget. The model
// maximum from /api/show is only one input: loaded /api/ps context_length,
// Modelfile num_ctx, worker defaults and per-worker caps are stronger runtime
// evidence for OpenAI-compatible requests that cannot set num_ctx themselves.
func (s *Server) preflightModel(w http.ResponseWriter, r *http.Request, body []byte, est cost.Estimate) (modelPreflight, bool) {
out := modelPreflight{}
cfg := s.cfg.ModelCapabilities
if cfg.Mode == "off" || strings.TrimSpace(est.Model) == "" {
return out, true
}
req, reqErr := requirementsFor(r.URL.Path, body)
if reqErr != nil {
writeProtocolError(w, r, http.StatusBadRequest, "invalid_num_ctx", reqErr.Error())
return out, false
}
out.RequestedContext = req.RequestedContext
meta, metaWorker, err := s.workers.Metadata(r.Context(), est.Model)
if err != nil {
w.Header().Set("X-Gateway-Model-Metadata", "unavailable")
s.log.Warn("model metadata unavailable; capability checks are best-effort", "model", est.Model, "error", err)
} else {
if metaWorker != "" {
w.Header().Set("X-Gateway-Model-Metadata-Worker", metaWorker)
}
if len(meta.Capabilities) > 0 {
w.Header().Set("X-Gateway-Model-Capabilities", strings.Join(meta.Capabilities, ","))
}
if meta.ContextLength > 0 {
w.Header().Set("X-Gateway-Model-Context", strconv.FormatInt(meta.ContextLength, 10))
}
if meta.ConfiguredContextLength > 0 {
w.Header().Set("X-Gateway-Model-Configured-Context", strconv.FormatInt(meta.ConfiguredContextLength, 10))
}
}
unsupported := make([]string, 0)
if err == nil && len(meta.Capabilities) > 0 {
for _, capability := range req.Capabilities {
if !worker.HasCapability(meta, capability) {
unsupported = append(unsupported, capability)
}
}
}
if len(unsupported) > 0 {
msg := fmt.Sprintf("model %s does not support required capability: %s", est.Model, strings.Join(unsupported, ", "))
if cfg.Mode == "enforce" {
writeProtocolError(w, r, http.StatusBadRequest, "unsupported_capability", msg)
return out, false
}
w.Header().Set("X-Gateway-Capability-Warning", strings.Join(unsupported, ","))
s.log.Warn("unsupported model capability observed", "model", est.Model, "required", unsupported)
}
if cfg.ContextGuard == "off" {
return out, true
}
required, visionReserve, margin := requiredContextTokens(est, body, cfg.Context.VisionReserveTokensPerImage, cfg.Context.EstimationMarginPercent)
out.RequiredContext = required
w.Header().Set("X-Gateway-Context-Required", strconv.FormatInt(required, 10))
if req.RequestedContext > 0 {
w.Header().Set("X-Gateway-Requested-Context", strconv.FormatInt(req.RequestedContext, 10))
}
if visionReserve > 0 {
w.Header().Set("X-Gateway-Context-Vision-Reserve", strconv.FormatInt(visionReserve, 10))
}
if margin > 0 {
w.Header().Set("X-Gateway-Context-Margin", strconv.FormatInt(margin, 10))
}
var warning string
gatewayCap := cfg.Context.MaxRequestedTokens
if gatewayCap > 0 && req.RequestedContext > gatewayCap {
warning = fmt.Sprintf("requested num_ctx %d exceeds gateway context cap %d", req.RequestedContext, gatewayCap)
} else if gatewayCap > 0 && required > gatewayCap {
warning = fmt.Sprintf("estimated request context %d tokens exceeds gateway context cap %d", required, gatewayCap)
} else if req.RequestedContext > 0 && meta.ContextLength > 0 && req.RequestedContext > meta.ContextLength {
warning = fmt.Sprintf("requested num_ctx %d exceeds model context length %d", req.RequestedContext, meta.ContextLength)
} else if req.RequestedContext > 0 && required > req.RequestedContext {
warning = fmt.Sprintf("estimated request context %d tokens exceeds explicit num_ctx %d", required, req.RequestedContext)
}
windows := s.workers.ContextWindows(r.Context(), est.Model)
allowed := make(map[string]bool, len(windows))
var maxEffective int64
for _, x := range windows {
eligible := false
if req.RequestedContext > 0 {
// Explicit native num_ctx can resize/reload the model, so current loaded
// context does not hard-block the worker. The theoretical model maximum
// and the operator's per-worker context cap still do.
capacity := minPositive64(x.ModelMaxTokens, x.WorkerLimitTokens)
if gatewayCap > 0 {
capacity = minPositive64(capacity, gatewayCap)
}
if capacity > maxEffective {
maxEffective = capacity
}
eligible = (x.ModelMaxTokens <= 0 || req.RequestedContext <= x.ModelMaxTokens) &&
(x.WorkerLimitTokens <= 0 || req.RequestedContext <= x.WorkerLimitTokens)
} else {
effective := x.EffectiveTokens
if gatewayCap > 0 {
effective = minPositive64(effective, gatewayCap)
}
if effective > maxEffective {
maxEffective = effective
}
eligible = effective > 0 && required <= effective
}
if eligible {
allowed[x.Worker] = true
}
}
if maxEffective > 0 {
w.Header().Set("X-Gateway-Effective-Context-Max", strconv.FormatInt(maxEffective, 10))
}
if len(windows) > 0 {
w.Header().Set("X-Gateway-Context-Eligible-Workers", fmt.Sprintf("%d/%d", len(allowed), len(windows)))
}
if warning == "" && len(windows) > 0 && len(allowed) == 0 {
if req.RequestedContext > 0 {
warning = fmt.Sprintf("requested num_ctx %d cannot be served by any eligible worker; contexts: %s", req.RequestedContext, contextWindowsSummary(windows))
} else {
warning = fmt.Sprintf("estimated request context %d tokens exceeds every effective worker context; contexts: %s", required, contextWindowsSummary(windows))
}
}
if warning != "" {
if cfg.ContextGuard == "reject" {
writeProtocolError(w, r, http.StatusBadRequest, "context_window_exceeded", warning)
return out, false
}
w.Header().Set("X-Gateway-Context-Warning", warning)
s.log.Warn("context guard warning", "model", est.Model, "warning", warning)
}
// Even in warn mode, prefer context-suitable workers when at least one exists.
// If none exists, warn mode intentionally preserves legacy routing behavior.
if len(allowed) > 0 {
out.AllowedWorkers = allowed
}
return out, true
}
+344
View File
@@ -0,0 +1,344 @@
package server
import (
"context"
"encoding/json"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"time"
"github.com/example/ollama-fair-gateway/internal/auth"
"github.com/example/ollama-fair-gateway/internal/config"
"github.com/example/ollama-fair-gateway/internal/cost"
"github.com/example/ollama-fair-gateway/internal/metrics"
px "github.com/example/ollama-fair-gateway/internal/proxy"
"github.com/example/ollama-fair-gateway/internal/quota"
"github.com/example/ollama-fair-gateway/internal/scheduler"
"github.com/example/ollama-fair-gateway/internal/usage"
"github.com/example/ollama-fair-gateway/internal/worker"
)
func newPreflightServer(t *testing.T, capabilities []string, contextLength int64) *httptest.Server {
t.Helper()
backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/ps":
_ = json.NewEncoder(w).Encode(map[string]any{"models": []any{map[string]any{"name": "codellama:7b", "model": "codellama:7b"}}})
case "/api/tags":
_ = json.NewEncoder(w).Encode(map[string]any{"models": []any{map[string]any{"name": "codellama:7b", "model": "codellama:7b"}}})
case "/api/show":
_ = json.NewEncoder(w).Encode(map[string]any{"capabilities": capabilities, "model_info": map[string]any{"llama.context_length": contextLength}})
case "/api/chat":
_ = json.NewEncoder(w).Encode(map[string]any{"done": true, "message": map[string]any{"content": "ok"}, "prompt_eval_count": 2, "eval_count": 1})
default:
http.NotFound(w, r)
}
}))
t.Cleanup(backend.Close)
cfg := &config.Config{
Server: config.ServerConfig{MaxBodyBytes: 1 << 20, MaxRequestDuration: config.Duration(time.Minute)},
Auth: config.AuthConfig{IPBypass: []config.IPBypassConfig{{CIDRs: []string{"127.0.0.1/32"}, Tenant: "t", Subject: "u"}}},
Scheduler: config.SchedulerConfig{GlobalConcurrency: 1, MaxQueue: 8, MaxQueuePerActor: 8, QueueTimeout: config.Duration(time.Second), DefaultTenantWeight: 1, DefaultActorWeight: 1, ComputePaths: []string{"/api/chat"}},
Cost: config.CostConfig{Default: config.ModelRate{InputCreditsPer1K: 1, OutputCreditsPer1K: 3, CachedInputFactor: 1}, DefaultMaxOutputTokens: 16},
Workers: []config.WorkerConfig{{Name: "w", URL: backend.URL, MaxConcurrent: 1, HealthInterval: config.Duration(time.Hour)}},
ModelCapabilities: config.ModelCapabilitiesConfig{Mode: "enforce", CacheTTL: config.Duration(time.Hour), ContextGuard: "reject"},
}
a, err := auth.New(context.Background(), cfg.Auth)
if err != nil {
t.Fatal(err)
}
wp := worker.New(cfg.Workers, "w")
wp.SetModelCapabilitiesConfig(cfg.ModelCapabilities)
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
wp.Start(ctx)
met := metrics.New()
rec, _ := usage.New("", 100, time.Second, nil)
sv := New(cfg, Dependencies{Auth: a, Scheduler: scheduler.NewLocal(1, 8, 8), Quota: quota.Disabled{}, Estimator: cost.New(cfg.Cost), Workers: wp, Proxy: px.New(), Usage: rec, Metrics: met, Logger: slog.Default()})
front := httptest.NewServer(sv.Handler())
t.Cleanup(front.Close)
return front
}
func TestRejectsUnsupportedToolsBeforeOllama(t *testing.T) {
front := newPreflightServer(t, []string{"completion"}, 16384)
resp, err := http.Post(front.URL+"/api/chat", "application/json", strings.NewReader(`{"model":"codellama:7b","messages":[{"role":"user","content":"x"}],"tools":[{"type":"function","function":{"name":"f"}}]}`))
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
b, _ := io.ReadAll(resp.Body)
if resp.StatusCode != 400 {
t.Fatalf("status=%d body=%s", resp.StatusCode, b)
}
if !strings.Contains(string(b), "does not support required capability: tools") {
t.Fatalf("unexpected body: %s", b)
}
}
func TestAllowsSupportedTools(t *testing.T) {
front := newPreflightServer(t, []string{"completion", "tools"}, 16384)
resp, err := http.Post(front.URL+"/api/chat", "application/json", strings.NewReader(`{"model":"codellama:7b","messages":[{"role":"user","content":"x"}],"tools":[{"type":"function","function":{"name":"f"}}]}`))
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != 200 {
b, _ := io.ReadAll(resp.Body)
t.Fatalf("status=%d body=%s", resp.StatusCode, b)
}
if got := resp.Header.Get("X-Gateway-Model-Capabilities"); !strings.Contains(got, "tools") {
t.Fatalf("capability header=%q", got)
}
}
func TestRejectsExplicitContextOverModelMaximum(t *testing.T) {
front := newPreflightServer(t, []string{"completion"}, 4096)
resp, err := http.Post(front.URL+"/api/chat", "application/json", strings.NewReader(`{"model":"codellama:7b","messages":[{"role":"user","content":"x"}],"options":{"num_ctx":8192}}`))
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
b, _ := io.ReadAll(resp.Body)
if resp.StatusCode != 400 || !strings.Contains(string(b), "exceeds model context length") {
t.Fatalf("status=%d body=%s", resp.StatusCode, b)
}
}
func newContextGuardServer(t *testing.T, modelMax, loadedContext int64, parameters string, ctxCfg config.ContextPolicyConfig, workerLimits map[string]int64) (*httptest.Server, *atomic.Int64) {
t.Helper()
var inferenceCalls atomic.Int64
backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/ps":
models := []any{}
if loadedContext > 0 {
models = append(models, map[string]any{"name": "ctx:latest", "model": "ctx:latest", "context_length": loadedContext})
}
_ = json.NewEncoder(w).Encode(map[string]any{"models": models})
case "/api/tags":
_ = json.NewEncoder(w).Encode(map[string]any{"models": []any{map[string]any{"name": "ctx:latest", "model": "ctx:latest"}}})
case "/api/show":
_ = json.NewEncoder(w).Encode(map[string]any{"capabilities": []string{"completion", "vision"}, "model_info": map[string]any{"ctx.context_length": modelMax}, "parameters": parameters})
case "/api/chat", "/v1/chat/completions", "/v1/responses", "/api/generate":
inferenceCalls.Add(1)
_ = json.NewEncoder(w).Encode(map[string]any{"done": true, "message": map[string]any{"content": "ok"}, "choices": []any{}, "usage": map[string]any{"prompt_tokens": 2, "completion_tokens": 1}, "prompt_eval_count": 2, "eval_count": 1})
default:
http.NotFound(w, r)
}
}))
t.Cleanup(backend.Close)
cfg := &config.Config{
Server: config.ServerConfig{MaxBodyBytes: 1 << 20, MaxRequestDuration: config.Duration(time.Minute)},
Auth: config.AuthConfig{IPBypass: []config.IPBypassConfig{{CIDRs: []string{"127.0.0.1/32"}, Tenant: "t", Subject: "u"}}},
Scheduler: config.SchedulerConfig{GlobalConcurrency: 1, MaxQueue: 8, MaxQueuePerActor: 8, QueueTimeout: config.Duration(time.Second), DefaultTenantWeight: 1, DefaultActorWeight: 1, ComputePaths: []string{"/api/chat", "/api/generate", "/v1/chat/completions", "/v1/responses"}},
Cost: config.CostConfig{Default: config.ModelRate{InputCreditsPer1K: 1, OutputCreditsPer1K: 3, CachedInputFactor: 1}, DefaultMaxOutputTokens: 16},
Workers: []config.WorkerConfig{{Name: "w", URL: backend.URL, MaxConcurrent: 1, HealthInterval: config.Duration(time.Hour), ContextLimits: workerLimits}},
ModelCapabilities: config.ModelCapabilitiesConfig{Mode: "enforce", CacheTTL: config.Duration(time.Hour), ContextGuard: "reject", Context: ctxCfg},
}
a, err := auth.New(context.Background(), cfg.Auth)
if err != nil {
t.Fatal(err)
}
wp := worker.New(cfg.Workers, "w")
wp.SetModelCapabilitiesConfig(cfg.ModelCapabilities)
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
wp.Start(ctx)
rec, _ := usage.New("", 100, time.Second, nil)
sv := New(cfg, Dependencies{Auth: a, Scheduler: scheduler.NewLocal(1, 8, 8), Quota: quota.Disabled{}, Estimator: cost.New(cfg.Cost), Workers: wp, Proxy: px.New(), Usage: rec, Metrics: metrics.New(), Logger: slog.Default()})
front := httptest.NewServer(sv.Handler())
t.Cleanup(front.Close)
return front, &inferenceCalls
}
func TestInvalidNumCtxRejectedBeforeBackend(t *testing.T) {
front, calls := newContextGuardServer(t, 131072, 0, "", config.ContextPolicyConfig{MaxRequestedTokens: 32768, DefaultWorkerTokens: 4096, EstimationMarginPercent: 15, VisionReserveTokensPerImage: 2048}, nil)
for _, raw := range []string{`-1`, `0`, `"8192"`, `1.5`, `null`} {
resp, err := http.Post(front.URL+"/api/chat", "application/json", strings.NewReader(`{"model":"ctx:latest","messages":[{"role":"user","content":"x"}],"options":{"num_ctx":`+raw+`}}`))
if err != nil {
t.Fatal(err)
}
b, _ := io.ReadAll(resp.Body)
resp.Body.Close()
if resp.StatusCode != 400 || !strings.Contains(string(b), "options.num_ctx must be a positive integer") {
t.Fatalf("num_ctx=%s status=%d body=%s", raw, resp.StatusCode, b)
}
}
if calls.Load() != 0 {
t.Fatalf("backend called %d times", calls.Load())
}
}
func TestResponsesInstructionsAndMaxOutputRespectEffectiveContext(t *testing.T) {
front, calls := newContextGuardServer(t, 131072, 4096, "", config.ContextPolicyConfig{MaxRequestedTokens: 32768, DefaultWorkerTokens: 4096, EstimationMarginPercent: 15, VisionReserveTokensPerImage: 2048}, nil)
body := `{"model":"ctx:latest","instructions":"` + strings.Repeat("x", 8000) + `","input":"hello","max_output_tokens":3000}`
resp, err := http.Post(front.URL+"/v1/responses", "application/json", strings.NewReader(body))
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
b, _ := io.ReadAll(resp.Body)
if resp.StatusCode != 400 || !strings.Contains(string(b), "effective worker context") {
t.Fatalf("status=%d body=%s headers=%v", resp.StatusCode, b, resp.Header)
}
if calls.Load() != 0 {
t.Fatalf("backend called %d times", calls.Load())
}
}
func TestOpenAIUsesModelfileNumCtxWhenModelNotLoaded(t *testing.T) {
front, calls := newContextGuardServer(t, 131072, 0, "num_ctx 16384\ntemperature 0.7", config.ContextPolicyConfig{MaxRequestedTokens: 32768, DefaultWorkerTokens: 4096, EstimationMarginPercent: 10, VisionReserveTokensPerImage: 2048}, nil)
body := `{"model":"ctx:latest","messages":[{"role":"user","content":"` + strings.Repeat("x", 12000) + `"}],"max_tokens":1024}`
resp, err := http.Post(front.URL+"/v1/chat/completions", "application/json", strings.NewReader(body))
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
b, _ := io.ReadAll(resp.Body)
if resp.StatusCode != 200 {
t.Fatalf("status=%d body=%s headers=%v", resp.StatusCode, b, resp.Header)
}
if calls.Load() != 1 {
t.Fatalf("backend calls=%d", calls.Load())
}
if got := resp.Header.Get("X-Gateway-Model-Configured-Context"); got != "16384" {
t.Fatalf("configured context header=%q", got)
}
}
func TestExplicitNumCtxHonorsGatewayAndWorkerCaps(t *testing.T) {
front, calls := newContextGuardServer(t, 131072, 4096, "", config.ContextPolicyConfig{MaxRequestedTokens: 32768, DefaultWorkerTokens: 4096, EstimationMarginPercent: 10, VisionReserveTokensPerImage: 2048}, map[string]int64{"*": 16384})
resp, err := http.Post(front.URL+"/api/chat", "application/json", strings.NewReader(`{"model":"ctx:latest","messages":[{"role":"user","content":"x"}],"options":{"num_ctx":20000}}`))
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
b, _ := io.ReadAll(resp.Body)
if resp.StatusCode != 400 || !strings.Contains(string(b), "cannot be served by any eligible worker") {
t.Fatalf("status=%d body=%s", resp.StatusCode, b)
}
if calls.Load() != 0 {
t.Fatalf("backend called %d times", calls.Load())
}
}
func TestVisionReserveParticipatesInContextGuard(t *testing.T) {
front, calls := newContextGuardServer(t, 131072, 4096, "", config.ContextPolicyConfig{MaxRequestedTokens: 32768, DefaultWorkerTokens: 4096, EstimationMarginPercent: 0, VisionReserveTokensPerImage: 3000}, nil)
body := `{"model":"ctx:latest","messages":[{"role":"user","content":[{"type":"text","text":"hello"},{"type":"image_url","image_url":{"url":"data:image/png;base64,AAAA"}}]}],"max_tokens":1500}`
resp, err := http.Post(front.URL+"/v1/chat/completions", "application/json", strings.NewReader(body))
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
b, _ := io.ReadAll(resp.Body)
if resp.StatusCode != 400 || resp.Header.Get("X-Gateway-Context-Vision-Reserve") != "3000" {
t.Fatalf("status=%d reserve=%q body=%s", resp.StatusCode, resp.Header.Get("X-Gateway-Context-Vision-Reserve"), b)
}
if calls.Load() != 0 {
t.Fatalf("backend called %d times", calls.Load())
}
}
func TestContextGuardRoutesToWorkerWithSufficientLoadedContext(t *testing.T) {
var smallCalls, largeCalls atomic.Int64
backend := func(ctxTokens int64, calls *atomic.Int64) *httptest.Server {
return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/ps":
_ = json.NewEncoder(w).Encode(map[string]any{"models": []any{map[string]any{"name": "ctx:latest", "model": "ctx:latest", "context_length": ctxTokens}}})
case "/api/tags":
_ = json.NewEncoder(w).Encode(map[string]any{"models": []any{map[string]any{"name": "ctx:latest", "model": "ctx:latest"}}})
case "/api/show":
_ = json.NewEncoder(w).Encode(map[string]any{"capabilities": []string{"completion"}, "model_info": map[string]any{"ctx.context_length": 131072}})
case "/v1/chat/completions":
calls.Add(1)
_ = json.NewEncoder(w).Encode(map[string]any{"choices": []any{}, "usage": map[string]any{"prompt_tokens": 3000, "completion_tokens": 100}})
default:
http.NotFound(w, r)
}
}))
}
small := backend(4096, &smallCalls)
defer small.Close()
large := backend(16384, &largeCalls)
defer large.Close()
cfg := &config.Config{
Server: config.ServerConfig{MaxBodyBytes: 1 << 20, MaxRequestDuration: config.Duration(time.Minute)},
Auth: config.AuthConfig{IPBypass: []config.IPBypassConfig{{CIDRs: []string{"127.0.0.1/32"}, Tenant: "t", Subject: "u"}}},
Scheduler: config.SchedulerConfig{GlobalConcurrency: 2, MaxQueue: 8, MaxQueuePerActor: 8, QueueTimeout: config.Duration(time.Second), DefaultTenantWeight: 1, DefaultActorWeight: 1, ComputePaths: []string{"/v1/chat/completions"}},
Cost: config.CostConfig{Default: config.ModelRate{InputCreditsPer1K: 1, OutputCreditsPer1K: 3, CachedInputFactor: 1}, DefaultMaxOutputTokens: 16},
Workers: []config.WorkerConfig{
{Name: "small", URL: small.URL, MaxConcurrent: 1, HealthInterval: config.Duration(time.Hour)},
{Name: "large", URL: large.URL, MaxConcurrent: 1, HealthInterval: config.Duration(time.Hour)},
},
ModelCapabilities: config.ModelCapabilitiesConfig{Mode: "enforce", CacheTTL: config.Duration(time.Hour), ContextGuard: "reject", Context: config.ContextPolicyConfig{MaxRequestedTokens: 32768, DefaultWorkerTokens: 4096, EstimationMarginPercent: 15, VisionReserveTokensPerImage: 2048}},
}
a, err := auth.New(context.Background(), cfg.Auth)
if err != nil {
t.Fatal(err)
}
wp := worker.New(cfg.Workers, "small")
wp.SetModelCapabilitiesConfig(cfg.ModelCapabilities)
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
wp.Start(ctx)
rec, _ := usage.New("", 100, time.Second, nil)
sv := New(cfg, Dependencies{Auth: a, Scheduler: scheduler.NewLocal(2, 8, 8), Quota: quota.Disabled{}, Estimator: cost.New(cfg.Cost), Workers: wp, Proxy: px.New(), Usage: rec, Metrics: metrics.New(), Logger: slog.Default()})
front := httptest.NewServer(sv.Handler())
defer front.Close()
body := `{"model":"ctx:latest","messages":[{"role":"user","content":"` + strings.Repeat("x", 12000) + `"}],"max_tokens":4096}`
resp, err := http.Post(front.URL+"/v1/chat/completions", "application/json", strings.NewReader(body))
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
b, _ := io.ReadAll(resp.Body)
if resp.StatusCode != 200 {
t.Fatalf("status=%d body=%s", resp.StatusCode, b)
}
if got := resp.Header.Get("X-Gateway-Worker"); got != "large" {
t.Fatalf("worker=%q, want large; small=%d large=%d", got, smallCalls.Load(), largeCalls.Load())
}
if smallCalls.Load() != 0 || largeCalls.Load() != 1 {
t.Fatalf("unexpected backend calls small=%d large=%d", smallCalls.Load(), largeCalls.Load())
}
}
func TestOpenAICannotSpoofEffectiveContextWithOptionsNumCtx(t *testing.T) {
front, calls := newContextGuardServer(t, 131072, 4096, "", config.ContextPolicyConfig{MaxRequestedTokens: 32768, DefaultWorkerTokens: 4096, EstimationMarginPercent: 15, VisionReserveTokensPerImage: 2048}, nil)
resp, err := http.Post(front.URL+"/v1/chat/completions", "application/json", strings.NewReader(`{"model":"ctx:latest","messages":[{"role":"user","content":"x"}],"options":{"num_ctx":16384}}`))
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
b, _ := io.ReadAll(resp.Body)
if resp.StatusCode != 400 || !strings.Contains(string(b), "only supported on native Ollama /api endpoints") {
t.Fatalf("status=%d body=%s", resp.StatusCode, b)
}
if calls.Load() != 0 {
t.Fatalf("backend called %d times", calls.Load())
}
}
func TestGenerateSuffixParticipatesInContextGuard(t *testing.T) {
front, calls := newContextGuardServer(t, 131072, 4096, "", config.ContextPolicyConfig{MaxRequestedTokens: 32768, DefaultWorkerTokens: 4096, EstimationMarginPercent: 10, VisionReserveTokensPerImage: 2048}, nil)
body := `{"model":"ctx:latest","prompt":"hello","suffix":"` + strings.Repeat("s", 14000) + `","options":{"num_predict":1024}}`
resp, err := http.Post(front.URL+"/api/generate", "application/json", strings.NewReader(body))
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
b, _ := io.ReadAll(resp.Body)
if resp.StatusCode != 400 || !strings.Contains(string(b), "effective worker context") {
t.Fatalf("status=%d body=%s", resp.StatusCode, b)
}
if calls.Load() != 0 {
t.Fatalf("backend called %d times", calls.Load())
}
}
+344
View File
@@ -0,0 +1,344 @@
package server
import (
"context"
"crypto/sha256"
"encoding/hex"
"net/http"
"sort"
"strconv"
"strings"
"time"
"github.com/example/ollama-fair-gateway/internal/liveflow"
"github.com/example/ollama-fair-gateway/internal/publicui"
"github.com/example/ollama-fair-gateway/internal/worker"
)
type publicDashboardCounts struct {
Workers int `json:"workers"`
HealthyWorkers int `json:"healthy_workers"`
Models int `json:"models"`
Active int `json:"active"`
Queued int64 `json:"queued"`
Routing int `json:"routing"`
Running int64 `json:"running"`
Streaming int `json:"streaming"`
}
type publicResourceMetrics struct {
MemoryUsedBytes int64 `json:"memory_used_bytes,omitempty"`
MemoryTotalBytes int64 `json:"memory_total_bytes,omitempty"`
VRAMUsedBytes int64 `json:"vram_used_bytes,omitempty"`
VRAMTotalBytes int64 `json:"vram_total_bytes,omitempty"`
GPUUtilizationPct float64 `json:"gpu_utilization_percent,omitempty"`
GPUTemperatureC float64 `json:"gpu_temperature_c,omitempty"`
GPUPowerWatts float64 `json:"gpu_power_watts,omitempty"`
}
type publicDashboardWorker struct {
Name string `json:"name"`
Healthy bool `json:"healthy"`
Active int64 `json:"active"`
MaxConcurrent int `json:"max_concurrent"`
AcceptingNew bool `json:"accepting_new"`
Maintenance string `json:"maintenance,omitempty"`
CircuitState string `json:"circuit_state,omitempty"`
LoadedModels []string `json:"loaded_models,omitempty"`
ResourceMetrics *publicResourceMetrics `json:"resource_metrics,omitempty"`
}
type publicDashboardRequest struct {
ID string `json:"id"`
State string `json:"state"`
Model string `json:"model,omitempty"`
Worker string `json:"worker,omitempty"`
QueueMS int64 `json:"queue_ms,omitempty"`
ServiceMS int64 `json:"service_ms,omitempty"`
PromptTokens int64 `json:"prompt_tokens,omitempty"`
CompletionTokens int64 `json:"completion_tokens,omitempty"`
}
type publicDashboardSnapshot struct {
SchemaVersion int `json:"schema_version"`
GeneratedAt time.Time `json:"generated_at"`
Title string `json:"title"`
Subtitle string `json:"subtitle"`
RefreshIntervalMS int64 `json:"refresh_interval_ms"`
UptimeSeconds float64 `json:"uptime_seconds"`
Counts publicDashboardCounts `json:"counts"`
Workers []publicDashboardWorker `json:"workers"`
Requests []publicDashboardRequest `json:"requests"`
}
// handlePublicDashboard serves a separate unauthenticated read-only surface.
// It is intentionally evaluated before authentication. The snapshot builder
// only copies an explicit allow-list of fields; never return admin snapshots
// directly from this handler.
func (s *Server) handlePublicDashboard(w http.ResponseWriter, r *http.Request) bool {
base := s.cfg.PublicDashboard.Path
if base == "" {
base = "/status"
}
if r.URL.Path != base && !strings.HasPrefix(r.URL.Path, base+"/") {
return false
}
if !s.cfg.PublicDashboard.Enabled {
http.NotFound(w, r)
return true
}
if r.URL.Path == base {
http.Redirect(w, r, base+"/", http.StatusTemporaryRedirect)
return true
}
setPublicDashboardHeaders(w)
rel := strings.TrimPrefix(r.URL.Path, base+"/")
if rel == "api/snapshot" {
if r.Method != http.MethodGet && r.Method != http.MethodHead {
w.WriteHeader(http.StatusMethodNotAllowed)
return true
}
w.Header().Set("Cache-Control", "public, max-age=1, stale-while-revalidate=2")
if r.Method == http.MethodHead {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
return true
}
writeJSON(w, http.StatusOK, s.publicDashboardSnapshot())
return true
}
if r.Method != http.MethodGet && r.Method != http.MethodHead {
w.WriteHeader(http.StatusMethodNotAllowed)
return true
}
w.Header().Set("Cache-Control", "public, max-age=300")
h := publicui.Handler()
r2 := r.Clone(r.Context())
r2.URL.Path = "/" + rel
h.ServeHTTP(w, r2)
return true
}
func setPublicDashboardHeaders(w http.ResponseWriter) {
w.Header().Set("X-Content-Type-Options", "nosniff")
w.Header().Set("Referrer-Policy", "no-referrer")
w.Header().Set("X-Frame-Options", "DENY")
w.Header().Set("Permissions-Policy", "camera=(), microphone=(), geolocation=(), payment=(), usb=()")
w.Header().Set("Content-Security-Policy", "default-src 'self'; connect-src 'self'; img-src 'self' data:; style-src 'self' 'unsafe-inline'; script-src 'self'; frame-ancestors 'none'; base-uri 'none'; form-action 'none'")
}
func (s *Server) publicDashboardSnapshot() publicDashboardSnapshot {
cfg := s.cfg.PublicDashboard
st := s.sched.Stats(context.Background())
ws := s.workers.Snapshots()
live := s.live.Snapshot()
sort.Slice(ws, func(i, j int) bool { return ws[i].Name < ws[j].Name })
workerNames := make(map[string]string, len(ws))
for i, w := range ws {
workerNames[w.Name] = publicWorkerDisplayName(cfg.ShowWorkerNames, cfg.WorkerDisplayNames, w.Name, i)
}
modelNames := publicModelDisplayNames(cfg.ShowModelNames, ws, live.Requests)
out := publicDashboardSnapshot{
SchemaVersion: 1,
GeneratedAt: time.Now().UTC(),
Title: cfg.Title,
Subtitle: cfg.Subtitle,
RefreshIntervalMS: cfg.RefreshInterval.Value().Milliseconds(),
UptimeSeconds: time.Since(s.startedAt).Seconds(),
Counts: publicDashboardCounts{
Workers: len(ws),
Active: live.Counts.Active,
Queued: st.Queued,
Routing: live.Counts.Routing,
Running: st.Running,
Streaming: live.Counts.Streaming,
},
}
modelSet := make(map[string]struct{})
for _, w := range ws {
pw := publicDashboardWorker{
Name: workerNames[w.Name],
Healthy: w.Healthy,
Active: w.Active,
MaxConcurrent: w.MaxConcurrent,
AcceptingNew: w.AcceptingNew,
Maintenance: publicMaintenance(w.Maintenance),
CircuitState: publicCircuitState(w.CircuitState),
}
if w.Healthy {
out.Counts.HealthyWorkers++
}
for _, m := range w.LoadedModels {
if m.Name == "" {
continue
}
modelSet[m.Name] = struct{}{}
pw.LoadedModels = append(pw.LoadedModels, modelNames[m.Name])
}
sort.Strings(pw.LoadedModels)
if cfg.ShowResourceMetrics {
memTotal := firstPositive(w.MemoryTotalBytes, w.MemoryCapacityBytes)
vramTotal := firstPositive(w.VRAMTotalBytes, w.VRAMCapacityBytes)
pw.ResourceMetrics = &publicResourceMetrics{
MemoryUsedBytes: w.MemoryUsedBytes,
MemoryTotalBytes: memTotal,
VRAMUsedBytes: w.VRAMUsedBytes,
VRAMTotalBytes: vramTotal,
GPUUtilizationPct: w.GPUUtilizationPct,
GPUTemperatureC: w.GPUTemperatureC,
GPUPowerWatts: w.GPUPowerWatts,
}
}
out.Workers = append(out.Workers, pw)
}
out.Counts.Models = len(modelSet)
requests := publicVisibleRequests(live.Requests, cfg.MaxLiveRequests)
for _, r := range requests {
pr := publicDashboardRequest{
ID: publicRequestID(r.ID),
State: publicRequestState(r.State),
QueueMS: maxInt64(0, r.QueueMS),
ServiceMS: maxInt64(0, r.ServiceMS),
PromptTokens: maxInt64(0, r.PromptTokens),
CompletionTokens: maxInt64(0, r.CompletionTokens),
}
if r.Model != "" {
pr.Model = modelNames[r.Model]
}
if r.Worker != "" {
pr.Worker = workerNames[r.Worker]
if pr.Worker == "" {
pr.Worker = "Worker"
}
}
out.Requests = append(out.Requests, pr)
}
return out
}
func publicWorkerDisplayName(show bool, aliases map[string]string, actual string, index int) string {
if alias := strings.TrimSpace(aliases[actual]); alias != "" {
return alias
}
if show && strings.TrimSpace(actual) != "" {
return actual
}
return "Worker " + twoDigits(index+1)
}
func publicModelDisplayNames(show bool, ws []worker.Snapshot, requests []liveflow.Request) map[string]string {
set := map[string]struct{}{}
for _, w := range ws {
for _, m := range w.LoadedModels {
if m.Name != "" {
set[m.Name] = struct{}{}
}
}
}
for _, r := range requests {
if r.Model != "" {
set[r.Model] = struct{}{}
}
}
names := make([]string, 0, len(set))
for name := range set {
names = append(names, name)
}
sort.Strings(names)
out := make(map[string]string, len(names))
for i, name := range names {
if show {
out[name] = name
} else {
out[name] = "Model " + twoDigits(i+1)
}
}
return out
}
func publicVisibleRequests(in []liveflow.Request, limit int) []liveflow.Request {
if limit <= 0 || len(in) <= limit {
return append([]liveflow.Request(nil), in...)
}
active := make([]liveflow.Request, 0, limit)
recent := make([]liveflow.Request, 0, limit)
for _, r := range in {
switch r.State {
case liveflow.StateCompleted, liveflow.StateCancelled, liveflow.StateFailed:
recent = append(recent, r)
default:
active = append(active, r)
}
}
if len(active) >= limit {
return active[:limit]
}
need := limit - len(active)
if need > len(recent) {
need = len(recent)
}
return append(active, recent[len(recent)-need:]...)
}
func publicRequestID(id string) string {
sum := sha256.Sum256([]byte(id))
return "REQ-" + strings.ToUpper(hex.EncodeToString(sum[:5]))
}
func publicRequestState(state string) string {
switch state {
case liveflow.StateQueued, liveflow.StateRouting, liveflow.StateRunning, liveflow.StateStreaming, liveflow.StateCompleted, liveflow.StateCancelled, liveflow.StateFailed:
return state
default:
return "running"
}
}
func publicMaintenance(v string) string {
switch strings.ToLower(strings.TrimSpace(v)) {
case "drain", "draining":
return "draining"
case "disabled", "offline":
return "disabled"
default:
return "active"
}
}
func publicCircuitState(v string) string {
switch strings.ToLower(strings.TrimSpace(v)) {
case "open", "half-open", "half_open":
return strings.ReplaceAll(v, "_", "-")
default:
return "closed"
}
}
func firstPositive(a, b int64) int64 {
if a > 0 {
return a
}
if b > 0 {
return b
}
return 0
}
func maxInt64(a, b int64) int64 {
if a > b {
return a
}
return b
}
func twoDigits(n int) string {
if n < 10 {
return "0" + string(rune('0'+n))
}
return strconv.Itoa(n)
}
+883
View File
@@ -0,0 +1,883 @@
package server
import (
"bytes"
"context"
"crypto/rand"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"net/url"
"sort"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/example/ollama-fair-gateway/internal/alerts"
"github.com/example/ollama-fair-gateway/internal/auth"
"github.com/example/ollama-fair-gateway/internal/autotune"
"github.com/example/ollama-fair-gateway/internal/batch"
"github.com/example/ollama-fair-gateway/internal/config"
"github.com/example/ollama-fair-gateway/internal/conversation"
"github.com/example/ollama-fair-gateway/internal/cost"
"github.com/example/ollama-fair-gateway/internal/infrastructure"
"github.com/example/ollama-fair-gateway/internal/liveflow"
"github.com/example/ollama-fair-gateway/internal/metrics"
"github.com/example/ollama-fair-gateway/internal/policy"
"github.com/example/ollama-fair-gateway/internal/proxy"
"github.com/example/ollama-fair-gateway/internal/quota"
"github.com/example/ollama-fair-gateway/internal/scheduler"
"github.com/example/ollama-fair-gateway/internal/session"
"github.com/example/ollama-fair-gateway/internal/telemetry"
"github.com/example/ollama-fair-gateway/internal/usage"
"github.com/example/ollama-fair-gateway/internal/warm"
"github.com/example/ollama-fair-gateway/internal/worker"
)
type ConfigStore interface {
Save(*config.Config) error
Delete() error
Path() string
}
type WorkerStateStore interface {
List(context.Context) (map[string]string, error)
Put(context.Context, string, string) error
Delete(context.Context, string) error
Health(context.Context) error
}
type ModelPlacementStore interface {
Get(context.Context, string) (config.ModelPlacementRule, bool, error)
List(context.Context) (map[string]config.ModelPlacementRule, error)
Put(context.Context, string, config.ModelPlacementRule) error
Delete(context.Context, string) error
Health(context.Context) error
}
func unixOrZero(t time.Time) int64 {
if t.IsZero() {
return 0
}
return t.Unix()
}
type Server struct {
cfg *config.Config
auth *auth.Authenticator
sched scheduler.Scheduler
quota quota.Ledger
estimator *cost.Estimator
workers *worker.Pool
proxy *proxy.Proxy
usage *usage.Recorder
metrics *metrics.Registry
live *liveflow.Tracker
infrastructure *infrastructure.Hub
policies policy.Store
sessions session.Store
ops *operationManager
jobs *jobManager
startedAt time.Time
log *slog.Logger
configStore ConfigStore
placementStore ModelPlacementStore
workerStateStore WorkerStateStore
autoTune *autotune.Manager
otel *telemetry.Exporter
warm *warm.Manager
alerts *alerts.Manager
conversations *conversation.Store
batchJobs *batch.Manager
aliasMu sync.Mutex
aliases atomic.Value // immutable map[string]config.ModelAliasConfig
modelAccessMu sync.Mutex
modelAccess atomic.Value // immutable config.ModelAccessConfig
}
type Dependencies struct {
Auth *auth.Authenticator
Scheduler scheduler.Scheduler
Quota quota.Ledger
Estimator *cost.Estimator
Workers *worker.Pool
Proxy *proxy.Proxy
Usage *usage.Recorder
Metrics *metrics.Registry
Live *liveflow.Tracker
Infrastructure *infrastructure.Hub
Policies policy.Store
Sessions session.Store
Logger *slog.Logger
ConfigStore ConfigStore
PlacementStore ModelPlacementStore
WorkerStateStore WorkerStateStore
AutoTune *autotune.Manager
OpenTelemetry *telemetry.Exporter
WarmModels *warm.Manager
Alerts *alerts.Manager
Conversations *conversation.Store
BatchJobs *batch.Manager
}
func New(cfg *config.Config, d Dependencies) *Server {
s := &Server{cfg: cfg, auth: d.Auth, sched: d.Scheduler, quota: d.Quota, estimator: d.Estimator, workers: d.Workers, proxy: d.Proxy, usage: d.Usage, metrics: d.Metrics, live: d.Live, infrastructure: d.Infrastructure, policies: d.Policies, sessions: d.Sessions, ops: newOperationManager(), jobs: newJobManager(), startedAt: time.Now(), log: d.Logger, configStore: d.ConfigStore, placementStore: d.PlacementStore, workerStateStore: d.WorkerStateStore, autoTune: d.AutoTune, otel: d.OpenTelemetry, warm: d.WarmModels, alerts: d.Alerts, conversations: d.Conversations, batchJobs: d.BatchJobs}
s.aliases.Store(cloneModelAliases(cfg.ModelAliases))
s.modelAccess.Store(cloneModelAccess(cfg.ModelAccess))
if s.live == nil {
s.live = liveflow.New(10*time.Second, 512)
}
if s.policies == nil {
s.policies = policy.NewMemory()
}
if s.sessions == nil {
s.sessions = session.NewMemory()
}
s.metrics.SetDynamic(func() metrics.Dynamic {
st := s.sched.Stats(context.Background())
ws := s.workers.Snapshots()
classes := make(map[string]metrics.ServiceClassMetric, len(st.Classes))
for name, cs := range st.Classes {
classes[name] = metrics.ServiceClassMetric{Queued: cs.Queued, Running: cs.Running}
}
wm := make([]metrics.WorkerMetric, 0, len(ws))
for _, w := range ws {
perf := make([]metrics.ModelPerformanceMetric, 0, len(w.Performance))
for _, p := range w.Performance {
perf = append(perf, metrics.ModelPerformanceMetric{Model: p.Model, PromptTPS: p.PromptTPS, OutputTPS: p.OutputTPS, Samples: p.Samples})
}
wm = append(wm, metrics.WorkerMetric{Name: w.Name, Healthy: w.Healthy, Active: w.Active, Max: w.MaxConcurrent, MemoryUsedBytes: w.MemoryUsedBytes, MemoryTotalBytes: w.MemoryTotalBytes, VRAMUsedBytes: w.VRAMUsedBytes, VRAMTotalBytes: w.VRAMTotalBytes, GPUUtilizationPct: w.GPUUtilizationPct, GPUTemperatureC: w.GPUTemperatureC, GPUPowerWatts: w.GPUPowerWatts, ModelActive: w.ModelActive, Performance: perf, CircuitState: w.CircuitState, Maintenance: w.Maintenance})
}
ret := s.usage.RetentionStatus()
d := metrics.Dynamic{Queued: st.Queued, Running: st.Running, OldestQueueWaitSeconds: st.OldestWait.Seconds(), Workers: wm, UsageRawFiles: ret.RawFiles, UsageDailyFiles: ret.DailyFiles, UsageMonthlyFiles: ret.MonthlyFiles, UsageRawBytes: ret.RawBytes, UsageDailyBytes: ret.DailyBytes, UsageMonthlyBytes: ret.MonthlyBytes, UsageLastCompactionUnix: unixOrZero(ret.LastCompaction), UsageLastReclaimedBytes: ret.LastReclaimedBytes, ServiceClasses: classes}
if s.otel != nil {
d.OTelExportedSpans = s.otel.Exported()
d.OTelFailedSpans = s.otel.Failed()
d.OTelDroppedSpans = s.otel.Dropped()
}
if s.warm != nil {
ws := s.warm.Status()
d.WarmEvictionSuggestions = len(ws.Suggestions)
for _, a := range ws.Actions {
if a.Status == "running" {
d.WarmActionsRunning++
}
}
}
if s.alerts != nil {
as := s.alerts.Status()
d.AlertsActive = len(as.Active)
d.AlertsLastEvaluateUnix = unixOrZero(as.LastEvaluate)
}
return d
})
return s
}
func (s *Server) Handler() http.Handler { return http.HandlerFunc(s.serveHTTP) }
func (s *Server) serveHTTP(w http.ResponseWriter, r *http.Request) {
if s.handlePublicDashboard(w, r) {
return
}
if s.handleUIPublic(w, r) {
return
}
switch r.URL.Path {
case "/healthz":
writeJSON(w, 200, map[string]any{"status": "ok"})
return
case "/readyz":
s.ready(w, r)
return
case "/metrics":
if s.cfg.Server.MetricsPublic {
w.Header().Set("Content-Type", "text/plain; version=0.0.4; charset=utf-8")
s.metrics.WritePrometheus(w)
return
}
}
sessionCookieAuth := s.injectUISession(r)
id, err := s.auth.Authenticate(r)
if err != nil {
clientIP := s.auth.ClientIP(r)
w.Header().Set("WWW-Authenticate", `Bearer realm="ollama-gateway"`)
w.Header().Set("X-Gateway-Client-IP", clientIP)
s.log.Warn("authentication rejected", "client_ip", clientIP, "remote_addr", r.RemoteAddr, "method", r.Method, "path", r.URL.Path, "user_agent", r.UserAgent())
writeProtocolError(w, r, 401, "unauthorized", "authentication required")
return
}
r = r.WithContext(auth.WithIdentity(r.Context(), id))
if s.handleModelDiscovery(w, r, id) {
return
}
if r.URL.Path == "/metrics" {
if !id.IsAdmin() {
writeProtocolError(w, r, 403, "forbidden", "gateway:admin scope required")
return
}
w.Header().Set("Content-Type", "text/plain; version=0.0.4; charset=utf-8")
s.metrics.WritePrometheus(w)
return
}
if strings.HasPrefix(r.URL.Path, "/gateway/ui-api/") {
s.uiAPI(w, r, id, sessionCookieAuth)
return
}
if strings.HasPrefix(r.URL.Path, "/gateway/v1/batches") {
s.batchAPI(w, r, id)
return
}
if strings.HasPrefix(r.URL.Path, "/gateway/v1/") {
s.gatewayAPI(w, r, id)
return
}
if !strings.HasPrefix(r.URL.Path, "/api/") && !strings.HasPrefix(r.URL.Path, "/v1/") {
writeProtocolError(w, r, 404, "not_found", "unknown endpoint")
return
}
if s.cfg.Native.ManagementRequiresAdmin && isManagement(r.Method, r.URL.Path) && !id.IsAdmin() {
writeProtocolError(w, r, 403, "forbidden", "model management requires gateway:admin scope")
return
}
s.forward(w, r, id)
}
func (s *Server) handleModelDiscovery(w http.ResponseWriter, r *http.Request, id auth.Identity) bool {
if r.Method != http.MethodGet {
return false
}
switch r.URL.Path {
case "/api/tags":
models, errs := s.workers.Tags(r.Context())
if len(errs) > 0 {
w.Header().Set("X-Gateway-Partial-Errors", strconv.Itoa(len(errs)))
s.log.Warn("partial model discovery", "endpoint", r.URL.Path, "errors", strings.Join(errs, "; "))
if len(models) == 0 {
writeProtocolError(w, r, 503, "worker_unavailable", "unable to retrieve Ollama model tags")
return true
}
}
filtered := make([]worker.TagModel, 0, len(models)+len(s.aliasSnapshot()))
for _, m := range models {
if s.modelAllowed(id, m.Model) {
filtered = append(filtered, m)
}
}
for _, alias := range s.visibleAliases(r.Context(), id) {
filtered = append(filtered, worker.TagModel{Name: alias, Model: alias, Digest: "virtual"})
}
sort.Slice(filtered, func(i, j int) bool { return filtered[i].Model < filtered[j].Model })
writeJSON(w, http.StatusOK, map[string]any{"models": filtered})
return true
case "/api/ps":
loaded := s.workers.Loaded()
filtered := loaded[:0]
for _, m := range loaded {
if s.modelAllowed(id, m.Model) {
filtered = append(filtered, m)
}
}
writeJSON(w, http.StatusOK, map[string]any{"models": filtered})
return true
case "/v1/models":
models, errs := s.workers.Tags(r.Context())
if len(errs) > 0 {
w.Header().Set("X-Gateway-Partial-Errors", strconv.Itoa(len(errs)))
s.log.Warn("partial model discovery", "endpoint", r.URL.Path, "errors", strings.Join(errs, "; "))
if len(models) == 0 {
writeProtocolError(w, r, 503, "worker_unavailable", "unable to retrieve Ollama models")
return true
}
}
data := make([]map[string]any, 0, len(models)+len(s.aliasSnapshot()))
for _, m := range models {
if !s.modelAllowed(id, m.Model) {
continue
}
created := int64(0)
if !m.ModifiedAt.IsZero() {
created = m.ModifiedAt.Unix()
}
data = append(data, map[string]any{"id": m.Model, "object": "model", "created": created, "owned_by": "ollama"})
}
for _, alias := range s.visibleAliases(r.Context(), id) {
data = append(data, map[string]any{"id": alias, "object": "model", "created": int64(0), "owned_by": "ollama-gateway"})
}
sort.Slice(data, func(i, j int) bool { return data[i]["id"].(string) < data[j]["id"].(string) })
writeJSON(w, http.StatusOK, map[string]any{"object": "list", "data": data})
return true
}
return false
}
func (s *Server) ready(w http.ResponseWriter, r *http.Request) {
ctx, cancel := context.WithTimeout(r.Context(), 2*time.Second)
defer cancel()
checks := map[string]string{}
ok := true
checksFn := map[string]func(context.Context) error{"workers": s.workers.Health, "scheduler": s.sched.Health, "quota": s.quota.Health, "usage": s.usage.Health, "policy": s.policies.Health, "api_keys": s.auth.RuntimeStoreHealth}
if s.placementStore != nil {
checksFn["model_placement"] = s.placementStore.Health
}
if s.workerStateStore != nil {
checksFn["worker_state"] = s.workerStateStore.Health
}
for name, fn := range checksFn {
if err := fn(ctx); err != nil {
checks[name] = err.Error()
ok = false
} else {
checks[name] = "ok"
}
}
status := 200
if !ok {
status = 503
}
writeJSON(w, status, map[string]any{"status": map[bool]string{true: "ready", false: "not_ready"}[ok], "checks": checks})
}
func (s *Server) gatewayAPI(w http.ResponseWriter, r *http.Request, id auth.Identity) {
switch r.URL.Path {
case "/gateway/v1/usage/me":
writeJSON(w, 200, s.usage.Actor(r.Context(), id.Tenant, id.Actor()))
case "/gateway/v1/status":
if !id.IsAdmin() {
writeProtocolError(w, r, 403, "forbidden", "gateway:admin scope required")
return
}
st := s.sched.Stats(r.Context())
writeJSON(w, 200, map[string]any{"scheduler": st, "workers": s.workers.Snapshots()})
case "/gateway/v1/usage/tenant":
if !id.IsAdmin() {
writeProtocolError(w, r, 403, "forbidden", "gateway:admin scope required")
return
}
t := r.URL.Query().Get("tenant")
if t == "" {
t = id.Tenant
}
writeJSON(w, 200, s.usage.Tenant(r.Context(), t))
default:
writeProtocolError(w, r, 404, "not_found", "unknown gateway endpoint")
}
}
func (s *Server) forward(w http.ResponseWriter, r *http.Request, id auth.Identity) {
started := time.Now()
ctx := r.Context()
if d := s.cfg.Server.MaxRequestDuration.Value(); d > 0 {
var cancel context.CancelFunc
ctx, cancel = context.WithTimeout(ctx, d)
defer cancel()
}
api := "ollama"
if r.URL.Path == "/v1/messages" {
api = "anthropic"
} else if strings.HasPrefix(r.URL.Path, "/v1/") {
api = "openai"
}
compute := s.isCompute(r.Method, r.URL.Path)
var body []byte
var outboundBody io.Reader
controlModel := ""
if compute || isModelRoutedControlRequest(r.Method, r.URL.Path) {
var err error
body, err = readBody(r, s.cfg.Server.MaxBodyBytes)
if err != nil {
writeProtocolError(w, r, 413, "request_too_large", err.Error())
return
}
if body != nil {
outboundBody = bytes.NewReader(body)
}
if !compute {
controlModel = modelFromBody(body)
}
} else if r.Body != nil && r.Body != http.NoBody {
// Large native management/blob endpoints remain true streaming
// passthroughs. Only compute and small model-introspection requests are
// buffered so they can be routed to a worker that actually owns the model.
outboundBody = r.Body
}
var conversationPlan *responseConversationPlan
if compute && r.URL.Path == "/v1/responses" && s.conversations != nil && s.conversations.Enabled() {
var err error
body, conversationPlan, err = s.prepareResponseConversation(body, id)
if err != nil {
writeProtocolError(w, r, 400, "invalid_previous_response_id", err.Error())
return
}
r.ContentLength = int64(len(body))
outboundBody = bytes.NewReader(body)
w.Header().Set("X-Gateway-Conversations", "enabled")
}
requestedModel := modelFromBody(body)
resolvedModel, aliasName, resolveErr := s.resolveModel(ctx, id, requestedModel)
if resolveErr != nil {
if errors.Is(resolveErr, ErrModelAccessDenied) {
writeProtocolError(w, r, 403, "model_access_denied", resolveErr.Error())
} else {
writeProtocolError(w, r, 404, "model_alias_unavailable", resolveErr.Error())
}
return
}
if resolvedModel != "" && resolvedModel != requestedModel {
var err error
body, err = rewriteModelBody(body, resolvedModel)
if err != nil {
writeProtocolError(w, r, 400, "bad_model_alias", err.Error())
return
}
r.ContentLength = int64(len(body))
outboundBody = bytes.NewReader(body)
if !compute {
controlModel = resolvedModel
}
w.Header().Set("X-Gateway-Model-Alias", aliasName)
w.Header().Set("X-Gateway-Resolved-Model", resolvedModel)
} else if !compute {
controlModel = resolvedModel
}
est := cost.Estimate{}
preflight := modelPreflight{}
serviceClass := ""
serviceCfg := config.ServiceClassConfig{}
if compute {
est = s.estimator.Estimate(r.URL.Path, body)
var ok bool
preflight, ok = s.preflightModel(w, r, body, est)
if !ok {
return
}
var classErr error
serviceClass, serviceCfg, classErr = s.serviceClassFor(r, id)
if classErr != nil {
writeProtocolError(w, r, 403, "service_class_denied", classErr.Error())
return
}
if h := strings.TrimSpace(s.cfg.ServiceClasses.Header); h != "" {
r.Header.Del(h)
}
w.Header().Set("X-Gateway-Service-Class", serviceClass)
}
requestID := newID()
w.Header().Set("X-Request-ID", requestID)
traceCtx := telemetry.Context{}
if compute && s.otel != nil {
traceCtx = s.otel.NewTrace(r.Header.Get("traceparent"))
if traceCtx.Sampled {
w.Header().Set("X-Gateway-Trace-ID", traceCtx.TraceID)
if tp := s.otel.TraceParent(traceCtx); tp != "" {
r.Header.Set("traceparent", tp)
}
}
}
var jobCancel context.CancelCauseFunc
if compute {
ctx, jobCancel = context.WithCancelCause(ctx)
defer jobCancel(nil)
}
var qlease *scheduler.Lease
var reservation quota.Reservation
queueDur := time.Duration(0)
workerName := ""
var workerLease *worker.Lease
var targetURL any
var queueStart, queueEnd, routeEnd time.Time
if compute {
pol := s.policyFor(ctx, id.Tenant)
limits := quota.Limits{ActorCreditsPerMinute: pol.ActorCreditsPerMinute, ActorBurstCredits: pol.ActorBurstCredits, TenantCreditsPerMinute: pol.TenantCreditsPerMinute, TenantBurstCredits: pol.TenantBurstCredits}
decision, err := s.quota.Reserve(ctx, id.Tenant, id.Actor(), est.Credits, limits)
if err != nil {
writeProtocolError(w, r, 503, "quota_unavailable", err.Error())
return
}
if !decision.Allowed {
w.Header().Set("Retry-After", proxy.FormatRetryAfter(decision.RetryAfter))
writeProtocolError(w, r, 429, "quota_exceeded", "compute-credit quota exceeded")
return
}
reservation = decision.Reservation
if s.alerts != nil {
s.alerts.ObserveQuota(id.Tenant, id.Actor(), decision.RemainingActor, decision.RemainingTenant, pol.ActorBurstCredits, pol.TenantBurstCredits)
}
s.live.Begin(liveflow.Request{ID: requestID, Tenant: id.Tenant, Actor: id.Actor(), Application: id.Application, ServiceClass: serviceClass, API: api, Path: r.URL.Path, Model: est.Model, EstimatedCredits: est.Credits, EstimatedPromptTokens: est.InputTokens})
s.jobs.register(jobEntry{ID: requestID, Tenant: id.Tenant, Actor: id.Actor(), Application: id.Application, ServiceClass: serviceClass, Model: est.Model, Path: r.URL.Path, API: api, CreatedAt: time.Now().UTC()}, jobCancel)
defer s.jobs.finish(requestID)
queueTimeout := s.cfg.Scheduler.QueueTimeout.Value()
if d := serviceCfg.MaxQueueWait.Value(); d > 0 && (queueTimeout <= 0 || d < queueTimeout) {
queueTimeout = d
}
queueStart = time.Now()
qlease, err = s.sched.Acquire(ctx, scheduler.Request{Tenant: id.Tenant, Actor: id.Actor(), Cost: est.Credits, TenantWeight: pol.TenantWeight, ActorWeight: pol.ActorWeight, Timeout: queueTimeout, ServiceClass: serviceClass, ClassWeight: serviceCfg.Weight, ClassMaxConcurrent: serviceCfg.MaxConcurrent})
if err != nil {
_ = s.quota.Reconcile(context.Background(), reservation, 0)
status := 503
if errors.Is(err, scheduler.ErrQueueFull) || errors.Is(err, scheduler.ErrActorQueueFull) {
status = 429
w.Header().Set("Retry-After", "1")
writeProtocolError(w, r, status, "queue_full", err.Error())
} else if errors.Is(err, context.DeadlineExceeded) {
writeProtocolError(w, r, status, "queue_timeout", "request exceeded queue timeout")
} else if errors.Is(err, context.Canceled) {
status = 499
if isAdminJobCancel(ctx) {
writeProtocolError(w, r, status, "request_cancelled", "request cancelled by administrator")
} else {
writeProtocolError(w, r, status, "request_cancelled", "request cancelled")
}
s.live.Cancel(requestID, status, 0, cost.Usage{}, 0)
return
} else {
writeProtocolError(w, r, status, "scheduler_unavailable", err.Error())
}
s.live.Drop(requestID, status)
return
}
defer qlease.Release()
queueEnd = time.Now()
s.live.MarkRouting(requestID, "", qlease.Wait)
workerLease, err = s.workers.AcquireAllowed(ctx, est.Model, preflight.AllowedWorkers, preflight.RequestedContext)
if err != nil {
_ = s.quota.Reconcile(context.Background(), reservation, 0)
if errors.Is(err, context.Canceled) {
writeProtocolError(w, r, 499, "request_cancelled", "request cancelled")
s.live.Cancel(requestID, 499, 0, cost.Usage{}, 0)
return
}
if errors.Is(err, worker.ErrModelPlacementBlocked) {
writeProtocolError(w, r, 403, "model_placement_denied", err.Error())
s.live.Drop(requestID, 403)
return
}
if errors.Is(err, worker.ErrModelNotInstalled) {
writeProtocolError(w, r, 404, "model_not_found", err.Error())
s.live.Drop(requestID, 404)
return
}
writeProtocolError(w, r, 503, "worker_unavailable", err.Error())
s.live.Drop(requestID, 503)
return
}
defer func() {
if workerLease != nil {
workerLease.Release()
}
}()
workerName = workerLease.Name()
if s.warm != nil {
s.warm.Touch(workerName, est.Model)
}
routeEnd = time.Now()
s.jobs.setWorker(requestID, workerName)
targetURL = workerLease.URL()
queueDur = time.Since(started)
s.live.MarkRouting(requestID, workerName, queueDur)
w.Header().Set("X-Gateway-Queue-Ms", strconv.FormatInt(queueDur.Milliseconds(), 10))
w.Header().Set("X-Gateway-Worker", workerName)
w.Header().Set("X-Gateway-Estimated-Credits", strconv.FormatFloat(est.Credits, 'f', 4, 64))
} else {
u, name, err := s.workers.ControlForModel(controlModel)
if err != nil {
if errors.Is(err, worker.ErrModelPlacementBlocked) {
writeProtocolError(w, r, 403, "model_placement_denied", err.Error())
return
}
if errors.Is(err, worker.ErrModelNotInstalled) {
writeProtocolError(w, r, 404, "model_not_found", err.Error())
return
}
writeProtocolError(w, r, 503, "worker_unavailable", err.Error())
return
}
targetURL = u
workerName = name
}
target := targetURL.(*url.URL)
serviceStart := time.Now()
var progress proxy.ProgressFunc
if compute {
s.live.MarkRunning(requestID)
progress = func(bytesOut int64, u cost.Usage) { s.live.Progress(requestID, bytesOut, u) }
}
attempts := 1
excluded := map[string]bool{}
reportedFailures := map[string]bool{}
forwardBody := outboundBody
if compute && body != nil {
forwardBody = bytes.NewReader(body)
}
forwardRequest := func(target *url.URL, requestBody io.Reader) proxy.Result {
if conversationPlan != nil && conversationPlan.Store {
return s.proxy.ForwardCapture(ctx, w, r, target, requestBody, api, est.InputTokens, s.cfg.Conversations.MaxContentBytes, progress)
}
return s.proxy.Forward(ctx, w, r, target, requestBody, api, est.InputTokens, progress)
}
res := forwardRequest(target, forwardBody)
for compute && s.cfg.Reliability.Enabled && res.Err != nil && !res.Started && attempts < s.cfg.Reliability.RetryAttempts && ctx.Err() == nil {
opened := s.workers.ReportResult(workerName, true, res.Err.Error())
s.metrics.RecordUpstreamFailure(workerName, "transport")
if opened {
s.metrics.RecordCircuitOpen(workerName)
}
reportedFailures[workerName] = true
excluded[workerName] = true
if workerLease != nil {
workerLease.Release()
workerLease = nil
}
if d := s.cfg.Reliability.RetryBackoff.Value(); d > 0 {
select {
case <-ctx.Done():
break
case <-time.After(d):
}
}
next, err := s.workers.AcquireAllowedExcluding(ctx, est.Model, preflight.AllowedWorkers, excluded, preflight.RequestedContext)
if err != nil {
break
}
workerLease = next
workerName = next.Name()
if s.warm != nil {
s.warm.Touch(workerName, est.Model)
}
target = next.URL()
s.jobs.setWorker(requestID, workerName)
s.live.MarkRouting(requestID, workerName, queueDur)
attempts++
s.metrics.RecordRetry(workerName)
w.Header().Set("X-Gateway-Retry-Count", strconv.Itoa(attempts-1))
w.Header().Set("X-Gateway-Worker", workerName)
res = forwardRequest(target, bytes.NewReader(body))
}
if compute && workerName != "" {
failedWorker := (res.Err != nil && !errors.Is(ctx.Err(), context.Canceled)) || res.Status >= 500
errText := ""
if res.Err != nil {
errText = res.Err.Error()
} else if res.Status >= 500 {
errText = fmt.Sprintf("HTTP %d", res.Status)
}
if !failedWorker || !reportedFailures[workerName] {
opened := s.workers.ReportResult(workerName, failedWorker, errText)
if failedWorker {
class := "http_5xx"
if res.Err != nil {
class = "transport"
}
s.metrics.RecordUpstreamFailure(workerName, class)
}
if opened {
s.metrics.RecordCircuitOpen(workerName)
}
}
}
serviceDur := time.Since(serviceStart)
cancelled := compute && errors.Is(ctx.Err(), context.Canceled)
if compute && res.Status < 400 && res.Usage.PromptEvalNS+res.Usage.EvalNS == 0 {
// Ollama's native API exposes exact eval durations. OpenAI-compatible
// responses generally do not, so use worker-slot wall time as a
// conservative compute-duration approximation when duration credits are enabled.
res.Usage.EvalNS = serviceDur.Nanoseconds()
res.Usage.Approximate = true
}
if res.Err != nil && !res.Started {
if cancelled {
writeProtocolError(w, r, 499, "request_cancelled", "request cancelled")
} else {
writeProtocolError(w, r, 502, "backend_error", proxy.BackendError(res.Err))
}
}
actual := 0.0
if compute && res.Status < 400 {
actual = s.estimator.Actual(est.Model, res.Usage)
if actual <= 0 {
actual = est.Credits
}
}
if compute {
if err := s.quota.Reconcile(context.Background(), reservation, actual); err != nil {
s.log.Warn("quota reconcile failed", "request_id", requestID, "error", err)
}
}
status := res.Status
if cancelled {
status = 499
} else if status == 0 {
status = 502
}
if compute {
if workerName != "" && status < 500 {
s.workers.Observe(workerName, est.Model, res.Usage.PromptTokens, res.Usage.CompletionTokens, res.Usage.PromptEvalNS, res.Usage.EvalNS, serviceDur)
}
if cancelled {
s.live.Cancel(requestID, status, actual, res.Usage, serviceDur)
} else {
s.live.Finish(requestID, status, actual, res.Usage, serviceDur)
}
}
if compute && r.URL.Path == "/v1/responses" && status < 400 {
s.persistResponseConversation(conversationPlan, res.Captured, res.CaptureTruncated, id, est.Model)
}
s.metrics.Record(api, status, queueDur, serviceDur, res.Usage.PromptTokens, res.Usage.CompletionTokens, actual, res.BytesIn, res.BytesOut)
if compute && s.otel != nil && traceCtx.Sampled {
intervals := []telemetry.Interval{}
if !queueStart.IsZero() {
intervals = append(intervals, telemetry.Interval{Name: "gateway.admission", Start: started, End: queueStart})
}
if !queueStart.IsZero() && !queueEnd.IsZero() {
intervals = append(intervals, telemetry.Interval{Name: "gateway.queue", Start: queueStart, End: queueEnd, Attrs: map[string]any{"ollama.gateway.service_class": serviceClass}})
}
if !queueEnd.IsZero() && !routeEnd.IsZero() {
intervals = append(intervals, telemetry.Interval{Name: "gateway.route", Start: queueEnd, End: routeEnd, Attrs: map[string]any{"server.address": workerName}})
}
intervals = append(intervals, telemetry.Interval{Name: "ollama.upstream", Start: serviceStart, End: serviceStart.Add(serviceDur), Attrs: map[string]any{"server.address": workerName}})
errText := ""
if res.Err != nil {
errText = res.Err.Error()
}
s.otel.Record(telemetry.Record{Trace: traceCtx, RequestID: requestID, API: api, Path: r.URL.Path, Tenant: id.Tenant, Actor: id.Actor(), Application: id.Application, ServiceClass: serviceClass, Model: est.Model, Alias: aliasName, Worker: workerName, Started: started, Finished: time.Now().UTC(), FirstByte: res.FirstByte, Status: status, PromptTokens: res.Usage.PromptTokens, OutputTokens: res.Usage.CompletionTokens, CachedTokens: res.Usage.CachedPromptTokens, Credits: actual, Error: errText, Intervals: intervals})
}
s.usage.Record(usage.Event{ID: requestID, Time: time.Now().UTC(), Tenant: id.Tenant, Subject: id.Subject, Actor: id.Actor(), Application: id.Application, ServiceClass: serviceClass, AuthType: id.AuthType, ClientIP: id.ClientIP, API: api, Path: r.URL.Path, Model: est.Model, Worker: workerName, Status: status, QueueMS: queueDur.Milliseconds(), ServiceMS: serviceDur.Milliseconds(), EstimatedCredits: est.Credits, ActualCredits: actual, Usage: res.Usage, BytesIn: res.BytesIn, BytesOut: res.BytesOut})
s.log.Info("request", "request_id", requestID, "tenant", id.Tenant, "subject", id.Subject, "api", api, "path", r.URL.Path, "model", est.Model, "worker", workerName, "service_class", serviceClass, "status", status, "queue_ms", queueDur.Milliseconds(), "service_ms", serviceDur.Milliseconds(), "credits", actual)
}
func (s *Server) serviceClassFor(r *http.Request, id auth.Identity) (string, config.ServiceClassConfig, error) {
name := strings.TrimSpace(id.ServiceClass)
if name == "" {
name = strings.TrimSpace(s.cfg.ServiceClasses.Default)
}
if name == "" {
name = "interactive"
}
if len(s.cfg.ServiceClasses.Classes) == 0 {
return name, config.ServiceClassConfig{Weight: 1, MaxQueueWait: s.cfg.Scheduler.QueueTimeout}, nil
}
header := strings.TrimSpace(s.cfg.ServiceClasses.Header)
if header != "" {
if requested := strings.TrimSpace(r.Header.Get(header)); requested != "" && requested != name {
if !id.HasScope(s.cfg.ServiceClasses.OverrideScope) {
return "", config.ServiceClassConfig{}, fmt.Errorf("service class override requires scope %s", s.cfg.ServiceClasses.OverrideScope)
}
name = requested
}
}
cfg, ok := s.cfg.ServiceClasses.Classes[name]
if !ok {
return "", config.ServiceClassConfig{}, fmt.Errorf("unknown service class %q", name)
}
return name, cfg, nil
}
func readBody(r *http.Request, maxBytes int64) ([]byte, error) {
if r.Body == nil {
return nil, nil
}
defer r.Body.Close()
b, err := io.ReadAll(io.LimitReader(r.Body, maxBytes+1))
if err != nil {
return nil, err
}
if int64(len(b)) > maxBytes {
return nil, fmt.Errorf("request body exceeds %d bytes", maxBytes)
}
return b, nil
}
func (s *Server) isCompute(method, path string) bool {
if method != http.MethodPost {
return false
}
for _, p := range s.cfg.Scheduler.ComputePaths {
if path == p {
return true
}
}
return false
}
func isModelRoutedControlRequest(method, path string) bool {
if method != http.MethodPost && method != http.MethodDelete {
return false
}
switch path {
case "/api/show", "/api/delete":
return true
default:
return false
}
}
func modelFromBody(body []byte) string {
if len(body) == 0 {
return ""
}
var v struct {
Model string `json:"model"`
Name string `json:"name"`
}
if json.Unmarshal(body, &v) != nil {
return ""
}
if strings.TrimSpace(v.Model) != "" {
return strings.TrimSpace(v.Model)
}
return strings.TrimSpace(v.Name)
}
func isManagement(method, path string) bool {
if !strings.HasPrefix(path, "/api/") {
return false
}
for _, p := range []string{"/api/pull", "/api/push", "/api/create", "/api/copy", "/api/delete", "/api/stop", "/api/blobs"} {
if path == p || strings.HasPrefix(path, p+"/") {
return true
}
}
return false
}
func writeProtocolError(w http.ResponseWriter, r *http.Request, status int, code, msg string) {
if r != nil && r.URL.Path == "/v1/messages" {
typ := "api_error"
switch status {
case 400:
typ = "invalid_request_error"
case 401:
typ = "authentication_error"
case 403:
typ = "permission_error"
case 404:
typ = "not_found_error"
case 429:
typ = "rate_limit_error"
case 500, 502, 503, 504:
typ = "api_error"
}
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(map[string]any{"type": "error", "error": map[string]any{"type": typ, "message": msg}})
return
}
// Ollama's native API uses a string-valued "error" field. OpenWebUI
// relies on that shape when verifying Ollama connections; returning the
// OpenAI-style nested object causes UI messages such as "[object Object]".
if r != nil && strings.HasPrefix(r.URL.Path, "/api/") {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(map[string]any{"error": msg})
return
}
proxy.WriteJSONError(w, status, code, msg)
}
func writeJSON(w http.ResponseWriter, status int, v any) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(v)
}
func newID() string { b := make([]byte, 16); _, _ = rand.Read(b); return hex.EncodeToString(b) }
+297
View File
@@ -0,0 +1,297 @@
package server
import (
"context"
"github.com/example/ollama-fair-gateway/internal/auth"
"github.com/example/ollama-fair-gateway/internal/config"
"github.com/example/ollama-fair-gateway/internal/cost"
"github.com/example/ollama-fair-gateway/internal/metrics"
px "github.com/example/ollama-fair-gateway/internal/proxy"
"github.com/example/ollama-fair-gateway/internal/quota"
"github.com/example/ollama-fair-gateway/internal/scheduler"
"github.com/example/ollama-fair-gateway/internal/usage"
"github.com/example/ollama-fair-gateway/internal/worker"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
)
func TestNativeStreamingPassthroughAndMetering(t *testing.T) {
backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path == "/api/ps" {
w.Header().Set("Content-Type", "application/json")
io.WriteString(w, `{"models":[{"name":"qwen3:8b"}]}`)
return
}
if r.URL.Path == "/api/chat" {
w.Header().Set("Content-Type", "application/x-ndjson")
io.WriteString(w, "{\"message\":{\"content\":\"hi\"},\"done\":false}\n{\"done\":true,\"prompt_eval_count\":10,\"eval_count\":2,\"eval_duration\":1000}\n")
return
}
w.WriteHeader(404)
}))
defer backend.Close()
cfg := &config.Config{Server: config.ServerConfig{MaxBodyBytes: 1 << 20, MaxRequestDuration: config.Duration(time.Minute), MetricsPublic: true}, Auth: config.AuthConfig{IPBypass: []config.IPBypassConfig{{CIDRs: []string{"127.0.0.1/32"}, Tenant: "test", Subject: "u"}}}, Scheduler: config.SchedulerConfig{GlobalConcurrency: 1, MaxQueue: 8, MaxQueuePerActor: 8, QueueTimeout: config.Duration(time.Second), DefaultTenantWeight: 1, DefaultActorWeight: 1}, Cost: config.CostConfig{Default: config.ModelRate{InputCreditsPer1K: 1, OutputCreditsPer1K: 3}, DefaultMaxOutputTokens: 16}, Workers: []config.WorkerConfig{{Name: "w", URL: backend.URL, MaxConcurrent: 1, HealthInterval: config.Duration(time.Hour)}}}
a, err := auth.New(context.Background(), cfg.Auth)
if err != nil {
t.Fatal(err)
}
wp := worker.New(cfg.Workers, "")
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
wp.Start(ctx)
met := metrics.New()
rec, _ := usage.New("", 100, time.Second, nil)
sv := New(cfg, Dependencies{Auth: a, Scheduler: scheduler.NewLocal(1, 8, 8), Quota: quota.Disabled{}, Estimator: cost.New(cfg.Cost), Workers: wp, Proxy: px.New(), Usage: rec, Metrics: met, Logger: slog.Default()})
front := httptest.NewServer(sv.Handler())
defer front.Close()
resp, err := http.Post(front.URL+"/api/chat", "application/json", strings.NewReader(`{"model":"qwen3:8b","messages":[{"role":"user","content":"x"}]}`))
if err != nil {
t.Fatal(err)
}
b, _ := io.ReadAll(resp.Body)
resp.Body.Close()
want := "{\"message\":{\"content\":\"hi\"},\"done\":false}\n{\"done\":true,\"prompt_eval_count\":10,\"eval_count\":2,\"eval_duration\":1000}\n"
if string(b) != want {
t.Fatalf("body changed:\n%s", b)
}
var s usage.Summary
deadline := time.Now().Add(time.Second)
for time.Now().Before(deadline) {
s = rec.Actor(context.Background(), "test", "u")
if s.PromptTokens == 10 && s.CompletionTokens == 2 {
break
}
time.Sleep(time.Millisecond)
}
if s.PromptTokens != 10 || s.CompletionTokens != 2 {
t.Fatalf("usage not metered: %#v", s)
}
}
func TestNativeNonComputeRequestBodyStreamsPastComputeLimit(t *testing.T) {
const bodySize = 256 << 10
gotSize := make(chan int64, 1)
backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/ps":
w.Header().Set("Content-Type", "application/json")
io.WriteString(w, `{"models":[]}`)
case "/api/blobs/sha256:test":
n, _ := io.Copy(io.Discard, r.Body)
gotSize <- n
w.WriteHeader(http.StatusCreated)
default:
w.WriteHeader(http.StatusNotFound)
}
}))
defer backend.Close()
cfg := &config.Config{
Server: config.ServerConfig{MaxBodyBytes: 32, MaxRequestDuration: config.Duration(time.Minute)},
Auth: config.AuthConfig{IPBypass: []config.IPBypassConfig{{CIDRs: []string{"127.0.0.1/32"}, Tenant: "test", Subject: "u", Scopes: []string{"gateway:admin"}}}},
Scheduler: config.SchedulerConfig{GlobalConcurrency: 1, MaxQueue: 8, MaxQueuePerActor: 8, QueueTimeout: config.Duration(time.Second), DefaultTenantWeight: 1, DefaultActorWeight: 1, ComputePaths: []string{"/api/chat"}},
Cost: config.CostConfig{Default: config.ModelRate{InputCreditsPer1K: 1, OutputCreditsPer1K: 3}, DefaultMaxOutputTokens: 16},
Workers: []config.WorkerConfig{{Name: "w", URL: backend.URL, MaxConcurrent: 1, HealthInterval: config.Duration(time.Hour)}},
Native: config.NativeConfig{ManagementRequiresAdmin: true, ControlWorker: "w"},
}
a, err := auth.New(context.Background(), cfg.Auth)
if err != nil {
t.Fatal(err)
}
wp := worker.New(cfg.Workers, "w")
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
wp.Start(ctx)
met := metrics.New()
rec, _ := usage.New("", 100, time.Second, nil)
sv := New(cfg, Dependencies{Auth: a, Scheduler: scheduler.NewLocal(1, 8, 8), Quota: quota.Disabled{}, Estimator: cost.New(cfg.Cost), Workers: wp, Proxy: px.New(), Usage: rec, Metrics: met, Logger: slog.Default()})
front := httptest.NewServer(sv.Handler())
defer front.Close()
req, _ := http.NewRequest(http.MethodPut, front.URL+"/api/blobs/sha256:test", io.LimitReader(strings.NewReader(strings.Repeat("x", bodySize)), bodySize))
resp, err := http.DefaultClient.Do(req)
if err != nil {
t.Fatal(err)
}
io.Copy(io.Discard, resp.Body)
resp.Body.Close()
if resp.StatusCode != http.StatusCreated {
t.Fatalf("status=%d", resp.StatusCode)
}
select {
case n := <-gotSize:
if n != bodySize {
t.Fatalf("backend received %d bytes, want %d", n, bodySize)
}
case <-time.After(time.Second):
t.Fatal("backend did not receive streamed body")
}
}
func TestModelAliasAndTenantACL(t *testing.T) {
seen := make(chan string, 1)
backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/ps":
io.WriteString(w, `{"models":[{"name":"real:1","model":"real:1"}]}`)
case "/api/tags":
io.WriteString(w, `{"models":[{"name":"real:1","model":"real:1"}]}`)
case "/api/show":
io.WriteString(w, `{"capabilities":["completion"]}`)
case "/api/chat":
b, _ := io.ReadAll(r.Body)
seen <- string(b)
w.Header().Set("Content-Type", "application/x-ndjson")
io.WriteString(w, `{"done":true,"prompt_eval_count":1,"eval_count":1}`+"\n")
default:
http.NotFound(w, r)
}
}))
defer backend.Close()
visible := true
cfg := &config.Config{
Server: config.ServerConfig{MaxBodyBytes: 1 << 20, MaxRequestDuration: config.Duration(time.Minute)},
Auth: config.AuthConfig{IPBypass: []config.IPBypassConfig{{CIDRs: []string{"127.0.0.1/32"}, Tenant: "t", Subject: "u"}}},
Scheduler: config.SchedulerConfig{GlobalConcurrency: 1, MaxQueue: 8, MaxQueuePerActor: 8, QueueTimeout: config.Duration(time.Second), DefaultTenantWeight: 1, DefaultActorWeight: 1, ComputePaths: []string{"/api/chat"}},
Cost: config.CostConfig{Default: config.ModelRate{InputCreditsPer1K: 1, OutputCreditsPer1K: 3}, DefaultMaxOutputTokens: 16},
Workers: []config.WorkerConfig{{Name: "w", URL: backend.URL, MaxConcurrent: 1, HealthInterval: config.Duration(time.Hour)}},
ModelAliases: map[string]config.ModelAliasConfig{"fast": {Models: []string{"real:1"}, Visible: &visible}},
ModelAccess: config.ModelAccessConfig{Default: config.ModelAccessRule{Mode: "whitelist", AllowedModels: []string{"fast"}}},
ModelCapabilities: config.ModelCapabilitiesConfig{Mode: "enforce", CacheTTL: config.Duration(time.Hour), ContextGuard: "off"},
}
a, err := auth.New(context.Background(), cfg.Auth)
if err != nil {
t.Fatal(err)
}
wp := worker.New(cfg.Workers, "")
wp.SetModelCapabilitiesConfig(cfg.ModelCapabilities)
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
wp.Start(ctx)
rec, _ := usage.New("", 100, time.Second, nil)
sv := New(cfg, Dependencies{Auth: a, Scheduler: scheduler.NewLocal(1, 8, 8), Quota: quota.Disabled{}, Estimator: cost.New(cfg.Cost), Workers: wp, Proxy: px.New(), Usage: rec, Metrics: metrics.New(), Logger: slog.Default()})
front := httptest.NewServer(sv.Handler())
defer front.Close()
resp, err := http.Get(front.URL + "/api/tags")
if err != nil {
t.Fatal(err)
}
b, _ := io.ReadAll(resp.Body)
resp.Body.Close()
if !strings.Contains(string(b), `"model":"fast"`) || strings.Contains(string(b), `"model":"real:1"`) {
t.Fatalf("unexpected discovery: %s", b)
}
resp, err = http.Post(front.URL+"/api/chat", "application/json", strings.NewReader(`{"model":"fast","messages":[{"role":"user","content":"x"}]}`))
if err != nil {
t.Fatal(err)
}
io.Copy(io.Discard, resp.Body)
resp.Body.Close()
if resp.StatusCode != 200 {
t.Fatalf("alias status=%d", resp.StatusCode)
}
select {
case body := <-seen:
if !strings.Contains(body, `"model":"real:1"`) {
t.Fatalf("backend body=%s", body)
}
case <-time.After(time.Second):
t.Fatal("backend not called")
}
resp, err = http.Post(front.URL+"/api/chat", "application/json", strings.NewReader(`{"model":"real:1","messages":[]}`))
if err != nil {
t.Fatal(err)
}
io.Copy(io.Discard, resp.Body)
resp.Body.Close()
if resp.StatusCode != 403 {
t.Fatalf("real model should be ACL denied, got %d", resp.StatusCode)
}
}
func TestSafeRetryBeforeResponseAndCircuitOpen(t *testing.T) {
bad := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/ps":
io.WriteString(w, `{"models":[]}`)
case "/api/tags":
io.WriteString(w, `{"models":[{"name":"m","model":"m"}]}`)
case "/api/show":
io.WriteString(w, `{"capabilities":["completion"]}`)
case "/api/chat":
c, _, _ := w.(http.Hijacker).Hijack()
_ = c.Close()
default:
http.NotFound(w, r)
}
}))
defer bad.Close()
good := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/ps":
io.WriteString(w, `{"models":[]}`)
case "/api/tags":
io.WriteString(w, `{"models":[{"name":"m","model":"m"}]}`)
case "/api/show":
io.WriteString(w, `{"capabilities":["completion"]}`)
case "/api/chat":
w.Header().Set("Content-Type", "application/x-ndjson")
io.WriteString(w, `{"done":true,"prompt_eval_count":2,"eval_count":1}`+"\n")
default:
http.NotFound(w, r)
}
}))
defer good.Close()
cfg := &config.Config{Server: config.ServerConfig{MaxBodyBytes: 1 << 20, MaxRequestDuration: config.Duration(time.Minute)}, Auth: config.AuthConfig{IPBypass: []config.IPBypassConfig{{CIDRs: []string{"127.0.0.1/32"}, Tenant: "t", Subject: "u"}}}, Scheduler: config.SchedulerConfig{GlobalConcurrency: 1, MaxQueue: 8, MaxQueuePerActor: 8, QueueTimeout: config.Duration(time.Second), DefaultTenantWeight: 1, DefaultActorWeight: 1, ComputePaths: []string{"/api/chat"}}, Cost: config.CostConfig{Default: config.ModelRate{InputCreditsPer1K: 1, OutputCreditsPer1K: 3}, DefaultMaxOutputTokens: 8}, Workers: []config.WorkerConfig{{Name: "bad", URL: bad.URL, MaxConcurrent: 1, HealthInterval: config.Duration(time.Hour)}, {Name: "good", URL: good.URL, MaxConcurrent: 1, HealthInterval: config.Duration(time.Hour)}}, Reliability: config.ReliabilityConfig{Enabled: true, FailureThreshold: 1, OpenDuration: config.Duration(time.Hour), RetryAttempts: 2}, ModelCapabilities: config.ModelCapabilitiesConfig{Mode: "off", ContextGuard: "off"}}
a, _ := auth.New(context.Background(), cfg.Auth)
wp := worker.New(cfg.Workers, "")
wp.SetReliabilityConfig(cfg.Reliability)
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
wp.Start(ctx)
rec, _ := usage.New("", 100, time.Second, nil)
sv := New(cfg, Dependencies{Auth: a, Scheduler: scheduler.NewLocal(1, 8, 8), Quota: quota.Disabled{}, Estimator: cost.New(cfg.Cost), Workers: wp, Proxy: px.New(), Usage: rec, Metrics: metrics.New(), Logger: slog.Default()})
front := httptest.NewServer(sv.Handler())
defer front.Close()
resp, err := http.Post(front.URL+"/api/chat", "application/json", strings.NewReader(`{"model":"m","messages":[]}`))
if err != nil {
t.Fatal(err)
}
io.Copy(io.Discard, resp.Body)
resp.Body.Close()
if resp.StatusCode != 200 {
t.Fatalf("status=%d", resp.StatusCode)
}
if got := resp.Header.Get("X-Gateway-Retry-Count"); got != "1" {
t.Fatalf("retry header=%q", got)
}
if got := resp.Header.Get("X-Gateway-Worker"); got != "good" {
t.Fatalf("worker=%q", got)
}
for _, snap := range wp.Snapshots() {
if snap.Name == "bad" && snap.CircuitState != "open" {
t.Fatalf("bad circuit=%s", snap.CircuitState)
}
}
}
func TestAPIKeyModelACLCanNarrowButNotWidenTenantACL(t *testing.T) {
cfg := &config.Config{}
cfg.ModelAccess = config.ModelAccessConfig{Default: config.ModelAccessRule{Mode: "whitelist", AllowedModels: []string{"fast", "qwen3:8b"}}}
s := &Server{cfg: cfg}
id := auth.Identity{Tenant: "team", ModelACLSet: true, ModelAccess: config.ModelAccessRule{Mode: "whitelist", AllowedModels: []string{"fast", "gemma4:*"}}}
if !s.modelAllowed(id, "fast") {
t.Fatal("expected intersection to allow fast")
}
if s.modelAllowed(id, "qwen3:8b") {
t.Fatal("API key ACL should narrow tenant ACL")
}
if s.modelAllowed(id, "gemma4:latest") {
t.Fatal("API key ACL must not widen tenant ACL")
}
}
+288
View File
@@ -0,0 +1,288 @@
package server
import (
"archive/zip"
"context"
"fmt"
"io"
"net/http"
"os"
"path/filepath"
"sort"
"strings"
"time"
"github.com/example/ollama-fair-gateway/internal/state"
)
type persistentSaver interface {
SavePersistent(string) error
}
type storageFileStatus struct {
Name string `json:"name"`
Kind string `json:"kind"`
Path string `json:"path"`
Exists bool `json:"exists"`
Size int64 `json:"size_bytes,omitempty"`
Modified time.Time `json:"modified_at,omitempty"`
}
func fileStatus(name, kind, path string) storageFileStatus {
x := storageFileStatus{Name: name, Kind: kind, Path: path}
st, err := os.Stat(path)
if err == nil && st.Mode().IsRegular() {
x.Exists = true
x.Size = st.Size()
x.Modified = st.ModTime().UTC()
}
return x
}
func (s *Server) storageStatus() map[string]any {
paths := state.Resolve(s.cfg.Storage)
files := []storageFileStatus{
fileStatus("configuration override", "config", paths.Config),
fileStatus("API keys", "security", paths.APIKeys),
fileStatus("tenant policies", "policy", paths.Policies),
fileStatus("metrics snapshot", "metrics", paths.Metrics),
fileStatus("quota buckets", "quota", paths.Quota),
fileStatus("worker performance", "routing", paths.WorkerPerformance),
fileStatus("model placement", "routing", paths.ModelPlacement),
fileStatus("worker runtime state", "routing", paths.WorkerState),
fileStatus("warm model policies", "capacity", paths.WarmModels),
fileStatus("alerts history", "alerts", paths.Alerts),
fileStatus("encrypted conversations", "content", paths.Conversations),
fileStatus("durable batch jobs", "batch", paths.BatchJobs),
}
usageDir := s.cfg.Usage.JournalDir
usageFiles := 0
var usageBytes int64
var usageNewest time.Time
if usageDir != "" {
matches, _ := filepath.Glob(filepath.Join(usageDir, "usage-*.jsonl"))
for _, p := range matches {
if st, err := os.Stat(p); err == nil && st.Mode().IsRegular() {
usageFiles++
usageBytes += st.Size()
if st.ModTime().After(usageNewest) {
usageNewest = st.ModTime().UTC()
}
}
}
}
batchFiles := 0
var batchBytes int64
var batchNewest time.Time
_ = filepath.Walk(paths.BatchDir, func(path string, info os.FileInfo, err error) error {
if err == nil && info.Mode().IsRegular() {
batchFiles++
batchBytes += info.Size()
if info.ModTime().After(batchNewest) {
batchNewest = info.ModTime().UTC()
}
}
return nil
})
retention := s.usage.RetentionStatus()
var total int64
for _, f := range files {
total += f.Size
}
total += usageBytes + retention.DailyBytes + retention.MonthlyBytes + batchBytes
return map[string]any{
"mode": "local-persistent",
"data_dir": paths.DataDir,
"flush_interval": s.cfg.Storage.FlushInterval.Value().String(),
"config_override_active": s.configStore != nil && fileStatus("", "", paths.Config).Exists,
"files": files,
"usage": map[string]any{
"directory": usageDir,
"files": usageFiles,
"size_bytes": usageBytes,
"modified_at": usageNewest,
"retention": retention,
},
"batch": map[string]any{
"enabled": s.batchJobs != nil && s.batchJobs.Enabled(),
"directory": paths.BatchDir,
"files": batchFiles,
"size_bytes": batchBytes,
"modified_at": batchNewest,
},
"conversations": func() any {
if s.conversations == nil {
return map[string]any{"enabled": false}
}
return s.conversations.Status()
}(),
"total_bytes": total,
"volatile": []string{
"active transient inference jobs and cancellation handles",
"active durable-batch attempt contexts (batch metadata remains persistent)",
"fair-queue heap and virtual clocks",
"active worker/model slots",
"browser OIDC sessions",
"live-flow animation state",
},
}
}
func (s *Server) flushPersistentState(ctx context.Context) error {
paths := state.Resolve(s.cfg.Storage)
var errs []string
if err := s.metrics.SavePersistent(paths.Metrics); err != nil {
errs = append(errs, "metrics: "+err.Error())
}
if saver, ok := s.quota.(persistentSaver); ok {
if err := saver.SavePersistent(paths.Quota); err != nil {
errs = append(errs, "quota: "+err.Error())
}
}
if err := s.workers.SavePerformance(paths.WorkerPerformance); err != nil {
errs = append(errs, "worker performance: "+err.Error())
}
if err := s.usage.Flush(ctx); err != nil {
errs = append(errs, "usage: "+err.Error())
}
if s.conversations != nil && s.conversations.Enabled() {
if err := s.conversations.Compact(); err != nil {
errs = append(errs, "conversations: "+err.Error())
}
}
if s.batchJobs != nil && s.batchJobs.Enabled() {
if err := s.batchJobs.Compact(); err != nil {
errs = append(errs, "batch jobs: "+err.Error())
}
}
if len(errs) > 0 {
return fmt.Errorf("persistent flush failed: %s", strings.Join(errs, "; "))
}
return nil
}
func (s *Server) uiStorageFlush(w http.ResponseWriter, r *http.Request) {
ctx, cancel := context.WithTimeout(r.Context(), 10*time.Second)
defer cancel()
if err := s.flushPersistentState(ctx); err != nil {
writeProtocolError(w, r, http.StatusServiceUnavailable, "storage_flush", err.Error())
return
}
writeJSON(w, http.StatusOK, map[string]any{"flushed": true, "at": time.Now().UTC(), "storage": s.storageStatus()})
}
func (s *Server) uiStorageCompact(w http.ResponseWriter, r *http.Request) {
ctx, cancel := context.WithTimeout(r.Context(), 30*time.Minute)
defer cancel()
if err := s.flushPersistentState(ctx); err != nil {
writeProtocolError(w, r, http.StatusServiceUnavailable, "storage_flush", err.Error())
return
}
status, err := s.usage.Compact(ctx)
if err != nil {
writeProtocolError(w, r, http.StatusInternalServerError, "usage_compaction", err.Error())
return
}
writeJSON(w, http.StatusOK, map[string]any{"compacted": true, "at": time.Now().UTC(), "retention": status, "storage": s.storageStatus()})
}
func (s *Server) uiStorageBackup(w http.ResponseWriter, r *http.Request) {
ctx, cancel := context.WithTimeout(r.Context(), 10*time.Second)
defer cancel()
if err := s.flushPersistentState(ctx); err != nil {
writeProtocolError(w, r, http.StatusServiceUnavailable, "storage_flush", err.Error())
return
}
name := "ollama-gateway-backup-" + time.Now().UTC().Format("20060102-150405") + ".zip"
w.Header().Set("Content-Type", "application/zip")
w.Header().Set("Content-Disposition", `attachment; filename="`+name+`"`)
w.Header().Set("Cache-Control", "no-store")
zw := zip.NewWriter(w)
defer zw.Close()
paths := state.Resolve(s.cfg.Storage)
known := []struct{ archive, path string }{
{"state/" + filepath.Base(paths.Config), paths.Config},
{"state/" + filepath.Base(paths.APIKeys), paths.APIKeys},
{"state/" + filepath.Base(paths.Policies), paths.Policies},
{"state/" + filepath.Base(paths.Metrics), paths.Metrics},
{"state/" + filepath.Base(paths.Quota), paths.Quota},
{"state/" + filepath.Base(paths.WorkerPerformance), paths.WorkerPerformance},
{"state/" + filepath.Base(paths.ModelPlacement), paths.ModelPlacement},
{"state/" + filepath.Base(paths.WorkerState), paths.WorkerState},
{"state/" + filepath.Base(paths.WarmModels), paths.WarmModels},
{"state/" + filepath.Base(paths.Alerts), paths.Alerts},
{"state/" + filepath.Base(paths.Conversations), paths.Conversations},
{"state/" + filepath.Base(paths.BatchJobs), paths.BatchJobs},
}
for _, f := range known {
if err := zipFile(zw, f.archive, f.path); err != nil && !os.IsNotExist(err) {
return
}
}
if paths.BatchDir != "" {
var matches []string
_ = filepath.Walk(paths.BatchDir, func(path string, info os.FileInfo, err error) error {
if err == nil && info.Mode().IsRegular() {
matches = append(matches, path)
}
return nil
})
sort.Strings(matches)
for _, p := range matches {
rel, err := filepath.Rel(paths.BatchDir, p)
if err != nil {
continue
}
if err := zipFile(zw, "batch/"+filepath.ToSlash(rel), p); err != nil && !os.IsNotExist(err) {
return
}
}
}
if s.cfg.Usage.JournalDir != "" {
var matches []string
_ = filepath.Walk(s.cfg.Usage.JournalDir, func(path string, info os.FileInfo, err error) error {
if err == nil && info.Mode().IsRegular() && (strings.HasSuffix(info.Name(), ".jsonl") || strings.HasSuffix(info.Name(), ".json")) {
matches = append(matches, path)
}
return nil
})
sort.Strings(matches)
for _, p := range matches {
rel, err := filepath.Rel(s.cfg.Usage.JournalDir, p)
if err != nil {
continue
}
if err := zipFile(zw, "usage/"+filepath.ToSlash(rel), p); err != nil && !os.IsNotExist(err) {
return
}
}
}
}
func zipFile(zw *zip.Writer, archiveName, path string) error {
f, err := os.Open(path)
if err != nil {
return err
}
defer f.Close()
st, err := f.Stat()
if err != nil {
return err
}
if !st.Mode().IsRegular() {
return nil
}
h, err := zip.FileInfoHeader(st)
if err != nil {
return err
}
h.Name = filepath.ToSlash(archiveName)
h.Method = zip.Deflate
dst, err := zw.CreateHeader(h)
if err != nil {
return err
}
_, err = io.Copy(dst, f)
return err
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff