1920 lines
56 KiB
Go
1920 lines
56 KiB
Go
// Dynagate LLM request proxying handlers
|
|
|
|
package main
|
|
|
|
import (
|
|
"bufio"
|
|
"bytes"
|
|
"encoding/json"
|
|
"fmt"
|
|
"io"
|
|
"log"
|
|
"net/http"
|
|
"strings"
|
|
"time"
|
|
)
|
|
|
|
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, retryBaseDelay ...time.Duration) http.HandlerFunc {
|
|
baseDelay := 100 * time.Millisecond
|
|
if len(retryBaseDelay) > 0 {
|
|
baseDelay = retryBaseDelay[0]
|
|
}
|
|
|
|
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
|
|
}
|
|
|
|
autodetectAndNormalizeMessages(bodyMap)
|
|
|
|
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 {
|
|
if i > 0 && baseDelay > 0 {
|
|
delay := time.Duration(fibonacci(i)) * baseDelay
|
|
log.Printf("Trial %d/%d: Fibonacci backoff delay of %v before retry...", i+1, len(trialConfigs), delay)
|
|
select {
|
|
case <-r.Context().Done():
|
|
log.Printf("Request context cancelled during retry delay before trial %d", i+1)
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(499)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": map[string]any{
|
|
"message": "Client closed request",
|
|
"type": "gateway_error",
|
|
},
|
|
})
|
|
return
|
|
case <-time.After(delay):
|
|
}
|
|
}
|
|
|
|
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, retryBaseDelay ...time.Duration) http.HandlerFunc {
|
|
baseDelay := 100 * time.Millisecond
|
|
if len(retryBaseDelay) > 0 {
|
|
baseDelay = retryBaseDelay[0]
|
|
}
|
|
|
|
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
|
|
}
|
|
|
|
if p, ok := bodyMap["prompt"]; ok && p != nil {
|
|
bodyMap["prompt"] = normalizeContent(p)
|
|
}
|
|
|
|
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 {
|
|
if i > 0 && baseDelay > 0 {
|
|
delay := time.Duration(fibonacci(i)) * baseDelay
|
|
log.Printf("Trial %d/%d: Fibonacci backoff delay of %v before retry...", i+1, len(trialConfigs), delay)
|
|
select {
|
|
case <-r.Context().Done():
|
|
log.Printf("Request context cancelled during retry delay before trial %d", i+1)
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(499)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": map[string]any{
|
|
"message": "Client closed request",
|
|
"type": "gateway_error",
|
|
},
|
|
})
|
|
return
|
|
case <-time.After(delay):
|
|
}
|
|
}
|
|
|
|
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
|
|
}
|
|
|
|
func isMediaContent(m map[string]any) bool {
|
|
if t, ok := m["type"].(string); ok {
|
|
tLower := strings.ToLower(t)
|
|
switch tLower {
|
|
case "image_url", "image", "input_audio", "audio", "file", "document", "video":
|
|
return true
|
|
}
|
|
}
|
|
for k := range m {
|
|
kLower := strings.ToLower(k)
|
|
switch kLower {
|
|
case "image_url", "input_audio", "inline_data", "file_data":
|
|
return true
|
|
}
|
|
}
|
|
return false
|
|
}
|
|
|
|
func normalizeContent(contentAny any) any {
|
|
if contentAny == nil {
|
|
return nil
|
|
}
|
|
|
|
switch v := contentAny.(type) {
|
|
case string:
|
|
return v
|
|
case []any:
|
|
if len(v) == 0 {
|
|
return ""
|
|
}
|
|
hasMedia := false
|
|
for _, item := range v {
|
|
if m, ok := item.(map[string]any); ok {
|
|
if isMediaContent(m) {
|
|
hasMedia = true
|
|
break
|
|
}
|
|
}
|
|
}
|
|
|
|
if hasMedia {
|
|
var normSlice []any
|
|
for _, item := range v {
|
|
if m, ok := item.(map[string]any); ok {
|
|
if t, ok := m["type"].(string); ok && strings.ToLower(t) == "text" {
|
|
textVal, _ := m["text"].(string)
|
|
normSlice = append(normSlice, map[string]any{
|
|
"type": "text",
|
|
"text": textVal,
|
|
})
|
|
} else {
|
|
normSlice = append(normSlice, m)
|
|
}
|
|
} else {
|
|
normSlice = append(normSlice, item)
|
|
}
|
|
}
|
|
return normSlice
|
|
}
|
|
|
|
var textParts []string
|
|
for _, item := range v {
|
|
switch elem := item.(type) {
|
|
case string:
|
|
textParts = append(textParts, elem)
|
|
case map[string]any:
|
|
if txt, ok := elem["text"].(string); ok {
|
|
textParts = append(textParts, txt)
|
|
} else if txt, ok := elem["content"].(string); ok {
|
|
textParts = append(textParts, txt)
|
|
}
|
|
}
|
|
}
|
|
|
|
var sb strings.Builder
|
|
for _, part := range textParts {
|
|
if part == "" {
|
|
continue
|
|
}
|
|
if sb.Len() > 0 {
|
|
lastChar := sb.String()[sb.Len()-1]
|
|
firstChar := part[0]
|
|
if lastChar != '\n' && lastChar != ' ' && firstChar != '\n' && firstChar != ' ' {
|
|
sb.WriteString("\n")
|
|
}
|
|
}
|
|
sb.WriteString(part)
|
|
}
|
|
return sb.String()
|
|
|
|
case map[string]any:
|
|
if isMediaContent(v) {
|
|
return []any{v}
|
|
}
|
|
if txt, ok := v["text"].(string); ok {
|
|
return txt
|
|
}
|
|
if txt, ok := v["content"].(string); ok {
|
|
return txt
|
|
}
|
|
return v
|
|
default:
|
|
return contentAny
|
|
}
|
|
}
|
|
|
|
func autodetectAndNormalizeMessages(bodyMap map[string]any) {
|
|
if bodyMap == nil {
|
|
return
|
|
}
|
|
|
|
var systemMsg map[string]any
|
|
if sysVal, ok := bodyMap["system"]; ok && sysVal != nil {
|
|
sysText := normalizeContent(sysVal)
|
|
if sysStr, isStr := sysText.(string); isStr && sysStr != "" {
|
|
systemMsg = map[string]any{
|
|
"role": "system",
|
|
"content": sysStr,
|
|
}
|
|
} else if sysSlice, isSlice := sysText.([]any); isSlice && len(sysSlice) > 0 {
|
|
systemMsg = map[string]any{
|
|
"role": "system",
|
|
"content": sysSlice,
|
|
}
|
|
}
|
|
delete(bodyMap, "system")
|
|
}
|
|
|
|
var rawMessages any
|
|
var sourceKey string
|
|
|
|
if msgs, ok := bodyMap["messages"]; ok && msgs != nil {
|
|
rawMessages = msgs
|
|
sourceKey = "messages"
|
|
} else if prompt, ok := bodyMap["prompt"]; ok && prompt != nil {
|
|
rawMessages = prompt
|
|
sourceKey = "prompt"
|
|
delete(bodyMap, "prompt")
|
|
} else if contents, ok := bodyMap["contents"]; ok && contents != nil {
|
|
rawMessages = contents
|
|
sourceKey = "contents"
|
|
delete(bodyMap, "contents")
|
|
} else if input, ok := bodyMap["input"]; ok && input != nil {
|
|
rawMessages = input
|
|
sourceKey = "input"
|
|
delete(bodyMap, "input")
|
|
}
|
|
|
|
if rawMessages == nil && systemMsg == nil {
|
|
return
|
|
}
|
|
|
|
var msgList []map[string]any
|
|
|
|
switch v := rawMessages.(type) {
|
|
case []any:
|
|
for _, item := range v {
|
|
switch elem := item.(type) {
|
|
case map[string]any:
|
|
msgList = append(msgList, elem)
|
|
case string:
|
|
msgList = append(msgList, map[string]any{
|
|
"role": "user",
|
|
"content": elem,
|
|
})
|
|
}
|
|
}
|
|
case map[string]any:
|
|
msgList = append(msgList, v)
|
|
case string:
|
|
if v != "" {
|
|
msgList = append(msgList, map[string]any{
|
|
"role": "user",
|
|
"content": v,
|
|
})
|
|
}
|
|
}
|
|
|
|
var normalizedList []any
|
|
hasSystemInList := false
|
|
|
|
for _, msg := range msgList {
|
|
role, _ := msg["role"].(string)
|
|
if role == "" {
|
|
role = "user"
|
|
}
|
|
if role == "system" {
|
|
hasSystemInList = true
|
|
}
|
|
|
|
normMsg := make(map[string]any)
|
|
for k, val := range msg {
|
|
normMsg[k] = val
|
|
}
|
|
normMsg["role"] = role
|
|
|
|
if cnt, ok := msg["content"]; ok {
|
|
normMsg["content"] = normalizeContent(cnt)
|
|
} else if parts, ok := msg["parts"]; ok {
|
|
normMsg["content"] = normalizeContent(parts)
|
|
delete(normMsg, "parts")
|
|
}
|
|
|
|
normalizedList = append(normalizedList, normMsg)
|
|
}
|
|
|
|
if systemMsg != nil {
|
|
if !hasSystemInList {
|
|
normalizedList = append([]any{systemMsg}, normalizedList...)
|
|
} else if len(normalizedList) > 0 {
|
|
if firstMsg, ok := normalizedList[0].(map[string]any); ok && firstMsg["role"] == "system" {
|
|
existingSys := normalizeContent(firstMsg["content"])
|
|
if sysStr, ok := systemMsg["content"].(string); ok {
|
|
if exStr, ok := existingSys.(string); ok && exStr != "" {
|
|
firstMsg["content"] = sysStr + "\n" + exStr
|
|
} else {
|
|
firstMsg["content"] = sysStr
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
if len(normalizedList) > 0 || sourceKey != "" {
|
|
bodyMap["messages"] = normalizedList
|
|
}
|
|
}
|
|
|
|
func fibonacci(n int) int64 {
|
|
if n <= 0 {
|
|
return 0
|
|
}
|
|
if n == 1 || n == 2 {
|
|
return 1
|
|
}
|
|
var a, b int64 = 1, 1
|
|
for i := 3; i <= n; i++ {
|
|
a, b = b, a+b
|
|
}
|
|
return b
|
|
}
|
|
|
|
func handleMessages(cm *ConfigManager, expectedTokens []string, retryBaseDelay ...time.Duration) http.HandlerFunc {
|
|
baseDelay := 100 * time.Millisecond
|
|
if len(retryBaseDelay) > 0 {
|
|
baseDelay = retryBaseDelay[0]
|
|
}
|
|
|
|
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
|
|
}
|
|
|
|
openAIBodyMap := convertAnthropicToOpenAI(bodyMap)
|
|
|
|
var requestedModel string
|
|
if m, ok := openAIBodyMap["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 := openAIBodyMap["stream"]; ok {
|
|
if b, ok := s.(bool); ok {
|
|
isStream = b
|
|
}
|
|
}
|
|
|
|
log.Printf("Received Anthropic messages request for model %q (stream=%t). Found %d config trials.", requestedModel, isStream, len(trialConfigs))
|
|
|
|
for i, trial := range trialConfigs {
|
|
if i > 0 && baseDelay > 0 {
|
|
delay := time.Duration(fibonacci(i)) * baseDelay
|
|
log.Printf("Trial %d/%d: Fibonacci backoff delay of %v before retry...", i+1, len(trialConfigs), delay)
|
|
select {
|
|
case <-r.Context().Done():
|
|
log.Printf("Request context cancelled during retry delay before trial %d", i+1)
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(499)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": map[string]any{
|
|
"message": "Client closed request",
|
|
"type": "gateway_error",
|
|
},
|
|
})
|
|
return
|
|
case <-time.After(delay):
|
|
}
|
|
}
|
|
|
|
log.Printf("Trial %d/%d: model=%s endpoint=%s key_len=%d", i+1, len(trialConfigs), trial.Model, trial.Endpoint, len(trial.Key))
|
|
|
|
openAIBodyMap["model"] = trial.Model
|
|
modifiedBody, err := json.Marshal(openAIBodyMap)
|
|
if err != nil {
|
|
log.Printf("Trial %d: Failed to marshal body for %s: %v", i+1, trial.Model, err)
|
|
continue
|
|
}
|
|
|
|
targetURL := buildURL(trial.Endpoint, "/v1/chat/completions")
|
|
|
|
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)
|
|
|
|
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
|
|
}
|
|
|
|
var openAIResp map[string]any
|
|
if err := json.Unmarshal(respBody, &openAIResp); err != nil {
|
|
log.Printf("Trial %d: Failed to unmarshal upstream response JSON: %v", i+1, err)
|
|
continue
|
|
}
|
|
|
|
anthropicResp := convertOpenAIToAnthropicResponse(openAIResp, requestedModel)
|
|
|
|
for k, vv := range resp.Header {
|
|
kLower := strings.ToLower(k)
|
|
if kLower == "content-length" || kLower == "content-type" {
|
|
continue
|
|
}
|
|
for _, v := range vv {
|
|
w.Header().Add(k, v)
|
|
}
|
|
}
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(resp.StatusCode)
|
|
_ = json.NewEncoder(w).Encode(anthropicResp)
|
|
return
|
|
}
|
|
|
|
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
|
|
}
|
|
|
|
proxyAnthropicStream(w, resp.Body, requestedModel, flusher)
|
|
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 handleResponses(cm *ConfigManager, expectedTokens []string, retryBaseDelay ...time.Duration) http.HandlerFunc {
|
|
baseDelay := 100 * time.Millisecond
|
|
if len(retryBaseDelay) > 0 {
|
|
baseDelay = retryBaseDelay[0]
|
|
}
|
|
|
|
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
|
|
}
|
|
|
|
openAIBodyMap := convertResponsesToOpenAI(bodyMap)
|
|
|
|
var requestedModel string
|
|
if m, ok := openAIBodyMap["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 := openAIBodyMap["stream"]; ok {
|
|
if b, ok := s.(bool); ok {
|
|
isStream = b
|
|
}
|
|
}
|
|
|
|
log.Printf("Received OpenAI Responses request for model %q (stream=%t). Found %d config trials.", requestedModel, isStream, len(trialConfigs))
|
|
|
|
for i, trial := range trialConfigs {
|
|
if i > 0 && baseDelay > 0 {
|
|
delay := time.Duration(fibonacci(i)) * baseDelay
|
|
log.Printf("Trial %d/%d: Fibonacci backoff delay of %v before retry...", i+1, len(trialConfigs), delay)
|
|
select {
|
|
case <-r.Context().Done():
|
|
log.Printf("Request context cancelled during retry delay before trial %d", i+1)
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(499)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"error": map[string]any{
|
|
"message": "Client closed request",
|
|
"type": "gateway_error",
|
|
},
|
|
})
|
|
return
|
|
case <-time.After(delay):
|
|
}
|
|
}
|
|
|
|
log.Printf("Trial %d/%d: model=%s endpoint=%s key_len=%d", i+1, len(trialConfigs), trial.Model, trial.Endpoint, len(trial.Key))
|
|
|
|
openAIBodyMap["model"] = trial.Model
|
|
modifiedBody, err := json.Marshal(openAIBodyMap)
|
|
if err != nil {
|
|
log.Printf("Trial %d: Failed to marshal body for %s: %v", i+1, trial.Model, err)
|
|
continue
|
|
}
|
|
|
|
targetURL := buildURL(trial.Endpoint, "/v1/chat/completions")
|
|
|
|
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)
|
|
|
|
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
|
|
}
|
|
|
|
var openAIResp map[string]any
|
|
if err := json.Unmarshal(respBody, &openAIResp); err != nil {
|
|
log.Printf("Trial %d: Failed to unmarshal upstream response JSON: %v", i+1, err)
|
|
continue
|
|
}
|
|
|
|
responsesResp := convertOpenAIToResponsesResponse(openAIResp, requestedModel)
|
|
|
|
for k, vv := range resp.Header {
|
|
kLower := strings.ToLower(k)
|
|
if kLower == "content-length" || kLower == "content-type" {
|
|
continue
|
|
}
|
|
for _, v := range vv {
|
|
w.Header().Add(k, v)
|
|
}
|
|
}
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(resp.StatusCode)
|
|
_ = json.NewEncoder(w).Encode(responsesResp)
|
|
return
|
|
}
|
|
|
|
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
|
|
}
|
|
|
|
proxyResponsesStream(w, resp.Body, requestedModel, flusher)
|
|
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 convertAnthropicToOpenAI(bodyMap map[string]any) map[string]any {
|
|
openAIBody := make(map[string]any)
|
|
|
|
if m, ok := bodyMap["model"]; ok {
|
|
openAIBody["model"] = m
|
|
}
|
|
if s, ok := bodyMap["stream"].(bool); ok {
|
|
openAIBody["stream"] = s
|
|
}
|
|
if t, ok := bodyMap["temperature"]; ok {
|
|
openAIBody["temperature"] = t
|
|
}
|
|
if p, ok := bodyMap["top_p"]; ok {
|
|
openAIBody["top_p"] = p
|
|
}
|
|
if mt, ok := bodyMap["max_tokens"]; ok {
|
|
openAIBody["max_tokens"] = mt
|
|
}
|
|
if stops, ok := bodyMap["stop_sequences"]; ok {
|
|
openAIBody["stop"] = stops
|
|
}
|
|
|
|
var openAIMessages []any
|
|
|
|
if sysVal, ok := bodyMap["system"]; ok && sysVal != nil {
|
|
sysContent := normalizeContent(sysVal)
|
|
if sysContent != nil {
|
|
openAIMessages = append(openAIMessages, map[string]any{
|
|
"role": "system",
|
|
"content": sysContent,
|
|
})
|
|
}
|
|
}
|
|
|
|
if msgs, ok := bodyMap["messages"].([]any); ok {
|
|
for _, item := range msgs {
|
|
if msgMap, ok := item.(map[string]any); ok {
|
|
role, _ := msgMap["role"].(string)
|
|
if role == "" {
|
|
role = "user"
|
|
}
|
|
cnt := msgMap["content"]
|
|
normCnt := convertAnthropicContentToOpenAI(cnt)
|
|
|
|
openAIMessages = append(openAIMessages, map[string]any{
|
|
"role": role,
|
|
"content": normCnt,
|
|
})
|
|
}
|
|
}
|
|
}
|
|
|
|
openAIBody["messages"] = openAIMessages
|
|
autodetectAndNormalizeMessages(openAIBody)
|
|
return openAIBody
|
|
}
|
|
|
|
func convertAnthropicContentToOpenAI(cntAny any) any {
|
|
if slice, ok := cntAny.([]any); ok {
|
|
var newSlice []any
|
|
for _, item := range slice {
|
|
if m, ok := item.(map[string]any); ok {
|
|
if t, ok := m["type"].(string); ok && t == "image" {
|
|
if src, ok := m["source"].(map[string]any); ok {
|
|
mediaType, _ := src["media_type"].(string)
|
|
if mediaType == "" {
|
|
mediaType = "image/png"
|
|
}
|
|
b64Data, _ := src["data"].(string)
|
|
dataURL := fmt.Sprintf("data:%s;base64,%s", mediaType, b64Data)
|
|
newSlice = append(newSlice, map[string]any{
|
|
"type": "image_url",
|
|
"image_url": map[string]any{
|
|
"url": dataURL,
|
|
},
|
|
})
|
|
continue
|
|
}
|
|
}
|
|
}
|
|
newSlice = append(newSlice, item)
|
|
}
|
|
return normalizeContent(newSlice)
|
|
}
|
|
return normalizeContent(cntAny)
|
|
}
|
|
|
|
func convertOpenAIToAnthropicResponse(openAIResp map[string]any, requestedModel string) map[string]any {
|
|
id, _ := openAIResp["id"].(string)
|
|
if id == "" {
|
|
id = fmt.Sprintf("msg_%d", time.Now().UnixNano())
|
|
} else if !strings.HasPrefix(id, "msg_") {
|
|
id = "msg_" + id
|
|
}
|
|
|
|
model, _ := openAIResp["model"].(string)
|
|
if model == "" {
|
|
model = requestedModel
|
|
}
|
|
|
|
var textContent string
|
|
var finishReason string
|
|
|
|
if choices, ok := openAIResp["choices"].([]any); ok && len(choices) > 0 {
|
|
if choice, ok := choices[0].(map[string]any); ok {
|
|
if fr, ok := choice["finish_reason"].(string); ok {
|
|
finishReason = fr
|
|
}
|
|
if msg, ok := choice["message"].(map[string]any); ok {
|
|
if cnt, ok := msg["content"].(string); ok {
|
|
textContent = cnt
|
|
} else if cntNorm := normalizeContent(msg["content"]); cntNorm != nil {
|
|
if s, ok := cntNorm.(string); ok {
|
|
textContent = s
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
stopReason := "end_turn"
|
|
switch finishReason {
|
|
case "length":
|
|
stopReason = "max_tokens"
|
|
case "tool_calls", "function_call":
|
|
stopReason = "tool_use"
|
|
case "stop":
|
|
stopReason = "end_turn"
|
|
}
|
|
|
|
inputTokens := 0
|
|
outputTokens := 0
|
|
if usage, ok := openAIResp["usage"].(map[string]any); ok {
|
|
if pt, ok := usage["prompt_tokens"].(float64); ok {
|
|
inputTokens = int(pt)
|
|
}
|
|
if ct, ok := usage["completion_tokens"].(float64); ok {
|
|
outputTokens = int(ct)
|
|
}
|
|
}
|
|
|
|
return map[string]any{
|
|
"id": id,
|
|
"type": "message",
|
|
"role": "assistant",
|
|
"model": model,
|
|
"content": []any{
|
|
map[string]any{
|
|
"type": "text",
|
|
"text": textContent,
|
|
},
|
|
},
|
|
"stop_reason": stopReason,
|
|
"stop_sequence": nil,
|
|
"usage": map[string]any{
|
|
"input_tokens": inputTokens,
|
|
"output_tokens": outputTokens,
|
|
},
|
|
}
|
|
}
|
|
|
|
func proxyAnthropicStream(w http.ResponseWriter, respBody io.ReadCloser, requestedModel string, flusher http.Flusher) {
|
|
defer respBody.Close()
|
|
|
|
w.Header().Set("Content-Type", "text/event-stream")
|
|
w.Header().Set("Cache-Control", "no-cache")
|
|
w.Header().Set("Connection", "keep-alive")
|
|
w.WriteHeader(http.StatusOK)
|
|
flusher.Flush()
|
|
|
|
reader := bufio.NewReader(respBody)
|
|
|
|
msgID := fmt.Sprintf("msg_%d", time.Now().UnixNano())
|
|
|
|
msgStartObj := map[string]any{
|
|
"type": "message_start",
|
|
"message": map[string]any{
|
|
"id": msgID,
|
|
"type": "message",
|
|
"role": "assistant",
|
|
"model": requestedModel,
|
|
"content": []any{},
|
|
"stop_reason": nil,
|
|
"stop_sequence": nil,
|
|
"usage": map[string]any{
|
|
"input_tokens": 0,
|
|
"output_tokens": 0,
|
|
},
|
|
},
|
|
}
|
|
msgStartBytes, _ := json.Marshal(msgStartObj)
|
|
_, _ = fmt.Fprintf(w, "event: message_start\ndata: %s\n\n", msgStartBytes)
|
|
|
|
blockStartObj := map[string]any{
|
|
"type": "content_block_start",
|
|
"index": 0,
|
|
"content_block": map[string]any{
|
|
"type": "text",
|
|
"text": "",
|
|
},
|
|
}
|
|
blockStartBytes, _ := json.Marshal(blockStartObj)
|
|
_, _ = fmt.Fprintf(w, "event: content_block_start\ndata: %s\n\n", blockStartBytes)
|
|
flusher.Flush()
|
|
|
|
var finalStopReason = "end_turn"
|
|
|
|
for {
|
|
lineBytes, err := reader.ReadBytes('\n')
|
|
if len(lineBytes) > 0 {
|
|
line := strings.TrimSpace(string(lineBytes))
|
|
if strings.HasPrefix(line, "data: ") {
|
|
dataStr := strings.TrimPrefix(line, "data: ")
|
|
dataStr = strings.TrimSpace(dataStr)
|
|
if dataStr == "[DONE]" {
|
|
break
|
|
}
|
|
var chunkMap map[string]any
|
|
if json.Unmarshal([]byte(dataStr), &chunkMap) == nil {
|
|
if choices, ok := chunkMap["choices"].([]any); ok && len(choices) > 0 {
|
|
if choice, ok := choices[0].(map[string]any); ok {
|
|
if fr, ok := choice["finish_reason"].(string); ok && fr != "" {
|
|
switch fr {
|
|
case "length":
|
|
finalStopReason = "max_tokens"
|
|
case "tool_calls", "function_call":
|
|
finalStopReason = "tool_use"
|
|
case "stop":
|
|
finalStopReason = "end_turn"
|
|
}
|
|
}
|
|
if delta, ok := choice["delta"].(map[string]any); ok {
|
|
if contentStr, ok := delta["content"].(string); ok && contentStr != "" {
|
|
deltaObj := map[string]any{
|
|
"type": "content_block_delta",
|
|
"index": 0,
|
|
"delta": map[string]any{
|
|
"type": "text_delta",
|
|
"text": contentStr,
|
|
},
|
|
}
|
|
deltaBytes, _ := json.Marshal(deltaObj)
|
|
_, _ = fmt.Fprintf(w, "event: content_block_delta\ndata: %s\n\n", deltaBytes)
|
|
flusher.Flush()
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
if err != nil {
|
|
break
|
|
}
|
|
}
|
|
|
|
blockStopObj := map[string]any{
|
|
"type": "content_block_stop",
|
|
"index": 0,
|
|
}
|
|
blockStopBytes, _ := json.Marshal(blockStopObj)
|
|
_, _ = fmt.Fprintf(w, "event: content_block_stop\ndata: %s\n\n", blockStopBytes)
|
|
|
|
msgDeltaObj := map[string]any{
|
|
"type": "message_delta",
|
|
"delta": map[string]any{
|
|
"stop_reason": finalStopReason,
|
|
"stop_sequence": nil,
|
|
},
|
|
"usage": map[string]any{
|
|
"output_tokens": 0,
|
|
},
|
|
}
|
|
msgDeltaBytes, _ := json.Marshal(msgDeltaObj)
|
|
_, _ = fmt.Fprintf(w, "event: message_delta\ndata: %s\n\n", msgDeltaBytes)
|
|
|
|
msgStopObj := map[string]any{
|
|
"type": "message_stop",
|
|
}
|
|
msgStopBytes, _ := json.Marshal(msgStopObj)
|
|
_, _ = fmt.Fprintf(w, "event: message_stop\ndata: %s\n\n", msgStopBytes)
|
|
flusher.Flush()
|
|
}
|
|
|
|
func convertResponsesToOpenAI(bodyMap map[string]any) map[string]any {
|
|
openAIBody := make(map[string]any)
|
|
|
|
if m, ok := bodyMap["model"]; ok {
|
|
openAIBody["model"] = m
|
|
}
|
|
if s, ok := bodyMap["stream"].(bool); ok {
|
|
openAIBody["stream"] = s
|
|
}
|
|
if t, ok := bodyMap["temperature"]; ok {
|
|
openAIBody["temperature"] = t
|
|
}
|
|
if p, ok := bodyMap["top_p"]; ok {
|
|
openAIBody["top_p"] = p
|
|
}
|
|
if mt, ok := bodyMap["max_output_tokens"]; ok {
|
|
openAIBody["max_tokens"] = mt
|
|
} else if mt, ok := bodyMap["max_tokens"]; ok {
|
|
openAIBody["max_tokens"] = mt
|
|
}
|
|
|
|
var openAIMessages []any
|
|
|
|
if instVal, ok := bodyMap["instructions"]; ok && instVal != nil {
|
|
sysContent := normalizeContent(instVal)
|
|
if sysContent != nil {
|
|
openAIMessages = append(openAIMessages, map[string]any{
|
|
"role": "system",
|
|
"content": sysContent,
|
|
})
|
|
}
|
|
}
|
|
|
|
if inputVal, ok := bodyMap["input"]; ok && inputVal != nil {
|
|
switch inp := inputVal.(type) {
|
|
case string:
|
|
openAIMessages = append(openAIMessages, map[string]any{
|
|
"role": "user",
|
|
"content": inp,
|
|
})
|
|
case []any:
|
|
for _, elem := range inp {
|
|
switch item := elem.(type) {
|
|
case string:
|
|
openAIMessages = append(openAIMessages, map[string]any{
|
|
"role": "user",
|
|
"content": item,
|
|
})
|
|
case map[string]any:
|
|
role, _ := item["role"].(string)
|
|
if role == "" {
|
|
role = "user"
|
|
}
|
|
cnt := item["content"]
|
|
if cnt == nil {
|
|
if txt, ok := item["text"].(string); ok {
|
|
cnt = txt
|
|
}
|
|
}
|
|
openAIMessages = append(openAIMessages, map[string]any{
|
|
"role": role,
|
|
"content": normalizeContent(cnt),
|
|
})
|
|
}
|
|
}
|
|
case map[string]any:
|
|
role, _ := inp["role"].(string)
|
|
if role == "" {
|
|
role = "user"
|
|
}
|
|
cnt := inp["content"]
|
|
if cnt == nil {
|
|
if txt, ok := inp["text"].(string); ok {
|
|
cnt = txt
|
|
}
|
|
}
|
|
openAIMessages = append(openAIMessages, map[string]any{
|
|
"role": role,
|
|
"content": normalizeContent(cnt),
|
|
})
|
|
}
|
|
} else if msgs, ok := bodyMap["messages"].([]any); ok {
|
|
for _, item := range msgs {
|
|
if msgMap, ok := item.(map[string]any); ok {
|
|
role, _ := msgMap["role"].(string)
|
|
if role == "" {
|
|
role = "user"
|
|
}
|
|
openAIMessages = append(openAIMessages, map[string]any{
|
|
"role": role,
|
|
"content": normalizeContent(msgMap["content"]),
|
|
})
|
|
}
|
|
}
|
|
}
|
|
|
|
openAIBody["messages"] = openAIMessages
|
|
autodetectAndNormalizeMessages(openAIBody)
|
|
return openAIBody
|
|
}
|
|
|
|
func convertOpenAIToResponsesResponse(openAIResp map[string]any, requestedModel string) map[string]any {
|
|
id, _ := openAIResp["id"].(string)
|
|
if id == "" {
|
|
id = fmt.Sprintf("resp_%d", time.Now().UnixNano())
|
|
} else if !strings.HasPrefix(id, "resp_") {
|
|
id = "resp_" + id
|
|
}
|
|
|
|
model, _ := openAIResp["model"].(string)
|
|
if model == "" {
|
|
model = requestedModel
|
|
}
|
|
|
|
var textContent string
|
|
if choices, ok := openAIResp["choices"].([]any); ok && len(choices) > 0 {
|
|
if choice, ok := choices[0].(map[string]any); ok {
|
|
if msg, ok := choice["message"].(map[string]any); ok {
|
|
if cnt, ok := msg["content"].(string); ok {
|
|
textContent = cnt
|
|
} else if cntNorm := normalizeContent(msg["content"]); cntNorm != nil {
|
|
if s, ok := cntNorm.(string); ok {
|
|
textContent = s
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
created := time.Now().Unix()
|
|
if c, ok := openAIResp["created"].(float64); ok {
|
|
created = int64(c)
|
|
}
|
|
|
|
promptTokens := 0
|
|
completionTokens := 0
|
|
totalTokens := 0
|
|
if usage, ok := openAIResp["usage"].(map[string]any); ok {
|
|
if pt, ok := usage["prompt_tokens"].(float64); ok {
|
|
promptTokens = int(pt)
|
|
}
|
|
if ct, ok := usage["completion_tokens"].(float64); ok {
|
|
completionTokens = int(ct)
|
|
}
|
|
if tt, ok := usage["total_tokens"].(float64); ok {
|
|
totalTokens = int(tt)
|
|
}
|
|
}
|
|
|
|
msgID := fmt.Sprintf("msg_%d", time.Now().UnixNano())
|
|
|
|
return map[string]any{
|
|
"id": id,
|
|
"object": "response",
|
|
"created_at": created,
|
|
"status": "completed",
|
|
"model": model,
|
|
"output": []any{
|
|
map[string]any{
|
|
"id": msgID,
|
|
"type": "message",
|
|
"status": "completed",
|
|
"role": "assistant",
|
|
"content": []any{
|
|
map[string]any{
|
|
"type": "output_text",
|
|
"text": textContent,
|
|
"annotations": []any{},
|
|
"logprobs": []any{},
|
|
},
|
|
},
|
|
},
|
|
},
|
|
"usage": map[string]any{
|
|
"prompt_tokens": promptTokens,
|
|
"completion_tokens": completionTokens,
|
|
"total_tokens": totalTokens,
|
|
},
|
|
}
|
|
}
|
|
|
|
func proxyResponsesStream(w http.ResponseWriter, respBody io.ReadCloser, requestedModel string, flusher http.Flusher) {
|
|
defer respBody.Close()
|
|
|
|
w.Header().Set("Content-Type", "text/event-stream")
|
|
w.Header().Set("Cache-Control", "no-cache")
|
|
w.Header().Set("Connection", "keep-alive")
|
|
w.WriteHeader(http.StatusOK)
|
|
flusher.Flush()
|
|
|
|
reader := bufio.NewReader(respBody)
|
|
|
|
respID := fmt.Sprintf("resp_%d", time.Now().UnixNano())
|
|
msgID := fmt.Sprintf("msg_%d", time.Now().UnixNano())
|
|
|
|
createdObj := map[string]any{
|
|
"type": "response.created",
|
|
"response": map[string]any{
|
|
"id": respID,
|
|
"object": "response",
|
|
"status": "in_progress",
|
|
"model": requestedModel,
|
|
},
|
|
}
|
|
createdBytes, _ := json.Marshal(createdObj)
|
|
_, _ = fmt.Fprintf(w, "event: response.created\ndata: %s\n\n", createdBytes)
|
|
|
|
partAddedObj := map[string]any{
|
|
"type": "response.content_part.added",
|
|
"part": map[string]any{
|
|
"type": "output_text",
|
|
"text": "",
|
|
},
|
|
}
|
|
partAddedBytes, _ := json.Marshal(partAddedObj)
|
|
_, _ = fmt.Fprintf(w, "event: response.content_part.added\ndata: %s\n\n", partAddedBytes)
|
|
flusher.Flush()
|
|
|
|
var fullTextBuf strings.Builder
|
|
|
|
for {
|
|
lineBytes, err := reader.ReadBytes('\n')
|
|
if len(lineBytes) > 0 {
|
|
line := strings.TrimSpace(string(lineBytes))
|
|
if strings.HasPrefix(line, "data: ") {
|
|
dataStr := strings.TrimPrefix(line, "data: ")
|
|
dataStr = strings.TrimSpace(dataStr)
|
|
if dataStr == "[DONE]" {
|
|
break
|
|
}
|
|
var chunkMap map[string]any
|
|
if json.Unmarshal([]byte(dataStr), &chunkMap) == nil {
|
|
if choices, ok := chunkMap["choices"].([]any); ok && len(choices) > 0 {
|
|
if choice, ok := choices[0].(map[string]any); ok {
|
|
if delta, ok := choice["delta"].(map[string]any); ok {
|
|
if contentStr, ok := delta["content"].(string); ok && contentStr != "" {
|
|
fullTextBuf.WriteString(contentStr)
|
|
deltaObj := map[string]any{
|
|
"type": "response.output_text.delta",
|
|
"delta": contentStr,
|
|
}
|
|
deltaBytes, _ := json.Marshal(deltaObj)
|
|
_, _ = fmt.Fprintf(w, "event: response.output_text.delta\ndata: %s\n\n", deltaBytes)
|
|
flusher.Flush()
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
if err != nil {
|
|
break
|
|
}
|
|
}
|
|
|
|
completedObj := map[string]any{
|
|
"type": "response.completed",
|
|
"response": map[string]any{
|
|
"id": respID,
|
|
"object": "response",
|
|
"status": "completed",
|
|
"model": requestedModel,
|
|
"created_at": time.Now().Unix(),
|
|
"output": []any{
|
|
map[string]any{
|
|
"id": msgID,
|
|
"type": "message",
|
|
"status": "completed",
|
|
"role": "assistant",
|
|
"content": []any{
|
|
map[string]any{
|
|
"type": "output_text",
|
|
"text": fullTextBuf.String(),
|
|
"annotations": []any{},
|
|
"logprobs": []any{},
|
|
},
|
|
},
|
|
},
|
|
},
|
|
},
|
|
}
|
|
completedBytes, _ := json.Marshal(completedObj)
|
|
_, _ = fmt.Fprintf(w, "event: response.completed\ndata: %s\n\n", completedBytes)
|
|
flusher.Flush()
|
|
}
|
|
|
|
|
|
|