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