package llm import ( "context" "encoding/json" "errors" "net/http" "net/http/httptest" "strings" "testing" "time" "gitea.maximumdirect.net/eric/scriptorium/internal/domain" ) func TestOpenAICompatibleClientGenerateSuccess(t *testing.T) { type observedRequest struct { Authorization string Body map[string]any } obs := &observedRequest{} ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { obs.Authorization = r.Header.Get("Authorization") if r.URL.Path != "/v1/chat/completions" { t.Fatalf("unexpected path: %s", r.URL.Path) } if ct := r.Header.Get("Content-Type"); ct != "application/json" { t.Fatalf("unexpected content type: %s", ct) } defer r.Body.Close() if err := json.NewDecoder(r.Body).Decode(&obs.Body); err != nil { t.Fatalf("failed to decode request body: %v", err) } w.Header().Set("Content-Type", "application/json") _, _ = w.Write([]byte(`{ "choices": [{"message": {"role": "assistant", "content": "hello from model"}}], "usage": {"prompt_tokens": 11, "completion_tokens": 22, "total_tokens": 33} }`)) })) defer ts.Close() client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ BaseURL: ts.URL + "/v1", APIKey: "secret-key", Timeout: 2 * time.Second, }) if err != nil { t.Fatalf("unexpected constructor error: %v", err) } resp, err := client.Generate(context.Background(), domain.GenerateRequest{ Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{ {Role: "system", Content: "You are helpful."}, {Role: "user", Content: "Say hello"}, }}, Target: domain.ModelTarget{ Model: "gpt-test", Temperature: 0.4, MaxTokens: 123, TopP: 0.7, }, }) if err != nil { t.Fatalf("expected no error, got %v", err) } if resp.Content != "hello from model" { t.Fatalf("unexpected content: %q", resp.Content) } if resp.Usage.PromptTokens != 11 || resp.Usage.CompletionTokens != 22 || resp.Usage.TotalTokens != 33 { t.Fatalf("unexpected usage: %+v", resp.Usage) } if obs.Authorization != "Bearer secret-key" { t.Fatalf("unexpected Authorization header: %q", obs.Authorization) } if got, ok := obs.Body["model"].(string); !ok || got != "gpt-test" { t.Fatalf("unexpected model payload: %#v", obs.Body["model"]) } msgs, ok := obs.Body["messages"].([]any) if !ok || len(msgs) != 2 { t.Fatalf("unexpected messages payload: %#v", obs.Body["messages"]) } msg0 := msgs[0].(map[string]any) if msg0["role"] != "system" || msg0["content"] != "You are helpful." { t.Fatalf("unexpected first message: %#v", msg0) } msg1 := msgs[1].(map[string]any) if msg1["role"] != "user" || msg1["content"] != "Say hello" { t.Fatalf("unexpected second message: %#v", msg1) } } func TestOpenAICompatibleClientNoAuthorizationHeaderWhenNoAPIKey(t *testing.T) { hadAuth := false ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { hadAuth = r.Header.Get("Authorization") != "" _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`)) })) defer ts.Close() client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1"}) if err != nil { t.Fatal(err) } _, err = client.Generate(context.Background(), domain.GenerateRequest{ Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}}, Target: domain.ModelTarget{Model: "model"}, }) if err != nil { t.Fatalf("expected no error, got %v", err) } if hadAuth { t.Fatal("did not expect Authorization header") } } func TestOpenAICompatibleClientModelFallbackFromConfig(t *testing.T) { gotModel := "" ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { var body map[string]any _ = json.NewDecoder(r.Body).Decode(&body) if m, ok := body["model"].(string); ok { gotModel = m } _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`)) })) defer ts.Close() client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1", Model: "default-model"}) if err != nil { t.Fatal(err) } _, err = client.Generate(context.Background(), domain.GenerateRequest{ Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}}, Target: domain.ModelTarget{}, }) if err != nil { t.Fatalf("expected no error, got %v", err) } if gotModel != "default-model" { t.Fatalf("expected default model, got %q", gotModel) } } func TestOpenAICompatibleClientEndpointOverride(t *testing.T) { defaultHit := false overrideHit := false defaultServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { defaultHit = true _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"default"}}]}`)) })) defer defaultServer.Close() overrideServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { overrideHit = true if r.URL.Path != "/v1/chat/completions" { t.Fatalf("unexpected path: %s", r.URL.Path) } _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"override"}}]}`)) })) defer overrideServer.Close() client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: defaultServer.URL + "/v1", Model: "m"}) if err != nil { t.Fatal(err) } resp, err := client.Generate(context.Background(), domain.GenerateRequest{ Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}}, Target: domain.ModelTarget{Endpoint: overrideServer.URL + "/v1"}, }) if err != nil { t.Fatalf("expected no error, got %v", err) } if resp.Content != "override" { t.Fatalf("expected override response, got %q", resp.Content) } if defaultHit { t.Fatal("default endpoint should not have been called") } if !overrideHit { t.Fatal("override endpoint should have been called") } } func TestOpenAICompatibleClientNon2xxError(t *testing.T) { ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusBadRequest) _, _ = w.Write([]byte(`{"error":"bad request payload"}`)) })) defer ts.Close() client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1", Model: "m"}) if err != nil { t.Fatal(err) } _, err = client.Generate(context.Background(), domain.GenerateRequest{ Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}}, }) if err == nil { t.Fatal("expected non-2xx error") } if !errors.Is(err, ErrUnexpectedStatus) { t.Fatalf("expected ErrUnexpectedStatus, got %v", err) } if !strings.Contains(err.Error(), "400") || !strings.Contains(err.Error(), "bad request payload") { t.Fatalf("expected status/body details, got %v", err) } } func TestOpenAICompatibleClientMalformedResponseInvalidJSON(t *testing.T) { ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { _, _ = w.Write([]byte(`{not valid json`)) })) defer ts.Close() client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1", Model: "m"}) if err != nil { t.Fatal(err) } _, err = client.Generate(context.Background(), domain.GenerateRequest{ Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}}, }) if err == nil { t.Fatal("expected malformed response error") } if !errors.Is(err, ErrMalformedResponse) { t.Fatalf("expected ErrMalformedResponse, got %v", err) } } func TestOpenAICompatibleClientMalformedResponseMissingChoices(t *testing.T) { ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { _, _ = w.Write([]byte(`{"choices": []}`)) })) defer ts.Close() client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1", Model: "m"}) if err != nil { t.Fatal(err) } _, err = client.Generate(context.Background(), domain.GenerateRequest{ Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}}, }) if err == nil { t.Fatal("expected malformed response error") } if !errors.Is(err, ErrMalformedResponse) { t.Fatalf("expected ErrMalformedResponse, got %v", err) } } func TestOpenAICompatibleClientTimeout(t *testing.T) { ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { time.Sleep(250 * time.Millisecond) _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`)) })) defer ts.Close() client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ BaseURL: ts.URL + "/v1", Model: "m", Timeout: 50 * time.Millisecond, }) if err != nil { t.Fatal(err) } _, err = client.Generate(context.Background(), domain.GenerateRequest{ Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}}, }) if err == nil { t.Fatal("expected timeout error") } if !errors.Is(err, ErrRequestFailed) { t.Fatalf("expected ErrRequestFailed, got %v", err) } }