package proxy import ( "encoding/json" "errors" "fmt" "io" "log/slog" "net/http" "strings" "time" "github.com/Atte149/ollama-proxy/internal/config" ) // Handler is the http.Handler that proxies incoming requests to upstream // providers (Ollama Cloud and OpenCode Go) through the balancer, applying // per-account Authorization and failover logic. type Handler struct { Balancer *Balancer Log *slog.Logger Client *http.Client Retries int // max attempts per request (= number of accounts to try) // DefaultBaseURL, when non-empty, overrides the provider-default upstream // root for accounts that have no explicit BaseURL. This is primarily used // by the --base-url flag for local testing; in production accounts rely on // the provider default. DefaultBaseURL string } // NewHandler builds a Handler with sensible HTTP client defaults. func NewHandler(b *Balancer, baseURL string, retries int, log *slog.Logger) *Handler { return &Handler{ Balancer: b, Log: log, Client: &http.Client{ Timeout: 0, CheckRedirect: func(req *http.Request, via []*http.Request) error { return http.ErrUseLastResponse }, }, Retries: retries, DefaultBaseURL: baseURL, } } // upstreamBaseURL resolves the canonical upstream root for an account: // 1. The account's own BaseURL override (highest priority) // 2. The handler's DefaultBaseURL — but ONLY for ollama-cloud accounts, // since --base-url is an Ollama-specific flag used for local testing // 3. The provider default func (h *Handler) upstreamBaseURL(acct config.Account) string { if acct.BaseURL != "" { return strings.TrimRight(acct.BaseURL, "/") } if h.DefaultBaseURL != "" && acct.Provider == config.ProviderOllamaCloud { return strings.TrimRight(h.DefaultBaseURL, "/") } return strings.TrimRight(config.DefaultProviderBaseURL(acct.Provider), "/") } // allowedPaths is the allowlist of paths the proxy will forward. The same set // is used for both providers; the /go/ prefix routes to OpenCode Go. 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 matches an allowed upstream endpoint. func isAllowed(path string) bool { if allowedPaths[path] { return true } for prefix := range allowedPaths { if strings.HasPrefix(path, prefix+"/") { return true } } return false } // routePath normalises the incoming path and returns (upstreamPath, forcedProvider). // Paths under /go/ force the OpenCode Go provider and have the /go prefix // stripped before forwarding. Otherwise the provider is chosen by the request // body model (or defaults to ollama-cloud when no model is present). func routePath(path string) (upstreamPath string, forced config.ProviderType, isGo bool) { if strings.HasPrefix(path, "/go/") || path == "/go" { return "/" + strings.TrimPrefix(path, "/go/"), config.ProviderOpenCodeGo, true } return path, "", false } // modelFromBody extracts the "model" field from a buffered request body. // Returns "" when the body is nil, empty, or does not contain a parseable // JSON object with a model field. Does not consume the buffer (resets position). func modelFromBody(body *bufferedBody) string { if body == nil || len(body.data) == 0 { return "" } var probe struct { Model string `json:"model"` } if err := json.Unmarshal(body.data, &probe); err != nil { return "" } return probe.Model } // bufferedBody wraps a []byte so it can be replayed across failover attempts. type bufferedBody struct { data []byte pos int } func (b *bufferedBody) Read(p []byte) (int, error) { if b.pos >= len(b.data) { return 0, io.EOF } n := copy(p, b.data[b.pos:]) b.pos += n return n, nil } func (b *bufferedBody) Close() error { return nil } func (b *bufferedBody) reset() { b.pos = 0 } // bufferRequestBody reads the request body (up to a cap) into a buffer so it // can be replayed across failover attempts. Bodies larger than the cap are // not buffered and the request is not retried. Returns nil for streaming-safe // GET/DELETE requests. func bufferRequestBody(r *http.Request) *bufferedBody { if r.Body == nil || r.Method == http.MethodGet || r.Method == http.MethodDelete || r.Method == http.MethodHead { return nil } const cap = 8 << 20 // 8 MiB data, err := io.ReadAll(io.LimitReader(r.Body, cap+1)) if err != nil || len(data) > cap { return nil // too large or unreadable; disable retries by returning nil } return &bufferedBody{data: data} } // ServeHTTP proxies a single client request to the upstream with failover. func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { upstreamPath, forced, isGo := routePath(r.URL.Path) _ = isGo if !isAllowed(upstreamPath) { http.NotFound(w, r) return } if h.Balancer.Len() == 0 { writeError(w, http.StatusServiceUnavailable, "no accounts configured") return } // Buffer the request body up front so we can both extract the model field // and replay the body across failover attempts. This must happen BEFORE // requestModel reads from r.Body. bodyBuf := bufferRequestBody(r) // Determine eligible providers for this request. var eligible []config.ProviderType switch { case forced != "": // /go/* path forces the OpenCode Go provider. eligible = []config.ProviderType{forced} case upstreamPath == "/v1/models" || upstreamPath == "/api/tags": // Listing endpoints: handled specially below (merged response). h.handleModelsEndpoint(w, r, upstreamPath) return case upstreamPath == "/api/version": // Version is Ollama-specific; route to Ollama Cloud. eligible = []config.ProviderType{config.ProviderOllamaCloud} default: // Chat/generate/etc.: route by the model in the body. model := modelFromBody(bodyBuf) if model == "" { eligible = nil // any provider } else { eligible = ProvidersForModel(model) } } maxAttempts := h.Retries if maxAttempts < 1 { maxAttempts = 1 } if n := h.Balancer.Len(); maxAttempts > n { maxAttempts = n } h.proxyWithFailover(w, r, upstreamPath, eligible, bodyBuf, maxAttempts) } // handleModelsEndpoint returns a merged list of all known models for // /v1/models (OpenAI format) or /api/tags (Ollama format). func (h *Handler) handleModelsEndpoint(w http.ResponseWriter, r *http.Request, path string) { now := time.Now().Unix() models := AllModels() if path == "/v1/models" { w.Header().Set("Content-Type", "application/json") w.WriteHeader(200) var sb strings.Builder sb.WriteString(`{"object":"list","data":[`) for i, m := range models { if i > 0 { sb.WriteByte(',') } // The id is the raw model ID (used in API requests). The name // includes a provider label for unique models so users can tell // them apart in the opencode model picker. name := m if !IsCommonModel(m) { name = m + " (" + ProviderLabel(ProvidersForModel(m)[0]) + ")" } fmt.Fprintf(&sb, `{"id":%q,"object":"model","created":%d,"owned_by":"ollama-proxy"}`, m, now) _ = name // name is conveyed via opencode config, not this endpoint } sb.WriteString(`]}`) _, _ = w.Write([]byte(sb.String())) return } // /api/tags — Ollama format w.Header().Set("Content-Type", "application/json") w.WriteHeader(200) var sb strings.Builder sb.WriteString(`{"models":[`) for i, m := range models { if i > 0 { sb.WriteByte(',') } fmt.Fprintf(&sb, `{"name":%q,"model":%q,"modified_at":"","size":0,"digest":""}`, m, m) } sb.WriteString(`]}`) _, _ = w.Write([]byte(sb.String())) } // proxyWithFailover tries up to maxAttempts accounts eligible for the request. func (h *Handler) proxyWithFailover(w http.ResponseWriter, r *http.Request, upstreamPath string, eligible []config.ProviderType, body *bufferedBody, maxAttempts int) { ctx := r.Context() var lastErr error for attempt := 0; attempt < maxAttempts; attempt++ { acct, err := h.Balancer.NextFor(eligible) if err != nil { switch err.(type) { case ErrAllCooldown: 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. base := h.upstreamBaseURL(acct) // Build the upstream URI from the routed path + original query string. uri := upstreamPath if r.URL.RawQuery != "" { uri += "?" + r.URL.RawQuery } upstreamURL := base + uri var bodyReader io.Reader if body != nil { body.reset() bodyReader = body } else if r.Body != nil { bodyReader = r.Body } req, err := http.NewRequestWithContext(ctx, r.Method, upstreamURL, bodyReader) if err != nil { lastErr = fmt.Errorf("build request: %w", err) continue } copyHeaders(req.Header, r.Header) req.Header.Set("Authorization", "Bearer "+acct.APIKey) // OpenCode Go uses OpenAI-compatible JSON; ensure content-type is set. if req.Header.Get("Content-Type") == "" { req.Header.Set("Content-Type", "application/json") } // Let the URL determine the Host header (do not clear it; clearing can // cause some CDNs to reject the request with 401/403). start := time.Now() resp, err := h.Client.Do(req) latency := time.Since(start) if err != nil { h.Log.Warn("upstream request failed", "account", acct.Name, "account_id", acct.ID, "provider", string(acct.Provider), "path", r.URL.Path, "latency_ms", latency.Milliseconds(), "err", err) lastErr = err continue } // Pre-stream decision: retriable errors → cooldown + failover. if isRetriable(resp.StatusCode) { bodyPreview := readPreview(resp.Body, 512) resp.Body.Close() h.Log.Warn("upstream rejected, failover", "account", acct.Name, "account_id", acct.ID, "provider", string(acct.Provider), "path", r.URL.Path, "status", resp.StatusCode, "latency_ms", latency.Milliseconds(), "preview", bodyPreview) if shouldCooldown(resp.StatusCode, bodyPreview) { 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, "provider", string(acct.Provider), "path", r.URL.Path, "status", resp.StatusCode, "latency_ms", latency.Milliseconds(), "stream", isStreaming(resp)) return } // Exhausted retries. if h.Balancer.Len() > 0 { allCooldown := true status := h.Balancer.Status() for _, a := range h.Balancer.Accounts() { if eligible != nil && len(eligible) > 0 { matched := false for _, p := range eligible { if a.Provider == p { matched = true break } } if !matched { continue } } 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. func (h *Handler) copyResponse(w http.ResponseWriter, resp *http.Response) { defer resp.Body.Close() 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 { 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 the client may have sent. func copyHeaders(dst, src http.Header) { hopByHop := []string{ "Connection", "Keep-Alive", "Proxy-Authenticate", "Proxy-Authorization", "Te", "Trailers", "Transfer-Encoding", "Upgrade", "Authorization", } 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) } } } func readPreview(r io.ReadCloser, n int) string { defer r.Close() buf := make([]byte, n) m, _ := r.Read(buf) return string(buf[:m]) } func isStreaming(resp *http.Response) bool { ct := resp.Header.Get("Content-Type") if strings.Contains(ct, "text/event-stream") || strings.Contains(ct, "application/x-ndjson") { return true } te := resp.Header.Get("Transfer-Encoding") return strings.Contains(strings.ToLower(te), "chunked") } 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) }