feat(tools): implement universal tool calling and multi-turn resolution for generic Gradio spaces
This commit is contained in:
+405
@@ -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())
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user