switched to fib backoff

This commit is contained in:
Luxferre
2026-07-30 09:53:36 +03:00
parent 1cc6ddb2e4
commit d49339d770
4 changed files with 519 additions and 5 deletions
+2 -1
View File
@@ -13,7 +13,7 @@ Dynagate is a lightweight, high-performance LLM gateway written in Go that acts
- Watches `models.csv` (overridable via command-line flags) continuously using a background thread and automatically reloads configuration updates without dropping active connections. - Watches `models.csv` (overridable via command-line flags) continuously using a background thread and automatically reloads configuration updates without dropping active connections.
3. **Fallback and retry orchestration**: 3. **Fallback and retry orchestration**:
- Attempts keys configured for the requested model sequentially. - Attempts keys configured for the requested model sequentially using Fibonacci backoff delays (`-retry-delay`).
- Cascades automatically to the next model in the priority chain (wrapping around to cover all entries) if all keys for the requested model fail. - Cascades automatically to the next model in the priority chain (wrapping around to cover all entries) if all keys for the requested model fail.
- Updates the `"model"` field dynamically in the outgoing request body before proxying, ensuring upstream compatibility. - Updates the `"model"` field dynamically in the outgoing request body before proxying, ensuring upstream compatibility.
@@ -98,6 +98,7 @@ my-secure-gateway-token-2
| `-key` | string | `""` | Path to a file containing the gateway's access tokens (one per line). If blank, authentication is disabled. | | `-key` | string | `""` | Path to a file containing the gateway's access tokens (one per line). If blank, authentication is disabled. |
| `-csv-updater` | string | `""` | Command/script to run periodically to update `models.csv`. | | `-csv-updater` | string | `""` | Command/script to run periodically to update `models.csv`. |
| `-csv-update-interval`| int | `10` | Frequency in minutes to invoke `-csv-updater`. | | `-csv-update-interval`| int | `10` | Frequency in minutes to invoke `-csv-updater`. |
| `-retry-delay` | int | `100` | Base delay in milliseconds for Fibonacci failover retries (0 to disable). |
### Model CSV format rules ### Model CSV format rules
The gateway maps columns dynamically by looking at the header row. If no header is present, it defaults to: The gateway maps columns dynamically by looking at the header row. If no header is present, it defaults to:
+211
View File
@@ -779,5 +779,216 @@ func TestMultipleTokensKeyFile(t *testing.T) {
} }
} }
func TestRequestMessagesAutodetectionAndContentArrays(t *testing.T) {
var receivedBody map[string]any
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
var b map[string]any
_ = json.NewDecoder(r.Body).Decode(&b)
receivedBody = b
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"OK"}}]} `))
}))
defer server.Close()
configs := []ModelConfig{
{Model: "test-model", Key: "test-key", Endpoint: server.URL},
}
cm := &ConfigManager{
configs: configs,
uniqueModels: []string{"test-model"},
}
handler := handleChatCompletions(cm, nil)
t.Run("Pi agent format - array content", func(t *testing.T) {
reqObj := map[string]any{
"model": "test-model",
"messages": []any{
map[string]any{
"role": "user",
"content": []any{
map[string]any{"type": "text", "text": "Hello from Pi agent"},
},
},
},
}
reqBytes, _ := json.Marshal(reqObj)
req := httptest.NewRequest("POST", "/v1/chat/completions", bytes.NewReader(reqBytes))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
if w.Result().StatusCode != http.StatusOK {
t.Fatalf("Expected status 200, got %d", w.Result().StatusCode)
}
msgs, ok := receivedBody["messages"].([]any)
if !ok || len(msgs) != 1 {
t.Fatalf("Expected 1 message in forwarded body, got %+v", receivedBody)
}
userMsg := msgs[0].(map[string]any)
if userMsg["role"] != "user" || userMsg["content"] != "Hello from Pi agent" {
t.Errorf("Expected content 'Hello from Pi agent', got %+v", userMsg)
}
})
t.Run("Top-level system prompt", func(t *testing.T) {
reqObj := map[string]any{
"model": "test-model",
"system": "You are a helpful coding assistant",
"messages": []any{
map[string]any{
"role": "user",
"content": "Write a test",
},
},
}
reqBytes, _ := json.Marshal(reqObj)
req := httptest.NewRequest("POST", "/v1/chat/completions", bytes.NewReader(reqBytes))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
msgs, ok := receivedBody["messages"].([]any)
if !ok || len(msgs) != 2 {
t.Fatalf("Expected 2 messages in forwarded body, got %+v", receivedBody)
}
sysMsg := msgs[0].(map[string]any)
userMsg := msgs[1].(map[string]any)
if sysMsg["role"] != "system" || sysMsg["content"] != "You are a helpful coding assistant" {
t.Errorf("Unexpected system message: %+v", sysMsg)
}
if userMsg["role"] != "user" || userMsg["content"] != "Write a test" {
t.Errorf("Unexpected user message: %+v", userMsg)
}
if _, exists := receivedBody["system"]; exists {
t.Errorf("Expected top-level system field to be removed")
}
})
t.Run("Prompt field autodetection", func(t *testing.T) {
reqObj := map[string]any{
"model": "test-model",
"prompt": "Explain recursion",
}
reqBytes, _ := json.Marshal(reqObj)
req := httptest.NewRequest("POST", "/v1/chat/completions", bytes.NewReader(reqBytes))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
msgs, ok := receivedBody["messages"].([]any)
if !ok || len(msgs) != 1 {
t.Fatalf("Expected 1 message in forwarded body, got %+v", receivedBody)
}
userMsg := msgs[0].(map[string]any)
if userMsg["role"] != "user" || userMsg["content"] != "Explain recursion" {
t.Errorf("Unexpected user message: %+v", userMsg)
}
})
t.Run("Multimodal array content preserved", func(t *testing.T) {
reqObj := map[string]any{
"model": "test-model",
"messages": []any{
map[string]any{
"role": "user",
"content": []any{
map[string]any{"type": "text", "text": "Describe this image"},
map[string]any{"type": "image_url", "image_url": map[string]any{"url": "http://example.com/img.png"}},
},
},
},
}
reqBytes, _ := json.Marshal(reqObj)
req := httptest.NewRequest("POST", "/v1/chat/completions", bytes.NewReader(reqBytes))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
msgs, ok := receivedBody["messages"].([]any)
if !ok || len(msgs) != 1 {
t.Fatalf("Expected 1 message in forwarded body, got %+v", receivedBody)
}
userMsg := msgs[0].(map[string]any)
cntArray, isArray := userMsg["content"].([]any)
if !isArray || len(cntArray) != 2 {
t.Fatalf("Expected multimodal array content to be preserved, got %+v", userMsg["content"])
}
})
}
func TestFibonacciBackoff(t *testing.T) {
expectedSeq := []struct {
n int
expected int64
}{
{0, 0},
{1, 1},
{2, 1},
{3, 2},
{4, 3},
{5, 5},
{6, 8},
{7, 13},
{8, 21},
{9, 34},
{10, 55},
}
for _, tc := range expectedSeq {
got := fibonacci(tc.n)
if got != tc.expected {
t.Errorf("fibonacci(%d) = %d; want %d", tc.n, got, tc.expected)
}
}
// Verify failover works with zero-delay in handleChatCompletions
serverA := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusInternalServerError)
}))
defer serverA.Close()
serverB := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"Success from B"}}]} `))
}))
defer serverB.Close()
configs := []ModelConfig{
{Model: "m1", Key: "k1", Endpoint: serverA.URL},
{Model: "m1", Key: "k2", Endpoint: serverB.URL},
}
cm := &ConfigManager{
configs: configs,
uniqueModels: []string{"m1"},
}
handler := handleChatCompletions(cm, nil, 0) // 0 delay for fast test
reqBodyObj := map[string]any{
"model": "m1",
"messages": []map[string]string{
{"role": "user", "content": "test"},
},
}
reqBytes, _ := json.Marshal(reqBodyObj)
req := httptest.NewRequest("POST", "/v1/chat/completions", bytes.NewReader(reqBytes))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
if w.Result().StatusCode != http.StatusOK {
t.Fatalf("Expected status 200 after Fibonacci retry, got %d", w.Result().StatusCode)
}
}
+301 -2
View File
@@ -10,6 +10,7 @@ import (
"log" "log"
"net/http" "net/http"
"strings" "strings"
"time"
) )
var httpClient = &http.Client{} var httpClient = &http.Client{}
@@ -96,7 +97,12 @@ func handleModels(cm *ConfigManager, expectedTokens []string) http.HandlerFunc {
} }
} }
func handleChatCompletions(cm *ConfigManager, expectedTokens []string) http.HandlerFunc { 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) { return func(w http.ResponseWriter, r *http.Request) {
if !checkAuth(expectedTokens, r) { if !checkAuth(expectedTokens, r) {
sendUnauthorized(w) sendUnauthorized(w)
@@ -141,6 +147,8 @@ func handleChatCompletions(cm *ConfigManager, expectedTokens []string) http.Hand
return return
} }
autodetectAndNormalizeMessages(bodyMap)
var requestedModel string var requestedModel string
if m, ok := bodyMap["model"]; ok { if m, ok := bodyMap["model"]; ok {
if s, ok := m.(string); ok { if s, ok := m.(string); ok {
@@ -186,6 +194,25 @@ func handleChatCompletions(cm *ConfigManager, expectedTokens []string) http.Hand
log.Printf("Received completion request for model %q (stream=%t). Found %d config trials.", requestedModel, isStream, len(trialConfigs)) log.Printf("Received completion request for model %q (stream=%t). Found %d config trials.", requestedModel, isStream, len(trialConfigs))
for i, trial := range 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)) 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 bodyMap["model"] = trial.Model
@@ -329,7 +356,12 @@ func handleChatCompletions(cm *ConfigManager, expectedTokens []string) http.Hand
} }
} }
func handleImageGenerations(cm *ConfigManager, expectedTokens []string) http.HandlerFunc { 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) { return func(w http.ResponseWriter, r *http.Request) {
if !checkAuth(expectedTokens, r) { if !checkAuth(expectedTokens, r) {
sendUnauthorized(w) sendUnauthorized(w)
@@ -374,6 +406,10 @@ func handleImageGenerations(cm *ConfigManager, expectedTokens []string) http.Han
return return
} }
if p, ok := bodyMap["prompt"]; ok && p != nil {
bodyMap["prompt"] = normalizeContent(p)
}
var requestedModel string var requestedModel string
if m, ok := bodyMap["model"]; ok { if m, ok := bodyMap["model"]; ok {
if s, ok := m.(string); ok { if s, ok := m.(string); ok {
@@ -412,6 +448,25 @@ func handleImageGenerations(cm *ConfigManager, expectedTokens []string) http.Han
log.Printf("Received image generation request for model %q. Found %d config trials.", requestedModel, len(trialConfigs)) log.Printf("Received image generation request for model %q. Found %d config trials.", requestedModel, len(trialConfigs))
for i, trial := range 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)) 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 bodyMap["model"] = trial.Model
@@ -555,3 +610,247 @@ func buildURL(endpoint string, reqPath string) string {
return endpoint + "/" + reqPath 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
}
+5 -2
View File
@@ -25,8 +25,11 @@ func main() {
keyPath := flag.String("key", "", "Path to plaintext file containing expected auth token") keyPath := flag.String("key", "", "Path to plaintext file containing expected auth token")
csvUpdater := flag.String("csv-updater", "", "Command to run to update the configuration CSV file") csvUpdater := flag.String("csv-updater", "", "Command to run to update the configuration CSV file")
csvUpdateInterval := flag.Int("csv-update-interval", 10, "Interval in minutes with which to run the CSV updater command") csvUpdateInterval := flag.Int("csv-update-interval", 10, "Interval in minutes with which to run the CSV updater command")
retryDelayMs := flag.Int("retry-delay", 100, "Base delay in milliseconds for Fibonacci failover retries (0 to disable)")
flag.Parse() flag.Parse()
retryDelay := time.Duration(*retryDelayMs) * time.Millisecond
var authTokens []string var authTokens []string
if *keyPath != "" { if *keyPath != "" {
content, err := os.ReadFile(*keyPath) content, err := os.ReadFile(*keyPath)
@@ -74,8 +77,8 @@ func main() {
// Register routes // Register routes
mux := http.NewServeMux() mux := http.NewServeMux()
mux.HandleFunc("/v1/models", handleModels(cm, authTokens)) mux.HandleFunc("/v1/models", handleModels(cm, authTokens))
mux.HandleFunc("/v1/chat/completions", handleChatCompletions(cm, authTokens)) mux.HandleFunc("/v1/chat/completions", handleChatCompletions(cm, authTokens, retryDelay))
mux.HandleFunc("/v1/images/generations", handleImageGenerations(cm, authTokens)) mux.HandleFunc("/v1/images/generations", handleImageGenerations(cm, authTokens, retryDelay))
server := &http.Server{ server := &http.Server{
Addr: fmt.Sprintf(":%d", *port), Addr: fmt.Sprintf(":%d", *port),