This commit is contained in:
2026-09-11 06:14:38 +02:00
parent bf64652300
commit e581949946
161 changed files with 31126 additions and 1 deletions
+339
View File
@@ -0,0 +1,339 @@
package main
import (
"bytes"
"context"
"encoding/json"
"flag"
"fmt"
"io"
"net/http"
"net/http/httptrace"
"os"
"sort"
"sync"
"time"
)
type result struct {
Latency time.Duration
TTFB time.Duration
Prompt int64
Completion int64
Bytes int64
Status int
Err error
}
type durationStats struct {
P50 float64 `json:"p50_ms"`
P95 float64 `json:"p95_ms"`
P99 float64 `json:"p99_ms"`
Max float64 `json:"max_ms"`
}
type summary struct {
StartedAt time.Time `json:"started_at"`
FinishedAt time.Time `json:"finished_at"`
BaseURL string `json:"base_url"`
Model string `json:"model"`
Concurrency int `json:"concurrency"`
Requests int `json:"requests"`
WarmupRequests int `json:"warmup_requests"`
Successful int `json:"successful"`
Errors int `json:"errors"`
StatusCounts map[string]int `json:"status_counts"`
WallMS float64 `json:"wall_ms"`
ThroughputRPS float64 `json:"throughput_rps"`
Latency durationStats `json:"latency"`
TTFB durationStats `json:"ttfb"`
PromptTokens int64 `json:"prompt_tokens"`
CompletionTokens int64 `json:"completion_tokens"`
CompletionTokPS float64 `json:"completion_tokens_per_second"`
BytesReceived int64 `json:"bytes_received"`
BytesPerSecond float64 `json:"bytes_per_second"`
Stream bool `json:"stream"`
KeepAlive bool `json:"keep_alive"`
ServiceClass string `json:"service_class,omitempty"`
ErrorsByMessage map[string]int `json:"errors_by_message,omitempty"`
}
type benchConfig struct {
Base string
Key string
Model string
Concurrency int
Requests int
Warmup int
MaxTokens int
Prompt string
Timeout time.Duration
Stream bool
DisableKeepAlive bool
ServiceClass string
JSONOut string
}
func main() {
base := flag.String("base-url", "http://127.0.0.1:8080", "gateway base URL")
key := flag.String("api-key", "", "API key; can also use GATEWAY_BENCH_API_KEY")
model := flag.String("model", "qwen3:8b", "model name")
conc := flag.Int("concurrency", 4, "parallel clients")
n := flag.Int("requests", 20, "total measured requests")
warmup := flag.Int("warmup", 0, "warmup requests before measurement")
maxTokens := flag.Int("max-tokens", 128, "max completion tokens")
prompt := flag.String("prompt", "Explain in three concise paragraphs why fair scheduling matters for shared LLM inference.", "prompt")
timeout := flag.Duration("timeout", 2*time.Minute, "per-request HTTP timeout")
stream := flag.Bool("stream", false, "request OpenAI streaming responses")
disableKeepAlive := flag.Bool("disable-keepalive", false, "disable HTTP connection reuse")
serviceClass := flag.String("service-class", "", "optional X-Gateway-Service-Class override")
jsonOut := flag.String("json-out", "", "optional path for machine-readable JSON summary; '-' writes JSON to stdout")
flag.Parse()
if *key == "" {
*key = os.Getenv("GATEWAY_BENCH_API_KEY")
}
cfg := benchConfig{Base: *base, Key: *key, Model: *model, Concurrency: *conc, Requests: *n, Warmup: *warmup, MaxTokens: *maxTokens, Prompt: *prompt, Timeout: *timeout, Stream: *stream, DisableKeepAlive: *disableKeepAlive, ServiceClass: *serviceClass, JSONOut: *jsonOut}
if cfg.Concurrency < 1 || cfg.Requests < 1 || cfg.Warmup < 0 {
fmt.Fprintln(os.Stderr, "concurrency and requests must be positive; warmup must be non-negative")
os.Exit(2)
}
if cfg.Timeout <= 0 {
fmt.Fprintln(os.Stderr, "timeout must be positive")
os.Exit(2)
}
client := newClient(cfg)
body, err := requestBody(cfg)
if err != nil {
fmt.Fprintln(os.Stderr, "request body:", err)
os.Exit(2)
}
ctx := context.Background()
if cfg.Warmup > 0 {
warm := run(ctx, client, cfg, body, cfg.Warmup)
for _, r := range warm {
if r.Err != nil {
fmt.Fprintln(os.Stderr, "warmup error:", r.Err)
}
}
}
started := time.Now().UTC()
results := run(ctx, client, cfg, body, cfg.Requests)
finished := time.Now().UTC()
s := summarize(cfg, started, finished, results)
printHuman(s)
if cfg.JSONOut != "" {
if err := writeJSONSummary(cfg.JSONOut, s); err != nil {
fmt.Fprintln(os.Stderr, "write JSON summary:", err)
os.Exit(1)
}
}
if s.Successful == 0 {
os.Exit(1)
}
}
func newClient(cfg benchConfig) *http.Client {
idle := max(256, cfg.Concurrency*2)
return &http.Client{
Timeout: cfg.Timeout,
Transport: &http.Transport{
MaxIdleConns: idle,
MaxIdleConnsPerHost: idle,
IdleConnTimeout: 90 * time.Second,
DisableKeepAlives: cfg.DisableKeepAlive,
},
}
}
func requestBody(cfg benchConfig) ([]byte, error) {
return json.Marshal(map[string]any{
"model": cfg.Model,
"messages": []map[string]string{{"role": "user", "content": cfg.Prompt}},
"max_tokens": cfg.MaxTokens,
"stream": cfg.Stream,
})
}
func run(ctx context.Context, client *http.Client, cfg benchConfig, body []byte, count int) []result {
jobs := make(chan struct{})
results := make(chan result, count)
var wg sync.WaitGroup
for i := 0; i < cfg.Concurrency; i++ {
wg.Add(1)
go func() {
defer wg.Done()
for range jobs {
results <- one(ctx, client, cfg, body)
}
}()
}
go func() {
for i := 0; i < count; i++ {
jobs <- struct{}{}
}
close(jobs)
wg.Wait()
close(results)
}()
out := make([]result, 0, count)
for r := range results {
out = append(out, r)
}
return out
}
func one(ctx context.Context, client *http.Client, cfg benchConfig, body []byte) result {
started := time.Now()
var firstByte time.Time
trace := &httptrace.ClientTrace{GotFirstResponseByte: func() { firstByte = time.Now() }}
reqCtx := httptrace.WithClientTrace(ctx, trace)
req, err := http.NewRequestWithContext(reqCtx, http.MethodPost, stringsTrimRightSlash(cfg.Base)+"/v1/chat/completions", bytes.NewReader(body))
if err != nil {
return result{Err: err}
}
req.Header.Set("Content-Type", "application/json")
if cfg.Key != "" {
req.Header.Set("Authorization", "Bearer "+cfg.Key)
}
if cfg.ServiceClass != "" {
req.Header.Set("X-Gateway-Service-Class", cfg.ServiceClass)
}
resp, err := client.Do(req)
if err != nil {
return result{Latency: time.Since(started), Err: err}
}
b, readErr := io.ReadAll(resp.Body)
closeErr := resp.Body.Close()
latency := time.Since(started)
ttfb := time.Duration(0)
if !firstByte.IsZero() {
ttfb = firstByte.Sub(started)
}
r := result{Latency: latency, TTFB: ttfb, Bytes: int64(len(b)), Status: resp.StatusCode}
if readErr != nil {
r.Err = readErr
return r
}
if closeErr != nil {
r.Err = closeErr
return r
}
if resp.StatusCode/100 != 2 {
r.Err = fmt.Errorf("HTTP %d: %s", resp.StatusCode, truncate(string(b), 240))
return r
}
if !cfg.Stream {
var doc struct {
Usage struct {
Prompt int64 `json:"prompt_tokens"`
Completion int64 `json:"completion_tokens"`
} `json:"usage"`
}
if err := json.Unmarshal(b, &doc); err == nil {
r.Prompt = doc.Usage.Prompt
r.Completion = doc.Usage.Completion
}
}
return r
}
func summarize(cfg benchConfig, started, finished time.Time, results []result) summary {
wall := finished.Sub(started)
latencies := make([]time.Duration, 0, len(results))
ttfbs := make([]time.Duration, 0, len(results))
statusCounts := map[string]int{}
errorsByMessage := map[string]int{}
var prompt, completion, received int64
errs := 0
for _, r := range results {
if r.Status != 0 {
statusCounts[fmt.Sprintf("%d", r.Status)]++
}
if r.Err != nil {
errs++
errorsByMessage[truncate(r.Err.Error(), 180)]++
continue
}
latencies = append(latencies, r.Latency)
if r.TTFB > 0 {
ttfbs = append(ttfbs, r.TTFB)
}
prompt += r.Prompt
completion += r.Completion
received += r.Bytes
}
success := len(latencies)
seconds := wall.Seconds()
if seconds <= 0 {
seconds = 1e-9
}
return summary{
StartedAt: started, FinishedAt: finished, BaseURL: cfg.Base, Model: cfg.Model,
Concurrency: cfg.Concurrency, Requests: cfg.Requests, WarmupRequests: cfg.Warmup,
Successful: success, Errors: errs, StatusCounts: statusCounts, WallMS: float64(wall) / float64(time.Millisecond),
ThroughputRPS: float64(success) / seconds, Latency: stats(latencies), TTFB: stats(ttfbs),
PromptTokens: prompt, CompletionTokens: completion, CompletionTokPS: float64(completion) / seconds,
BytesReceived: received, BytesPerSecond: float64(received) / seconds, Stream: cfg.Stream,
KeepAlive: !cfg.DisableKeepAlive, ServiceClass: cfg.ServiceClass, ErrorsByMessage: errorsByMessage,
}
}
func stats(ds []time.Duration) durationStats {
if len(ds) == 0 {
return durationStats{}
}
sort.Slice(ds, func(i, j int) bool { return ds[i] < ds[j] })
pct := func(p float64) time.Duration {
idx := int(float64(len(ds)-1) * p)
return ds[idx]
}
return durationStats{P50: msFloat(pct(.50)), P95: msFloat(pct(.95)), P99: msFloat(pct(.99)), Max: msFloat(ds[len(ds)-1])}
}
func printHuman(s summary) {
fmt.Printf("successful=%d errors=%d concurrency=%d wall=%s keepalive=%t stream=%t\n", s.Successful, s.Errors, s.Concurrency, time.Duration(s.WallMS*float64(time.Millisecond)).Round(time.Millisecond), s.KeepAlive, s.Stream)
fmt.Printf("latency p50=%s p95=%s p99=%s max=%s\n", fmtMS(s.Latency.P50), fmtMS(s.Latency.P95), fmtMS(s.Latency.P99), fmtMS(s.Latency.Max))
fmt.Printf("ttfb p50=%s p95=%s p99=%s max=%s\n", fmtMS(s.TTFB.P50), fmtMS(s.TTFB.P95), fmtMS(s.TTFB.P99), fmtMS(s.TTFB.Max))
fmt.Printf("throughput=%.2f req/s bytes=%d bytes/s=%.0f prompt_tokens=%d completion_tokens=%d completion_tok/s=%.2f\n", s.ThroughputRPS, s.BytesReceived, s.BytesPerSecond, s.PromptTokens, s.CompletionTokens, s.CompletionTokPS)
if len(s.StatusCounts) > 0 {
b, _ := json.Marshal(s.StatusCounts)
fmt.Printf("status=%s\n", b)
}
if len(s.ErrorsByMessage) > 0 {
for msg, n := range s.ErrorsByMessage {
fmt.Fprintf(os.Stderr, "error x%d: %s\n", n, msg)
}
}
}
func writeJSONSummary(path string, s summary) error {
b, err := json.MarshalIndent(s, "", " ")
if err != nil {
return err
}
b = append(b, '\n')
if path == "-" {
_, err = os.Stdout.Write(b)
return err
}
return os.WriteFile(path, b, 0644)
}
func msFloat(d time.Duration) float64 { return float64(d) / float64(time.Millisecond) }
func fmtMS(v float64) string {
return (time.Duration(v * float64(time.Millisecond))).Round(time.Microsecond).String()
}
func truncate(s string, n int) string {
if len(s) <= n {
return s
}
return s[:n] + "…"
}
func stringsTrimRightSlash(s string) string {
for len(s) > 0 && s[len(s)-1] == '/' {
s = s[:len(s)-1]
}
return s
}
+68
View File
@@ -0,0 +1,68 @@
package main
import (
"context"
"fmt"
"net/http"
"net/http/httptest"
"testing"
"time"
)
func TestStats(t *testing.T) {
d := []time.Duration{10 * time.Millisecond, 50 * time.Millisecond, 20 * time.Millisecond, 40 * time.Millisecond, 30 * time.Millisecond}
s := stats(d)
if s.P50 != 30 || s.P95 != 40 || s.P99 != 40 || s.Max != 50 {
t.Fatalf("unexpected stats: %+v", s)
}
}
func TestOneCapturesStatusUsageBytesAndHeaders(t *testing.T) {
var gotAuth, gotClass string
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
gotAuth = r.Header.Get("Authorization")
gotClass = r.Header.Get("X-Gateway-Service-Class")
w.Header().Set("Content-Type", "application/json")
fmt.Fprint(w, `{"choices":[{"message":{"role":"assistant","content":"ok"}}],"usage":{"prompt_tokens":9,"completion_tokens":4}}`)
}))
defer ts.Close()
cfg := benchConfig{Base: ts.URL + "/", Key: "secret", Model: "m", ServiceClass: "batch", Timeout: time.Second}
body, err := requestBody(cfg)
if err != nil {
t.Fatal(err)
}
r := one(context.Background(), newClient(cfg), cfg, body)
if r.Err != nil {
t.Fatal(r.Err)
}
if r.Status != http.StatusOK || r.Prompt != 9 || r.Completion != 4 || r.Bytes == 0 || r.Latency <= 0 || r.TTFB <= 0 {
t.Fatalf("unexpected result: %+v", r)
}
if gotAuth != "Bearer secret" || gotClass != "batch" {
t.Fatalf("headers auth=%q class=%q", gotAuth, gotClass)
}
}
func TestSummarizeSeparatesErrors(t *testing.T) {
cfg := benchConfig{Base: "http://gateway", Model: "m", Concurrency: 2, Requests: 3, Warmup: 1}
start := time.Unix(1, 0).UTC()
finish := start.Add(time.Second)
s := summarize(cfg, start, finish, []result{
{Latency: 10 * time.Millisecond, TTFB: 5 * time.Millisecond, Prompt: 8, Completion: 2, Bytes: 100, Status: 200},
{Latency: 20 * time.Millisecond, TTFB: 7 * time.Millisecond, Prompt: 8, Completion: 3, Bytes: 110, Status: 200},
{Latency: 2 * time.Millisecond, Status: 503, Err: fmt.Errorf("HTTP 503: unavailable")},
})
if s.Successful != 2 || s.Errors != 1 || s.StatusCounts["200"] != 2 || s.StatusCounts["503"] != 1 {
t.Fatalf("unexpected counts: %+v", s)
}
if s.ThroughputRPS != 2 || s.PromptTokens != 16 || s.CompletionTokens != 5 || s.BytesReceived != 210 {
t.Fatalf("unexpected throughput/usage: %+v", s)
}
}
func TestStringsTrimRightSlash(t *testing.T) {
if got := stringsTrimRightSlash("http://x///"); got != "http://x" {
t.Fatalf("got %q", got)
}
}
+396
View File
@@ -0,0 +1,396 @@
package main
import (
"bufio"
"encoding/json"
"errors"
"flag"
"fmt"
"os"
"path/filepath"
"sort"
"strconv"
"strings"
"time"
"github.com/example/ollama-fair-gateway/internal/haresource"
)
type durationStats struct {
P50 float64 `json:"p50_ms"`
P95 float64 `json:"p95_ms"`
P99 float64 `json:"p99_ms"`
Max float64 `json:"max_ms"`
}
type benchSummary struct {
StartedAt time.Time `json:"started_at"`
FinishedAt time.Time `json:"finished_at"`
BaseURL string `json:"base_url"`
Model string `json:"model"`
Concurrency int `json:"concurrency"`
Requests int `json:"requests"`
WarmupRequests int `json:"warmup_requests"`
Successful int `json:"successful"`
Errors int `json:"errors"`
StatusCounts map[string]int `json:"status_counts"`
WallMS float64 `json:"wall_ms"`
ThroughputRPS float64 `json:"throughput_rps"`
Latency durationStats `json:"latency"`
TTFB durationStats `json:"ttfb"`
ServiceClass string `json:"service_class,omitempty"`
ErrorsByMessage map[string]int `json:"errors_by_message,omitempty"`
}
type metricSample struct {
Name string
Labels string
Value float64
}
type resourceSnapshot = haresource.Snapshot
type resourceSamplingReport = haresource.SamplingReport
type levelReport struct {
Concurrency int `json:"concurrency"`
Requests int `json:"requests"`
WarmupRequests int `json:"warmup_requests"`
Successful int `json:"successful"`
Errors int `json:"errors"`
ThroughputRPS float64 `json:"throughput_rps"`
Latency durationStats `json:"latency"`
TTFB durationStats `json:"ttfb"`
GatewayRequestsDelta float64 `json:"gateway_requests_delta"`
Gateway2xxDelta float64 `json:"gateway_2xx_delta"`
GatewayErrorDelta float64 `json:"gateway_error_delta"`
RetriesDelta float64 `json:"retries_delta"`
CircuitOpensDelta float64 `json:"circuit_opens_delta"`
PromptTokensDelta float64 `json:"prompt_tokens_delta"`
CompletionTokensDelta float64 `json:"completion_tokens_delta"`
QueueObservationsDelta float64 `json:"queue_observations_delta"`
ServiceObservationsDelta float64 `json:"service_observations_delta"`
AccountingMatches bool `json:"accounting_matches"`
ResourcesBefore *resourceSnapshot `json:"resources_before,omitempty"`
ResourcesAfter *resourceSnapshot `json:"resources_after,omitempty"`
ResourceSamples *resourceSamplingReport `json:"resource_samples,omitempty"`
}
type evidenceReport struct {
GeneratedAt time.Time `json:"generated_at"`
InputDir string `json:"input_dir"`
Levels []levelReport `json:"levels"`
EvidenceComplete bool `json:"evidence_complete"`
ResourceEvidenceComplete bool `json:"resource_evidence_complete"`
SustainedResourceEvidenceComplete bool `json:"sustained_resource_evidence_complete"`
GateStatus string `json:"gate_status"`
GateReason string `json:"gate_reason"`
}
func main() {
input := flag.String("input", "", "HA-readiness result directory")
jsonOut := flag.String("json-out", "", "JSON report path; defaults to <input>/report.json")
markdownOut := flag.String("markdown-out", "", "Markdown report path; defaults to <input>/report.md")
flag.Parse()
if *input == "" {
fmt.Fprintln(os.Stderr, "-input is required")
os.Exit(2)
}
if *jsonOut == "" {
*jsonOut = filepath.Join(*input, "report.json")
}
if *markdownOut == "" {
*markdownOut = filepath.Join(*input, "report.md")
}
r, err := buildReport(*input)
if err != nil {
fmt.Fprintln(os.Stderr, "build report:", err)
os.Exit(1)
}
if err := writeJSON(*jsonOut, r); err != nil {
fmt.Fprintln(os.Stderr, "write JSON report:", err)
os.Exit(1)
}
if err := os.WriteFile(*markdownOut, []byte(renderMarkdown(r)), 0o644); err != nil {
fmt.Fprintln(os.Stderr, "write Markdown report:", err)
os.Exit(1)
}
fmt.Printf("report: %s\nreport: %s\ngate: %s — %s\n", *jsonOut, *markdownOut, r.GateStatus, r.GateReason)
}
func buildReport(dir string) (evidenceReport, error) {
matches, err := filepath.Glob(filepath.Join(dir, "concurrency-*.json"))
if err != nil {
return evidenceReport{}, err
}
if len(matches) == 0 {
return evidenceReport{}, errors.New("no concurrency-*.json files found")
}
levels := make([]levelReport, 0, len(matches))
complete := true
resourceComplete := true
resourceSeen := false
sustainedComplete := true
sustainedSeen := false
for _, p := range matches {
b, err := os.ReadFile(p)
if err != nil {
return evidenceReport{}, err
}
var s benchSummary
if err := json.Unmarshal(b, &s); err != nil {
return evidenceReport{}, fmt.Errorf("%s: %w", p, err)
}
beforePath := filepath.Join(dir, fmt.Sprintf("metrics-before-c%d.prom", s.Concurrency))
afterPath := filepath.Join(dir, fmt.Sprintf("metrics-after-c%d.prom", s.Concurrency))
before, errBefore := parsePrometheusFile(beforePath)
after, errAfter := parsePrometheusFile(afterPath)
if errBefore != nil || errAfter != nil {
complete = false
}
l := levelReport{
Concurrency: s.Concurrency, Requests: s.Requests, WarmupRequests: s.WarmupRequests,
Successful: s.Successful, Errors: s.Errors, ThroughputRPS: s.ThroughputRPS,
Latency: s.Latency, TTFB: s.TTFB,
}
resourceBefore, errResourceBefore := readResourceSnapshot(filepath.Join(dir, fmt.Sprintf("resources-before-c%d.json", s.Concurrency)))
resourceAfter, errResourceAfter := readResourceSnapshot(filepath.Join(dir, fmt.Sprintf("resources-after-c%d.json", s.Concurrency)))
if errResourceBefore == nil && errResourceAfter == nil {
l.ResourcesBefore = &resourceBefore
l.ResourcesAfter = &resourceAfter
resourceSeen = true
} else {
resourceComplete = false
}
resourceSamples, errResourceSamples := readResourceSamples(filepath.Join(dir, fmt.Sprintf("resources-samples-c%d.json", s.Concurrency)))
if errResourceSamples == nil && resourceSamples.Complete() {
l.ResourceSamples = &resourceSamples
sustainedSeen = true
} else {
sustainedComplete = false
}
if errBefore == nil && errAfter == nil {
l.GatewayRequestsDelta = deltaByName(before, after, "ollama_gateway_requests_total")
l.Gateway2xxDelta = deltaByNameLabelContains(before, after, "ollama_gateway_requests_total", `status_class="2xx"`)
l.GatewayErrorDelta = deltaByName(before, after, "ollama_gateway_errors_total")
// Older/current builds expose errors via requests_total status class rather than a dedicated counter.
if l.GatewayErrorDelta == 0 {
l.GatewayErrorDelta = l.GatewayRequestsDelta - l.Gateway2xxDelta
}
l.RetriesDelta = deltaByName(before, after, "ollama_gateway_retries_total")
l.CircuitOpensDelta = deltaByName(before, after, "ollama_gateway_circuit_opens_total")
l.PromptTokensDelta = deltaByName(before, after, "ollama_gateway_prompt_tokens_total")
l.CompletionTokensDelta = deltaByName(before, after, "ollama_gateway_completion_tokens_total")
l.QueueObservationsDelta = deltaByName(before, after, "ollama_gateway_queue_seconds_count")
l.ServiceObservationsDelta = deltaByName(before, after, "ollama_gateway_service_seconds_count")
expected := float64(s.Requests + s.WarmupRequests)
l.AccountingMatches = almostEqual(l.GatewayRequestsDelta, expected) && almostEqual(l.QueueObservationsDelta, expected) && almostEqual(l.ServiceObservationsDelta, expected)
if !l.AccountingMatches {
complete = false
}
}
levels = append(levels, l)
}
sort.Slice(levels, func(i, j int) bool { return levels[i].Concurrency < levels[j].Concurrency })
if !resourceSeen {
resourceComplete = false
}
if !sustainedSeen {
sustainedComplete = false
}
r := evidenceReport{
GeneratedAt: time.Now().UTC(), InputDir: filepath.Clean(dir), Levels: levels,
EvidenceComplete: complete, ResourceEvidenceComplete: resourceComplete,
SustainedResourceEvidenceComplete: sustainedComplete,
}
if !complete {
r.GateStatus = "incomplete"
r.GateReason = "benchmark and gateway metrics evidence are missing or do not reconcile"
} else if sustainedComplete {
r.GateStatus = "not-proven"
r.GateReason = "client/gateway accounting reconciles and sustained host/process resource sampling is complete; HA still requires an operator-demonstrated capacity, availability, or topology need"
} else if resourceComplete {
r.GateStatus = "not-proven"
r.GateReason = "client/gateway accounting reconciles and before/after host/process resource snapshots are present; sustained peak resource evidence is incomplete, and HA still requires an operator-demonstrated capacity, availability, or topology need"
} else {
r.GateStatus = "not-proven"
r.GateReason = "client and gateway counters reconcile; host/process resource evidence is incomplete, and HA still requires demonstrated capacity, availability, or topology need"
}
return r, nil
}
func readResourceSnapshot(path string) (resourceSnapshot, error) {
b, err := os.ReadFile(path)
if err != nil {
return resourceSnapshot{}, err
}
var s resourceSnapshot
if err := json.Unmarshal(b, &s); err != nil {
return resourceSnapshot{}, err
}
if s.PID <= 0 {
return resourceSnapshot{}, fmt.Errorf("%s: invalid pid", path)
}
return s, nil
}
func readResourceSamples(path string) (resourceSamplingReport, error) {
b, err := os.ReadFile(path)
if err != nil {
return resourceSamplingReport{}, err
}
var r resourceSamplingReport
if err := json.Unmarshal(b, &r); err != nil {
return resourceSamplingReport{}, err
}
if !r.Complete() {
return r, fmt.Errorf("%s: incomplete sustained resource report", path)
}
return r, nil
}
func parsePrometheusFile(path string) ([]metricSample, error) {
f, err := os.Open(path)
if err != nil {
return nil, err
}
defer f.Close()
var out []metricSample
s := bufio.NewScanner(f)
buf := make([]byte, 64*1024)
s.Buffer(buf, 4*1024*1024)
for s.Scan() {
line := strings.TrimSpace(s.Text())
if line == "" || strings.HasPrefix(line, "#") {
continue
}
parts := strings.Fields(line)
if len(parts) < 2 {
continue
}
v, err := strconv.ParseFloat(parts[1], 64)
if err != nil {
continue
}
key := parts[0]
name, labels := key, ""
if i := strings.IndexByte(key, '{'); i >= 0 {
name = key[:i]
labels = strings.TrimSuffix(key[i+1:], "}")
}
out = append(out, metricSample{Name: name, Labels: labels, Value: v})
}
return out, s.Err()
}
func deltaByName(before, after []metricSample, name string) float64 {
return sumBy(before, name, "", false, after)
}
func deltaByNameLabelContains(before, after []metricSample, name, label string) float64 {
return sumBy(before, name, label, true, after)
}
func sumBy(before []metricSample, name, label string, filter bool, after []metricSample) float64 {
sum := func(xs []metricSample) float64 {
var n float64
for _, x := range xs {
if x.Name != name {
continue
}
if filter && !strings.Contains(x.Labels, label) {
continue
}
n += x.Value
}
return n
}
return sum(after) - sum(before)
}
func almostEqual(a, b float64) bool {
d := a - b
if d < 0 {
d = -d
}
return d < 0.000001
}
func writeJSON(path string, v any) error {
b, err := json.MarshalIndent(v, "", " ")
if err != nil {
return err
}
b = append(b, '\n')
return os.WriteFile(path, b, 0o644)
}
func renderMarkdown(r evidenceReport) string {
var b strings.Builder
fmt.Fprintln(&b, "# HA readiness evidence report")
fmt.Fprintln(&b)
fmt.Fprintf(&b, "Generated: `%s` \n", r.GeneratedAt.Format(time.RFC3339))
fmt.Fprintf(&b, "Gateway accounting complete: **%t** \n", r.EvidenceComplete)
fmt.Fprintf(&b, "Resource snapshots complete: **%t** \n", r.ResourceEvidenceComplete)
fmt.Fprintf(&b, "Sustained resource sampling complete: **%t** \n", r.SustainedResourceEvidenceComplete)
fmt.Fprintf(&b, "P3.2 gate: **%s** — %s\n\n", r.GateStatus, r.GateReason)
fmt.Fprintln(&b, "| Concurrency | Success/Error | req/s | p95 latency | p95 TTFB | Gateway requests Δ | 2xx Δ | retries Δ | circuit opens Δ | counters reconcile | after RSS | sampled peak RSS | sampled peak CPU | sampled peak FDs | samples |")
fmt.Fprintln(&b, "|---:|---:|---:|---:|---:|---:|---:|---:|---:|:---:|---:|---:|---:|---:|---:|")
for _, l := range r.Levels {
afterRSS, _, _, _, _ := resourceCells(l.ResourcesAfter)
peakRSS, peakCPU, peakFDs, sampleCount := sustainedCells(l.ResourceSamples)
fmt.Fprintf(&b, "| %d | %d/%d | %.2f | %.2f ms | %.2f ms | %.0f | %.0f | %.0f | %.0f | %t | %s | %s | %s | %s | %s |\n", l.Concurrency, l.Successful, l.Errors, l.ThroughputRPS, l.Latency.P95, l.TTFB.P95, l.GatewayRequestsDelta, l.Gateway2xxDelta, l.RetriesDelta, l.CircuitOpensDelta, l.AccountingMatches, afterRSS, peakRSS, peakCPU, peakFDs, sampleCount)
}
fmt.Fprintln(&b, "\n## Interpretation")
fmt.Fprintln(&b)
fmt.Fprintln(&b, "This report verifies that client-side benchmark counts reconcile with gateway-side Prometheus counters. Before/after resource files remain endpoint evidence. When `resources-samples-cN.json` is present and complete, the sampled peak columns come from measurements taken throughout that benchmark level; on Linux process CPU is calculated from `/proc` process/host tick deltas and may exceed 100% when multiple logical CPUs are used.")
fmt.Fprintln(&b)
fmt.Fprintln(&b, "The report intentionally does not prove that HA is required. Correlate sustained process/host evidence with worker/GPU saturation and an explicit capacity, availability, or topology requirement before changing the P3.2 gate.")
return b.String()
}
func resourceCells(s *resourceSnapshot) (rss, cpu, fds, threads, load1 string) {
if s == nil {
return "-", "-", "-", "-", "-"
}
if s.ProcessRSSBytes > 0 {
rss = fmt.Sprintf("%.1f MiB", float64(s.ProcessRSSBytes)/(1024*1024))
} else {
rss = "-"
}
if s.ProcessCPUPercent > 0 {
cpu = fmt.Sprintf("%.1f%%", s.ProcessCPUPercent)
} else {
cpu = "0.0%"
}
if s.ProcessOpenFDs > 0 {
fds = strconv.FormatInt(s.ProcessOpenFDs, 10)
} else {
fds = "-"
}
if s.ProcessThreads > 0 {
threads = strconv.FormatInt(s.ProcessThreads, 10)
} else {
threads = "-"
}
load1 = fmt.Sprintf("%.2f", s.HostLoad1)
return
}
func sustainedCells(r *resourceSamplingReport) (rss, cpu, fds, samples string) {
if r == nil {
return "-", "-", "-", "-"
}
if r.PeakProcessRSSBytes > 0 {
rss = fmt.Sprintf("%.1f MiB", float64(r.PeakProcessRSSBytes)/(1024*1024))
} else {
rss = "-"
}
cpu = fmt.Sprintf("%.1f%%", r.PeakProcessCPUPercent)
if r.PeakProcessOpenFDs > 0 {
fds = strconv.FormatInt(r.PeakProcessOpenFDs, 10)
} else {
fds = "-"
}
samples = strconv.Itoa(len(r.Samples))
return
}
+207
View File
@@ -0,0 +1,207 @@
package main
import (
"encoding/json"
"os"
"path/filepath"
"strings"
"testing"
)
func TestBuildReportReconcilesWarmupAndMeasuredRequests(t *testing.T) {
d := t.TempDir()
s := benchSummary{Concurrency: 4, Requests: 20, WarmupRequests: 3, Successful: 20, ThroughputRPS: 123.4, Latency: durationStats{P95: 8}, TTFB: durationStats{P95: 3}}
b, _ := json.Marshal(s)
if err := os.WriteFile(filepath.Join(d, "concurrency-4.json"), b, 0o600); err != nil {
t.Fatal(err)
}
before := `ollama_gateway_requests_total{api="openai",status_class="2xx"} 10
ollama_gateway_queue_seconds_count 10
ollama_gateway_service_seconds_count 10
ollama_gateway_prompt_tokens_total 100
ollama_gateway_completion_tokens_total 50
`
after := `ollama_gateway_requests_total{api="openai",status_class="2xx"} 33
ollama_gateway_queue_seconds_count 33
ollama_gateway_service_seconds_count 33
ollama_gateway_prompt_tokens_total 330
ollama_gateway_completion_tokens_total 165
ollama_gateway_retries_total{worker="w"} 0
`
if err := os.WriteFile(filepath.Join(d, "metrics-before-c4.prom"), []byte(before), 0o600); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(d, "metrics-after-c4.prom"), []byte(after), 0o600); err != nil {
t.Fatal(err)
}
r, err := buildReport(d)
if err != nil {
t.Fatal(err)
}
if !r.EvidenceComplete || len(r.Levels) != 1 {
t.Fatalf("unexpected report: %+v", r)
}
l := r.Levels[0]
if l.GatewayRequestsDelta != 23 || !l.AccountingMatches {
t.Fatalf("unexpected level: %+v", l)
}
if r.GateStatus != "not-proven" {
t.Fatalf("gate=%q", r.GateStatus)
}
if got := renderMarkdown(r); !strings.Contains(got, "| 4 | 20/0 |") {
t.Fatalf("markdown missing row: %s", got)
}
}
func TestBuildReportMarksMissingMetricsIncomplete(t *testing.T) {
d := t.TempDir()
b, _ := json.Marshal(benchSummary{Concurrency: 1, Requests: 1, Successful: 1})
if err := os.WriteFile(filepath.Join(d, "concurrency-1.json"), b, 0o600); err != nil {
t.Fatal(err)
}
r, err := buildReport(d)
if err != nil {
t.Fatal(err)
}
if r.EvidenceComplete || r.GateStatus != "incomplete" {
t.Fatalf("unexpected report: %+v", r)
}
}
func TestPrometheusParserUsesFullOllamaGatewayNames(t *testing.T) {
p := filepath.Join(t.TempDir(), "m.prom")
content := "# HELP x x\nollama_gateway_requests_total{api=\"openai\",status_class=\"2xx\"} 7\n"
if err := os.WriteFile(p, []byte(content), 0o600); err != nil {
t.Fatal(err)
}
xs, err := parsePrometheusFile(p)
if err != nil {
t.Fatal(err)
}
if got := deltaByName(nil, xs, "ollama_gateway_requests_total"); got != 7 {
t.Fatalf("got %v", got)
}
if got := deltaByName(nil, xs, "gateway_requests_total"); got != 0 {
t.Fatalf("short metric name unexpectedly matched: %v", got)
}
}
func TestBuildReportIncludesCompleteResourceSnapshots(t *testing.T) {
d := t.TempDir()
s := benchSummary{Concurrency: 8, Requests: 5, WarmupRequests: 2, Successful: 5, ThroughputRPS: 50, Latency: durationStats{P95: 10}, TTFB: durationStats{P95: 4}}
b, _ := json.Marshal(s)
if err := os.WriteFile(filepath.Join(d, "concurrency-8.json"), b, 0o600); err != nil {
t.Fatal(err)
}
beforeMetrics := "ollama_gateway_requests_total{status_class=\"2xx\"} 10\nollama_gateway_queue_seconds_count 10\nollama_gateway_service_seconds_count 10\n"
afterMetrics := "ollama_gateway_requests_total{status_class=\"2xx\"} 17\nollama_gateway_queue_seconds_count 17\nollama_gateway_service_seconds_count 17\n"
if err := os.WriteFile(filepath.Join(d, "metrics-before-c8.prom"), []byte(beforeMetrics), 0o600); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(d, "metrics-after-c8.prom"), []byte(afterMetrics), 0o600); err != nil {
t.Fatal(err)
}
before := resourceSnapshot{PID: 4242, ProcessRSSBytes: 64 * 1024 * 1024, ProcessCPUPercent: 2.5, ProcessThreads: 8, ProcessOpenFDs: 21, HostLogicalCPUs: 8, HostLoad1: 0.5}
after := resourceSnapshot{PID: 4242, ProcessRSSBytes: 96 * 1024 * 1024, ProcessCPUPercent: 33.3, ProcessThreads: 12, ProcessOpenFDs: 40, HostLogicalCPUs: 8, HostLoad1: 1.25}
for name, snap := range map[string]resourceSnapshot{"resources-before-c8.json": before, "resources-after-c8.json": after} {
data, _ := json.Marshal(snap)
if err := os.WriteFile(filepath.Join(d, name), data, 0o600); err != nil {
t.Fatal(err)
}
}
r, err := buildReport(d)
if err != nil {
t.Fatal(err)
}
if !r.EvidenceComplete || !r.ResourceEvidenceComplete {
t.Fatalf("unexpected completeness: %+v", r)
}
if r.Levels[0].ResourcesAfter == nil || r.Levels[0].ResourcesAfter.ProcessOpenFDs != 40 {
t.Fatalf("resource snapshot missing: %+v", r.Levels[0])
}
md := renderMarkdown(r)
if !strings.Contains(md, "96.0 MiB") || !strings.Contains(md, "Sustained resource sampling complete: **false**") || !strings.Contains(md, "Before/after resource files remain endpoint evidence") {
t.Fatalf("resource evidence missing from markdown: %s", md)
}
}
func TestBuildReportMarksPartialResourceSnapshotsIncompleteWithoutBreakingAccounting(t *testing.T) {
d := t.TempDir()
b, _ := json.Marshal(benchSummary{Concurrency: 2, Requests: 1, Successful: 1})
if err := os.WriteFile(filepath.Join(d, "concurrency-2.json"), b, 0o600); err != nil {
t.Fatal(err)
}
metrics := "ollama_gateway_requests_total{status_class=\"2xx\"} 1\nollama_gateway_queue_seconds_count 1\nollama_gateway_service_seconds_count 1\n"
if err := os.WriteFile(filepath.Join(d, "metrics-before-c2.prom"), []byte("ollama_gateway_requests_total{status_class=\"2xx\"} 0\nollama_gateway_queue_seconds_count 0\nollama_gateway_service_seconds_count 0\n"), 0o600); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(d, "metrics-after-c2.prom"), []byte(metrics), 0o600); err != nil {
t.Fatal(err)
}
data, _ := json.Marshal(resourceSnapshot{PID: 99, HostLogicalCPUs: 4})
if err := os.WriteFile(filepath.Join(d, "resources-before-c2.json"), data, 0o600); err != nil {
t.Fatal(err)
}
r, err := buildReport(d)
if err != nil {
t.Fatal(err)
}
if !r.EvidenceComplete || r.ResourceEvidenceComplete || r.GateStatus != "not-proven" {
t.Fatalf("unexpected report: %+v", r)
}
if !strings.Contains(r.GateReason, "resource evidence is incomplete") {
t.Fatalf("unexpected gate reason: %s", r.GateReason)
}
}
func TestBuildReportIncludesSustainedResourceSampling(t *testing.T) {
d := t.TempDir()
s := benchSummary{Concurrency: 16, Requests: 10, WarmupRequests: 2, Successful: 10, ThroughputRPS: 75, Latency: durationStats{P95: 12}, TTFB: durationStats{P95: 5}}
b, _ := json.Marshal(s)
if err := os.WriteFile(filepath.Join(d, "concurrency-16.json"), b, 0o600); err != nil {
t.Fatal(err)
}
beforeMetrics := "ollama_gateway_requests_total{status_class=\"2xx\"} 4\nollama_gateway_queue_seconds_count 4\nollama_gateway_service_seconds_count 4\n"
afterMetrics := "ollama_gateway_requests_total{status_class=\"2xx\"} 16\nollama_gateway_queue_seconds_count 16\nollama_gateway_service_seconds_count 16\n"
if err := os.WriteFile(filepath.Join(d, "metrics-before-c16.prom"), []byte(beforeMetrics), 0o600); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(d, "metrics-after-c16.prom"), []byte(afterMetrics), 0o600); err != nil {
t.Fatal(err)
}
for name, snap := range map[string]resourceSnapshot{
"resources-before-c16.json": {PID: 77, ProcessRSSBytes: 50 * 1024 * 1024, HostLogicalCPUs: 8},
"resources-after-c16.json": {PID: 77, ProcessRSSBytes: 60 * 1024 * 1024, HostLogicalCPUs: 8},
} {
data, _ := json.Marshal(snap)
if err := os.WriteFile(filepath.Join(d, name), data, 0o600); err != nil {
t.Fatal(err)
}
}
samples := resourceSamplingReport{
Version: 1, PID: 77, GOOS: "linux", PeakProcessRSSBytes: 80 * 1024 * 1024,
PeakProcessCPUPercent: 145.5, PeakProcessOpenFDs: 55,
Samples: []resourceSnapshot{{PID: 77}, {PID: 77}}, StopReason: "stop-file",
}
data, _ := json.Marshal(samples)
if err := os.WriteFile(filepath.Join(d, "resources-samples-c16.json"), data, 0o600); err != nil {
t.Fatal(err)
}
r, err := buildReport(d)
if err != nil {
t.Fatal(err)
}
if !r.EvidenceComplete || !r.ResourceEvidenceComplete || !r.SustainedResourceEvidenceComplete {
t.Fatalf("unexpected completeness: %+v", r)
}
if r.Levels[0].ResourceSamples == nil || r.Levels[0].ResourceSamples.PeakProcessOpenFDs != 55 {
t.Fatalf("missing sustained samples: %+v", r.Levels[0])
}
if !strings.Contains(r.GateReason, "sustained host/process resource sampling is complete") {
t.Fatalf("unexpected gate reason: %s", r.GateReason)
}
md := renderMarkdown(r)
if !strings.Contains(md, "80.0 MiB") || !strings.Contains(md, "145.5%") || !strings.Contains(md, "| 16 | 10/0 |") {
t.Fatalf("sustained evidence missing from markdown: %s", md)
}
}
+50
View File
@@ -0,0 +1,50 @@
package main
import (
"context"
"encoding/json"
"flag"
"fmt"
"os"
"os/signal"
"syscall"
"time"
"github.com/example/ollama-fair-gateway/internal/haresource"
)
func main() {
pid := flag.Int("pid", 0, "gateway process PID")
interval := flag.Duration("interval", 250*time.Millisecond, "resource sample interval (minimum 50ms)")
maxDuration := flag.Duration("max-duration", 15*time.Minute, "safety limit for one sampling run")
stopFile := flag.String("stop-file", "", "stop sampling after this file appears")
out := flag.String("json-out", "-", "output JSON path; '-' writes stdout")
flag.Parse()
if *pid <= 0 || *stopFile == "" {
fmt.Fprintln(os.Stderr, "-pid must be positive and -stop-file is required")
os.Exit(2)
}
ctx, cancel := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
defer cancel()
r, err := haresource.Sample(ctx, *pid, *interval, *maxDuration, *stopFile)
if err != nil {
fmt.Fprintln(os.Stderr, err)
os.Exit(1)
}
b, err := json.MarshalIndent(r, "", " ")
if err != nil {
fmt.Fprintln(os.Stderr, err)
os.Exit(1)
}
b = append(b, '\n')
if *out == "-" {
_, _ = os.Stdout.Write(b)
} else if err := os.WriteFile(*out, b, 0o644); err != nil {
fmt.Fprintln(os.Stderr, err)
os.Exit(1)
}
if !r.Complete() {
fmt.Fprintln(os.Stderr, "sampling completed but sustained resource evidence is incomplete")
os.Exit(3)
}
}
+40
View File
@@ -0,0 +1,40 @@
package main
import (
"encoding/json"
"flag"
"fmt"
"os"
"path/filepath"
"github.com/example/ollama-fair-gateway/internal/haresource"
)
func main() {
pid := flag.Int("pid", 0, "gateway process PID")
out := flag.String("json-out", "-", "output JSON path; '-' writes stdout")
flag.Parse()
if *pid <= 0 {
fmt.Fprintln(os.Stderr, "-pid must be positive")
os.Exit(2)
}
s := haresource.Collect(*pid)
b, err := json.MarshalIndent(s, "", " ")
if err != nil {
fmt.Fprintln(os.Stderr, err)
os.Exit(1)
}
b = append(b, '\n')
if *out == "-" {
_, _ = os.Stdout.Write(b)
return
}
if err := os.WriteFile(*out, b, 0o644); err != nil {
fmt.Fprintln(os.Stderr, err)
os.Exit(1)
}
}
func outputName(dir string, concurrency int, phase string) string {
return filepath.Join(dir, fmt.Sprintf("resources-%s-c%d.json", phase, concurrency))
}
+10
View File
@@ -0,0 +1,10 @@
package main
import "testing"
func TestOutputName(t *testing.T) {
got := outputName("out", 8, "before")
if got != "out/resources-before-c8.json" {
t.Fatalf("outputName=%q", got)
}
}
+191
View File
@@ -0,0 +1,191 @@
package main
import (
"bufio"
"encoding/json"
"flag"
"fmt"
"io"
"log"
"net/http"
"strings"
"sync/atomic"
"time"
)
type server struct {
model string
delay time.Duration
streamDelay time.Duration
responseBytes int
promptTokens int64
outputTokens int64
failEvery int64
requests atomic.Int64
}
func main() {
listen := flag.String("listen", "127.0.0.1:11435", "listen address")
model := flag.String("model", "qwen3:8b", "mock model name")
delay := flag.Duration("delay", 0, "delay before response headers")
streamDelay := flag.Duration("stream-delay", 0, "delay between streaming chunks")
responseBytes := flag.Int("response-bytes", 128, "approximate generated content bytes")
promptTokens := flag.Int64("prompt-tokens", 16, "reported prompt token count")
outputTokens := flag.Int64("output-tokens", 8, "reported completion token count")
failEvery := flag.Int64("fail-every", 0, "return HTTP 503 for every Nth inference request; 0 disables")
flag.Parse()
if *responseBytes < 0 || *promptTokens < 0 || *outputTokens < 0 || *failEvery < 0 {
log.Fatal("numeric flags must be non-negative")
}
s := &server{model: *model, delay: *delay, streamDelay: *streamDelay, responseBytes: *responseBytes, promptTokens: *promptTokens, outputTokens: *outputTokens, failEvery: *failEvery}
mux := http.NewServeMux()
mux.HandleFunc("/api/ps", s.ps)
mux.HandleFunc("/api/tags", s.tags)
mux.HandleFunc("/api/show", s.show)
mux.HandleFunc("/api/chat", s.nativeChat)
mux.HandleFunc("/api/generate", s.nativeGenerate)
mux.HandleFunc("/v1/chat/completions", s.openAIChat)
mux.HandleFunc("/healthz", func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusNoContent) })
h := &http.Server{Addr: *listen, Handler: mux, ReadHeaderTimeout: 5 * time.Second, IdleTimeout: 2 * time.Minute}
log.Printf("mock Ollama listening on http://%s model=%s delay=%s response_bytes=%d", *listen, *model, *delay, *responseBytes)
log.Fatal(h.ListenAndServe())
}
func (s *server) ps(w http.ResponseWriter, r *http.Request) {
writeJSON(w, map[string]any{"models": []map[string]any{{"name": s.model, "model": s.model, "size": 1 << 30, "size_vram": 1 << 30, "context_length": 32768}}})
}
func (s *server) tags(w http.ResponseWriter, r *http.Request) {
writeJSON(w, map[string]any{"models": []map[string]any{{"name": s.model, "model": s.model, "size": 1 << 30, "details": map[string]any{"parameter_size": "8B", "quantization_level": "Q4_K_M"}}}})
}
func (s *server) show(w http.ResponseWriter, r *http.Request) {
writeJSON(w, map[string]any{"capabilities": []string{"completion", "tools", "thinking", "vision"}, "model_info": map[string]any{"mock.context_length": 32768}})
}
func (s *server) shouldFail() bool {
n := s.requests.Add(1)
return s.failEvery > 0 && n%s.failEvery == 0
}
func (s *server) openAIChat(w http.ResponseWriter, r *http.Request) {
if s.delay > 0 {
time.Sleep(s.delay)
}
if s.shouldFail() {
writeJSONStatus(w, http.StatusServiceUnavailable, map[string]any{"error": map[string]any{"message": "mock failure", "type": "server_error"}})
return
}
var req struct {
Stream bool `json:"stream"`
}
_ = json.NewDecoder(io.LimitReader(r.Body, 8<<20)).Decode(&req)
content := strings.Repeat("x", s.responseBytes)
if req.Stream {
w.Header().Set("Content-Type", "text/event-stream")
w.WriteHeader(http.StatusOK)
f, _ := w.(http.Flusher)
parts := split(content, 4)
for i, p := range parts {
chunk := map[string]any{"id": "chatcmpl-mock", "object": "chat.completion.chunk", "choices": []map[string]any{{"index": 0, "delta": map[string]any{"content": p}, "finish_reason": nil}}}
if i == len(parts)-1 {
chunk["choices"] = []map[string]any{{"index": 0, "delta": map[string]any{}, "finish_reason": "stop"}}
}
b, _ := json.Marshal(chunk)
fmt.Fprintf(w, "data: %s\n\n", b)
if f != nil {
f.Flush()
}
if s.streamDelay > 0 {
time.Sleep(s.streamDelay)
}
}
io.WriteString(w, "data: [DONE]\n\n")
if f != nil {
f.Flush()
}
return
}
writeJSON(w, map[string]any{
"id": "chatcmpl-mock", "object": "chat.completion", "created": time.Now().Unix(), "model": s.model,
"choices": []map[string]any{{"index": 0, "message": map[string]any{"role": "assistant", "content": content}, "finish_reason": "stop"}},
"usage": map[string]any{"prompt_tokens": s.promptTokens, "completion_tokens": s.outputTokens, "total_tokens": s.promptTokens + s.outputTokens},
})
}
func (s *server) nativeChat(w http.ResponseWriter, r *http.Request) {
s.native(w, r, "message")
}
func (s *server) nativeGenerate(w http.ResponseWriter, r *http.Request) {
s.native(w, r, "response")
}
func (s *server) native(w http.ResponseWriter, r *http.Request, field string) {
if s.delay > 0 {
time.Sleep(s.delay)
}
if s.shouldFail() {
writeJSONStatus(w, http.StatusServiceUnavailable, map[string]any{"error": "mock failure"})
return
}
var req map[string]any
_ = json.NewDecoder(io.LimitReader(r.Body, 8<<20)).Decode(&req)
stream, _ := req["stream"].(bool)
content := strings.Repeat("x", s.responseBytes)
w.Header().Set("Content-Type", "application/x-ndjson")
bw := bufio.NewWriter(w)
if !stream {
doc := map[string]any{"model": s.model, "done": true, "prompt_eval_count": s.promptTokens, "eval_count": s.outputTokens}
if field == "message" {
doc[field] = map[string]any{"role": "assistant", "content": content}
} else {
doc[field] = content
}
_ = json.NewEncoder(bw).Encode(doc)
_ = bw.Flush()
return
}
f, _ := w.(http.Flusher)
for _, p := range split(content, 4) {
doc := map[string]any{"model": s.model, "done": false}
if field == "message" {
doc[field] = map[string]any{"role": "assistant", "content": p}
} else {
doc[field] = p
}
_ = json.NewEncoder(bw).Encode(doc)
_ = bw.Flush()
if f != nil {
f.Flush()
}
if s.streamDelay > 0 {
time.Sleep(s.streamDelay)
}
}
_ = json.NewEncoder(bw).Encode(map[string]any{"model": s.model, "done": true, "prompt_eval_count": s.promptTokens, "eval_count": s.outputTokens})
_ = bw.Flush()
if f != nil {
f.Flush()
}
}
func split(s string, n int) []string {
if n <= 1 || len(s) == 0 {
return []string{s}
}
out := make([]string, 0, n)
for i := 0; i < n; i++ {
start := len(s) * i / n
end := len(s) * (i + 1) / n
out = append(out, s[start:end])
}
return out
}
func writeJSON(w http.ResponseWriter, v any) { writeJSONStatus(w, http.StatusOK, v) }
func writeJSONStatus(w http.ResponseWriter, status int, v any) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(v)
}
+108
View File
@@ -0,0 +1,108 @@
package main
import (
"bufio"
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"sync"
"testing"
)
func TestOpenAIChatNonStreaming(t *testing.T) {
s := &server{model: "mock:latest", responseBytes: 12, promptTokens: 7, outputTokens: 3}
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(`{"stream":false}`))
rr := httptest.NewRecorder()
s.openAIChat(rr, req)
if rr.Code != http.StatusOK {
t.Fatalf("status=%d body=%s", rr.Code, rr.Body.String())
}
var doc struct {
Model string `json:"model"`
Choices []struct {
Message struct {
Content string `json:"content"`
} `json:"message"`
} `json:"choices"`
Usage struct {
Prompt int64 `json:"prompt_tokens"`
Output int64 `json:"completion_tokens"`
} `json:"usage"`
}
if err := json.Unmarshal(rr.Body.Bytes(), &doc); err != nil {
t.Fatal(err)
}
if doc.Model != "mock:latest" || len(doc.Choices) != 1 || len(doc.Choices[0].Message.Content) != 12 {
t.Fatalf("unexpected response: %+v", doc)
}
if doc.Usage.Prompt != 7 || doc.Usage.Output != 3 {
t.Fatalf("unexpected usage: %+v", doc.Usage)
}
}
func TestOpenAIChatStreaming(t *testing.T) {
s := &server{model: "mock:latest", responseBytes: 16}
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(`{"stream":true}`))
rr := httptest.NewRecorder()
s.openAIChat(rr, req)
if rr.Code != http.StatusOK {
t.Fatalf("status=%d body=%s", rr.Code, rr.Body.String())
}
if ct := rr.Header().Get("Content-Type"); !strings.Contains(ct, "text/event-stream") {
t.Fatalf("content-type=%q", ct)
}
if !strings.Contains(rr.Body.String(), "data: [DONE]") {
t.Fatalf("missing done marker: %s", rr.Body.String())
}
}
func TestNativeStreamingEndsWithUsage(t *testing.T) {
s := &server{model: "mock:latest", responseBytes: 8, promptTokens: 11, outputTokens: 5}
req := httptest.NewRequest(http.MethodPost, "/api/chat", strings.NewReader(`{"stream":true}`))
rr := httptest.NewRecorder()
s.nativeChat(rr, req)
if rr.Code != http.StatusOK {
t.Fatalf("status=%d body=%s", rr.Code, rr.Body.String())
}
scanner := bufio.NewScanner(strings.NewReader(rr.Body.String()))
var last map[string]any
for scanner.Scan() {
if err := json.Unmarshal(scanner.Bytes(), &last); err != nil {
t.Fatal(err)
}
}
if err := scanner.Err(); err != nil {
t.Fatal(err)
}
if done, _ := last["done"].(bool); !done {
t.Fatalf("last chunk not done: %#v", last)
}
if last["prompt_eval_count"] != float64(11) || last["eval_count"] != float64(5) {
t.Fatalf("unexpected usage: %#v", last)
}
}
func TestShouldFailConcurrent(t *testing.T) {
const total = 1000
s := &server{failEvery: 5}
var wg sync.WaitGroup
failures := make(chan struct{}, total)
for i := 0; i < total; i++ {
wg.Add(1)
go func() {
defer wg.Done()
if s.shouldFail() {
failures <- struct{}{}
}
}()
}
wg.Wait()
close(failures)
if got, want := len(failures), total/5; got != want {
t.Fatalf("failures=%d want=%d", got, want)
}
if got := s.requests.Load(); got != total {
t.Fatalf("requests=%d want=%d", got, total)
}
}
+255
View File
@@ -0,0 +1,255 @@
package main
import (
"context"
"flag"
"log/slog"
"net/http"
"os"
"os/signal"
"syscall"
"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/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/proxy"
"github.com/example/ollama-fair-gateway/internal/quota"
"github.com/example/ollama-fair-gateway/internal/scheduler"
"github.com/example/ollama-fair-gateway/internal/server"
"github.com/example/ollama-fair-gateway/internal/session"
"github.com/example/ollama-fair-gateway/internal/state"
"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"
)
func main() {
configPath := flag.String("config", "config.json", "configuration file")
checkConfig := flag.Bool("check-config", false, "validate effective configuration and storage, then exit")
probeURL := flag.String("probe", "", "probe an HTTP health/readiness URL, then exit")
probeTimeout := flag.Duration("probe-timeout", 2*time.Second, "timeout for -probe")
flag.Parse()
log := slog.New(slog.NewJSONHandler(os.Stdout, &slog.HandlerOptions{Level: slog.LevelInfo}))
slog.SetDefault(log)
if *probeURL != "" {
if err := runProbe(*probeURL, *probeTimeout); err != nil {
log.Error("probe failed", "url", *probeURL, "error", err)
os.Exit(1)
}
return
}
if *checkConfig {
if err := runConfigCheck(*configPath, os.Stdout); err != nil {
log.Error("configuration preflight failed", "error", err)
os.Exit(2)
}
return
}
loaded, err := loadEffectiveConfig(*configPath)
if err != nil {
log.Error("configuration error", "error", err)
os.Exit(2)
}
cfg := loaded.Config
paths := loaded.Paths
configStore := loaded.Store
if loaded.Persistent {
log.Info("loaded persistent UI configuration", "path", configStore.Path())
}
rootCtx, cancel := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
defer cancel()
apiKeyStore, err := state.NewAPIKeyStore(paths.APIKeys)
if err != nil {
log.Error("API key store initialization failed", "error", err)
os.Exit(2)
}
authenticator, err := auth.NewWithRuntimeStore(rootCtx, cfg.Auth, apiKeyStore)
if err != nil {
log.Error("authentication initialization failed", "error", err)
os.Exit(2)
}
sched := scheduler.NewLocal(cfg.Scheduler.GlobalConcurrency, cfg.Scheduler.MaxQueue, cfg.Scheduler.MaxQueuePerActor)
var ledger quota.Ledger = quota.Disabled{}
var quotaMemory *quota.Memory
if cfg.Quota.Enabled {
quotaMemory = quota.NewMemory()
if err := quotaMemory.LoadPersistent(paths.Quota); err != nil {
log.Error("quota state initialization failed", "error", err)
os.Exit(2)
}
quotaMemory.StartPersistence(rootCtx, paths.Quota, cfg.Storage.FlushInterval.Value(), func(err error) { log.Error("quota persistence failed", "error", err) })
ledger = quotaMemory
}
policyStore, err := state.NewPolicyStore(paths.Policies)
if err != nil {
log.Error("policy store initialization failed", "error", err)
os.Exit(2)
}
placementStore, err := state.NewModelPlacementStore(paths.ModelPlacement)
if err != nil {
log.Error("model placement store initialization failed", "error", err)
os.Exit(2)
}
workerStateStore, err := state.NewWorkerRuntimeStore(paths.WorkerState)
if err != nil {
log.Error("worker state store initialization failed", "error", err)
os.Exit(2)
}
uiSessions := session.NewMemory()
conversationStore, err := conversation.New(cfg.Conversations, paths.Conversations)
if err != nil {
log.Error("conversation store initialization failed", "error", err)
os.Exit(2)
}
conversationStore.StartCleanup(rootCtx, func(err error) { log.Error("conversation retention cleanup failed", "error", err) })
batchManager, err := batch.New(cfg.BatchJobs, paths.BatchJobs, paths.BatchDir)
if err != nil {
log.Error("batch job manager initialization failed", "error", err)
os.Exit(2)
}
pool := worker.New(cfg.Workers, cfg.Native.ControlWorker)
pool.SetRoutingConfig(cfg.Routing)
pool.SetModelCapabilitiesConfig(cfg.ModelCapabilities)
pool.SetReliabilityConfig(cfg.Reliability)
if overrides, err := placementStore.List(rootCtx); err != nil {
log.Error("model placement overrides load failed", "error", err)
os.Exit(2)
} else {
for workerName, rule := range overrides {
if err := pool.SetPlacement(workerName, rule, true); err != nil {
// Keep orphaned rules durable when a worker is temporarily removed;
// they become active again if that worker name returns later.
log.Warn("ignoring model placement override for unknown worker", "worker", workerName, "error", err)
}
}
}
if modes, err := workerStateStore.List(rootCtx); err != nil {
log.Error("worker runtime state load failed", "error", err)
os.Exit(2)
} else {
for name, mode := range modes {
if err := pool.SetMaintenance(name, mode); err != nil {
log.Warn("ignoring worker state for unknown worker", "worker", name, "error", err)
}
}
}
if err := pool.LoadPerformance(paths.WorkerPerformance); err != nil {
log.Error("worker performance state initialization failed", "error", err)
os.Exit(2)
}
pool.StartPerformancePersistence(rootCtx, paths.WorkerPerformance, cfg.Storage.FlushInterval.Value(), func(err error) { log.Error("worker performance persistence failed", "error", err) })
tuner, err := autotune.New(cfg.AutoTuning, pool, paths.AutoTune)
if err != nil {
log.Error("auto tuning state initialization failed", "error", err)
os.Exit(2)
}
for workerName, models := range tuner.Applied() {
for model, limit := range models {
if err := pool.SetModelConcurrency(workerName, model, limit); err != nil {
log.Warn("ignoring auto-tune override", "worker", workerName, "model", model, "error", err)
}
}
}
pool.Start(rootCtx)
warmManager, err := warm.New(cfg.WarmModels, pool, paths.WarmModels)
if err != nil {
log.Error("warm model manager initialization failed", "error", err)
os.Exit(2)
}
warmManager.Start(rootCtx)
alertManager, err := alerts.New(cfg.Alerts, paths.Alerts, func() alerts.Snapshot {
st := sched.Stats(context.Background())
ws := pool.Snapshots()
aw := make([]alerts.Worker, 0, len(ws))
for _, w := range ws {
aw = append(aw, alerts.Worker{Name: w.Name, Healthy: w.Healthy, CircuitState: w.CircuitState, LastCircuitError: w.LastCircuitError, VRAMUsedBytes: w.VRAMUsedBytes, VRAMTotalBytes: w.VRAMTotalBytes})
}
return alerts.Snapshot{QueueDepth: int(st.Queued), QueueWait: st.OldestWait, Workers: aw, StorageBytes: alerts.DirSize(paths.DataDir)}
})
if err != nil {
log.Error("alerts manager initialization failed", "error", err)
os.Exit(2)
}
alertManager.Start(rootCtx)
live := liveflow.New(10*time.Second, max(512, cfg.Infrastructure.MaxRequests))
infra := infrastructure.New(cfg.Infrastructure, live, sched, pool)
infra.Start(rootCtx)
met := metrics.New()
if err := met.LoadPersistent(paths.Metrics); err != nil {
log.Error("metrics state initialization failed", "error", err)
os.Exit(2)
}
met.StartPersistence(rootCtx, paths.Metrics, cfg.Storage.FlushInterval.Value(), func(err error) { log.Error("metrics persistence failed", "error", err) })
rec, err := usage.NewWithRetention(cfg.Usage.JournalDir, cfg.Usage.Buffer, cfg.Usage.FlushInterval.Value(), usage.RetentionConfig{DetailDays: cfg.Usage.Retention.DetailDays, DailyDays: cfg.Usage.Retention.DailyDays, MonthlyMonths: cfg.Usage.Retention.MonthlyMonths, CompactionInterval: cfg.Usage.Retention.CompactionInterval.Value()}, met.DropUsage)
if err != nil {
log.Error("usage recorder initialization failed", "error", err)
os.Exit(2)
}
rec.SetRecentCapacity(cfg.UI.RecentEvents)
defer rec.Close()
otelExporter := telemetry.New(cfg.OpenTelemetry)
if otelExporter != nil {
defer func() {
cctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if err := otelExporter.Close(cctx); err != nil {
log.Warn("OpenTelemetry shutdown", "error", err)
}
}()
}
srvHandler := server.New(cfg, server.Dependencies{Auth: authenticator, Scheduler: sched, Quota: ledger, Estimator: cost.New(cfg.Cost), Workers: pool, Proxy: proxy.New(), Usage: rec, Metrics: met, Policies: policyStore, Sessions: uiSessions, Live: live, Infrastructure: infra, Logger: log, ConfigStore: configStore, PlacementStore: placementStore, WorkerStateStore: workerStateStore, AutoTune: tuner, OpenTelemetry: otelExporter, WarmModels: warmManager, Alerts: alertManager, Conversations: conversationStore, BatchJobs: batchManager})
batchManager.Start(rootCtx, srvHandler.ExecuteBatch)
hs := &http.Server{Addr: cfg.Server.Listen, Handler: srvHandler.Handler(), ReadHeaderTimeout: cfg.Server.ReadHeaderTimeout.Value(), IdleTimeout: cfg.Server.IdleTimeout.Value(), MaxHeaderBytes: 1 << 20}
go func() {
<-rootCtx.Done()
ctx, c := context.WithTimeout(context.Background(), 15*time.Second)
defer c()
_ = hs.Shutdown(ctx)
}()
log.Info("ollama fair gateway starting", "listen", cfg.Server.Listen, "workers", len(cfg.Workers), "coordination", "in-memory", "node", infra.NodeName(), "quota", cfg.Quota.Enabled, "storage", cfg.Storage.DataDir)
if cfg.Server.TLSCert != "" || cfg.Server.TLSKey != "" {
err = hs.ListenAndServeTLS(cfg.Server.TLSCert, cfg.Server.TLSKey)
} else {
err = hs.ListenAndServe()
}
if err != nil && err != http.ErrServerClosed {
log.Error("server failed", "error", err)
os.Exit(1)
}
if warmManager != nil {
wctx, wcancel := context.WithTimeout(context.Background(), 5*time.Second)
if err := warmManager.Wait(wctx); err != nil {
log.Warn("warm model actions still active during shutdown", "error", err)
}
wcancel()
}
if batchManager != nil {
bctx, bcancel := context.WithTimeout(context.Background(), 10*time.Second)
if err := batchManager.Wait(bctx); err != nil {
log.Warn("batch attempts still active during shutdown", "error", err)
}
bcancel()
}
rec.Close()
if quotaMemory != nil {
if err := quotaMemory.SavePersistent(paths.Quota); err != nil {
log.Error("final quota persistence failed", "error", err)
}
}
if err := met.SavePersistent(paths.Metrics); err != nil {
log.Error("final metrics persistence failed", "error", err)
}
if err := pool.SavePerformance(paths.WorkerPerformance); err != nil {
log.Error("final worker performance persistence failed", "error", err)
}
log.Info("gateway stopped")
}
+222
View File
@@ -0,0 +1,222 @@
package main
import (
"encoding/json"
"fmt"
"io"
"net"
"net/http"
"net/url"
"os"
"path/filepath"
"strings"
"time"
"github.com/example/ollama-fair-gateway/internal/config"
"github.com/example/ollama-fair-gateway/internal/state"
)
type effectiveConfig struct {
Bootstrap *config.Config
Config *config.Config
Paths state.Paths
Store *state.ConfigStore
Persistent bool
}
func loadEffectiveConfig(configPath string) (*effectiveConfig, error) {
bootstrapCfg, err := config.Load(configPath)
if err != nil {
return nil, fmt.Errorf("configuration error: %w", err)
}
paths := state.Resolve(bootstrapCfg.Storage)
configStore := state.NewConfigStore(paths.Config)
cfg := bootstrapCfg
persistentCfg, ok, err := configStore.LoadIfExists(bootstrapCfg)
if err != nil {
return nil, fmt.Errorf("persistent configuration error (%s): %w", configStore.Path(), err)
}
if ok {
if persistentCfg.Storage != bootstrapCfg.Storage {
return nil, fmt.Errorf("persistent configuration changes bootstrap-only storage settings: bootstrap=%+v persistent=%+v", bootstrapCfg.Storage, persistentCfg.Storage)
}
cfg = persistentCfg
}
return &effectiveConfig{
Bootstrap: bootstrapCfg,
Config: cfg,
Paths: state.Resolve(cfg.Storage),
Store: configStore,
Persistent: ok,
}, nil
}
type configCheckResult struct {
Status string `json:"status"`
ConfigPath string `json:"config_path"`
PersistentOverride bool `json:"persistent_override"`
PersistentPath string `json:"persistent_path"`
Workers int `json:"workers"`
DataDir string `json:"data_dir"`
DataDirAbsolute string `json:"data_dir_absolute"`
StorageWritable bool `json:"storage_writable"`
Warnings []string `json:"warnings,omitempty"`
}
func runConfigCheck(configPath string, out io.Writer) error {
loaded, err := loadEffectiveConfig(configPath)
if err != nil {
return err
}
if err := checkWritableDir(loaded.Paths.DataDir); err != nil {
return fmt.Errorf("storage data_dir %q is not writable: %w", loaded.Paths.DataDir, err)
}
absDir, err := filepath.Abs(loaded.Paths.DataDir)
if err != nil {
absDir = loaded.Paths.DataDir
}
res := configCheckResult{
Status: "ok",
ConfigPath: configPath,
PersistentOverride: loaded.Persistent,
PersistentPath: loaded.Store.Path(),
Workers: len(loaded.Config.Workers),
DataDir: loaded.Paths.DataDir,
DataDirAbsolute: absDir,
StorageWritable: true,
Warnings: configWarnings(loaded.Config),
}
enc := json.NewEncoder(out)
enc.SetIndent("", " ")
return enc.Encode(res)
}
func checkWritableDir(dir string) error {
if err := os.MkdirAll(dir, 0700); err != nil {
return err
}
f, err := os.CreateTemp(dir, ".gateway-preflight-*")
if err != nil {
return err
}
name := f.Name()
defer os.Remove(name)
if err := f.Chmod(0600); err != nil {
_ = f.Close()
return err
}
if _, err := f.WriteString("ok\n"); err != nil {
_ = f.Close()
return err
}
if err := f.Sync(); err != nil {
_ = f.Close()
return err
}
return f.Close()
}
func configWarnings(cfg *config.Config) []string {
var warnings []string
if control := strings.TrimSpace(cfg.Native.ControlWorker); control != "" {
found := false
for _, w := range cfg.Workers {
if w.Name == control {
found = true
break
}
}
if !found {
warnings = append(warnings, fmt.Sprintf("native.control_worker %q does not match any configured worker; management requests will fall back to another healthy worker", control))
}
}
for _, w := range cfg.Workers {
if !w.LocalSystemStats {
continue
}
u, err := url.Parse(w.URL)
if err != nil {
continue
}
host := u.Hostname()
if host == "" || isLoopbackHost(host) {
continue
}
warnings = append(warnings, fmt.Sprintf("worker %q has local_system_stats=true but URL host %q is remote; local system stats describe the gateway host, not that worker (use telemetry_url or disable local_system_stats)", w.Name, host))
}
if cfg.Auth.IPBypassUseForwardedIP && len(cfg.Auth.IPBypass) > 0 {
warnings = append(warnings, "auth.ip_bypass_use_forwarded_ip=true allows X-Forwarded-For-derived addresses to satisfy credential-free IP bypass; use only with tightly restricted trusted_proxies and network ACLs")
}
for _, raw := range cfg.Auth.TrustedProxies {
if broadTrustedProxyCIDR(raw) {
warnings = append(warnings, fmt.Sprintf("auth.trusted_proxies contains broad CIDR %q; any directly reachable peer in that range can influence X-Forwarded-For-derived client_ip, so prefer exact proxy addresses", raw))
}
}
if cfg.UI.Enabled && !cfg.UI.OIDC.Enabled && len(cfg.Auth.APIKeys) == 0 {
warnings = append(warnings, "UI is enabled without UI OIDC and without bootstrap API keys; remote UI login depends entirely on IP bypass rules")
}
if cfg.UI.Enabled && !cfg.UI.SecureCookies && cfg.UI.OIDC.Enabled {
warnings = append(warnings, "ui.secure_cookies=false while UI OIDC is enabled; enable secure cookies when the browser reaches the gateway over HTTPS")
}
if cfg.Server.MetricsPublic {
warnings = append(warnings, "server.metrics_public=true exposes gateway metrics without authentication")
}
if cfg.ModelCapabilities.Context.MaxRequestedTokens == -1 {
warnings = append(warnings, "model_capabilities.context.max_requested_tokens=-1 removes the gateway-side context cap; large native num_ctx requests can cause substantial KV-cache/VRAM pressure")
}
if cfg.ModelCapabilities.Context.DefaultWorkerTokens == -1 {
warnings = append(warnings, "model_capabilities.context.default_worker_tokens=-1 falls back to the theoretical model maximum for unloaded models without Modelfile num_ctx; configure an explicit worker default for predictable memory use")
}
if !filepath.IsAbs(cfg.Storage.DataDir) {
warnings = append(warnings, fmt.Sprintf("storage.data_dir %q is relative and therefore depends on the process working directory", cfg.Storage.DataDir))
}
return warnings
}
func broadTrustedProxyCIDR(raw string) bool {
ip, n, err := net.ParseCIDR(strings.TrimSpace(raw))
if err != nil || ip == nil || n == nil || ip.IsLoopback() {
return false
}
ones, bits := n.Mask.Size()
if bits == 32 {
return ones < 24
}
if bits == 128 {
return ones < 64
}
return false
}
func isLoopbackHost(host string) bool {
if strings.EqualFold(host, "localhost") {
return true
}
ip := net.ParseIP(host)
return ip != nil && ip.IsLoopback()
}
func runProbe(rawURL string, timeout time.Duration) error {
if timeout <= 0 {
return fmt.Errorf("probe timeout must be > 0")
}
u, err := url.Parse(rawURL)
if err != nil || u.Scheme == "" || u.Host == "" {
return fmt.Errorf("invalid probe URL %q", rawURL)
}
client := &http.Client{Timeout: timeout}
req, err := http.NewRequest(http.MethodGet, rawURL, nil)
if err != nil {
return err
}
resp, err := client.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
_, _ = io.Copy(io.Discard, resp.Body)
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return fmt.Errorf("probe returned HTTP %d", resp.StatusCode)
}
return nil
}
+154
View File
@@ -0,0 +1,154 @@
package main
import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"testing"
"time"
"github.com/example/ollama-fair-gateway/internal/config"
"github.com/example/ollama-fair-gateway/internal/state"
)
func writeBootstrapConfig(t *testing.T, dataDir string) string {
t.Helper()
path := filepath.Join(t.TempDir(), "config.json")
body := `{
"auth":{"api_keys":[{"name":"admin","key":"01234567890123456789012345678901","tenant":"t","subject":"s","scopes":["gateway:admin"]}]},
"workers":[{"name":"remote","url":"http://10.2.10.48:11434","local_system_stats":true}],
"ui":{"enabled":true,"path":"/admin","title":"Bootstrap"},
"storage":{"data_dir":` + quoteJSON(dataDir) + `,"config_file":"gateway-config.json"}
}`
if err := os.WriteFile(path, []byte(body), 0600); err != nil {
t.Fatal(err)
}
return path
}
func quoteJSON(s string) string {
b, _ := json.Marshal(s)
return string(b)
}
func TestLoadEffectiveConfigUsesPersistentOverrideButBootstrapSecrets(t *testing.T) {
dataDir := t.TempDir()
path := writeBootstrapConfig(t, dataDir)
base, err := config.Load(path)
if err != nil {
t.Fatal(err)
}
persistent := *base
persistent.UI.Title = "Persistent"
store := state.NewConfigStore(state.Resolve(base.Storage).Config)
if err := store.Save(&persistent); err != nil {
t.Fatal(err)
}
loaded, err := loadEffectiveConfig(path)
if err != nil {
t.Fatal(err)
}
if !loaded.Persistent || loaded.Config.UI.Title != "Persistent" {
t.Fatalf("persistent override not loaded: %#v", loaded)
}
if got := loaded.Config.Auth.APIKeys[0].Key; got != "01234567890123456789012345678901" {
t.Fatalf("bootstrap secret was not restored, got %q", got)
}
}
func TestRunConfigCheckReportsWritableStorageAndWarning(t *testing.T) {
dataDir := filepath.Join(t.TempDir(), "state")
path := writeBootstrapConfig(t, dataDir)
var out bytes.Buffer
if err := runConfigCheck(path, &out); err != nil {
t.Fatal(err)
}
var got configCheckResult
if err := json.Unmarshal(out.Bytes(), &got); err != nil {
t.Fatal(err)
}
if got.Status != "ok" || !got.StorageWritable || got.Workers != 1 {
t.Fatalf("unexpected result: %+v", got)
}
if len(got.Warnings) == 0 || !strings.Contains(strings.Join(got.Warnings, "\n"), "local_system_stats=true") {
t.Fatalf("expected remote local-system-stats warning, got %#v", got.Warnings)
}
}
func TestRunProbe(t *testing.T) {
ok := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusNoContent)
}))
defer ok.Close()
if err := runProbe(ok.URL, time.Second); err != nil {
t.Fatal(err)
}
bad := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
http.Error(w, "not ready", http.StatusServiceUnavailable)
}))
defer bad.Close()
if err := runProbe(bad.URL, time.Second); err == nil || !strings.Contains(err.Error(), "503") {
t.Fatalf("expected HTTP 503 probe error, got %v", err)
}
}
func TestComposeDoesNotRepeatEntrypointAndHasHealthcheck(t *testing.T) {
b, err := os.ReadFile("../../docker-compose.yml")
if err != nil {
t.Fatal(err)
}
s := string(b)
if strings.Contains(s, `command: ["/ollama-gateway"`) {
t.Fatal("compose command must not repeat the image ENTRYPOINT")
}
if !strings.Contains(s, `command: ["-config"`) {
t.Fatal("compose command must pass -config as an ENTRYPOINT argument")
}
if !strings.Contains(s, "healthcheck:") || !strings.Contains(s, "/healthz") {
t.Fatal("compose must define a liveness healthcheck")
}
}
func TestConfigWarningsFlagForwardedBypassAndBroadTrustedProxy(t *testing.T) {
cfg := &config.Config{
Auth: config.AuthConfig{
IPBypassUseForwardedIP: true,
TrustedProxies: []string{"10.0.0.0/8", "127.0.0.1/8"},
IPBypass: []config.IPBypassConfig{{
CIDRs: []string{"127.0.0.1/32"},
Tenant: "local",
}},
},
Storage: config.StorageConfig{DataDir: "/var/lib/ollama-gateway"},
}
warnings := strings.Join(configWarnings(cfg), "\n")
if !strings.Contains(warnings, "ip_bypass_use_forwarded_ip=true") {
t.Fatalf("missing forwarded bypass warning: %s", warnings)
}
if !strings.Contains(warnings, `broad CIDR "10.0.0.0/8"`) {
t.Fatalf("missing broad trusted proxy warning: %s", warnings)
}
if strings.Contains(warnings, `broad CIDR "127.0.0.1/8"`) {
t.Fatalf("loopback trusted proxy should not be flagged as broad: %s", warnings)
}
}
func TestConfigWarningsFlagUnlimitedContextPolicy(t *testing.T) {
cfg := &config.Config{
ModelCapabilities: config.ModelCapabilitiesConfig{Context: config.ContextPolicyConfig{MaxRequestedTokens: -1, DefaultWorkerTokens: -1}},
Storage: config.StorageConfig{DataDir: "/var/lib/ollama-gateway"},
}
warnings := strings.Join(configWarnings(cfg), "\n")
if !strings.Contains(warnings, "max_requested_tokens=-1") {
t.Fatalf("missing unlimited context cap warning: %s", warnings)
}
if !strings.Contains(warnings, "default_worker_tokens=-1") {
t.Fatalf("missing model-max fallback warning: %s", warnings)
}
}
+180
View File
@@ -0,0 +1,180 @@
package main
import (
"context"
"encoding/json"
"flag"
"fmt"
"log"
"net"
"net/http"
"os"
"os/signal"
"strings"
"syscall"
"time"
"github.com/example/ollama-fair-gateway/internal/hoststats"
)
type telemetry 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"`
Source string `json:"source,omitempty"`
UpdatedAt time.Time `json:"updated_at,omitempty"`
Error string `json:"error,omitempty"`
}
type collector struct {
nvidia bool
nvidiaGPU string
amd bool
amdDevice string
}
func (c collector) collect(ctx context.Context) telemetry {
out := telemetry{UpdatedAt: time.Now().UTC()}
var errs []string
mctx, cancel := context.WithTimeout(ctx, 1500*time.Millisecond)
m, err := hoststats.ReadMemory(mctx)
cancel()
if err != nil {
errs = append(errs, "memory: "+err.Error())
} else {
out.MemoryTotalBytes, out.MemoryUsedBytes = m.TotalBytes, m.UsedBytes
out.Source = appendSource(out.Source, "host-memory")
}
if c.nvidia {
gctx, cancel := context.WithTimeout(ctx, 1500*time.Millisecond)
g, err := hoststats.ReadNVIDIA(gctx, c.nvidiaGPU)
cancel()
if err != nil {
errs = append(errs, "nvidia: "+err.Error())
} else {
out.VRAMUsedBytes, out.VRAMTotalBytes = g.MemoryUsedBytes, g.MemoryTotalBytes
out.GPUUtilizationPct, out.GPUTemperatureC, out.GPUPowerWatts = g.UtilizationPercent, g.TemperatureC, g.PowerWatts
out.Source = appendSource(out.Source, "nvidia-smi")
}
}
if c.amd {
gctx, cancel := context.WithTimeout(ctx, 1500*time.Millisecond)
g, err := hoststats.ReadAMD(gctx, c.amdDevice)
cancel()
if err != nil {
errs = append(errs, "amd: "+err.Error())
} else {
out.VRAMUsedBytes, out.VRAMTotalBytes = g.MemoryUsedBytes, g.MemoryTotalBytes
out.GPUUtilizationPct, out.GPUTemperatureC, out.GPUPowerWatts = g.UtilizationPercent, g.TemperatureC, g.PowerWatts
out.Source = appendSource(out.Source, "amdgpu-sysfs")
}
}
out.Error = strings.Join(errs, "; ")
return out
}
func appendSource(cur, next string) string {
if cur == "" {
return next
}
return cur + "+" + next
}
type cidrAllowlist struct{ nets []*net.IPNet }
func parseCIDRs(raw string) (cidrAllowlist, error) {
var out cidrAllowlist
for _, part := range strings.Split(raw, ",") {
part = strings.TrimSpace(part)
if part == "" {
continue
}
_, n, err := net.ParseCIDR(part)
if err != nil {
return out, fmt.Errorf("invalid allow CIDR %q: %w", part, err)
}
out.nets = append(out.nets, n)
}
if len(out.nets) == 0 {
return out, fmt.Errorf("at least one allow CIDR is required")
}
return out, nil
}
func (a cidrAllowlist) allowed(remote string) bool {
host, _, err := net.SplitHostPort(remote)
if err != nil {
host = remote
}
ip := net.ParseIP(strings.Trim(host, "[]"))
if ip == nil {
return false
}
for _, n := range a.nets {
if n.Contains(ip) {
return true
}
}
return false
}
func main() {
listen := flag.String("listen", "127.0.0.1:11500", "listen address")
path := flag.String("path", "/telemetry", "telemetry path")
allow := flag.String("allow-cidrs", "127.0.0.1/32,::1/128", "comma-separated client CIDRs allowed to read telemetry")
nvidia := flag.Bool("nvidia-smi", false, "collect NVIDIA telemetry with nvidia-smi")
nvidiaGPU := flag.String("nvidia-gpu", "", "optional nvidia-smi GPU selector")
amd := flag.Bool("amd-sysfs", false, "collect Linux AMDGPU telemetry from sysfs")
amdDevice := flag.String("amd-device", "", "optional AMDGPU device path such as /sys/class/drm/card0/device; empty auto-detects")
once := flag.Bool("once", false, "print one telemetry sample as JSON and exit")
flag.Parse()
if !strings.HasPrefix(*path, "/") {
log.Fatal("-path must begin with /")
}
acl, err := parseCIDRs(*allow)
if err != nil {
log.Fatal(err)
}
c := collector{nvidia: *nvidia, nvidiaGPU: *nvidiaGPU, amd: *amd, amdDevice: *amdDevice}
if *once {
enc := json.NewEncoder(os.Stdout)
enc.SetIndent("", " ")
if err := enc.Encode(c.collect(context.Background())); err != nil {
log.Fatal(err)
}
return
}
mux := http.NewServeMux()
mux.HandleFunc("/healthz", func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusNoContent) })
mux.HandleFunc(*path, func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
w.WriteHeader(http.StatusMethodNotAllowed)
return
}
if !acl.allowed(r.RemoteAddr) {
http.Error(w, "forbidden", http.StatusForbidden)
return
}
w.Header().Set("Content-Type", "application/json")
w.Header().Set("Cache-Control", "no-store")
w.Header().Set("X-Content-Type-Options", "nosniff")
_ = json.NewEncoder(w).Encode(c.collect(r.Context()))
})
srv := &http.Server{Addr: *listen, Handler: mux, ReadHeaderTimeout: 5 * time.Second, IdleTimeout: time.Minute}
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
defer stop()
go func() {
<-ctx.Done()
shutdown, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
_ = srv.Shutdown(shutdown)
}()
log.Printf("worker telemetry listening on http://%s%s allowed=%s nvidia=%t amd=%t", *listen, *path, *allow, *nvidia, *amd)
if err := srv.ListenAndServe(); err != nil && err != http.ErrServerClosed {
log.Fatal(err)
}
}
+19
View File
@@ -0,0 +1,19 @@
package main
import "testing"
func TestCIDRAllowlist(t *testing.T) {
a, err := parseCIDRs("127.0.0.1/32,10.2.19.0/24")
if err != nil {
t.Fatal(err)
}
if !a.allowed("10.2.19.42:1234") || !a.allowed("127.0.0.1:1") || a.allowed("10.2.18.1:5") {
t.Fatal("unexpected allowlist result")
}
}
func TestAppendSource(t *testing.T) {
if got := appendSource("host-memory", "amdgpu-sysfs"); got != "host-memory+amdgpu-sysfs" {
t.Fatal(got)
}
}