diff --git a/internal/llm/openai_compatible_client.go b/internal/llm/openai_compatible_client.go new file mode 100644 index 0000000..900912a --- /dev/null +++ b/internal/llm/openai_compatible_client.go @@ -0,0 +1,181 @@ +package llm + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/url" + "strings" + "time" + + "gitea.maximumdirect.net/eric/scriptorium/internal/domain" +) + +var ( + ErrInvalidConfig = errors.New("invalid llm client configuration") + ErrInvalidRequest = errors.New("invalid generate request") + ErrRequestFailed = errors.New("llm request failed") + ErrUnexpectedStatus = errors.New("llm returned non-success status") + ErrMalformedResponse = errors.New("malformed llm response") +) + +type OpenAICompatibleConfig struct { + BaseURL string + APIKey string + Model string + Timeout time.Duration + HTTPClient *http.Client +} + +type OpenAICompatibleClient struct { + baseURL string + apiKey string + defaultModel string + timeout time.Duration + httpClient *http.Client +} + +func NewOpenAICompatibleClient(cfg OpenAICompatibleConfig) (*OpenAICompatibleClient, error) { + if strings.TrimSpace(cfg.BaseURL) == "" { + return nil, fmt.Errorf("%w: base URL is required", ErrInvalidConfig) + } + if _, err := url.ParseRequestURI(cfg.BaseURL); err != nil { + return nil, fmt.Errorf("%w: invalid base URL: %v", ErrInvalidConfig, err) + } + + timeout := cfg.Timeout + if timeout <= 0 { + timeout = 30 * time.Second + } + + var client *http.Client + if cfg.HTTPClient != nil { + client = cfg.HTTPClient + if client.Timeout == 0 { + client.Timeout = timeout + } + } else { + client = &http.Client{Timeout: timeout} + } + + return &OpenAICompatibleClient{ + baseURL: strings.TrimRight(cfg.BaseURL, "/"), + apiKey: cfg.APIKey, + defaultModel: cfg.Model, + timeout: timeout, + httpClient: client, + }, nil +} + +func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.GenerateRequest) (*domain.GenerateResponse, error) { + model := strings.TrimSpace(req.Target.Model) + if model == "" { + model = strings.TrimSpace(c.defaultModel) + } + if model == "" { + return nil, fmt.Errorf("%w: model is required", ErrInvalidRequest) + } + + endpoint := strings.TrimSpace(req.Target.Endpoint) + if endpoint == "" { + endpoint = c.baseURL + } + endpoint = strings.TrimRight(endpoint, "/") + "/chat/completions" + + wireReq := openAIChatRequest{ + Model: model, + } + + wireReq.Messages = make([]openAIChatMessage, 0, len(req.Prompt.Messages)) + for _, msg := range req.Prompt.Messages { + wireReq.Messages = append(wireReq.Messages, openAIChatMessage{ + Role: msg.Role, + Content: msg.Content, + }) + } + + if req.Target.Temperature != 0 { + wireReq.Temperature = &req.Target.Temperature + } + if req.Target.MaxTokens != 0 { + wireReq.MaxTokens = &req.Target.MaxTokens + } + if req.Target.TopP != 0 { + wireReq.TopP = &req.Target.TopP + } + + payload, err := json.Marshal(wireReq) + if err != nil { + return nil, fmt.Errorf("%w: failed to encode request: %v", ErrRequestFailed, err) + } + + httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(payload)) + if err != nil { + return nil, fmt.Errorf("%w: failed to create request: %v", ErrRequestFailed, err) + } + httpReq.Header.Set("Content-Type", "application/json") + if strings.TrimSpace(c.apiKey) != "" { + httpReq.Header.Set("Authorization", "Bearer "+c.apiKey) + } + + httpResp, err := c.httpClient.Do(httpReq) + if err != nil { + return nil, fmt.Errorf("%w: %v", ErrRequestFailed, err) + } + defer httpResp.Body.Close() + + if httpResp.StatusCode < 200 || httpResp.StatusCode >= 300 { + body, _ := io.ReadAll(io.LimitReader(httpResp.Body, 4096)) + return nil, fmt.Errorf("%w: status=%d body=%q", ErrUnexpectedStatus, httpResp.StatusCode, strings.TrimSpace(string(body))) + } + + var wireResp openAIChatResponse + if err := json.NewDecoder(httpResp.Body).Decode(&wireResp); err != nil { + return nil, fmt.Errorf("%w: failed to decode response: %v", ErrMalformedResponse, err) + } + + if len(wireResp.Choices) == 0 { + return nil, fmt.Errorf("%w: no choices returned", ErrMalformedResponse) + } + content := wireResp.Choices[0].Message.Content + if content == "" { + return nil, fmt.Errorf("%w: first choice has empty message content", ErrMalformedResponse) + } + + return &domain.GenerateResponse{ + Content: content, + Usage: domain.TokenUsage{ + PromptTokens: wireResp.Usage.PromptTokens, + CompletionTokens: wireResp.Usage.CompletionTokens, + TotalTokens: wireResp.Usage.TotalTokens, + }, + }, nil +} + +type openAIChatRequest struct { + Model string `json:"model"` + Messages []openAIChatMessage `json:"messages"` + Temperature *float64 `json:"temperature,omitempty"` + MaxTokens *int `json:"max_tokens,omitempty"` + TopP *float64 `json:"top_p,omitempty"` +} + +type openAIChatMessage struct { + Role string `json:"role"` + Content string `json:"content"` +} + +type openAIChatResponse struct { + Choices []struct { + Message openAIChatMessage `json:"message"` + } `json:"choices"` + Usage struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + TotalTokens int `json:"total_tokens"` + } `json:"usage"` +} diff --git a/internal/llm/openai_compatible_client_test.go b/internal/llm/openai_compatible_client_test.go new file mode 100644 index 0000000..0aa65cb --- /dev/null +++ b/internal/llm/openai_compatible_client_test.go @@ -0,0 +1,289 @@ +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) + } +}