ollama-proxy/internal/proxy/handler.go
Atte149 3963eede70 feat: ollama-proxy — multi-account Ollama Cloud reverse proxy
Go reverse proxy for ollama.com that balances requests across multiple
API keys with round-robin and failover on 429/5xx. Exposes both native
Ollama API (/api/*) and OpenAI-compatible (/v1/*) passthrough.

- Round-robin balancer with per-account cooldown (60s default)
- Pre-stream failover: 429 → cooldown + next account; 5xx → next account
- Streaming invariant: once 2xx starts streaming, no account switch
- SSE (text/event-stream) and NDJSON passthrough with http.Flusher
- CLI: accounts add/list/remove/set-base-url, serve, version
- Accounts stored in ~/.config/ollama-proxy/accounts.json (chmod 0600)
- systemd unit (User=dueattendant149, 127.0.0.1:11435, Restart=always)
- 43 tests (unit + integration with httptest upstream)
- opencode integration: custom provider 'ocp' with explicit model list
- Requires NO_PROXY=127.0.0.1,localhost when HTTP_PROXY is set
2026-06-19 13:58:13 +03:00

294 lines
8.6 KiB
Go

package proxy
import (
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"strings"
"time"
)
// Handler is the http.Handler that proxies incoming requests to Ollama Cloud
// through the balancer, applying per-account Authorization and failover logic.
type Handler struct {
Balancer *Balancer
BaseURL string // upstream root, e.g. "https://ollama.com"
Client *http.Client
Log *slog.Logger
Retries int // max attempts per request (= number of accounts to try)
}
// NewHandler builds a Handler with sensible HTTP client defaults (no timeout
// on the overall request — streaming responses can be long; per-read timeout
// is governed by the caller's context).
func NewHandler(b *Balancer, baseURL string, retries int, log *slog.Logger) *Handler {
return &Handler{
Balancer: b,
BaseURL: strings.TrimRight(baseURL, "/"),
Client: &http.Client{
// No overall timeout: streaming chat may take minutes. The caller's
// request context (cancelled when the client disconnects) still
// propagates through to upstream via NewWithContext.
Timeout: 0,
CheckRedirect: func(req *http.Request, via []*http.Request) error {
return http.ErrUseLastResponse // don't follow redirects automatically
},
},
Log: log,
Retries: retries,
}
}
// allowedPaths is the allowlist of upstream paths the proxy will forward.
// Anything else returns 404 — we never proxy arbitrary paths.
var allowedPaths = map[string]bool{
"/api/chat": true,
"/api/generate": true,
"/api/tags": true,
"/api/show": true,
"/api/ps": true,
"/api/version": true,
"/api/delete": true,
"/v1/chat/completions": true,
"/v1/completions": true,
"/v1/models": true,
"/v1/embeddings": true,
"/v1/files": true,
}
// isAllowed reports whether a path (possibly with a trailing slash or query)
// matches one of the allowed upstream endpoints.
func isAllowed(path string) bool {
if allowedPaths[path] {
return true
}
// allow sub-paths like /v1/files/<id> under an allowed prefix.
for prefix := range allowedPaths {
if strings.HasPrefix(path, prefix+"/") {
return true
}
}
return false
}
// ServeHTTP proxies a single client request to Ollama Cloud with failover.
func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
if !isAllowed(r.URL.Path) {
http.NotFound(w, r)
return
}
if h.Balancer.Len() == 0 {
writeError(w, http.StatusServiceUnavailable, "no accounts configured")
return
}
maxAttempts := h.Retries
if maxAttempts < 1 {
maxAttempts = 1
}
if maxAttempts > h.Balancer.Len() {
maxAttempts = h.Balancer.Len()
}
h.proxyWithFailover(w, r, maxAttempts)
}
// proxyWithFailover tries up to maxAttempts accounts. The first attempt that
// begins streaming (returns headers + a 2xx) commits: we copy the rest of the
// response to the client verbatim and stop retrying. Pre-stream 429/5xx move
// on to the next account.
func (h *Handler) proxyWithFailover(w http.ResponseWriter, r *http.Request, maxAttempts int) {
ctx := r.Context()
var lastErr error
for attempt := 0; attempt < maxAttempts; attempt++ {
acct, err := h.Balancer.Next()
if err != nil {
// No accounts available.
switch err.(type) {
case ErrAllCooldown:
// All accounts rate-limited upstream → propagate 429 to client.
msg := err.Error()
if lastErr != nil {
msg = msg + "; last upstream error: " + lastErr.Error()
}
writeError(w, http.StatusTooManyRequests, msg)
return
case ErrNoAccounts:
writeError(w, http.StatusServiceUnavailable, err.Error())
return
default:
writeError(w, http.StatusServiceUnavailable, err.Error())
return
}
}
// Build the upstream request, substituting the account's key.
upstreamURL := h.BaseURL + r.URL.RequestURI()
// Per-account BaseURL override wins when set.
if acct.BaseURL != "" {
upstreamURL = strings.TrimRight(acct.BaseURL, "/") + r.URL.RequestURI()
}
req, err := http.NewRequestWithContext(ctx, r.Method, upstreamURL, r.Body)
if err != nil {
lastErr = fmt.Errorf("build request: %w", err)
continue
}
// Copy headers, replacing Authorization with the account key.
copyHeaders(req.Header, r.Header)
req.Header.Set("Authorization", "Bearer "+acct.APIKey)
req.Host = "" // let the URL determine Host
start := time.Now()
resp, err := h.Client.Do(req)
latency := time.Since(start)
if err != nil {
// Network/timeout error: log and try next account.
h.Log.Warn("upstream request failed",
"account", acct.Name, "account_id", acct.ID,
"path", r.URL.Path, "latency_ms", latency.Milliseconds(), "err", err)
lastErr = err
continue
}
// Pre-stream decision: 429 or 5xx → cooldown + failover.
if resp.StatusCode == http.StatusTooManyRequests || resp.StatusCode >= 500 {
bodyPreview := readPreview(resp.Body, 256)
resp.Body.Close()
h.Log.Warn("upstream rejected, failover",
"account", acct.Name, "account_id", acct.ID,
"path", r.URL.Path, "status", resp.StatusCode,
"latency_ms", latency.Milliseconds(), "preview", bodyPreview)
if resp.StatusCode == http.StatusTooManyRequests {
h.Balancer.MarkCooldown(acct.ID)
}
lastErr = fmt.Errorf("upstream %d: %s", resp.StatusCode, bodyPreview)
continue
}
// 2xx (or other non-retriable): commit and stream.
h.copyResponse(w, resp)
h.Log.Info("proxied",
"account", acct.Name, "account_id", acct.ID,
"path", r.URL.Path, "status", resp.StatusCode,
"latency_ms", latency.Milliseconds(), "stream", isStreaming(resp))
return
}
// Exhausted retries. Distinguish "everything is rate-limited" (→ 429) from
// "mixed upstream errors" (→ 502).
if h.Balancer.Len() > 0 {
allCooldown := true
status := h.Balancer.Status()
for _, a := range h.Balancer.Accounts() {
if !status[a.ID].InCooldown {
allCooldown = false
break
}
}
if allCooldown {
msg := "all accounts rate-limited"
if lastErr != nil {
msg = msg + ": " + lastErr.Error()
}
writeError(w, http.StatusTooManyRequests, msg)
return
}
}
if lastErr != nil {
writeError(w, http.StatusBadGateway, "all accounts failed: "+lastErr.Error())
return
}
writeError(w, http.StatusBadGateway, "all accounts failed")
}
// copyResponse streams the upstream response body to the client. For chunked /
// SSE responses we flush after every read so tokens reach the client immediately.
func (h *Handler) copyResponse(w http.ResponseWriter, resp *http.Response) {
defer resp.Body.Close()
// Copy headers (except hop-by-hop ones).
for k, vs := range resp.Header {
for _, v := range vs {
w.Header().Add(k, v)
}
}
w.WriteHeader(resp.StatusCode)
flusher, _ := w.(http.Flusher)
buf := make([]byte, 4096)
for {
n, err := resp.Body.Read(buf)
if n > 0 {
if _, werr := w.Write(buf[:n]); werr != nil {
// client went away; stop copying silently
return
}
if flusher != nil {
flusher.Flush()
}
}
if err != nil {
if !errors.Is(err, io.EOF) {
h.Log.Debug("upstream body read ended", "err", err)
}
return
}
}
}
// copyHeaders duplicates src into dst, dropping hop-by-hop headers and any
// Authorization that the client may have sent (we always set our own).
func copyHeaders(dst, src http.Header) {
hopByHop := []string{
"Connection", "Keep-Alive", "Proxy-Authenticate", "Proxy-Authorization",
"Te", "Trailers", "Transfer-Encoding", "Upgrade",
"Authorization", // always overwritten by the account key
}
for k, vs := range src {
skip := false
for _, h := range hopByHop {
if strings.EqualFold(k, h) {
skip = true
break
}
}
if skip {
continue
}
for _, v := range vs {
dst.Add(k, v)
}
}
}
// readPreview reads up to n bytes from r and returns them as a string, always
// closing the reader.
func readPreview(r io.ReadCloser, n int) string {
defer r.Close()
buf := make([]byte, n)
m, _ := r.Read(buf)
return string(buf[:m])
}
// isStreaming reports whether the response is streaming (SSE or NDJSON/chunked).
func isStreaming(resp *http.Response) bool {
ct := resp.Header.Get("Content-Type")
if strings.Contains(ct, "text/event-stream") {
return true
}
if strings.Contains(ct, "application/x-ndjson") {
return true
}
te := resp.Header.Get("Transfer-Encoding")
return strings.Contains(strings.ToLower(te), "chunked")
}
// writeError emits a JSON-formatted error to the client, mirroring Ollama's
// error shape so OpenAI-compatible clients can parse it.
func writeError(w http.ResponseWriter, status int, msg string) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(status)
_, _ = fmt.Fprintf(w, `{"error":{"message":%q,"type":"ollama_proxy"}}`, msg)
}