feat(tools): implement universal tool calling and multi-turn resolution for generic Gradio spaces

This commit is contained in:
Luxferre
2026-09-07 08:24:07 +03:00
parent d3be7d4e80
commit de19bab4eb
3 changed files with 825 additions and 182 deletions
+405
View File
@@ -492,3 +492,408 @@ func TestHunyuan3MockServerCompletion(t *testing.T) {
}
}
func TestUniversalToolCallingTransformMessages(t *testing.T) {
req := ChatCompletionRequest{
Tools: []Tool{
{
Type: "function",
Function: map[string]interface{}{
"name": "get_weather",
"description": "Get current weather",
},
},
},
Messages: []ChatMessage{
{Role: "user", Content: "What is the weather in Tokyo and Paris?"},
{
Role: "assistant",
ToolCalls: []ToolCall{
{ID: "call_tokyo", Type: "function", Function: ToolCallFunction{Name: "get_weather", Arguments: `{"city":"Tokyo"}`}},
{ID: "call_paris", Type: "function", Function: ToolCallFunction{Name: "get_weather", Arguments: `{"city":"Paris"}`}},
},
},
{Role: "tool", ToolCallID: "call_tokyo", Content: `{"temp": 20}`},
{Role: "tool", ToolCallID: "call_paris", Content: `{"temp": 15}`},
},
}
processed, toolInstruction, hasSystem := TransformMessages(req)
if !hasSystem {
t.Errorf("expected hasSystem to be true after injecting tool instructions")
}
if toolInstruction == "" {
t.Errorf("expected non-empty toolInstruction")
}
// Expect:
// [0] System message with tool instructions
// [1] User message: "What is the weather in Tokyo and Paris?"
// [2] Assistant message with <tool_call> blocks
// [3] User message with coalesced <tool_response> blocks
if len(processed) != 4 {
t.Fatalf("expected 4 processed messages, got %d", len(processed))
}
if processed[0].Role != "system" || !strings.Contains(processed[0].GetContentString(), "Tool Calling Instructions") {
t.Errorf("unexpected message 0: %+v", processed[0])
}
if processed[1].Role != "user" || processed[1].GetContentString() != "What is the weather in Tokyo and Paris?" {
t.Errorf("unexpected message 1: %+v", processed[1])
}
if processed[2].Role != "assistant" || !strings.Contains(processed[2].GetContentString(), "get_weather") {
t.Errorf("unexpected message 2: %+v", processed[2])
}
respContent := processed[3].GetContentString()
if processed[3].Role != "user" {
t.Errorf("expected coalesced message 3 to have role user, got %q", processed[3].Role)
}
if !strings.Contains(respContent, `{"name": "get_weather", "content": {"temp": 20}}`) {
t.Errorf("expected resolved function name get_weather for tokyo, got:\n%s", respContent)
}
if !strings.Contains(respContent, `{"name": "get_weather", "content": {"temp": 15}}`) {
t.Errorf("expected resolved function name get_weather for paris, got:\n%s", respContent)
}
if !strings.Contains(respContent, "Please answer the user's request based on the tool results.") {
t.Errorf("expected continuation prompt in coalesced message, got:\n%s", respContent)
}
}
func TestUniversalToolCallDetectionVariants(t *testing.T) {
// 1. Array of tool calls inside <tool_calls> tag
multiXML := `<tool_calls>
[
{"name": "get_weather", "arguments": {"city": "Tokyo"}},
{"name": "get_weather", "arguments": {"city": "Paris"}}
]
</tool_calls>`
calls1, rem1, ok1 := DetectToolCalls(multiXML)
if !ok1 || len(calls1) != 2 {
t.Fatalf("expected 2 tool calls from <tool_calls>, got %d", len(calls1))
}
if calls1[0].Function.Name != "get_weather" || calls1[1].Function.Name != "get_weather" {
t.Errorf("unexpected function names: %+v", calls1)
}
if rem1 != "" {
t.Errorf("expected empty remaining, got %q", rem1)
}
// 2. <function_call> tag
fnCallXML := `Some preamble before call.
<function_call>
{"name": "search", "arguments": {"q": "golang"}}
</function_call>
Some postamble.`
calls2, rem2, ok2 := DetectToolCalls(fnCallXML)
if !ok2 || len(calls2) != 1 {
t.Fatalf("expected 1 tool call from <function_call>, got %d", len(calls2))
}
if calls2[0].Function.Name != "search" {
t.Errorf("expected function search, got %q", calls2[0].Function.Name)
}
if strings.Contains(rem2, "function_call") {
t.Errorf("expected tag stripped from remaining, got %q", rem2)
}
if !strings.Contains(rem2, "Some preamble") || !strings.Contains(rem2, "Some postamble") {
t.Errorf("expected surrounding text preserved in remaining, got %q", rem2)
}
// 3. [TOOL_CALLS] bracket syntax
bracketXML := `[TOOL_CALLS]
{"name": "calculate", "arguments": {"x": 42}}
[/TOOL_CALLS]`
calls3, rem3, ok3 := DetectToolCalls(bracketXML)
if !ok3 || len(calls3) != 1 {
t.Fatalf("expected 1 tool call from [TOOL_CALLS], got %d", len(calls3))
}
if calls3[0].Function.Name != "calculate" {
t.Errorf("expected function calculate, got %q", calls3[0].Function.Name)
}
if rem3 != "" {
t.Errorf("expected empty remaining, got %q", rem3)
}
// 4. Raw JSON array without tags
rawArray := `[{"name": "f1", "arguments": {}}, {"name": "f2", "arguments": {}}]`
calls4, rem4, ok4 := DetectToolCalls(rawArray)
if !ok4 || len(calls4) != 2 {
t.Fatalf("expected 2 calls from raw array, got %d", len(calls4))
}
if calls4[0].Function.Name != "f1" || calls4[1].Function.Name != "f2" {
t.Errorf("unexpected names from raw array: %+v", calls4)
}
if rem4 != "" {
t.Errorf("expected empty remaining, got %q", rem4)
}
}
func TestUniversalStreamToolCallFilterVariants(t *testing.T) {
filter := NewStreamToolCallFilter()
var contentParts []string
var toolCalls []ToolCall
onContent := func(s string) { contentParts = append(contentParts, s) }
onToolCall := func(tc ToolCall) { toolCalls = append(toolCalls, tc) }
// Stream using [TOOL_CALLS] across multiple chunk boundaries
chunks := []string{
"Preamble text: ",
"[TOOL_",
"CALLS]\n{\"name\": \"browse\", \"arguments\": {\"url\": \"example.com\"}}\n[/TOOL_",
"CALLS]",
" Completed.",
}
for _, c := range chunks {
filter.Feed(c, onContent, onToolCall)
}
filter.Flush(onContent, onToolCall)
if len(toolCalls) != 1 {
t.Fatalf("expected 1 tool call from stream filter, got %d", len(toolCalls))
}
if toolCalls[0].Function.Name != "browse" {
t.Errorf("expected function browse, got %q", toolCalls[0].Function.Name)
}
if !filter.emittedCall {
t.Errorf("expected emittedCall to be true")
}
fullContent := strings.Join(contentParts, "")
if strings.Contains(fullContent, "TOOL_CALLS") {
t.Errorf("tag leaked into stream content: %q", fullContent)
}
if fullContent != "Preamble text: Completed." {
t.Errorf("unexpected streamed content: %q", fullContent)
}
}
func TestBuildGradioPayloadGenericSpaces(t *testing.T) {
gw := &GradioGateway{}
req := ChatCompletionRequest{
Tools: []Tool{
{Type: "function", Function: map[string]interface{}{"name": "lookup"}},
},
Messages: []ChatMessage{
{Role: "user", Content: "What is 10+10?"},
{
Role: "assistant",
ToolCalls: []ToolCall{
{ID: "c1", Type: "function", Function: ToolCallFunction{Name: "lookup", Arguments: `{"q":"10+10"}`}},
},
},
{Role: "tool", ToolCallID: "c1", Content: `{"result": 20}`},
},
}
// 1. Space with native system prompt input (SystemIndex: 0, MessageIndex: 1, HistoryIndex: 2)
discWithSystem := NewDefaultSpaceDiscovery("https://space-1.hf.space")
discWithSystem.TotalInputs = 3
discWithSystem.SystemIndex = 0
discWithSystem.MessageIndex = 1
discWithSystem.HistoryIndex = 2
discWithSystem.HistoryFormat = "pairs"
data1, err := gw.BuildGradioPayload(discWithSystem, req)
if err != nil {
t.Fatalf("failed to build payload 1: %v", err)
}
sysStr, ok := data1[0].(string)
if !ok || !strings.Contains(sysStr, "Tool Calling Instructions") {
t.Errorf("expected system prompt at index 0, got %v", data1[0])
}
msgStr, ok := data1[1].(string)
if !ok || !strings.Contains(msgStr, "Please answer the user's request based on the tool result.") {
t.Errorf("expected coalesced tool prompt at index 1, got %v", data1[1])
}
pairs1, ok := data1[2].([][]string)
if !ok || len(pairs1) != 1 {
t.Fatalf("expected 1 history pair at index 2, got %T (%v)", data1[2], data1[2])
}
if pairs1[0][0] != "What is 10+10?" || !strings.Contains(pairs1[0][1], "lookup") {
t.Errorf("unexpected history pair: %+v", pairs1[0])
}
// 2. Space without system prompt (SystemIndex: -1, MessageIndex: 0, HistoryIndex: 1)
discNoSystem := NewDefaultSpaceDiscovery("https://space-2.hf.space")
discNoSystem.TotalInputs = 2
discNoSystem.SystemIndex = -1
discNoSystem.MessageIndex = 0
discNoSystem.HistoryIndex = 1
discNoSystem.HistoryFormat = "pairs"
data2, err := gw.BuildGradioPayload(discNoSystem, req)
if err != nil {
t.Fatalf("failed to build payload 2: %v", err)
}
pairs2, ok := data2[1].([][]string)
if !ok || len(pairs2) != 1 {
t.Fatalf("expected 1 history pair at index 1, got %T (%v)", data2[1], data2[1])
}
// Instructions prepended to the first user turn:
if !strings.Contains(pairs2[0][0], "Tool Calling Instructions") || !strings.Contains(pairs2[0][0], "What is 10+10?") {
t.Errorf("expected system instructions prepended to first pair user message, got: %q", pairs2[0][0])
}
if !strings.Contains(pairs2[0][1], "lookup") {
t.Errorf("expected assistant tool call in pair bot turn, got: %q", pairs2[0][1])
}
// 3. Single-textbox space (SystemIndex: -1, MessageIndex: 0, HistoryIndex: -1)
discSingleInput := NewDefaultSpaceDiscovery("https://space-3.hf.space")
discSingleInput.TotalInputs = 1
discSingleInput.SystemIndex = -1
discSingleInput.MessageIndex = 0
discSingleInput.HistoryIndex = -1
data3, err := gw.BuildGradioPayload(discSingleInput, req)
if err != nil {
t.Fatalf("failed to build payload 3: %v", err)
}
transcript, ok := data3[0].(string)
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>") {
t.Errorf("unexpected single-input transcript: %s", transcript)
}
}
func TestGenericSpaceMockServerToolCalling(t *testing.T) {
var lastReceivedData []interface{}
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: "system_prompt"},
{ParameterName: "message"},
{ParameterName: "history"},
},
},
},
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(resp)
return
}
if r.URL.Path == "/gradio_api/call/chat_fn" {
var body map[string]interface{}
json.NewDecoder(r.Body).Decode(&body)
if dataSlice, ok := body["data"].([]interface{}); ok {
lastReceivedData = dataSlice
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(GradioJoinResponse{EventID: "evt_generic"})
return
}
if r.URL.Path == "/gradio_api/call/chat_fn/evt_generic" {
w.Header().Set("Content-Type", "text/event-stream")
flusher, ok := w.(http.Flusher)
if !ok {
t.Fatal("expected flusher")
}
msgStr := ""
if len(lastReceivedData) > 1 {
msgStr, _ = lastReceivedData[1].(string)
}
if strings.Contains(msgStr, "<tool_response>") {
// Turn 2: answer
fmt.Fprintf(w, "event: generating\ndata: [\"The weather in Tokyo is 20 C.\", null]\n\n")
flusher.Flush()
fmt.Fprintf(w, "event: complete\ndata: [\"The weather in Tokyo is 20 C.\", null]\n\n")
flusher.Flush()
} else {
// Turn 1: tool call
fmt.Fprintf(w, "event: generating\ndata: [\"<tool_call>\\n{\\\"name\\\": \\\"get_weather\\\", \\\"arguments\\\": {\\\"city\\\": \\\"Tokyo\\\"}}\\n</tool_call>\", null]\n\n")
flusher.Flush()
fmt.Fprintf(w, "event: complete\ndata: [\"<tool_call>\\n{\\\"name\\\": \\\"get_weather\\\", \\\"arguments\\\": {\\\"city\\\": \\\"Tokyo\\\"}}\\n</tool_call>\", null]\n\n")
flusher.Flush()
}
return
}
http.NotFound(w, r)
}))
defer ts.Close()
gw := NewGradioGateway(ts.URL, "", 10*time.Second)
// Turn 1: User question with tools
req1 := ChatCompletionRequest{
Model: "generic-bot",
Tools: []Tool{
{Type: "function", Function: map[string]interface{}{"name": "get_weather"}},
},
Messages: []ChatMessage{
{Role: "user", Content: "Weather in Tokyo?"},
},
Stream: false,
}
b1, _ := json.Marshal(req1)
httpReq1 := httptest.NewRequest("POST", "/v1/chat/completions", bytes.NewBuffer(b1))
rec1 := httptest.NewRecorder()
err := gw.ExecuteChatCompletion(rec1, httpReq1, req1)
if err != nil {
t.Fatalf("Turn 1 execution failed: %v", err)
}
var resp1 ChatCompletionResponse
if err := json.NewDecoder(rec1.Body).Decode(&resp1); err != nil {
t.Fatalf("Turn 1 decode failed: %v", err)
}
if resp1.Choices[0].FinishReason != "tool_calls" {
t.Fatalf("expected finish_reason 'tool_calls', got %q", resp1.Choices[0].FinishReason)
}
if len(resp1.Choices[0].Message.ToolCalls) != 1 {
t.Fatalf("expected 1 tool call, got %d", len(resp1.Choices[0].Message.ToolCalls))
}
tc := resp1.Choices[0].Message.ToolCalls[0]
if tc.Function.Name != "get_weather" {
t.Fatalf("expected function name get_weather, got %q", tc.Function.Name)
}
// Turn 2: Send tool response
req2 := ChatCompletionRequest{
Model: "generic-bot",
Tools: []Tool{
{Type: "function", Function: map[string]interface{}{"name": "get_weather"}},
},
Messages: []ChatMessage{
{Role: "user", Content: "Weather in Tokyo?"},
resp1.Choices[0].Message,
{Role: "tool", ToolCallID: tc.ID, Content: `{"temp": 20}`},
},
Stream: false,
}
b2, _ := json.Marshal(req2)
httpReq2 := httptest.NewRequest("POST", "/v1/chat/completions", bytes.NewBuffer(b2))
rec2 := httptest.NewRecorder()
err = gw.ExecuteChatCompletion(rec2, httpReq2, req2)
if err != nil {
t.Fatalf("Turn 2 execution failed: %v", err)
}
var resp2 ChatCompletionResponse
if err := json.NewDecoder(rec2.Body).Decode(&resp2); err != nil {
t.Fatalf("Turn 2 decode failed: %v", err)
}
if resp2.Choices[0].FinishReason != "stop" {
t.Errorf("expected finish_reason 'stop', got %q", resp2.Choices[0].FinishReason)
}
if resp2.Choices[0].Message.GetContentString() != "The weather in Tokyo is 20 C." {
t.Errorf("expected final answer, got %q", resp2.Choices[0].Message.GetContentString())
}
}