558 lines
16 KiB
Go
558 lines
16 KiB
Go
// Dynagate LLM request proxying handlers
|
|
|
|
package main
|
|
|
|
import (
|
|
"bytes"
|
|
"encoding/json"
|
|
"fmt"
|
|
"io"
|
|
"log"
|
|
"net/http"
|
|
"strings"
|
|
)
|
|
|
|
var httpClient = &http.Client{}
|
|
|
|
func checkAuth(expectedTokens []string, r *http.Request) bool {
|
|
if expectedTokens == nil {
|
|
return true
|
|
}
|
|
authHeader := r.Header.Get("Authorization")
|
|
if !strings.HasPrefix(authHeader, "Bearer ") {
|
|
return false
|
|
}
|
|
token := strings.TrimPrefix(authHeader, "Bearer ")
|
|
for _, expected := range expectedTokens {
|
|
if token == expected {
|
|
return true
|
|
}
|
|
}
|
|
return false
|
|
}
|
|
|
|
func sendUnauthorized(w http.ResponseWriter) {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(http.StatusUnauthorized)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": map[string]any{
|
|
"message": "Incorrect API key provided.",
|
|
"type": "invalid_request_error",
|
|
"param": nil,
|
|
"code": "invalid_api_key",
|
|
},
|
|
})
|
|
}
|
|
|
|
func handleModels(cm *ConfigManager, expectedTokens []string) http.HandlerFunc {
|
|
type ModelData struct {
|
|
ID string `json:"id"`
|
|
Object string `json:"object"`
|
|
Created int64 `json:"created"`
|
|
OwnedBy string `json:"owned_by"`
|
|
}
|
|
|
|
type ModelsResponse struct {
|
|
Object string `json:"object"`
|
|
Data []ModelData `json:"data"`
|
|
}
|
|
|
|
return func(w http.ResponseWriter, r *http.Request) {
|
|
if !checkAuth(expectedTokens, r) {
|
|
sendUnauthorized(w)
|
|
return
|
|
}
|
|
|
|
if r.Method != http.MethodGet {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(http.StatusMethodNotAllowed)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": map[string]any{
|
|
"message": "Method not allowed",
|
|
"type": "invalid_request_error",
|
|
},
|
|
})
|
|
return
|
|
}
|
|
|
|
models := cm.GetUniqueModels()
|
|
data := make([]ModelData, len(models))
|
|
for i, m := range models {
|
|
data[i] = ModelData{
|
|
ID: m,
|
|
Object: "model",
|
|
Created: 1686935002,
|
|
OwnedBy: "dynagate",
|
|
}
|
|
}
|
|
|
|
resp := ModelsResponse{
|
|
Object: "list",
|
|
Data: data,
|
|
}
|
|
|
|
w.Header().Set("Content-Type", "application/json")
|
|
_ = json.NewEncoder(w).Encode(resp)
|
|
}
|
|
}
|
|
|
|
func handleChatCompletions(cm *ConfigManager, expectedTokens []string) http.HandlerFunc {
|
|
return func(w http.ResponseWriter, r *http.Request) {
|
|
if !checkAuth(expectedTokens, r) {
|
|
sendUnauthorized(w)
|
|
return
|
|
}
|
|
|
|
if r.Method != http.MethodPost {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(http.StatusMethodNotAllowed)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": map[string]any{
|
|
"message": "Method not allowed",
|
|
"type": "invalid_request_error",
|
|
},
|
|
})
|
|
return
|
|
}
|
|
|
|
bodyBytes, err := io.ReadAll(r.Body)
|
|
if err != nil {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(http.StatusBadRequest)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": map[string]any{
|
|
"message": "Failed to read request body",
|
|
"type": "invalid_request_error",
|
|
},
|
|
})
|
|
return
|
|
}
|
|
|
|
var bodyMap map[string]any
|
|
if err := json.Unmarshal(bodyBytes, &bodyMap); err != nil {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(http.StatusBadRequest)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": map[string]any{
|
|
"message": "Invalid JSON in request body",
|
|
"type": "invalid_request_error",
|
|
},
|
|
})
|
|
return
|
|
}
|
|
|
|
var requestedModel string
|
|
if m, ok := bodyMap["model"]; ok {
|
|
if s, ok := m.(string); ok {
|
|
requestedModel = s
|
|
}
|
|
}
|
|
|
|
configs := cm.GetConfigs()
|
|
uniqueModels := cm.GetUniqueModels()
|
|
|
|
if len(configs) == 0 {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(http.StatusServiceUnavailable)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": map[string]any{
|
|
"message": "No model configurations loaded",
|
|
"type": "gateway_error",
|
|
},
|
|
})
|
|
return
|
|
}
|
|
|
|
trialConfigs := getTrialConfigs(configs, uniqueModels, requestedModel)
|
|
if len(trialConfigs) == 0 {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(http.StatusServiceUnavailable)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": map[string]any{
|
|
"message": "No valid trial configuration candidates",
|
|
"type": "gateway_error",
|
|
},
|
|
})
|
|
return
|
|
}
|
|
|
|
var isStream bool
|
|
if s, ok := bodyMap["stream"]; ok {
|
|
if b, ok := s.(bool); ok {
|
|
isStream = b
|
|
}
|
|
}
|
|
|
|
log.Printf("Received completion request for model %q (stream=%t). Found %d config trials.", requestedModel, isStream, len(trialConfigs))
|
|
|
|
for i, trial := range trialConfigs {
|
|
log.Printf("Trial %d/%d: model=%s endpoint=%s key_len=%d", i+1, len(trialConfigs), trial.Model, trial.Endpoint, len(trial.Key))
|
|
|
|
bodyMap["model"] = trial.Model
|
|
modifiedBody, err := json.Marshal(bodyMap)
|
|
if err != nil {
|
|
log.Printf("Trial %d: Failed to marshal body for %s: %v", i+1, trial.Model, err)
|
|
continue
|
|
}
|
|
|
|
targetURL := buildURL(trial.Endpoint, r.URL.Path)
|
|
|
|
outReq, err := http.NewRequestWithContext(r.Context(), "POST", targetURL, bytes.NewReader(modifiedBody))
|
|
if err != nil {
|
|
log.Printf("Trial %d: Failed to create outgoing request to %s: %v", i+1, targetURL, err)
|
|
continue
|
|
}
|
|
|
|
// Copy headers from incoming request, excluding host and auth
|
|
for k, vv := range r.Header {
|
|
kLower := strings.ToLower(k)
|
|
if kLower == "authorization" || kLower == "host" || kLower == "content-length" {
|
|
continue
|
|
}
|
|
for _, v := range vv {
|
|
outReq.Header.Add(k, v)
|
|
}
|
|
}
|
|
if trial.Key == "-blank-" {
|
|
outReq.Header.Set("Authorization", "Bearer")
|
|
} else if trial.Key != "" && trial.Key != "-" {
|
|
outReq.Header.Set("Authorization", "Bearer "+trial.Key)
|
|
}
|
|
outReq.Header.Set("Content-Type", "application/json")
|
|
|
|
if trial.Extra != "" {
|
|
var extraHeaders map[string]any
|
|
if err := json.Unmarshal([]byte(trial.Extra), &extraHeaders); err != nil {
|
|
log.Printf("Trial %d: Failed to parse extra headers JSON: %v", i+1, err)
|
|
} else {
|
|
for hk, hv := range extraHeaders {
|
|
var valStr string
|
|
switch v := hv.(type) {
|
|
case string:
|
|
valStr = v
|
|
default:
|
|
valStr = fmt.Sprintf("%v", v)
|
|
}
|
|
outReq.Header.Set(hk, valStr)
|
|
}
|
|
}
|
|
}
|
|
|
|
resp, err := httpClient.Do(outReq)
|
|
if err != nil {
|
|
log.Printf("Trial %d: Request to %s failed: %v", i+1, targetURL, err)
|
|
continue
|
|
}
|
|
|
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
|
errBody, _ := io.ReadAll(io.LimitReader(resp.Body, 1024))
|
|
resp.Body.Close()
|
|
log.Printf("Trial %d: Request to %s returned error status %d: %s", i+1, targetURL, resp.StatusCode, strings.TrimSpace(string(errBody)))
|
|
continue
|
|
}
|
|
|
|
log.Printf("Trial %d: Connection established with status %d. Proxying response.", i+1, resp.StatusCode)
|
|
|
|
if !isStream {
|
|
respBody, readErr := io.ReadAll(resp.Body)
|
|
resp.Body.Close()
|
|
if readErr != nil {
|
|
log.Printf("Trial %d: Failed to read non-streaming response body: %v", i+1, readErr)
|
|
continue
|
|
}
|
|
|
|
for k, vv := range resp.Header {
|
|
for _, v := range vv {
|
|
w.Header().Add(k, v)
|
|
}
|
|
}
|
|
w.WriteHeader(resp.StatusCode)
|
|
_, _ = w.Write(respBody)
|
|
return
|
|
}
|
|
|
|
// Handle streaming
|
|
defer resp.Body.Close()
|
|
flusher, ok := w.(http.Flusher)
|
|
if !ok {
|
|
log.Printf("Trial %d: Flusher not supported on current ResponseWriter", i+1)
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(http.StatusInternalServerError)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": map[string]any{
|
|
"message": "Response flusher not supported",
|
|
"type": "gateway_error",
|
|
},
|
|
})
|
|
return
|
|
}
|
|
|
|
for k, vv := range resp.Header {
|
|
for _, v := range vv {
|
|
w.Header().Add(k, v)
|
|
}
|
|
}
|
|
w.WriteHeader(resp.StatusCode)
|
|
flusher.Flush()
|
|
|
|
buf := make([]byte, 4096)
|
|
for {
|
|
n, rErr := resp.Body.Read(buf)
|
|
if n > 0 {
|
|
if _, wErr := w.Write(buf[:n]); wErr != nil {
|
|
log.Printf("Trial %d: Client disconnected or write error: %v", i+1, wErr)
|
|
return
|
|
}
|
|
flusher.Flush()
|
|
}
|
|
if rErr != nil {
|
|
if rErr != io.EOF {
|
|
log.Printf("Trial %d: Error reading response stream: %v", i+1, rErr)
|
|
}
|
|
break
|
|
}
|
|
}
|
|
return
|
|
}
|
|
|
|
log.Printf("All trials failed. Returning Bad Gateway.")
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(http.StatusBadGateway)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": map[string]any{
|
|
"message": "All configured models and keys failed to respond.",
|
|
"type": "gateway_error",
|
|
"param": nil,
|
|
"code": "all_endpoints_failed",
|
|
},
|
|
})
|
|
}
|
|
}
|
|
|
|
func handleImageGenerations(cm *ConfigManager, expectedTokens []string) http.HandlerFunc {
|
|
return func(w http.ResponseWriter, r *http.Request) {
|
|
if !checkAuth(expectedTokens, r) {
|
|
sendUnauthorized(w)
|
|
return
|
|
}
|
|
|
|
if r.Method != http.MethodPost {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(http.StatusMethodNotAllowed)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": map[string]any{
|
|
"message": "Method not allowed",
|
|
"type": "invalid_request_error",
|
|
},
|
|
})
|
|
return
|
|
}
|
|
|
|
bodyBytes, err := io.ReadAll(r.Body)
|
|
if err != nil {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(http.StatusBadRequest)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": map[string]any{
|
|
"message": "Failed to read request body",
|
|
"type": "invalid_request_error",
|
|
},
|
|
})
|
|
return
|
|
}
|
|
|
|
var bodyMap map[string]any
|
|
if err := json.Unmarshal(bodyBytes, &bodyMap); err != nil {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(http.StatusBadRequest)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": map[string]any{
|
|
"message": "Invalid JSON in request body",
|
|
"type": "invalid_request_error",
|
|
},
|
|
})
|
|
return
|
|
}
|
|
|
|
var requestedModel string
|
|
if m, ok := bodyMap["model"]; ok {
|
|
if s, ok := m.(string); ok {
|
|
requestedModel = s
|
|
}
|
|
}
|
|
|
|
configs := cm.GetConfigs()
|
|
uniqueModels := cm.GetUniqueModels()
|
|
|
|
if len(configs) == 0 {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(http.StatusServiceUnavailable)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": map[string]any{
|
|
"message": "No model configurations loaded",
|
|
"type": "gateway_error",
|
|
},
|
|
})
|
|
return
|
|
}
|
|
|
|
trialConfigs := getTrialConfigs(configs, uniqueModels, requestedModel)
|
|
if len(trialConfigs) == 0 {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(http.StatusServiceUnavailable)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": map[string]any{
|
|
"message": "No valid trial configuration candidates",
|
|
"type": "gateway_error",
|
|
},
|
|
})
|
|
return
|
|
}
|
|
|
|
log.Printf("Received image generation request for model %q. Found %d config trials.", requestedModel, len(trialConfigs))
|
|
|
|
for i, trial := range trialConfigs {
|
|
log.Printf("Trial %d/%d: model=%s endpoint=%s key_len=%d", i+1, len(trialConfigs), trial.Model, trial.Endpoint, len(trial.Key))
|
|
|
|
bodyMap["model"] = trial.Model
|
|
modifiedBody, err := json.Marshal(bodyMap)
|
|
if err != nil {
|
|
log.Printf("Trial %d: Failed to marshal body for %s: %v", i+1, trial.Model, err)
|
|
continue
|
|
}
|
|
|
|
targetURL := buildURL(trial.Endpoint, r.URL.Path)
|
|
|
|
outReq, err := http.NewRequestWithContext(r.Context(), "POST", targetURL, bytes.NewReader(modifiedBody))
|
|
if err != nil {
|
|
log.Printf("Trial %d: Failed to create outgoing request to %s: %v", i+1, targetURL, err)
|
|
continue
|
|
}
|
|
|
|
for k, vv := range r.Header {
|
|
kLower := strings.ToLower(k)
|
|
if kLower == "authorization" || kLower == "host" || kLower == "content-length" {
|
|
continue
|
|
}
|
|
for _, v := range vv {
|
|
outReq.Header.Add(k, v)
|
|
}
|
|
}
|
|
if trial.Key == "-blank-" {
|
|
outReq.Header.Set("Authorization", "Bearer")
|
|
} else if trial.Key != "" && trial.Key != "-" {
|
|
outReq.Header.Set("Authorization", "Bearer "+trial.Key)
|
|
}
|
|
outReq.Header.Set("Content-Type", "application/json")
|
|
|
|
if trial.Extra != "" {
|
|
var extraHeaders map[string]any
|
|
if err := json.Unmarshal([]byte(trial.Extra), &extraHeaders); err != nil {
|
|
log.Printf("Trial %d: Failed to parse extra headers JSON: %v", i+1, err)
|
|
} else {
|
|
for hk, hv := range extraHeaders {
|
|
var valStr string
|
|
switch v := hv.(type) {
|
|
case string:
|
|
valStr = v
|
|
default:
|
|
valStr = fmt.Sprintf("%v", v)
|
|
}
|
|
outReq.Header.Set(hk, valStr)
|
|
}
|
|
}
|
|
}
|
|
|
|
resp, err := httpClient.Do(outReq)
|
|
if err != nil {
|
|
log.Printf("Trial %d: Request to %s failed: %v", i+1, targetURL, err)
|
|
continue
|
|
}
|
|
|
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
|
errBody, _ := io.ReadAll(io.LimitReader(resp.Body, 1024))
|
|
resp.Body.Close()
|
|
log.Printf("Trial %d: Request to %s returned error status %d: %s", i+1, targetURL, resp.StatusCode, strings.TrimSpace(string(errBody)))
|
|
continue
|
|
}
|
|
|
|
log.Printf("Trial %d: Connection established with status %d. Proxying response.", i+1, resp.StatusCode)
|
|
|
|
respBody, readErr := io.ReadAll(resp.Body)
|
|
resp.Body.Close()
|
|
if readErr != nil {
|
|
log.Printf("Trial %d: Failed to read response body: %v", i+1, readErr)
|
|
continue
|
|
}
|
|
|
|
for k, vv := range resp.Header {
|
|
for _, v := range vv {
|
|
w.Header().Add(k, v)
|
|
}
|
|
}
|
|
w.WriteHeader(resp.StatusCode)
|
|
_, _ = w.Write(respBody)
|
|
return
|
|
}
|
|
|
|
log.Printf("All trials failed. Returning Bad Gateway.")
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(http.StatusBadGateway)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": map[string]any{
|
|
"message": "All configured models and keys failed to respond.",
|
|
"type": "gateway_error",
|
|
"param": nil,
|
|
"code": "all_endpoints_failed",
|
|
},
|
|
})
|
|
}
|
|
}
|
|
|
|
|
|
func getTrialConfigs(configs []ModelConfig, uniqueModels []string, requestedModel string) []ModelConfig {
|
|
if len(configs) == 0 {
|
|
return nil
|
|
}
|
|
|
|
reqIdx := -1
|
|
for i, m := range uniqueModels {
|
|
if m == requestedModel {
|
|
reqIdx = i
|
|
break
|
|
}
|
|
}
|
|
|
|
var orderedModels []string
|
|
if reqIdx != -1 {
|
|
orderedModels = append(orderedModels, uniqueModels[reqIdx:]...)
|
|
orderedModels = append(orderedModels, uniqueModels[:reqIdx]...)
|
|
} else {
|
|
orderedModels = uniqueModels
|
|
}
|
|
|
|
var trialConfigs []ModelConfig
|
|
for _, modelName := range orderedModels {
|
|
for _, cfg := range configs {
|
|
if cfg.Model == modelName {
|
|
trialConfigs = append(trialConfigs, cfg)
|
|
}
|
|
}
|
|
}
|
|
|
|
return trialConfigs
|
|
}
|
|
|
|
func buildURL(endpoint string, reqPath string) string {
|
|
endpoint = strings.TrimSuffix(endpoint, "/")
|
|
reqPath = strings.TrimPrefix(reqPath, "/")
|
|
|
|
if strings.HasSuffix(endpoint, "/v1") {
|
|
if strings.HasPrefix(reqPath, "v1/") {
|
|
reqPath = strings.TrimPrefix(reqPath, "v1/")
|
|
}
|
|
}
|
|
|
|
return endpoint + "/" + reqPath
|
|
}
|