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