250 lines
7.7 KiB
Go
250 lines
7.7 KiB
Go
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 := "<think>Let me calculate 2+2.</think>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 := `<tool_call>
|
|
{"name": "get_weather", "arguments": {"city": "Paris"}}
|
|
</tool_call>`
|
|
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{"<thi", "nk>Thinking de", "eply</th", "ink>Here 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: ",
|
|
"<tool_c",
|
|
"all>\n{\"name\": \"search_web\", \"arguments\": {\"query\": \"golang\"}}\n</tool_",
|
|
"call>",
|
|
" 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)
|
|
}
|
|
}
|