feat(tools): implement universal tool calling and multi-turn resolution for generic Gradio spaces
This commit is contained in:
@@ -171,7 +171,7 @@ curl -N http://localhost:8080/v1/chat/completions \
|
|||||||
}'
|
}'
|
||||||
```
|
```
|
||||||
|
|
||||||
### Tool calling
|
### Tool calling (turn 1)
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
curl http://localhost:8080/v1/chat/completions \
|
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
|
### Dynamic target space override
|
||||||
|
|
||||||
Override the target space per request without restarting the server:
|
Override the target space per request without restarting the server:
|
||||||
|
|||||||
@@ -438,18 +438,33 @@ func BuildToolInstruction(tools []Tool) string {
|
|||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
toolsBytes, _ := json.MarshalIndent(tools, "", " ")
|
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\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<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))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TransformMessages(req ChatCompletionRequest) (processed []ChatMessage, toolInstruction string, hasSystem bool) {
|
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 _, 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()
|
contentStr := msg.GetContentString()
|
||||||
m := ChatMessage{Role: msg.Role, Content: contentStr}
|
|
||||||
switch msg.Role {
|
switch msg.Role {
|
||||||
case "system":
|
case "system":
|
||||||
hasSystem = true
|
hasSystem = true
|
||||||
m.Content = contentStr
|
staged = append(staged, ChatMessage{Role: "system", Content: contentStr})
|
||||||
|
|
||||||
case "assistant":
|
case "assistant":
|
||||||
var sb strings.Builder
|
var sb strings.Builder
|
||||||
if contentStr != "" {
|
if contentStr != "" {
|
||||||
@@ -465,40 +480,79 @@ func TransformMessages(req ChatCompletionRequest) (processed []ChatMessage, tool
|
|||||||
}
|
}
|
||||||
sb.WriteString(fmt.Sprintf("<tool_call>\n{\"name\": %q, \"arguments\": %s}\n</tool_call>", tc.Function.Name, args))
|
sb.WriteString(fmt.Sprintf("<tool_call>\n{\"name\": %q, \"arguments\": %s}\n</tool_call>", 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":
|
case "tool", "function":
|
||||||
m.Role = "user"
|
// Gather consecutive tool returns into a coalesced turn
|
||||||
toolName := msg.Name
|
var toolResponses []string
|
||||||
if toolName == "" {
|
j := i
|
||||||
toolName = msg.ToolCallID
|
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("<tool_response>\n{\"name\": %q, \"content\": %s}\n</tool_response>", tName, string(contentJSON)))
|
||||||
|
j++
|
||||||
}
|
}
|
||||||
var contentJSON []byte
|
i = j - 1 // advance loop
|
||||||
if json.Valid([]byte(contentStr)) {
|
|
||||||
contentJSON = []byte(contentStr)
|
promptSuffix := "Please answer the user's request based on the tool result."
|
||||||
} else {
|
if len(toolResponses) > 1 {
|
||||||
contentJSON, _ = json.Marshal(contentStr)
|
promptSuffix = "Please answer the user's request based on the tool results."
|
||||||
}
|
}
|
||||||
m.Content = fmt.Sprintf("<tool_response>\n{\"name\": %q, \"content\": %s}\n</tool_response>", 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 toolInstruction != "" {
|
||||||
if hasSystem {
|
if hasSystem {
|
||||||
for i, m := range processed {
|
systemInjected := false
|
||||||
|
for i, m := range staged {
|
||||||
if m.Role == "system" {
|
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
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
if !systemInjected {
|
||||||
|
staged = append([]ChatMessage{
|
||||||
|
{Role: "system", Content: strings.TrimSpace(toolInstruction)},
|
||||||
|
}, staged...)
|
||||||
|
}
|
||||||
} else {
|
} else {
|
||||||
processed = append([]ChatMessage{
|
staged = append([]ChatMessage{
|
||||||
{Role: "user", Content: strings.TrimSpace(toolInstruction)},
|
{Role: "system", Content: strings.TrimSpace(toolInstruction)},
|
||||||
}, processed...)
|
}, staged...)
|
||||||
|
hasSystem = true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return processed, toolInstruction, hasSystem
|
return staged, toolInstruction, hasSystem
|
||||||
}
|
}
|
||||||
|
|
||||||
func cleanJSONBlock(input string) string {
|
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: "<tool_call>", End: "</tool_call>"},
|
||||||
|
{Start: "<tool_calls>", End: "</tool_calls>"},
|
||||||
|
{Start: "<function_call>", End: "</function_call>"},
|
||||||
|
{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) {
|
func repairToolCallJSON(input string) (ToolCall, bool) {
|
||||||
s := strings.TrimSpace(input)
|
s := strings.TrimSpace(input)
|
||||||
reName := regexp.MustCompile(`"(?:name|function|action|call)"\s*:\s*"([^"]+)"`)
|
reName := regexp.MustCompile(`"(?:name|function|action|call)"\s*:\s*"([^"]+)"`)
|
||||||
@@ -547,8 +630,8 @@ func repairToolCallJSON(input string) (ToolCall, bool) {
|
|||||||
}
|
}
|
||||||
fnName := matches[1]
|
fnName := matches[1]
|
||||||
|
|
||||||
reArgs := regexp.MustCompile(`"(?:arguments|parameters|args|input)"\s*:\s*(\{[\s\S]*\})`)
|
reArgsObj := regexp.MustCompile(`"(?:arguments|parameters|args|input)"\s*:\s*(\{[\s\S]*\})`)
|
||||||
argMatches := reArgs.FindStringSubmatch(s)
|
argMatches := reArgsObj.FindStringSubmatch(s)
|
||||||
argsStr := "{}"
|
argsStr := "{}"
|
||||||
if len(argMatches) >= 2 {
|
if len(argMatches) >= 2 {
|
||||||
candidate := argMatches[1]
|
candidate := argMatches[1]
|
||||||
@@ -556,6 +639,15 @@ func repairToolCallJSON(input string) (ToolCall, bool) {
|
|||||||
if json.Unmarshal([]byte(candidate), &dummy) == nil {
|
if json.Unmarshal([]byte(candidate), &dummy) == nil {
|
||||||
argsStr = candidate
|
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{
|
return ToolCall{
|
||||||
@@ -671,18 +763,68 @@ func parseSingleToolCall(jsonStr string) (ToolCall, bool) {
|
|||||||
return repairToolCallJSON(cleaned)
|
return repairToolCallJSON(cleaned)
|
||||||
}
|
}
|
||||||
|
|
||||||
func parseXMLToolCall(block string) (ToolCall, bool) {
|
func parseMultipleToolCalls(raw string) ([]ToolCall, bool) {
|
||||||
inner := strings.TrimSpace(block)
|
cleaned := cleanJSONBlock(raw)
|
||||||
if strings.HasPrefix(inner, "<tool_call>") {
|
if cleaned == "" {
|
||||||
inner = strings.TrimPrefix(inner, "<tool_call>")
|
return nil, false
|
||||||
}
|
}
|
||||||
if strings.HasSuffix(inner, "</tool_call>") {
|
|
||||||
inner = strings.TrimSuffix(inner, "</tool_call>")
|
// 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)
|
inner = cleanJSONBlock(inner)
|
||||||
|
|
||||||
if tc, ok := parseSingleToolCall(inner); ok {
|
if calls, ok := parseMultipleToolCalls(inner); ok && len(calls) > 0 {
|
||||||
return tc, true
|
return calls, true
|
||||||
}
|
}
|
||||||
|
|
||||||
var fnName string
|
var fnName string
|
||||||
@@ -707,88 +849,58 @@ func parseXMLToolCall(block string) (ToolCall, bool) {
|
|||||||
if argsStr == "" {
|
if argsStr == "" {
|
||||||
argsStr = "{}"
|
argsStr = "{}"
|
||||||
}
|
}
|
||||||
return ToolCall{
|
return []ToolCall{{
|
||||||
ID: "call_" + GenerateUUID()[:8],
|
ID: "call_" + GenerateUUID()[:8],
|
||||||
Type: "function",
|
Type: "function",
|
||||||
Function: ToolCallFunction{
|
Function: ToolCallFunction{
|
||||||
Name: fnName,
|
Name: fnName,
|
||||||
Arguments: argsStr,
|
Arguments: argsStr,
|
||||||
},
|
},
|
||||||
}, true
|
}}, true
|
||||||
}
|
}
|
||||||
|
|
||||||
return ToolCall{}, false
|
return nil, false
|
||||||
}
|
}
|
||||||
|
|
||||||
func ExtractToolCallBlocks(content string) (blocks []string, remaining string) {
|
func ExtractToolCallBlocks(content string) (blocks []string, remaining string) {
|
||||||
s := content
|
|
||||||
remaining = content
|
remaining = content
|
||||||
|
|
||||||
for strings.Contains(s, "<tool_call>") {
|
for _, pair := range ToolTagPairs {
|
||||||
sIdx := strings.Index(s, "<tool_call>")
|
for strings.Contains(remaining, pair.Start) {
|
||||||
rest := s[sIdx+len("<tool_call>"):]
|
sIdx := strings.Index(remaining, pair.Start)
|
||||||
|
rest := remaining[sIdx+len(pair.Start):]
|
||||||
|
|
||||||
relNextSIdx := strings.Index(rest, "<tool_call>")
|
relNextSIdx := strings.Index(rest, pair.Start)
|
||||||
var nextSIdx int
|
var nextSIdx int
|
||||||
if relNextSIdx != -1 {
|
if relNextSIdx != -1 {
|
||||||
nextSIdx = sIdx + len("<tool_call>") + relNextSIdx
|
nextSIdx = sIdx + len(pair.Start) + relNextSIdx
|
||||||
} else {
|
} else {
|
||||||
nextSIdx = -1
|
nextSIdx = -1
|
||||||
}
|
}
|
||||||
|
|
||||||
relEIdx := strings.Index(rest, "</tool_call>")
|
relEIdx := strings.Index(rest, pair.End)
|
||||||
var eIdx int
|
var eIdx int
|
||||||
if relEIdx != -1 {
|
if relEIdx != -1 {
|
||||||
eIdx = sIdx + len("<tool_call>") + relEIdx
|
eIdx = sIdx + len(pair.Start) + relEIdx
|
||||||
} else {
|
} else {
|
||||||
eIdx = -1
|
eIdx = -1
|
||||||
}
|
}
|
||||||
|
|
||||||
var blockText string
|
var blockText string
|
||||||
var blockEndPos int
|
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) {
|
blocks = append(blocks, blockText)
|
||||||
blockEndPos = eIdx + len("</tool_call>")
|
|
||||||
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, "<tool_call>") {
|
|
||||||
st := strings.Index(remaining, "<tool_call>")
|
|
||||||
rest := remaining[st+len("<tool_call>"):]
|
|
||||||
|
|
||||||
relNext := strings.Index(rest, "<tool_call>")
|
|
||||||
var nextSt int
|
|
||||||
if relNext != -1 {
|
|
||||||
nextSt = st + len("<tool_call>") + relNext
|
|
||||||
} else {
|
|
||||||
nextSt = -1
|
|
||||||
}
|
|
||||||
|
|
||||||
relEn := strings.Index(rest, "</tool_call>")
|
|
||||||
var en int
|
|
||||||
if relEn != -1 {
|
|
||||||
en = st + len("<tool_call>") + relEn
|
|
||||||
} else {
|
|
||||||
en = -1
|
|
||||||
}
|
|
||||||
|
|
||||||
if en != -1 && (nextSt == -1 || en < nextSt) {
|
|
||||||
remaining = strings.TrimSpace(remaining[:st] + remaining[en+len("</tool_call>"):])
|
|
||||||
} else if nextSt != -1 {
|
|
||||||
remaining = strings.TrimSpace(remaining[:st] + remaining[nextSt:])
|
|
||||||
} else {
|
|
||||||
remaining = strings.TrimSpace(remaining[:st])
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -800,8 +912,8 @@ func DetectToolCalls(content string) ([]ToolCall, string, bool) {
|
|||||||
var calls []ToolCall
|
var calls []ToolCall
|
||||||
|
|
||||||
for _, block := range blocks {
|
for _, block := range blocks {
|
||||||
if toolCall, ok := parseXMLToolCall(block); ok {
|
if tcs, ok := parseXMLToolCall(block); ok {
|
||||||
calls = append(calls, toolCall)
|
calls = append(calls, tcs...)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -809,8 +921,8 @@ func DetectToolCalls(content string) ([]ToolCall, string, bool) {
|
|||||||
return calls, remaining, true
|
return calls, remaining, true
|
||||||
}
|
}
|
||||||
|
|
||||||
if tc, ok := parseSingleToolCall(strings.TrimSpace(content)); ok {
|
if tcs, ok := parseMultipleToolCalls(strings.TrimSpace(content)); ok && len(tcs) > 0 {
|
||||||
return []ToolCall{tc}, "", true
|
return tcs, "", true
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil, content, false
|
return nil, content, false
|
||||||
@@ -955,12 +1067,24 @@ func NewStreamThinkingFilter() *StreamThinkingFilter {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func hasPrefixOf(target string, prefixes []string) int {
|
func hasPrefixOf(target string, prefixes []string) int {
|
||||||
|
maxMatch := 0
|
||||||
for _, p := range prefixes {
|
for _, p := range prefixes {
|
||||||
if strings.HasSuffix(target, p) {
|
if strings.HasSuffix(target, p) && len(p) > maxMatch {
|
||||||
return len(p)
|
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)) {
|
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 {
|
type StreamToolCallFilter struct {
|
||||||
inToolCall bool
|
inToolCall bool
|
||||||
buf string
|
buf string
|
||||||
toolCallBuf string
|
toolCallBuf string
|
||||||
toolIndex int
|
toolIndex int
|
||||||
emittedCall bool
|
emittedCall bool
|
||||||
|
activePair ToolTagPair
|
||||||
|
activeEndTag string
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewStreamToolCallFilter() *StreamToolCallFilter {
|
func NewStreamToolCallFilter() *StreamToolCallFilter {
|
||||||
@@ -1047,22 +1173,31 @@ func NewStreamToolCallFilter() *StreamToolCallFilter {
|
|||||||
|
|
||||||
func (f *StreamToolCallFilter) Feed(chunk string, onContent func(string), onToolCall func(ToolCall)) {
|
func (f *StreamToolCallFilter) Feed(chunk string, onContent func(string), onToolCall func(ToolCall)) {
|
||||||
f.buf += chunk
|
f.buf += chunk
|
||||||
toolStartTag := "<tool_call>"
|
|
||||||
toolEndTag := "</tool_call>"
|
|
||||||
|
|
||||||
startPrefixes := []string{"<", "<t", "<to", "<too", "<tool", "<tool_", "<tool_c", "<tool_ca", "<tool_cal", "<tool_call"}
|
|
||||||
endPrefixes := []string{"<", "</", "</t", "</to", "</too", "</tool", "</tool_", "</tool_c", "</tool_ca", "</tool_cal", "</tool_call"}
|
|
||||||
|
|
||||||
for len(f.buf) > 0 {
|
for len(f.buf) > 0 {
|
||||||
if !f.inToolCall {
|
if !f.inToolCall {
|
||||||
if idx := strings.Index(f.buf, toolStartTag); idx != -1 {
|
earliestIdx := -1
|
||||||
before := f.buf[:idx]
|
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 != "" {
|
if before != "" {
|
||||||
onContent(before)
|
onContent(before)
|
||||||
}
|
}
|
||||||
f.inToolCall = true
|
f.inToolCall = true
|
||||||
f.buf = f.buf[idx+len(toolStartTag):]
|
f.activePair = matchedPair
|
||||||
} else if matchLen := hasPrefixOf(f.buf, startPrefixes); matchLen > 0 {
|
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]
|
safe := f.buf[:len(f.buf)-matchLen]
|
||||||
if safe != "" {
|
if safe != "" {
|
||||||
onContent(safe)
|
onContent(safe)
|
||||||
@@ -1075,28 +1210,32 @@ func (f *StreamToolCallFilter) Feed(chunk string, onContent func(string), onTool
|
|||||||
break
|
break
|
||||||
}
|
}
|
||||||
} else {
|
} 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.toolCallBuf += f.buf[:idx]
|
||||||
f.buf = f.buf[idx+len(toolEndTag):]
|
f.buf = f.buf[idx+len(f.activeEndTag):]
|
||||||
f.inToolCall = false
|
f.inToolCall = false
|
||||||
|
|
||||||
if tc, ok := parseSingleToolCall(f.toolCallBuf); ok {
|
if tcs, ok := parseMultipleToolCalls(f.toolCallBuf); ok && len(tcs) > 0 {
|
||||||
idxCopy := f.toolIndex
|
for _, tc := range tcs {
|
||||||
tc.Index = &idxCopy
|
idxCopy := f.toolIndex
|
||||||
f.toolIndex++
|
tc.Index = &idxCopy
|
||||||
f.emittedCall = true
|
f.toolIndex++
|
||||||
onToolCall(tc)
|
f.emittedCall = true
|
||||||
} else if tc2, ok2 := parseXMLToolCall("<tool_call>" + f.toolCallBuf + "</tool_call>"); ok2 {
|
onToolCall(tc)
|
||||||
idxCopy := f.toolIndex
|
}
|
||||||
tc2.Index = &idxCopy
|
} else if tcs2, ok2 := parseXMLToolCall(f.activePair.Start + f.toolCallBuf + f.activePair.End); ok2 && len(tcs2) > 0 {
|
||||||
f.toolIndex++
|
for _, tc := range tcs2 {
|
||||||
f.emittedCall = true
|
idxCopy := f.toolIndex
|
||||||
onToolCall(tc2)
|
tc.Index = &idxCopy
|
||||||
|
f.toolIndex++
|
||||||
|
f.emittedCall = true
|
||||||
|
onToolCall(tc)
|
||||||
|
}
|
||||||
} else {
|
} else {
|
||||||
onContent("<tool_call>" + f.toolCallBuf + "</tool_call>")
|
onContent(f.activePair.Start + f.toolCallBuf + f.activePair.End)
|
||||||
}
|
}
|
||||||
f.toolCallBuf = ""
|
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]
|
safe := f.buf[:len(f.buf)-matchLen]
|
||||||
f.toolCallBuf += safe
|
f.toolCallBuf += safe
|
||||||
f.buf = f.buf[len(f.buf)-matchLen:]
|
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)) {
|
func (f *StreamToolCallFilter) Flush(onContent func(string), onToolCall func(ToolCall)) {
|
||||||
if f.inToolCall && len(f.toolCallBuf) > 0 {
|
if f.inToolCall && len(f.toolCallBuf) > 0 {
|
||||||
if tc, ok := parseSingleToolCall(f.toolCallBuf); ok {
|
if tcs, ok := parseMultipleToolCalls(f.toolCallBuf); ok && len(tcs) > 0 {
|
||||||
idxCopy := f.toolIndex
|
for _, tc := range tcs {
|
||||||
tc.Index = &idxCopy
|
idxCopy := f.toolIndex
|
||||||
f.emittedCall = true
|
tc.Index = &idxCopy
|
||||||
onToolCall(tc)
|
f.toolIndex++
|
||||||
} else if tc2, ok2 := parseXMLToolCall("<tool_call>" + f.toolCallBuf + "</tool_call>"); ok2 {
|
f.emittedCall = true
|
||||||
idxCopy := f.toolIndex
|
onToolCall(tc)
|
||||||
tc2.Index = &idxCopy
|
}
|
||||||
f.emittedCall = true
|
} else if tcs2, ok2 := parseXMLToolCall(f.activePair.Start + f.toolCallBuf + f.activePair.End); ok2 && len(tcs2) > 0 {
|
||||||
onToolCall(tc2)
|
for _, tc := range tcs2 {
|
||||||
|
idxCopy := f.toolIndex
|
||||||
|
tc.Index = &idxCopy
|
||||||
|
f.toolIndex++
|
||||||
|
f.emittedCall = true
|
||||||
|
onToolCall(tc)
|
||||||
|
}
|
||||||
} else {
|
} else {
|
||||||
onContent("<tool_call>" + f.toolCallBuf)
|
onContent(f.activePair.Start + f.toolCallBuf)
|
||||||
}
|
}
|
||||||
f.toolCallBuf = ""
|
f.toolCallBuf = ""
|
||||||
}
|
}
|
||||||
@@ -1268,21 +1413,18 @@ func (d *SpaceDiscovery) GetModelList() []ModelItem {
|
|||||||
return items
|
return items
|
||||||
}
|
}
|
||||||
|
|
||||||
// InspectSpace queries Gradio's /gradio_api/info, /config, and HuggingFace Space APIs
|
func NewDefaultSpaceDiscovery(spaceURL string) *SpaceDiscovery {
|
||||||
// to build an adaptive schema mapping for any Gradio space.
|
cleanURL := strings.TrimRight(spaceURL, "/")
|
||||||
func InspectSpace(client *http.Client, rawURL, userAgent string) (*SpaceDiscovery, error) {
|
if cleanURL != "" && !strings.HasPrefix(cleanURL, "http://") && !strings.HasPrefix(cleanURL, "https://") {
|
||||||
cleanURL := strings.TrimRight(rawURL, "/")
|
|
||||||
if !strings.HasPrefix(cleanURL, "http://") && !strings.HasPrefix(cleanURL, "https://") {
|
|
||||||
cleanURL = "https://" + cleanURL
|
cleanURL = "https://" + cleanURL
|
||||||
}
|
}
|
||||||
|
return &SpaceDiscovery{
|
||||||
discovery := &SpaceDiscovery{
|
SpaceURL: cleanURL,
|
||||||
SpaceURL: cleanURL,
|
APIPrefix: "/gradio_api",
|
||||||
APIPrefix: "/gradio_api",
|
Endpoint: "/chat_fn",
|
||||||
Endpoint: "/chat_fn",
|
CleanEndpoint: "chat_fn",
|
||||||
CleanEndpoint: "chat_fn",
|
Protocol: "call",
|
||||||
Protocol: "call",
|
TotalInputs: 1,
|
||||||
TotalInputs: 1,
|
|
||||||
HistoryIndex: -1,
|
HistoryIndex: -1,
|
||||||
MessageIndex: 0,
|
MessageIndex: 0,
|
||||||
SystemIndex: -1,
|
SystemIndex: -1,
|
||||||
@@ -1295,6 +1437,17 @@ func InspectSpace(client *http.Client, rawURL, userAgent string) (*SpaceDiscover
|
|||||||
HistoryFormat: "messages",
|
HistoryFormat: "messages",
|
||||||
LastDiscovered: time.Now(),
|
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
|
// 1. Try fetching /gradio_api/info or /info
|
||||||
var infoResp GradioAPIInfoResponse
|
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 bestEndpointInfo != nil {
|
||||||
|
if discovery.TotalInputs < len(bestEndpointInfo.Parameters) {
|
||||||
|
discovery.TotalInputs = len(bestEndpointInfo.Parameters)
|
||||||
|
}
|
||||||
for idx, p := range bestEndpointInfo.Parameters {
|
for idx, p := range bestEndpointInfo.Parameters {
|
||||||
pName := strings.ToLower(p.ParameterName)
|
pName := strings.ToLower(p.ParameterName)
|
||||||
if strings.Contains(pName, "message") && discovery.MessageIndex == 0 {
|
if strings.Contains(pName, "system") {
|
||||||
discovery.MessageIndex = idx
|
|
||||||
} else if strings.Contains(pName, "history") {
|
|
||||||
discovery.HistoryIndex = idx
|
|
||||||
} else if strings.Contains(pName, "system") {
|
|
||||||
discovery.SystemIndex = idx
|
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") {
|
} else if strings.Contains(pName, "think_level") || strings.Contains(pName, "thinking_level") {
|
||||||
discovery.ThinkLevelIndex = idx
|
discovery.ThinkLevelIndex = idx
|
||||||
} else if strings.Contains(pName, "functions") || strings.Contains(pName, "tools") {
|
} 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 discovery
|
||||||
fallback := &SpaceDiscovery{
|
fallback := NewDefaultSpaceDiscovery(cleanTarget)
|
||||||
SpaceURL: cleanTarget,
|
fallback.TotalInputs = 2
|
||||||
APIPrefix: "/gradio_api",
|
fallback.PrimaryModel = "gradio-chat"
|
||||||
Endpoint: "/chat_fn",
|
fallback.Models = []string{"gradio-chat"}
|
||||||
CleanEndpoint: "chat_fn",
|
|
||||||
Protocol: "call",
|
|
||||||
TotalInputs: 2,
|
|
||||||
HistoryIndex: -1,
|
|
||||||
MessageIndex: 0,
|
|
||||||
SystemIndex: -1,
|
|
||||||
PrimaryModel: "gradio-chat",
|
|
||||||
Models: []string{"gradio-chat"},
|
|
||||||
LastDiscovered: time.Now(),
|
|
||||||
}
|
|
||||||
g.discoveries[cleanTarget] = fallback
|
g.discoveries[cleanTarget] = fallback
|
||||||
return fallback
|
return fallback
|
||||||
}
|
}
|
||||||
@@ -1703,13 +1849,27 @@ func (g *GradioGateway) BuildGradioPayload(disc *SpaceDiscovery, req ChatComplet
|
|||||||
var nonSystem []ChatMessage
|
var nonSystem []ChatMessage
|
||||||
for _, m := range transformed {
|
for _, m := range transformed {
|
||||||
cStr := m.GetContentString()
|
cStr := m.GetContentString()
|
||||||
if m.Role == "system" && systemPromptStr == "" {
|
if m.Role == "system" {
|
||||||
systemPromptStr = cStr
|
if systemPromptStr == "" {
|
||||||
|
systemPromptStr = cStr
|
||||||
|
} else {
|
||||||
|
systemPromptStr += "\n\n" + cStr
|
||||||
|
}
|
||||||
} else {
|
} else {
|
||||||
nonSystem = append(nonSystem, m)
|
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 {
|
if len(nonSystem) > 0 {
|
||||||
for i := 0; i < len(nonSystem)-1; i++ {
|
for i := 0; i < len(nonSystem)-1; i++ {
|
||||||
m := nonSystem[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(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()
|
promptMessageText = sb.String()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+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