package main
import (
"bytes"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
)
func TestParseSOCKS5URL(t *testing.T) {
tests := []struct {
input string
expected *SOCKS5Config
}{
{"", nil},
{"127.0.0.1:1080", &SOCKS5Config{Address: "127.0.0.1:1080"}},
{"socks5://127.0.0.1:9050", &SOCKS5Config{Address: "127.0.0.1:9050"}},
{"socks5h://user:pass@10.0.0.1:1080", &SOCKS5Config{Address: "10.0.0.1:1080", Username: "user", Password: "pass"}},
}
for _, tc := range tests {
cfg, err := ParseSOCKS5URL(tc.input)
if err != nil {
t.Fatalf("unexpected error for %q: %v", tc.input, err)
}
if tc.expected == nil {
if cfg != nil {
t.Errorf("expected nil config, got %+v", cfg)
}
continue
}
if cfg.Address != tc.expected.Address || cfg.Username != tc.expected.Username || cfg.Password != tc.expected.Password {
t.Errorf("for %q, expected %+v, got %+v", tc.input, tc.expected, cfg)
}
}
}
func TestChatMessageGetContentString(t *testing.T) {
m1 := ChatMessage{Role: "user", Content: "hello world"}
if m1.GetContentString() != "hello world" {
t.Errorf("expected 'hello world', got %q", m1.GetContentString())
}
m2 := ChatMessage{
Role: "user",
Content: []interface{}{
map[string]interface{}{"type": "text", "text": "part 1 "},
map[string]interface{}{"type": "text", "text": "part 2"},
},
}
if m2.GetContentString() != "part 1 part 2" {
t.Errorf("expected 'part 1 part 2', got %q", m2.GetContentString())
}
}
func TestExtractThinking(t *testing.T) {
content := "Let me calculate 2+2.The answer is 4."
clean, reasoning := ExtractThinking(content)
if reasoning != "Let me calculate 2+2." {
t.Errorf("expected reasoning 'Let me calculate 2+2.', got %q", reasoning)
}
if clean != "The answer is 4." {
t.Errorf("expected clean 'The answer is 4.', got %q", clean)
}
}
func TestDetectToolCalls(t *testing.T) {
xmlContent := `
{"name": "get_weather", "arguments": {"city": "Paris"}}
`
calls, rem, ok := DetectToolCalls(xmlContent)
if !ok || len(calls) != 1 {
t.Fatalf("expected 1 tool call, got %d (ok: %v)", len(calls), ok)
}
if calls[0].Function.Name != "get_weather" {
t.Errorf("expected function name get_weather, got %q", calls[0].Function.Name)
}
if rem != "" {
t.Errorf("expected empty remaining content, got %q", rem)
}
jsonContent := `{"name": "calculator", "arguments": {"expr": "1+1"}}`
calls2, rem2, ok2 := DetectToolCalls(jsonContent)
if !ok2 || len(calls2) != 1 {
t.Fatalf("expected 1 tool call from JSON, got %d", len(calls2))
}
if calls2[0].Function.Name != "calculator" {
t.Errorf("expected function calculator, got %q", calls2[0].Function.Name)
}
if rem2 != "" {
t.Errorf("expected empty remaining, got %q", rem2)
}
}
func TestStreamThinkingFilter(t *testing.T) {
filter := NewStreamThinkingFilter()
var contentParts []string
var reasoningParts []string
onContent := func(s string) { contentParts = append(contentParts, s) }
onReasoning := func(s string) { reasoningParts = append(reasoningParts, s) }
chunks := []string{"Thinking de", "eplyHere is your answer."}
for _, c := range chunks {
filter.Feed(c, onContent, onReasoning)
}
filter.Flush(onContent, onReasoning)
fullReasoning := strings.Join(reasoningParts, "")
fullContent := strings.Join(contentParts, "")
if fullReasoning != "Thinking deeply" {
t.Errorf("expected reasoning 'Thinking deeply', got %q", fullReasoning)
}
if fullContent != "Here is your answer." {
t.Errorf("expected content 'Here is your answer.', got %q", fullContent)
}
}
func TestStreamToolCallFilter(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) }
chunks := []string{
"Searching now: ",
"\n{\"name\": \"search_web\", \"arguments\": {\"query\": \"golang\"}}\n",
" Done.",
}
for _, c := range chunks {
filter.Feed(c, onContent, onToolCall)
}
filter.Flush(onContent, onToolCall)
if len(toolCalls) != 1 {
t.Fatalf("expected 1 emitted tool call, got %d", len(toolCalls))
}
if toolCalls[0].Function.Name != "search_web" {
t.Errorf("expected tool name 'search_web', got %q", toolCalls[0].Function.Name)
}
fullContent := strings.Join(contentParts, "")
if fullContent != "Searching now: Done." {
t.Errorf("expected 'Searching now: Done.', got %q", fullContent)
}
}
func TestMockGradioServerCompletion(t *testing.T) {
// Setup a mock Gradio server
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path == "/gradio_api/info" {
resp := GradioAPIInfoResponse{
NamedEndpoints: map[string]GradioEndpointInfo{
"/chat_fn": {
Parameters: []GradioParamInfo{
{ParameterName: "message", Component: "Textbox"},
},
Returns: []GradioParamInfo{
{ParameterName: "response", Component: "Json"},
},
},
},
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(resp)
return
}
if r.URL.Path == "/gradio_api/call/chat_fn" {
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(GradioJoinResponse{EventID: "evt_123"})
return
}
if r.URL.Path == "/gradio_api/call/chat_fn/evt_123" {
w.Header().Set("Content-Type", "text/event-stream")
flusher, ok := w.(http.Flusher)
if !ok {
t.Fatal("expected flusher")
}
fmt.Fprintf(w, "event: generating\ndata: [\"Hello \", null]\n\n")
flusher.Flush()
fmt.Fprintf(w, "event: generating\ndata: [\"Hello world!\", null]\n\n")
flusher.Flush()
fmt.Fprintf(w, "event: complete\ndata: [\"Hello world!\", null]\n\n")
flusher.Flush()
return
}
http.NotFound(w, r)
}))
defer ts.Close()
gw := NewGradioGateway(ts.URL, "", 10*time.Second)
// 1. Test Non-streaming request
reqBody := ChatCompletionRequest{
Model: "test-model",
Messages: []ChatMessage{
{Role: "user", Content: "Hi"},
},
Stream: false,
}
b, _ := json.Marshal(reqBody)
httpReq := httptest.NewRequest("POST", "/v1/chat/completions", bytes.NewBuffer(b))
httpReq.Header.Set("Content-Type", "application/json")
rec := httptest.NewRecorder()
err := gw.ExecuteChatCompletion(rec, httpReq, reqBody)
if err != nil {
t.Fatalf("unexpected completion error: %v", err)
}
var resp ChatCompletionResponse
if err := json.NewDecoder(rec.Body).Decode(&resp); err != nil {
t.Fatalf("failed to decode completion response: %v", err)
}
if len(resp.Choices) != 1 {
t.Fatalf("expected 1 choice, got %d", len(resp.Choices))
}
if resp.Choices[0].Message.GetContentString() != "Hello world!" {
t.Errorf("expected 'Hello world!', got %q", resp.Choices[0].Message.GetContentString())
}
// 2. Test Streaming request
reqBodyStream := reqBody
reqBodyStream.Stream = true
recStream := httptest.NewRecorder()
err = gw.ExecuteChatCompletion(recStream, httpReq, reqBodyStream)
if err != nil {
t.Fatalf("unexpected streaming error: %v", err)
}
streamOutput := recStream.Body.String()
if !strings.Contains(streamOutput, "data: [DONE]") {
t.Errorf("expected stream to contain [DONE], got:\n%s", streamOutput)
}
if !strings.Contains(streamOutput, "Hello world!") && !strings.Contains(streamOutput, "world!") {
t.Errorf("expected stream output to contain delta tokens, got:\n%s", streamOutput)
}
}
func TestParseGradioStreamOutput(t *testing.T) {
// 1. Standard 1D Gradio array
frame1 := ParseGradioStreamOutput(`["Hello from 1D", null]`)
if !frame1.OK || frame1.Content != "Hello from 1D" || frame1.Reasoning != "" || len(frame1.ToolCalls) != 0 {
t.Errorf("unexpected frame1: %+v", frame1)
}
// 2. Hy3 2D array frame with reasoning
hy3Raw := `[["Hello answer", "Let me think deeply...", [], [{"role": "user", "content": "hi"}]]]`
frame2 := ParseGradioStreamOutput(hy3Raw)
if !frame2.OK || frame2.Content != "Hello answer" || frame2.Reasoning != "Let me think deeply..." || len(frame2.ToolCalls) != 0 {
t.Errorf("unexpected frame2: %+v", frame2)
}
// 3. Hy3 2D array frame with tool calls
hy3ToolRaw := `[["", "Calling weather tool", [{"id": "call_abc", "type": "function", "function": {"name": "get_weather", "arguments": "{\"city\": \"Tokyo\"}"}}], []]]`
frame3 := ParseGradioStreamOutput(hy3ToolRaw)
if !frame3.OK || frame3.Content != "" || frame3.Reasoning != "Calling weather tool" || len(frame3.ToolCalls) != 1 {
t.Fatalf("unexpected frame3: %+v", frame3)
}
if frame3.ToolCalls[0].ID != "call_abc" || frame3.ToolCalls[0].Function.Name != "get_weather" {
t.Errorf("unexpected tool call in frame3: %+v", frame3.ToolCalls[0])
}
// 4. Chat pairs
pairRaw := `[[["user prompt", "assistant answer"]]]`
frame4 := ParseGradioStreamOutput(pairRaw)
if !frame4.OK || frame4.Content != "assistant answer" {
t.Errorf("unexpected frame4: %+v", frame4)
}
// 5. Messages array
msgRaw := `[[{"role": "assistant", "content": "msg answer", "reasoning_content": "msg think"}]]`
frame5 := ParseGradioStreamOutput(msgRaw)
if !frame5.OK || frame5.Content != "msg answer" || frame5.Reasoning != "msg think" {
t.Errorf("unexpected frame5: %+v", frame5)
}
}
func TestHunyuan3BuildPayload(t *testing.T) {
gw := &GradioGateway{}
disc := &SpaceDiscovery{
TotalInputs: 9,
MessageIndex: 0,
SystemIndex: 1,
HistoryIndex: 2,
ThinkLevelIndex: 3,
TempIndex: 4,
MaxTokensIndex: 5,
TopPIndex: 6,
FunctionsJSONIndex: 8,
IsHunyuan3: true,
}
temp := 0.2
req := ChatCompletionRequest{
Model: "hy3",
ReasoningEffort: "low",
Temperature: &temp,
Tools: []Tool{
{
Type: "function",
Function: map[string]interface{}{
"name": "calc",
},
},
},
Messages: []ChatMessage{
{Role: "system", Content: "Be helpful"},
{Role: "user", Content: "2+2"},
{
Role: "assistant",
ReasoningContent: "Thinking...",
ToolCalls: []ToolCall{
{ID: "c1", Type: "function", Function: ToolCallFunction{Name: "calc", Arguments: `{"expr":"2+2"}`}},
},
},
{Role: "tool", ToolCallID: "c1", Content: "4"},
},
}
data, err := gw.BuildGradioPayload(disc, req)
if err != nil {
t.Fatalf("BuildGradioPayload failed: %v", err)
}
if len(data) != 9 {
t.Fatalf("expected 9 payload items, got %d", len(data))
}
// Message parameter (0): should be prompt continuation since last was tool
if msg, ok := data[0].(string); !ok || msg != "Please proceed based on the tool results." {
t.Errorf("expected continuation prompt, got %v", data[0])
}
// System parameter (1)
if sys, ok := data[1].(string); !ok || sys != "Be helpful" {
t.Errorf("expected 'Be helpful', got %v", data[1])
}
// History parameter (2): should contain all messages including the tool turn
hist, ok := data[2].([]map[string]interface{})
if !ok {
t.Fatalf("expected history slice of maps, got %T", data[2])
}
if len(hist) != 3 {
t.Fatalf("expected 3 history items (user, assistant, tool), got %d", len(hist))
}
if hist[2]["role"] != "tool" || hist[2]["content"] != "4" || hist[2]["tool_call_id"] != "c1" {
t.Errorf("unexpected tool history entry: %+v", hist[2])
}
// ThinkLevel parameter (3)
if data[3] != "low" {
t.Errorf("expected think_level 'low', got %v", data[3])
}
// Temp parameter (4)
if data[4] != 0.2 {
t.Errorf("expected temp 0.2, got %v", data[4])
}
// FunctionsJSON parameter (8)
fnStr, ok := data[8].(string)
if !ok || !strings.Contains(fnStr, "calc") {
t.Errorf("expected functions_json_str to contain 'calc', got %v", data[8])
}
}
func TestHunyuan3MockServerCompletion(t *testing.T) {
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path == "/gradio_api/info" {
resp := GradioAPIInfoResponse{
NamedEndpoints: map[string]GradioEndpointInfo{
"/chat": {
Parameters: []GradioParamInfo{
{ParameterName: "message"},
{ParameterName: "system_prompt"},
{ParameterName: "history"},
{ParameterName: "think_level"},
{ParameterName: "temperature"},
{ParameterName: "max_tokens"},
{ParameterName: "top_p"},
{ParameterName: "preserved_thinking"},
{ParameterName: "functions_json_str"},
},
},
},
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(resp)
return
}
if r.URL.Path == "/gradio_api/call/chat" {
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(GradioJoinResponse{EventID: "evt_hy3"})
return
}
if r.URL.Path == "/gradio_api/call/chat/evt_hy3" {
w.Header().Set("Content-Type", "text/event-stream")
flusher, ok := w.(http.Flusher)
if !ok {
t.Fatal("expected flusher")
}
// Frame 1: Reasoning delta
fmt.Fprintf(w, "event: generating\ndata: [[\"\", \"Reasoning part 1 \", [], []]]\n\n")
flusher.Flush()
// Frame 2: Tool call initiated
fmt.Fprintf(w, "event: generating\ndata: [[\"\", \"Reasoning part 1 and 2\", [{\"id\": \"call_hy3\", \"type\": \"function\", \"function\": {\"name\": \"search\", \"arguments\": \"{\\\"q\\\": \\\"tencent\\\"}\"}}], []]]\n\n")
flusher.Flush()
// Frame 3: Completion
fmt.Fprintf(w, "event: complete\ndata: [[\"\", \"Reasoning part 1 and 2\", [{\"id\": \"call_hy3\", \"type\": \"function\", \"function\": {\"name\": \"search\", \"arguments\": \"{\\\"q\\\": \\\"tencent\\\"}\"}}], []]]\n\n")
flusher.Flush()
return
}
http.NotFound(w, r)
}))
defer ts.Close()
gw := NewGradioGateway(ts.URL, "", 10*time.Second)
// 1. Non-streaming tool call test
reqBody := ChatCompletionRequest{
Model: "hy3",
Messages: []ChatMessage{
{Role: "user", Content: "search for tencent"},
},
Stream: false,
}
b, _ := json.Marshal(reqBody)
httpReq := httptest.NewRequest("POST", "/v1/chat/completions", bytes.NewBuffer(b))
httpReq.Header.Set("Content-Type", "application/json")
rec := httptest.NewRecorder()
err := gw.ExecuteChatCompletion(rec, httpReq, reqBody)
if err != nil {
t.Fatalf("unexpected completion error: %v", err)
}
var resp ChatCompletionResponse
if err := json.NewDecoder(rec.Body).Decode(&resp); err != nil {
t.Fatalf("failed to decode response: %v", err)
}
if resp.Choices[0].FinishReason != "tool_calls" {
t.Errorf("expected finish_reason 'tool_calls', got %q", resp.Choices[0].FinishReason)
}
if resp.Choices[0].Message.ReasoningContent != "Reasoning part 1 and 2" {
t.Errorf("expected native reasoning, got %q", resp.Choices[0].Message.ReasoningContent)
}
if len(resp.Choices[0].Message.ToolCalls) != 1 {
t.Fatalf("expected 1 tool call, got %d", len(resp.Choices[0].Message.ToolCalls))
}
if resp.Choices[0].Message.ToolCalls[0].Function.Name != "search" {
t.Errorf("expected function 'search', got %q", resp.Choices[0].Message.ToolCalls[0].Function.Name)
}
// 2. Streaming tool call test
reqBodyStream := reqBody
reqBodyStream.Stream = true
recStream := httptest.NewRecorder()
err = gw.ExecuteChatCompletion(recStream, httpReq, reqBodyStream)
if err != nil {
t.Fatalf("unexpected streaming error: %v", err)
}
streamOut := recStream.Body.String()
if !strings.Contains(streamOut, "reasoning_content") {
t.Errorf("expected stream to contain reasoning_content, got:\n%s", streamOut)
}
if !strings.Contains(streamOut, "tool_calls") {
t.Errorf("expected stream to contain tool_calls, got:\n%s", streamOut)
}
if !strings.Contains(streamOut, "call_hy3") {
t.Errorf("expected stream to contain tool call ID call_hy3, got:\n%s", streamOut)
}
if !strings.Contains(streamOut, "\"finish_reason\":\"tool_calls\"") {
t.Errorf("expected stream finish_reason tool_calls, got:\n%s", streamOut)
}
}