mirror of
https://github.com/Warky-Devs/vecna.git
synced 2026-08-07 10:47:38 +00:00
1953a4f4f9
* implement proxy handler for various OpenAI API routes * add error handling for request body and response streaming * introduce new error response format for API compatibility * add tests for recover middleware to handle panics gracefully
183 lines
5.4 KiB
Go
183 lines
5.4 KiB
Go
package server
|
|
|
|
import (
|
|
"encoding/json"
|
|
"io"
|
|
"net/http"
|
|
"strings"
|
|
"time"
|
|
|
|
"github.com/uptrace/bunrouter"
|
|
|
|
"github.com/Warky-Devs/vecna.git/pkg/forwardproxy"
|
|
)
|
|
|
|
// passthroughRoute describes one OpenAI-compatible endpoint forwarded verbatim.
|
|
type passthroughRoute struct {
|
|
method string
|
|
path string // route pattern registered with bunrouter (":model" etc.)
|
|
}
|
|
|
|
// passthroughRoutes lists the OpenAI-compatible endpoints vecna forwards without
|
|
// any body processing, beyond /v1/embeddings (which gets dimension adaptation).
|
|
var passthroughRoutes = []passthroughRoute{
|
|
{http.MethodGet, "/v1/models"},
|
|
{http.MethodGet, "/v1/models/:model"},
|
|
{http.MethodPost, "/v1/chat/completions"},
|
|
{http.MethodPost, "/v1/completions"},
|
|
{http.MethodPost, "/v1/moderations"},
|
|
{http.MethodPost, "/v1/images/generations"},
|
|
{http.MethodPost, "/v1/images/edits"},
|
|
{http.MethodPost, "/v1/images/variations"},
|
|
{http.MethodPost, "/v1/audio/speech"},
|
|
{http.MethodPost, "/v1/audio/transcriptions"},
|
|
{http.MethodPost, "/v1/audio/translations"},
|
|
{http.MethodPost, "/v1/rerank"},
|
|
}
|
|
|
|
// registerPassthroughRoutes wires up both the unmapped and /map/:mapping variants
|
|
// of every entry in passthroughRoutes.
|
|
func registerPassthroughRoutes(authed *bunrouter.Group, h *handler) {
|
|
for _, r := range passthroughRoutes {
|
|
upstreamPath := strings.Replace(r.path, ":model", "{model}", 1)
|
|
authed.Handle(r.method, r.path, h.proxyHandler(r.method, upstreamPath))
|
|
authed.Handle(r.method, "/map/:mapping"+r.path, h.proxyHandlerMapped(r.method, upstreamPath))
|
|
}
|
|
}
|
|
|
|
// proxyHandler returns a handler that forwards method/path verbatim to the
|
|
// target resolved from the request body's "model" field (or forward.default).
|
|
// No vecna-specific processing (adaptation, etc.) is applied to the body.
|
|
func (h *handler) proxyHandler(method, path string) bunrouter.HandlerFunc {
|
|
return func(w http.ResponseWriter, req bunrouter.Request) error {
|
|
return h.doProxy(w, req, method, path, "")
|
|
}
|
|
}
|
|
|
|
// proxyHandlerMapped is the /map/:mapping variant: the target is always the
|
|
// extra_map's forward_target, regardless of the body's "model" field.
|
|
func (h *handler) proxyHandlerMapped(method, path string) bunrouter.HandlerFunc {
|
|
return func(w http.ResponseWriter, req bunrouter.Request) error {
|
|
em, err := h.resolveExtraMap(req.Param("mapping"))
|
|
if err != nil {
|
|
return writeError(w, http.StatusNotFound, err.Error())
|
|
}
|
|
return h.doProxy(w, req, method, path, em.ForwardTarget)
|
|
}
|
|
}
|
|
|
|
func (h *handler) doProxy(w http.ResponseWriter, req bunrouter.Request, method, pathTemplate, targetOverride string) error {
|
|
model := req.Param("model")
|
|
upstreamPath := strings.Replace(pathTemplate, "{model}", model, 1)
|
|
|
|
var body []byte
|
|
if method != http.MethodGet && method != http.MethodDelete {
|
|
b, err := io.ReadAll(req.Body)
|
|
if err != nil {
|
|
return writeError(w, http.StatusBadRequest, "invalid request body")
|
|
}
|
|
body = b
|
|
if model == "" && strings.HasPrefix(req.Header.Get("Content-Type"), "application/json") {
|
|
model = modelPeek(body)
|
|
}
|
|
}
|
|
|
|
fwd, targetName, err := h.resolveForwarderOverride(targetOverride, model)
|
|
if err != nil {
|
|
return writeError(w, http.StatusBadGateway, err.Error())
|
|
}
|
|
|
|
trace := TraceFromContext(req.Context())
|
|
trace.ForwardTarget = targetName
|
|
|
|
if isStreaming(body) {
|
|
return h.streamProxy(w, req, fwd, method, upstreamPath, body, trace)
|
|
}
|
|
|
|
t0 := time.Now()
|
|
res, err := fwd.Do(req.Context(), method, upstreamPath, body, req.Header)
|
|
trace.ForwardDuration = time.Since(t0)
|
|
if err != nil {
|
|
return writeError(w, http.StatusBadGateway, err.Error())
|
|
}
|
|
trace.ForwardURL = res.Target
|
|
|
|
writeTraceHeaders(w, trace)
|
|
writeUpstreamResult(w, res)
|
|
return nil
|
|
}
|
|
|
|
// streamProxy relays a single-attempt streaming response (e.g. SSE chat
|
|
// completions) directly to the client as bytes arrive.
|
|
func (h *handler) streamProxy(w http.ResponseWriter, req bunrouter.Request, fwd *forwardproxy.Router, method, path string, body []byte, trace *RequestTrace) error {
|
|
t0 := time.Now()
|
|
status, header, respBody, err := fwd.DoStream(req.Context(), method, path, body, req.Header)
|
|
trace.ForwardDuration = time.Since(t0)
|
|
if err != nil {
|
|
return writeError(w, http.StatusBadGateway, err.Error())
|
|
}
|
|
defer func() { _ = respBody.Close() }()
|
|
|
|
dst := w.Header()
|
|
for k, vv := range header {
|
|
if forwardproxy.SkipResponseHeader(k) {
|
|
continue
|
|
}
|
|
for _, v := range vv {
|
|
dst.Add(k, v)
|
|
}
|
|
}
|
|
w.WriteHeader(status)
|
|
|
|
flusher, _ := w.(http.Flusher)
|
|
buf := make([]byte, 4096)
|
|
for {
|
|
n, rerr := respBody.Read(buf)
|
|
if n > 0 {
|
|
if _, werr := w.Write(buf[:n]); werr != nil {
|
|
return nil
|
|
}
|
|
if flusher != nil {
|
|
flusher.Flush()
|
|
}
|
|
}
|
|
if rerr != nil {
|
|
break
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func writeUpstreamResult(w http.ResponseWriter, res forwardproxy.Result) {
|
|
dst := w.Header()
|
|
for k, vv := range res.Header {
|
|
if forwardproxy.SkipResponseHeader(k) {
|
|
continue
|
|
}
|
|
for _, v := range vv {
|
|
dst.Add(k, v)
|
|
}
|
|
}
|
|
w.WriteHeader(res.Status)
|
|
_, _ = w.Write(res.Body)
|
|
}
|
|
|
|
// modelPeek extracts the "model" field from a JSON body without full decoding,
|
|
// used to select a forward target the same way embeddings does.
|
|
func modelPeek(body []byte) string {
|
|
var v struct {
|
|
Model string `json:"model"`
|
|
}
|
|
_ = json.Unmarshal(body, &v)
|
|
return v.Model
|
|
}
|
|
|
|
// isStreaming reports whether a JSON request body sets "stream": true.
|
|
func isStreaming(body []byte) bool {
|
|
var v struct {
|
|
Stream bool `json:"stream"`
|
|
}
|
|
_ = json.Unmarshal(body, &v)
|
|
return v.Stream
|
|
}
|