fix(tools): eliminate zero-token suppression and propagate upstream errors cleanly
This commit is contained in:
@@ -438,7 +438,23 @@ func BuildToolInstruction(tools []Tool) string {
|
||||
return ""
|
||||
}
|
||||
toolsBytes, _ := json.MarshalIndent(tools, "", " ")
|
||||
return fmt.Sprintf("\n\n# Tool Calling Instructions\n\nYou have access to the following functions:\n<tools>\n%s\n</tools>\n\nWhen you need to call a function, respond ONLY with a <tool_call> block formatted exactly as follows:\n<tool_call>\n{\"name\": \"<function-name>\", \"arguments\": {<args-json-object>}}\n</tool_call>\n\nWhen you receive a <tool_response>, use the provided information to answer the user's request, or call further tools if needed.\nDo not include conversational filler before or after the tool call.", string(toolsBytes))
|
||||
return fmt.Sprintf(`# Tool Calling Instructions
|
||||
|
||||
You have access to the following tools:
|
||||
<tools>
|
||||
%s
|
||||
</tools>
|
||||
|
||||
To call a tool, you MUST output a <tool_call> block directly in your text response formatted exactly as follows:
|
||||
<tool_call>
|
||||
{"name": "<function-name>", "arguments": {<args-json-object>}}
|
||||
</tool_call>
|
||||
|
||||
Rules:
|
||||
- If you need to call a tool, respond ONLY with the <tool_call> block. Do not include introductory text, explanations, or commentary around the block.
|
||||
- If you need to call multiple tools, provide each tool call in its own <tool_call> block.
|
||||
- If no tool call is needed, answer the user's request directly and normally without using tool tags.
|
||||
- When you receive a <tool_response>, answer the user's request using the information provided in the response, or call another tool if additional information is required.`, string(toolsBytes))
|
||||
}
|
||||
|
||||
func TransformMessages(req ChatCompletionRequest) (processed []ChatMessage, toolInstruction string, hasSystem bool) {
|
||||
@@ -986,48 +1002,105 @@ func WriteCompletionResponse(w http.ResponseWriter, completionID string, created
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
}
|
||||
|
||||
func extractGradioErrorMessage(dataStr string) string {
|
||||
var errObj map[string]interface{}
|
||||
if err := json.Unmarshal([]byte(dataStr), &errObj); err == nil {
|
||||
if e, ok := errObj["error"].(string); ok && e != "" {
|
||||
return e
|
||||
}
|
||||
if m, ok := errObj["message"].(string); ok && m != "" {
|
||||
return m
|
||||
}
|
||||
if eNull, ok := errObj["error"]; ok && eNull == nil {
|
||||
if t, ok := errObj["title"].(string); ok && t != "" {
|
||||
return t
|
||||
}
|
||||
return "internal space error (check Gradio inputs/types)"
|
||||
}
|
||||
}
|
||||
clean := strings.TrimSpace(dataStr)
|
||||
if clean != "" && clean != "null" {
|
||||
return clean
|
||||
}
|
||||
return "unknown upstream Gradio error"
|
||||
}
|
||||
|
||||
type Streamer struct {
|
||||
w http.ResponseWriter
|
||||
flusher http.Flusher
|
||||
id string
|
||||
created int64
|
||||
model string
|
||||
started bool
|
||||
}
|
||||
|
||||
func NewStreamer(w http.ResponseWriter, flusher http.Flusher, id string, created int64, model string) *Streamer {
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
w.Header().Set("Cache-Control", "no-cache")
|
||||
w.Header().Set("Connection", "keep-alive")
|
||||
return &Streamer{w: w, flusher: flusher, id: id, created: created, model: model}
|
||||
}
|
||||
|
||||
func (s *Streamer) EnsureStarted() {
|
||||
if s.started {
|
||||
return
|
||||
}
|
||||
s.w.Header().Set("Content-Type", "text/event-stream")
|
||||
s.w.Header().Set("Cache-Control", "no-cache")
|
||||
s.w.Header().Set("Connection", "keep-alive")
|
||||
s.started = true
|
||||
s.Role()
|
||||
}
|
||||
|
||||
func (s *Streamer) Role() {
|
||||
sendStreamDelta(s.w, s.flusher, s.id, s.created, s.model, StreamDelta{Role: "assistant"})
|
||||
}
|
||||
|
||||
func (s *Streamer) Reasoning(text string) {
|
||||
s.EnsureStarted()
|
||||
sendStreamDelta(s.w, s.flusher, s.id, s.created, s.model, StreamDelta{ReasoningContent: text})
|
||||
}
|
||||
|
||||
func (s *Streamer) Content(text string) {
|
||||
s.EnsureStarted()
|
||||
sendStreamDelta(s.w, s.flusher, s.id, s.created, s.model, StreamDelta{Content: text})
|
||||
}
|
||||
|
||||
func (s *Streamer) ToolCallDelta(tc ToolCall) {
|
||||
s.EnsureStarted()
|
||||
sendStreamDelta(s.w, s.flusher, s.id, s.created, s.model, StreamDelta{ToolCalls: []ToolCall{tc}})
|
||||
}
|
||||
|
||||
func (s *Streamer) Finish(reason string) {
|
||||
if !s.started {
|
||||
return
|
||||
}
|
||||
sendStreamChunk(s.w, s.flusher, s.id, s.created, s.model, StreamDelta{}, &reason)
|
||||
}
|
||||
|
||||
func (s *Streamer) Done() {
|
||||
if !s.started {
|
||||
return
|
||||
}
|
||||
fmt.Fprintf(s.w, "data: [DONE]\n\n")
|
||||
if s.flusher != nil {
|
||||
s.flusher.Flush()
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Streamer) Error(errMsg string) {
|
||||
s.EnsureStarted()
|
||||
errChunk := map[string]interface{}{
|
||||
"error": map[string]interface{}{
|
||||
"message": errMsg,
|
||||
"type": "upstream_error",
|
||||
"code": 502,
|
||||
},
|
||||
}
|
||||
b, _ := json.Marshal(errChunk)
|
||||
fmt.Fprintf(s.w, "data: %s\n\n", b)
|
||||
if s.flusher != nil {
|
||||
s.flusher.Flush()
|
||||
}
|
||||
}
|
||||
|
||||
func sendStreamDelta(w http.ResponseWriter, flusher http.Flusher, completionID string, createdTime int64, modelName string, delta StreamDelta) {
|
||||
sendStreamChunk(w, flusher, completionID, createdTime, modelName, delta, nil)
|
||||
}
|
||||
@@ -1948,13 +2021,20 @@ func (g *GradioGateway) BuildGradioPayload(disc *SpaceDiscovery, req ChatComplet
|
||||
promptMessageText = lastUserMessage
|
||||
} else {
|
||||
// Single message space: compose multi-turn history into the prompt
|
||||
if len(transformed) <= 1 && systemPromptStr == "" {
|
||||
if len(nonSystem) <= 1 {
|
||||
if systemPromptStr != "" && len(nonSystem) == 1 {
|
||||
promptMessageText = systemPromptStr + "\n\n" + lastUserMessage
|
||||
} else if systemPromptStr != "" {
|
||||
promptMessageText = systemPromptStr
|
||||
} else {
|
||||
promptMessageText = lastUserMessage
|
||||
}
|
||||
} else {
|
||||
var sb strings.Builder
|
||||
if systemPromptStr != "" {
|
||||
sb.WriteString("System: " + systemPromptStr + "\n\n")
|
||||
sb.WriteString("# Instructions\n" + systemPromptStr + "\n\n")
|
||||
}
|
||||
sb.WriteString("# Conversation History\n")
|
||||
for i := 0; i < len(nonSystem)-1; i++ {
|
||||
m := nonSystem[i]
|
||||
roleLabel := "User"
|
||||
@@ -1967,11 +2047,7 @@ func (g *GradioGateway) BuildGradioPayload(disc *SpaceDiscovery, req ChatComplet
|
||||
if len(nonSystem) > 0 && nonSystem[len(nonSystem)-1].Role == "assistant" {
|
||||
lastRoleLabel = "Assistant"
|
||||
}
|
||||
if sb.Len() > 0 {
|
||||
sb.WriteString(fmt.Sprintf("%s: %s", lastRoleLabel, lastUserMessage))
|
||||
} else {
|
||||
sb.WriteString(lastUserMessage)
|
||||
}
|
||||
sb.WriteString(fmt.Sprintf("# Current Request\n%s: %s", lastRoleLabel, lastUserMessage))
|
||||
promptMessageText = sb.String()
|
||||
}
|
||||
}
|
||||
@@ -2342,7 +2418,9 @@ func (g *GradioGateway) ExecuteChatCompletion(w http.ResponseWriter, r *http.Req
|
||||
if strings.HasPrefix(line, "data: ") {
|
||||
dataStr := strings.TrimPrefix(line, "data: ")
|
||||
if currentEvent == "error" {
|
||||
return fmt.Errorf("gradio stream error: %s", dataStr)
|
||||
errMsg := extractGradioErrorMessage(dataStr)
|
||||
log.Printf("Upstream Gradio error: %s", errMsg)
|
||||
return fmt.Errorf("upstream Gradio error: %s", errMsg)
|
||||
}
|
||||
if frame := ParseGradioStreamOutput(dataStr); frame.OK {
|
||||
latestFrame = frame
|
||||
@@ -2353,6 +2431,10 @@ func (g *GradioGateway) ExecuteChatCompletion(w http.ResponseWriter, r *http.Req
|
||||
}
|
||||
}
|
||||
|
||||
if !latestFrame.OK {
|
||||
return fmt.Errorf("upstream Gradio space returned empty or unparseable response")
|
||||
}
|
||||
|
||||
cleanText := latestFrame.Content
|
||||
reasoning := latestFrame.Reasoning
|
||||
toolCalls := latestFrame.ToolCalls
|
||||
@@ -2386,7 +2468,6 @@ func (g *GradioGateway) ExecuteChatCompletion(w http.ResponseWriter, r *http.Req
|
||||
// 4. Streaming Mode
|
||||
flusher, _ := w.(http.Flusher)
|
||||
streamer := NewStreamer(w, flusher, completionID, createdTime, modelName)
|
||||
streamer.Role()
|
||||
|
||||
thinkFilter := NewStreamThinkingFilter()
|
||||
toolFilter := NewStreamToolCallFilter()
|
||||
@@ -2398,6 +2479,7 @@ func (g *GradioGateway) ExecuteChatCompletion(w http.ResponseWriter, r *http.Req
|
||||
nativeReasoningSeen := false
|
||||
nativeToolCallsSeen := false
|
||||
currentEvent := ""
|
||||
var streamErr error
|
||||
|
||||
for {
|
||||
line, err := reader.ReadString('\n')
|
||||
@@ -2414,6 +2496,13 @@ func (g *GradioGateway) ExecuteChatCompletion(w http.ResponseWriter, r *http.Req
|
||||
if strings.HasPrefix(line, "data: ") {
|
||||
dataStr := strings.TrimPrefix(line, "data: ")
|
||||
if currentEvent == "error" {
|
||||
errMsg := extractGradioErrorMessage(dataStr)
|
||||
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
|
||||
}
|
||||
|
||||
@@ -2535,6 +2624,10 @@ func (g *GradioGateway) ExecuteChatCompletion(w http.ResponseWriter, r *http.Req
|
||||
}
|
||||
}
|
||||
|
||||
if streamErr != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Flush remaining tokens in filters if used
|
||||
if !nativeReasoningSeen {
|
||||
thinkFilter.Flush(func(contentChunk string) {
|
||||
@@ -2566,6 +2659,10 @@ func (g *GradioGateway) ExecuteChatCompletion(w http.ResponseWriter, r *http.Req
|
||||
})
|
||||
}
|
||||
|
||||
if !streamer.started {
|
||||
return fmt.Errorf("upstream Gradio space closed stream without sending content")
|
||||
}
|
||||
|
||||
if nativeToolCallsSeen || toolFilter.emittedCall {
|
||||
streamer.Finish("tool_calls")
|
||||
} else {
|
||||
@@ -2684,7 +2781,15 @@ func main() {
|
||||
|
||||
if err := gateway.ExecuteChatCompletion(w, r, req); err != nil {
|
||||
log.Printf("Chat completion error: %v", err)
|
||||
http.Error(w, fmt.Sprintf("Gateway error: %v", err), http.StatusBadGateway)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusBadGateway)
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"error": map[string]interface{}{
|
||||
"message": err.Error(),
|
||||
"type": "upstream_error",
|
||||
"code": http.StatusBadGateway,
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
+129
-1
@@ -755,7 +755,7 @@ func TestBuildGradioPayloadGenericSpaces(t *testing.T) {
|
||||
if !ok {
|
||||
t.Fatalf("expected string transcript, got %T", data3[0])
|
||||
}
|
||||
if !strings.Contains(transcript, "System: ") || !strings.Contains(transcript, "User: What is 10+10?") || !strings.Contains(transcript, "Assistant: <tool_call>") {
|
||||
if !strings.Contains(transcript, "# Instructions") || !strings.Contains(transcript, "User: What is 10+10?") || !strings.Contains(transcript, "Assistant: <tool_call>") || !strings.Contains(transcript, "# Current Request") {
|
||||
t.Errorf("unexpected single-input transcript: %s", transcript)
|
||||
}
|
||||
}
|
||||
@@ -896,4 +896,132 @@ func TestGenericSpaceMockServerToolCalling(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractGradioErrorMessage(t *testing.T) {
|
||||
cases := []struct {
|
||||
input string
|
||||
expected string
|
||||
}{
|
||||
{
|
||||
input: `{"error": "Client error '402 Payment Required'"}`,
|
||||
expected: "Client error '402 Payment Required'",
|
||||
},
|
||||
{
|
||||
input: `{"message": "Rate limit exceeded"}`,
|
||||
expected: "Rate limit exceeded",
|
||||
},
|
||||
{
|
||||
input: `{"error": null, "title": "Validation Error"}`,
|
||||
expected: "Validation Error",
|
||||
},
|
||||
{
|
||||
input: `{"error": null}`,
|
||||
expected: "internal space error (check Gradio inputs/types)",
|
||||
},
|
||||
{
|
||||
input: `raw server failure`,
|
||||
expected: "raw server failure",
|
||||
},
|
||||
}
|
||||
|
||||
for _, c := range cases {
|
||||
got := extractGradioErrorMessage(c.input)
|
||||
if got != c.expected {
|
||||
t.Errorf("extractGradioErrorMessage(%q) = %q, expected %q", c.input, got, c.expected)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestStreamerLazyStartAndErrors(t *testing.T) {
|
||||
rec := httptest.NewRecorder()
|
||||
s := NewStreamer(rec, nil, "cmpl-1", 12345, "test-model")
|
||||
if s.started {
|
||||
t.Errorf("expected streamer to start as not started")
|
||||
}
|
||||
|
||||
// Ensure Role/headers only sent on first write
|
||||
s.Content("Hello")
|
||||
if !s.started {
|
||||
t.Errorf("expected streamer to be started after Content")
|
||||
}
|
||||
s.Finish("stop")
|
||||
s.Done()
|
||||
|
||||
body := rec.Body.String()
|
||||
if !strings.Contains(body, `"role":"assistant"`) {
|
||||
t.Errorf("expected role chunk in body: %s", body)
|
||||
}
|
||||
if !strings.Contains(body, `"content":"Hello"`) {
|
||||
t.Errorf("expected content chunk in body: %s", body)
|
||||
}
|
||||
if !strings.Contains(body, "[DONE]") {
|
||||
t.Errorf("expected [DONE] in body: %s", body)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpstreamGradioErrorPropagation(t *testing.T) {
|
||||
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == "/gradio_api/info" {
|
||||
resp := GradioAPIInfoResponse{
|
||||
NamedEndpoints: map[string]GradioEndpointInfo{
|
||||
"/chat_fn": {
|
||||
Parameters: []GradioParamInfo{
|
||||
{ParameterName: "message", Component: "Textbox"},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
return
|
||||
}
|
||||
if r.URL.Path == "/gradio_api/call/chat_fn" {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(GradioJoinResponse{EventID: "err-event"})
|
||||
return
|
||||
}
|
||||
if r.URL.Path == "/gradio_api/call/chat_fn/err-event" {
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
flusher := w.(http.Flusher)
|
||||
fmt.Fprintf(w, "event: error\ndata: {\"error\": \"Quota exceeded: 402 Payment Required\"}\n\n")
|
||||
flusher.Flush()
|
||||
return
|
||||
}
|
||||
http.NotFound(w, r)
|
||||
}))
|
||||
defer ts.Close()
|
||||
|
||||
gw := NewGradioGateway(ts.URL, "", 10*time.Second)
|
||||
|
||||
// 1. Non-streaming error should return error with upstream message
|
||||
recNonStream := httptest.NewRecorder()
|
||||
httpReqNonStream := httptest.NewRequest("POST", "/v1/chat/completions", nil)
|
||||
err := gw.ExecuteChatCompletion(recNonStream, httpReqNonStream, ChatCompletionRequest{
|
||||
Messages: []ChatMessage{{Role: "user", Content: "Hello"}},
|
||||
Stream: false,
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatalf("expected error from non-streaming upstream failure, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "402 Payment Required") {
|
||||
t.Errorf("expected error to mention 402 Payment Required, got %v", err)
|
||||
}
|
||||
|
||||
// 2. Streaming error before start should return error with upstream message
|
||||
recStream := httptest.NewRecorder()
|
||||
httpReqStream := httptest.NewRequest("POST", "/v1/chat/completions", nil)
|
||||
err = gw.ExecuteChatCompletion(recStream, httpReqStream, ChatCompletionRequest{
|
||||
Messages: []ChatMessage{{Role: "user", Content: "Hello"}},
|
||||
Stream: true,
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatalf("expected error from streaming upstream failure, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "402 Payment Required") {
|
||||
t.Errorf("expected error to mention 402 Payment Required, got %v", err)
|
||||
}
|
||||
// And nothing should have been written to body
|
||||
if strings.Contains(recStream.Body.String(), "[DONE]") {
|
||||
t.Errorf("did not expect [DONE] on upstream error: %s", recStream.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user