feat: add non-ZeroGPU endpoints, automatic failover, and endpoint cooldown
This commit is contained in:
@@ -15,9 +15,11 @@ import (
|
||||
"log"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
@@ -29,6 +31,12 @@ var (
|
||||
ConfiguredBaseURL string
|
||||
ConfiguredModel string
|
||||
EnableThinkingDefault = true
|
||||
DefaultEndpoints = []string{
|
||||
"https://microhero-qwen3-8-27b-uncensored-chat.hf.space",
|
||||
"https://wanyamaelis-qwen3-8-27b.hf.space",
|
||||
"https://apathy-exe-qwen3-8-27b.hf.space",
|
||||
"https://apathy-exe-qwen3-8-flash-next.hf.space",
|
||||
}
|
||||
)
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
@@ -1195,19 +1203,52 @@ func parseAssistantText(dataJSON string) (string, bool) {
|
||||
return "", false
|
||||
}
|
||||
|
||||
type EndpointNode struct {
|
||||
URL string
|
||||
Mode string
|
||||
CooldownUntil time.Time
|
||||
FailureCount int
|
||||
}
|
||||
|
||||
type QwenService struct {
|
||||
endpoint string
|
||||
endpoints []*EndpointNode
|
||||
modelName string
|
||||
mode string
|
||||
token string
|
||||
apiKey string
|
||||
baseURL string
|
||||
enableThinking bool
|
||||
autoFailover bool
|
||||
mu sync.Mutex
|
||||
client *http.Client
|
||||
}
|
||||
|
||||
func NewQwenService(endpoint, modelName, mode, token, apiKey, baseURL, socksProxy string, enableThinking bool) *QwenService {
|
||||
cleanEndpoint := strings.TrimRight(endpoint, "/")
|
||||
func parseEndpointList(rawList []string) []*EndpointNode {
|
||||
var nodes []*EndpointNode
|
||||
seen := make(map[string]bool)
|
||||
for _, item := range rawList {
|
||||
parts := strings.Split(item, ",")
|
||||
for _, p := range parts {
|
||||
clean := strings.TrimRight(strings.TrimSpace(p), "/")
|
||||
if clean != "" && !seen[clean] {
|
||||
seen[clean] = true
|
||||
nodes = append(nodes, &EndpointNode{
|
||||
URL: clean,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(nodes) == 0 {
|
||||
for _, ep := range DefaultEndpoints {
|
||||
nodes = append(nodes, &EndpointNode{
|
||||
URL: ep,
|
||||
})
|
||||
}
|
||||
}
|
||||
return nodes
|
||||
}
|
||||
|
||||
func NewQwenService(endpoints []string, modelName, mode, token, apiKey, baseURL, socksProxy string, enableThinking, autoFailover bool) *QwenService {
|
||||
if modelName == "" {
|
||||
modelName = "Qwen/Qwen3.8-27B-Uncensored"
|
||||
}
|
||||
@@ -1227,14 +1268,17 @@ func NewQwenService(endpoint, modelName, mode, token, apiKey, baseURL, socksProx
|
||||
}
|
||||
}
|
||||
|
||||
nodes := parseEndpointList(endpoints)
|
||||
|
||||
return &QwenService{
|
||||
endpoint: cleanEndpoint,
|
||||
endpoints: nodes,
|
||||
modelName: modelName,
|
||||
mode: mode,
|
||||
token: token,
|
||||
apiKey: apiKey,
|
||||
baseURL: baseURL,
|
||||
enableThinking: enableThinking,
|
||||
autoFailover: autoFailover,
|
||||
client: &http.Client{Transport: transport, Timeout: 300 * time.Second},
|
||||
}
|
||||
}
|
||||
@@ -1274,24 +1318,24 @@ func (s *QwenService) ListModels() []ModelItem {
|
||||
return models
|
||||
}
|
||||
|
||||
func (s *QwenService) detectEndpointMode() string {
|
||||
if s.mode != "" && s.mode != "auto" {
|
||||
return s.mode
|
||||
func (s *QwenService) detectEndpointModeFor(epURL, defaultMode string) string {
|
||||
if defaultMode != "" && defaultMode != "auto" {
|
||||
return defaultMode
|
||||
}
|
||||
|
||||
ep := strings.ToLower(s.endpoint)
|
||||
ep := strings.ToLower(epURL)
|
||||
if strings.Contains(ep, "microhero") || strings.HasSuffix(ep, "/respond") {
|
||||
return "respond"
|
||||
}
|
||||
if strings.Contains(ep, "halvo78") || strings.HasSuffix(ep, "/chat_response") {
|
||||
return "chat_response"
|
||||
}
|
||||
if strings.Contains(ep, "apathy-exe") || strings.HasSuffix(ep, "/v1") {
|
||||
if strings.Contains(ep, "wanyamaelis") || strings.Contains(ep, "apathy-exe") || strings.HasSuffix(ep, "/v1") {
|
||||
return "openai"
|
||||
}
|
||||
|
||||
// Probe /gradio_api/info
|
||||
infoURL := s.endpoint + "/gradio_api/info"
|
||||
infoURL := epURL + "/gradio_api/info"
|
||||
req, err := http.NewRequest("GET", infoURL, nil)
|
||||
if err == nil {
|
||||
req.Header.Set("User-Agent", DefaultUserAgent)
|
||||
@@ -1316,28 +1360,144 @@ func (s *QwenService) detectEndpointMode() string {
|
||||
}
|
||||
}
|
||||
|
||||
return "respond"
|
||||
return "openai"
|
||||
}
|
||||
|
||||
func (s *QwenService) getEligibleEndpoints() []*EndpointNode {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
var ready []*EndpointNode
|
||||
var cooling []*EndpointNode
|
||||
|
||||
for _, node := range s.endpoints {
|
||||
if now.After(node.CooldownUntil) {
|
||||
ready = append(ready, node)
|
||||
} else {
|
||||
cooling = append(cooling, node)
|
||||
}
|
||||
}
|
||||
|
||||
if len(ready) > 0 {
|
||||
return ready
|
||||
}
|
||||
|
||||
// If all are cooling down, return all so we still attempt
|
||||
return cooling
|
||||
}
|
||||
|
||||
func (s *QwenService) markEndpointFailure(node *EndpointNode, err error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
node.FailureCount++
|
||||
errStr := strings.ToLower(err.Error())
|
||||
// If it's a quota or rate-limit error, cool down for 5 minutes
|
||||
if strings.Contains(errStr, "zerogpu") || strings.Contains(errStr, "quota") || strings.Contains(errStr, "429") {
|
||||
node.CooldownUntil = time.Now().Add(5 * time.Minute)
|
||||
log.Printf("Endpoint %s hit quota/rate limit, cooling down until %s", node.URL, node.CooldownUntil.Format("15:04:05"))
|
||||
} else {
|
||||
// For other transient errors, cool down for 30 seconds
|
||||
node.CooldownUntil = time.Now().Add(30 * time.Second)
|
||||
log.Printf("Endpoint %s failed (%v), cooling down for 30s", node.URL, err)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *QwenService) markEndpointSuccess(node *EndpointNode) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
node.FailureCount = 0
|
||||
node.CooldownUntil = time.Time{}
|
||||
}
|
||||
|
||||
func (s *QwenService) Chat(w http.ResponseWriter, r *http.Request, req ChatCompletionRequest) error {
|
||||
resolvedModel := EffectiveModelID(req.Model, s.modelName)
|
||||
maxTokens := ResolveMaxTokens(req)
|
||||
mode := s.detectEndpointMode()
|
||||
|
||||
switch mode {
|
||||
case "openai":
|
||||
return s.chatDirectOpenAI(w, r, req, resolvedModel, maxTokens)
|
||||
case "chat_response":
|
||||
return s.chatGradioChatResponse(w, r, req, resolvedModel, maxTokens)
|
||||
case "respond":
|
||||
fallthrough
|
||||
default:
|
||||
return s.chatGradioRespond(w, r, req, resolvedModel, maxTokens)
|
||||
endpoints := s.getEligibleEndpoints()
|
||||
if len(endpoints) == 0 {
|
||||
return fmt.Errorf("no upstream endpoints configured")
|
||||
}
|
||||
|
||||
var lastErr error
|
||||
for i, node := range endpoints {
|
||||
mode := node.Mode
|
||||
if mode == "" || mode == "auto" {
|
||||
mode = s.detectEndpointModeFor(node.URL, s.mode)
|
||||
node.Mode = mode
|
||||
}
|
||||
|
||||
if !req.Stream {
|
||||
rec := httptest.NewRecorder()
|
||||
var err error
|
||||
switch mode {
|
||||
case "openai":
|
||||
err = s.chatDirectOpenAI(node.URL, rec, r, req, resolvedModel, maxTokens)
|
||||
case "chat_response":
|
||||
err = s.chatGradioChatResponse(node.URL, rec, r, req, resolvedModel, maxTokens)
|
||||
case "respond":
|
||||
fallthrough
|
||||
default:
|
||||
err = s.chatGradioRespond(node.URL, rec, r, req, resolvedModel, maxTokens)
|
||||
}
|
||||
|
||||
if err == nil && rec.Code == http.StatusOK {
|
||||
s.markEndpointSuccess(node)
|
||||
for k, vv := range rec.Header() {
|
||||
for _, v := range vv {
|
||||
w.Header().Add(k, v)
|
||||
}
|
||||
}
|
||||
w.WriteHeader(rec.Code)
|
||||
w.Write(rec.Body.Bytes())
|
||||
return nil
|
||||
}
|
||||
|
||||
if err == nil && rec.Code != http.StatusOK {
|
||||
err = fmt.Errorf("HTTP %d: %s", rec.Code, rec.Body.String())
|
||||
}
|
||||
|
||||
lastErr = err
|
||||
s.markEndpointFailure(node, err)
|
||||
if !s.autoFailover || i == len(endpoints)-1 {
|
||||
break
|
||||
}
|
||||
log.Printf("Endpoint %s failed (%v), failing over to next endpoint...", node.URL, err)
|
||||
continue
|
||||
}
|
||||
|
||||
// Streaming request
|
||||
var err error
|
||||
switch mode {
|
||||
case "openai":
|
||||
err = s.chatDirectOpenAI(node.URL, w, r, req, resolvedModel, maxTokens)
|
||||
case "chat_response":
|
||||
err = s.chatGradioChatResponse(node.URL, w, r, req, resolvedModel, maxTokens)
|
||||
case "respond":
|
||||
fallthrough
|
||||
default:
|
||||
err = s.chatGradioRespond(node.URL, w, r, req, resolvedModel, maxTokens)
|
||||
}
|
||||
|
||||
if err == nil {
|
||||
s.markEndpointSuccess(node)
|
||||
return nil
|
||||
}
|
||||
|
||||
lastErr = err
|
||||
s.markEndpointFailure(node, err)
|
||||
if !s.autoFailover || i == len(endpoints)-1 {
|
||||
break
|
||||
}
|
||||
log.Printf("Endpoint %s streaming failed (%v), failing over to next endpoint...", node.URL, err)
|
||||
}
|
||||
|
||||
return fmt.Errorf("all endpoints failed (last error: %v)", lastErr)
|
||||
}
|
||||
|
||||
func (s *QwenService) chatDirectOpenAI(w http.ResponseWriter, r *http.Request, req ChatCompletionRequest, resolvedModel string, maxTokens int) error {
|
||||
targetURL := s.endpoint
|
||||
func (s *QwenService) chatDirectOpenAI(endpointURL string, w http.ResponseWriter, r *http.Request, req ChatCompletionRequest, resolvedModel string, maxTokens int) error {
|
||||
targetURL := endpointURL
|
||||
if !strings.HasSuffix(targetURL, "/v1/chat/completions") && !strings.HasSuffix(targetURL, "/chat/completions") {
|
||||
if strings.HasSuffix(targetURL, "/v1") {
|
||||
targetURL += "/chat/completions"
|
||||
@@ -1375,6 +1535,11 @@ func (s *QwenService) chatDirectOpenAI(w http.ResponseWriter, r *http.Request, r
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
bodyBytes, _ := io.ReadAll(resp.Body)
|
||||
return fmt.Errorf("upstream HTTP %d: %s", resp.StatusCode, string(bodyBytes))
|
||||
}
|
||||
|
||||
for k, vv := range resp.Header {
|
||||
for _, v := range vv {
|
||||
w.Header().Add(k, v)
|
||||
@@ -1404,7 +1569,7 @@ func (s *QwenService) chatDirectOpenAI(w http.ResponseWriter, r *http.Request, r
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *QwenService) chatGradioRespond(w http.ResponseWriter, r *http.Request, req ChatCompletionRequest, resolvedModel string, maxTokens int) error {
|
||||
func (s *QwenService) chatGradioRespond(endpointURL string, w http.ResponseWriter, r *http.Request, req ChatCompletionRequest, resolvedModel string, maxTokens int) error {
|
||||
var promptText string
|
||||
toolsPrompt := FormatToolsPrompt(req.Tools)
|
||||
|
||||
@@ -1496,7 +1661,7 @@ func (s *QwenService) chatGradioRespond(w http.ResponseWriter, r *http.Request,
|
||||
return fmt.Errorf("failed to encode request: %w", err)
|
||||
}
|
||||
|
||||
callURL := s.endpoint + "/gradio_api/call/respond"
|
||||
callURL := endpointURL + "/gradio_api/call/respond"
|
||||
makeCallReq := func() (*http.Request, error) {
|
||||
reqObj, err := http.NewRequest("POST", callURL, bytes.NewBuffer(jsonPayload))
|
||||
if err != nil {
|
||||
@@ -1521,7 +1686,7 @@ func (s *QwenService) chatGradioRespond(w http.ResponseWriter, r *http.Request,
|
||||
return fmt.Errorf("failed to parse Gradio event ID")
|
||||
}
|
||||
|
||||
streamURL := fmt.Sprintf("%s/gradio_api/call/respond/%s", s.endpoint, joinRes.EventID)
|
||||
streamURL := fmt.Sprintf("%s/gradio_api/call/respond/%s", endpointURL, joinRes.EventID)
|
||||
makeStreamReq := func() (*http.Request, error) {
|
||||
reqObj, err := http.NewRequest("GET", streamURL, nil)
|
||||
if err != nil {
|
||||
@@ -1624,7 +1789,7 @@ func (s *QwenService) chatGradioRespond(w http.ResponseWriter, r *http.Request,
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *QwenService) chatGradioChatResponse(w http.ResponseWriter, r *http.Request, req ChatCompletionRequest, resolvedModel string, maxTokens int) error {
|
||||
func (s *QwenService) chatGradioChatResponse(endpointURL string, w http.ResponseWriter, r *http.Request, req ChatCompletionRequest, resolvedModel string, maxTokens int) error {
|
||||
var systemPromptStr string
|
||||
var historyArray []map[string]interface{}
|
||||
var messageStr string
|
||||
@@ -1791,7 +1956,7 @@ func (s *QwenService) chatGradioChatResponse(w http.ResponseWriter, r *http.Requ
|
||||
return fmt.Errorf("failed to encode request: %w", err)
|
||||
}
|
||||
|
||||
callURL := s.endpoint + "/gradio_api/call/chat_response"
|
||||
callURL := endpointURL + "/gradio_api/call/chat_response"
|
||||
makeCallReq := func() (*http.Request, error) {
|
||||
reqObj, err := http.NewRequest("POST", callURL, bytes.NewBuffer(jsonPayload))
|
||||
if err != nil {
|
||||
@@ -1816,7 +1981,7 @@ func (s *QwenService) chatGradioChatResponse(w http.ResponseWriter, r *http.Requ
|
||||
return fmt.Errorf("failed to parse Gradio event ID")
|
||||
}
|
||||
|
||||
streamURL := fmt.Sprintf("%s/gradio_api/call/chat_response/%s", s.endpoint, joinRes.EventID)
|
||||
streamURL := fmt.Sprintf("%s/gradio_api/call/chat_response/%s", endpointURL, joinRes.EventID)
|
||||
makeStreamReq := func() (*http.Request, error) {
|
||||
reqObj, err := http.NewRequest("GET", streamURL, nil)
|
||||
if err != nil {
|
||||
@@ -1977,11 +2142,18 @@ func (s *QwenService) chatGradioChatResponse(w http.ResponseWriter, r *http.Requ
|
||||
|
||||
func main() {
|
||||
port := flag.Int("port", 8080, "Port to listen on")
|
||||
defaultEndpoint := "https://microhero-qwen3-8-27b-uncensored-chat.hf.space"
|
||||
if envEP := os.Getenv("QFLASH_ENDPOINT"); envEP != "" {
|
||||
defaultEndpoint = envEP
|
||||
defaultEndpointsStr := strings.Join(DefaultEndpoints, ",")
|
||||
if envEP := os.Getenv("QFLASH_ENDPOINTS"); envEP != "" {
|
||||
defaultEndpointsStr = envEP
|
||||
} else if envEP := os.Getenv("QFLASH_ENDPOINT"); envEP != "" {
|
||||
defaultEndpointsStr = envEP
|
||||
}
|
||||
endpoint := flag.String("endpoint", defaultEndpoint, "Upstream Hugging Face Space or OpenAI URL")
|
||||
endpointsFlag := flag.String("endpoints", defaultEndpointsStr, "Comma-separated upstream Hugging Face Spaces or OpenAI URLs")
|
||||
flag.StringVar(endpointsFlag, "endpoint", defaultEndpointsStr, "Alias for -endpoints")
|
||||
|
||||
autoFailover := flag.Bool("failover", true, "Enable automatic failover across endpoints on error or quota limit")
|
||||
flag.BoolVar(autoFailover, "auto-failover", true, "Alias for -failover")
|
||||
|
||||
defaultModelVal := "Qwen/Qwen3.8-27B-Uncensored"
|
||||
if envModel := os.Getenv("QFLASH_MODEL"); envModel != "" {
|
||||
defaultModelVal = envModel
|
||||
@@ -2005,6 +2177,11 @@ func main() {
|
||||
if envMode := os.Getenv("QFLASH_MODE"); envMode != "" && *mode == "auto" {
|
||||
*mode = envMode
|
||||
}
|
||||
if envFailover := os.Getenv("QFLASH_FAILOVER"); envFailover != "" {
|
||||
if v, err := strconv.ParseBool(envFailover); err == nil {
|
||||
*autoFailover = v
|
||||
}
|
||||
}
|
||||
|
||||
if *userAgent != "" {
|
||||
ConfiguredUserAgent = *userAgent
|
||||
@@ -2031,7 +2208,17 @@ func main() {
|
||||
}
|
||||
}
|
||||
|
||||
svc := NewQwenService(*endpoint, *defaultModel, *mode, *hfToken, *apiKey, *baseURL, proxyURL, *thinking)
|
||||
var rawEndpoints []string
|
||||
if *endpointsFlag != "" {
|
||||
for _, ep := range strings.Split(*endpointsFlag, ",") {
|
||||
ep = strings.TrimSpace(ep)
|
||||
if ep != "" {
|
||||
rawEndpoints = append(rawEndpoints, ep)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
svc := NewQwenService(rawEndpoints, *defaultModel, *mode, *hfToken, *apiKey, *baseURL, proxyURL, *thinking, *autoFailover)
|
||||
|
||||
mux := http.NewServeMux()
|
||||
|
||||
@@ -2098,7 +2285,10 @@ func main() {
|
||||
})
|
||||
|
||||
addr := fmt.Sprintf(":%d", *port)
|
||||
log.Printf("Starting qflash gateway on %s -> %s", addr, *endpoint)
|
||||
log.Printf("Starting qflash gateway on %s with %d endpoints (auto-failover: %v)", addr, len(svc.endpoints), *autoFailover)
|
||||
for _, ep := range svc.endpoints {
|
||||
log.Printf(" - Upstream endpoint: %s", ep.URL)
|
||||
}
|
||||
if proxyURL != "" {
|
||||
log.Printf("Routing through SOCKS5 proxy: %s", proxyURL)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user