diff --git a/CHANGELOG.md b/CHANGELOG.md index 7401930..78dc237 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -18,6 +18,17 @@ and the project adheres to [Semantic Versioning 2.0.0](https://semver.org/spec/v comments is now derived from the PR's `url` field (which always points at the base repo, including cross-fork PRs). ### Added +- **Three new providers: DeepSeek, Mistral, Cohere.** Each is a standalone + provider package reusing the `openai-go` SDK pointed at the provider's + OpenAI-compatible endpoint — **no new dependency** (DeepSeek + `api.deepseek.com`, Mistral `api.mistral.ai/v1`, Cohere's + `compatibility/v1`). API keys via config or `DEEPSEEK_API_KEY` / + `MISTRAL_API_KEY` / `COHERE_API_KEY`; all three appear in + `commitbrief setup`. Structured output is prompt-driven (no + `response_format`) since these providers' strict-JSON support varies — + the retry-once-then-degrade pipeline (ADR-0014) covers non-conforming + output, same as Ollama. Total live providers: **9** (4 API + these 3 + + 2 CLI-backed). - **`--suggest-commit`.** After the review, makes a second free-form provider call and prints a single Conventional Commit message for the staged diff to stdout. Read-only — it suggests, never writes git diff --git a/README.md b/README.md index e5c8bbb..646004a 100644 --- a/README.md +++ b/README.md @@ -198,6 +198,9 @@ Four API providers + two CLI-tool-backed providers ship in the box: | **Anthropic** | Claude Opus 4.7, Sonnet 4.6, Haiku 4.5 | Ephemeral prompt caching (5 m TTL) cuts repeated input cost ~10×. | | **OpenAI** | GPT-4o, GPT-4o-mini | Automatic prompt caching at ≥1024-token prefixes. | | **Google Gemini** | Gemini 2.5 Pro (2 M context!), 2.5 Flash, 1.5 Flash | Largest free-tier context windows. | +| **DeepSeek** | deepseek-chat, deepseek-reasoner | OpenAI-compatible API (`DEEPSEEK_API_KEY`); JSON is prompt-driven (degrades gracefully). | +| **Mistral** | Mistral Large / Small, Codestral | OpenAI-compatible API (`MISTRAL_API_KEY`). | +| **Cohere** | Command R+ / R, Command A | Cohere's OpenAI-compatibility endpoint (`COHERE_API_KEY`). | | **Ollama** | Whatever you've `ollama pull`'d | Local-only, no API key, no per-token cost. | | **`claude-cli`** | Whatever your local Claude Code uses | Subprocess of `claude -p -` — no API key on our side; reuses your Claude Code subscription. `commitbrief --cli claude --staged`. | | **`gemini-cli`** | Whatever your local Gemini CLI uses | Subprocess of `gemini -p` — no API key on our side; reuses your Gemini CLI auth. `commitbrief --cli gemini --staged`. | diff --git a/cmd/commitbrief/main.go b/cmd/commitbrief/main.go index 2ac3f7a..e9b40c4 100644 --- a/cmd/commitbrief/main.go +++ b/cmd/commitbrief/main.go @@ -13,8 +13,11 @@ import ( // a local subprocess rather than an HTTPS API. _ "github.com/CommitBrief/commitbrief/internal/provider/anthropic" _ "github.com/CommitBrief/commitbrief/internal/provider/claude-cli" + _ "github.com/CommitBrief/commitbrief/internal/provider/cohere" + _ "github.com/CommitBrief/commitbrief/internal/provider/deepseek" _ "github.com/CommitBrief/commitbrief/internal/provider/gemini" _ "github.com/CommitBrief/commitbrief/internal/provider/gemini-cli" + _ "github.com/CommitBrief/commitbrief/internal/provider/mistral" _ "github.com/CommitBrief/commitbrief/internal/provider/ollama" _ "github.com/CommitBrief/commitbrief/internal/provider/openai" ) diff --git a/internal/config/env.go b/internal/config/env.go index e7720a2..4f4a657 100644 --- a/internal/config/env.go +++ b/internal/config/env.go @@ -22,6 +22,15 @@ func ApplyEnv(c *Config) { if v := os.Getenv("GEMINI_API_KEY"); v != "" { setProviderField(c, "gemini", func(p *ProviderConfig) { p.APIKey = v }) } + if v := os.Getenv("DEEPSEEK_API_KEY"); v != "" { + setProviderField(c, "deepseek", func(p *ProviderConfig) { p.APIKey = v }) + } + if v := os.Getenv("MISTRAL_API_KEY"); v != "" { + setProviderField(c, "mistral", func(p *ProviderConfig) { p.APIKey = v }) + } + if v := os.Getenv("COHERE_API_KEY"); v != "" { + setProviderField(c, "cohere", func(p *ProviderConfig) { p.APIKey = v }) + } if v := os.Getenv("OLLAMA_HOST"); v != "" { setProviderField(c, "ollama", func(p *ProviderConfig) { p.BaseURL = v }) } diff --git a/internal/provider/cohere/client.go b/internal/provider/cohere/client.go new file mode 100644 index 0000000..008e558 --- /dev/null +++ b/internal/provider/cohere/client.go @@ -0,0 +1,157 @@ +// SPDX-License-Identifier: GPL-3.0-or-later + +// Package cohere implements the Provider interface against Cohere's +// OpenAI-compatibility endpoint (https://api.cohere.ai/compatibility/v1), +// reusing the openai-go SDK — no new dependency. Structured output is +// prompt-driven (no response_format); JSON shape comes from the system +// prompt's contract plus the retry-once-then-degrade pipeline (ADR-0014 +// §4), the same way Ollama works. +package cohere + +import ( + "context" + "errors" + "fmt" + + sdk "github.com/openai/openai-go" + "github.com/openai/openai-go/option" + "github.com/openai/openai-go/shared" + + "github.com/CommitBrief/commitbrief/internal/config" + "github.com/CommitBrief/commitbrief/internal/provider" + "github.com/CommitBrief/commitbrief/internal/tokens" +) + +const ( + defaultBaseURL = "https://api.cohere.ai/compatibility/v1" + defaultMaxTokens = 4096 + testPingPrompt = "ping" + testPingMaxTok = 8 +) + +type Client struct { + sdk sdk.Client + model string +} + +func New(cfg config.ProviderConfig) (provider.Provider, error) { + if cfg.APIKey == "" { + return nil, fmt.Errorf("cohere: %w", provider.ErrUnauthorized) + } + baseURL := cfg.BaseURL + if baseURL == "" { + baseURL = defaultBaseURL + } + return &Client{ + sdk: sdk.NewClient(option.WithAPIKey(cfg.APIKey), option.WithBaseURL(baseURL)), + model: cfg.Model, + }, nil +} + +func (c *Client) Name() string { return Name } + +func (c *Client) DefaultModel() string { + if c.model != "" { + return c.model + } + return DefaultModel +} + +func (c *Client) ContextWindow(model string) int { + if model == "" { + model = c.DefaultModel() + } + return contextWindowFor(model) +} + +func (c *Client) EstimateTokens(s string) int { return tokens.Estimate(s) } + +func (c *Client) Pricing(model string) provider.Pricing { + if model == "" { + model = c.DefaultModel() + } + return pricingFor(model) +} + +func (c *Client) Review(ctx context.Context, req provider.Request) (provider.Response, error) { + completion, err := c.sdk.Chat.Completions.New(ctx, c.buildParams(req)) + if err != nil { + return provider.Response{}, mapError(err) + } + return provider.Response{ + Content: extractText(completion), + Model: completion.Model, + Usage: mapUsage(completion.Usage), + }, nil +} + +func (c *Client) TestConnection(ctx context.Context) error { + _, err := c.sdk.Chat.Completions.New(ctx, sdk.ChatCompletionNewParams{ + Model: shared.ChatModel(c.DefaultModel()), + MaxCompletionTokens: sdk.Int(testPingMaxTok), + Messages: []sdk.ChatCompletionMessageParamUnion{sdk.UserMessage(testPingPrompt)}, + }) + return mapError(err) +} + +func (c *Client) buildParams(req provider.Request) sdk.ChatCompletionNewParams { + model := req.Model + if model == "" { + model = c.DefaultModel() + } + maxTokens := int64(req.MaxTokens) + if maxTokens <= 0 { + maxTokens = defaultMaxTokens + } + messages := make([]sdk.ChatCompletionMessageParamUnion, 0, 2) + if req.SystemPrompt != "" { + messages = append(messages, sdk.SystemMessage(req.SystemPrompt)) + } + messages = append(messages, sdk.UserMessage(req.UserPrompt)) + // No response_format — JSON is prompt-driven (retry/degrade covers + // non-conforming output). FreeForm (ADR-0015) is naturally satisfied. + return sdk.ChatCompletionNewParams{ + Model: shared.ChatModel(model), + MaxCompletionTokens: sdk.Int(maxTokens), + Messages: messages, + } +} + +func extractText(c *sdk.ChatCompletion) string { + if c == nil || len(c.Choices) == 0 { + return "" + } + return c.Choices[0].Message.Content +} + +func mapUsage(u sdk.CompletionUsage) provider.Usage { + return provider.Usage{ + InputTokens: int(u.PromptTokens), + OutputTokens: int(u.CompletionTokens), + CachedInputTokens: int(u.PromptTokensDetails.CachedTokens), + } +} + +func mapError(err error) error { + if err == nil { + return nil + } + var apiErr *sdk.Error + if errors.As(err, &apiErr) { + switch apiErr.StatusCode { + case 401, 403: + return fmt.Errorf("cohere: %w: %s", provider.ErrUnauthorized, apiErr.Error()) + case 429: + return fmt.Errorf("cohere: %w: %s", provider.ErrRateLimit, apiErr.Error()) + case 404: + return fmt.Errorf("cohere: %w: %s", provider.ErrModelNotSupported, apiErr.Error()) + } + } + return fmt.Errorf("cohere: %w", err) +} + +func init() { + provider.Register(Name, New) +} + +var _ provider.Provider = (*Client)(nil) diff --git a/internal/provider/cohere/cohere_test.go b/internal/provider/cohere/cohere_test.go new file mode 100644 index 0000000..74f873a --- /dev/null +++ b/internal/provider/cohere/cohere_test.go @@ -0,0 +1,95 @@ +// SPDX-License-Identifier: GPL-3.0-or-later + +package cohere + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/CommitBrief/commitbrief/internal/config" + "github.com/CommitBrief/commitbrief/internal/provider" +) + +func TestModelsAndSupport(t *testing.T) { + if len(Models()) != 3 { + t.Errorf("Models() = %v, want 3", Models()) + } + if !IsModelSupported(ModelCommandRPlus) || IsModelSupported("gpt-4o") { + t.Error("model support check wrong") + } + Models()[0] = "tampered" + if Models()[0] == "tampered" { + t.Error("Models() must return a defensive copy") + } +} + +func TestPricingAndContextWindow(t *testing.T) { + if p := pricingFor(ModelCommandRPlus); p.InputPer1M == 0 || p.OutputPer1M == 0 { + t.Errorf("command-r-plus pricing missing: %+v", p) + } + if pricingFor("unknown").InputPer1M != 0 { + t.Error("unknown model should yield zero pricing") + } + if contextWindowFor(ModelCommandA) != 256_000 { + t.Errorf("command-a context window wrong: %d", contextWindowFor(ModelCommandA)) + } + if contextWindowFor("unknown") != defaultContextWindow { + t.Error("unknown model should fall back to default context window") + } +} + +func TestNewMissingAPIKey(t *testing.T) { + if _, err := New(config.ProviderConfig{}); !errors.Is(err, provider.ErrUnauthorized) { + t.Errorf("err = %v, want ErrUnauthorized", err) + } +} + +func TestNewDefaults(t *testing.T) { + c, err := New(config.ProviderConfig{APIKey: "k"}) + if err != nil { + t.Fatal(err) + } + if c.Name() != Name || c.DefaultModel() != DefaultModel { + t.Errorf("Name/DefaultModel wrong: %q / %q", c.Name(), c.DefaultModel()) + } +} + +func TestRegisteredViaInit(t *testing.T) { + for _, n := range provider.Names() { + if n == Name { + return + } + } + t.Errorf("cohere not registered; Names() = %v", provider.Names()) +} + +func TestReviewWithFakeServer(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if !strings.Contains(r.URL.Path, "/chat/completions") { + http.NotFound(w, r) + return + } + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{ + "id": "x", "object": "chat.completion", "created": 1, "model": ModelCommandRPlus, + "choices": []map[string]any{{"index": 0, "finish_reason": "stop", + "message": map[string]any{"role": "assistant", "content": "cohere review"}}}, + "usage": map[string]any{"prompt_tokens": 25, "completion_tokens": 9, "total_tokens": 34}, + }) + })) + defer srv.Close() + + c, _ := New(config.ProviderConfig{APIKey: "k", BaseURL: srv.URL}) + resp, err := c.Review(context.Background(), provider.Request{UserPrompt: "diff", MaxTokens: 64}) + if err != nil { + t.Fatalf("Review: %v", err) + } + if resp.Content != "cohere review" || resp.Usage.InputTokens != 25 { + t.Errorf("resp = %q / %+v", resp.Content, resp.Usage) + } +} diff --git a/internal/provider/cohere/context_window.go b/internal/provider/cohere/context_window.go new file mode 100644 index 0000000..31a375d --- /dev/null +++ b/internal/provider/cohere/context_window.go @@ -0,0 +1,18 @@ +// SPDX-License-Identifier: GPL-3.0-or-later + +package cohere + +const defaultContextWindow = 128_000 + +var contextWindows = map[string]int{ + ModelCommandRPlus: 128_000, + ModelCommandR: 128_000, + ModelCommandA: 256_000, +} + +func contextWindowFor(model string) int { + if w, ok := contextWindows[model]; ok { + return w + } + return defaultContextWindow +} diff --git a/internal/provider/cohere/models.go b/internal/provider/cohere/models.go new file mode 100644 index 0000000..aa0e043 --- /dev/null +++ b/internal/provider/cohere/models.go @@ -0,0 +1,30 @@ +// SPDX-License-Identifier: GPL-3.0-or-later + +package cohere + +const ( + Name = "cohere" + + ModelCommandRPlus = "command-r-plus" + ModelCommandR = "command-r" + ModelCommandA = "command-a-03-2025" + + DefaultModel = ModelCommandRPlus +) + +var supportedModels = []string{ModelCommandRPlus, ModelCommandR, ModelCommandA} + +func Models() []string { + out := make([]string, len(supportedModels)) + copy(out, supportedModels) + return out +} + +func IsModelSupported(model string) bool { + for _, m := range supportedModels { + if m == model { + return true + } + } + return false +} diff --git a/internal/provider/cohere/pricing.go b/internal/provider/cohere/pricing.go new file mode 100644 index 0000000..c19cf03 --- /dev/null +++ b/internal/provider/cohere/pricing.go @@ -0,0 +1,22 @@ +// SPDX-License-Identifier: GPL-3.0-or-later + +package cohere + +import "github.com/CommitBrief/commitbrief/internal/provider" + +// Cohere per-1M-token pricing snapshot (USD). Source: +// https://cohere.com/pricing — refresh on price change. No automatic +// prompt-cache discount is surfaced through the compatibility endpoint, +// so CachedInputPer1M is left at 0. +var pricingTable = map[string]provider.Pricing{ + ModelCommandRPlus: {InputPer1M: 2.50, OutputPer1M: 10.00}, + ModelCommandR: {InputPer1M: 0.15, OutputPer1M: 0.60}, + ModelCommandA: {InputPer1M: 2.50, OutputPer1M: 10.00}, +} + +func pricingFor(model string) provider.Pricing { + if p, ok := pricingTable[model]; ok { + return p + } + return provider.Pricing{} +} diff --git a/internal/provider/deepseek/client.go b/internal/provider/deepseek/client.go new file mode 100644 index 0000000..662bb1a --- /dev/null +++ b/internal/provider/deepseek/client.go @@ -0,0 +1,159 @@ +// SPDX-License-Identifier: GPL-3.0-or-later + +// Package deepseek implements the Provider interface against DeepSeek's +// OpenAI-compatible Chat Completions API. It reuses the openai-go SDK +// pointed at DeepSeek's base URL — no new dependency. Structured output +// is prompt-driven (no response_format): DeepSeek's strict json_schema +// support is uneven, so JSON shape comes from the system prompt's +// contract plus the retry-once-then-degrade pipeline (ADR-0014 §4), the +// same way Ollama works. +package deepseek + +import ( + "context" + "errors" + "fmt" + + sdk "github.com/openai/openai-go" + "github.com/openai/openai-go/option" + "github.com/openai/openai-go/shared" + + "github.com/CommitBrief/commitbrief/internal/config" + "github.com/CommitBrief/commitbrief/internal/provider" + "github.com/CommitBrief/commitbrief/internal/tokens" +) + +const ( + defaultBaseURL = "https://api.deepseek.com" + defaultMaxTokens = 4096 + testPingPrompt = "ping" + testPingMaxTok = 8 +) + +type Client struct { + sdk sdk.Client + model string +} + +func New(cfg config.ProviderConfig) (provider.Provider, error) { + if cfg.APIKey == "" { + return nil, fmt.Errorf("deepseek: %w", provider.ErrUnauthorized) + } + baseURL := cfg.BaseURL + if baseURL == "" { + baseURL = defaultBaseURL + } + return &Client{ + sdk: sdk.NewClient(option.WithAPIKey(cfg.APIKey), option.WithBaseURL(baseURL)), + model: cfg.Model, + }, nil +} + +func (c *Client) Name() string { return Name } + +func (c *Client) DefaultModel() string { + if c.model != "" { + return c.model + } + return DefaultModel +} + +func (c *Client) ContextWindow(model string) int { + if model == "" { + model = c.DefaultModel() + } + return contextWindowFor(model) +} + +func (c *Client) EstimateTokens(s string) int { return tokens.Estimate(s) } + +func (c *Client) Pricing(model string) provider.Pricing { + if model == "" { + model = c.DefaultModel() + } + return pricingFor(model) +} + +func (c *Client) Review(ctx context.Context, req provider.Request) (provider.Response, error) { + completion, err := c.sdk.Chat.Completions.New(ctx, c.buildParams(req)) + if err != nil { + return provider.Response{}, mapError(err) + } + return provider.Response{ + Content: extractText(completion), + Model: completion.Model, + Usage: mapUsage(completion.Usage), + }, nil +} + +func (c *Client) TestConnection(ctx context.Context) error { + _, err := c.sdk.Chat.Completions.New(ctx, sdk.ChatCompletionNewParams{ + Model: shared.ChatModel(c.DefaultModel()), + MaxCompletionTokens: sdk.Int(testPingMaxTok), + Messages: []sdk.ChatCompletionMessageParamUnion{sdk.UserMessage(testPingPrompt)}, + }) + return mapError(err) +} + +func (c *Client) buildParams(req provider.Request) sdk.ChatCompletionNewParams { + model := req.Model + if model == "" { + model = c.DefaultModel() + } + maxTokens := int64(req.MaxTokens) + if maxTokens <= 0 { + maxTokens = defaultMaxTokens + } + messages := make([]sdk.ChatCompletionMessageParamUnion, 0, 2) + if req.SystemPrompt != "" { + messages = append(messages, sdk.SystemMessage(req.SystemPrompt)) + } + messages = append(messages, sdk.UserMessage(req.UserPrompt)) + // No response_format — JSON is prompt-driven (retry/degrade covers + // non-conforming output). FreeForm (ADR-0015) is naturally satisfied: + // there is no structured-output enforcement to skip. + return sdk.ChatCompletionNewParams{ + Model: shared.ChatModel(model), + MaxCompletionTokens: sdk.Int(maxTokens), + Messages: messages, + } +} + +func extractText(c *sdk.ChatCompletion) string { + if c == nil || len(c.Choices) == 0 { + return "" + } + return c.Choices[0].Message.Content +} + +func mapUsage(u sdk.CompletionUsage) provider.Usage { + return provider.Usage{ + InputTokens: int(u.PromptTokens), + OutputTokens: int(u.CompletionTokens), + CachedInputTokens: int(u.PromptTokensDetails.CachedTokens), + } +} + +func mapError(err error) error { + if err == nil { + return nil + } + var apiErr *sdk.Error + if errors.As(err, &apiErr) { + switch apiErr.StatusCode { + case 401, 403: + return fmt.Errorf("deepseek: %w: %s", provider.ErrUnauthorized, apiErr.Error()) + case 429: + return fmt.Errorf("deepseek: %w: %s", provider.ErrRateLimit, apiErr.Error()) + case 404: + return fmt.Errorf("deepseek: %w: %s", provider.ErrModelNotSupported, apiErr.Error()) + } + } + return fmt.Errorf("deepseek: %w", err) +} + +func init() { + provider.Register(Name, New) +} + +var _ provider.Provider = (*Client)(nil) diff --git a/internal/provider/deepseek/context_window.go b/internal/provider/deepseek/context_window.go new file mode 100644 index 0000000..c588035 --- /dev/null +++ b/internal/provider/deepseek/context_window.go @@ -0,0 +1,17 @@ +// SPDX-License-Identifier: GPL-3.0-or-later + +package deepseek + +const defaultContextWindow = 64_000 + +var contextWindows = map[string]int{ + ModelChat: 64_000, + ModelReasoner: 64_000, +} + +func contextWindowFor(model string) int { + if w, ok := contextWindows[model]; ok { + return w + } + return defaultContextWindow +} diff --git a/internal/provider/deepseek/deepseek_test.go b/internal/provider/deepseek/deepseek_test.go new file mode 100644 index 0000000..307d9d6 --- /dev/null +++ b/internal/provider/deepseek/deepseek_test.go @@ -0,0 +1,99 @@ +// SPDX-License-Identifier: GPL-3.0-or-later + +package deepseek + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/CommitBrief/commitbrief/internal/config" + "github.com/CommitBrief/commitbrief/internal/provider" +) + +func TestModelsAndSupport(t *testing.T) { + if len(Models()) != 2 { + t.Errorf("Models() = %v, want 2", Models()) + } + if !IsModelSupported(ModelChat) || IsModelSupported("gpt-4o") { + t.Error("model support check wrong") + } + Models()[0] = "tampered" + if Models()[0] == "tampered" { + t.Error("Models() must return a defensive copy") + } +} + +func TestPricingAndContextWindow(t *testing.T) { + p := pricingFor(ModelChat) + if p.InputPer1M == 0 || p.OutputPer1M == 0 { + t.Errorf("deepseek-chat pricing missing: %+v", p) + } + if pricingFor("unknown").InputPer1M != 0 { + t.Error("unknown model should yield zero pricing") + } + if contextWindowFor("unknown") != defaultContextWindow { + t.Error("unknown model should fall back to default context window") + } +} + +func TestNewMissingAPIKey(t *testing.T) { + if _, err := New(config.ProviderConfig{}); !errors.Is(err, provider.ErrUnauthorized) { + t.Errorf("err = %v, want ErrUnauthorized", err) + } +} + +func TestNewDefaultsBaseURLAndModel(t *testing.T) { + c, err := New(config.ProviderConfig{APIKey: "k"}) + if err != nil { + t.Fatal(err) + } + if c.Name() != Name { + t.Errorf("Name = %q", c.Name()) + } + if c.DefaultModel() != DefaultModel { + t.Errorf("DefaultModel = %q, want %q", c.DefaultModel(), DefaultModel) + } +} + +func TestRegisteredViaInit(t *testing.T) { + for _, n := range provider.Names() { + if n == Name { + return + } + } + t.Errorf("deepseek not registered; Names() = %v", provider.Names()) +} + +func TestReviewWithFakeServer(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if !strings.Contains(r.URL.Path, "/chat/completions") { + http.NotFound(w, r) + return + } + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{ + "id": "x", "object": "chat.completion", "created": 1, "model": ModelChat, + "choices": []map[string]any{{"index": 0, "finish_reason": "stop", + "message": map[string]any{"role": "assistant", "content": "deepseek review"}}}, + "usage": map[string]any{"prompt_tokens": 30, "completion_tokens": 12, "total_tokens": 42}, + }) + })) + defer srv.Close() + + c, _ := New(config.ProviderConfig{APIKey: "k", BaseURL: srv.URL}) + resp, err := c.Review(context.Background(), provider.Request{UserPrompt: "diff", MaxTokens: 64}) + if err != nil { + t.Fatalf("Review: %v", err) + } + if resp.Content != "deepseek review" { + t.Errorf("Content = %q", resp.Content) + } + if resp.Usage.InputTokens != 30 || resp.Usage.OutputTokens != 12 { + t.Errorf("Usage = %+v", resp.Usage) + } +} diff --git a/internal/provider/deepseek/models.go b/internal/provider/deepseek/models.go new file mode 100644 index 0000000..3c94aae --- /dev/null +++ b/internal/provider/deepseek/models.go @@ -0,0 +1,29 @@ +// SPDX-License-Identifier: GPL-3.0-or-later + +package deepseek + +const ( + Name = "deepseek" + + ModelChat = "deepseek-chat" + ModelReasoner = "deepseek-reasoner" + + DefaultModel = ModelChat +) + +var supportedModels = []string{ModelChat, ModelReasoner} + +func Models() []string { + out := make([]string, len(supportedModels)) + copy(out, supportedModels) + return out +} + +func IsModelSupported(model string) bool { + for _, m := range supportedModels { + if m == model { + return true + } + } + return false +} diff --git a/internal/provider/deepseek/pricing.go b/internal/provider/deepseek/pricing.go new file mode 100644 index 0000000..b51198a --- /dev/null +++ b/internal/provider/deepseek/pricing.go @@ -0,0 +1,21 @@ +// SPDX-License-Identifier: GPL-3.0-or-later + +package deepseek + +import "github.com/CommitBrief/commitbrief/internal/provider" + +// DeepSeek per-1M-token pricing snapshot (USD, standard/cache-miss rates). +// Source: https://api-docs.deepseek.com — refresh on price change. +// CachedInputPer1M reflects DeepSeek's cache-hit discount; whether the +// OpenAI-compatible usage payload reports cached tokens varies by model. +var pricingTable = map[string]provider.Pricing{ + ModelChat: {InputPer1M: 0.27, OutputPer1M: 1.10, CachedInputPer1M: 0.07}, + ModelReasoner: {InputPer1M: 0.55, OutputPer1M: 2.19, CachedInputPer1M: 0.14}, +} + +func pricingFor(model string) provider.Pricing { + if p, ok := pricingTable[model]; ok { + return p + } + return provider.Pricing{} +} diff --git a/internal/provider/mistral/client.go b/internal/provider/mistral/client.go new file mode 100644 index 0000000..f4fe85d --- /dev/null +++ b/internal/provider/mistral/client.go @@ -0,0 +1,157 @@ +// SPDX-License-Identifier: GPL-3.0-or-later + +// Package mistral implements the Provider interface against Mistral's +// OpenAI-compatible Chat Completions API (https://api.mistral.ai/v1), +// reusing the openai-go SDK — no new dependency. Structured output is +// prompt-driven (no response_format); JSON shape comes from the system +// prompt's contract plus the retry-once-then-degrade pipeline (ADR-0014 +// §4), the same way Ollama works. +package mistral + +import ( + "context" + "errors" + "fmt" + + sdk "github.com/openai/openai-go" + "github.com/openai/openai-go/option" + "github.com/openai/openai-go/shared" + + "github.com/CommitBrief/commitbrief/internal/config" + "github.com/CommitBrief/commitbrief/internal/provider" + "github.com/CommitBrief/commitbrief/internal/tokens" +) + +const ( + defaultBaseURL = "https://api.mistral.ai/v1" + defaultMaxTokens = 4096 + testPingPrompt = "ping" + testPingMaxTok = 8 +) + +type Client struct { + sdk sdk.Client + model string +} + +func New(cfg config.ProviderConfig) (provider.Provider, error) { + if cfg.APIKey == "" { + return nil, fmt.Errorf("mistral: %w", provider.ErrUnauthorized) + } + baseURL := cfg.BaseURL + if baseURL == "" { + baseURL = defaultBaseURL + } + return &Client{ + sdk: sdk.NewClient(option.WithAPIKey(cfg.APIKey), option.WithBaseURL(baseURL)), + model: cfg.Model, + }, nil +} + +func (c *Client) Name() string { return Name } + +func (c *Client) DefaultModel() string { + if c.model != "" { + return c.model + } + return DefaultModel +} + +func (c *Client) ContextWindow(model string) int { + if model == "" { + model = c.DefaultModel() + } + return contextWindowFor(model) +} + +func (c *Client) EstimateTokens(s string) int { return tokens.Estimate(s) } + +func (c *Client) Pricing(model string) provider.Pricing { + if model == "" { + model = c.DefaultModel() + } + return pricingFor(model) +} + +func (c *Client) Review(ctx context.Context, req provider.Request) (provider.Response, error) { + completion, err := c.sdk.Chat.Completions.New(ctx, c.buildParams(req)) + if err != nil { + return provider.Response{}, mapError(err) + } + return provider.Response{ + Content: extractText(completion), + Model: completion.Model, + Usage: mapUsage(completion.Usage), + }, nil +} + +func (c *Client) TestConnection(ctx context.Context) error { + _, err := c.sdk.Chat.Completions.New(ctx, sdk.ChatCompletionNewParams{ + Model: shared.ChatModel(c.DefaultModel()), + MaxCompletionTokens: sdk.Int(testPingMaxTok), + Messages: []sdk.ChatCompletionMessageParamUnion{sdk.UserMessage(testPingPrompt)}, + }) + return mapError(err) +} + +func (c *Client) buildParams(req provider.Request) sdk.ChatCompletionNewParams { + model := req.Model + if model == "" { + model = c.DefaultModel() + } + maxTokens := int64(req.MaxTokens) + if maxTokens <= 0 { + maxTokens = defaultMaxTokens + } + messages := make([]sdk.ChatCompletionMessageParamUnion, 0, 2) + if req.SystemPrompt != "" { + messages = append(messages, sdk.SystemMessage(req.SystemPrompt)) + } + messages = append(messages, sdk.UserMessage(req.UserPrompt)) + // No response_format — JSON is prompt-driven (retry/degrade covers + // non-conforming output). FreeForm (ADR-0015) is naturally satisfied. + return sdk.ChatCompletionNewParams{ + Model: shared.ChatModel(model), + MaxCompletionTokens: sdk.Int(maxTokens), + Messages: messages, + } +} + +func extractText(c *sdk.ChatCompletion) string { + if c == nil || len(c.Choices) == 0 { + return "" + } + return c.Choices[0].Message.Content +} + +func mapUsage(u sdk.CompletionUsage) provider.Usage { + return provider.Usage{ + InputTokens: int(u.PromptTokens), + OutputTokens: int(u.CompletionTokens), + CachedInputTokens: int(u.PromptTokensDetails.CachedTokens), + } +} + +func mapError(err error) error { + if err == nil { + return nil + } + var apiErr *sdk.Error + if errors.As(err, &apiErr) { + switch apiErr.StatusCode { + case 401, 403: + return fmt.Errorf("mistral: %w: %s", provider.ErrUnauthorized, apiErr.Error()) + case 429: + return fmt.Errorf("mistral: %w: %s", provider.ErrRateLimit, apiErr.Error()) + case 404: + return fmt.Errorf("mistral: %w: %s", provider.ErrModelNotSupported, apiErr.Error()) + } + } + return fmt.Errorf("mistral: %w", err) +} + +func init() { + provider.Register(Name, New) +} + +var _ provider.Provider = (*Client)(nil) diff --git a/internal/provider/mistral/context_window.go b/internal/provider/mistral/context_window.go new file mode 100644 index 0000000..0cf9432 --- /dev/null +++ b/internal/provider/mistral/context_window.go @@ -0,0 +1,18 @@ +// SPDX-License-Identifier: GPL-3.0-or-later + +package mistral + +const defaultContextWindow = 128_000 + +var contextWindows = map[string]int{ + ModelLarge: 128_000, + ModelSmall: 32_000, + ModelCodestral: 256_000, +} + +func contextWindowFor(model string) int { + if w, ok := contextWindows[model]; ok { + return w + } + return defaultContextWindow +} diff --git a/internal/provider/mistral/mistral_test.go b/internal/provider/mistral/mistral_test.go new file mode 100644 index 0000000..c4a5d59 --- /dev/null +++ b/internal/provider/mistral/mistral_test.go @@ -0,0 +1,95 @@ +// SPDX-License-Identifier: GPL-3.0-or-later + +package mistral + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/CommitBrief/commitbrief/internal/config" + "github.com/CommitBrief/commitbrief/internal/provider" +) + +func TestModelsAndSupport(t *testing.T) { + if len(Models()) != 3 { + t.Errorf("Models() = %v, want 3", Models()) + } + if !IsModelSupported(ModelLarge) || IsModelSupported("gpt-4o") { + t.Error("model support check wrong") + } + Models()[0] = "tampered" + if Models()[0] == "tampered" { + t.Error("Models() must return a defensive copy") + } +} + +func TestPricingAndContextWindow(t *testing.T) { + if p := pricingFor(ModelLarge); p.InputPer1M == 0 || p.OutputPer1M == 0 { + t.Errorf("mistral-large pricing missing: %+v", p) + } + if pricingFor("unknown").InputPer1M != 0 { + t.Error("unknown model should yield zero pricing") + } + if contextWindowFor(ModelCodestral) != 256_000 { + t.Errorf("codestral context window wrong: %d", contextWindowFor(ModelCodestral)) + } + if contextWindowFor("unknown") != defaultContextWindow { + t.Error("unknown model should fall back to default context window") + } +} + +func TestNewMissingAPIKey(t *testing.T) { + if _, err := New(config.ProviderConfig{}); !errors.Is(err, provider.ErrUnauthorized) { + t.Errorf("err = %v, want ErrUnauthorized", err) + } +} + +func TestNewDefaults(t *testing.T) { + c, err := New(config.ProviderConfig{APIKey: "k"}) + if err != nil { + t.Fatal(err) + } + if c.Name() != Name || c.DefaultModel() != DefaultModel { + t.Errorf("Name/DefaultModel wrong: %q / %q", c.Name(), c.DefaultModel()) + } +} + +func TestRegisteredViaInit(t *testing.T) { + for _, n := range provider.Names() { + if n == Name { + return + } + } + t.Errorf("mistral not registered; Names() = %v", provider.Names()) +} + +func TestReviewWithFakeServer(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if !strings.Contains(r.URL.Path, "/chat/completions") { + http.NotFound(w, r) + return + } + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{ + "id": "x", "object": "chat.completion", "created": 1, "model": ModelLarge, + "choices": []map[string]any{{"index": 0, "finish_reason": "stop", + "message": map[string]any{"role": "assistant", "content": "mistral review"}}}, + "usage": map[string]any{"prompt_tokens": 40, "completion_tokens": 20, "total_tokens": 60}, + }) + })) + defer srv.Close() + + c, _ := New(config.ProviderConfig{APIKey: "k", BaseURL: srv.URL}) + resp, err := c.Review(context.Background(), provider.Request{UserPrompt: "diff", MaxTokens: 64}) + if err != nil { + t.Fatalf("Review: %v", err) + } + if resp.Content != "mistral review" || resp.Usage.OutputTokens != 20 { + t.Errorf("resp = %q / %+v", resp.Content, resp.Usage) + } +} diff --git a/internal/provider/mistral/models.go b/internal/provider/mistral/models.go new file mode 100644 index 0000000..dbb73b3 --- /dev/null +++ b/internal/provider/mistral/models.go @@ -0,0 +1,30 @@ +// SPDX-License-Identifier: GPL-3.0-or-later + +package mistral + +const ( + Name = "mistral" + + ModelLarge = "mistral-large-latest" + ModelSmall = "mistral-small-latest" + ModelCodestral = "codestral-latest" + + DefaultModel = ModelLarge +) + +var supportedModels = []string{ModelLarge, ModelSmall, ModelCodestral} + +func Models() []string { + out := make([]string, len(supportedModels)) + copy(out, supportedModels) + return out +} + +func IsModelSupported(model string) bool { + for _, m := range supportedModels { + if m == model { + return true + } + } + return false +} diff --git a/internal/provider/mistral/pricing.go b/internal/provider/mistral/pricing.go new file mode 100644 index 0000000..6636b1b --- /dev/null +++ b/internal/provider/mistral/pricing.go @@ -0,0 +1,22 @@ +// SPDX-License-Identifier: GPL-3.0-or-later + +package mistral + +import "github.com/CommitBrief/commitbrief/internal/provider" + +// Mistral per-1M-token pricing snapshot (USD). Source: +// https://mistral.ai/pricing — refresh on price change. Mistral does not +// expose an automatic prompt-cache discount in the OpenAI-compatible +// usage payload, so CachedInputPer1M is left at 0. +var pricingTable = map[string]provider.Pricing{ + ModelLarge: {InputPer1M: 2.00, OutputPer1M: 6.00}, + ModelSmall: {InputPer1M: 0.20, OutputPer1M: 0.60}, + ModelCodestral: {InputPer1M: 0.30, OutputPer1M: 0.90}, +} + +func pricingFor(model string) provider.Pricing { + if p, ok := pricingTable[model]; ok { + return p + } + return provider.Pricing{} +} diff --git a/internal/setup/wizard.go b/internal/setup/wizard.go index d2bbc67..065e9b6 100644 --- a/internal/setup/wizard.go +++ b/internal/setup/wizard.go @@ -61,6 +61,27 @@ var DefaultSpecs = []ProviderSpec{ Models: []string{"gemini-2.5-pro", "gemini-2.5-flash", "gemini-1.5-flash"}, APIKeyHelp: "Get an API key from https://aistudio.google.com/", }, + { + Name: "deepseek", + Label: "DeepSeek", + NeedsKey: true, + Models: []string{"deepseek-chat", "deepseek-reasoner"}, + APIKeyHelp: "Get an API key from https://platform.deepseek.com/", + }, + { + Name: "mistral", + Label: "Mistral", + NeedsKey: true, + Models: []string{"mistral-large-latest", "mistral-small-latest", "codestral-latest"}, + APIKeyHelp: "Get an API key from https://console.mistral.ai/", + }, + { + Name: "cohere", + Label: "Cohere", + NeedsKey: true, + Models: []string{"command-r-plus", "command-r", "command-a-03-2025"}, + APIKeyHelp: "Get an API key from https://dashboard.cohere.com/", + }, { Name: "ollama", Label: "Ollama (local, no API key needed)",