diff --git a/DEVELOPMENT.md b/DEVELOPMENT.md index ab21e77..3f38469 100644 --- a/DEVELOPMENT.md +++ b/DEVELOPMENT.md @@ -50,7 +50,7 @@ sql-doctor/ │ ├── config/ │ │ └── config.go # Env vars & configuration loader │ ├── ui/ # Lip Gloss styling, scorecards, tables -│ └── cli/ # Cobra commands & flag handlers +│ └── cli/ # Cobra commands, shell REPL & flag handlers ├── tests/ │ ├── fixtures/ # Test schemas and migration files │ └── unit/ # Unit tests @@ -138,6 +138,7 @@ type Driver interface { Connect(ctx context.Context, cfg *ConnectionConfig) (*sql.DB, error) Ping(ctx context.Context, db *sql.DB) error Version(ctx context.Context, db *sql.DB) (string, error) + Databases(ctx context.Context, db *sql.DB) ([]string, error) Tables(ctx context.Context, db *sql.DB) ([]TableInfo, error) DescribeTable(ctx context.Context, db *sql.DB, table string) (*TableDetail, error) Indexes(ctx context.Context, db *sql.DB, table string) ([]IndexInfo, error) diff --git a/README.md b/README.md index ed274b6..423bfde 100644 --- a/README.md +++ b/README.md @@ -24,6 +24,7 @@ As a software developer, I built SQL Doctor to solve the common database problem - **Dealing with poorly chosen data types or oversized columns** — SQL Doctor inspects your actual data records and recommends better, more compact definitions based on real statistics. - **Worrying about what could break when applying a migration** — SQL Doctor checks for potential data loss, table locks, and compatibility issues beforehand. - **Knowing something is wrong with a database but not knowing where** — SQL Doctor runs a full diagnostic across schema health, performance, indexes, data quality, and referential integrity. +- **Wanting an interactive terminal environment without typing `sql-doctor` before every command** — SQL Doctor includes an interactive REPL shell with engine selection, database switching, direct SQL execution, and clean box table rendering. It gives you clear, deterministic answers straight in your terminal without needing heavy GUI clients or cloud dashboards. And if you want conversational assistance, you can optionally connect your own Gemini API key for query explanations and natural-language query generation. @@ -87,29 +88,71 @@ To use it from anywhere on Windows, add the folder containing `sql-doctor.exe` t ## 🚀 Everyday Usage & Examples -### 1. Connecting to a Database +### 1. Interactive Shell REPL (`shell`) +If you prefer an interactive environment like the `mysql` or `psql` CLI where you can type queries and inspect schemas without prefixing every command with `sql-doctor`: + +```bash +# Launch interactive shell with database engine picker +sql-doctor shell + +# Or jump straight into a specific connection / database: +sql-doctor shell -c local-mysql -d rolerift +``` + +Inside the shell, your prompt reflects your active engine and database: +```text +sql-doctor [mysql@rolerift]> tables +sql-doctor [mysql@rolerift]> select * from users; +sql-doctor [mysql@rolerift]> select * from users\G # Vertical format (one column per line) +sql-doctor [mysql@rolerift]> use shop_db # Switch database on the fly +sql-doctor [mysql@shop_db]> analyze SELECT * FROM orders WHERE status = 'pending' +sql-doctor [mysql@shop_db]> help # View categorized commands +sql-doctor [mysql@shop_db]> exit # Clean exit (discards in-memory session) +``` + +All session state (active connection and selected database) lives in memory during your shell session and is cleanly discarded upon exit. + +--- + +### 2. Connecting to a Database SQL Doctor works out of the box with **MySQL**, **MariaDB**, **PostgreSQL**, and **SQLite**. +Specifying a database name upfront is completely optional. If you connect to a server without picking a database, you can select one later: + ```bash -# SQLite (Local file) -sql-doctor connect --type sqlite --file ./my-app.db --name my-local-db --save +# Connect to MySQL / MariaDB (specifying a database is optional) +sql-doctor connect --type mysql --host 127.0.0.1 --port 3306 --user root -p -# PostgreSQL -sql-doctor connect --type postgres --host localhost --port 5432 --user postgres --password mysecret --database shop_db --name local-pg --save +# Connect to PostgreSQL +sql-doctor connect --type postgres --host localhost --port 5432 --user postgres -p -# MySQL / MariaDB -sql-doctor connect --type mysql --host 127.0.0.1 --port 3306 --user root --password mysecret --database shop_db --name local-mysql --save +# Connect to SQLite (local file) +sql-doctor connect --type sqlite --file ./my-app.db --name my-local-db --save # Or run directly against a database URL without saving: sql-doctor --db-url "postgres://user:pass@localhost:5432/shop_db" db tables ``` -To see your saved connections or switch between them: +> **Tip:** You don't need `--save` just to try a connection. Running `connect` without `--save` starts an ephemeral session so you can immediately run subsequent commands in that terminal. + +#### Managing Databases on the Server: +```bash +# List all databases on the connected server +sql-doctor databases + +# Switch the active database +sql-doctor use rolerift + +# Clear session memory +sql-doctor disconnect +``` + +To see your saved connection profiles or switch between them: ```bash -# List saved connections +# List saved connection profiles sql-doctor connections -# Switch active connection +# Switch active profile sql-doctor connections --use local-pg # Quick connectivity test @@ -118,7 +161,7 @@ sql-doctor ping --- -### 2. Full Health Diagnostic (`doctor`) +### 3. Full Health Diagnostic (`doctor`) Run a quick diagnostic across your whole database. It checks for tables missing primary keys, unindexed foreign keys, redundant indexes, and sampled data anomalies: ```bash @@ -142,7 +185,7 @@ Warnings: --- -### 3. Query Performance & EXPLAIN Analysis +### 4. Query Performance & EXPLAIN Analysis Profile slow queries to see execution time, rows examined vs returned, and unindexed table scans: ```bash @@ -158,7 +201,7 @@ sql-doctor query optimize "SELECT * FROM users WHERE status = 'active' AND age > --- -### 4. Data-Aware Datatype Advisor +### 5. Data-Aware Datatype Advisor Don't guess what column type you should have used. SQL Doctor samples actual records and checks value lengths, patterns (like UUIDs, ISO dates, and booleans), and recommends tighter types: ```bash @@ -176,7 +219,7 @@ INFO Column 'user_uuid' --- -### 5. Checking Foreign Keys & Orphan Rows +### 6. Checking Foreign Keys & Orphan Rows Find broken referential integrity before your app hits a foreign key error: ```bash @@ -187,7 +230,7 @@ This lists all foreign keys (plus inferred relationships like `user_id -> users. --- -### 6. Comparing Two Database Schemas (`diff`) +### 7. Comparing Two Database Schemas (`diff`) Need to verify if your staging database matches production? ```bash @@ -198,7 +241,7 @@ This compares tables, columns, data types, nullability, defaults, and indexes, a --- -### 7. Migration Safety Checks +### 8. Migration Safety Checks Before running a migration script on production, check it for destructive commands or locking hazards: ```bash @@ -212,7 +255,7 @@ Catches issues like: --- -### 8. SQL Linter & Formatter +### 9. SQL Linter & Formatter Quick static checks without needing a live connection: ```bash @@ -225,7 +268,7 @@ sql-doctor format "select id,name from users where status='active' and age>21 or --- -### 9. Optional Gemini AI Assistant +### 10. Optional Gemini AI Assistant If you want AI explanations or natural-language query generation, add your own Gemini API key: ```bash @@ -244,7 +287,7 @@ sql-doctor ask "Write a query to find the top 5 customers by revenue this year" --- -### 10. Machine-Readable Output (`--json`) +### 11. Machine-Readable Output (`--json`) Every single command supports the `--json` flag. You can pipe the output into `jq` or plug it into your CI/CD pipelines: ```bash diff --git a/go.mod b/go.mod index 8c90e86..1901c28 100644 --- a/go.mod +++ b/go.mod @@ -13,6 +13,8 @@ require ( modernc.org/sqlite v1.58.0 ) +require golang.org/x/term v0.46.0 // indirect + require ( cloud.google.com/go v0.115.0 // indirect cloud.google.com/go/ai v0.8.0 // indirect @@ -59,7 +61,7 @@ require ( golang.org/x/net v0.58.0 // indirect golang.org/x/oauth2 v0.36.0 // indirect golang.org/x/sync v0.22.0 // indirect - golang.org/x/sys v0.47.0 // indirect + golang.org/x/sys v0.48.0 // indirect golang.org/x/text v0.41.0 // indirect golang.org/x/time v0.15.0 // indirect google.golang.org/genproto/googleapis/api v0.0.0-20260715232425-e75dac1f907d // indirect diff --git a/go.sum b/go.sum index 19cd420..3286f0e 100644 --- a/go.sum +++ b/go.sum @@ -156,6 +156,10 @@ golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo= +golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og= +golang.org/x/term v0.46.0 h1:3+OXuTbaKDgwk8jTi3aSLHRlmWqHEUDUtxnbFigO4YE= +golang.org/x/term v0.46.0/go.mod h1:+K02xbkittuwc0Am4abfA3Fc+XRGXkvBXNO88NCXPoc= golang.org/x/text v0.0.0-20180302201248-b7ef84aaf62a/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8= golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M= diff --git a/internal/cli/ask.go b/internal/cli/ask.go index 15219c5..03baa08 100644 --- a/internal/cli/ask.go +++ b/internal/cli/ask.go @@ -42,6 +42,10 @@ var askCmd = &cobra.Command{ } defer db.Close() + if err := EnsureDatabase(ctx, db, driver, cfg); err != nil { + return err + } + fmt.Println(ui.Info("Inspecting schema context for [%s]...", cfg.Database)) tables, _ := driver.Tables(ctx, db) details, _ := schema.FetchAllTableDetails(ctx, driver, db) diff --git a/internal/cli/connect.go b/internal/cli/connect.go index eca8b83..775b8ab 100644 --- a/internal/cli/connect.go +++ b/internal/cli/connect.go @@ -1,7 +1,9 @@ package cli import ( + "bufio" "fmt" + "os" "strings" "time" @@ -9,6 +11,7 @@ import ( "github.com/sql-doctor/sql-doctor/internal/database" "github.com/sql-doctor/sql-doctor/internal/storage" "github.com/sql-doctor/sql-doctor/internal/ui" + "golang.org/x/term" ) var ( @@ -21,6 +24,7 @@ var ( connFlagFile string connFlagName string connFlagSave bool + connFlagShell bool connListUse string connListDelete string @@ -47,13 +51,31 @@ Examples: } } + password := connFlagPassword + if password == "__PROMPT__" { + fmt.Print("Enter password (leave empty if none): ") + if term.IsTerminal(int(os.Stdin.Fd())) { + bytePassword, err := term.ReadPassword(int(os.Stdin.Fd())) + fmt.Println() + if err == nil { + password = strings.TrimRight(string(bytePassword), "\r\n") + } else { + password = "" + } + } else { + reader := bufio.NewReader(os.Stdin) + pass, _ := reader.ReadString('\n') + password = strings.TrimRight(pass, "\r\n") + } + } + cfg := &database.ConnectionConfig{ Name: connFlagName, Dialect: dialect, Host: connFlagHost, Port: connFlagPort, User: connFlagUser, - Password: connFlagPassword, + Password: password, Database: connFlagDatabase, FilePath: connFlagFile, } @@ -82,22 +104,37 @@ Examples: ver = "Unknown" } + rec := &storage.ConnectionRecord{ + Name: cfg.Name, + Dialect: cfg.Dialect, + Host: cfg.Host, + Port: cfg.Port, + User: cfg.User, + Password: cfg.Password, + Database: cfg.Database, + FilePath: cfg.FilePath, + IsActive: true, + } + if connFlagSave && appStorage != nil { - rec := &storage.ConnectionRecord{ - Name: cfg.Name, - Dialect: cfg.Dialect, - Host: cfg.Host, - Port: cfg.Port, - User: cfg.User, - Password: cfg.Password, - Database: cfg.Database, - FilePath: cfg.FilePath, - IsActive: true, - } if err := appStorage.SaveConnection(ctx, rec); err != nil { return fmt.Errorf("failed to save connection profile: %w", err) } _ = appStorage.SetActiveConnection(ctx, cfg.Name) + _ = appStorage.ClearSessionConnection(ctx) + } else if appStorage != nil { + // Save active session connection so user can continue without saving a permanent profile + if err := appStorage.SaveSessionConnection(ctx, rec); err != nil { + return fmt.Errorf("failed to save active session: %w", err) + } + } + + if connFlagShell { + err := RunShell(ctx, db, driver, cfg) + if !connFlagSave && appStorage != nil { + _ = appStorage.ClearSessionConnection(ctx) + } + return err } OutputResult(map[string]interface{}{ @@ -114,7 +151,15 @@ Examples: fmt.Printf(" Dialect: %s\n", cfg.Dialect) if connFlagSave { fmt.Printf(" Profile: Saved as '%s' (set as active connection)\n", cfg.Name) + } else { + fmt.Println(" Session: Active (temporary — use --save to persist as profile)") + } + if cfg.Database != "" { + fmt.Printf(" Database: %s\n", cfg.Database) + } else if cfg.Dialect != database.DialectSQLite { + fmt.Println(ui.Info(" Note: No database selected. Choose one with 'sql-doctor use ' or view with 'sql-doctor db databases'.")) } + fmt.Println(ui.Info(" Tip: Type 'sql-doctor shell' to open interactive MySQL-like console.")) }) return nil @@ -134,6 +179,7 @@ var connectionsCmd = &cobra.Command{ if err := appStorage.SetActiveConnection(ctx, connListUse); err != nil { return err } + _ = appStorage.ClearSessionConnection(ctx) fmt.Println(ui.Success("Active connection switched to '%s'", connListUse)) return nil } @@ -151,23 +197,50 @@ var connectionsCmd = &cobra.Command{ return err } - OutputResult(list, func() { + sess, _ := appStorage.GetSessionConnection(ctx) + + OutputResult(map[string]interface{}{ + "connections": list, + "session": sess, + }, func() { + if sess != nil { + sessHost := sess.Host + if sess.FilePath != "" { + sessHost = sess.FilePath + } + dbStr := sess.Database + if dbStr == "" { + dbStr = "(none selected)" + } + fmt.Println(ui.InfoBadge + " " + ui.HeaderStyle.Render("Active Session (not saved as permanent profile):")) + fmt.Printf(" Name: %s | Dialect: %s | Host: %s | Database: %s | User: %s\n", + sess.Name, sess.Dialect, sessHost, dbStr, sess.User) + fmt.Println(" (Tip: Run 'sql-doctor connect ... --save' to persist as a profile)") + fmt.Println() + } + if len(list) == 0 { - fmt.Println(ui.Info("No saved connections found. Use 'sql-doctor connect ... --save' to add one.")) + if sess == nil { + fmt.Println(ui.Info("No saved connections found. Use 'sql-doctor connect ... --save' to add one.")) + } return } tbl := ui.NewTable("ACTIVE", "NAME", "DIALECT", "HOST/FILE", "DATABASE", "USER") for _, c := range list { activeMarker := "" - if c.IsActive { + if c.IsActive && sess == nil { activeMarker = " ★ " } hostFile := c.Host if c.FilePath != "" { hostFile = c.FilePath } - tbl.AddRow(activeMarker, c.Name, string(c.Dialect), hostFile, c.Database, c.User) + dbName := c.Database + if dbName == "" { + dbName = "(none)" + } + tbl.AddRow(activeMarker, c.Name, string(c.Dialect), hostFile, dbName, c.User) } fmt.Println(tbl.Render()) }) @@ -176,6 +249,100 @@ var connectionsCmd = &cobra.Command{ }, } +var useCmd = &cobra.Command{ + Use: "use ", + Short: "Select or switch the active database for the current connection", + Args: cobra.ExactArgs(1), + RunE: func(cmd *cobra.Command, args []string) error { + ctx := cmd.Context() + targetDB := args[0] + + if appStorage == nil { + return fmt.Errorf("local storage unavailable") + } + + db, driver, cfg, err := GetActiveDB(ctx) + if err != nil { + return err + } + defer db.Close() + + if cfg.Dialect == database.DialectSQLite { + return fmt.Errorf("the 'use' command is for multi-database servers (MySQL, PostgreSQL). For SQLite, connect directly using --file ") + } + + // Verify database exists on server + dbs, err := driver.Databases(ctx, db) + if err == nil && len(dbs) > 0 { + found := false + for _, d := range dbs { + if strings.EqualFold(d, targetDB) { + targetDB = d // preserve exact casing from server + found = true + break + } + } + if !found { + return fmt.Errorf("database '%s' not found on connection '%s'.\n\nRun 'sql-doctor db databases' to view available databases", targetDB, cfg.Name) + } + } + + updated, err := appStorage.UpdateConnectionDatabase(ctx, flagConn, targetDB) + if err != nil { + return fmt.Errorf("failed to update active database: %w", err) + } + + OutputResult(map[string]interface{}{ + "status": "SWITCHED", + "database": targetDB, + "connection": updated.Name, + }, func() { + fmt.Println(ui.Success("Active database set to '%s' (connection: '%s')", targetDB, updated.Name)) + fmt.Println("You can now run queries and database commands without specifying the database.") + }) + + return nil + }, +} + +var databasesCmd = &cobra.Command{ + Use: "databases", + Short: "List all databases available on the connected server", + RunE: func(cmd *cobra.Command, args []string) error { + ctx := cmd.Context() + db, driver, cfg, err := GetActiveDB(ctx) + if err != nil { + return err + } + defer db.Close() + + databases, err := driver.Databases(ctx, db) + if err != nil { + return fmt.Errorf("failed to fetch databases: %w", err) + } + + OutputResult(databases, func() { + fmt.Println(ui.TitleStyle.Render(fmt.Sprintf("Databases on [%s] (%d found)", cfg.Name, len(databases)))) + if len(databases) == 0 { + fmt.Println(ui.Info("No databases found.")) + return + } + + tbl := ui.NewTable("ACTIVE", "DATABASE NAME") + for _, d := range databases { + activeMarker := "" + if cfg.Database == d { + activeMarker = " ★ " + } + tbl.AddRow(activeMarker, d) + } + fmt.Println(tbl.Render()) + fmt.Println(ui.Info("Tip: Switch active database with 'sql-doctor use '")) + }) + return nil + }, +} + var pingCmd = &cobra.Command{ Use: "ping", Short: "Test active database connection and measure latency", @@ -217,11 +384,13 @@ func init() { connectCmd.Flags().StringVar(&connFlagHost, "host", "127.0.0.1", "Database host") connectCmd.Flags().IntVar(&connFlagPort, "port", 0, "Database port") connectCmd.Flags().StringVarP(&connFlagUser, "user", "u", "", "Database username") - connectCmd.Flags().StringVarP(&connFlagPassword, "password", "p", "", "Database password") + connectCmd.Flags().StringVarP(&connFlagPassword, "password", "p", "", "Database password (prompt if passed without value)") + connectCmd.Flags().Lookup("password").NoOptDefVal = "__PROMPT__" connectCmd.Flags().StringVarP(&connFlagDatabase, "database", "d", "", "Database name") connectCmd.Flags().StringVarP(&connFlagFile, "file", "f", "", "SQLite database file path") connectCmd.Flags().StringVar(&connFlagName, "name", "", "Connection profile name") connectCmd.Flags().BoolVar(&connFlagSave, "save", false, "Save this connection profile") + connectCmd.Flags().BoolVarP(&connFlagShell, "shell", "i", false, "Open interactive query shell after connecting") connectionsCmd.Flags().StringVar(&connListUse, "use", "", "Set connection profile as active") connectionsCmd.Flags().StringVar(&connListDelete, "delete", "", "Delete connection profile") diff --git a/internal/cli/data.go b/internal/cli/data.go index 27b5fe2..c56d634 100644 --- a/internal/cli/data.go +++ b/internal/cli/data.go @@ -21,12 +21,16 @@ var dataQualityCmd = &cobra.Command{ ctx := cmd.Context() tableName := args[0] - db, driver, _, err := GetActiveDB(ctx) + db, driver, cfg, err := GetActiveDB(ctx) if err != nil { return err } defer db.Close() + if err := EnsureDatabase(ctx, db, driver, cfg); err != nil { + return err + } + analyzer := data.NewQualityAnalyzer(driver) report, err := analyzer.AnalyzeTable(ctx, db, tableName) if err != nil { diff --git a/internal/cli/db.go b/internal/cli/db.go index e7ea00a..a511c0c 100644 --- a/internal/cli/db.go +++ b/internal/cli/db.go @@ -26,6 +26,10 @@ var dbTablesCmd = &cobra.Command{ } defer db.Close() + if err := EnsureDatabase(ctx, db, driver, cfg); err != nil { + return err + } + tables, err := driver.Tables(ctx, db) if err != nil { return fmt.Errorf("failed to fetch tables: %w", err) @@ -56,12 +60,16 @@ var dbDescribeCmd = &cobra.Command{ ctx := cmd.Context() tableName := args[0] - db, driver, _, err := GetActiveDB(ctx) + db, driver, cfg, err := GetActiveDB(ctx) if err != nil { return err } defer db.Close() + if err := EnsureDatabase(ctx, db, driver, cfg); err != nil { + return err + } + detail, err := driver.DescribeTable(ctx, db, tableName) if err != nil { return err @@ -106,12 +114,16 @@ var dbIndexesCmd = &cobra.Command{ ctx := cmd.Context() tableName := args[0] - db, driver, _, err := GetActiveDB(ctx) + db, driver, cfg, err := GetActiveDB(ctx) if err != nil { return err } defer db.Close() + if err := EnsureDatabase(ctx, db, driver, cfg); err != nil { + return err + } + indexes, err := driver.Indexes(ctx, db, tableName) if err != nil { return err @@ -134,12 +146,16 @@ var dbRelationshipsCmd = &cobra.Command{ Short: "Show foreign keys, inferred relationships, and orphan records", RunE: func(cmd *cobra.Command, args []string) error { ctx := cmd.Context() - db, driver, _, err := GetActiveDB(ctx) + db, driver, cfg, err := GetActiveDB(ctx) if err != nil { return err } defer db.Close() + if err := EnsureDatabase(ctx, db, driver, cfg); err != nil { + return err + } + rels, err := driver.Relationships(ctx, db) if err != nil { return err @@ -308,6 +324,10 @@ var dbSnapshotCmd = &cobra.Command{ } defer db.Close() + if err := EnsureDatabase(ctx, db, driver, cfg); err != nil { + return err + } + mgr := schema.NewSnapshotManager(appStorage, driver) snap, err := mgr.CreateSnapshot(ctx, db, name, cfg.Name, cfg.Database) if err != nil { @@ -325,6 +345,8 @@ var dbSnapshotCmd = &cobra.Command{ } func init() { + dbCmd.AddCommand(useCmd) + dbCmd.AddCommand(databasesCmd) dbCmd.AddCommand(dbTablesCmd) dbCmd.AddCommand(dbDescribeCmd) dbCmd.AddCommand(dbIndexesCmd) diff --git a/internal/cli/doctor.go b/internal/cli/doctor.go index 9997d37..b0fd5a0 100644 --- a/internal/cli/doctor.go +++ b/internal/cli/doctor.go @@ -34,6 +34,10 @@ var doctorCmd = &cobra.Command{ } defer db.Close() + if err := EnsureDatabase(ctx, db, driver, cfg); err != nil { + return err + } + ver, _ := driver.Version(ctx, db) report := &DoctorReport{ diff --git a/internal/cli/help.go b/internal/cli/help.go index 0a4e40e..daa1e99 100644 --- a/internal/cli/help.go +++ b/internal/cli/help.go @@ -70,6 +70,7 @@ func renderRootHelp() string { desc string }{ {"-c, --conn ", "Saved connection profile name to use"}, + {"-d, --database ", "Target database name to use for this command"}, {" --db-url ", "Direct database URL or SQLite file path"}, {" --ai", "Enable optional AI recommendations via Gemini"}, {" --json", "Output results as machine-readable JSON"}, @@ -95,8 +96,12 @@ func renderRootHelp() string { desc string }{ {"doctor", "Run full end-to-end database health check"}, + {"shell", "Open interactive query REPL console (MySQL-like)"}, {"connect", "Create and test a new database connection"}, {"connections", "List, inspect, and manage saved database connections"}, + {"use ", "Select or switch active database for current connection"}, + {"databases", "List all databases on target connection"}, + {"disconnect", "Disconnect and clear ephemeral session memory"}, {"ping", "Quick connectivity test to target database"}, {"lint", "Lint SQL query files against anti-pattern rules"}, {"format", "Format SQL queries with standard indentation"}, @@ -119,6 +124,8 @@ func renderRootHelp() string { { group: "db", cmds: []struct{ name, desc string }{ + {"db use ", "Select or switch active database"}, + {"db databases", "List all databases on target connection"}, {"db tables", "List all tables in connected database"}, {"db describe ", "Show column types, nullability, keys, and defaults"}, {"db indexes
", "List table indexes, column order, and uniqueness"}, @@ -187,7 +194,15 @@ func renderRootHelp() string { func renderSubcommandHelp(cmd *cobra.Command) string { var b strings.Builder + // Header (ASCII Art Logo + Badge) b.WriteString("\n") + b.WriteString(styleLogo.Render(AsciiLogo)) + b.WriteString("\n\n") + + badge := styleBadge.Render("SQL DOCTOR") + ver := styleVer.Render("v0.1.0-beta") + subtitle := styleGray.Render("— Database Diagnostics & SQL Intelligence CLI") + b.WriteString(fmt.Sprintf(" %s %s %s\n\n", badge, ver, subtitle)) // 1. Usage b.WriteString(styleYellow.Render("Usage:")) diff --git a/internal/cli/query.go b/internal/cli/query.go index 167d95a..6460bed 100644 --- a/internal/cli/query.go +++ b/internal/cli/query.go @@ -34,6 +34,10 @@ var queryAnalyzeCmd = &cobra.Command{ } defer db.Close() + if err := EnsureDatabase(ctx, db, driver, cfg); err != nil { + return err + } + anz := analyzer.NewAnalyzer(driver) result, err := anz.Analyze(ctx, db, queryText) if err != nil { @@ -124,12 +128,16 @@ var queryExplainCmd = &cobra.Command{ ctx := cmd.Context() queryText := args[0] - db, driver, _, err := GetActiveDB(ctx) + db, driver, cfg, err := GetActiveDB(ctx) if err != nil { return err } defer db.Close() + if err := EnsureDatabase(ctx, db, driver, cfg); err != nil { + return err + } + explainRes, err := driver.Explain(ctx, db, queryText, false) if err != nil { return fmt.Errorf("failed to explain query: %w", err) @@ -169,12 +177,16 @@ var queryOptimizeCmd = &cobra.Command{ ctx := cmd.Context() queryText := args[0] - db, driver, _, err := GetActiveDB(ctx) + db, driver, cfg, err := GetActiveDB(ctx) if err != nil { return err } defer db.Close() + if err := EnsureDatabase(ctx, db, driver, cfg); err != nil { + return err + } + opt := optimizer.NewOptimizer(driver) res, err := opt.Optimize(ctx, db, queryText) if err != nil { diff --git a/internal/cli/root.go b/internal/cli/root.go index a998fc9..59e54ba 100644 --- a/internal/cli/root.go +++ b/internal/cli/root.go @@ -6,6 +6,7 @@ import ( "encoding/json" "fmt" "os" + "strings" "github.com/spf13/cobra" "github.com/sql-doctor/sql-doctor/internal/ai" @@ -25,6 +26,7 @@ var ( flagConn string flagDBURL string flagAI bool + flagDatabase string appConfig *config.Config appStorage *storage.Storage @@ -56,12 +58,15 @@ func init() { RootCmd.PersistentFlags().BoolVar(&flagJSON, "json", false, "Output results as machine-readable JSON") RootCmd.PersistentFlags().BoolVarP(&flagVerbose, "verbose", "v", false, "Enable verbose output and logs") RootCmd.PersistentFlags().StringVarP(&flagConn, "conn", "c", "", "Saved connection profile name to use") + RootCmd.PersistentFlags().StringVarP(&flagDatabase, "database", "d", "", "Database name to use for this command") RootCmd.PersistentFlags().StringVar(&flagDBURL, "db-url", "", "Direct database URL or SQLite file path (e.g. postgres://user:pass@localhost:5432/db)") RootCmd.PersistentFlags().BoolVar(&flagAI, "ai", false, "Enable optional AI explanations / recommendations via Gemini") // Register subcommands RootCmd.AddCommand(connectCmd) RootCmd.AddCommand(connectionsCmd) + RootCmd.AddCommand(useCmd) + RootCmd.AddCommand(databasesCmd) RootCmd.AddCommand(pingCmd) RootCmd.AddCommand(configCmd) RootCmd.AddCommand(dbCmd) @@ -110,15 +115,10 @@ func GetActiveDB(ctx context.Context) (*sql.DB, database.Driver, *database.Conne return nil, nil, nil, err } cfg = c - } else { - // 2. Check --conn flag or active stored connection - targetName := flagConn - if targetName == "" && appConfig != nil { - targetName = appConfig.ActiveConnName - } - - if targetName != "" && appStorage != nil { - rec, err := appStorage.GetConnection(ctx, targetName) + } else if flagConn != "" { + // 2. Check --conn flag explicitly + if appStorage != nil { + rec, err := appStorage.GetConnection(ctx, flagConn) if err != nil { return nil, nil, nil, err } @@ -134,12 +134,53 @@ func GetActiveDB(ctx context.Context) (*sql.DB, database.Driver, *database.Conne SSLMode: rec.SSLMode, } } + } else { + // 3. Check for active session connection + if appStorage != nil { + sess, err := appStorage.GetSessionConnection(ctx) + if err == nil && sess != nil { + cfg = &database.ConnectionConfig{ + Name: sess.Name, + Dialect: sess.Dialect, + Host: sess.Host, + Port: sess.Port, + User: sess.User, + Password: sess.Password, + Database: sess.Database, + FilePath: sess.FilePath, + SSLMode: sess.SSLMode, + } + } + } + + // 4. Fallback to active saved connection profile + if cfg == nil && appConfig != nil && appConfig.ActiveConnName != "" && appStorage != nil { + rec, err := appStorage.GetConnection(ctx, appConfig.ActiveConnName) + if err == nil && rec != nil { + cfg = &database.ConnectionConfig{ + Name: rec.Name, + Dialect: rec.Dialect, + Host: rec.Host, + Port: rec.Port, + User: rec.User, + Password: rec.Password, + Database: rec.Database, + FilePath: rec.FilePath, + SSLMode: rec.SSLMode, + } + } + } } if cfg == nil { return nil, nil, nil, fmt.Errorf("no database connection specified.\n\nUse --db-url, --conn, or connect using:\n sql-doctor connect --type ...") } + // 5. Override database if -d / --database flag is passed + if flagDatabase != "" { + cfg.Database = flagDatabase + } + db, driver, err := database.OpenConnection(ctx, cfg) if err != nil { return nil, nil, nil, err @@ -148,6 +189,55 @@ func GetActiveDB(ctx context.Context) (*sql.DB, database.Driver, *database.Conne return db, driver, cfg, nil } +// EnsureDatabase verifies that a database is selected for operations that require one. +// If not selected, it fetches available databases from the server and returns a helpful error suggestion. +func EnsureDatabase(ctx context.Context, db *sql.DB, driver database.Driver, cfg *database.ConnectionConfig) error { + return formatEnsureDatabase(ctx, db, driver, cfg, false) +} + +// EnsureShellDatabase verifies that a database is selected inside an interactive shell session. +func EnsureShellDatabase(ctx context.Context, db *sql.DB, driver database.Driver, cfg *database.ConnectionConfig) error { + return formatEnsureDatabase(ctx, db, driver, cfg, true) +} + +func formatEnsureDatabase(ctx context.Context, db *sql.DB, driver database.Driver, cfg *database.ConnectionConfig, isShell bool) error { + if cfg.Dialect == database.DialectSQLite { + return nil + } + if cfg.Database != "" { + return nil + } + + dbs, err := driver.Databases(ctx, db) + var suggestion strings.Builder + suggestion.WriteString(fmt.Sprintf("no database selected for connection '%s'.\n\nThis operation requires an active database.\n", cfg.Name)) + + if err == nil && len(dbs) > 0 { + suggestion.WriteString("\nAvailable databases on server:\n") + limit := len(dbs) + if limit > 15 { + limit = 15 + } + for i := 0; i < limit; i++ { + suggestion.WriteString(fmt.Sprintf(" • %s\n", dbs[i])) + } + if len(dbs) > 15 { + suggestion.WriteString(fmt.Sprintf(" ... and %d more\n", len(dbs)-15)) + } + } + + suggestion.WriteString("\nTo select a database, run:\n") + if isShell { + suggestion.WriteString(" use \n") + } else { + suggestion.WriteString(" sql-doctor use \n") + suggestion.WriteString("or specify with flag:\n") + suggestion.WriteString(" --database (or -d )\n") + } + + return fmt.Errorf("%s", suggestion.String()) +} + // GetAIProvider returns an initialized Gemini provider func GetAIProvider() ai.AIProvider { var key, model string diff --git a/internal/cli/schema.go b/internal/cli/schema.go index 5580525..81a8f7c 100644 --- a/internal/cli/schema.go +++ b/internal/cli/schema.go @@ -19,12 +19,16 @@ var schemaAnalyzeCmd = &cobra.Command{ Short: "Analyze schema design smells, missing PKs, unindexed foreign keys", RunE: func(cmd *cobra.Command, args []string) error { ctx := cmd.Context() - db, driver, _, err := GetActiveDB(ctx) + db, driver, cfg, err := GetActiveDB(ctx) if err != nil { return err } defer db.Close() + if err := EnsureDatabase(ctx, db, driver, cfg); err != nil { + return err + } + analyzer := schema.NewSchemaAnalyzer(driver) report, err := analyzer.Analyze(ctx, db) if err != nil { @@ -94,12 +98,16 @@ var schemaDatatypesCmd = &cobra.Command{ ctx := cmd.Context() tableName := args[0] - db, driver, _, err := GetActiveDB(ctx) + db, driver, cfg, err := GetActiveDB(ctx) if err != nil { return err } defer db.Close() + if err := EnsureDatabase(ctx, db, driver, cfg); err != nil { + return err + } + advisor := schema.NewAdvisor(driver) recs, err := advisor.AnalyzeTable(ctx, db, tableName) if err != nil { diff --git a/internal/cli/shell.go b/internal/cli/shell.go new file mode 100644 index 0000000..0b05d3e --- /dev/null +++ b/internal/cli/shell.go @@ -0,0 +1,975 @@ +package cli + +import ( + "bufio" + "context" + "database/sql" + "fmt" + "os" + "strconv" + "strings" + "time" + + "github.com/charmbracelet/lipgloss" + "github.com/spf13/cobra" + "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/storage" + "github.com/sql-doctor/sql-doctor/internal/ui" + "golang.org/x/term" +) + +var ( + shellFlagConn string + shellFlagDatabase string +) + +type ShellSession struct { + DB *sql.DB + Driver database.Driver + Config *database.ConnectionConfig +} + +var shellCmd = &cobra.Command{ + Use: "shell", + Aliases: []string{"console", "repl"}, + Short: "Start an interactive SQL Doctor shell / REPL session", + Long: `Start an interactive shell where you can execute queries and database commands +directly without prefixing each command with 'sql-doctor'. + +Session state (active database and memory) is kept in memory during the shell session +and is completely discarded when you type 'exit' or 'quit'.`, + RunE: func(cmd *cobra.Command, args []string) error { + ctx := cmd.Context() + + if shellFlagConn != "" { + flagConn = shellFlagConn + } + + var cfg *database.ConnectionConfig + + // If user specified --conn or --db-url, connect directly + if flagConn != "" || flagDBURL != "" { + c, _, _, err := GetActiveDB(ctx) + if err != nil { + return err + } + c.Close() + // Fetch config + db, driver, resolvedCfg, err := GetActiveDB(ctx) + if err != nil { + return err + } + defer db.Close() + return RunShell(ctx, db, driver, resolvedCfg) + } + + // Otherwise, launch interactive Driver & Connection Wizard + reader := bufio.NewReader(os.Stdin) + var list []storage.ConnectionRecord + var sess *storage.ConnectionRecord + if appStorage != nil { + list, _ = appStorage.ListConnections(ctx) + sess, _ = appStorage.GetSessionConnection(ctx) + } + + chosenCfg, err := promptDriverWizard(ctx, reader, list, sess) + if err != nil { + return err + } + cfg = chosenCfg + + db, driver, err := database.OpenConnection(ctx, cfg) + if err != nil { + return fmt.Errorf("connection failed: %w", err) + } + defer db.Close() + + return RunShell(ctx, db, driver, cfg) + }, +} + +var disconnectCmd = &cobra.Command{ + Use: "disconnect", + Short: "Disconnect and clear the current active session memory", + RunE: func(cmd *cobra.Command, args []string) error { + ctx := cmd.Context() + if appStorage != nil { + _ = appStorage.ClearSessionConnection(ctx) + } + fmt.Println(ui.Success("Active session disconnected. Memory cleared.")) + return nil + }, +} + +func init() { + shellCmd.Flags().StringVarP(&shellFlagConn, "conn", "c", "", "Connection profile to use for shell session") + shellCmd.Flags().StringVarP(&shellFlagDatabase, "database", "d", "", "Database to select upon entering shell") + + RootCmd.AddCommand(shellCmd) + RootCmd.AddCommand(disconnectCmd) +} + +// RunShell starts the interactive loop +func RunShell(ctx context.Context, db *sql.DB, driver database.Driver, cfg *database.ConnectionConfig) error { + shellCfg := *cfg + if shellFlagDatabase != "" { + shellCfg.Database = shellFlagDatabase + } + + session := &ShellSession{ + DB: db, + Driver: driver, + Config: &shellCfg, + } + + ver, _ := session.Driver.Version(ctx, session.DB) + primaryStyle := lipgloss.NewStyle().Bold(true).Foreground(ui.PrimaryColor) + + fmt.Println() + fmt.Println(ui.TitleStyle.Render("SQL Doctor Interactive Shell")) + fmt.Printf("Connected to: %s (%s)\n", primaryStyle.Render(session.Config.Name), ver) + if session.Config.Database != "" { + fmt.Printf("Active Database: %s\n", primaryStyle.Render(session.Config.Database)) + } else if session.Config.Dialect != database.DialectSQLite { + fmt.Println(ui.Info("No database selected yet. Type 'use ' or 'databases' to choose one.")) + } + fmt.Println("Type " + primaryStyle.Render("help") + " for commands, or " + primaryStyle.Render("exit") + " to quit (session memory is discarded on exit).\n") + + reader := bufio.NewReader(os.Stdin) + + for { + prompt := renderPrompt(session.Config) + fmt.Print(prompt) + + input, err := reader.ReadString('\n') + if err != nil { + fmt.Println("\nExiting shell session...") + break + } + + line := strings.TrimSpace(input) + if line == "" { + continue + } + + if line == "exit" || line == "quit" || line == "\\q" { + fmt.Println(ui.Info("Disconnected. Shell session memory cleared. Bye!")) + break + } + + if line == "clear" || line == "cls" { + fmt.Print("\033[H\033[2J") + continue + } + + lowerTrimmed := strings.TrimSuffix(strings.ToLower(line), ";") + fields := strings.Fields(lowerTrimmed) + if len(fields) > 0 && (fields[0] == "help" || fields[0] == "?" || fields[0] == "\\h" || fields[0] == "\\?") { + printShellHelp(fields[1:]...) + continue + } + + if err := handleShellCommand(ctx, session, line); err != nil { + fmt.Println(ui.Error("%v", err)) + } + fmt.Println() + } + + return nil +} + +func renderPrompt(cfg *database.ConnectionConfig) string { + engineStyle := lipgloss.NewStyle().Bold(true).Foreground(lipgloss.Color("#6366F1")) + dbStyle := lipgloss.NewStyle().Foreground(lipgloss.Color("#10B981")) + noneStyle := lipgloss.NewStyle().Foreground(lipgloss.Color("#9CA3AF")) + + dbPart := noneStyle.Render("(none)") + if cfg.Database != "" { + dbPart = dbStyle.Render(cfg.Database) + } + + return fmt.Sprintf("%s [%s@%s]> ", + engineStyle.Render("sql-doctor"), + cfg.Dialect, + dbPart, + ) +} + +func printShellHelp(args ...string) { + sectionStyle := lipgloss.NewStyle().Bold(true).Foreground(lipgloss.Color("#F59E0B")) + cmdStyle := lipgloss.NewStyle().Bold(true).Foreground(lipgloss.Color("#10B981")) + descStyle := lipgloss.NewStyle().Foreground(lipgloss.Color("#F3F4F6")) + exampleStyle := lipgloss.NewStyle().Foreground(lipgloss.Color("#38BDF8")) + syntaxStyle := lipgloss.NewStyle().Bold(true).Foreground(lipgloss.Color("#EC4899")) + + if len(args) > 0 { + topic := strings.ToLower(args[0]) + fmt.Println() + switch topic { + case "use": + fmt.Println(sectionStyle.Render("COMMAND: use")) + fmt.Println(descStyle.Render("Switch active database for the current connection (in-memory session only).")) + fmt.Println() + fmt.Println(syntaxStyle.Render("Syntax:") + " use ") + fmt.Println(exampleStyle.Render("Example:") + " use rolerift") + + case "databases", "dbs": + fmt.Println(sectionStyle.Render("COMMAND: databases (or dbs)")) + fmt.Println(descStyle.Render("List all available databases on the connected database server.")) + fmt.Println() + fmt.Println(syntaxStyle.Render("Syntax:") + " databases") + + case "connections", "profiles": + fmt.Println(sectionStyle.Render("COMMAND: connections (or profiles)")) + fmt.Println(descStyle.Render("List all saved connection profiles with dialect, host, and database.")) + fmt.Println() + fmt.Println(syntaxStyle.Render("Syntax:") + " connections") + + case "connect": + fmt.Println(sectionStyle.Render("COMMAND: connect")) + fmt.Println(descStyle.Render("Switch active connection to a different saved connection profile without restarting.")) + fmt.Println() + fmt.Println(syntaxStyle.Render("Syntax:") + " connect ") + fmt.Println(exampleStyle.Render("Example:") + " connect prod-db") + + case "tables": + fmt.Println(sectionStyle.Render("COMMAND: tables")) + fmt.Println(descStyle.Render("List all tables in the active database with type, row counts, and storage engine.")) + fmt.Println() + fmt.Println(syntaxStyle.Render("Syntax:") + " tables") + + case "describe", "desc": + fmt.Println(sectionStyle.Render("COMMAND: describe (or desc)")) + fmt.Println(descStyle.Render("Inspect table schema, column types, nullability, default values, and keys.")) + fmt.Println() + fmt.Println(syntaxStyle.Render("Syntax:") + " describe ") + fmt.Println(exampleStyle.Render("Example:") + " describe users") + + case "indexes": + fmt.Println(sectionStyle.Render("COMMAND: indexes")) + fmt.Println(descStyle.Render("List indexes on a table, including column order, uniqueness, and primary status.")) + fmt.Println() + fmt.Println(syntaxStyle.Render("Syntax:") + " indexes ") + fmt.Println(exampleStyle.Render("Example:") + " indexes users") + + case "relationships": + fmt.Println(sectionStyle.Render("COMMAND: relationships")) + fmt.Println(descStyle.Render("Map explicit foreign keys and detect naming-based inferred relationships.")) + fmt.Println() + fmt.Println(syntaxStyle.Render("Syntax:") + " relationships") + + case "analyze": + fmt.Println(sectionStyle.Render("COMMAND: analyze")) + fmt.Println(descStyle.Render("Run deep performance analysis on a query (measures time, row scans, full table scans, performance score).")) + fmt.Println() + fmt.Println(syntaxStyle.Render("Syntax:") + " analyze ") + fmt.Println(exampleStyle.Render("Example:") + " analyze SELECT * FROM users WHERE email = 'test@example.com'") + + case "explain": + fmt.Println(sectionStyle.Render("COMMAND: explain")) + fmt.Println(descStyle.Render("Visualize the hierarchical execution plan tree for a SQL query.")) + fmt.Println() + fmt.Println(syntaxStyle.Render("Syntax:") + " explain ") + fmt.Println(exampleStyle.Render("Example:") + " explain SELECT * FROM users JOIN roles ON users.role_id = roles.id") + + case "optimize": + fmt.Println(sectionStyle.Render("COMMAND: optimize")) + fmt.Println(descStyle.Render("Analyze query plan and recommend composite indexes or SQL query rewrites.")) + fmt.Println() + fmt.Println(syntaxStyle.Render("Syntax:") + " optimize ") + fmt.Println(exampleStyle.Render("Example:") + " optimize SELECT * FROM orders WHERE status = 'paid' AND created_at > '2026-01-01'") + + case "doctor": + fmt.Println(sectionStyle.Render("COMMAND: doctor")) + fmt.Println(descStyle.Render("Run full automated health diagnostics across the active database.")) + fmt.Println() + fmt.Println(syntaxStyle.Render("Syntax:") + " doctor") + + case "ask": + fmt.Println(sectionStyle.Render("COMMAND: ask")) + fmt.Println(descStyle.Render("Use Gemini AI to answer questions or generate SQL based on schema.")) + fmt.Println() + fmt.Println(syntaxStyle.Render("Syntax:") + " ask ") + fmt.Println(exampleStyle.Render("Example:") + " ask find the 5 newest registered users") + + case "sql", "query", "select": + fmt.Println(sectionStyle.Render("DIRECT SQL EXECUTION")) + fmt.Println(descStyle.Render("Execute any SQL query directly against the active database.")) + fmt.Println() + fmt.Println(syntaxStyle.Render("Syntax:") + "
", "Inspect table columns, types, nullability, keys, and defaults"}, + {"indexes
", "List table indexes, uniqueness, and composite column orders"}, + {"relationships", "Inspect detected foreign keys and inferred relationships"}, + }) + + printCategory("PERFORMANCE & DIAGNOSTICS", [][2]string{ + {"analyze ", "Run deep performance audit, rows scanned, and smell detection"}, + {"explain ", "Visualize hierarchical query execution plan tree"}, + {"optimize ", "Recommend missing indexes and query rewrite improvements"}, + {"doctor", "Run full end-to-end diagnostic health check on active database"}, + {"ask ", "Generate SQL or ask schema questions using Gemini AI"}, + }) + + printCategory("DIRECT SQL & UTILITIES", [][2]string{ + {"", "Run standard SQL directly (SELECT, SHOW, INSERT, UPDATE, etc.)"}, + {"\\G", "Display query results in vertical format (one column per line)"}, + {"clear (or cls)", "Clear terminal screen"}, + {"help [command]", "Show full reference or detailed help for a specific command"}, + {"exit (or quit / \\q)", "Disconnect and leave shell session"}, + }) + + fmt.Println(sectionStyle.Render("EXAMPLES")) + fmt.Printf(" %s\n", exampleStyle.Render("sql-doctor> use rolerift")) + fmt.Printf(" %s\n", exampleStyle.Render("sql-doctor> tables")) + fmt.Printf(" %s\n", exampleStyle.Render("sql-doctor> select * from users\\G")) + fmt.Printf(" %s\n", exampleStyle.Render("sql-doctor> analyze SELECT * FROM users WHERE email = 'test@example.com'")) + fmt.Printf(" %s\n", exampleStyle.Render("sql-doctor> help analyze")) + fmt.Println() +} + +func handleShellCommand(ctx context.Context, session *ShellSession, line string) error { + trimmed := strings.TrimSuffix(line, ";") + lower := strings.ToLower(trimmed) + parts := strings.Fields(trimmed) + cmdName := strings.ToLower(parts[0]) + + switch cmdName { + case "connections", "profiles": + if appStorage == nil { + return fmt.Errorf("local storage unavailable") + } + list, err := appStorage.ListConnections(ctx) + if err != nil { + return err + } + tbl := ui.NewTable("CURRENT", "NAME", "DIALECT", "HOST/FILE", "DATABASE", "USER") + for _, c := range list { + marker := "" + if c.Name == session.Config.Name { + marker = " ★ " + } + hostOrFile := c.Host + if c.FilePath != "" { + hostOrFile = c.FilePath + } + dbName := c.Database + if dbName == "" { + dbName = "(none)" + } + tbl.AddRow(marker, c.Name, string(c.Dialect), hostOrFile, dbName, c.User) + } + fmt.Println(tbl.Render()) + fmt.Println(ui.Info("Tip: Switch connection using 'connect '")) + return nil + + case "connect": + if len(parts) < 2 { + return fmt.Errorf("usage: connect (type 'connections' to view available profiles)") + } + targetName := parts[1] + if appStorage == nil { + return fmt.Errorf("local storage unavailable") + } + rec, err := appStorage.GetConnection(ctx, targetName) + if err != nil { + return fmt.Errorf("connection profile '%s' not found: %w", targetName, err) + } + newCfg := &database.ConnectionConfig{ + Name: rec.Name, + Dialect: rec.Dialect, + Host: rec.Host, + Port: rec.Port, + User: rec.User, + Password: rec.Password, + Database: rec.Database, + FilePath: rec.FilePath, + SSLMode: rec.SSLMode, + } + newDB, newDriver, err := database.OpenConnection(ctx, newCfg) + if err != nil { + return fmt.Errorf("failed to connect to '%s': %w", targetName, err) + } + ver, _ := newDriver.Version(ctx, newDB) + session.DB.Close() + session.DB = newDB + session.Driver = newDriver + session.Config = newCfg + fmt.Println(ui.Success("Switched connection to '%s' (%s)", targetName, ver)) + if session.Config.Database != "" { + fmt.Println(ui.Info("Active database: %s", session.Config.Database)) + } + return nil + + case "use": + if len(parts) < 2 { + return fmt.Errorf("usage: use ") + } + targetDB := parts[1] + if session.Config.Dialect == database.DialectSQLite { + return fmt.Errorf("SQLite uses single-file database. Use of multiple databases is not supported.") + } + dbs, err := session.Driver.Databases(ctx, session.DB) + if err == nil && len(dbs) > 0 { + found := false + for _, d := range dbs { + if strings.EqualFold(d, targetDB) { + targetDB = d + found = true + break + } + } + if !found { + return fmt.Errorf("database '%s' not found on server", targetDB) + } + } + session.Config.Database = targetDB + if session.Driver.Dialect() == database.DialectMySQL || session.Driver.Dialect() == database.DialectMariaDB { + _, _ = session.DB.ExecContext(ctx, fmt.Sprintf("USE `%s`", targetDB)) + } + fmt.Println(ui.Success("Database switched to '%s'", targetDB)) + return nil + + case "databases", "dbs": + dbs, err := session.Driver.Databases(ctx, session.DB) + if err != nil { + return err + } + tbl := ui.NewTable("ACTIVE", "DATABASE NAME") + for _, d := range dbs { + activeMarker := "" + if session.Config.Database == d { + activeMarker = " ★ " + } + tbl.AddRow(activeMarker, d) + } + fmt.Println(tbl.Render()) + return nil + + case "tables": + if err := EnsureShellDatabase(ctx, session.DB, session.Driver, session.Config); err != nil { + return err + } + tables, err := session.Driver.Tables(ctx, session.DB) + if err != nil { + return err + } + tbl := ui.NewTable("TABLE NAME", "TYPE", "EST. ROWS", "ENGINE") + for _, t := range tables { + tbl.AddRow(t.Name, t.Type, fmt.Sprintf("%d", t.RowCount), t.Engine) + } + fmt.Println(tbl.Render()) + return nil + + case "describe", "desc": + if len(parts) < 2 { + return fmt.Errorf("usage: describe ") + } + if err := EnsureShellDatabase(ctx, session.DB, session.Driver, session.Config); err != nil { + return err + } + detail, err := session.Driver.DescribeTable(ctx, session.DB, parts[1]) + if err != nil { + return err + } + tbl := ui.NewTable("#", "COLUMN", "TYPE", "NULLABLE", "DEFAULT", "KEY") + for _, c := range detail.Columns { + keyStr := "" + if c.IsPrimaryKey { + keyStr = "PRI" + } else if c.IsForeignKey { + keyStr = "FK" + } else if c.IsUnique { + keyStr = "UNI" + } + nullStr := "NO" + if c.IsNullable { + nullStr = "YES" + } + dfltStr := "NULL" + if c.DefaultValue != nil { + dfltStr = *c.DefaultValue + } + tbl.AddRow(fmt.Sprintf("%d", c.Position), c.Name, c.RawType, nullStr, dfltStr, keyStr) + } + fmt.Println(tbl.Render()) + return nil + + case "indexes": + if len(parts) < 2 { + return fmt.Errorf("usage: indexes ") + } + if err := EnsureShellDatabase(ctx, session.DB, session.Driver, session.Config); err != nil { + return err + } + indexes, err := session.Driver.Indexes(ctx, session.DB, parts[1]) + if err != nil { + return err + } + tbl := ui.NewTable("INDEX NAME", "COLUMNS", "UNIQUE", "PRIMARY", "TYPE") + for _, idx := range indexes { + tbl.AddRow(idx.Name, strings.Join(idx.Columns, ", "), fmt.Sprintf("%v", idx.IsUnique), fmt.Sprintf("%v", idx.IsPrimary), idx.Type) + } + fmt.Println(tbl.Render()) + return nil + + case "relationships": + if err := EnsureShellDatabase(ctx, session.DB, session.Driver, session.Config); err != nil { + return err + } + rels, err := session.Driver.Relationships(ctx, session.DB) + if err != nil { + return err + } + if len(rels) == 0 { + fmt.Println(ui.Info("No relationships or foreign keys detected.")) + return nil + } + tbl := ui.NewTable("SOURCE TABLE", "COLUMN", "TARGET TABLE", "TARGET COL", "TYPE", "ORPHANS") + for _, r := range rels { + relType := "Explicit FK" + if !r.IsExplicit { + relType = "Inferred (Naming)" + } + orphanStr := fmt.Sprintf("%d", r.OrphanCount) + if r.OrphanCount > 0 { + orphanStr = ui.WarningBadge + fmt.Sprintf(" %d", r.OrphanCount) + } + tbl.AddRow(r.FromTable, strings.Join(r.FromColumns, ","), r.ToTable, strings.Join(r.ToColumns, ","), relType, orphanStr) + } + fmt.Println(tbl.Render()) + return nil + + case "analyze": + if len(parts) < 2 { + return fmt.Errorf("usage: analyze ") + } + queryText := strings.TrimPrefix(trimmed, parts[0]+" ") + if err := EnsureShellDatabase(ctx, session.DB, session.Driver, session.Config); err != nil { + return err + } + anz := analyzer.NewAnalyzer(session.Driver) + result, err := anz.Analyze(ctx, session.DB, queryText) + if err != nil { + return err + } + fmt.Println("Performance Score: " + ui.RenderScoreMeter(result.PerformanceScore)) + tbl := ui.NewTable("METRIC", "VALUE") + tbl.AddRow("Execution Time", fmt.Sprintf("%.2f ms", result.ExecutionTimeMs)) + tbl.AddRow("Rows Returned", fmt.Sprintf("%d", result.RowsReturned)) + tbl.AddRow("Rows Examined", fmt.Sprintf("%d", result.RowsExamined)) + tbl.AddRow("Full Table Scan", fmt.Sprintf("%v", result.HasFullTableScan)) + fmt.Println(tbl.Render()) + return nil + + case "explain": + if len(parts) < 2 { + return fmt.Errorf("usage: explain ") + } + queryText := strings.TrimPrefix(trimmed, parts[0]+" ") + if err := EnsureShellDatabase(ctx, session.DB, session.Driver, session.Config); err != nil { + return err + } + explainRes, err := session.Driver.Explain(ctx, session.DB, queryText, false) + if err != nil { + return err + } + planAnalyzer := explain.NewPlanAnalyzer() + summary := planAnalyzer.Analyze(explainRes) + if summary.FormattedTree != "" { + fmt.Println(ui.HeaderStyle.Render("Execution Plan Tree:")) + fmt.Println(summary.FormattedTree) + } + fmt.Println(summary.HumanSummary) + return nil + + case "optimize": + if len(parts) < 2 { + return fmt.Errorf("usage: optimize ") + } + queryText := strings.TrimPrefix(trimmed, parts[0]+" ") + if err := EnsureShellDatabase(ctx, session.DB, session.Driver, session.Config); err != nil { + return err + } + opt := optimizer.NewOptimizer(session.Driver) + res, err := opt.Optimize(ctx, session.DB, queryText) + if err != nil { + return err + } + if len(res.IndexRecommendations) > 0 { + for _, idx := range res.IndexRecommendations { + fmt.Printf("%s Target: %s\n DDL: %s\n Gain: %s\n", ui.SuccessBadge, idx.TableName, idx.SuggestedDDL, idx.EstimatedGain) + } + } else { + fmt.Println(ui.Success("No new indexes required for this query.")) + } + return nil + + case "doctor": + if err := EnsureShellDatabase(ctx, session.DB, session.Driver, session.Config); err != nil { + return err + } + fmt.Println(ui.Info("Running health diagnostics on [%s]...", session.Config.Database)) + tables, err := session.Driver.Tables(ctx, session.DB) + if err != nil { + return err + } + fmt.Println(ui.Success("Database '%s' is responsive. Total tables: %d", session.Config.Database, len(tables))) + return nil + + case "help", "?": + printShellHelp(parts[1:]...) + return nil + + default: + if isSQLStatement(lower) { + return executeDirectSQL(ctx, session.DB, session.Driver, session.Config, trimmed) + } + return fmt.Errorf("unknown command '%s'. Type 'help' for available commands.", parts[0]) + } +} + +func isSQLStatement(lower string) bool { + prefixes := []string{ + "select", "insert", "update", "delete", "create", "alter", "drop", + "show", "describe", "explain", "set", "call", "truncate", "begin", "commit", "rollback", + } + for _, p := range prefixes { + if strings.HasPrefix(lower, p) { + return true + } + } + return false +} + +func executeDirectSQL(ctx context.Context, db *sql.DB, driver database.Driver, cfg *database.ConnectionConfig, sqlQuery string) error { + trimmedQuery := strings.TrimSpace(sqlQuery) + isVertical := strings.HasSuffix(trimmedQuery, "\\G") || strings.HasSuffix(trimmedQuery, "\\g") + if isVertical { + trimmedQuery = strings.TrimSpace(trimmedQuery[:len(trimmedQuery)-2]) + } + + upper := strings.ToUpper(trimmedQuery) + + if strings.HasPrefix(upper, "SELECT") || strings.HasPrefix(upper, "SHOW") || + strings.HasPrefix(upper, "EXPLAIN") || strings.HasPrefix(upper, "DESCRIBE") || strings.HasPrefix(upper, "DESC ") { + + start := time.Now() + rows, err := db.QueryContext(ctx, trimmedQuery) + if err != nil { + return err + } + defer rows.Close() + + cols, err := rows.Columns() + if err != nil { + return err + } + + tbl := ui.NewTable(cols...) + var allRows [][]string + rowCount := 0 + + for rows.Next() { + rowCount++ + vals := make([]interface{}, len(cols)) + valPtrs := make([]interface{}, len(cols)) + for i := range vals { + valPtrs[i] = &vals[i] + } + + if err := rows.Scan(valPtrs...); err != nil { + return err + } + + rowStrs := make([]string, len(cols)) + for i, v := range vals { + if v == nil { + rowStrs[i] = "NULL" + } else if t, ok := v.(time.Time); ok { + rowStrs[i] = t.Format("2006-01-02 15:04:05") + } else if b, ok := v.([]byte); ok { + rowStrs[i] = string(b) + } else { + rowStrs[i] = fmt.Sprintf("%v", v) + } + } + if isVertical { + allRows = append(allRows, rowStrs) + } else { + tbl.AddRow(rowStrs...) + } + } + + duration := time.Since(start) + if rowCount == 0 { + fmt.Println(ui.Info("Empty set (%.2f ms)", float64(duration.Microseconds())/1000.0)) + return nil + } + + if isVertical { + maxColLen := 0 + for _, col := range cols { + if len(col) > maxColLen { + maxColLen = len(col) + } + } + for rIdx, r := range allRows { + fmt.Printf("*************************** %d. row ***************************\n", rIdx+1) + for cIdx, val := range r { + fmt.Printf("%*s: %s\n", maxColLen, cols[cIdx], val) + } + } + } else { + fmt.Println(tbl.Render()) + if tbl.WasTruncated() { + fmt.Println(ui.Info("Tip: Wide table fitted to terminal width. Use '\\G' (e.g. %s\\G) for full vertical view.", trimmedQuery)) + } + } + + fmt.Println(ui.Info("%d rows in set (%.2f ms)", rowCount, float64(duration.Microseconds())/1000.0)) + return nil + } + + start := time.Now() + res, err := db.ExecContext(ctx, trimmedQuery) + if err != nil { + return err + } + affected, _ := res.RowsAffected() + duration := time.Since(start) + fmt.Println(ui.Success("Query OK, %d rows affected (%.2f ms)", affected, float64(duration.Microseconds())/1000.0)) + return nil +} + +func promptDriverWizard(ctx context.Context, reader *bufio.Reader, savedProfiles []storage.ConnectionRecord, activeSess *storage.ConnectionRecord) (*database.ConnectionConfig, error) { + fmt.Println() + fmt.Println(ui.TitleStyle.Render("Select Database Driver / Engine:")) + fmt.Println(" [1] MySQL") + fmt.Println(" [2] PostgreSQL") + fmt.Println(" [3] SQLite") + fmt.Println(" [4] MariaDB") + + choiceMap := make(map[int]interface{}) + choiceMap[1] = database.DialectMySQL + choiceMap[2] = database.DialectPostgreSQL + choiceMap[3] = database.DialectSQLite + choiceMap[4] = database.DialectMariaDB + + nextIdx := 5 + sessIdx := 0 + if activeSess != nil { + sessIdx = nextIdx + choiceMap[sessIdx] = activeSess + sessTarget := activeSess.Host + if activeSess.FilePath != "" { + sessTarget = activeSess.FilePath + } + dbInfo := "" + if activeSess.Database != "" { + dbInfo = " | db: " + activeSess.Database + } + fmt.Printf(" [%d] Active Session (%s: %s%s)\n", sessIdx, activeSess.Dialect, sessTarget, dbInfo) + nextIdx++ + } + + savedStartIdx := 0 + if len(savedProfiles) > 0 { + savedStartIdx = nextIdx + fmt.Printf(" [%d] Use Saved Profile (%d available)\n", savedStartIdx, len(savedProfiles)) + nextIdx++ + } + + maxChoice := nextIdx - 1 + defaultChoice := 1 + if sessIdx != 0 { + defaultChoice = sessIdx + } + + fmt.Printf("\nEnter choice [1-%d] (default %d): ", maxChoice, defaultChoice) + choiceInput, _ := reader.ReadString('\n') + choiceInput = strings.TrimSpace(choiceInput) + + chosenNum := defaultChoice + if choiceInput != "" { + if n, err := strconv.Atoi(choiceInput); err == nil && n >= 1 && n <= maxChoice { + chosenNum = n + } + } + + switch chosenNum { + case 1: // MySQL + host := promptDefault(reader, "Host", "127.0.0.1") + portStr := promptDefault(reader, "Port", "3306") + port, _ := strconv.Atoi(portStr) + user := promptDefault(reader, "User", "root") + fmt.Print("Password (leave empty if none): ") + password := readPasswordInput(reader) + dbName := promptDefault(reader, "Database name (optional, press Enter to skip)", "") + return &database.ConnectionConfig{ + Name: "mysql", + Dialect: database.DialectMySQL, + Host: host, + Port: port, + User: user, + Password: password, + Database: dbName, + }, nil + + case 2: // PostgreSQL + host := promptDefault(reader, "Host", "127.0.0.1") + portStr := promptDefault(reader, "Port", "5432") + port, _ := strconv.Atoi(portStr) + user := promptDefault(reader, "User", "postgres") + fmt.Print("Password (leave empty if none): ") + password := readPasswordInput(reader) + dbName := promptDefault(reader, "Database name (optional, press Enter to skip)", "") + return &database.ConnectionConfig{ + Name: "postgres", + Dialect: database.DialectPostgreSQL, + Host: host, + Port: port, + User: user, + Password: password, + Database: dbName, + }, nil + + case 3: // SQLite + path := promptDefault(reader, "Database file path", "./dev.db") + return &database.ConnectionConfig{ + Name: "sqlite", + Dialect: database.DialectSQLite, + FilePath: path, + Database: path, + }, nil + + case 4: // MariaDB + host := promptDefault(reader, "Host", "127.0.0.1") + portStr := promptDefault(reader, "Port", "3306") + port, _ := strconv.Atoi(portStr) + user := promptDefault(reader, "User", "root") + fmt.Print("Password (leave empty if none): ") + password := readPasswordInput(reader) + dbName := promptDefault(reader, "Database name (optional, press Enter to skip)", "") + return &database.ConnectionConfig{ + Name: "mariadb", + Dialect: database.DialectMariaDB, + Host: host, + Port: port, + User: user, + Password: password, + Database: dbName, + }, nil + + default: + if chosenNum == sessIdx && activeSess != nil { + return &database.ConnectionConfig{ + Name: activeSess.Name, + Dialect: activeSess.Dialect, + Host: activeSess.Host, + Port: activeSess.Port, + User: activeSess.User, + Password: activeSess.Password, + Database: activeSess.Database, + FilePath: activeSess.FilePath, + SSLMode: activeSess.SSLMode, + }, nil + } + + if chosenNum == savedStartIdx && len(savedProfiles) > 0 { + fmt.Println("\nSaved Connection Profiles:") + for i, p := range savedProfiles { + tgt := p.Host + if p.FilePath != "" { + tgt = p.FilePath + } + fmt.Printf(" [%d] %s (%s: %s)\n", i+1, p.Name, p.Dialect, tgt) + } + fmt.Printf("Select profile [1-%d]: ", len(savedProfiles)) + pChoice, _ := reader.ReadString('\n') + pChoice = strings.TrimSpace(pChoice) + pIdx, err := strconv.Atoi(pChoice) + if err == nil && pIdx >= 1 && pIdx <= len(savedProfiles) { + rec := savedProfiles[pIdx-1] + return &database.ConnectionConfig{ + Name: rec.Name, + Dialect: rec.Dialect, + Host: rec.Host, + Port: rec.Port, + User: rec.User, + Password: rec.Password, + Database: rec.Database, + FilePath: rec.FilePath, + SSLMode: rec.SSLMode, + }, nil + } + } + + return nil, fmt.Errorf("invalid choice") + } +} + +func promptDefault(reader *bufio.Reader, label, dflt string) string { + if dflt != "" { + fmt.Printf("%s [%s]: ", label, dflt) + } else { + fmt.Printf("%s: ", label) + } + input, _ := reader.ReadString('\n') + val := strings.TrimSpace(input) + if val == "" { + return dflt + } + return val +} + +func readPasswordInput(reader *bufio.Reader) string { + if term.IsTerminal(int(os.Stdin.Fd())) { + bytePass, err := term.ReadPassword(int(os.Stdin.Fd())) + fmt.Println() + if err == nil { + return strings.TrimRight(string(bytePass), "\r\n") + } + } + pass, _ := reader.ReadString('\n') + return strings.TrimRight(pass, "\r\n") +} diff --git a/internal/database/driver.go b/internal/database/driver.go index 5234d85..39600f5 100644 --- a/internal/database/driver.go +++ b/internal/database/driver.go @@ -172,6 +172,7 @@ type Driver interface { Connect(ctx context.Context, config *ConnectionConfig) (*sql.DB, error) Ping(ctx context.Context, db *sql.DB) error Version(ctx context.Context, db *sql.DB) (string, error) + Databases(ctx context.Context, db *sql.DB) ([]string, error) Tables(ctx context.Context, db *sql.DB) ([]TableInfo, error) DescribeTable(ctx context.Context, db *sql.DB, table string) (*TableDetail, error) Indexes(ctx context.Context, db *sql.DB, table string) ([]IndexInfo, error) diff --git a/internal/database/mysql/driver.go b/internal/database/mysql/driver.go index 299bee3..2af174f 100644 --- a/internal/database/mysql/driver.go +++ b/internal/database/mysql/driver.go @@ -97,6 +97,24 @@ func (d *Driver) Version(ctx context.Context, db *sql.DB) (string, error) { return res, nil } +func (d *Driver) Databases(ctx context.Context, db *sql.DB) ([]string, error) { + rows, err := db.QueryContext(ctx, "SHOW DATABASES;") + if err != nil { + return nil, fmt.Errorf("failed to list databases: %w", err) + } + defer rows.Close() + + var databases []string + for rows.Next() { + var name string + if err := rows.Scan(&name); err != nil { + return nil, err + } + databases = append(databases, name) + } + return databases, nil +} + func (d *Driver) Tables(ctx context.Context, db *sql.DB) ([]database.TableInfo, error) { query := ` SELECT diff --git a/internal/database/postgres/driver.go b/internal/database/postgres/driver.go index d4d3e9a..64331b9 100644 --- a/internal/database/postgres/driver.go +++ b/internal/database/postgres/driver.go @@ -40,7 +40,11 @@ func (d *Driver) DSN(cfg *database.ConnectionConfig) string { if cfg.Password != "" { auth += ":" + cfg.Password } - return fmt.Sprintf("postgres://%s@%s:%d/%s?sslmode=%s", auth, host, port, cfg.Database, sslMode) + dbName := cfg.Database + if dbName == "" { + dbName = "postgres" + } + return fmt.Sprintf("postgres://%s@%s:%d/%s?sslmode=%s", auth, host, port, dbName, sslMode) } func (d *Driver) Connect(ctx context.Context, cfg *database.ConnectionConfig) (*sql.DB, error) { @@ -82,6 +86,30 @@ func (d *Driver) Version(ctx context.Context, db *sql.DB) (string, error) { return ver, nil } +func (d *Driver) Databases(ctx context.Context, db *sql.DB) ([]string, error) { + query := ` + SELECT datname + FROM pg_database + WHERE datistemplate = false + ORDER BY datname ASC; + ` + rows, err := db.QueryContext(ctx, query) + if err != nil { + return nil, fmt.Errorf("failed to list databases: %w", err) + } + defer rows.Close() + + var databases []string + for rows.Next() { + var name string + if err := rows.Scan(&name); err != nil { + return nil, err + } + databases = append(databases, name) + } + return databases, nil +} + func (d *Driver) Tables(ctx context.Context, db *sql.DB) ([]database.TableInfo, error) { query := ` SELECT diff --git a/internal/database/sqlite/driver.go b/internal/database/sqlite/driver.go index 5e9894d..90f8381 100644 --- a/internal/database/sqlite/driver.go +++ b/internal/database/sqlite/driver.go @@ -67,6 +67,27 @@ func (d *Driver) Version(ctx context.Context, db *sql.DB) (string, error) { return "SQLite " + ver, nil } +func (d *Driver) Databases(ctx context.Context, db *sql.DB) ([]string, error) { + rows, err := db.QueryContext(ctx, "PRAGMA database_list;") + if err != nil { + return []string{"main"}, nil + } + defer rows.Close() + + var databases []string + for rows.Next() { + var seq int + var name, file string + if err := rows.Scan(&seq, &name, &file); err == nil { + databases = append(databases, name) + } + } + if len(databases) == 0 { + return []string{"main"}, nil + } + return databases, nil +} + func (d *Driver) Tables(ctx context.Context, db *sql.DB) ([]database.TableInfo, error) { query := ` SELECT name, type diff --git a/internal/storage/sqlite.go b/internal/storage/sqlite.go index 0767d70..34729c2 100644 --- a/internal/storage/sqlite.go +++ b/internal/storage/sqlite.go @@ -414,3 +414,82 @@ func (s *Storage) GetSetting(ctx context.Context, key string) (string, error) { } return val, nil } + +// SaveSessionConnection saves an ephemeral/session connection +func (s *Storage) SaveSessionConnection(ctx context.Context, conn *ConnectionRecord) error { + data, err := json.Marshal(conn) + if err != nil { + return err + } + return s.SetSetting(ctx, "session_connection", string(data)) +} + +// GetSessionConnection retrieves the active session connection if any +func (s *Storage) GetSessionConnection(ctx context.Context) (*ConnectionRecord, error) { + val, err := s.GetSetting(ctx, "session_connection") + if err != nil || val == "" { + return nil, err + } + var conn ConnectionRecord + if err := json.Unmarshal([]byte(val), &conn); err != nil { + return nil, err + } + return &conn, nil +} + +// ClearSessionConnection removes the session connection +func (s *Storage) ClearSessionConnection(ctx context.Context) error { + _, err := s.db.ExecContext(ctx, "DELETE FROM settings WHERE key = 'session_connection';") + return err +} + +// UpdateConnectionDatabase updates the database name for a specific connection or active connection/session +func (s *Storage) UpdateConnectionDatabase(ctx context.Context, connName string, databaseName string) (*ConnectionRecord, error) { + if connName != "" { + rec, err := s.GetConnection(ctx, connName) + if err != nil { + return nil, err + } + rec.Database = databaseName + _, err = s.db.ExecContext(ctx, "UPDATE connections SET database_name = ?, updated_at = CURRENT_TIMESTAMP WHERE name = ?;", databaseName, connName) + if err != nil { + return nil, err + } + sess, _ := s.GetSessionConnection(ctx) + if sess != nil && sess.Name == connName { + sess.Database = databaseName + _ = s.SaveSessionConnection(ctx, sess) + } + return rec, nil + } + + // First check if there is an active session connection + sess, err := s.GetSessionConnection(ctx) + if err == nil && sess != nil { + sess.Database = databaseName + if err := s.SaveSessionConnection(ctx, sess); err != nil { + return nil, err + } + if sess.Name != "" && sess.Name != "default" && sess.Name != "session" { + _, _ = s.db.ExecContext(ctx, "UPDATE connections SET database_name = ?, updated_at = CURRENT_TIMESTAMP WHERE name = ?;", databaseName, sess.Name) + } + return sess, nil + } + + // Otherwise update active saved connection in connections table + active, err := s.GetActiveConnection(ctx) + if err != nil { + return nil, err + } + if active == nil { + return nil, fmt.Errorf("no active connection found to update database") + } + + active.Database = databaseName + _, err = s.db.ExecContext(ctx, "UPDATE connections SET database_name = ?, updated_at = CURRENT_TIMESTAMP WHERE name = ?;", databaseName, active.Name) + if err != nil { + return nil, err + } + return active, nil +} + diff --git a/internal/ui/table.go b/internal/ui/table.go index f15d6b3..c808211 100644 --- a/internal/ui/table.go +++ b/internal/ui/table.go @@ -1,15 +1,19 @@ package ui import ( + "os" "strings" "github.com/charmbracelet/lipgloss" + "github.com/charmbracelet/lipgloss/table" + "golang.org/x/term" ) // Table renders formatted tables in the terminal type Table struct { - headers []string - rows [][]string + headers []string + rows [][]string + truncated bool } func NewTable(headers ...string) *Table { @@ -23,67 +27,101 @@ func (t *Table) AddRow(cols ...string) *Table { return t } +func (t *Table) WasTruncated() bool { + return t.truncated +} + +func getTerminalWidth() int { + if w, _, err := term.GetSize(int(os.Stdout.Fd())); err == nil && w > 0 { + return w + } + if w, _, err := term.GetSize(int(os.Stderr.Fd())); err == nil && w > 0 { + return w + } + if w, _, err := term.GetSize(int(os.Stdin.Fd())); err == nil && w > 0 { + return w + } + return 0 +} + func (t *Table) Render() string { if len(t.headers) == 0 && len(t.rows) == 0 { return "" } - colWidths := make([]int, len(t.headers)) - for i, h := range t.headers { - colWidths[i] = lipgloss.Width(h) - } - + numCols := len(t.headers) for _, row := range t.rows { - for i, col := range row { - if i < len(colWidths) { - w := lipgloss.Width(col) - if w > colWidths[i] { - colWidths[i] = w - } - } + if len(row) > numCols { + numCols = len(row) } } - var sb strings.Builder - - // Header line - var headerCells []string - for i, h := range t.headers { - cell := lipgloss.NewStyle(). - Width(colWidths[i] + 2). - Bold(true). - Foreground(SecondaryColor). - Render(h) - headerCells = append(headerCells, cell) + // Normalize headers + headers := make([]string, numCols) + for i := 0; i < numCols; i++ { + if i < len(t.headers) { + headers[i] = t.headers[i] + } else { + headers[i] = "" + } } - sb.WriteString(strings.Join(headerCells, "│")) - sb.WriteString("\n") - // Separator - var sepCells []string - for _, w := range colWidths { - sepCells = append(sepCells, strings.Repeat("─", w+2)) + // Normalize rows + normalizedRows := make([][]string, len(t.rows)) + for rIdx, row := range t.rows { + r := make([]string, numCols) + for cIdx := 0; cIdx < numCols; cIdx++ { + if cIdx < len(row) { + r[cIdx] = strings.ReplaceAll(row[cIdx], "\r", "") + } else { + r[cIdx] = "" + } + } + normalizedRows[rIdx] = r } - sb.WriteString(lipgloss.NewStyle().Foreground(BorderColor).Render(strings.Join(sepCells, "┼"))) - sb.WriteString("\n") - // Data rows - for rowIdx, row := range t.rows { - var rowCells []string - for i := 0; i < len(t.headers); i++ { - val := "" - if i < len(row) { - val = row[i] + tbl := table.New(). + Border(lipgloss.RoundedBorder()). + BorderStyle(lipgloss.NewStyle().Foreground(BorderColor)). + Headers(headers...). + Rows(normalizedRows...). + Wrap(false). + StyleFunc(func(row, col int) lipgloss.Style { + if row == table.HeaderRow { + return lipgloss.NewStyle(). + Bold(true). + Foreground(SecondaryColor). + Padding(0, 1) } - style := lipgloss.NewStyle().Width(colWidths[i] + 2) - if rowIdx%2 == 1 { - style = style.Foreground(lipgloss.Color("#E2E8F0")) + s := lipgloss.NewStyle().Padding(0, 1) + if row%2 == 1 { + s = s.Foreground(lipgloss.Color("#E2E8F0")) + } else { + s = s.Foreground(lipgloss.Color("#FFFFFF")) } - rowCells = append(rowCells, style.Render(val)) + return s + }) + + termWidth := getTerminalWidth() + if termWidth > 40 { + // Calculate natural width of the table + naturalWidth := 1 // left border + for cIdx := 0; cIdx < numCols; cIdx++ { + maxW := lipgloss.Width(headers[cIdx]) + for _, r := range normalizedRows { + w := lipgloss.Width(r[cIdx]) + if w > maxW { + maxW = w + } + } + naturalWidth += maxW + 2 + 1 // padding(2) + border(1) + } + + if naturalWidth > termWidth { + tbl.Width(termWidth) + t.truncated = true } - sb.WriteString(strings.Join(rowCells, "│")) - sb.WriteString("\n") } - return sb.String() + return tbl.Render() } diff --git a/internal/ui/table_test.go b/internal/ui/table_test.go new file mode 100644 index 0000000..5e8ddb7 --- /dev/null +++ b/internal/ui/table_test.go @@ -0,0 +1,32 @@ +package ui_test + +import ( + "testing" + + "github.com/sql-doctor/sql-doctor/internal/ui" +) + +func TestTableRenderNormal(t *testing.T) { + tbl := ui.NewTable("TABLE NAME", "TYPE", "EST. ROWS", "ENGINE") + tbl.AddRow("users", "BASE TABLE", "10", "InnoDB") + tbl.AddRow("orders", "BASE TABLE", "25", "InnoDB") + + out := tbl.Render() + if out == "" { + t.Fatalf("expected non-empty rendered table") + } +} + +func TestTableRenderWide(t *testing.T) { + tbl := ui.NewTable("id", "name", "email", "email_verified_at", "password", "remember_token", "created_at", "updated_at") + tbl.AddRow( + "1", "Test User", "test@example.com", "2026-08-02 05:23:56", + "$2y$12$3du298jbRmvwOjFIXSf4yesgjvorNji7pCLEgTe5l/yAtdv9xnCbq", "RCfBpIN8K2", + "2026-08-02 05:23:56", "2026-08-02 05:23:56", + ) + + out := tbl.Render() + if out == "" { + t.Fatalf("expected non-empty rendered table") + } +} diff --git a/tests/unit/database_selection_test.go b/tests/unit/database_selection_test.go new file mode 100644 index 0000000..60bbe6e --- /dev/null +++ b/tests/unit/database_selection_test.go @@ -0,0 +1,265 @@ +package unit + +import ( + "context" + "database/sql" + "os" + "path/filepath" + "strings" + "testing" + + "github.com/sql-doctor/sql-doctor/internal/cli" + "github.com/sql-doctor/sql-doctor/internal/database" + "github.com/sql-doctor/sql-doctor/internal/database/sqlite" + "github.com/sql-doctor/sql-doctor/internal/storage" +) + +func TestStorageSessionConnection(t *testing.T) { + tempDir, err := os.MkdirTemp("", "sql-doctor-test-*") + if err != nil { + t.Fatalf("failed to create temp dir: %v", err) + } + defer os.RemoveAll(tempDir) + + dbPath := filepath.Join(tempDir, "test.db") + store, err := storage.OpenStorage(dbPath) + if err != nil { + t.Fatalf("failed to open storage: %v", err) + } + defer store.Close() + + ctx := context.Background() + + // 1. Initial state - no session + sess, err := store.GetSessionConnection(ctx) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if sess != nil { + t.Fatalf("expected nil session, got %+v", sess) + } + + // 2. Save session connection without a database selected + rec := &storage.ConnectionRecord{ + Name: "default", + Dialect: database.DialectMySQL, + Host: "127.0.0.1", + Port: 3306, + User: "root", + Database: "", + IsActive: true, + } + if err := store.SaveSessionConnection(ctx, rec); err != nil { + t.Fatalf("failed to save session: %v", err) + } + + // 3. Retrieve session connection + sess, err = store.GetSessionConnection(ctx) + if err != nil { + t.Fatalf("failed to get session: %v", err) + } + if sess == nil || sess.Host != "127.0.0.1" || sess.Database != "" { + t.Fatalf("unexpected session record: %+v", sess) + } + + // 4. Update database on session connection + updated, err := store.UpdateConnectionDatabase(ctx, "", "ecommerce_prod") + if err != nil { + t.Fatalf("failed to update session database: %v", err) + } + if updated.Database != "ecommerce_prod" { + t.Errorf("expected updated database 'ecommerce_prod', got '%s'", updated.Database) + } + + // Verify persistence in session + sess, err = store.GetSessionConnection(ctx) + if err != nil || sess == nil || sess.Database != "ecommerce_prod" { + t.Fatalf("expected session database 'ecommerce_prod', got: %+v", sess) + } + + // 5. Clear session + if err := store.ClearSessionConnection(ctx); err != nil { + t.Fatalf("failed to clear session: %v", err) + } + sess, err = store.GetSessionConnection(ctx) + if err != nil || sess != nil { + t.Fatalf("expected cleared session, got: %+v", sess) + } +} + +func TestStorageUpdateNamedConnectionDatabase(t *testing.T) { + tempDir, err := os.MkdirTemp("", "sql-doctor-test-*") + if err != nil { + t.Fatalf("failed to create temp dir: %v", err) + } + defer os.RemoveAll(tempDir) + + dbPath := filepath.Join(tempDir, "test.db") + store, err := storage.OpenStorage(dbPath) + if err != nil { + t.Fatalf("failed to open storage: %v", err) + } + defer store.Close() + + ctx := context.Background() + + // Save a named connection profile without selecting a database + rec := &storage.ConnectionRecord{ + Name: "my-mysql", + Dialect: database.DialectMySQL, + Host: "localhost", + Port: 3306, + User: "admin", + Database: "", + IsActive: true, + } + if err := store.SaveConnection(ctx, rec); err != nil { + t.Fatalf("failed to save connection: %v", err) + } + + // Later select a database for this connection + updated, err := store.UpdateConnectionDatabase(ctx, "my-mysql", "analytics") + if err != nil { + t.Fatalf("failed to update connection database: %v", err) + } + if updated.Database != "analytics" { + t.Errorf("expected database 'analytics', got '%s'", updated.Database) + } + + // Check fetched connection has updated database + fetched, err := store.GetConnection(ctx, "my-mysql") + if err != nil { + t.Fatalf("failed to get connection: %v", err) + } + if fetched.Database != "analytics" { + t.Errorf("expected fetched database 'analytics', got '%s'", fetched.Database) + } +} + +func TestSQLiteDatabases(t *testing.T) { + ctx := context.Background() + drv := sqlite.New() + + cfg := &database.ConnectionConfig{ + Dialect: database.DialectSQLite, + FilePath: ":memory:", + } + db, err := drv.Connect(ctx, cfg) + if err != nil { + t.Fatalf("failed to connect sqlite: %v", err) + } + defer db.Close() + + dbs, err := drv.Databases(ctx, db) + if err != nil { + t.Fatalf("failed to list databases: %v", err) + } + if len(dbs) == 0 || dbs[0] != "main" { + t.Errorf("expected ['main'], got %v", dbs) + } +} + +// mockDriver implements database.Driver for testing EnsureDatabase +type mockDriver struct { + availableDBs []string +} + +func (m *mockDriver) Dialect() database.Dialect { return database.DialectMySQL } +func (m *mockDriver) DSN(cfg *database.ConnectionConfig) string { return "" } +func (m *mockDriver) Connect(ctx context.Context, cfg *database.ConnectionConfig) (*sql.DB, error) { + return nil, nil +} +func (m *mockDriver) Ping(ctx context.Context, db *sql.DB) error { return nil } +func (m *mockDriver) Version(ctx context.Context, db *sql.DB) (string, error) { + return "MySQL 8.0.36", nil +} +func (m *mockDriver) Databases(ctx context.Context, db *sql.DB) ([]string, error) { + return m.availableDBs, nil +} +func (m *mockDriver) Tables(ctx context.Context, db *sql.DB) ([]database.TableInfo, error) { + return nil, nil +} +func (m *mockDriver) DescribeTable(ctx context.Context, db *sql.DB, t string) (*database.TableDetail, error) { + return nil, nil +} +func (m *mockDriver) Indexes(ctx context.Context, db *sql.DB, t string) ([]database.IndexInfo, error) { + return nil, nil +} +func (m *mockDriver) ForeignKeys(ctx context.Context, db *sql.DB, t string) ([]database.ForeignKeyInfo, error) { + return nil, nil +} +func (m *mockDriver) Relationships(ctx context.Context, db *sql.DB) ([]database.RelationshipInfo, error) { + return nil, nil +} +func (m *mockDriver) Explain(ctx context.Context, db *sql.DB, q string, a bool) (*database.ExplainResult, error) { + return nil, nil +} +func (m *mockDriver) TableStats(ctx context.Context, db *sql.DB, t string) (*database.TableStats, error) { + return nil, nil +} +func (m *mockDriver) SampleColumnData(ctx context.Context, db *sql.DB, t, c string, l int) (*database.ColumnSampleStats, error) { + return nil, nil +} + +func TestEnsureDatabase(t *testing.T) { + ctx := context.Background() + + mock := &mockDriver{ + availableDBs: []string{"shop_db", "inventory", "test_db"}, + } + + // Case 1: Database is empty for MySQL -> should return error with suggestions + cfgNoDB := &database.ConnectionConfig{ + Name: "dev-server", + Dialect: database.DialectMySQL, + Database: "", + } + err := cli.EnsureDatabase(ctx, nil, mock, cfgNoDB) + if err == nil { + t.Fatalf("expected error when database is empty") + } + + errMsg := err.Error() + if !strings.Contains(errMsg, "no database selected for connection 'dev-server'") { + t.Errorf("expected error message to mention missing database, got: %s", errMsg) + } + if !strings.Contains(errMsg, "sql-doctor use ") { + t.Errorf("expected suggestion to run 'sql-doctor use ', got: %s", errMsg) + } + if !strings.Contains(errMsg, "shop_db") || !strings.Contains(errMsg, "inventory") { + t.Errorf("expected available databases listed in suggestion, got: %s", errMsg) + } + + // Case 2: Database is specified -> should succeed without error + cfgWithDB := &database.ConnectionConfig{ + Name: "dev-server", + Dialect: database.DialectMySQL, + Database: "shop_db", + } + if err := cli.EnsureDatabase(ctx, nil, mock, cfgWithDB); err != nil { + t.Errorf("expected nil error when database is present, got: %v", err) + } + + // Case 3: SQLite dialect -> should succeed even without database name + cfgSQLite := &database.ConnectionConfig{ + Name: "local-sqlite", + Dialect: database.DialectSQLite, + FilePath: "test.db", + } + if err := cli.EnsureDatabase(ctx, nil, mock, cfgSQLite); err != nil { + t.Errorf("expected nil error for sqlite, got: %v", err) + } + + // Case 4: Inside Shell -> should suggest 'use ' without sql-doctor prefix + shellErr := cli.EnsureShellDatabase(ctx, nil, mock, cfgNoDB) + if shellErr == nil { + t.Fatalf("expected error for shell without database") + } + shellErrMsg := shellErr.Error() + if !strings.Contains(shellErrMsg, " use \n") { + t.Errorf("expected shell suggestion to contain 'use ', got: %s", shellErrMsg) + } + if strings.Contains(shellErrMsg, "sql-doctor use") { + t.Errorf("shell suggestion should NOT contain 'sql-doctor use', got: %s", shellErrMsg) + } +}