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) } }