Support Gradio queue protocol with diff streaming and call fallback

This commit is contained in:
Luxferre
2026-09-07 14:25:06 +03:00
parent defc727c95
commit 7998912fed
2 changed files with 694 additions and 8 deletions
+504 -8
View File
@@ -1205,6 +1205,18 @@ func extractGradioErrorMessage(dataStr string) string {
return clean
}
func isProtocolOrInputError(errMsg string) bool {
lower := strings.ToLower(errMsg)
return strings.Contains(lower, "check gradio inputs/types") ||
strings.Contains(lower, "null error") ||
strings.Contains(lower, "typeerror") ||
strings.Contains(lower, "missing") ||
strings.Contains(lower, "required positional argument") ||
strings.Contains(lower, "internal server error") ||
strings.Contains(lower, "not found") ||
strings.Contains(lower, "unknown upstream gradio error")
}
type Streamer struct {
w http.ResponseWriter
flusher http.Flusher
@@ -1622,6 +1634,7 @@ type SpaceDiscovery struct {
Models []string `json:"models"`
PrimaryModel string `json:"primary_model"`
TotalInputs int `json:"total_inputs"`
RawTotalInputs int `json:"raw_total_inputs,omitempty"`
ParamMappings []SpaceParamMapping `json:"param_mappings"`
DefaultInputs []interface{} `json:"default_inputs"`
MessageIsMultimodal bool `json:"message_is_multimodal"`
@@ -2270,6 +2283,7 @@ func InspectSpace(client *http.Client, rawURL, userAgent string) (*SpaceDiscover
// 6. Correlate with config.dependencies to determine exact input count & state padding
if bestMatchingDep != nil {
discovery.TotalInputs = len(bestMatchingDep.Inputs)
discovery.RawTotalInputs = len(bestMatchingDep.Inputs)
discovery.DefaultInputs = make([]interface{}, len(bestMatchingDep.Inputs))
discovery.ParamMappings = nil
discovery.MessageIndex = -1
@@ -3062,9 +3076,45 @@ type GradioOutputFrame struct {
Content string
Reasoning string
ToolCalls []ToolCall
IsDelta bool
OK bool
}
// extractGradioDiffDelta extracts delta string from Gradio 6 streaming diff operations
func extractGradioDiffDelta(v []interface{}) (string, bool) {
if len(v) == 0 {
return "", false
}
if len(v) >= 3 {
if op, ok := v[0].(string); ok && (op == "append" || op == "add") {
if delta, ok := v[2].(string); ok {
return delta, true
}
}
}
for _, item := range v {
if innerList, ok := item.([]interface{}); ok && len(innerList) > 0 {
for _, opItem := range innerList {
if opArr, ok := opItem.([]interface{}); ok && len(opArr) >= 3 {
if op, ok := opArr[0].(string); ok && (op == "append" || op == "add") {
if delta, ok := opArr[2].(string); ok {
return delta, true
}
}
}
}
if len(innerList) >= 3 {
if op, ok := innerList[0].(string); ok && (op == "append" || op == "add") {
if delta, ok := innerList[2].(string); ok {
return delta, true
}
}
}
}
}
return "", false
}
// ParseGradioStreamOutput extracts structured content, reasoning, and tool calls from Gradio output
func ParseGradioStreamOutput(rawJSON string) (frame GradioOutputFrame) {
defer func() {
@@ -3093,6 +3143,13 @@ func ParseGradioStreamOutput(rawJSON string) (frame GradioOutputFrame) {
return frame
}
if delta, ok := extractGradioDiffDelta(v); ok {
frame.Content = delta
frame.IsDelta = true
frame.OK = true
return frame
}
// 1. Check if v[0] is an inner slice with len >= 3 (e.g. Hy3: [[content, reasoning, tool_calls, history]])
if inner, ok := v[0].([]interface{}); ok && len(inner) >= 3 {
s0, _ := inner[0].(string)
@@ -3344,8 +3401,20 @@ func (g *GradioGateway) executePredictCompletion(w http.ResponseWriter, r *http.
fnIndex = 0
}
predictData := make([]interface{}, len(gradioData))
copy(predictData, gradioData)
if disc.RawTotalInputs > len(predictData) {
for i := len(predictData); i < disc.RawTotalInputs; i++ {
var defVal interface{}
if i < len(disc.DefaultInputs) {
defVal = disc.DefaultInputs[i]
}
predictData = append(predictData, defVal)
}
}
payloadMap := map[string]interface{}{
"data": gradioData,
"data": predictData,
"fn_index": fnIndex,
"session_hash": GenerateUUID(),
}
@@ -3453,6 +3522,417 @@ func (g *GradioGateway) executePredictCompletion(w http.ResponseWriter, r *http.
return nil
}
// executeQueueCompletion handles completions using Gradio queue SSE protocol (/queue/join + /queue/data).
func (g *GradioGateway) executeQueueCompletion(w http.ResponseWriter, r *http.Request, disc *SpaceDiscovery, gradioData []interface{}, req ChatCompletionRequest, completionID string, createdTime int64, modelName, effUA string) error {
fnIndex := disc.FnIndex
if fnIndex < 0 {
fnIndex = 0
}
queueData := make([]interface{}, len(gradioData))
copy(queueData, gradioData)
if disc.RawTotalInputs > len(queueData) {
for i := len(queueData); i < disc.RawTotalInputs; i++ {
var defVal interface{}
if i < len(disc.DefaultInputs) {
defVal = disc.DefaultInputs[i]
}
queueData = append(queueData, defVal)
}
}
sessionHash := GenerateUUID()
payloadMap := map[string]interface{}{
"data": queueData,
"fn_index": fnIndex,
"session_hash": sessionHash,
}
jsonPayload, err := json.Marshal(payloadMap)
if err != nil {
return fmt.Errorf("failed to marshal queue payload: %w", err)
}
var joinURLs []string
if disc.APIPrefix != "" {
joinURLs = append(joinURLs, fmt.Sprintf("%s%s/queue/join", disc.SpaceURL, disc.APIPrefix))
}
joinURLs = append(joinURLs, fmt.Sprintf("%s/queue/join", disc.SpaceURL))
var resp *http.Response
var lastErr error
for _, targetURL := range joinURLs {
curURL := targetURL
makeReq := func() (*http.Request, error) {
req, err := http.NewRequest("POST", curURL, bytes.NewBuffer(jsonPayload))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("User-Agent", effUA)
return req, nil
}
resp, lastErr = DoWithFibonacciRetry(g.client, makeReq, 3)
if lastErr == nil && resp != nil && resp.StatusCode == http.StatusOK {
break
}
if resp != nil {
resp.Body.Close()
resp = nil
}
}
if resp == nil {
log.Printf("Gradio /queue/join unavailable (%v), falling back to /run/predict protocol...", lastErr)
return g.executePredictCompletion(w, r, disc, gradioData, req, completionID, createdTime, modelName, effUA)
}
resp.Body.Close()
// 2. Connect to Gradio SSE queue data stream
var dataURLs []string
if disc.APIPrefix != "" {
dataURLs = append(dataURLs, fmt.Sprintf("%s%s/queue/data?session_hash=%s", disc.SpaceURL, disc.APIPrefix, sessionHash))
}
dataURLs = append(dataURLs, fmt.Sprintf("%s/queue/data?session_hash=%s", disc.SpaceURL, sessionHash))
var streamResp *http.Response
for _, targetURL := range dataURLs {
curURL := targetURL
makeStreamReq := func() (*http.Request, error) {
req, err := http.NewRequest("GET", curURL, nil)
if err != nil {
return nil, err
}
req.Header.Set("Accept", "text/event-stream")
req.Header.Set("User-Agent", effUA)
return req, nil
}
streamResp, lastErr = DoWithFibonacciRetry(g.client, makeStreamReq, 3)
if lastErr == nil && streamResp != nil && streamResp.StatusCode == http.StatusOK {
break
}
if streamResp != nil {
streamResp.Body.Close()
streamResp = nil
}
}
if streamResp == nil {
log.Printf("Gradio /queue/data unavailable (%v), falling back to /run/predict protocol...", lastErr)
return g.executePredictCompletion(w, r, disc, gradioData, req, completionID, createdTime, modelName, effUA)
}
defer streamResp.Body.Close()
type queueMsg struct {
Msg string `json:"msg"`
EventID string `json:"event_id,omitempty"`
Success *bool `json:"success,omitempty"`
Output map[string]interface{} `json:"output,omitempty"`
}
// 3. Handle Non-Streaming vs Streaming
if !req.Stream {
reader := bufio.NewReader(streamResp.Body)
var latestFrame GradioOutputFrame
for {
line, err := reader.ReadString('\n')
if err != nil {
break
}
line = strings.TrimRight(line, "\r\n")
if strings.HasPrefix(line, "data: ") {
dataStr := strings.TrimPrefix(line, "data: ")
var qMsg queueMsg
if err := json.Unmarshal([]byte(dataStr), &qMsg); err != nil {
continue
}
if qMsg.Msg == "close_stream" {
break
}
if qMsg.Msg == "process_generating" || qMsg.Msg == "process_completed" {
if qMsg.Success != nil && !*qMsg.Success {
errBytes, _ := json.Marshal(qMsg.Output)
errMsg := extractGradioErrorMessage(string(errBytes))
log.Printf("Upstream Gradio error: %s", errMsg)
return fmt.Errorf("upstream Gradio error: %s", errMsg)
}
if qMsg.Output != nil && qMsg.Output["data"] != nil {
dataBytes, _ := json.Marshal(qMsg.Output["data"])
frame := ParseGradioStreamOutput(string(dataBytes))
if frame.OK {
if frame.IsDelta {
latestFrame.Content += frame.Content
} else {
if frame.Content != "" || latestFrame.Content == "" {
latestFrame.Content = frame.Content
}
}
if frame.Reasoning != "" || latestFrame.Reasoning == "" {
latestFrame.Reasoning = frame.Reasoning
}
if len(frame.ToolCalls) > 0 {
latestFrame.ToolCalls = frame.ToolCalls
}
latestFrame.OK = true
}
}
if qMsg.Msg == "process_completed" {
break
}
}
}
}
if !latestFrame.OK {
return fmt.Errorf("upstream Gradio space returned empty or unparseable response")
}
cleanText := latestFrame.Content
reasoning := latestFrame.Reasoning
toolCalls := latestFrame.ToolCalls
hasTools := len(toolCalls) > 0
if reasoning == "" {
cleanText, reasoning = ExtractThinking(cleanText)
}
if !hasTools {
toolCalls, cleanText, hasTools = DetectToolCalls(cleanText)
}
finishReason := "stop"
var finalContent interface{} = cleanText
if hasTools && len(toolCalls) > 0 {
finishReason = "tool_calls"
if strings.TrimSpace(cleanText) == "" {
finalContent = nil
}
}
WriteCompletionResponse(w, completionID, createdTime, modelName, FinalOutput{
Content: finalContent,
ReasoningContent: reasoning,
ToolCalls: toolCalls,
FinishReason: finishReason,
})
disc.Protocol = "queue"
return nil
}
// 4. Streaming Mode
flusher, _ := w.(http.Flusher)
streamer := NewStreamer(w, flusher, completionID, createdTime, modelName)
thinkFilter := NewStreamThinkingFilter()
toolFilter := NewStreamToolCallFilter()
reader := bufio.NewReader(streamResp.Body)
var prevContent string
var prevReasoning string
prevToolArgs := make(map[int]string)
nativeReasoningSeen := false
nativeToolCallsSeen := false
var streamErr error
for {
line, err := reader.ReadString('\n')
if err != nil {
break
}
line = strings.TrimRight(line, "\r\n")
if strings.HasPrefix(line, "data: ") {
dataStr := strings.TrimPrefix(line, "data: ")
var qMsg queueMsg
if err := json.Unmarshal([]byte(dataStr), &qMsg); err != nil {
continue
}
if qMsg.Msg == "close_stream" {
break
}
if qMsg.Msg == "process_generating" || qMsg.Msg == "process_completed" {
if qMsg.Success != nil && !*qMsg.Success {
errBytes, _ := json.Marshal(qMsg.Output)
errMsg := extractGradioErrorMessage(string(errBytes))
log.Printf("Upstream Gradio error: %s", errMsg)
if !streamer.started {
return fmt.Errorf("upstream Gradio error: %s", errMsg)
}
streamer.Error(errMsg)
streamErr = fmt.Errorf("upstream Gradio error: %s", errMsg)
break
}
if qMsg.Output != nil && qMsg.Output["data"] != nil {
dataBytes, _ := json.Marshal(qMsg.Output["data"])
frame := ParseGradioStreamOutput(string(dataBytes))
if frame.OK {
// 1. Native reasoning handling
if frame.Reasoning != "" || nativeReasoningSeen {
nativeReasoningSeen = true
if len(frame.Reasoning) > len(prevReasoning) && strings.HasPrefix(frame.Reasoning, prevReasoning) {
delta := frame.Reasoning[len(prevReasoning):]
streamer.Reasoning(delta)
prevReasoning = frame.Reasoning
} else if prevReasoning == "" && frame.Reasoning != "" {
streamer.Reasoning(frame.Reasoning)
prevReasoning = frame.Reasoning
}
}
// 2. Native tool call handling
if len(frame.ToolCalls) > 0 || nativeToolCallsSeen {
nativeToolCallsSeen = true
for i, tc := range frame.ToolCalls {
prevArgs := prevToolArgs[i]
fullArgs := tc.Function.Arguments
if len(fullArgs) > len(prevArgs) && strings.HasPrefix(fullArgs, prevArgs) {
deltaArgs := fullArgs[len(prevArgs):]
iCopy := i
streamer.ToolCallDelta(ToolCall{
Index: &iCopy,
ID: tc.ID,
Type: tc.Type,
Function: ToolCallFunction{
Name: tc.Function.Name,
Arguments: deltaArgs,
},
})
prevToolArgs[i] = fullArgs
} else if prevArgs == "" {
iCopy := i
streamer.ToolCallDelta(ToolCall{
Index: &iCopy,
ID: tc.ID,
Type: tc.Type,
Function: ToolCallFunction{
Name: tc.Function.Name,
Arguments: fullArgs,
},
})
prevToolArgs[i] = fullArgs
}
}
}
// 3. Content handling
currentText := frame.Content
var delta string
if frame.IsDelta {
delta = frame.Content
prevContent += delta
} else {
if strings.HasPrefix(currentText, prevContent) {
delta = currentText[len(prevContent):]
} else if prevContent == "" {
delta = currentText
} else {
delta = currentText
}
prevContent = currentText
}
if delta != "" {
if nativeReasoningSeen || nativeToolCallsSeen {
if nativeReasoningSeen && nativeToolCallsSeen {
streamer.Content(delta)
} else if nativeReasoningSeen {
toolFilter.Feed(delta, func(cleanChunk string) {
if cleanChunk != "" {
streamer.Content(cleanChunk)
}
}, func(tc ToolCall) {
streamer.ToolCallDelta(tc)
})
} else {
thinkFilter.Feed(delta, func(contentChunk string) {
if contentChunk != "" {
streamer.Content(contentChunk)
}
}, func(reasoningChunk string) {
if reasoningChunk != "" {
streamer.Reasoning(reasoningChunk)
}
})
}
} else {
thinkFilter.Feed(delta, func(contentChunk string) {
toolFilter.Feed(contentChunk, func(cleanChunk string) {
if cleanChunk != "" {
streamer.Content(cleanChunk)
}
}, func(tc ToolCall) {
streamer.ToolCallDelta(tc)
})
}, func(reasoningChunk string) {
if reasoningChunk != "" {
streamer.Reasoning(reasoningChunk)
}
})
}
}
}
}
if qMsg.Msg == "process_completed" {
break
}
}
}
}
if streamErr != nil {
return streamErr
}
// Flush remaining tokens in filters if used
if !nativeReasoningSeen {
thinkFilter.Flush(func(contentChunk string) {
if !nativeToolCallsSeen {
toolFilter.Feed(contentChunk, func(cleanChunk string) {
if cleanChunk != "" {
streamer.Content(cleanChunk)
}
}, func(tc ToolCall) {
streamer.ToolCallDelta(tc)
})
} else if contentChunk != "" {
streamer.Content(contentChunk)
}
}, func(reasoningChunk string) {
if reasoningChunk != "" {
streamer.Reasoning(reasoningChunk)
}
})
}
if !nativeToolCallsSeen {
toolFilter.Flush(func(cleanChunk string) {
if cleanChunk != "" {
streamer.Content(cleanChunk)
}
}, func(tc ToolCall) {
streamer.ToolCallDelta(tc)
})
}
finishReason := "stop"
if nativeToolCallsSeen || toolFilter.emittedCall {
finishReason = "tool_calls"
}
streamer.Finish(finishReason)
streamer.Done()
disc.Protocol = "queue"
return nil
}
// ExecuteChatCompletion handles both streaming and non-streaming requests.
func (g *GradioGateway) ExecuteChatCompletion(w http.ResponseWriter, r *http.Request, req ChatCompletionRequest) error {
effUA := EffectiveUserAgent(r)
@@ -3488,12 +3968,18 @@ func (g *GradioGateway) ExecuteChatCompletion(w http.ResponseWriter, r *http.Req
if disc.Protocol == "predict" {
return g.executePredictCompletion(w, r, disc, gradioData, req, completionID, createdTime, modelName, effUA)
}
if disc.Protocol == "queue" {
return g.executeQueueCompletion(w, r, disc, gradioData, req, completionID, createdTime, modelName, effUA)
}
isV2 := disc.Protocol == "call_v2"
buildPayload := func(v2 bool) ([]byte, error) {
if v2 {
v2Payload := make(map[string]interface{})
for idx, val := range gradioData {
if idx < len(disc.ParamMappings) && disc.ParamMappings[idx].ParamType == "state" {
continue
}
pName := ""
if idx < len(disc.ParamMappings) {
if disc.ParamMappings[idx].ParamName != "" {
@@ -3602,9 +4088,9 @@ func (g *GradioGateway) ExecuteChatCompletion(w http.ResponseWriter, r *http.Req
return nil
}
}
if strings.Contains(err.Error(), "404") || strings.Contains(err.Error(), "405") {
log.Printf("Gradio /call endpoint unavailable (%v), falling back to /run/predict protocol...", err)
return g.executePredictCompletion(w, r, disc, gradioData, req, completionID, createdTime, modelName, effUA)
if strings.Contains(err.Error(), "404") || strings.Contains(err.Error(), "405") || strings.Contains(err.Error(), "500") {
log.Printf("Gradio /call endpoint unavailable (%v), falling back to /queue/join protocol...", err)
return g.executeQueueCompletion(w, r, disc, gradioData, req, completionID, createdTime, modelName, effUA)
}
return fmt.Errorf("upstream Gradio call error: %w", err)
}
@@ -3613,8 +4099,8 @@ func (g *GradioGateway) ExecuteChatCompletion(w http.ResponseWriter, r *http.Req
var joinRes GradioJoinResponse
if err := json.NewDecoder(resp.Body).Decode(&joinRes); err != nil || joinRes.EventID == "" {
log.Printf("Gradio /call returned non-SSE response, falling back to /run/predict protocol...")
return g.executePredictCompletion(w, r, disc, gradioData, req, completionID, createdTime, modelName, effUA)
log.Printf("Gradio /call returned non-SSE response, falling back to /queue/join protocol...")
return g.executeQueueCompletion(w, r, disc, gradioData, req, completionID, createdTime, modelName, effUA)
}
// 2. Connect to Gradio SSE EventStream
@@ -3652,7 +4138,8 @@ func (g *GradioGateway) ExecuteChatCompletion(w http.ResponseWriter, r *http.Req
return nil
}
}
return fmt.Errorf("upstream Gradio stream error: %w", err)
log.Printf("Gradio /call stream connection failed (%v), falling back to /queue/join protocol...", err)
return g.executeQueueCompletion(w, r, disc, gradioData, req, completionID, createdTime, modelName, effUA)
}
defer streamResp.Body.Close()
@@ -3687,6 +4174,10 @@ func (g *GradioGateway) ExecuteChatCompletion(w http.ResponseWriter, r *http.Req
}
}
errMsg := extractGradioErrorMessage(dataStr)
if isProtocolOrInputError(errMsg) {
log.Printf("Upstream Gradio /call stream error (%s), falling back to /queue/join protocol...", errMsg)
return g.executeQueueCompletion(w, r, disc, gradioData, req, completionID, createdTime, modelName, effUA)
}
log.Printf("Upstream Gradio error: %s", errMsg)
return fmt.Errorf("upstream Gradio error: %s", errMsg)
}
@@ -3785,10 +4276,15 @@ func (g *GradioGateway) ExecuteChatCompletion(w http.ResponseWriter, r *http.Req
}
}
errMsg := extractGradioErrorMessage(dataStr)
log.Printf("Upstream Gradio error: %s", errMsg)
if !streamer.started {
if isProtocolOrInputError(errMsg) {
log.Printf("Upstream Gradio /call stream error (%s), falling back to /queue/join protocol...", errMsg)
return g.executeQueueCompletion(w, r, disc, gradioData, req, completionID, createdTime, modelName, effUA)
}
log.Printf("Upstream Gradio error: %s", errMsg)
return fmt.Errorf("upstream Gradio error: %s", errMsg)
}
log.Printf("Upstream Gradio error: %s", errMsg)
streamer.Error(errMsg)
streamErr = fmt.Errorf("upstream Gradio error: %s", errMsg)
break