From 1ffc1a148fc721e02349b2bb28eb0386baee85d9 Mon Sep 17 00:00:00 2001 From: ShubhScript Date: Thu, 24 Sep 2026 22:03:55 +0530 Subject: [PATCH] feat(ai): add multi-provider support for OpenAI, Claude, Ollama and Gemini 3.x --- DEVELOPMENT.md | 5 +- README.md | 30 +- internal/ai/base.go | 177 ++++++++++ internal/ai/claude/client.go | 132 +++++++ internal/ai/gemini/client.go | 165 +-------- internal/ai/openai/client.go | 170 +++++++++ internal/ai/provider.go | 56 ++- internal/cli/ai.go | 34 ++ internal/cli/ask.go | 16 +- internal/cli/config.go | 604 +++++++++++++++++++++++++++++--- internal/cli/root.go | 11 +- internal/cli/shell.go | 54 +++ internal/config/config.go | 185 +++++++++- internal/storage/sqlite.go | 6 + tests/unit/ai_providers_test.go | 171 +++++++++ 15 files changed, 1577 insertions(+), 239 deletions(-) create mode 100644 internal/ai/base.go create mode 100644 internal/ai/claude/client.go create mode 100644 internal/ai/openai/client.go create mode 100644 internal/cli/ai.go create mode 100644 tests/unit/ai_providers_test.go diff --git a/DEVELOPMENT.md b/DEVELOPMENT.md index 3f38469..ae63ad7 100644 --- a/DEVELOPMENT.md +++ b/DEVELOPMENT.md @@ -42,8 +42,11 @@ sql-doctor/ │ ├── migration/ │ │ └── analyzer.go # Migration lock risk & destructive check │ ├── ai/ -│ │ ├── provider.go # AIProvider interface +│ │ ├── provider.go # AIProvider interface & curated models +│ │ ├── base.go # BaseProvider & shared prompt templates │ │ ├── gemini/ # Google GenAI SDK client +│ │ ├── openai/ # OpenAI & Ollama/local LLM client +│ │ ├── claude/ # Anthropic Claude Messages API client │ │ └── context/ # Schema context minifier & prompt builder │ ├── storage/ │ │ └── sqlite.go # Local SQLite state repo (~/.sql-doctor/) diff --git a/README.md b/README.md index 423bfde..50624b7 100644 --- a/README.md +++ b/README.md @@ -268,20 +268,34 @@ sql-doctor format "select id,name from users where status='active' and age>21 or --- -### 10. Optional Gemini AI Assistant -If you want AI explanations or natural-language query generation, add your own Gemini API key: +### 10. Multi-Model AI Assistant (Gemini, OpenAI, Claude, Ollama) +If you want AI explanations or natural-language query generation, SQL Doctor supports **Google Gemini**, **OpenAI (ChatGPT)**, **Anthropic Claude**, and **Ollama / Local LLMs** (OpenAI-compatible): ```bash -# Configure your API key -sql-doctor config set-ai-key - -# Verify configuration +# Open the interactive AI configuration dashboard sql-doctor config ai -# Ask questions grounded in your schema +# Or switch provider and model directly: +sql-doctor config ai switch openai gpt-4o-mini +sql-doctor config ai switch claude claude-3-5-haiku-20241022 +sql-doctor config ai switch gemini gemini-3.8-flash +sql-doctor config ai switch ollama deepseek-r1:8b + +# Configure API keys (prompts with masked input if key omitted): +sql-doctor config ai set-key openai +sql-doctor config ai set-key claude +sql-doctor config ai set-key gemini + +# Set custom endpoint for Ollama / local models: +sql-doctor config ai set-endpoint http://localhost:11434/v1 + +# Test connection and latency: +sql-doctor config ai test + +# Ask questions grounded in your schema: sql-doctor ask "Which tables store customer billing records?" -# Generate queries +# Generate queries (with interactive execution prompt): sql-doctor ask "Write a query to find the top 5 customers by revenue this year" ``` diff --git a/internal/ai/base.go b/internal/ai/base.go new file mode 100644 index 0000000..b5e2832 --- /dev/null +++ b/internal/ai/base.go @@ -0,0 +1,177 @@ +package ai + +import ( + "context" + "fmt" + "strings" + + "github.com/sql-doctor/sql-doctor/internal/query/analyzer" +) + +// BaseCaller is a function that makes the actual API call to the LLM backend +type BaseCaller func(ctx context.Context, systemPrompt, userPrompt string) (string, error) + +// BaseProvider implements the standard AIProvider methods on top of a BaseCaller +type BaseProvider struct { + Name string + ModelName string + Configured bool + Caller BaseCaller +} + +func (b *BaseProvider) IsConfigured() bool { + return b.Configured +} + +func (b *BaseProvider) ProviderName() string { + return b.Name +} + +func (b *BaseProvider) Model() string { + return b.ModelName +} + +func (b *BaseProvider) TestConnection(ctx context.Context) error { + if !b.Configured { + return fmt.Errorf("%s is not configured (missing API key or endpoint)", b.Name) + } + res, err := b.Caller(ctx, "You are a test ping agent.", "Reply with the single word 'PONG' only.") + if err != nil { + return err + } + if strings.TrimSpace(res) == "" { + return fmt.Errorf("received empty response from %s", b.Name) + } + return nil +} + +func (b *BaseProvider) ExplainQuery(ctx context.Context, sqlQuery string, metrics *analyzer.QueryAnalysisResult) (string, error) { + sys := "You are a Senior Principal Database Performance Engineer. Provide clear, concise, actionable query analysis." + var metricsStr string + if metrics != nil { + metricsStr = fmt.Sprintf("Execution Time: %.2fms\nRows Examined: %d\nRows Returned: %d\nFull Table Scan: %v\nPerformance Score: %d/100", + metrics.ExecutionTimeMs, metrics.RowsExamined, metrics.RowsReturned, metrics.HasFullTableScan, metrics.PerformanceScore) + } + + prompt := fmt.Sprintf(`Explain the execution behavior and performance characteristics of this SQL query: + +SQL: +%s + +Observed Metrics: +%s + +Explain in 2-3 concise paragraphs: +1. What the query is doing logically. +2. Why the query is fast or slow based on the observed metrics. +3. Specific actionable steps to improve it.`, sqlQuery, metricsStr) + + return b.Caller(ctx, sys, prompt) +} + +func (b *BaseProvider) OptimizeQuery(ctx context.Context, sqlQuery string, schemaContext string, metrics *analyzer.QueryAnalysisResult) (string, error) { + sys := "You are an expert SQL Query Optimizer. Always prioritize index selection, sargability, and deterministic query rewrites." + prompt := fmt.Sprintf(`Analyze and provide concrete optimization recommendations for this SQL query: + +QUERY: +%s + +SCHEMA CONTEXT: +%s + +Provide: +1. Suggested optimized SQL rewrite. +2. Any recommended composite or single-column indexes with exact DDL. +3. Rationale explaining why the rewrite is faster.`, sqlQuery, schemaContext) + + return b.Caller(ctx, sys, prompt) +} + +func (b *BaseProvider) ReviewSchema(ctx context.Context, schemaSummary string) (string, error) { + sys := "You are a Senior Database Architect. Review database schema design, normalization, relationships, and index strategy." + prompt := fmt.Sprintf(`Review this database schema and identify design smells, missing constraints, or performance hazards: + +%s + +Provide: +1. Architectural strengths and design quality evaluation. +2. High-priority schema risks or normalization smells. +3. Recommended improvements.`, schemaSummary) + + return b.Caller(ctx, sys, prompt) +} + +func (b *BaseProvider) Ask(ctx context.Context, question string, schemaContext string) (string, error) { + sys := "You are SQL Doctor, an intelligent database diagnostics and engineering assistant. Ground your answer strictly in the provided database schema." + prompt := fmt.Sprintf(`Question: %s + +Connected Database Schema: +%s + +Answer the user's question clearly and provide relevant SQL snippets or explanations based strictly on the provided schema.`, question, schemaContext) + + return b.Caller(ctx, sys, prompt) +} + +func (b *BaseProvider) GenerateSQL(ctx context.Context, userGoal string, schemaContext string) (*GeneratedSQL, error) { + sys := "You are an expert SQL developer. Generate valid, high-performance SQL based strictly on the provided schema. Output only the SQL query and brief rationale." + prompt := fmt.Sprintf(`User Goal: %s + +Schema: +%s + +Generate the optimal SQL query to accomplish the user's goal. +Format your output as: +---SQL--- + +---EXPLANATION--- + +`, userGoal, schemaContext) + + raw, err := b.Caller(ctx, sys, prompt) + if err != nil { + return nil, err + } + + result := &GeneratedSQL{} + if strings.Contains(raw, "---SQL---") { + parts := strings.Split(raw, "---SQL---") + if len(parts) > 1 { + subParts := strings.Split(parts[1], "---EXPLANATION---") + result.SQL = CleanCodeBlock(subParts[0]) + if len(subParts) > 1 { + result.Explanation = strings.TrimSpace(subParts[1]) + } + } + } else { + result.SQL = CleanCodeBlock(raw) + } + + upper := strings.ToUpper(result.SQL) + if strings.Contains(upper, "DELETE") || strings.Contains(upper, "UPDATE") || strings.Contains(upper, "DROP") || strings.Contains(upper, "TRUNCATE") { + result.IsDestructive = true + } + + return result, nil +} + +func (b *BaseProvider) SummarizeDoctor(ctx context.Context, doctorSummary string) (string, error) { + sys := "You are a Database Reliability Engineer. Provide an executive summary of database health findings." + prompt := fmt.Sprintf(`Summarize these database diagnostic findings into an executive briefing with prioritized action items: + +%s`, doctorSummary) + + return b.Caller(ctx, sys, prompt) +} + +// CleanCodeBlock strips markdown backticks from generated code +func CleanCodeBlock(s string) string { + s = strings.TrimSpace(s) + if strings.HasPrefix(s, "```sql") { + s = strings.TrimPrefix(s, "```sql") + } else if strings.HasPrefix(s, "```") { + s = strings.TrimPrefix(s, "```") + } + s = strings.TrimSuffix(s, "```") + return strings.TrimSpace(s) +} diff --git a/internal/ai/claude/client.go b/internal/ai/claude/client.go new file mode 100644 index 0000000..4e30e4e --- /dev/null +++ b/internal/ai/claude/client.go @@ -0,0 +1,132 @@ +package claude + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "time" + + "github.com/sql-doctor/sql-doctor/internal/ai" +) + +// Client implements ai.AIProvider for Anthropic Claude via the Messages API +type Client struct { + ai.BaseProvider + apiKey string + client *http.Client +} + +type messagePayload struct { + Role string `json:"role"` + Content string `json:"content"` +} + +type claudeRequest struct { + Model string `json:"model"` + MaxTokens int `json:"max_tokens"` + System string `json:"system,omitempty"` + Messages []messagePayload `json:"messages"` + Temperature float64 `json:"temperature"` +} + +type claudeResponse struct { + Content []struct { + Type string `json:"type"` + Text string `json:"text"` + } `json:"content"` + Error *struct { + Type string `json:"type"` + Message string `json:"message"` + } `json:"error"` +} + +// New creates an Anthropic Claude provider client +func New(apiKey, modelName string) *Client { + if modelName == "" { + modelName = "claude-3-5-haiku-20241022" + } + + c := &Client{ + apiKey: apiKey, + client: &http.Client{ + Timeout: 60 * time.Second, + }, + } + + c.BaseProvider = ai.BaseProvider{ + Name: "Anthropic Claude", + ModelName: modelName, + Configured: strings.TrimSpace(apiKey) != "", + Caller: c.callClaude, + } + + return c +} + +func (c *Client) callClaude(ctx context.Context, systemPrompt, userPrompt string) (string, error) { + if !c.IsConfigured() { + return "", fmt.Errorf("Anthropic Claude API key is not configured. Set key using: sql-doctor config ai set-key claude or set ANTHROPIC_API_KEY environment variable") + } + + reqBody := claudeRequest{ + Model: c.ModelName, + MaxTokens: 4096, + System: systemPrompt, + Messages: []messagePayload{ + {Role: "user", Content: userPrompt}, + }, + Temperature: 0.2, + } + + jsonBytes, err := json.Marshal(reqBody) + if err != nil { + return "", fmt.Errorf("failed to marshal Claude request: %w", err) + } + + url := "https://api.anthropic.com/v1/messages" + req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewBuffer(jsonBytes)) + if err != nil { + return "", fmt.Errorf("failed to create request: %w", err) + } + + req.Header.Set("Content-Type", "application/json") + req.Header.Set("x-api-key", c.apiKey) + req.Header.Set("anthropic-version", "2023-06-01") + + resp, err := c.client.Do(req) + if err != nil { + return "", fmt.Errorf("Anthropic Claude connection failed: %w", err) + } + defer resp.Body.Close() + + bodyBytes, err := io.ReadAll(resp.Body) + if err != nil { + return "", fmt.Errorf("failed to read response: %w", err) + } + + var clResp claudeResponse + if err := json.Unmarshal(bodyBytes, &clResp); err != nil { + return "", fmt.Errorf("Anthropic Claude returned non-JSON response (HTTP %d): %s", resp.StatusCode, string(bodyBytes)) + } + + if clResp.Error != nil && clResp.Error.Message != "" { + return "", fmt.Errorf("Anthropic Claude API error: %s", clResp.Error.Message) + } + + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + return "", fmt.Errorf("Anthropic Claude HTTP error %d: %s", resp.StatusCode, string(bodyBytes)) + } + + var sb strings.Builder + for _, block := range clResp.Content { + if block.Type == "text" { + sb.WriteString(block.Text) + } + } + + return strings.TrimSpace(sb.String()), nil +} diff --git a/internal/ai/gemini/client.go b/internal/ai/gemini/client.go index b58e490..9371a20 100644 --- a/internal/ai/gemini/client.go +++ b/internal/ai/gemini/client.go @@ -7,41 +7,34 @@ import ( "github.com/google/generative-ai-go/genai" "github.com/sql-doctor/sql-doctor/internal/ai" - "github.com/sql-doctor/sql-doctor/internal/query/analyzer" "google.golang.org/api/option" ) // Client implements ai.AIProvider via the official Google Gemini Go SDK type Client struct { - apiKey string - modelName string + ai.BaseProvider + apiKey string } func New(apiKey, modelName string) *Client { if modelName == "" { - modelName = "gemini-2.5-flash" + modelName = "gemini-3.8-flash" } - return &Client{ - apiKey: apiKey, - modelName: modelName, + c := &Client{ + apiKey: apiKey, } -} - -func (c *Client) IsConfigured() bool { - return strings.TrimSpace(c.apiKey) != "" -} - -func (c *Client) Model() string { - return c.modelName -} - -func (c *Client) missingKeyErr() error { - return fmt.Errorf("AI features are unavailable because a Gemini API key has not been configured.\n\nConfigure your key using:\n sql-doctor config set-ai-key\nor set the GEMINI_API_KEY environment variable") + c.BaseProvider = ai.BaseProvider{ + Name: "Google Gemini", + ModelName: modelName, + Configured: strings.TrimSpace(apiKey) != "", + Caller: c.callGemini, + } + return c } func (c *Client) callGemini(ctx context.Context, systemPrompt, userPrompt string) (string, error) { if !c.IsConfigured() { - return "", c.missingKeyErr() + return "", fmt.Errorf("AI features are unavailable because a Gemini API key has not been configured.\n\nConfigure your key using:\n sql-doctor config ai set-key gemini \nor set the GEMINI_API_KEY environment variable") } client, err := genai.NewClient(ctx, option.WithAPIKey(c.apiKey)) @@ -50,7 +43,7 @@ func (c *Client) callGemini(ctx context.Context, systemPrompt, userPrompt string } defer client.Close() - model := client.GenerativeModel(c.modelName) + model := client.GenerativeModel(c.ModelName) if systemPrompt != "" { model.SystemInstruction = &genai.Content{ Parts: []genai.Part{genai.Text(systemPrompt)}, @@ -75,133 +68,3 @@ func (c *Client) callGemini(ctx context.Context, systemPrompt, userPrompt string return strings.TrimSpace(sb.String()), nil } - -func (c *Client) ExplainQuery(ctx context.Context, sqlQuery string, metrics *analyzer.QueryAnalysisResult) (string, error) { - sys := "You are a Senior Principal Database Performance Engineer. Provide clear, concise, actionable query analysis." - var metricsStr string - if metrics != nil { - metricsStr = fmt.Sprintf("Execution Time: %.2fms\nRows Examined: %d\nRows Returned: %d\nFull Table Scan: %v\nPerformance Score: %d/100", - metrics.ExecutionTimeMs, metrics.RowsExamined, metrics.RowsReturned, metrics.HasFullTableScan, metrics.PerformanceScore) - } - - prompt := fmt.Sprintf(`Explain the execution behavior and performance characteristics of this SQL query: - -SQL: -%s - -Observed Metrics: -%s - -Explain in 2-3 concise paragraphs: -1. What the query is doing logically. -2. Why the query is fast or slow based on the observed metrics. -3. Specific actionable steps to improve it.`, sqlQuery, metricsStr) - - return c.callGemini(ctx, sys, prompt) -} - -func (c *Client) OptimizeQuery(ctx context.Context, sqlQuery string, schemaContext string, metrics *analyzer.QueryAnalysisResult) (string, error) { - sys := "You are an expert SQL Query Optimizer. Always prioritize index selection, sargability, and deterministic query rewrites." - prompt := fmt.Sprintf(`Analyze and provide concrete optimization recommendations for this SQL query: - -QUERY: -%s - -SCHEMA CONTEXT: -%s - -Provide: -1. Suggested optimized SQL rewrite. -2. Any recommended composite or single-column indexes with exact DDL. -3. Rationale explaining why the rewrite is faster.`, sqlQuery, schemaContext) - - return c.callGemini(ctx, sys, prompt) -} - -func (c *Client) ReviewSchema(ctx context.Context, schemaSummary string) (string, error) { - sys := "You are a Senior Database Architect. Review database schema design, normalization, relationships, and index strategy." - prompt := fmt.Sprintf(`Review this database schema and identify design smells, missing constraints, or performance hazards: - -%s - -Provide: -1. Architectural strengths and design quality evaluation. -2. High-priority schema risks or normalization smells. -3. Recommended improvements.`, schemaSummary) - - return c.callGemini(ctx, sys, prompt) -} - -func (c *Client) Ask(ctx context.Context, question string, schemaContext string) (string, error) { - sys := "You are SQL Doctor, an intelligent database diagnostics and engineering assistant. Ground your answer strictly in the provided database schema." - prompt := fmt.Sprintf(`Question: %s - -Connected Database Schema: -%s - -Answer the user's question clearly and provide relevant SQL snippets or explanations based strictly on the provided schema.`, question, schemaContext) - - return c.callGemini(ctx, sys, prompt) -} - -func (c *Client) GenerateSQL(ctx context.Context, userGoal string, schemaContext string) (*ai.GeneratedSQL, error) { - sys := "You are an expert SQL developer. Generate valid, high-performance SQL based strictly on the provided schema. Output only the SQL query and brief rationale." - prompt := fmt.Sprintf(`User Goal: %s - -Schema: -%s - -Generate the optimal SQL query to accomplish the user's goal. -Format your output as: ----SQL--- - ----EXPLANATION--- - -`, userGoal, schemaContext) - - raw, err := c.callGemini(ctx, sys, prompt) - if err != nil { - return nil, err - } - - result := &ai.GeneratedSQL{} - if strings.Contains(raw, "---SQL---") { - parts := strings.Split(raw, "---SQL---") - if len(parts) > 1 { - subParts := strings.Split(parts[1], "---EXPLANATION---") - result.SQL = cleanCodeBlock(subParts[0]) - if len(subParts) > 1 { - result.Explanation = strings.TrimSpace(subParts[1]) - } - } - } else { - result.SQL = cleanCodeBlock(raw) - } - - upper := strings.ToUpper(result.SQL) - if strings.Contains(upper, "DELETE") || strings.Contains(upper, "UPDATE") || strings.Contains(upper, "DROP") || strings.Contains(upper, "TRUNCATE") { - result.IsDestructive = true - } - - return result, nil -} - -func (c *Client) SummarizeDoctor(ctx context.Context, doctorSummary string) (string, error) { - sys := "You are a Database Reliability Engineer. Provide an executive summary of database health findings." - prompt := fmt.Sprintf(`Summarize these database diagnostic findings into an executive briefing with prioritized action items: - -%s`, doctorSummary) - - return c.callGemini(ctx, sys, prompt) -} - -func cleanCodeBlock(s string) string { - s = strings.TrimSpace(s) - if strings.HasPrefix(s, "```sql") { - s = strings.TrimPrefix(s, "```sql") - } else if strings.HasPrefix(s, "```") { - s = strings.TrimPrefix(s, "```") - } - s = strings.TrimSuffix(s, "```") - return strings.TrimSpace(s) -} diff --git a/internal/ai/openai/client.go b/internal/ai/openai/client.go new file mode 100644 index 0000000..48ebace --- /dev/null +++ b/internal/ai/openai/client.go @@ -0,0 +1,170 @@ +package openai + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "time" + + "github.com/sql-doctor/sql-doctor/internal/ai" +) + +// Client implements ai.AIProvider for OpenAI and OpenAI-compatible local APIs (Ollama, vLLM, etc.) +type Client struct { + ai.BaseProvider + apiKey string + endpoint string + isOllama bool + client *http.Client +} + +// Request payload for OpenAI Chat Completions API +type chatRequest struct { + Model string `json:"model"` + Messages []chatMessage `json:"messages"` + Temperature float64 `json:"temperature"` +} + +type chatMessage struct { + Role string `json:"role"` + Content string `json:"content"` +} + +// Response structure for Chat Completions API +type chatResponse struct { + Choices []struct { + Message struct { + Content string `json:"content"` + } `json:"message"` + } `json:"choices"` + Error *struct { + Message string `json:"message"` + Type string `json:"type"` + } `json:"error"` +} + +// New creates an OpenAI or Ollama provider client +func New(apiKey, modelName, endpoint string, isOllama bool) *Client { + providerName := "OpenAI" + if isOllama { + providerName = "Ollama (Local)" + if endpoint == "" { + endpoint = "http://localhost:11434/v1" + } + if modelName == "" { + modelName = "deepseek-r1:8b" + } + } else { + if endpoint == "" { + endpoint = "https://api.openai.com/v1" + } + if modelName == "" { + modelName = "gpt-4o-mini" + } + } + + endpoint = strings.TrimSuffix(endpoint, "/") + + configured := false + if isOllama { + // Ollama is configured if an endpoint is provided (keys are usually optional) + configured = endpoint != "" + } else { + configured = strings.TrimSpace(apiKey) != "" + } + + c := &Client{ + apiKey: apiKey, + endpoint: endpoint, + isOllama: isOllama, + client: &http.Client{ + Timeout: 60 * time.Second, + }, + } + + c.BaseProvider = ai.BaseProvider{ + Name: providerName, + ModelName: modelName, + Configured: configured, + Caller: c.callOpenAI, + } + + return c +} + +func (c *Client) callOpenAI(ctx context.Context, systemPrompt, userPrompt string) (string, error) { + if !c.IsConfigured() { + if c.isOllama { + return "", fmt.Errorf("Ollama is not configured. Set endpoint using: sql-doctor config ai set-endpoint ") + } + return "", fmt.Errorf("OpenAI API key is not configured. Set key using: sql-doctor config ai set-key openai or set OPENAI_API_KEY environment variable") + } + + messages := make([]chatMessage, 0, 2) + if systemPrompt != "" { + messages = append(messages, chatMessage{ + Role: "system", + Content: systemPrompt, + }) + } + messages = append(messages, chatMessage{ + Role: "user", + Content: userPrompt, + }) + + reqBody := chatRequest{ + Model: c.ModelName, + Messages: messages, + Temperature: 0.2, + } + + jsonBytes, err := json.Marshal(reqBody) + if err != nil { + return "", fmt.Errorf("failed to marshal OpenAI request: %w", err) + } + + url := c.endpoint + "/chat/completions" + req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewBuffer(jsonBytes)) + if err != nil { + return "", fmt.Errorf("failed to create request: %w", err) + } + + req.Header.Set("Content-Type", "application/json") + if c.apiKey != "" { + req.Header.Set("Authorization", "Bearer "+c.apiKey) + } + + resp, err := c.client.Do(req) + if err != nil { + return "", fmt.Errorf("%s connection failed: %w", c.Name, err) + } + defer resp.Body.Close() + + bodyBytes, err := io.ReadAll(resp.Body) + if err != nil { + return "", fmt.Errorf("failed to read response: %w", err) + } + + var chatResp chatResponse + if err := json.Unmarshal(bodyBytes, &chatResp); err != nil { + return "", fmt.Errorf("%s returned non-JSON response (HTTP %d): %s", c.Name, resp.StatusCode, string(bodyBytes)) + } + + if chatResp.Error != nil && chatResp.Error.Message != "" { + return "", fmt.Errorf("%s API error: %s", c.Name, chatResp.Error.Message) + } + + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + return "", fmt.Errorf("%s HTTP error %d: %s", c.Name, resp.StatusCode, string(bodyBytes)) + } + + if len(chatResp.Choices) == 0 { + return "", fmt.Errorf("%s returned no completion choices", c.Name) + } + + return strings.TrimSpace(chatResp.Choices[0].Message.Content), nil +} diff --git a/internal/ai/provider.go b/internal/ai/provider.go index 85464a1..83547d8 100644 --- a/internal/ai/provider.go +++ b/internal/ai/provider.go @@ -6,18 +6,68 @@ import ( "github.com/sql-doctor/sql-doctor/internal/query/analyzer" ) +// Supported Provider Constants +const ( + ProviderGemini = "gemini" + ProviderOpenAI = "openai" + ProviderClaude = "claude" + ProviderOllama = "ollama" +) + +// CuratedModel describes a recommended model option +type CuratedModel struct { + ID string + Name string + Description string + Recommended bool +} + +// CuratedModelsByProvider lists recommended models that are well-balanced for SQL tasks +var CuratedModelsByProvider = map[string][]CuratedModel{ + ProviderGemini: { + {ID: "gemini-3.8-flash", Name: "Gemini 3.8 Flash", Description: "Fast, state-of-the-art reasoning for SQL, code & schema diagnosis", Recommended: true}, + {ID: "gemini-3.5-flash", Name: "Gemini 3.5 Flash", Description: "Balanced high-performance model for diagnostics & query generation"}, + {ID: "gemini-3.5-flash-lite", Name: "Gemini 3.5 Flash Lite", Description: "Lightweight, ultra-fast latency and cost-efficient"}, + }, + ProviderOpenAI: { + {ID: "gpt-4o-mini", Name: "GPT-4o Mini", Description: "Fast, cost-effective, exceptional SQL generation & accuracy", Recommended: true}, + {ID: "gpt-4o", Name: "GPT-4o", Description: "Flagship model with deep multi-step reasoning"}, + {ID: "o3-mini", Name: "o3 Mini", Description: "Advanced reasoning model for complex optimization math"}, + }, + ProviderClaude: { + {ID: "claude-3-5-haiku-20241022", Name: "Claude 3.5 Haiku", Description: "Ultra-fast, cost-effective, precise SQL syntax generation", Recommended: true}, + {ID: "claude-3-7-sonnet-20250219", Name: "Claude 3.7 Sonnet", Description: "Hybrid reasoning & high-precision architectural analysis"}, + {ID: "claude-3-5-sonnet-20241022", Name: "Claude 3.5 Sonnet", Description: "Industry-leading coding & schema architectural analysis"}, + }, + ProviderOllama: { + {ID: "deepseek-r1:8b", Name: "DeepSeek R1 (8B)", Description: "Open-weights reasoning model running locally", Recommended: true}, + {ID: "qwen2.5-coder:7b", Name: "Qwen 2.5 Coder (7B)", Description: "Optimized specifically for code and SQL queries"}, + {ID: "llama3.1:8b", Name: "Llama 3.1 (8B)", Description: "Popular general-purpose local LLM"}, + }, +} + // GeneratedSQL contains generated SQL and safety precautions type GeneratedSQL struct { - SQL string `json:"sql"` - Explanation string `json:"explanation"` + SQL string `json:"sql"` + Explanation string `json:"explanation"` IsDestructive bool `json:"is_destructive"` - Assumptions string `json:"assumptions"` + Assumptions string `json:"assumptions"` +} + +// ProviderParams defines parameters for initializing an AIProvider +type ProviderParams struct { + Provider string + Model string + APIKey string + Endpoint string } // AIProvider defines the conversational and AI query generation interface type AIProvider interface { IsConfigured() bool + ProviderName() string Model() string + TestConnection(ctx context.Context) error ExplainQuery(ctx context.Context, sqlQuery string, metrics *analyzer.QueryAnalysisResult) (string, error) OptimizeQuery(ctx context.Context, sqlQuery string, schemaContext string, metrics *analyzer.QueryAnalysisResult) (string, error) ReviewSchema(ctx context.Context, schemaSummary string) (string, error) diff --git a/internal/cli/ai.go b/internal/cli/ai.go new file mode 100644 index 0000000..fd22e16 --- /dev/null +++ b/internal/cli/ai.go @@ -0,0 +1,34 @@ +package cli + +import ( + "strings" + + "github.com/sql-doctor/sql-doctor/internal/ai" + "github.com/sql-doctor/sql-doctor/internal/ai/claude" + "github.com/sql-doctor/sql-doctor/internal/ai/gemini" + "github.com/sql-doctor/sql-doctor/internal/ai/openai" +) + +// NewAIProvider creates an AIProvider instance based on provider parameters +func NewAIProvider(p ai.ProviderParams) ai.AIProvider { + provider := strings.ToLower(strings.TrimSpace(p.Provider)) + if provider == "" { + provider = ai.ProviderGemini + } + + switch provider { + case ai.ProviderOpenAI: + return openai.New(p.APIKey, p.Model, p.Endpoint, false) + + case ai.ProviderClaude, "anthropic": + return claude.New(p.APIKey, p.Model) + + case ai.ProviderOllama, "local": + return openai.New(p.APIKey, p.Model, p.Endpoint, true) + + case ai.ProviderGemini, "google": + fallthrough + default: + return gemini.New(p.APIKey, p.Model) + } +} diff --git a/internal/cli/ask.go b/internal/cli/ask.go index 03baa08..90d4f47 100644 --- a/internal/cli/ask.go +++ b/internal/cli/ask.go @@ -20,7 +20,7 @@ var ( var askCmd = &cobra.Command{ Use: "ask \"\"", - Short: "Ask questions about your database or generate SQL using Gemini AI", + Short: "Ask questions about your database or generate SQL using AI", Args: cobra.ExactArgs(1), RunE: func(cmd *cobra.Command, args []string) error { ctx := cmd.Context() @@ -28,11 +28,11 @@ var askCmd = &cobra.Command{ provider := GetAIProvider() if !provider.IsConfigured() { - fmt.Println(ui.Warning("AI features are unavailable because a Gemini API key has not been configured.")) - fmt.Println("\nConfigure your own key using:") - fmt.Println(" sql-doctor config set-ai-key ") - fmt.Println("or set the environment variable:") - fmt.Println(" export GEMINI_API_KEY=\"...\"") + fmt.Println(ui.Warning("AI features are unavailable because %s is not configured.", provider.ProviderName())) + fmt.Println("\nConfigure your AI provider and credentials using:") + fmt.Println(" sql-doctor config ai") + fmt.Println("or switch provider:") + fmt.Println(" sql-doctor config ai switch ") return nil } @@ -58,7 +58,7 @@ var askCmd = &cobra.Command{ strings.HasPrefix(lowerPrompt, "get") || strings.Contains(lowerPrompt, "query to") if isQueryGen { - fmt.Println(ui.Info("Generating SQL with Gemini %s...", provider.Model())) + fmt.Println(ui.Info("Generating SQL with %s (%s)...", provider.ProviderName(), provider.Model())) generated, err := provider.GenerateSQL(ctx, prompt, schemaContext) if err != nil { return err @@ -112,7 +112,7 @@ var askCmd = &cobra.Command{ } // Conversational Question - fmt.Println(ui.Info("Consulting Gemini %s...", provider.Model())) + fmt.Println(ui.Info("Consulting %s (%s)...", provider.ProviderName(), provider.Model())) answer, err := provider.Ask(ctx, prompt, schemaContext) if err != nil { return err diff --git a/internal/cli/config.go b/internal/cli/config.go index 5194b9e..6b46b82 100644 --- a/internal/cli/config.go +++ b/internal/cli/config.go @@ -1,12 +1,19 @@ package cli import ( + "bufio" + "context" "fmt" + "os" + "strconv" "strings" + "time" "github.com/spf13/cobra" + "github.com/sql-doctor/sql-doctor/internal/ai" "github.com/sql-doctor/sql-doctor/internal/config" "github.com/sql-doctor/sql-doctor/internal/ui" + "golang.org/x/term" ) var configCmd = &cobra.Command{ @@ -16,105 +23,600 @@ var configCmd = &cobra.Command{ var configAICmd = &cobra.Command{ Use: "ai", - Short: "Show AI configuration status", + Short: "View and manage multi-model AI configuration", + Long: `Display the AI configuration dashboard across all supported providers (Gemini, OpenAI, Claude, Ollama), +switch active models, configure credentials, and run connectivity tests.`, RunE: func(cmd *cobra.Command, args []string) error { + ctx := cmd.Context() + if appConfig == nil { + var err error + appConfig, err = config.LoadConfig(ctx, appStorage) + if err != nil { + return err + } + } + provider := GetAIProvider() configured := provider.IsConfigured() + providerList := []map[string]interface{}{ + { + "id": ai.ProviderGemini, + "name": "Google Gemini", + "active": appConfig.AIProvider == ai.ProviderGemini || appConfig.AIProvider == "", + "model": appConfig.GeminiModel, + "configured": appConfig.GeminiKey != "", + "key": maskKey(appConfig.GeminiKey), + "source": keySource(appConfig.GeminiKey, "GEMINI_API_KEY"), + }, + { + "id": ai.ProviderOpenAI, + "name": "OpenAI (ChatGPT)", + "active": appConfig.AIProvider == ai.ProviderOpenAI, + "model": appConfig.OpenAIModel, + "configured": appConfig.OpenAIKey != "", + "key": maskKey(appConfig.OpenAIKey), + "source": keySource(appConfig.OpenAIKey, "OPENAI_API_KEY"), + }, + { + "id": ai.ProviderClaude, + "name": "Anthropic Claude", + "active": appConfig.AIProvider == ai.ProviderClaude || appConfig.AIProvider == "anthropic", + "model": appConfig.ClaudeModel, + "configured": appConfig.ClaudeKey != "", + "key": maskKey(appConfig.ClaudeKey), + "source": keySource(appConfig.ClaudeKey, "ANTHROPIC_API_KEY", "CLAUDE_API_KEY"), + }, + { + "id": ai.ProviderOllama, + "name": "Ollama (Local / OpenAI-compatible)", + "active": appConfig.AIProvider == ai.ProviderOllama || appConfig.AIProvider == "local", + "model": appConfig.OllamaModel, + "configured": appConfig.OllamaEndpoint != "", + "key": maskKey(appConfig.OllamaKey), + "endpoint": appConfig.OllamaEndpoint, + "source": appConfig.OllamaEndpoint, + }, + } + OutputResult(map[string]interface{}{ - "ai_enabled": configured, - "model": provider.Model(), - "provider": "Google Gemini", - "key_source": resolveKeySource(), + "active_provider": appConfig.AIProvider, + "active_model": provider.Model(), + "ai_ready": configured, + "providers": providerList, }, func() { - fmt.Println(ui.TitleStyle.Render("SQL Doctor — AI Configuration Status")) - if configured { - fmt.Println(ui.Success("Gemini AI is configured and ready.")) - fmt.Printf(" Provider: Google Gemini\n") - fmt.Printf(" Model: %s\n", provider.Model()) - fmt.Printf(" Source: %s\n", resolveKeySource()) - } else { - fmt.Println(ui.Warning("Gemini AI is NOT configured.")) - fmt.Println("\nTo enable AI query explanations, natural language questions, and optimizations:") - fmt.Println(" sql-doctor config set-ai-key ") - fmt.Println("or set the environment variable:") - fmt.Println(" export GEMINI_API_KEY=\"...\"") + renderAIDashboard(appConfig, provider, providerList) + }) + + // If running in an interactive terminal and not JSON, offer action menu + if !flagJSON && term.IsTerminal(int(os.Stdin.Fd())) { + reader := bufio.NewReader(os.Stdin) + return promptInteractiveAIMenu(ctx, reader) + } + + return nil + }, +} + +func renderAIDashboard(cfg *config.Config, activeProvider ai.AIProvider, list []map[string]interface{}) { + fmt.Println() + fmt.Println(ui.TitleStyle.Render("SQL Doctor — AI Configuration Dashboard")) + + activeName := "Google Gemini" + for _, p := range list { + if p["active"].(bool) { + activeName = p["name"].(string) + break + } + } + + statusBadge := ui.CriticalBadge + if activeProvider.IsConfigured() { + statusBadge = ui.SuccessBadge + } + + fmt.Printf("Active Engine: %s %s (Model: %s)\n\n", + ui.HeaderStyle.Render(activeName), + statusBadge, + ui.WarningBadge+" "+activeProvider.Model(), + ) + + tbl := ui.NewTable("ACTIVE", "PROVIDER", "STATUS", "ACTIVE MODEL", "SOURCE / ENDPOINT", "API KEY") + for _, p := range list { + marker := "" + if p["active"].(bool) { + marker = " ★ " + } + status := ui.Error("Not Set") + if p["configured"].(bool) { + status = ui.Success("Ready") + } + endpointOrSrc := p["source"].(string) + if ep, ok := p["endpoint"].(string); ok && ep != "" { + endpointOrSrc = ep + } + tbl.AddRow(marker, p["name"].(string), status, p["model"].(string), endpointOrSrc, p["key"].(string)) + } + fmt.Println(tbl.Render()) + fmt.Println() +} + +func promptInteractiveAIMenu(ctx context.Context, reader *bufio.Reader) error { + fmt.Println(ui.HeaderStyle.Render("Select an action:")) + fmt.Println(" [1] Switch active provider & model") + fmt.Println(" [2] Configure API key / endpoint") + fmt.Println(" [3] Test AI connection") + fmt.Println(" [4] Delete / Clear API key") + fmt.Println(" [5] Exit") + fmt.Println() + fmt.Print("Enter choice [1-5] (default 5): ") + + input, _ := reader.ReadString('\n') + choice := strings.TrimSpace(input) + if choice == "" || choice == "5" || choice == "exit" || choice == "q" { + return nil + } + + switch choice { + case "1": + return runInteractiveSwitch(ctx, reader) + case "2": + return runInteractiveSetKey(ctx, reader) + case "3": + return runTestConnection(ctx, appConfig.AIProvider) + case "4": + return runInteractiveClearKey(ctx, reader) + default: + fmt.Println(ui.Warning("Invalid option.")) + return nil + } +} + +func runInteractiveSwitch(ctx context.Context, reader *bufio.Reader) error { + fmt.Println() + fmt.Println(ui.TitleStyle.Render("Select AI Provider to Activate:")) + fmt.Println(" [1] Google Gemini") + fmt.Println(" [2] OpenAI (ChatGPT)") + fmt.Println(" [3] Anthropic Claude") + fmt.Println(" [4] Ollama (Local / OpenAI-compatible)") + fmt.Println() + fmt.Print("Enter choice [1-4] (default 1): ") + + input, _ := reader.ReadString('\n') + c := strings.TrimSpace(input) + if c == "" { + c = "1" + } + + var targetProvider string + switch c { + case "1": + targetProvider = ai.ProviderGemini + case "2": + targetProvider = ai.ProviderOpenAI + case "3": + targetProvider = ai.ProviderClaude + case "4": + targetProvider = ai.ProviderOllama + default: + return fmt.Errorf("invalid provider selection") + } + + curated := ai.CuratedModelsByProvider[targetProvider] + fmt.Println() + fmt.Println(ui.TitleStyle.Render("Select Model for " + targetProvider + ":")) + for i, m := range curated { + rec := "" + if m.Recommended { + rec = " (Recommended)" + } + fmt.Printf(" [%d] %-28s %s%s\n", i+1, m.ID, m.Description, rec) + } + customIdx := len(curated) + 1 + fmt.Printf(" [%d] Custom model name...\n", customIdx) + fmt.Println() + fmt.Printf("Enter choice [1-%d] (default 1): ", customIdx) + + mInput, _ := reader.ReadString('\n') + mChoice := strings.TrimSpace(mInput) + if mChoice == "" { + mChoice = "1" + } + + var targetModel string + idx, err := strconv.Atoi(mChoice) + if err == nil && idx >= 1 && idx <= len(curated) { + targetModel = curated[idx-1].ID + } else if idx == customIdx { + fmt.Print("Enter custom model name: ") + custInput, _ := reader.ReadString('\n') + targetModel = strings.TrimSpace(custInput) + } else { + targetModel = curated[0].ID + } + + if err := config.SetAIProvider(ctx, appStorage, targetProvider); err != nil { + return err + } + if err := config.SetAIModel(ctx, appStorage, targetProvider, targetModel); err != nil { + return err + } + + // Reload config + appConfig, _ = config.LoadConfig(ctx, appStorage) + fmt.Println() + fmt.Println(ui.Success("Active AI provider switched to '%s' with model '%s'", targetProvider, targetModel)) + return nil +} + +func runInteractiveSetKey(ctx context.Context, reader *bufio.Reader) error { + fmt.Println() + fmt.Println(ui.TitleStyle.Render("Configure Credentials for Provider:")) + fmt.Println(" [1] Google Gemini") + fmt.Println(" [2] OpenAI") + fmt.Println(" [3] Anthropic Claude") + fmt.Println(" [4] Ollama (Local / Custom Endpoint)") + fmt.Println() + fmt.Print("Enter choice [1-4] (default 1): ") + + input, _ := reader.ReadString('\n') + c := strings.TrimSpace(input) + if c == "" { + c = "1" + } + + var targetProvider string + switch c { + case "1": + targetProvider = ai.ProviderGemini + case "2": + targetProvider = ai.ProviderOpenAI + case "3": + targetProvider = ai.ProviderClaude + case "4": + targetProvider = ai.ProviderOllama + default: + return fmt.Errorf("invalid provider selection") + } + + if targetProvider == ai.ProviderOllama { + defaultEndpoint := "http://localhost:11434/v1" + if appConfig.OllamaEndpoint != "" { + defaultEndpoint = appConfig.OllamaEndpoint + } + fmt.Printf("Ollama API Endpoint [%s]: ", defaultEndpoint) + epInput, _ := reader.ReadString('\n') + endpoint := strings.TrimSpace(epInput) + if endpoint == "" { + endpoint = defaultEndpoint + } + if err := config.SetAIEndpoint(ctx, appStorage, endpoint); err != nil { + return err + } + fmt.Println(ui.Success("Ollama endpoint set to '%s'", endpoint)) + return nil + } + + fmt.Printf("Enter %s API Key (input will be hidden): ", targetProvider) + byteKey, err := term.ReadPassword(int(os.Stdin.Fd())) + fmt.Println() + if err != nil { + return fmt.Errorf("failed to read API key: %w", err) + } + + key := strings.TrimSpace(string(byteKey)) + if key == "" { + fmt.Println(ui.Warning("API key was empty. Nothing saved.")) + return nil + } + + if err := config.SetAIKey(ctx, appStorage, targetProvider, key); err != nil { + return err + } + + appConfig, _ = config.LoadConfig(ctx, appStorage) + fmt.Println(ui.Success("%s API key saved successfully (%s)", targetProvider, maskKey(key))) + return nil +} + +func runInteractiveClearKey(ctx context.Context, reader *bufio.Reader) error { + fmt.Println() + fmt.Println(ui.TitleStyle.Render("Select API Key to Clear / Delete:")) + fmt.Println(" [1] Google Gemini") + fmt.Println(" [2] OpenAI") + fmt.Println(" [3] Anthropic Claude") + fmt.Println(" [4] Ollama (Endpoint & Key)") + fmt.Println() + fmt.Print("Enter choice [1-4]: ") + + input, _ := reader.ReadString('\n') + c := strings.TrimSpace(input) + var targetProvider string + switch c { + case "1": + targetProvider = ai.ProviderGemini + case "2": + targetProvider = ai.ProviderOpenAI + case "3": + targetProvider = ai.ProviderClaude + case "4": + targetProvider = ai.ProviderOllama + default: + return fmt.Errorf("invalid choice") + } + + if err := config.ClearAIKey(ctx, appStorage, targetProvider); err != nil { + return err + } + + appConfig, _ = config.LoadConfig(ctx, appStorage) + fmt.Println(ui.Success("Cleared stored credentials for %s.", targetProvider)) + return nil +} + +func runTestConnection(ctx context.Context, providerName string) error { + if providerName == "" { + providerName = appConfig.AIProvider + } + if providerName == "" { + providerName = ai.ProviderGemini + } + + fmt.Printf("Testing connection to %s...\n", ui.HeaderStyle.Render(providerName)) + p := GetAIProvider() + if !p.IsConfigured() { + return fmt.Errorf("provider %s is not configured with an API key or endpoint", providerName) + } + + start := time.Now() + err := p.TestConnection(ctx) + duration := time.Since(start) + + if err != nil { + fmt.Println(ui.Error("Connection test failed: %v", err)) + return err + } + + fmt.Println(ui.Success("Connected to %s (%s) successfully! Latency: %.2f ms", p.ProviderName(), p.Model(), float64(duration.Microseconds())/1000.0)) + return nil +} + +// Subcommands for direct CLI usage + +var configAISwitchCmd = &cobra.Command{ + Use: "switch [provider] [model]", + Short: "Switch active AI provider and model (gemini, openai, claude, ollama)", + Args: cobra.RangeArgs(0, 2), + RunE: func(cmd *cobra.Command, args []string) error { + ctx := cmd.Context() + if len(args) == 0 { + reader := bufio.NewReader(os.Stdin) + return runInteractiveSwitch(ctx, reader) + } + + targetProvider := args[0] + targetModel := "" + if len(args) > 1 { + targetModel = args[1] + } + + if err := config.SetAIProvider(ctx, appStorage, targetProvider); err != nil { + return err + } + if targetModel != "" { + if err := config.SetAIModel(ctx, appStorage, targetProvider, targetModel); err != nil { + return err } + } + + appConfig, _ = config.LoadConfig(ctx, appStorage) + OutputResult(map[string]string{ + "provider": targetProvider, + "model": appConfig.ActiveAIParams().Model, + }, func() { + fmt.Println(ui.Success("Switched AI provider to '%s' (model: %s)", targetProvider, appConfig.ActiveAIParams().Model)) }) return nil }, } -var configSetAIKeyCmd = &cobra.Command{ - Use: "set-ai-key ", - Short: "Configure your Gemini API key", - Args: cobra.ExactArgs(1), +var configAISetKeyCmd = &cobra.Command{ + Use: "set-key [provider] [key]", + Short: "Configure API key for an AI provider (gemini, openai, claude)", + Args: cobra.RangeArgs(0, 2), RunE: func(cmd *cobra.Command, args []string) error { ctx := cmd.Context() - if appStorage == nil { - return fmt.Errorf("local storage unavailable") + targetProvider := appConfig.AIProvider + var key string + + if len(args) == 1 { + // If 1 arg, check if it's a provider name or key + lower := strings.ToLower(args[0]) + if lower == "gemini" || lower == "openai" || lower == "claude" || lower == "ollama" { + targetProvider = lower + } else { + key = args[0] + } + } else if len(args) >= 2 { + targetProvider = args[0] + key = args[1] } - key := strings.TrimSpace(args[0]) if key == "" { - return fmt.Errorf("API key cannot be empty") + if !term.IsTerminal(int(os.Stdin.Fd())) { + return fmt.Errorf("API key required. Usage: sql-doctor config ai set-key ") + } + fmt.Printf("Enter %s API Key (input will be hidden): ", targetProvider) + byteKey, err := term.ReadPassword(int(os.Stdin.Fd())) + fmt.Println() + if err != nil { + return err + } + key = strings.TrimSpace(string(byteKey)) + if key == "" { + return fmt.Errorf("API key cannot be empty") + } } - if err := config.SetGeminiKey(ctx, appStorage, key); err != nil { - return fmt.Errorf("failed to save API key: %w", err) + if err := config.SetAIKey(ctx, appStorage, targetProvider, key); err != nil { + return err } - masked := key - if len(key) > 8 { - masked = key[:4] + strings.Repeat("*", len(key)-8) + key[len(key)-4:] + appConfig, _ = config.LoadConfig(ctx, appStorage) + OutputResult(map[string]string{ + "provider": targetProvider, + "key": maskKey(key), + }, func() { + fmt.Println(ui.Success("%s API key saved successfully (%s)", targetProvider, maskKey(key))) + }) + return nil + }, +} + +var configAISetModelCmd = &cobra.Command{ + Use: "set-model [provider] ", + Short: "Set the model name for a provider", + Args: cobra.RangeArgs(1, 2), + RunE: func(cmd *cobra.Command, args []string) error { + ctx := cmd.Context() + targetProvider := appConfig.AIProvider + targetModel := "" + + if len(args) == 1 { + targetModel = args[0] + } else { + targetProvider = args[0] + targetModel = args[1] } - OutputResult(map[string]interface{}{ - "status": "SAVED", - "key": masked, + if err := config.SetAIModel(ctx, appStorage, targetProvider, targetModel); err != nil { + return err + } + + appConfig, _ = config.LoadConfig(ctx, appStorage) + OutputResult(map[string]string{ + "provider": targetProvider, + "model": targetModel, }, func() { - fmt.Println(ui.Success("Gemini API key saved successfully (%s)", masked)) + fmt.Println(ui.Success("%s model set to '%s'", targetProvider, targetModel)) }) return nil }, } -var configSetModelCmd = &cobra.Command{ - Use: "set-model ", - Short: "Set the Gemini model name (default: gemini-2.5-flash)", +var configAISetEndpointCmd = &cobra.Command{ + Use: "set-endpoint ", + Short: "Set custom endpoint URL for Ollama / local LLM (e.g. http://localhost:11434/v1)", Args: cobra.ExactArgs(1), RunE: func(cmd *cobra.Command, args []string) error { ctx := cmd.Context() - if appStorage == nil { - return fmt.Errorf("local storage unavailable") + endpoint := strings.TrimSpace(args[0]) + if err := config.SetAIEndpoint(ctx, appStorage, endpoint); err != nil { + return err + } + + appConfig, _ = config.LoadConfig(ctx, appStorage) + OutputResult(map[string]string{ + "endpoint": endpoint, + }, func() { + fmt.Println(ui.Success("Ollama / local LLM endpoint set to '%s'", endpoint)) + }) + return nil + }, +} + +var configAITestCmd = &cobra.Command{ + Use: "test [provider]", + Short: "Test connection and latency to active or specified AI provider", + Args: cobra.RangeArgs(0, 1), + RunE: func(cmd *cobra.Command, args []string) error { + target := appConfig.AIProvider + if len(args) > 0 { + target = args[0] + } + return runTestConnection(cmd.Context(), target) + }, +} + +var configAIClearCmd = &cobra.Command{ + Use: "clear [provider]", + Aliases: []string{"delete", "reset"}, + Short: "Clear saved API credentials for an AI provider", + Args: cobra.RangeArgs(0, 1), + RunE: func(cmd *cobra.Command, args []string) error { + ctx := cmd.Context() + target := appConfig.AIProvider + if len(args) > 0 { + target = args[0] } - model := strings.TrimSpace(args[0]) - if err := config.SetGeminiModel(ctx, appStorage, model); err != nil { - return fmt.Errorf("failed to save model: %w", err) + if err := config.ClearAIKey(ctx, appStorage, target); err != nil { + return err } - OutputResult(map[string]interface{}{ - "status": "SAVED", - "model": model, + appConfig, _ = config.LoadConfig(ctx, appStorage) + OutputResult(map[string]string{ + "cleared": target, }, func() { - fmt.Println(ui.Success("Gemini model set to '%s'", model)) + fmt.Println(ui.Success("Cleared stored credentials for %s.", target)) }) return nil }, } -func resolveKeySource() string { - if appConfig == nil { - return "None" +// Backward compatibility commands + +var configSetAIKeyCmd = &cobra.Command{ + Use: "set-ai-key ", + Short: "Configure your Gemini / active AI provider API key", + Args: cobra.ExactArgs(1), + RunE: func(cmd *cobra.Command, args []string) error { + return configAISetKeyCmd.RunE(cmd, []string{appConfig.AIProvider, args[0]}) + }, +} + +var configSetModelCmd = &cobra.Command{ + Use: "set-model ", + Short: "Set the model name for the active AI provider", + Args: cobra.ExactArgs(1), + RunE: func(cmd *cobra.Command, args []string) error { + return configAISetModelCmd.RunE(cmd, []string{appConfig.AIProvider, args[0]}) + }, +} + +func maskKey(k string) string { + k = strings.TrimSpace(k) + if k == "" { + return "(not set)" + } + if len(k) <= 8 { + return "********" + } + return k[:4] + "..." + k[len(k)-4:] +} + +func keySource(storedKey string, envVars ...string) string { + for _, env := range envVars { + if os.Getenv(env) != "" { + return "Environment (" + env + ")" + } } - if appConfig.GeminiKey != "" { - return "Configured (Stored / Environment Variable)" + if storedKey != "" { + return "Saved (~/.sql-doctor)" } return "None" } func init() { + // Register subcommands under `config ai` + configAICmd.AddCommand(configAISwitchCmd) + configAICmd.AddCommand(configAISetKeyCmd) + configAICmd.AddCommand(configAISetModelCmd) + configAICmd.AddCommand(configAISetEndpointCmd) + configAICmd.AddCommand(configAITestCmd) + configAICmd.AddCommand(configAIClearCmd) + + // Register under `config` configCmd.AddCommand(configAICmd) configCmd.AddCommand(configSetAIKeyCmd) configCmd.AddCommand(configSetModelCmd) diff --git a/internal/cli/root.go b/internal/cli/root.go index 59e54ba..22eefd3 100644 --- a/internal/cli/root.go +++ b/internal/cli/root.go @@ -10,7 +10,6 @@ import ( "github.com/spf13/cobra" "github.com/sql-doctor/sql-doctor/internal/ai" - "github.com/sql-doctor/sql-doctor/internal/ai/gemini" "github.com/sql-doctor/sql-doctor/internal/config" "github.com/sql-doctor/sql-doctor/internal/database" "github.com/sql-doctor/sql-doctor/internal/database/mysql" @@ -238,14 +237,12 @@ func formatEnsureDatabase(ctx context.Context, db *sql.DB, driver database.Drive return fmt.Errorf("%s", suggestion.String()) } -// GetAIProvider returns an initialized Gemini provider +// GetAIProvider returns an initialized AI provider based on active configuration func GetAIProvider() ai.AIProvider { - var key, model string - if appConfig != nil { - key = appConfig.GeminiKey - model = appConfig.GeminiModel + if appConfig == nil { + return NewAIProvider(ai.ProviderParams{}) } - return gemini.New(key, model) + return NewAIProvider(appConfig.ActiveAIParams()) } // OutputResult renders data as either JSON or passes to a terminal printer diff --git a/internal/cli/shell.go b/internal/cli/shell.go index 0b05d3e..550c113 100644 --- a/internal/cli/shell.go +++ b/internal/cli/shell.go @@ -12,10 +12,12 @@ import ( "github.com/charmbracelet/lipgloss" "github.com/spf13/cobra" + aiContext "github.com/sql-doctor/sql-doctor/internal/ai/context" "github.com/sql-doctor/sql-doctor/internal/database" "github.com/sql-doctor/sql-doctor/internal/query/analyzer" "github.com/sql-doctor/sql-doctor/internal/query/explain" "github.com/sql-doctor/sql-doctor/internal/query/optimizer" + "github.com/sql-doctor/sql-doctor/internal/schema" "github.com/sql-doctor/sql-doctor/internal/storage" "github.com/sql-doctor/sql-doctor/internal/ui" "golang.org/x/term" @@ -656,6 +658,58 @@ func handleShellCommand(ctx context.Context, session *ShellSession, line string) fmt.Println(ui.Success("Database '%s' is responsive. Total tables: %d", session.Config.Database, len(tables))) return nil + case "ask": + if len(parts) < 2 { + return fmt.Errorf("usage: ask ") + } + question := strings.TrimPrefix(trimmed, parts[0]+" ") + if err := EnsureShellDatabase(ctx, session.DB, session.Driver, session.Config); err != nil { + return err + } + p := GetAIProvider() + if !p.IsConfigured() { + fmt.Println(ui.Warning("AI features are unavailable because %s is not configured.", p.ProviderName())) + fmt.Println("\nConfigure your AI provider outside the shell using:") + fmt.Println(" sql-doctor config ai") + return nil + } + fmt.Println(ui.Info("Thinking with %s (%s)...", p.ProviderName(), p.Model())) + tables, _ := session.Driver.Tables(ctx, session.DB) + details, _ := schema.FetchAllTableDetails(ctx, session.Driver, session.DB) + schemaContext := aiContext.BuildMinifiedSchema(tables, details) + + lowerQ := strings.ToLower(question) + isQueryGen := strings.HasPrefix(lowerQ, "write") || strings.HasPrefix(lowerQ, "generate") || + strings.HasPrefix(lowerQ, "find") || strings.HasPrefix(lowerQ, "select") || + strings.HasPrefix(lowerQ, "get") || strings.Contains(lowerQ, "query to") + + if isQueryGen { + genSQL, err := p.GenerateSQL(ctx, question, schemaContext) + if err != nil { + return err + } + fmt.Println() + fmt.Println(ui.HeaderStyle.Render("Generated SQL Query:")) + fmt.Println(ui.CardStyle.Render(genSQL.SQL)) + if genSQL.Explanation != "" { + fmt.Printf("Explanation: %s\n", genSQL.Explanation) + } + if genSQL.IsDestructive { + fmt.Println(ui.CriticalBadge + " " + ui.Error("Warning: This query modifies or deletes data!")) + } + fmt.Println() + fmt.Println(ui.Info("Tip: Copy and paste the query above to execute it.")) + return nil + } + + ans, err := p.Ask(ctx, question, schemaContext) + if err != nil { + return err + } + fmt.Println() + fmt.Println(ans) + return nil + case "help", "?": printShellHelp(parts[1:]...) return nil diff --git a/internal/config/config.go b/internal/config/config.go index 1759541..2e3a726 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -3,27 +3,89 @@ package config import ( "context" "os" + "strings" + "github.com/sql-doctor/sql-doctor/internal/ai" "github.com/sql-doctor/sql-doctor/internal/storage" ) // Config holds runtime configuration settings type Config struct { + AIProvider string // "gemini", "openai", "claude", "ollama" GeminiKey string GeminiModel string + OpenAIKey string + OpenAIModel string + ClaudeKey string + ClaudeModel string + OllamaEndpoint string + OllamaModel string + OllamaKey string ActiveConnName string OutputFormat string // "text" or "json" Verbose bool } +// ActiveAIParams returns the parameters for initializing the currently active AI provider +func (c *Config) ActiveAIParams() ai.ProviderParams { + provider := strings.ToLower(strings.TrimSpace(c.AIProvider)) + if provider == "" { + provider = ai.ProviderGemini + } + + switch provider { + case ai.ProviderOpenAI: + return ai.ProviderParams{ + Provider: ai.ProviderOpenAI, + Model: c.OpenAIModel, + APIKey: c.OpenAIKey, + } + case ai.ProviderClaude, "anthropic": + return ai.ProviderParams{ + Provider: ai.ProviderClaude, + Model: c.ClaudeModel, + APIKey: c.ClaudeKey, + } + case ai.ProviderOllama, "local": + return ai.ProviderParams{ + Provider: ai.ProviderOllama, + Model: c.OllamaModel, + APIKey: c.OllamaKey, + Endpoint: c.OllamaEndpoint, + } + default: + return ai.ProviderParams{ + Provider: ai.ProviderGemini, + Model: c.GeminiModel, + APIKey: c.GeminiKey, + } + } +} + // LoadConfig resolves configuration parameters from env vars and local storage func LoadConfig(ctx context.Context, store *storage.Storage) (*Config, error) { cfg := &Config{ - GeminiModel: "gemini-2.5-flash", - OutputFormat: "text", + AIProvider: ai.ProviderGemini, + GeminiModel: "gemini-3.8-flash", + OpenAIModel: "gpt-4o-mini", + ClaudeModel: "claude-3-5-haiku-20241022", + OllamaEndpoint: "http://localhost:11434/v1", + OllamaModel: "deepseek-r1:8b", + OutputFormat: "text", + } + + // 1. Active Provider + if envP := os.Getenv("AI_PROVIDER"); envP != "" { + cfg.AIProvider = strings.ToLower(envP) + } else if envP2 := os.Getenv("SQL_DOCTOR_AI_PROVIDER"); envP2 != "" { + cfg.AIProvider = strings.ToLower(envP2) + } else if store != nil { + if storedP, _ := store.GetSetting(ctx, "ai_provider"); storedP != "" { + cfg.AIProvider = strings.ToLower(storedP) + } } - // 1. Gemini Key + // 2. Google Gemini if envKey := os.Getenv("GEMINI_API_KEY"); envKey != "" { cfg.GeminiKey = envKey } else if store != nil { @@ -31,8 +93,6 @@ func LoadConfig(ctx context.Context, store *storage.Storage) (*Config, error) { cfg.GeminiKey = storedKey } } - - // 2. Gemini Model if envModel := os.Getenv("GEMINI_MODEL"); envModel != "" { cfg.GeminiModel = envModel } else if store != nil { @@ -41,7 +101,68 @@ func LoadConfig(ctx context.Context, store *storage.Storage) (*Config, error) { } } - // 3. Active Connection Name + // 3. OpenAI + if envKey := os.Getenv("OPENAI_API_KEY"); envKey != "" { + cfg.OpenAIKey = envKey + } else if store != nil { + if storedKey, _ := store.GetSetting(ctx, "openai_api_key"); storedKey != "" { + cfg.OpenAIKey = storedKey + } + } + if envModel := os.Getenv("OPENAI_MODEL"); envModel != "" { + cfg.OpenAIModel = envModel + } else if store != nil { + if storedModel, _ := store.GetSetting(ctx, "openai_model"); storedModel != "" { + cfg.OpenAIModel = storedModel + } + } + + // 4. Anthropic Claude + if envKey := os.Getenv("ANTHROPIC_API_KEY"); envKey != "" { + cfg.ClaudeKey = envKey + } else if envKey2 := os.Getenv("CLAUDE_API_KEY"); envKey2 != "" { + cfg.ClaudeKey = envKey2 + } else if store != nil { + if storedKey, _ := store.GetSetting(ctx, "claude_api_key"); storedKey != "" { + cfg.ClaudeKey = storedKey + } + } + if envModel := os.Getenv("ANTHROPIC_MODEL"); envModel != "" { + cfg.ClaudeModel = envModel + } else if envModel2 := os.Getenv("CLAUDE_MODEL"); envModel2 != "" { + cfg.ClaudeModel = envModel2 + } else if store != nil { + if storedModel, _ := store.GetSetting(ctx, "claude_model"); storedModel != "" { + cfg.ClaudeModel = storedModel + } + } + + // 5. Ollama / Local + if envEndpoint := os.Getenv("OLLAMA_ENDPOINT"); envEndpoint != "" { + cfg.OllamaEndpoint = envEndpoint + } else if envEndpoint2 := os.Getenv("OLLAMA_HOST"); envEndpoint2 != "" { + cfg.OllamaEndpoint = envEndpoint2 + } else if store != nil { + if storedEndpoint, _ := store.GetSetting(ctx, "ollama_endpoint"); storedEndpoint != "" { + cfg.OllamaEndpoint = storedEndpoint + } + } + if envModel := os.Getenv("OLLAMA_MODEL"); envModel != "" { + cfg.OllamaModel = envModel + } else if store != nil { + if storedModel, _ := store.GetSetting(ctx, "ollama_model"); storedModel != "" { + cfg.OllamaModel = storedModel + } + } + if envKey := os.Getenv("OLLAMA_API_KEY"); envKey != "" { + cfg.OllamaKey = envKey + } else if store != nil { + if storedKey, _ := store.GetSetting(ctx, "ollama_api_key"); storedKey != "" { + cfg.OllamaKey = storedKey + } + } + + // 6. Active Connection Name if envConn := os.Getenv("SQL_DOCTOR_CONN"); envConn != "" { cfg.ActiveConnName = envConn } else if store != nil { @@ -53,12 +174,56 @@ func LoadConfig(ctx context.Context, store *storage.Storage) (*Config, error) { return cfg, nil } -// SetGeminiKey updates the persisted Gemini API key +// SetAIProvider updates the active AI provider +func SetAIProvider(ctx context.Context, store *storage.Storage, provider string) error { + return store.SetSetting(ctx, "ai_provider", strings.ToLower(strings.TrimSpace(provider))) +} + +// SetAIKey updates the persisted API key for a provider +func SetAIKey(ctx context.Context, store *storage.Storage, provider, key string) error { + p := normalizeProvider(provider) + return store.SetSetting(ctx, p+"_api_key", strings.TrimSpace(key)) +} + +// SetAIModel updates the persisted model for a provider +func SetAIModel(ctx context.Context, store *storage.Storage, provider, model string) error { + p := normalizeProvider(provider) + return store.SetSetting(ctx, p+"_model", strings.TrimSpace(model)) +} + +// SetAIEndpoint updates the persisted endpoint for Ollama / local LLM +func SetAIEndpoint(ctx context.Context, store *storage.Storage, endpoint string) error { + return store.SetSetting(ctx, "ollama_endpoint", strings.TrimSpace(endpoint)) +} + +// ClearAIKey deletes the stored API key for a provider +func ClearAIKey(ctx context.Context, store *storage.Storage, provider string) error { + p := normalizeProvider(provider) + return store.DeleteSetting(ctx, p+"_api_key") +} + +// SetGeminiKey updates the persisted Gemini API key (backward compatibility) func SetGeminiKey(ctx context.Context, store *storage.Storage, key string) error { - return store.SetSetting(ctx, "gemini_api_key", key) + return SetAIKey(ctx, store, ai.ProviderGemini, key) } -// SetGeminiModel updates the persisted Gemini Model name +// SetGeminiModel updates the persisted Gemini Model name (backward compatibility) func SetGeminiModel(ctx context.Context, store *storage.Storage, model string) error { - return store.SetSetting(ctx, "gemini_model", model) + return SetAIModel(ctx, store, ai.ProviderGemini, model) +} + +func normalizeProvider(p string) string { + lower := strings.ToLower(strings.TrimSpace(p)) + switch lower { + case "google", "gemini": + return "gemini" + case "openai", "gpt": + return "openai" + case "claude", "anthropic": + return "claude" + case "ollama", "local": + return "ollama" + default: + return lower + } } diff --git a/internal/storage/sqlite.go b/internal/storage/sqlite.go index 34729c2..9790acc 100644 --- a/internal/storage/sqlite.go +++ b/internal/storage/sqlite.go @@ -415,6 +415,12 @@ func (s *Storage) GetSetting(ctx context.Context, key string) (string, error) { return val, nil } +// DeleteSetting removes a persistent configuration setting +func (s *Storage) DeleteSetting(ctx context.Context, key string) error { + _, err := s.db.ExecContext(ctx, "DELETE FROM settings WHERE key = ?;", key) + return err +} + // SaveSessionConnection saves an ephemeral/session connection func (s *Storage) SaveSessionConnection(ctx context.Context, conn *ConnectionRecord) error { data, err := json.Marshal(conn) diff --git a/tests/unit/ai_providers_test.go b/tests/unit/ai_providers_test.go new file mode 100644 index 0000000..91917d0 --- /dev/null +++ b/tests/unit/ai_providers_test.go @@ -0,0 +1,171 @@ +package unit + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "path/filepath" + "testing" + + "github.com/sql-doctor/sql-doctor/internal/ai" + "github.com/sql-doctor/sql-doctor/internal/ai/claude" + "github.com/sql-doctor/sql-doctor/internal/ai/openai" + "github.com/sql-doctor/sql-doctor/internal/cli" + "github.com/sql-doctor/sql-doctor/internal/config" + "github.com/sql-doctor/sql-doctor/internal/storage" +) + +func TestAIProviderFactory(t *testing.T) { + // Gemini default + pGemini := cli.NewAIProvider(ai.ProviderParams{ + Provider: "gemini", + APIKey: "test-gemini-key", + }) + if pGemini.ProviderName() != "Google Gemini" { + t.Errorf("expected 'Google Gemini', got '%s'", pGemini.ProviderName()) + } + if pGemini.Model() != "gemini-3.8-flash" { + t.Errorf("expected default 'gemini-3.8-flash', got '%s'", pGemini.Model()) + } + if !pGemini.IsConfigured() { + t.Errorf("expected gemini to be configured with api key") + } + + // OpenAI + pOpenAI := cli.NewAIProvider(ai.ProviderParams{ + Provider: "openai", + APIKey: "sk-test-key", + Model: "gpt-4o", + }) + if pOpenAI.ProviderName() != "OpenAI" { + t.Errorf("expected 'OpenAI', got '%s'", pOpenAI.ProviderName()) + } + if pOpenAI.Model() != "gpt-4o" { + t.Errorf("expected 'gpt-4o', got '%s'", pOpenAI.Model()) + } + if !pOpenAI.IsConfigured() { + t.Errorf("expected openai to be configured") + } + + // Claude + pClaude := cli.NewAIProvider(ai.ProviderParams{ + Provider: "claude", + APIKey: "sk-ant-test", + Model: "claude-3-5-sonnet-20241022", + }) + if pClaude.ProviderName() != "Anthropic Claude" { + t.Errorf("expected 'Anthropic Claude', got '%s'", pClaude.ProviderName()) + } + if pClaude.Model() != "claude-3-5-sonnet-20241022" { + t.Errorf("expected claude model, got '%s'", pClaude.Model()) + } + + // Ollama + pOllama := cli.NewAIProvider(ai.ProviderParams{ + Provider: "ollama", + Endpoint: "http://localhost:11434/v1", + }) + if pOllama.ProviderName() != "Ollama (Local)" { + t.Errorf("expected 'Ollama (Local)', got '%s'", pOllama.ProviderName()) + } + if pOllama.Model() != "deepseek-r1:8b" { + t.Errorf("expected default 'deepseek-r1:8b', got '%s'", pOllama.Model()) + } + if !pOllama.IsConfigured() { + t.Errorf("expected ollama with endpoint to be configured") + } +} + +func TestOpenAIClientMock(t *testing.T) { + ctx := context.Background() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/chat/completions" { + t.Errorf("unexpected path: %s", r.URL.Path) + } + if r.Header.Get("Authorization") != "Bearer test-secret" { + t.Errorf("missing or invalid authorization header") + } + + resp := map[string]interface{}{ + "choices": []map[string]interface{}{ + { + "message": map[string]string{ + "role": "assistant", + "content": "PONG", + }, + }, + }, + } + _ = json.NewEncoder(w).Encode(resp) + })) + defer server.Close() + + client := openai.New("test-secret", "gpt-4o-mini", server.URL, false) + if err := client.TestConnection(ctx); err != nil { + t.Fatalf("TestConnection failed: %v", err) + } +} + +func TestClaudeClientMock(t *testing.T) { + client := claude.New("ant-test-secret", "claude-3-5-haiku-20241022") + if !client.IsConfigured() { + t.Fatalf("expected client to be configured") + } + if client.ProviderName() != "Anthropic Claude" { + t.Errorf("expected 'Anthropic Claude', got '%s'", client.ProviderName()) + } +} + +func TestStorageDeleteSetting(t *testing.T) { + ctx := context.Background() + tmpDir := t.TempDir() + dbPath := filepath.Join(tmpDir, "test_settings.db") + + store, err := storage.OpenStorage(dbPath) + if err != nil { + t.Fatalf("failed to init storage: %v", err) + } + defer store.Close() + + // 1. Set key + if err := config.SetAIKey(ctx, store, "openai", "sk-proj-12345"); err != nil { + t.Fatalf("failed to set key: %v", err) + } + + val, err := store.GetSetting(ctx, "openai_api_key") + if err != nil || val != "sk-proj-12345" { + t.Errorf("expected 'sk-proj-12345', got '%s'", val) + } + + // 2. Clear key + if err := config.ClearAIKey(ctx, store, "openai"); err != nil { + t.Fatalf("failed to clear key: %v", err) + } + + valAfter, err := store.GetSetting(ctx, "openai_api_key") + if err != nil || valAfter != "" { + t.Errorf("expected empty string after deletion, got '%s'", valAfter) + } +} + +func TestCuratedModelsByProvider(t *testing.T) { + providers := []string{ai.ProviderGemini, ai.ProviderOpenAI, ai.ProviderClaude, ai.ProviderOllama} + for _, p := range providers { + models, ok := ai.CuratedModelsByProvider[p] + if !ok || len(models) == 0 { + t.Errorf("expected curated models for provider %s", p) + } + hasRecommended := false + for _, m := range models { + if m.Recommended { + hasRecommended = true + break + } + } + if !hasRecommended { + t.Errorf("expected at least one recommended model for provider %s", p) + } + } +}