diff --git a/README.md b/README.md index 3f1fc7c..1b65327 100644 --- a/README.md +++ b/README.md @@ -171,7 +171,7 @@ curl -N http://localhost:8080/v1/chat/completions \ }' ``` -### Tool calling +### Tool calling (turn 1) ```bash curl http://localhost:8080/v1/chat/completions \ @@ -233,6 +233,76 @@ Response: } ``` +### Tool response submission (turn 2) + +```bash +curl http://localhost:8080/v1/chat/completions \ + -H "Content-Type: application/json" \ + -d '{ + "model": "openai/gpt-oss-120b", + "messages": [ + {"role": "user", "content": "What is the weather in Tokyo?"}, + { + "role": "assistant", + "tool_calls": [ + { + "id": "call_3d4c016a", + "type": "function", + "function": { + "name": "get_weather", + "arguments": "{\"location\":\"Tokyo\"}" + } + } + ] + }, + { + "role": "tool", + "tool_call_id": "call_3d4c016a", + "content": "{\"temperature\": 20, \"condition\": \"sunny\"}" + } + ], + "tools": [{ + "type": "function", + "function": { + "name": "get_weather", + "description": "Get current weather in a location", + "parameters": { + "type": "object", + "properties": { + "location": {"type": "string", "description": "City name"} + }, + "required": ["location"] + } + } + }] + }' +``` + +Response: +```json +{ + "id": "chatcmpl-4903ba12-f12b-4cd3-a801-7290bc91a421", + "object": "chat.completion", + "created": 1788756335, + "model": "openai/gpt-oss-120b", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "The current weather in Tokyo is sunny with a temperature of 20 °C." + }, + "finish_reason": "stop" + } + ], + "usage": { + "prompt_tokens": 0, + "completion_tokens": 0, + "total_tokens": 0 + } +} +``` + ### Dynamic target space override Override the target space per request without restarting the server: diff --git a/gr2gw.go b/gr2gw.go index 4a8aff6..5cd7a9a 100644 --- a/gr2gw.go +++ b/gr2gw.go @@ -438,18 +438,33 @@ 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\n%s\n\n\nWhen you need to call a function, respond ONLY with a block formatted exactly as follows:\n\n{\"name\": \"\", \"arguments\": {}}\n\n\nDo not include conversational filler before or after the tool call.", string(toolsBytes)) + return fmt.Sprintf("\n\n# Tool Calling Instructions\n\nYou have access to the following functions:\n\n%s\n\n\nWhen you need to call a function, respond ONLY with a block formatted exactly as follows:\n\n{\"name\": \"\", \"arguments\": {}}\n\n\nWhen you receive a , 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)) } func TransformMessages(req ChatCompletionRequest) (processed []ChatMessage, toolInstruction string, hasSystem bool) { - toolInstruction = BuildToolInstruction(req.Tools) + // 1. Build lookup map from tool_call_id to function name across all assistant messages + toolIDToName := make(map[string]string) for _, msg := range req.Messages { + for _, tc := range msg.ToolCalls { + if tc.ID != "" && tc.Function.Name != "" { + toolIDToName[tc.ID] = tc.Function.Name + } + } + } + + toolInstruction = BuildToolInstruction(req.Tools) + + // 2. Process and coalesce messages preserving turn parity + var staged []ChatMessage + for i := 0; i < len(req.Messages); i++ { + msg := req.Messages[i] contentStr := msg.GetContentString() - m := ChatMessage{Role: msg.Role, Content: contentStr} + switch msg.Role { case "system": hasSystem = true - m.Content = contentStr + staged = append(staged, ChatMessage{Role: "system", Content: contentStr}) + case "assistant": var sb strings.Builder if contentStr != "" { @@ -465,40 +480,79 @@ func TransformMessages(req ChatCompletionRequest) (processed []ChatMessage, tool } sb.WriteString(fmt.Sprintf("\n{\"name\": %q, \"arguments\": %s}\n", tc.Function.Name, args)) } - m.Content = sb.String() + staged = append(staged, ChatMessage{ + Role: "assistant", + Content: sb.String(), + ReasoningContent: msg.ReasoningContent, + ToolCalls: msg.ToolCalls, + }) + case "tool", "function": - m.Role = "user" - toolName := msg.Name - if toolName == "" { - toolName = msg.ToolCallID + // Gather consecutive tool returns into a coalesced turn + var toolResponses []string + j := i + for j < len(req.Messages) && (req.Messages[j].Role == "tool" || req.Messages[j].Role == "function") { + tMsg := req.Messages[j] + tContent := tMsg.GetContentString() + tName := tMsg.Name + if tName == "" && tMsg.ToolCallID != "" { + if mapped, ok := toolIDToName[tMsg.ToolCallID]; ok { + tName = mapped + } else { + tName = tMsg.ToolCallID + } + } + var contentJSON []byte + if json.Valid([]byte(tContent)) { + contentJSON = []byte(tContent) + } else { + contentJSON, _ = json.Marshal(tContent) + } + toolResponses = append(toolResponses, fmt.Sprintf("\n{\"name\": %q, \"content\": %s}\n", tName, string(contentJSON))) + j++ } - var contentJSON []byte - if json.Valid([]byte(contentStr)) { - contentJSON = []byte(contentStr) - } else { - contentJSON, _ = json.Marshal(contentStr) + i = j - 1 // advance loop + + promptSuffix := "Please answer the user's request based on the tool result." + if len(toolResponses) > 1 { + promptSuffix = "Please answer the user's request based on the tool results." } - m.Content = fmt.Sprintf("\n{\"name\": %q, \"content\": %s}\n", toolName, string(contentJSON)) + coalesced := strings.Join(toolResponses, "\n") + "\n\n" + promptSuffix + staged = append(staged, ChatMessage{ + Role: "user", + Content: coalesced, + }) + + default: // "user" or other roles + staged = append(staged, ChatMessage{Role: msg.Role, Content: contentStr}) } - processed = append(processed, m) } + // 3. Inject tool instructions into system prompt if toolInstruction != "" { if hasSystem { - for i, m := range processed { + systemInjected := false + for i, m := range staged { if m.Role == "system" { - processed[i].Content = m.GetContentString() + "\n" + strings.TrimSpace(toolInstruction) + staged[i].Content = m.GetContentString() + "\n\n" + strings.TrimSpace(toolInstruction) + systemInjected = true break } } + if !systemInjected { + staged = append([]ChatMessage{ + {Role: "system", Content: strings.TrimSpace(toolInstruction)}, + }, staged...) + } } else { - processed = append([]ChatMessage{ - {Role: "user", Content: strings.TrimSpace(toolInstruction)}, - }, processed...) + staged = append([]ChatMessage{ + {Role: "system", Content: strings.TrimSpace(toolInstruction)}, + }, staged...) + hasSystem = true } } - return processed, toolInstruction, hasSystem + return staged, toolInstruction, hasSystem } func cleanJSONBlock(input string) string { @@ -538,6 +592,35 @@ func sanitizeJSONValue(v interface{}) interface{} { } } +type ToolTagPair struct { + Start string + End string +} + +var ToolTagPairs = []ToolTagPair{ + {Start: "", End: ""}, + {Start: "", End: ""}, + {Start: "", End: ""}, + {Start: "[TOOL_CALLS]", End: "[/TOOL_CALLS]"}, +} + +func getToolStartPrefixes() []string { + seen := make(map[string]bool) + var prefixes []string + for _, pair := range ToolTagPairs { + for i := 1; i <= len(pair.Start); i++ { + pref := pair.Start[:i] + if !seen[pref] { + seen[pref] = true + prefixes = append(prefixes, pref) + } + } + } + return prefixes +} + +var toolStartPrefixes = getToolStartPrefixes() + func repairToolCallJSON(input string) (ToolCall, bool) { s := strings.TrimSpace(input) reName := regexp.MustCompile(`"(?:name|function|action|call)"\s*:\s*"([^"]+)"`) @@ -547,8 +630,8 @@ func repairToolCallJSON(input string) (ToolCall, bool) { } fnName := matches[1] - reArgs := regexp.MustCompile(`"(?:arguments|parameters|args|input)"\s*:\s*(\{[\s\S]*\})`) - argMatches := reArgs.FindStringSubmatch(s) + reArgsObj := regexp.MustCompile(`"(?:arguments|parameters|args|input)"\s*:\s*(\{[\s\S]*\})`) + argMatches := reArgsObj.FindStringSubmatch(s) argsStr := "{}" if len(argMatches) >= 2 { candidate := argMatches[1] @@ -556,6 +639,15 @@ func repairToolCallJSON(input string) (ToolCall, bool) { if json.Unmarshal([]byte(candidate), &dummy) == nil { argsStr = candidate } + } else { + reArgsStr := regexp.MustCompile(`"(?:arguments|parameters|args|input)"\s*:\s*"((?:\\.|[^"\\])*)"`) + strMatches := reArgsStr.FindStringSubmatch(s) + if len(strMatches) >= 2 { + var unescaped string + if json.Unmarshal([]byte(`"`+strMatches[1]+`"`), &unescaped) == nil { + argsStr = unescaped + } + } } return ToolCall{ @@ -671,18 +763,68 @@ func parseSingleToolCall(jsonStr string) (ToolCall, bool) { return repairToolCallJSON(cleaned) } -func parseXMLToolCall(block string) (ToolCall, bool) { - inner := strings.TrimSpace(block) - if strings.HasPrefix(inner, "") { - inner = strings.TrimPrefix(inner, "") +func parseMultipleToolCalls(raw string) ([]ToolCall, bool) { + cleaned := cleanJSONBlock(raw) + if cleaned == "" { + return nil, false } - if strings.HasSuffix(inner, "") { - inner = strings.TrimSuffix(inner, "") + + // 1. Direct JSON array: [{"name":...}, ...] + var rawList []interface{} + if err := json.Unmarshal([]byte(cleaned), &rawList); err == nil { + var calls []ToolCall + for _, item := range rawList { + b, err := json.Marshal(item) + if err == nil { + if tc, ok := parseSingleToolCall(string(b)); ok { + calls = append(calls, tc) + } + } + } + if len(calls) > 0 { + return calls, true + } + } + + // 2. Wrapper object with tool_calls / calls array + var rawMap map[string]interface{} + if err := json.Unmarshal([]byte(cleaned), &rawMap); err == nil { + for _, listKey := range []string{"tool_calls", "calls", "functions"} { + if subArr, ok := rawMap[listKey].([]interface{}); ok && len(subArr) > 0 { + var calls []ToolCall + for _, item := range subArr { + b, err := json.Marshal(item) + if err == nil { + if tc, ok := parseSingleToolCall(string(b)); ok { + calls = append(calls, tc) + } + } + } + if len(calls) > 0 { + return calls, true + } + } + } + } + + // 3. Single tool call + if tc, ok := parseSingleToolCall(cleaned); ok { + return []ToolCall{tc}, true + } + + return nil, false +} + +func parseXMLToolCall(block string) ([]ToolCall, bool) { + inner := strings.TrimSpace(block) + for _, pair := range ToolTagPairs { + inner = strings.ReplaceAll(inner, pair.Start, "") + inner = strings.ReplaceAll(inner, pair.End, "") } inner = cleanJSONBlock(inner) - if tc, ok := parseSingleToolCall(inner); ok { - return tc, true + if calls, ok := parseMultipleToolCalls(inner); ok && len(calls) > 0 { + return calls, true } var fnName string @@ -707,88 +849,58 @@ func parseXMLToolCall(block string) (ToolCall, bool) { if argsStr == "" { argsStr = "{}" } - return ToolCall{ + return []ToolCall{{ ID: "call_" + GenerateUUID()[:8], Type: "function", Function: ToolCallFunction{ Name: fnName, Arguments: argsStr, }, - }, true + }}, true } - return ToolCall{}, false + return nil, false } func ExtractToolCallBlocks(content string) (blocks []string, remaining string) { - s := content remaining = content - for strings.Contains(s, "") { - sIdx := strings.Index(s, "") - rest := s[sIdx+len(""):] + for _, pair := range ToolTagPairs { + for strings.Contains(remaining, pair.Start) { + sIdx := strings.Index(remaining, pair.Start) + rest := remaining[sIdx+len(pair.Start):] - relNextSIdx := strings.Index(rest, "") - var nextSIdx int - if relNextSIdx != -1 { - nextSIdx = sIdx + len("") + relNextSIdx - } else { - nextSIdx = -1 - } + relNextSIdx := strings.Index(rest, pair.Start) + var nextSIdx int + if relNextSIdx != -1 { + nextSIdx = sIdx + len(pair.Start) + relNextSIdx + } else { + nextSIdx = -1 + } - relEIdx := strings.Index(rest, "") - var eIdx int - if relEIdx != -1 { - eIdx = sIdx + len("") + relEIdx - } else { - eIdx = -1 - } + relEIdx := strings.Index(rest, pair.End) + var eIdx int + if relEIdx != -1 { + eIdx = sIdx + len(pair.Start) + relEIdx + } else { + eIdx = -1 + } - var blockText string - var blockEndPos int + var blockText string + if eIdx != -1 && (nextSIdx == -1 || eIdx < nextSIdx) { + blockEndPos := eIdx + len(pair.End) + blockText = remaining[sIdx:blockEndPos] + remaining = strings.TrimSpace(remaining[:sIdx] + remaining[blockEndPos:]) + } else if nextSIdx != -1 { + blockEndPos := nextSIdx + blockText = remaining[sIdx:blockEndPos] + remaining = strings.TrimSpace(remaining[:sIdx] + remaining[blockEndPos:]) + } else { + blockText = remaining[sIdx:] + remaining = strings.TrimSpace(remaining[:sIdx]) + } - if eIdx != -1 && (nextSIdx == -1 || eIdx < nextSIdx) { - blockEndPos = eIdx + len("") - blockText = s[sIdx:blockEndPos] - s = s[blockEndPos:] - } else if nextSIdx != -1 { - blockEndPos = nextSIdx - blockText = s[sIdx:blockEndPos] - s = s[blockEndPos:] - } else { - blockText = s[sIdx:] - s = "" - } - - blocks = append(blocks, blockText) - } - - for strings.Contains(remaining, "") { - st := strings.Index(remaining, "") - rest := remaining[st+len(""):] - - relNext := strings.Index(rest, "") - var nextSt int - if relNext != -1 { - nextSt = st + len("") + relNext - } else { - nextSt = -1 - } - - relEn := strings.Index(rest, "") - var en int - if relEn != -1 { - en = st + len("") + relEn - } else { - en = -1 - } - - if en != -1 && (nextSt == -1 || en < nextSt) { - remaining = strings.TrimSpace(remaining[:st] + remaining[en+len(""):]) - } else if nextSt != -1 { - remaining = strings.TrimSpace(remaining[:st] + remaining[nextSt:]) - } else { - remaining = strings.TrimSpace(remaining[:st]) + blocks = append(blocks, blockText) } } @@ -800,8 +912,8 @@ func DetectToolCalls(content string) ([]ToolCall, string, bool) { var calls []ToolCall for _, block := range blocks { - if toolCall, ok := parseXMLToolCall(block); ok { - calls = append(calls, toolCall) + if tcs, ok := parseXMLToolCall(block); ok { + calls = append(calls, tcs...) } } @@ -809,8 +921,8 @@ func DetectToolCalls(content string) ([]ToolCall, string, bool) { return calls, remaining, true } - if tc, ok := parseSingleToolCall(strings.TrimSpace(content)); ok { - return []ToolCall{tc}, "", true + if tcs, ok := parseMultipleToolCalls(strings.TrimSpace(content)); ok && len(tcs) > 0 { + return tcs, "", true } return nil, content, false @@ -955,12 +1067,24 @@ func NewStreamThinkingFilter() *StreamThinkingFilter { } func hasPrefixOf(target string, prefixes []string) int { + maxMatch := 0 for _, p := range prefixes { - if strings.HasSuffix(target, p) { - return len(p) + if strings.HasSuffix(target, p) && len(p) > maxMatch { + maxMatch = len(p) } } - return 0 + return maxMatch +} + +func hasSuffixPrefixOf(target string, tag string) int { + maxMatch := 0 + for i := 1; i < len(tag); i++ { + p := tag[:i] + if strings.HasSuffix(target, p) && len(p) > maxMatch { + maxMatch = len(p) + } + } + return maxMatch } func (f *StreamThinkingFilter) Feed(chunk string, onContent func(string), onReasoning func(string)) { @@ -1034,11 +1158,13 @@ func (f *StreamThinkingFilter) Flush(onContent func(string), onReasoning func(st // --------------------------------------------------------------------------- type StreamToolCallFilter struct { - inToolCall bool - buf string - toolCallBuf string - toolIndex int - emittedCall bool + inToolCall bool + buf string + toolCallBuf string + toolIndex int + emittedCall bool + activePair ToolTagPair + activeEndTag string } func NewStreamToolCallFilter() *StreamToolCallFilter { @@ -1047,22 +1173,31 @@ func NewStreamToolCallFilter() *StreamToolCallFilter { func (f *StreamToolCallFilter) Feed(chunk string, onContent func(string), onToolCall func(ToolCall)) { f.buf += chunk - toolStartTag := "" - toolEndTag := "" - - startPrefixes := []string{"<", " 0 { if !f.inToolCall { - if idx := strings.Index(f.buf, toolStartTag); idx != -1 { - before := f.buf[:idx] + earliestIdx := -1 + var matchedPair ToolTagPair + + for _, pair := range ToolTagPairs { + if idx := strings.Index(f.buf, pair.Start); idx != -1 { + if earliestIdx == -1 || idx < earliestIdx { + earliestIdx = idx + matchedPair = pair + } + } + } + + if earliestIdx != -1 { + before := f.buf[:earliestIdx] if before != "" { onContent(before) } f.inToolCall = true - f.buf = f.buf[idx+len(toolStartTag):] - } else if matchLen := hasPrefixOf(f.buf, startPrefixes); matchLen > 0 { + f.activePair = matchedPair + f.activeEndTag = matchedPair.End + f.buf = f.buf[earliestIdx+len(matchedPair.Start):] + } else if matchLen := hasPrefixOf(f.buf, toolStartPrefixes); matchLen > 0 { safe := f.buf[:len(f.buf)-matchLen] if safe != "" { onContent(safe) @@ -1075,28 +1210,32 @@ func (f *StreamToolCallFilter) Feed(chunk string, onContent func(string), onTool break } } else { - if idx := strings.Index(f.buf, toolEndTag); idx != -1 { + if idx := strings.Index(f.buf, f.activeEndTag); idx != -1 { f.toolCallBuf += f.buf[:idx] - f.buf = f.buf[idx+len(toolEndTag):] + f.buf = f.buf[idx+len(f.activeEndTag):] f.inToolCall = false - if tc, ok := parseSingleToolCall(f.toolCallBuf); ok { - idxCopy := f.toolIndex - tc.Index = &idxCopy - f.toolIndex++ - f.emittedCall = true - onToolCall(tc) - } else if tc2, ok2 := parseXMLToolCall("" + f.toolCallBuf + ""); ok2 { - idxCopy := f.toolIndex - tc2.Index = &idxCopy - f.toolIndex++ - f.emittedCall = true - onToolCall(tc2) + if tcs, ok := parseMultipleToolCalls(f.toolCallBuf); ok && len(tcs) > 0 { + for _, tc := range tcs { + idxCopy := f.toolIndex + tc.Index = &idxCopy + f.toolIndex++ + f.emittedCall = true + onToolCall(tc) + } + } else if tcs2, ok2 := parseXMLToolCall(f.activePair.Start + f.toolCallBuf + f.activePair.End); ok2 && len(tcs2) > 0 { + for _, tc := range tcs2 { + idxCopy := f.toolIndex + tc.Index = &idxCopy + f.toolIndex++ + f.emittedCall = true + onToolCall(tc) + } } else { - onContent("" + f.toolCallBuf + "") + onContent(f.activePair.Start + f.toolCallBuf + f.activePair.End) } f.toolCallBuf = "" - } else if matchLen := hasPrefixOf(f.buf, endPrefixes); matchLen > 0 { + } else if matchLen := hasSuffixPrefixOf(f.buf, f.activeEndTag); matchLen > 0 { safe := f.buf[:len(f.buf)-matchLen] f.toolCallBuf += safe f.buf = f.buf[len(f.buf)-matchLen:] @@ -1112,18 +1251,24 @@ func (f *StreamToolCallFilter) Feed(chunk string, onContent func(string), onTool func (f *StreamToolCallFilter) Flush(onContent func(string), onToolCall func(ToolCall)) { if f.inToolCall && len(f.toolCallBuf) > 0 { - if tc, ok := parseSingleToolCall(f.toolCallBuf); ok { - idxCopy := f.toolIndex - tc.Index = &idxCopy - f.emittedCall = true - onToolCall(tc) - } else if tc2, ok2 := parseXMLToolCall("" + f.toolCallBuf + ""); ok2 { - idxCopy := f.toolIndex - tc2.Index = &idxCopy - f.emittedCall = true - onToolCall(tc2) + if tcs, ok := parseMultipleToolCalls(f.toolCallBuf); ok && len(tcs) > 0 { + for _, tc := range tcs { + idxCopy := f.toolIndex + tc.Index = &idxCopy + f.toolIndex++ + f.emittedCall = true + onToolCall(tc) + } + } else if tcs2, ok2 := parseXMLToolCall(f.activePair.Start + f.toolCallBuf + f.activePair.End); ok2 && len(tcs2) > 0 { + for _, tc := range tcs2 { + idxCopy := f.toolIndex + tc.Index = &idxCopy + f.toolIndex++ + f.emittedCall = true + onToolCall(tc) + } } else { - onContent("" + f.toolCallBuf) + onContent(f.activePair.Start + f.toolCallBuf) } f.toolCallBuf = "" } @@ -1268,21 +1413,18 @@ func (d *SpaceDiscovery) GetModelList() []ModelItem { return items } -// InspectSpace queries Gradio's /gradio_api/info, /config, and HuggingFace Space APIs -// to build an adaptive schema mapping for any Gradio space. -func InspectSpace(client *http.Client, rawURL, userAgent string) (*SpaceDiscovery, error) { - cleanURL := strings.TrimRight(rawURL, "/") - if !strings.HasPrefix(cleanURL, "http://") && !strings.HasPrefix(cleanURL, "https://") { +func NewDefaultSpaceDiscovery(spaceURL string) *SpaceDiscovery { + cleanURL := strings.TrimRight(spaceURL, "/") + if cleanURL != "" && !strings.HasPrefix(cleanURL, "http://") && !strings.HasPrefix(cleanURL, "https://") { cleanURL = "https://" + cleanURL } - - discovery := &SpaceDiscovery{ - SpaceURL: cleanURL, - APIPrefix: "/gradio_api", - Endpoint: "/chat_fn", - CleanEndpoint: "chat_fn", - Protocol: "call", - TotalInputs: 1, + return &SpaceDiscovery{ + SpaceURL: cleanURL, + APIPrefix: "/gradio_api", + Endpoint: "/chat_fn", + CleanEndpoint: "chat_fn", + Protocol: "call", + TotalInputs: 1, HistoryIndex: -1, MessageIndex: 0, SystemIndex: -1, @@ -1295,6 +1437,17 @@ func InspectSpace(client *http.Client, rawURL, userAgent string) (*SpaceDiscover HistoryFormat: "messages", LastDiscovered: time.Now(), } +} + +// InspectSpace queries Gradio's /gradio_api/info, /config, and HuggingFace Space APIs +// to build an adaptive schema mapping for any Gradio space. +func InspectSpace(client *http.Client, rawURL, userAgent string) (*SpaceDiscovery, error) { + cleanURL := strings.TrimRight(rawURL, "/") + if !strings.HasPrefix(cleanURL, "http://") && !strings.HasPrefix(cleanURL, "https://") { + cleanURL = "https://" + cleanURL + } + + discovery := NewDefaultSpaceDiscovery(cleanURL) // 1. Try fetching /gradio_api/info or /info var infoResp GradioAPIInfoResponse @@ -1532,16 +1685,19 @@ func InspectSpace(client *http.Client, rawURL, userAgent string) (*SpaceDiscover } } - // Check parameters in bestEndpointInfo for history support + // Check parameters in bestEndpointInfo for input indices and history support if bestEndpointInfo != nil { + if discovery.TotalInputs < len(bestEndpointInfo.Parameters) { + discovery.TotalInputs = len(bestEndpointInfo.Parameters) + } for idx, p := range bestEndpointInfo.Parameters { pName := strings.ToLower(p.ParameterName) - if strings.Contains(pName, "message") && discovery.MessageIndex == 0 { - discovery.MessageIndex = idx - } else if strings.Contains(pName, "history") { - discovery.HistoryIndex = idx - } else if strings.Contains(pName, "system") { + if strings.Contains(pName, "system") { discovery.SystemIndex = idx + } else if strings.Contains(pName, "history") || strings.Contains(pName, "chat") { + discovery.HistoryIndex = idx + } else if strings.Contains(pName, "message") || (strings.Contains(pName, "prompt") && !strings.Contains(pName, "system")) || strings.Contains(pName, "text") { + discovery.MessageIndex = idx } else if strings.Contains(pName, "think_level") || strings.Contains(pName, "thinking_level") { discovery.ThinkLevelIndex = idx } else if strings.Contains(pName, "functions") || strings.Contains(pName, "tools") { @@ -1660,20 +1816,10 @@ func (g *GradioGateway) GetDiscovery(spaceURL, userAgent string) *SpaceDiscovery } // Fallback discovery - fallback := &SpaceDiscovery{ - SpaceURL: cleanTarget, - APIPrefix: "/gradio_api", - Endpoint: "/chat_fn", - CleanEndpoint: "chat_fn", - Protocol: "call", - TotalInputs: 2, - HistoryIndex: -1, - MessageIndex: 0, - SystemIndex: -1, - PrimaryModel: "gradio-chat", - Models: []string{"gradio-chat"}, - LastDiscovered: time.Now(), - } + fallback := NewDefaultSpaceDiscovery(cleanTarget) + fallback.TotalInputs = 2 + fallback.PrimaryModel = "gradio-chat" + fallback.Models = []string{"gradio-chat"} g.discoveries[cleanTarget] = fallback return fallback } @@ -1703,13 +1849,27 @@ func (g *GradioGateway) BuildGradioPayload(disc *SpaceDiscovery, req ChatComplet var nonSystem []ChatMessage for _, m := range transformed { cStr := m.GetContentString() - if m.Role == "system" && systemPromptStr == "" { - systemPromptStr = cStr + if m.Role == "system" { + if systemPromptStr == "" { + systemPromptStr = cStr + } else { + systemPromptStr += "\n\n" + cStr + } } else { nonSystem = append(nonSystem, m) } } + // If the space has NO native system prompt input (disc.SystemIndex == -1), + // but we have system instructions (from system message or tool instructions): + if disc.SystemIndex == -1 && systemPromptStr != "" && len(nonSystem) > 0 { + // If the space supports conversation history, prepend system instructions to the first turn + if disc.HistoryIndex != -1 { + nonSystem[0].Content = systemPromptStr + "\n\n" + nonSystem[0].GetContentString() + systemPromptStr = "" + } + } + if len(nonSystem) > 0 { for i := 0; i < len(nonSystem)-1; i++ { m := nonSystem[i] @@ -1803,7 +1963,15 @@ func (g *GradioGateway) BuildGradioPayload(disc *SpaceDiscovery, req ChatComplet } sb.WriteString(fmt.Sprintf("%s: %s\n\n", roleLabel, m.GetContentString())) } - sb.WriteString(lastUserMessage) + lastRoleLabel := "User" + 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) + } promptMessageText = sb.String() } } diff --git a/gr2gw_test.go b/gr2gw_test.go index aa7ae06..5c7c015 100644 --- a/gr2gw_test.go +++ b/gr2gw_test.go @@ -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 blocks + // [3] User message with coalesced 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 tag + multiXML := ` +[ + {"name": "get_weather", "arguments": {"city": "Tokyo"}}, + {"name": "get_weather", "arguments": {"city": "Paris"}} +] +` + calls1, rem1, ok1 := DetectToolCalls(multiXML) + if !ok1 || len(calls1) != 2 { + t.Fatalf("expected 2 tool calls from , 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. tag + fnCallXML := `Some preamble before call. + +{"name": "search", "arguments": {"q": "golang"}} + +Some postamble.` + calls2, rem2, ok2 := DetectToolCalls(fnCallXML) + if !ok2 || len(calls2) != 1 { + t.Fatalf("expected 1 tool call from , 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: ") { + 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, "") { + // 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: [\"\\n{\\\"name\\\": \\\"get_weather\\\", \\\"arguments\\\": {\\\"city\\\": \\\"Tokyo\\\"}}\\n\", null]\n\n") + flusher.Flush() + fmt.Fprintf(w, "event: complete\ndata: [\"\\n{\\\"name\\\": \\\"get_weather\\\", \\\"arguments\\\": {\\\"city\\\": \\\"Tokyo\\\"}}\\n\", 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()) + } +} + +