227 lines
6.9 KiB
Go
227 lines
6.9 KiB
Go
package proxy
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"net/http"
|
|
"net/http/httputil"
|
|
"net/url"
|
|
"strings"
|
|
"time"
|
|
|
|
"git.wntrmute.dev/kyle/crossbar/internal/store"
|
|
)
|
|
|
|
// forward builds the reverse proxy for one host, tees the response, records the accounting row, and
|
|
// logs. leaseState is "new" or "reused"; waited is the time spent in the queue. ctxEst is the
|
|
// prompt size the context guard estimated (0 when the guard did not run); ctxHeader is the
|
|
// "moved:…<host>" header to set when the guard relocated the conversation.
|
|
func (p *Handler) forward(w http.ResponseWriter, r *http.Request, route, host, leaseState, rest, fp, model string, started time.Time, waited time.Duration, ctxEst int, ctxHeader string) {
|
|
hostCfg, ok := p.cfg.Hosts[host]
|
|
if !ok {
|
|
p.writeError(w, http.StatusBadGateway, "upstream failed")
|
|
return
|
|
}
|
|
target, err := url.Parse(hostCfg.BaseURL)
|
|
if err != nil {
|
|
p.writeError(w, http.StatusBadGateway, "upstream failed")
|
|
return
|
|
}
|
|
|
|
rev := &forwardState{started: started}
|
|
rp := newReverseProxy(p.health, host, leaseState, ctxHeader, target, rest, rev)
|
|
rec := &statusRecorder{ResponseWriter: w, status: http.StatusOK}
|
|
|
|
// ServeHTTP unwinds with http.ErrAbortHandler when a client leaves mid-stream; recover so the
|
|
// row the request earned is still written, then re-panic so the server keeps its semantics.
|
|
defer func() {
|
|
if pv := recover(); pv != nil {
|
|
total := time.Since(started).Milliseconds()
|
|
req := forwardRow(route, fp, model, host, started, waited, rev, rec.status, total)
|
|
if perr, ok := pv.(error); ok && errors.Is(perr, http.ErrAbortHandler) {
|
|
req.Status = 499
|
|
req.Err = "client cancelled"
|
|
} else {
|
|
req.Err = "upstream error"
|
|
}
|
|
p.writeRecord(req)
|
|
panic(pv)
|
|
}
|
|
}()
|
|
|
|
rp.ServeHTTP(rec, r)
|
|
total := time.Since(started)
|
|
|
|
req := forwardRow(route, fp, model, host, started, waited, rev, rec.status, total.Milliseconds())
|
|
if r.Context().Err() != nil {
|
|
req.Status = 499
|
|
req.Err = "client cancelled"
|
|
}
|
|
p.writeRecord(req)
|
|
|
|
fp8 := fp
|
|
if len(fp8) > 8 {
|
|
fp8 = fp8[:8]
|
|
}
|
|
p.log.Info("request",
|
|
"route", route,
|
|
"host", host,
|
|
"method", r.Method,
|
|
"path", rest,
|
|
"status", rec.status,
|
|
"lease", leaseState,
|
|
"queued_ms", waited.Milliseconds(),
|
|
"fp", fp8,
|
|
"ctx_est", ctxEst,
|
|
"ms", total.Milliseconds(),
|
|
)
|
|
}
|
|
|
|
// forwardRow builds the accounting row from the state a forward gathered: what the tee scanned and
|
|
// what the recorder captured. totalMs is measured from start to the caller's exit, so the forward
|
|
// path and the recovery path above build identical rows.
|
|
func forwardRow(route, fp, model, host string, start time.Time, waited time.Duration, rev *forwardState, status int, totalMs int64) store.Request {
|
|
req := store.Request{
|
|
Route: route,
|
|
FP: fp,
|
|
Model: model,
|
|
Host: host,
|
|
Started: start,
|
|
QueuedMs: waited.Milliseconds(),
|
|
TTFBMs: ttfbMs(rev),
|
|
TotalMs: totalMs,
|
|
Status: status,
|
|
Streamed: rev.streamed,
|
|
}
|
|
if rev.tee != nil {
|
|
prompt, cached, completion := rev.tee.tokens()
|
|
req.PromptTokens = int64(prompt)
|
|
req.CachedTokens = int64(cached)
|
|
req.CompletionTokens = int64(completion)
|
|
}
|
|
return req
|
|
}
|
|
|
|
// leaseState is "reused" when the lease already held the conversation, else "new".
|
|
func leaseState(reused bool) string {
|
|
if reused {
|
|
return "reused"
|
|
}
|
|
return "new"
|
|
}
|
|
|
|
// ttfbMs is the time from request start to the response head; zero when the head never arrived.
|
|
func ttfbMs(rev *forwardState) int64 {
|
|
if rev.ttfb.IsZero() || rev.ttfb.Before(rev.started) {
|
|
return 0
|
|
}
|
|
return rev.ttfb.Sub(rev.started).Milliseconds()
|
|
}
|
|
|
|
// forwardState carries, across one forward, when the request started, when the head arrived, whether
|
|
// the response streamed, and the tee that scanned it.
|
|
type forwardState struct {
|
|
started time.Time
|
|
ttfb time.Time
|
|
streamed bool
|
|
tee *tee
|
|
}
|
|
|
|
// statusRecorder records the status written and forwards Flush so the reverse proxy can stream.
|
|
type statusRecorder struct {
|
|
http.ResponseWriter
|
|
status int
|
|
}
|
|
|
|
func (r *statusRecorder) WriteHeader(code int) {
|
|
r.status = code
|
|
r.ResponseWriter.WriteHeader(code)
|
|
}
|
|
|
|
func (r *statusRecorder) Flush() {
|
|
if f, ok := r.ResponseWriter.(http.Flusher); ok {
|
|
f.Flush()
|
|
}
|
|
}
|
|
|
|
// writeError answers with a JSON {"error":"…"} body.
|
|
func (p *Handler) writeError(w http.ResponseWriter, status int, msg string) {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(status)
|
|
_ = json.NewEncoder(w).Encode(map[string]string{"error": msg})
|
|
}
|
|
|
|
// writeNoHealthyHost answers the 503 when no candidate woke. The body names the
|
|
// hosts that were asked to wake (an empty list, never null), and the row
|
|
// records the miss.
|
|
func (p *Handler) writeNoHealthyHost(w http.ResponseWriter, route, model, fp string, started time.Time, tried []string) {
|
|
if tried == nil {
|
|
tried = []string{}
|
|
}
|
|
p.writeRecord(store.Request{
|
|
Route: route,
|
|
FP: fp,
|
|
Model: model,
|
|
Started: started,
|
|
TotalMs: time.Since(started).Milliseconds(),
|
|
Status: http.StatusServiceUnavailable,
|
|
Err: "no healthy host",
|
|
})
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(http.StatusServiceUnavailable)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": "no healthy host",
|
|
"woke": tried,
|
|
})
|
|
}
|
|
|
|
// writeRecord writes one accounting row, logging (never returning) a recorder error.
|
|
func (p *Handler) writeRecord(req store.Request) {
|
|
if p.rec == nil {
|
|
return
|
|
}
|
|
if err := p.rec.RecordRequest(req); err != nil {
|
|
p.log.Error("record request", "err", err)
|
|
}
|
|
}
|
|
|
|
// newReverseProxy forwards to a single host, rewriting the path to target.Path+rest and keeping the
|
|
// original query string. It flushes after every write so long server-sent-event streams are not
|
|
// buffered, tees the response for usage/timings, and marks the host down on any transport error
|
|
// other than a client disconnect.
|
|
func newReverseProxy(h Health, host, leaseState, ctxHeader string, target *url.URL, rest string, rev *forwardState) *httputil.ReverseProxy {
|
|
return &httputil.ReverseProxy{
|
|
Rewrite: func(pr *httputil.ProxyRequest) {
|
|
pr.SetURL(target)
|
|
pr.Out.URL.Path = target.Path + rest
|
|
pr.Out.URL.RawPath = ""
|
|
pr.Out.Host = target.Host
|
|
pr.SetXForwarded()
|
|
},
|
|
FlushInterval: -1,
|
|
ModifyResponse: func(resp *http.Response) error {
|
|
resp.Header.Set(HostHeader, host)
|
|
resp.Header.Set(LeaseHeader, leaseState)
|
|
if ctxHeader != "" {
|
|
resp.Header.Set(CtxHeader, ctxHeader)
|
|
}
|
|
rev.ttfb = time.Now()
|
|
rev.streamed = strings.HasPrefix(resp.Header.Get("Content-Type"), "text/event-stream")
|
|
t := newTee(resp.Body, rev.streamed)
|
|
resp.Body = t
|
|
rev.tee = t
|
|
return nil
|
|
},
|
|
ErrorHandler: func(w http.ResponseWriter, req *http.Request, err error) {
|
|
if errors.Is(err, context.Canceled) {
|
|
return
|
|
}
|
|
h.MarkDown(host, err.Error())
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(http.StatusBadGateway)
|
|
_ = json.NewEncoder(w).Encode(map[string]string{"error": "upstream failed", "host": host})
|
|
},
|
|
}
|
|
}
|