-
This commit is contained in:
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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"])
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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) }
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
Reference in New Issue
Block a user