diff --git a/.changeset/v2-2-1-providers-and-fixes.md b/.changeset/v2-2-1-providers-and-fixes.md new file mode 100644 index 00000000..b7883a9d --- /dev/null +++ b/.changeset/v2-2-1-providers-and-fixes.md @@ -0,0 +1,29 @@ +### New Features +- Connect TiDB Cloud, Turso, Nile and Upstash from the provider picker. TiDB shows its sign-in code right in the app +- Alt+wheel scrolls the table a few rows at a time, for landing on a precise row in a huge table +- Sign out of a provider from its card, with a confirm first +- Saved connections show their provider's logo and a small chip naming the engine +- Keyboard shortcuts in the connections dialog: new connection, paste a URL, filter and resume the last one +- The MCP dialog copies a ready config for every client +- A Claude font preset in Settings → Appearance + +### Bug Fixes +- A SQL run that ends in a SELECT keeps its earlier writes instead of rolling them back +- Picking a saved provider database connects to that database, never one with a similar name +- Passwords with a raw %, #, ? or @ import from a connection URL whole +- Stopping a query in the editor stops it on the server too +- Deleted connections stay deleted, and deleting one asks first +- Provider sign-in survives a restart on Windows +- PlanetScale lists databases across organizations and creates passwords again +- Upstash connections keep their provider after a restart +- One organization the token can't read no longer empties the whole database list +- The splash screen no longer flashes one colour and then another +- A chart on a windowed table uses only the rows that are loaded + +### Changes +- Provider databases open on the last list instead of a spinner, and reconnecting skips the API calls entirely +- Connecting to far away databases is faster: no restarted handshakes, no double connect, one Redis connection reused +- Geist is the default font again +- The connect page uses the full width with four cards per row, and the New table dialog reads as one table +- An empty sidebar tab shows one empty state instead of two +- Failed database picks show as a toast and leave the picker open diff --git a/DESIGN_SYSTEM.md b/DESIGN_SYSTEM.md index a0060ae5..72afdc02 100644 --- a/DESIGN_SYSTEM.md +++ b/DESIGN_SYSTEM.md @@ -36,11 +36,14 @@ mismatched control heights, or one-off selected states. ### Fonts (`src/app.css`) | Role | Variable | Stack | |---|---|---| -| UI / sans | `--font-sans` | Geist Variable → system-ui | -| Data / mono | `--font-mono` | Geist Mono Variable → ui-monospace | +| UI / sans | `--font-sans` | Geist Variable (default) → system-ui | +| Headings | `--heading-font` (utility `font-heading`) | Source Serif 4 Variable (Claude preset), else follows `--font-sans` | +| Data / mono | `--font-mono` | Geist Mono Variable (default) → ui-monospace | | Code editor | `--editor-font-family` | JetBrains Mono → Geist Mono | - **Sans** for all chrome, labels, prose, buttons. +- **Heading** (`font-heading`) for dialog titles, via `Dialog.Title`. Nothing + smaller than a title uses it. - **Mono** (`font-mono`) for data: identifiers, values, counts, IDs, SQL, hostnames, table/column names, timings. Numbers in stat cards use `font-mono tabular-nums`. diff --git a/README.md b/README.md index 1c7e268c..e0f7fd0e 100644 --- a/README.md +++ b/README.md @@ -2,443 +2,96 @@ # Stroke -**A fast, minimal desktop database client.** +**A fast desktop database client, for people and for their AI tools.** -Eleven engines and four one-click providers, in one window. Browse and edit data, write SQL, draw -diagrams and maps of what you find, back it up, and let your AI tools query it through a built-in -MCP server. - -[Download](#install) · [Features](#features) · [Build from source](#build-from-source) · [Website](https://stroke.click) +[Download](#install) · [What's inside](#whats-inside) · [Build from source](#build-from-source) · [stroke.click](https://stroke.click) ---- - -## Why Stroke - -Built with **Rust + Svelte** — a native backend paired with a reactive UI — so the app is snappy, memory-efficient, and never feels like a browser tab. - -- **Fast, intuitive, easy** -- **Handles millions of rows** without slowing down -- **Keyboard-first** — every feature is reachable without a mouse -- **Read-only mode** — safely browse production without risking accidental writes -- **AI-native** — a built-in assistant and an MCP server your other AI tools can connect to -- **Extensions** — expand functionality with plugins -- **Themes** — dark and light, out of the box - ---- - -## Supported databases - -**Core engines** - -| Database | Notes | -|----------|-------| -| **PostgreSQL** | Full schema, enums, sequences, triggers, indexes, extensions | -| **MySQL** | Standard host/port connections | -| **MariaDB** | Its own driver, not a MySQL alias | -| **CockroachDB** | Distributed SQL, Postgres wire protocol | -| **SQLite** | Local file or `:memory:` | -| **Turso / LibSQL** | Serverless SQLite at the edge | -| **Cloudflare D1** | OAuth sign-in or API token; local `wrangler` databases are found automatically | -| **ClickHouse** | Columnar OLAP over HTTP(S) | -| **DuckDB** | Embedded analytical database (local file) | -| **Microsoft SQL Server** | Host/port connections | -| **Redis** | Keyspace browser + command console (key-value; SQL surfaces hidden) | - -**One-click providers** - -Sign in and pick a database — no connection string to assemble: - -| Provider | Based on | -|----------|----------| -| **Neon** | Serverless Postgres | -| **Supabase** | Postgres | -| **Prisma Postgres** | Serverless Postgres | -| **PlanetScale** | MySQL-compatible serverless | - -Other PostgreSQL-compatible databases connect through the PostgreSQL option. - -Connect directly or through an **SSH tunnel** for databases behind a bastion host or private network. -Stroke also **finds databases already running on your machine** — Docker containers, local -Postgres/MySQL instances, and the SQLite files your project's ORM config points at — so a local -connection is usually one click, not a form. - ---- - -## Features - -### Data grid - -- Paginated browsing with configurable page size -- Rows paint immediately; the total count (`… of N`) fills in the background -- Resizable columns, saved per table -- Pin columns to keep them in view while scrolling -- Show/hide columns and reset column width to default -- Column statistics — min, max, avg, nulls, distinct count -- Multi-column sort — shift-click headers to add secondary keys -- Full-text search and a visual filter builder, with relative date-range presets and enum value pickers -- Click any foreign-key value to jump to the referenced row -- Choose from six grid styles (Lines, Bordered, Striped, Dotted, Dots, Minimal) in **Settings → Appearance** - -**Cell right-click menu** - -Open · Edit · Copy · Filter by value · Exclude this value · Set NULL · Expand · Select row · Delete row -Copy row as → JSON · CSV · Plain text · Markdown table · INSERT statement - -**Column header right-click menu** - -Sort ascending · Sort descending · Filter by this column · Pin column · Hide column · Reset column width · Column stats - -### Inline editing - -Edit text, numbers, booleans, enums, dates, UUIDs, and JSON in place. -Insert rows with smart defaults, delete rows, and set any cell to NULL in one action. -Postgres array columns (`text[]`, `int[]`, …) get a dedicated add / remove / reorder editor. -Optionally preview the generated SQL for any change before it's applied. - -### Row inspector - -View any row in formatted, raw JSON, or preview mode — with a rich viewer for nested JSON, arrays, and media/URL values. - -### Schema explorer - -- Tables, views, materialized views, and foreign tables -- Live row counts in the sidebar -- Index browser and DDL viewer for any object -- Create tables, schemas, sequences, triggers, enums, and foreign keys through guided dialogs -- Manage PostgreSQL enum types - -### Relation tree - -Walk foreign-key relationships as an expandable tree, drilling from a row into everything it references — row counts stream in without blocking navigation. - -### Global search - -Search across tables and objects to find what you need without hunting through the sidebar. - -### Live mode - -Watch a table for real-time changes — inserts, updates, and deletes stream into the grid as they happen on the database. - -### SQL console - -- Syntax highlighting and schema-aware autocomplete -- Query formatting and execution timing -- **EXPLAIN plans** — visualize the query planner's execution plan -- CSV, JSON, and SQL export of results -- Automatic query history — everything is saved -- Saved queries — bookmark the ones you keep -- SQL Notebooks (`.sqlnb`) — multi-cell notebooks for documenting and replaying query sequences - -### Quick access - -Open the quick-access screen to jump anywhere in one keystroke: - -| | | -|--|--| -| **SQL** | Console with history and saved queries | -| **Dashboard** | Pinned charts and saved views | -| **AI** | AI chat assistant | -| **ORM** | ORM query runner | -| **Schema** | Schema and index explorer | -| **Security** | Roles, grants, and privilege viewer | -| **Logs** | Query and audit logs | -| **Charts** | Turn query results into charts | -| **Diagrams** | Auto-generated ERDs | -| **Timeline** | Schema drift detection across snapshots | -| **Data Diff** | Row-level diff between snapshots or queries | -| **Search** | Global search across the database | -| **Extensions** | Install and manage plugins | -| **Connect** | Manage and switch connections | - -### Charts & dashboards - -Turn any query result into a chart — bar, line, area, pie, and geographic **choropleth** maps (powered by ECharts) — and pin it to a dashboard. - -### Schema diagrams - -Auto-generated entity-relationship diagrams (ERDs) built from your live schema, with an interactive canvas. - -### Schema timeline - -Track schema changes over time. See exactly which columns, indexes, or constraints were added or removed between snapshots. - -### Data diff - -Compare two query results or table snapshots row-by-row — added, removed, and modified rows are highlighted clearly. - -### ORM runner - -Write ORM-style queries, preview the generated SQL, and run them against the connected database. - -### Security viewer - -Inspect roles, grants, and privileges — see who can read or write what, at a glance. - -### Backup & restore - -Export a full **SQL dump** — schema DDL, data, views, triggers, functions, sequences, and enums — with granular include toggles and per-table selection. A live log streams progress, and any export can be stopped mid-run. Restore by executing a `.sql` file, with per-statement results and clear error reporting. - -Available for **PostgreSQL, MySQL, SQLite, and Cloudflare D1**. - -### Docker launch - -Spin up a local PostgreSQL or MySQL container in one click without leaving the app. - -### Extensions - -Install community and first-party extensions to add new panels, query tools, or integrations. - -### Cell viewers - -Values that don't fit a grid cell get a real viewer instead of being truncated into nonsense: - -- **JSON / JSONB** — a collapsible tree, with search and a full-screen editor -- **Vectors** (`pgvector`) — dimension, norm, mean, standard deviation, a per-dimension strip - and a value histogram, so you can see whether an embedding is shaped the way you expect -- **Geometry** (PostGIS) — the shape drawn on a pannable, zoomable map, with type, SRID and - vertex list. Offline by default; tiled basemaps are one click away -- **Arrays**, long text, and oversized values (multi-MB cells are capped before they reach the - UI, so one big blob can't freeze the window) - -### Views of the same rows - -Every table tab can render its rows as a **grid**, **JSON**, a **record card**, plain **text**, -a **chart**, or an **ER diagram** — switchable per tab, with a default you can set globally. +Stroke is a Rust + Svelte app for browsing, editing and querying databases. It stays quick on tables with millions of rows, works from the keyboard, and ships an MCP server so Claude, Cursor and other agents can query the same databases you do. -### Map view +## Databases -Spatial columns drawn on a map: every PostGIS layer in the database, clustered when there are -too many features to draw individually, filterable with the same operators as the grid. The -basemap ships with the app, so the default view makes no network requests. +**Engines:** PostgreSQL, MySQL, MariaDB, CockroachDB, SQLite, Turso / libSQL, Cloudflare D1, ClickHouse, DuckDB, SQL Server and Redis. -### Instance insights +**Sign in and pick a database**, no connection string needed: -What the server itself is doing — version, uptime, connections, replication state, cache hit -rates, and the configuration values that matter. Settings that are safe to change can be edited -in place. +| Provider | Engine | +|---|---| +| Neon, Supabase, Prisma Postgres, Nile | Postgres | +| PlanetScale, TiDB Cloud | MySQL | +| Turso | libSQL | +| Cloudflare D1 | SQLite | +| Upstash | Redis | -### Database objects +Railway is next. -Every enum, sequence, trigger, function and index in one browsable list, rather than scattered -across the schema tree. +Stroke also finds databases already running on your machine (Docker containers, local Postgres and MySQL, the SQLite file your ORM points at), so a local connection is usually one click. Anything can go through an SSH tunnel. -### Notebooks +## What's inside -Interleave SQL cells and Markdown in a `.sqlnb` file — run cells independently, keep the results -with the prose. Useful for an analysis you want to hand to someone else. +- **Data grid** that paints millions of rows, with filters, multi-column sort, pinned columns, foreign-key jumps and inline editing, including Postgres arrays and JSON +- **SQL console** with schema-aware autocomplete, EXPLAIN plans, history, saved queries and `.sqlnb` notebooks +- **Viewers** for JSON, pgvector embeddings and PostGIS geometry, plus a map view of every spatial layer +- **Schema tools:** ER diagrams, a relation tree, schema timeline, data diff, and Prisma or Drizzle codegen +- **Charts and dashboards** from any query result +- **Backup and restore** as SQL dumps for Postgres, MySQL, SQLite and D1 +- **Read-only mode** for browsing production without the risk +- **AI chat** that runs queries and draws charts, with a free tier, your own key, local models (Ollama, LM Studio) or GitHub Copilot. Destructive statements always ask first +- **MCP server:** start it in Settings, copy the config for your client, done +- **Command palette** on `Cmd/Ctrl+K`, split panes, themes and extensions -### Codegen - -Read the live schema back out as **Prisma** or **Drizzle** source. Introspects once and re-renders -locally, so switching between the two is instant even on a large schema. - -### Split panes - -Drag a tab to either edge to split the window. Compare two tables, or keep a query beside the -rows it returns. - -### Activity log - -Every statement the app has run, with duration and outcome — including the ones it ran on your -behalf, so nothing the UI does is invisible. - -### Read-only mode - -Lock any connection so writes are blocked entirely — safe for browsing production. - -### AI chat - -An assistant with real database access: it runs queries, reads schemas, explains what it found, -and renders charts and diagrams inline. Destructive statements always ask first. - -- **Free tier built in** — a shared daily allowance, no key required, nothing to configure -- **Your own key** — any OpenAI-compatible provider (OpenAI, Anthropic, Google, OpenRouter, …) -- **Local models** — Ollama and LM Studio, including Ollama Cloud models that run on their - hardware with nothing to download. The picker lists what your server actually has, so it can - never suggest a model you haven't installed -- **GitHub Copilot** — sign in with your existing subscription -- **OmniRoute** — install and start the local gateway from inside the app -- **Skills** — Markdown files that shape how the agent works on your schema -- **Web search** — off by default; when on, the agent can look up an error code or a function's - syntax that your database can't answer - -### MCP server - -Stroke ships a built-in **Model Context Protocol** server so Claude, Cursor, and other MCP clients can query your database directly. - -1. Open **Settings → MCP Server** -2. Click **Start** -3. Copy the one-click config for your AI tool - -Uses a stable bearer token — configure once and it just works. - -### Command palette - -`Cmd/Ctrl+K` — type anything (table name, feature, shortcut) to jump there instantly. - -### Status bar - -Active connection, database name, row counts, query execution time, and connection health — always visible at the bottom. - -### Security & storage - -AI keys and provider OAuth tokens are stored in the **OS keychain** (macOS Keychain, Windows Credential Manager, or the Linux Secret Service), and migrated automatically from any older plaintext store. - ---- - -## License - -Stroke is **source-available** under the [Stroke Sustainable Use License](LICENSE). -The entire source — Pro features included — lives in this repository: - -- **Free to use** — personally, in your company, anywhere -- **Read, modify, and contribute** to all of it -- **Not for resale** — you can't sell Stroke, rebrand it, or offer it as a paid product or hosted service -- **Pro features** are key-gated in official builds — after the built-in trial they - require a [Stroke Pro](https://stroke.click/pricing) license, which is what funds development - -Official builds and updates ship from this repository and [stroke.click](https://stroke.click). - ---- +Keys and provider tokens live in the OS keychain. ## Install -### Recommended: package managers - -**macOS** — [Homebrew](https://brew.sh) +**macOS** ([Homebrew](https://brew.sh)) ```bash brew install --cask stroke-app/tap/stroke ``` -**Windows** — [Scoop](https://scoop.sh) +**Windows** ([Scoop](https://scoop.sh)) ```powershell scoop bucket add stroke https://github.com/stroke-app/stroke scoop install stroke ``` -### Or download directly - -Grab the installer from the [Releases](https://github.com/stroke-app/stroke/releases) page. - -| Platform | File | -|----------|------| -| macOS (Apple Silicon) | `stroke_x.x.x_aarch64.dmg` | -| macOS (Intel) | `stroke_x.x.x_x64.dmg` | -| Windows | `stroke_x.x.x_x64-setup.exe` | -| Linux (Debian/Ubuntu) | `stroke_x.x.x_amd64.deb` | -| Linux (AppImage) | `stroke_x.x.x_amd64.AppImage` | - -**macOS** — open the `.dmg`, drag Stroke to Applications. If macOS blocks it on first launch: - -```bash -xattr -cr /Applications/Stroke.app -``` - -**Windows** — run the `.exe`. If SmartScreen shows a warning, click **More info → Run anyway**. (Scoop skips this.) - -**Linux (Debian/Ubuntu)** - -```bash -sudo dpkg -i stroke_*_amd64.deb -``` - -**Linux (AppImage)** - -```bash -chmod +x stroke_*_amd64.AppImage -./stroke_*_amd64.AppImage -``` - ---- - -## Connecting - -Pick a database type and fill in your credentials. Hit **Test connection**, then **Connect**. - -- **Discovered automatically** — Docker containers, local Postgres/MySQL, and the SQLite or D1 database your project's ORM config points at. Pick it from the list; no form. -- **PostgreSQL / MySQL / MariaDB / CockroachDB** — host, port, database, user, password, optional SSL. Paste a full connection string and click **Parse** to fill the form automatically. -- **SQLite / DuckDB** — point to a database file (`.db`, `.sqlite`, `.duckdb`), or use `:memory:` for SQLite. -- **Turso / LibSQL** — database URL and optional auth token. -- **Cloudflare D1** — sign in with Cloudflare, or enter Account ID, Database ID, and an API token. -- **ClickHouse** — host, port, database, user, password, optional TLS. -- **SQL Server** — host, port, database, user, password. -- **Redis** — host, port, optional password, database number, optional TLS. -- **Neon / Supabase / Prisma / PlanetScale** — sign in with the provider and pick a database. -- **SSH tunnel** — any connection type can go through an SSH tunnel for databases in private networks. +Or grab an installer from [Releases](https://github.com/stroke-app/stroke/releases): `.dmg` for macOS (Apple Silicon or Intel), `-setup.exe` for Windows, `.deb` or `.AppImage` for Linux. -Connections are saved locally. Stroke reopens your last connection on launch. +If macOS blocks the first launch, run `xattr -cr /Applications/Stroke.app`. If Windows SmartScreen warns, choose **More info → Run anyway**. ---- +## Shortcuts -## Keyboard shortcuts - -| Shortcut | Action | -|----------|--------| +| Keys | Action | +|---|---| | `Cmd/Ctrl+K` | Command palette | -| `Cmd/Ctrl+N` | Quick access | +| `Cmd/Ctrl+Enter` | Run the query | | `Cmd/Ctrl+T` | Search tables | -| `Cmd/Ctrl+Shift+D` | Table data view | -| `Cmd/Ctrl+Shift+S` | SQL console | -| `Cmd/Ctrl+Shift+E` | AI chat | -| `Cmd/Ctrl+Enter` | Run SQL query | -| `Cmd/Ctrl+Tab` | Switch tabs | -| `Alt+←/→` | Back / forward | -| `Cmd/Ctrl+B` | Toggle sidebar | -| `Cmd/Ctrl+W` | Close tab | -| `Cmd/Ctrl+?` | Show all shortcuts | - ---- +| `Cmd/Ctrl+B` | Toggle the sidebar | +| `Cmd/Ctrl+W` | Close the tab | +| `Alt+wheel` | Scroll the grid a few rows at a time | +| `Cmd/Ctrl+?` | Every shortcut | ## Build from source -This repository is the complete source. Requires [Node.js](https://nodejs.org) 20.19+ -(or 22.12+; CI builds on 22) and the [Rust toolchain](https://rustup.rs). +Needs [Node.js](https://nodejs.org) 20.19+ (CI uses 22) and the [Rust toolchain](https://rustup.rs). ```bash git clone https://github.com/stroke-app/stroke cd stroke npm install -npm run tauri # dev -npm run tauri:build # release binary -``` - -**Arch Linux** — build `.deb` only to avoid an AppImage linker issue: - -```bash -npm run tauri:build -- --bundles deb +npm run tauri # dev +npm run tauri:build # release build ``` -Or use the included Arch helper: +On Arch, `npm run tauri:build:arch` avoids an AppImage linker issue. `npm run tauri:build:mac` and `npm run tauri:build:win` build release installers locally; `scripts/build-local.mjs` lists what each one needs. -```bash -npm run tauri:build:arch -``` - -**Release installers on your own machine** — no CI needed: - -```bash -npm run tauri:build:mac # macOS .app + .dmg (on a Mac) -npm run tauri:build:win # Windows x64 setup .exe - native on Windows, - # cross-compiled from macOS/Linux -npm run tauri:build:local -- mac-x64 # also: host (default), linux -``` - -The Windows cross-build needs NSIS, LLVM and `cargo-xwin` (macOS: -`brew install nsis llvm && cargo install --locked cargo-xwin`); the script checks -for each and prints what is missing. macOS and Linux bundles only build on their -own OS. Details and the DuckDB workaround are in `scripts/build-local.mjs`. +## License ---- +Source-available under the [Stroke Sustainable Use License](LICENSE). Use it anywhere, personally or at work, and read, modify and contribute to all of it. Reselling or rebranding it, or offering it as a hosted service, is not allowed. Pro features are key-gated in official builds and a [Stroke Pro](https://stroke.click/pricing) license funds development. ## Contributing -Bug reports, fixes, features, and docs are all welcome — see -[CONTRIBUTING.md](CONTRIBUTING.md) for the dev setup (including one-command -Docker test databases), code conventions, and the PR checklist. - -Found a bug or have a feature request? Open an issue at -[github.com/stroke-app/stroke/issues](https://github.com/stroke-app/stroke/issues). +Fixes, features and docs are welcome. [CONTRIBUTING.md](CONTRIBUTING.md) covers the dev setup, including one-command Docker test databases, and the PR checklist. Bugs and ideas go in [issues](https://github.com/stroke-app/stroke/issues). diff --git a/index.html b/index.html index 02489bec..c0d930d8 100644 --- a/index.html +++ b/index.html @@ -37,7 +37,15 @@ // These two are the base light/dark --background, and they match // LIGHT_SURFACE / DARK_SURFACE in src-tauri/src/lib.rs so the window, // the webview and the page all come up the same colour. - document.documentElement.style.backgroundColor = isLight ? '#f7f7f7' : '#080808' + // The theme's own background when a previous run recorded it (see + // rememberBootSurface in settings.js); the base colours only on a first + // launch or a theme never shown before. + let surface = isLight ? '#f7f7f7' : '#080808' + try { + const boot = JSON.parse(localStorage.getItem('stroke:boot-surface') || 'null') + if (boot && boot.theme === theme && typeof boot.color === 'string') surface = boot.color + } catch (_) {} + document.documentElement.style.backgroundColor = surface let zoom = 1 if (parsed?.zoom != null) zoom = Number(parsed.zoom) else if (parsed?.fontSize != null) zoom = Number(parsed.fontSize) / 14 diff --git a/package-lock.json b/package-lock.json index 1004f990..4ed03e1c 100644 --- a/package-lock.json +++ b/package-lock.json @@ -1,12 +1,12 @@ { "name": "stroke", - "version": "2.0.0", + "version": "2.2.0", "lockfileVersion": 3, "requires": true, "packages": { "": { "name": "stroke", - "version": "2.0.0", + "version": "2.2.0", "license": "SEE LICENSE IN LICENSE", "dependencies": { "@codemirror/commands": "^6.11.1", @@ -23,6 +23,7 @@ "@fontsource-variable/inter": "^5.2.8", "@fontsource-variable/jetbrains-mono": "^5.2.8", "@fontsource-variable/source-code-pro": "^5.2.7", + "@fontsource-variable/source-serif-4": "^5.3.0", "@fontsource-variable/space-grotesk": "^5.2.10", "@fontsource/ibm-plex-mono": "^5.2.7", "@fontsource/ibm-plex-sans": "^5.2.8", @@ -372,6 +373,15 @@ "url": "https://github.com/sponsors/ayuhito" } }, + "node_modules/@fontsource-variable/source-serif-4": { + "version": "5.3.0", + "resolved": "https://registry.npmjs.org/@fontsource-variable/source-serif-4/-/source-serif-4-5.3.0.tgz", + "integrity": "sha512-9vch9WqxjaaA+1o9Ur8pOgIGbCYLjRReOUel23A6lOpD1syptgjtkORevvNmldJ5kGXQL29onQUqI5Ltz0s3bQ==", + "license": "OFL-1.1", + "funding": { + "url": "https://github.com/sponsors/ayuhito" + } + }, "node_modules/@fontsource-variable/space-grotesk": { "version": "5.2.10", "resolved": "https://registry.npmjs.org/@fontsource-variable/space-grotesk/-/space-grotesk-5.2.10.tgz", diff --git a/package.json b/package.json index d19f3236..d3e7d2e2 100644 --- a/package.json +++ b/package.json @@ -60,6 +60,7 @@ "@fontsource-variable/inter": "^5.2.8", "@fontsource-variable/jetbrains-mono": "^5.2.8", "@fontsource-variable/source-code-pro": "^5.2.7", + "@fontsource-variable/source-serif-4": "^5.3.0", "@fontsource-variable/space-grotesk": "^5.2.10", "@fontsource/ibm-plex-mono": "^5.2.7", "@fontsource/ibm-plex-sans": "^5.2.8", diff --git a/src-tauri/src/app_lock.rs b/src-tauri/src/app_lock.rs index bcc1b95f..bc1a7898 100644 --- a/src-tauri/src/app_lock.rs +++ b/src-tauri/src/app_lock.rs @@ -126,30 +126,31 @@ where .map_err(|e| e.to_string()) } -fn read_config(app: &tauri::AppHandle) -> LockConfig { - crate::secrets::read_all(app) - .get(VAULT_KEY) +fn config_from(map: &std::collections::HashMap) -> LockConfig { + map.get(VAULT_KEY) .and_then(|json| serde_json::from_str::(json).ok()) .unwrap_or_default() } -fn write_config(app: &tauri::AppHandle, cfg: &LockConfig) -> Result<(), String> { - let mut map = crate::secrets::read_all(app); - let json = serde_json::to_string(cfg).map_err(|e| e.to_string())?; - map.insert(VAULT_KEY.to_string(), json); - crate::secrets::write_all(app, &map) +fn read_config(app: &tauri::AppHandle) -> LockConfig { + config_from(&crate::secrets::read_all(app)) } -/// Read, mutate, write in a single hop so the pair shares one keychain unlock. +/// Read, mutate, write as one locked vault update, so the pair shares one +/// keychain unlock and a concurrent secret write can't be lost. A rejected +/// change (`f` errs) leaves the stored config as it was. async fn edit(app: tauri::AppHandle, f: F) -> Result where F: FnOnce(&mut LockConfig) -> Result<(), String> + Send + 'static, { off_thread(move || { - let mut cfg = read_config(&app); - f(&mut cfg)?; - write_config(&app, &cfg)?; - Ok(LockStatus::from(&cfg)) + crate::secrets::update(&app, move |map| { + let mut cfg = config_from(map); + f(&mut cfg)?; + let json = serde_json::to_string(&cfg).map_err(|e| e.to_string())?; + map.insert(VAULT_KEY.to_string(), json); + Ok(LockStatus::from(&cfg)) + })? }) .await? } diff --git a/src-tauri/src/cloudflare.rs b/src-tauri/src/cloudflare.rs index b6cf2cc2..e867119a 100644 --- a/src-tauri/src/cloudflare.rs +++ b/src-tauri/src/cloudflare.rs @@ -52,6 +52,10 @@ fn http() -> &'static reqwest::Client { .user_agent("stroke/1.0") .tcp_keepalive(std::time::Duration::from_secs(60)) .pool_max_idle_per_host(4) + // Idle sockets go before an upstream load balancer's 60s cutoff, so a + // request after a pause doesn't go out on a connection already closed + // (same fix as the provider client in providers/mod.rs). + .pool_idle_timeout(std::time::Duration::from_secs(20)) // Bounded on purpose. `reqwest` has no default timeout, so a request // that never answers - captive portal, dropped route, a stalled edge - // leaves the command awaiting forever and the UI on its spinner with no @@ -125,27 +129,8 @@ async fn await_oauth_callback( listener: TcpListener, expected_state: &str, ) -> Result { - let success_html = r#" - -Stroke - authorized - -
-

Authorization successful

-

You can close this tab and return to Stroke.

-
"#; - - let error_html = r#" - -Stroke - error - -
-

Authorization failed

-

You can close this tab and try again in Stroke.

-
"#; + let success_html = crate::oauth_page::page(true, "Cloudflare"); + let error_html = crate::oauth_page::page(false, "Cloudflare"); let send_html = |html: &str| -> String { format!( @@ -198,7 +183,7 @@ h2{color:#ef4444;margin-bottom:12px}p{color:#888;margin:0} if let Some(err) = &error { let _ = stream - .write_all(send_html(error_html).as_bytes()) + .write_all(send_html(&error_html).as_bytes()) .await; return Err(format!("Cloudflare denied authorization: {err}")); } @@ -207,7 +192,7 @@ h2{color:#ef4444;margin-bottom:12px}p{color:#888;margin:0} Some(c) if !c.is_empty() => c, _ => { let _ = stream - .write_all(send_html(error_html).as_bytes()) + .write_all(send_html(&error_html).as_bytes()) .await; return Err("No authorization code in callback".to_string()); } @@ -215,13 +200,13 @@ h2{color:#ef4444;margin-bottom:12px}p{color:#888;margin:0} if state.as_deref() != Some(expected_state) { let _ = stream - .write_all(send_html(error_html).as_bytes()) + .write_all(send_html(&error_html).as_bytes()) .await; return Err("OAuth state mismatch - possible CSRF".to_string()); } let _ = stream - .write_all(send_html(success_html).as_bytes()) + .write_all(send_html(&success_html).as_bytes()) .await; let _ = stream.flush().await; @@ -347,33 +332,37 @@ async fn store_tokens( expires_in: Option, email: Option<&str>, ) -> Result<(), String> { - let mut map = crate::secrets::read_all_async(app).await; - map.insert(KEY_ACCESS.to_string(), access.to_string()); - if let Some(r) = refresh { - map.insert(KEY_REFRESH.to_string(), r.to_string()); - } - if let Some(exp) = expires_in { - let expires_at = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_secs() - + exp - - 30; // 30s buffer - map.insert(KEY_EXPIRES.to_string(), expires_at.to_string()); - } - if let Some(e) = email { - map.insert(KEY_EMAIL.to_string(), e.to_string()); - } - crate::secrets::write_all_async(app, map).await + let access = access.to_string(); + let refresh = refresh.map(str::to_string); + let email = email.map(str::to_string); + crate::secrets::update_async(app, move |map| { + map.insert(KEY_ACCESS.to_string(), access); + if let Some(r) = refresh { + map.insert(KEY_REFRESH.to_string(), r); + } + if let Some(exp) = expires_in { + let expires_at = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_secs() + + exp.saturating_sub(30); // 30s buffer + map.insert(KEY_EXPIRES.to_string(), expires_at.to_string()); + } + if let Some(e) = email { + map.insert(KEY_EMAIL.to_string(), e); + } + }) + .await } async fn clear_tokens(app: &tauri::AppHandle) -> Result<(), String> { - let mut map = crate::secrets::read_all_async(app).await; - map.remove(KEY_ACCESS); - map.remove(KEY_REFRESH); - map.remove(KEY_EXPIRES); - map.remove(KEY_EMAIL); - crate::secrets::write_all_async(app, map).await + crate::secrets::update_async(app, |map| { + map.remove(KEY_ACCESS); + map.remove(KEY_REFRESH); + map.remove(KEY_EXPIRES); + map.remove(KEY_EMAIL); + }) + .await } fn now_secs() -> u64 { @@ -471,13 +460,25 @@ pub fn set_app_handle(app: tauri::AppHandle) { /// Returns None rather than an error because the caller's job is to report the /// *original* failure when no refresh is possible - "session expired" would be /// a misleading thing to show someone using a manual API token. +// One refresh at a time. Cloudflare rotates refresh tokens, so two concurrent +// refreshes (the D1 driver's 401 recovery and a panel's token check) with the +// same one left the loser rejected and the session looking signed out. +static REFRESH_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(()); + pub async fn refreshed_token() -> Option { let app = APP.get()?.clone(); + let before = crate::secrets::read_all_async(&app).await.get(KEY_REFRESH).cloned(); + let _guard = REFRESH_LOCK.lock().await; let map = crate::secrets::read_all_async(&app).await; let refresh = map.get(KEY_REFRESH).cloned().unwrap_or_default(); if refresh.is_empty() { return None; } + // Refreshed by someone else while we waited: that token is the fresh one, + // and reusing the rotated-out refresh token would be rejected. + if before.as_deref() != Some(refresh.as_str()) { + return map.get(KEY_ACCESS).cloned(); + } let new_token = refresh_access_token(&refresh).await.ok()?; let email = map.get(KEY_EMAIL).cloned(); store_tokens( @@ -511,6 +512,19 @@ pub async fn cloudflare_get_valid_token(app: tauri::AppHandle) -> Result().ok()) + .unwrap_or(u64::MAX); + if let Some(access) = map.get(KEY_ACCESS).filter(|a| !a.is_empty()) { + if now_secs() < expires_at { + return Ok(access.clone()); + } + } + let refresh = map.get(KEY_REFRESH).cloned().unwrap_or_default(); if refresh.is_empty() { return Err( diff --git a/src-tauri/src/commands.rs b/src-tauri/src/commands.rs index 1d50b54f..9a932b5b 100644 --- a/src-tauri/src/commands.rs +++ b/src-tauri/src/commands.rs @@ -1123,6 +1123,21 @@ pub async fn cancel_query( Ok(()) } +// ── Saved connections ───────────────────────────────────────────────────────── + +/// The durable saved-connections payload, or `null` before the first write. +#[tauri::command(async)] +pub fn connections_store_read(app: tauri::AppHandle) -> Result, String> { + let dir = app.path().app_data_dir().map_err(|e| e.to_string())?; + crate::connection_store::read(&dir) +} + +#[tauri::command(async)] +pub fn connections_store_write(app: tauri::AppHandle, json: String) -> Result<(), String> { + let dir = app.path().app_data_dir().map_err(|e| e.to_string())?; + crate::connection_store::write(&dir, &json) +} + // ── License ─────────────────────────────────────────────────────────────────── // `async` so the first call (device-fingerprint subprocess + trial-file I/O) diff --git a/src-tauri/src/connection_store.rs b/src-tauri/src/connection_store.rs new file mode 100644 index 00000000..51d44c99 --- /dev/null +++ b/src-tauri/src/connection_store.rs @@ -0,0 +1,53 @@ +//! Durable copy of the saved-connections list. +//! +//! The frontend keeps connections in `localStorage`, but the webview flushes +//! that to disk on its own schedule. WebView2 on Windows in particular can drop +//! the last writes when the app closes soon after them, so a deleted connection +//! came back on the next launch (and a freshly added one could vanish). This +//! file is written synchronously and fsynced on every change, and the frontend +//! loads it before anything reads the list. + +use std::path::{Path, PathBuf}; + +const FILE_NAME: &str = "connections.json"; + +fn store_path(data_dir: &Path) -> PathBuf { + data_dir.join(FILE_NAME) +} + +/// The stored payload, or `None` before the first write. +pub fn read(data_dir: &Path) -> Result, String> { + match std::fs::read_to_string(store_path(data_dir)) { + Ok(s) => Ok(Some(s)), + Err(e) if e.kind() == std::io::ErrorKind::NotFound => Ok(None), + Err(e) => Err(format!("Could not read saved connections: {e}")), + } +} + +/// Replace the stored payload. Written to a sibling temp file, fsynced, then +/// renamed over the old one, so a crash mid-write leaves the previous list +/// intact instead of a truncated file. +pub fn write(data_dir: &Path, json: &str) -> Result<(), String> { + use std::io::Write; + + serde_json::from_str::(json) + .map_err(|e| format!("Refusing to save connections: payload is not JSON ({e})"))?; + std::fs::create_dir_all(data_dir).map_err(|e| e.to_string())?; + + let path = store_path(data_dir); + let tmp = path.with_extension("json.tmp"); + let mut opts = std::fs::OpenOptions::new(); + opts.write(true).create(true).truncate(true); + // Saved connections include passwords: keep the file private to this user. + #[cfg(unix)] + { + use std::os::unix::fs::OpenOptionsExt; + opts.mode(0o600); + } + let mut f = opts.open(&tmp).map_err(|e| format!("Could not save connections: {e}"))?; + f.write_all(json.as_bytes()) + .and_then(|_| f.sync_all()) + .map_err(|e| format!("Could not save connections: {e}"))?; + drop(f); + std::fs::rename(&tmp, &path).map_err(|e| format!("Could not save connections: {e}")) +} diff --git a/src-tauri/src/db/connection.rs b/src-tauri/src/db/connection.rs index dee03d8b..7fde28eb 100644 --- a/src-tauri/src/db/connection.rs +++ b/src-tauri/src/db/connection.rs @@ -73,7 +73,9 @@ impl PgConfig { urlencoding::encode(&self.password), self.host, self.port, - self.database, + // Encoded like the credentials: sqlx percent-decodes the path, so a + // database named `sales%2024` or `a#b` was read as something else. + urlencoding::encode(&self.database), ssl ) } @@ -142,7 +144,7 @@ impl MysqlConfig { urlencoding::encode(&self.password), self.host, self.port, - self.database, + urlencoding::encode(&self.database), params.join("&") ) } @@ -889,8 +891,17 @@ where // TCP demonstrably works. Stop bounding the handshake and let it land. break; } - match tokio::time::timeout(Duration::from_millis(*ms), attempt()).await { + let fut = attempt(); + tokio::pin!(fut); + match tokio::time::timeout(Duration::from_millis(*ms), &mut fut).await { Ok(res) => return res, + // The probe proved TCP WHILE this attempt was running. Checking only + // before an attempt missed that case, which is the common one against + // a distant host: a Sydney pooler probed reachable at 376ms, then + // attempt 1 was binned at 800ms and the connect landed at 4560ms, + // having paid TCP + TLS + auth twice. The stall is slowness, so keep + // the handshake that already has a head start. + Err(_) if tcp_ok.load(std::sync::atomic::Ordering::Relaxed) => return fut.await, Err(_) => log::info!("connect attempt {} exceeded {ms}ms, retrying with a fresh SYN", i + 1), } } @@ -1302,6 +1313,19 @@ mod tests { } } + /// A password ending in `%` and a database name with `%`/`#` must reach + /// sqlx intact: it percent-decodes the whole URL, so anything left raw was + /// cut or misread. + #[test] + fn pg_url_round_trips_special_characters_through_sqlx() { + let mut c = pg(false, None, None); + c.password = "Lm$$pR0D54%".into(); + c.database = "sales%2024#a".into(); + let opts: PgConnectOptions = c.connection_url().parse().expect("url parses"); + assert_eq!(opts.get_database(), Some("sales%2024#a")); + assert_eq!(opts.get_username(), "ada"); + } + #[test] fn pg_url_omits_tls_params_when_tls_is_off() { let url = pg(false, None, None).connection_url(); @@ -1454,6 +1478,27 @@ mod tests { assert_eq!(tries.load(Ordering::Relaxed), 1, "a slow but healthy handshake was restarted"); } + /// The probe usually lands WHILE the first attempt is running. That attempt + /// already has TCP and TLS behind it and must be kept, not restarted. + #[tokio::test(start_paused = true)] + async fn tcp_proven_mid_attempt_keeps_that_attempt() { + let tries = AtomicUsize::new(0); + let tcp_ok = std::sync::Arc::new(AtomicBool::new(false)); + let flag = tcp_ok.clone(); + tokio::spawn(async move { + tokio::time::sleep(Duration::from_millis(376)).await; + flag.store(true, Ordering::Relaxed); + }); + let got = retry_fast(&tcp_ok, || async { + tries.fetch_add(1, Ordering::Relaxed); + tokio::time::sleep(Duration::from_secs(4)).await; + Ok::(7) + }) + .await; + assert_eq!(got.unwrap(), 7); + assert_eq!(tries.load(Ordering::Relaxed), 1, "a handshake past a proven TCP path was restarted"); + } + /// While TCP is still unproven a stall really might be a lost SYN, and a fresh /// SYN is the only thing that helps - so there the ladder still fires. #[tokio::test(start_paused = true)] diff --git a/src-tauri/src/db/query.rs b/src-tauri/src/db/query.rs index 324067ad..e47ce811 100644 --- a/src-tauri/src/db/query.rs +++ b/src-tauri/src/db/query.rs @@ -7,6 +7,8 @@ use serde::{Deserialize, Serialize}; use serde_json::{json, Value}; use sqlx::{Column, Decode, Postgres, Row, TypeInfo, ValueRef}; use std::collections::HashMap; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::Arc; use tauri::State; use uuid::Uuid; @@ -2156,7 +2158,9 @@ pub async fn count_table_rows( .execute(&mut *tx) .await; let count = count_query.fetch_one(&mut *tx).await; - let _ = tx.rollback().await; + // COMMIT resets SET LOCAL exactly like ROLLBACK does, and unlike ROLLBACK it + // is accepted by Nile's proxy (see execute_sql_pg). + let _ = tx.commit().await; match count { Ok(n) => Ok(n), // 57014 = query_canceled (statement timeout). A count that can't finish @@ -2806,7 +2810,7 @@ pub async fn execute_sql( ActiveConnection::Postgres(_) | ActiveConnection::Mysql(_) => unreachable!(), } } => r, - _ = async { let _ = cancel_rx.await; } => Err("Query cancelled".to_string()), + _ = async { let _ = cancel_rx.await; } => Err(QUERY_CANCELLED.to_string()), }, }; super::connection::unregister_cancel(&state, &cancel_key); @@ -2875,14 +2879,63 @@ const EXECUTE_SQL_MAX_ROWS: usize = 1_000_000_000; /// heavier scans (e.g. tables with large TOASTed JSON columns) to finish. pub(crate) const EXECUTE_SQL_TIMEOUT_MS: i64 = 60_000; +/// What a stopped run reports. The editor matches on this text to show the run +/// as stopped rather than failed. +pub(crate) const QUERY_CANCELLED: &str = "Query cancelled"; + +/// Arm Stop for a Postgres run on `tx`'s connection. When `rx` fires, `cancelled` +/// is raised first so a row loop draining already-buffered rows bails on its next +/// row, then `pg_cancel_backend` stops the statement on the server. If the run +/// finishes first the sender is dropped, `rx.await` errors and the watcher exits. +async fn arm_pg_cancel( + tx: &mut sqlx::Transaction<'_, Postgres>, + pool: &sqlx::PgPool, + rx: Option>, + cancelled: Arc, +) { + let Some(rx) = rx else { return }; + let pid = sqlx::query_scalar::<_, i32>("SELECT pg_backend_pid()") + .fetch_one(&mut **tx) + .await + .ok(); + let cancel_pool = pool.clone(); + tokio::spawn(async move { + if rx.await.is_ok() { + cancelled.store(true, Ordering::Relaxed); + if let Some(pid) = pid { + let _ = sqlx::query("SELECT pg_cancel_backend($1)") + .bind(pid) + .execute(&cancel_pool) + .await; + } + } + }); +} + async fn execute_sql_pg( pool: &sqlx::PgPool, sql: &str, // When `Some`, real cancellation is armed: if the receiver fires we ask the // server to cancel this query's backend (so the statement actually stops) // rather than merely abandoning the future while the server keeps working. - // Callers that can't be cancelled (diff/multi paths) pass `None`. + // Callers that can't be cancelled (diff paths) pass `None`. cancel_rx: Option>, +) -> Result { + // Whatever error the cancelled statement surfaced ("canceling statement due + // to user request", a dropped stream), a stopped run reads as stopped. + let cancelled = Arc::new(AtomicBool::new(false)); + let result = run_sql_pg(pool, sql, cancel_rx, cancelled.clone()).await; + if cancelled.load(Ordering::Relaxed) { + return Err(QUERY_CANCELLED.into()); + } + result +} + +async fn run_sql_pg( + pool: &sqlx::PgPool, + sql: &str, + cancel_rx: Option>, + cancelled: Arc, ) -> Result { let started = std::time::Instant::now(); let query_ms = || started.elapsed().as_millis() as u64; @@ -2910,27 +2963,7 @@ async fn execute_sql_pg( .execute(&mut *tx) .await; - // Arm server-side cancellation. Capture the backend PID for THIS connection, - // then watch the cancel channel on a background task; on cancel we run - // pg_cancel_backend on a *separate* pooled connection, which makes the - // in-flight statement error out. If the query finishes first the sender is - // dropped and `rx.await` errors, so the watcher is a no-op. - if let Some(rx) = cancel_rx { - if let Ok(pid) = sqlx::query_scalar::<_, i32>("SELECT pg_backend_pid()") - .fetch_one(&mut *tx) - .await - { - let cancel_pool = pool.clone(); - tokio::spawn(async move { - if rx.await.is_ok() { - let _ = sqlx::query("SELECT pg_cancel_backend($1)") - .bind(pid) - .execute(&cancel_pool) - .await; - } - }); - } - } + arm_pg_cancel(&mut tx, pool, cancel_rx, cancelled.clone()).await; let last_idx = stmts.len() - 1; @@ -2954,6 +2987,9 @@ async fn execute_sql_pg( loop { match stream.try_next().await { Ok(Some(row)) => { + if cancelled.load(Ordering::Relaxed) { + return Err(QUERY_CANCELLED.into()); + } if data.is_empty() { columns = row .columns() @@ -2979,7 +3015,12 @@ async fn execute_sql_pg( let Some(msg) = failure else { break }; // The transaction is poisoned by the failed statement either way. let _ = tx.rollback().await; - if rewritten || !is_missing_binary_output(&msg) { + // The retry below re-runs only this last statement in a fresh + // transaction, so it is only sound when there is nothing before + // it: with `UPDATE …; SELECT …` the rollback above has already + // undone the UPDATE, and retrying just the SELECT would report + // success over a write that never happened. + if rewritten || last_idx != 0 || !is_missing_binary_output(&msg) { return Err(format!("Query failed: {msg}")); } let Some(wrapped) = text_safe_wrap(pool, stmt).await else { @@ -3000,7 +3041,15 @@ async fn execute_sql_pg( .execute(&mut *tx) .await; } - let _ = tx.rollback().await; + // COMMIT, not ROLLBACK. "Ends in a SELECT" does not mean "read-only": + // `UPDATE …; SELECT …`, a data-modifying CTE (`WITH d AS (DELETE … + // RETURNING *) SELECT …`) and `SELECT nextval(…)` all write, and a + // rollback here threw those writes away after showing their result. + // On a transaction that really was read-only, COMMIT costs the same. + // It also matters for Nile, whose proxy rejects this ROLLBACK and + // leaves the connection marked in-transaction, so sqlx closed it and + // the next query paid a fresh handshake. + let _ = tx.commit().await; let row_count = data.len() as i64; return Ok(SqlResult { @@ -3148,12 +3197,29 @@ pub(crate) fn split_sql_statements(sql: &str) -> Vec { out } -async fn execute_sql_multi_pg(pool: &sqlx::PgPool, stmts: &[String]) -> Result, String> { +async fn execute_sql_multi_pg( + pool: &sqlx::PgPool, + stmts: &[String], + cancel_rx: Option>, +) -> Result, String> { // Single statement - delegate to existing path (avoids code duplication) if stmts.len() == 1 { - return execute_sql_pg(pool, &stmts[0], None).await.map(|r| vec![r]); + return execute_sql_pg(pool, &stmts[0], cancel_rx).await.map(|r| vec![r]); } + let cancelled = Arc::new(AtomicBool::new(false)); + let result = run_sql_multi_pg(pool, stmts, cancel_rx, cancelled.clone()).await; + if cancelled.load(Ordering::Relaxed) { + return Err(QUERY_CANCELLED.into()); + } + result +} +async fn run_sql_multi_pg( + pool: &sqlx::PgPool, + stmts: &[String], + cancel_rx: Option>, + cancelled: Arc, +) -> Result, String> { let mut tx = pool .begin() .await @@ -3163,6 +3229,8 @@ async fn execute_sql_multi_pg(pool: &sqlx::PgPool, stmts: &[String]) -> Result = Vec::new(); for stmt in stmts { @@ -3180,6 +3248,9 @@ async fn execute_sql_multi_pg(pool: &sqlx::PgPool, stmts: &[String]) -> Result { + if cancelled.load(Ordering::Relaxed) { + return Err(QUERY_CANCELLED.into()); + } if data.is_empty() { columns = row .columns() @@ -3259,12 +3330,17 @@ pub async fn execute_sql_multi( super::connection::unregister_cancel(&state, &cancel_key); return Err("Query is empty".into()); } + // Postgres runs multi-statement scripts inside a single transaction and + // cancels them server-side, so Stop both frees the UI and ends the statement. + // Abandoning the future alone left the server scanning and the connection + // busy until the statement timeout. + if let ActiveConnection::Postgres(pool) = &conn { + let result = execute_sql_multi_pg(pool, &stmts, Some(cancel_rx)).await; + super::connection::unregister_cancel(&state, &cancel_key); + return result; + } let result = tokio::select! { r = async move { - // Postgres runs multi-statement scripts inside a single transaction - if let ActiveConnection::Postgres(pool) = &conn { - return execute_sql_multi_pg(pool, &stmts).await; - } // Other engines: execute sequentially, one result set per statement. // Cancellation happens at the outer select! (the future is dropped), // so per-statement executors get no cancel receiver - same as the @@ -3291,7 +3367,7 @@ pub async fn execute_sql_multi( } Ok(results) } => r, - _ = async { let _ = cancel_rx.await; } => Err("Query cancelled".to_string()), + _ = async { let _ = cancel_rx.await; } => Err(QUERY_CANCELLED.to_string()), }; super::connection::unregister_cancel(&state, &cancel_key); result diff --git a/src-tauri/src/db/redis.rs b/src-tauri/src/db/redis.rs index 69cef272..4409ce32 100644 --- a/src-tauri/src/db/redis.rs +++ b/src-tauri/src/db/redis.rs @@ -12,32 +12,90 @@ use std::time::Instant; // ── Connection ──────────────────────────────────────────────────────────────── -/// Open a multiplexed async connection from a Redis config. Builds a -/// `redis://[:password@]host:port/db` URL (`rediss://` when TLS is enabled) and -/// hands back a pooled multiplexed connection. -async fn open(cfg: &RedisConfig) -> Result<::redis::aio::MultiplexedConnection, String> { +/// Connections kept for reuse, keyed by URL (host, port, db, TLS and password). +/// +/// Every command used to open its own connection. Against a hosted Redis +/// (Upstash, Railway) that is DNS + TCP + TLS + AUTH per key browsed, three or +/// more round trips before the command itself. A multiplexed connection is +/// cheap to clone and safe to share, so one per config serves every command. +struct Cached { + conn: ::redis::aio::MultiplexedConnection, + last_used: Instant, +} +static POOL: std::sync::LazyLock>> = + std::sync::LazyLock::new(Default::default); + +/// Past this much idle time a cached connection is PINGed before reuse: hosted +/// Redis closes idle clients, and one round trip to find out beats a failed +/// command. +const IDLE_CHECK: std::time::Duration = std::time::Duration::from_secs(30); + +fn url_for(cfg: &RedisConfig) -> String { let scheme = if cfg.tls { "rediss" } else { "redis" }; let auth = match cfg.password.as_deref().map(str::trim).filter(|p| !p.is_empty()) { Some(pw) => format!(":{}@", urlencoding::encode(pw)), None => String::new(), }; - let url = format!("{scheme}://{auth}{}:{}/{}", cfg.host, cfg.port, cfg.db); + format!("{scheme}://{auth}{}:{}/{}", cfg.host, cfg.port, cfg.db) +} +async fn dial(url: &str) -> Result<::redis::aio::MultiplexedConnection, String> { let client = ::redis::Client::open(url).map_err(|e| format!("Redis connection failed: {e}"))?; + // Bounded: with no timeout an unreachable host (or a TLS port spoken to in + // plain TCP) left the connect awaiting until the OS gave up. + let config = ::redis::AsyncConnectionConfig::new() + .set_connection_timeout(Some(std::time::Duration::from_secs(10))) + .set_response_timeout(Some(std::time::Duration::from_secs(30))); client - .get_multiplexed_async_connection() + .get_multiplexed_async_connection_with_config(&config) .await .map_err(|e| format!("Redis connection failed: {e}")) } +/// A multiplexed connection for this config: the cached one when it is still +/// alive, otherwise a new one (which then becomes the cached one). +async fn open(cfg: &RedisConfig) -> Result<::redis::aio::MultiplexedConnection, String> { + let url = url_for(cfg); + let mut pool = POOL.lock().await; + if let Some(entry) = pool.get_mut(&url) { + let fresh = entry.last_used.elapsed() < IDLE_CHECK; + let alive = fresh + || ::redis::cmd("PING") + .query_async::<::redis::Value>(&mut entry.conn) + .await + .is_ok(); + if alive { + entry.last_used = Instant::now(); + return Ok(entry.conn.clone()); + } + pool.remove(&url); + } + // Dial without holding the lock for the whole handshake would let two + // commands race to open two connections; holding it is the simpler bound, + // and it only ever waits on the first connect for a given host. + let conn = dial(&url).await?; + pool.insert(url, Cached { conn: conn.clone(), last_used: Instant::now() }); + Ok(conn) +} + +/// Drop the cached connection for a config (disconnect, or a failed PING in +/// `ping`), so the next command dials fresh. +pub async fn forget(cfg: &RedisConfig) { + POOL.lock().await.remove(&url_for(cfg)); +} + /// Connectivity/credential check - issues a `PING`. pub async fn ping(cfg: &RedisConfig) -> Result<(), String> { let mut conn = open(cfg).await?; - ::redis::cmd("PING") + let r = ::redis::cmd("PING") .query_async::<::redis::Value>(&mut conn) .await .map(|_| ()) - .map_err(|e| format!("Redis PING failed: {e}")) + .map_err(|e| format!("Redis PING failed: {e}")); + if r.is_err() { + forget(cfg).await; + } + r } // ── Raw command execution ───────────────────────────────────────────────────── diff --git a/src-tauri/src/db/schema.rs b/src-tauri/src/db/schema.rs index 183f0af8..d43b0a53 100644 --- a/src-tauri/src/db/schema.rs +++ b/src-tauri/src/db/schema.rs @@ -167,9 +167,11 @@ async fn exact_row_count(pool: &PgPool, schema: &str, table: &str) -> Result = sqlx::query_scalar(&sql).fetch_one(&mut *tx).await; - // Read-only, so rollback is the cheap way back; failure to roll back doesn't - // change the answer. - let _ = tx.rollback().await; + // Read-only, so COMMIT is as cheap as ROLLBACK and resets SET LOCAL the same + // way. Not ROLLBACK: Nile's proxy rejects it here, the connection comes back + // marked in-transaction, sqlx closes it, and every row count on a Nile table + // cost a fresh ~4s handshake. + let _ = tx.commit().await; count.map_err(|e| format!("Failed to count rows for {table}: {e}")) } diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 62709b47..64835e31 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -12,11 +12,13 @@ static GLOBAL: mimalloc::MiMalloc = mimalloc::MiMalloc; mod app_lock; mod cloudflare; mod commands; +mod connection_store; mod copilot; mod db; mod docker; mod license; mod mcp; +mod oauth_page; mod omniroute; mod plugins; mod metrics; @@ -606,6 +608,8 @@ pub fn run() { commands::deactivate_license, commands::run_license_check, commands::init_sample_db, + commands::connections_store_read, + commands::connections_store_write, metrics::get_app_metrics, metrics::set_process_title, commands::enable_autostart, diff --git a/src-tauri/src/oauth_page.rs b/src-tauri/src/oauth_page.rs new file mode 100644 index 00000000..c5953b91 --- /dev/null +++ b/src-tauri/src/oauth_page.rs @@ -0,0 +1,188 @@ +/*! + * The page the browser lands on after a provider sign-in (Neon, Supabase, + * PlanetScale, Prisma, TiDB, Turso, Cloudflare). + * + * It is served by Stroke's own localhost callback listener, so it can't load + * anything: no network assets, no fonts from a CDN. The logo is inlined as a data + * URI and the page follows the browser's light/dark preference. It says who is + * talking (Stroke, with a link to the site), what happened, and what to do next, + * which is usually nothing but switching back to the app. + */ + +use base64::{engine::general_purpose::STANDARD, Engine}; +use std::sync::OnceLock; + +const WEBSITE: &str = "https://stroke.click"; + +/// White mark for dark pages, dark mark for light ones (see Logo.svelte). +fn logos() -> &'static (String, String) { + static LOGOS: OnceLock<(String, String)> = OnceLock::new(); + LOGOS.get_or_init(|| { + let uri = |png: &[u8]| format!("data:image/png;base64,{}", STANDARD.encode(png)); + ( + uri(include_bytes!("../../public/stroke_white.png")), + uri(include_bytes!("../../public/stroke_light.png")), + ) + }) +} + +fn escape(s: &str) -> String { + s.replace('&', "&").replace('<', "<").replace('>', ">").replace('"', """) +} + +/// `utm_source` says the visit came from the app, `utm_medium` which surface +/// sent it, `utm_campaign` which provider's sign-in, `utm_content` how it went. +fn website_link(provider: &str, ok: bool) -> String { + let slug: String = provider + .chars() + .filter_map(|c| if c.is_ascii_alphanumeric() { Some(c.to_ascii_lowercase()) } else if c == ' ' { Some('-') } else { None }) + .collect(); + format!( + "{WEBSITE}/?utm_source=stroke-app&utm_medium=oauth-callback&utm_campaign={slug}&utm_content={}", + if ok { "success" } else { "error" } + ) +} + +/// The full HTML for the callback tab. `provider` is the display name +/// ("Neon", "Cloudflare"); `ok` picks the success or failure copy. +pub fn page(ok: bool, provider: &str) -> String { + let (logo_dark_ui, logo_light_ui) = logos(); + let link = website_link(provider, ok); + let provider = escape(provider); + let (title, status, heading, body, tone, icon) = if ok { + ( + format!("Signed in to {provider} · Stroke"), + "Signed in", + format!("You're connected to {provider}"), + "Switch back to Stroke to pick a database. This tab can be closed.", + "ok", + r#""#, + ) + } else { + ( + format!("{provider} sign-in failed · Stroke"), + "Not signed in", + format!("{provider} sign-in didn't finish"), + "Nothing was saved. Close this tab and start the sign-in again from Stroke.", + "err", + r#""#, + ) + }; + + // One card, one leading edge: brand row, a status pill, the heading and the + // next step, all left-aligned. The first version centred a lone logo above + // a lone check circle, two marks stacked with nothing tying them together. + format!( + r##" + + + + + +{title} + + + +
+
+ + + Stroke +
+
+ {status} +

{heading}

+

{body}

+
+
+
+ The database studio for agents and humans + stroke.click +
+ +"## + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn link_carries_utm_for_the_provider_and_outcome() { + let html = page(true, "TiDB Cloud"); + assert!(html.contains("utm_source=stroke-app&utm_medium=oauth-callback&utm_campaign=tidb-cloud&utm_content=success")); + assert!(page(false, "Neon").contains("utm_campaign=neon&utm_content=error")); + } + + #[test] + fn provider_name_is_escaped() { + assert!(page(true, "").contains("<x>")); + } +} diff --git a/src-tauri/src/providers/mod.rs b/src-tauri/src/providers/mod.rs index 215a485d..6c94f8fd 100644 --- a/src-tauri/src/providers/mod.rs +++ b/src-tauri/src/providers/mod.rs @@ -17,6 +17,11 @@ mod neon; mod planetscale; mod prisma; mod supabase; +mod nile; +mod railway; +mod tidb; +mod turso; +mod upstash; use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine}; use serde::{Deserialize, Serialize}; @@ -56,6 +61,24 @@ pub enum Provider { Supabase, PlanetScale, Prisma, + TiDB, + Turso, + Railway, + Nile, + Upstash, +} + +/// How a provider signs the user in. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum SignIn { + /// Browser authorization code + localhost redirect (Neon, Supabase, …). + AuthCode, + /// OAuth 2.0 device code grant, RFC 8628 (TiDB Cloud): the browser confirms + /// a code and the app polls for the token, so no callback port is involved. + DeviceCode, + /// A CLI login page that redirects to localhost with the API token itself in + /// `?jwt=` (Turso). No code exchange and no refresh token. + TokenRedirect, } impl Provider { @@ -65,6 +88,11 @@ impl Provider { "supabase" => Ok(Self::Supabase), "planetscale" => Ok(Self::PlanetScale), "prisma" => Ok(Self::Prisma), + "tidb" => Ok(Self::TiDB), + "turso" => Ok(Self::Turso), + "railway" => Ok(Self::Railway), + "nile" => Ok(Self::Nile), + "upstash" => Ok(Self::Upstash), other => Err(format!("Unknown provider: {other}")), } } @@ -76,6 +104,34 @@ impl Provider { Self::Supabase => "supabase", Self::PlanetScale => "planetscale", Self::Prisma => "prisma", + Self::TiDB => "tidb", + Self::Turso => "turso", + Self::Railway => "railway", + Self::Nile => "nile", + Self::Upstash => "upstash", + } + } + + /// Display name, for the page the browser lands on after sign-in. + fn label(&self) -> &'static str { + match self { + Self::Neon => "Neon", + Self::Supabase => "Supabase", + Self::PlanetScale => "PlanetScale", + Self::Prisma => "Prisma", + Self::TiDB => "TiDB Cloud", + Self::Turso => "Turso", + Self::Railway => "Railway", + Self::Nile => "Nile", + Self::Upstash => "Upstash", + } + } + + fn sign_in(&self) -> SignIn { + match self { + Self::TiDB => SignIn::DeviceCode, + Self::Turso => SignIn::TokenRedirect, + _ => SignIn::AuthCode, } } @@ -85,14 +141,18 @@ impl Provider { Self::Supabase => supabase::OAUTH, Self::PlanetScale => planetscale::OAUTH, Self::Prisma => prisma::OAUTH, + Self::TiDB => tidb::OAUTH, + Self::Turso => turso::OAUTH, + Self::Railway => railway::OAUTH, + Self::Nile => nile::OAUTH, + Self::Upstash => upstash::OAUTH, } } - /// Whether a provider uses a pasted API token instead of the browser OAuth - /// dance. None currently do (Prisma moved to Management-API OAuth), but the - /// hook stays so a future token-only provider can opt in. + /// Whether a provider uses a pasted API credential instead of the browser + /// OAuth dance: Upstash, which offers no OAuth to third-party apps. fn is_token_based(&self) -> bool { - false + matches!(self, Self::Upstash) } /// Localhost callback ports to try, in order. PlanetScale accepts only ONE @@ -101,7 +161,9 @@ impl Provider { /// redirects use the full range so a busy port can fall back. fn callback_ports(&self) -> &'static [u16] { match self { - Self::PlanetScale => &[8989], + // Railway, like PlanetScale, matches the redirect URI exactly and + // the app registers one: http://localhost:8989/oauth/callback. + Self::PlanetScale | Self::Railway => &[8989], _ => CALLBACK_PORTS, } } @@ -116,7 +178,9 @@ impl Provider { /// Public PKCE clients ship no secret, so the token exchange goes DIRECTLY to /// the provider (no stroke.click proxy). Neon reuses neonctl's public client. fn is_public_client(&self) -> bool { - matches!(self, Self::Neon) + // Neon and Nile reuse their CLIs' public clients; Railway is registered + // as a native (public) app. None of them has a secret to inject. + matches!(self, Self::Neon | Self::Nile | Self::Railway) } /// The loopback redirect URI. Most providers registered @@ -125,6 +189,8 @@ impl Provider { fn redirect_uri(&self, port: u16) -> String { match self { Self::Neon => format!("http://127.0.0.1:{port}/callback"), + // nilecli's client allows any localhost port on /callback. + Self::Nile => format!("http://localhost:{port}/callback"), _ => format!("http://localhost:{port}/oauth/callback"), } } @@ -135,6 +201,11 @@ impl Provider { Self::Supabase => supabase::list_databases(token).await, Self::PlanetScale => planetscale::list_databases(token).await, Self::Prisma => prisma::list_databases(token).await, + Self::TiDB => tidb::list_databases(token).await, + Self::Turso => turso::list_databases(token).await, + Self::Railway => railway::list_databases(token).await, + Self::Nile => nile::list_databases(token).await, + Self::Upstash => upstash::list_databases(token).await, } } @@ -148,6 +219,11 @@ impl Provider { Self::Supabase => supabase::build_connection(token, db_ref).await, Self::PlanetScale => planetscale::build_connection(token, db_ref).await, Self::Prisma => prisma::build_connection(token, db_ref).await, + Self::TiDB => tidb::build_connection(token, db_ref).await, + Self::Turso => turso::build_connection(token, db_ref).await, + Self::Railway => railway::build_connection(token, db_ref).await, + Self::Nile => nile::build_connection(token, db_ref).await, + Self::Upstash => upstash::build_connection(token, db_ref).await, } } } @@ -181,10 +257,12 @@ pub struct ProviderDatabase { /// Everything the frontend needs to construct a `SavedConnection` and connect. #[derive(Debug, Serialize, Deserialize, Clone, Default)] pub struct ProviderConnection { - pub db_type: String, // "postgres" | "mysql" + pub db_type: String, // "postgres" | "mysql" | "libsql" + /// For libsql, the full `libsql://` URL (there is no host/port split). pub host: String, pub port: u16, pub username: String, + /// For libsql, the database auth token. pub password: String, pub database: String, pub ssl: bool, @@ -210,6 +288,12 @@ pub(crate) fn http() -> &'static reqwest::Client { .user_agent("stroke/1.0") .tcp_keepalive(std::time::Duration::from_secs(60)) .pool_max_idle_per_host(4) + // Drop idle sockets before the far end does. reqwest keeps them 90s + // by default, while the AWS load balancers in front of TiDB Cloud's + // API (and most others here) close idle connections at 60s: the + // next request went out on a dead socket and failed with a bare + // "error sending request", sometimes twice in a row. + .pool_idle_timeout(std::time::Duration::from_secs(20)) // Same reason as the Cloudflare client: Neon, Supabase, PlanetScale // and Prisma discovery all run through here, and an unbounded call // shows as a provider panel that never finishes loading. @@ -220,6 +304,42 @@ pub(crate) fn http() -> &'static reqwest::Client { }) } +/// Merge per-organization (or per-workspace) results where some may fail. +/// +/// A token often covers only some of the orgs a user belongs to: PlanetScale +/// answers 403 for an org not picked on the consent screen, Railway errors for a +/// workspace the token wasn't shared. One of those must not fail the whole list. +/// Rules: an ended session (401) always wins, since nothing else will work +/// either; otherwise skip failures, and only fail when every page did, with the +/// first error plus `hint`. +pub(crate) fn merge_partial(pages: Vec>, hint: &str) -> Result, String> { + if pages.iter().any(|r| matches!(r, Err(e) if e == UNAUTHORIZED)) { + return Err(UNAUTHORIZED.into()); + } + if !pages.is_empty() && pages.iter().all(Result::is_err) { + let first = pages.into_iter().find_map(Result::err).unwrap_or_default(); + return Err(if hint.is_empty() { first } else { format!("{first}. {hint}") }); + } + Ok(pages.into_iter().filter_map(Result::ok).collect()) +} + +/// A reqwest error with its causes. reqwest's own message stops at "error +/// sending request for url (…)", which hides whether it was DNS, a refused +/// connection, TLS, or a reset socket. +pub(crate) fn describe(e: &reqwest::Error) -> String { + let mut out = e.to_string(); + let mut src = std::error::Error::source(e); + while let Some(cause) = src { + let c = cause.to_string(); + if !out.contains(&c) { + out.push_str(": "); + out.push_str(&c); + } + src = cause.source(); + } + out +} + // ── PKCE helpers ───────────────────────────────────────────────────────────────── fn random_base64url(n: usize) -> String { @@ -264,18 +384,18 @@ async fn bind_callback_listener(ports: &[u16]) -> Result<(TcpListener, u16), Str )) } -const OK_HTML: &str = r#"Stroke - authorized - -

Authorization successful

You can close this tab and return to Stroke.

"#; -const ERR_HTML: &str = r#"Stroke - error - -

Authorization failed

You can close this tab and try again in Stroke.

"#; - -/// Wait for one OAuth redirect and return the authorization code. -async fn await_oauth_callback(listener: TcpListener, expected_state: &str) -> Result { +/// Wait for one OAuth redirect and return the value of `value_key` from its +/// query: the authorization code (`code`), or for a token redirect the token +/// itself (`jwt`). +async fn await_oauth_callback( + listener: TcpListener, + expected_state: &str, + value_key: &str, + provider_label: &str, +) -> Result { + let ok_page = crate::oauth_page::page(true, provider_label); + let err_page = crate::oauth_page::page(false, provider_label); let send = |html: &str| { format!( "HTTP/1.1 200 OK\r\nContent-Type: text/html; charset=utf-8\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}", @@ -284,26 +404,48 @@ async fn await_oauth_callback(listener: TcpListener, expected_state: &str) -> Re ) }; - let (mut stream, _) = listener - .accept() - .await - .map_err(|e| format!("Callback accept failed: {e}"))?; - - let mut buf = vec![0u8; 8192]; - let n = stream - .read(&mut buf) - .await - .map_err(|e| format!("Callback read failed: {e}"))?; - let req = String::from_utf8_lossy(&buf[..n]); - - let first_line = req.lines().next().unwrap_or(""); - let query = first_line - .split_whitespace() - .nth(1) - .unwrap_or("") - .split('?') - .nth(1) - .unwrap_or(""); + // Keep accepting until a request that is actually the redirect arrives. The + // listener used to take exactly one connection, so anything that reached the + // port first spent it: a browser's speculative preconnect (a socket that + // never sends), a `/favicon.ico` fetch, or any other local process. Every + // connection is read concurrently, so a silent socket can't hold up the real + // redirect behind it; the rest get a 404. AUTH_TIMEOUT_SECS still bounds it. + async fn read_request(mut stream: tokio::net::TcpStream, value_key: String) -> Option<(tokio::net::TcpStream, String)> { + let mut buf = vec![0u8; 8192]; + let n = match tokio::time::timeout(std::time::Duration::from_secs(10), stream.read(&mut buf)).await { + Ok(Ok(n)) if n > 0 => n, + _ => return None, + }; + let req = String::from_utf8_lossy(&buf[..n]); + let target = req.lines().next().unwrap_or("").split_whitespace().nth(1).unwrap_or("").to_string(); + let query = target.split_once('?').map(|(_, q)| q.to_string()).unwrap_or_default(); + let has_answer = query.split('&').any(|kv| { + let k = kv.split('=').next().unwrap_or(""); + k == value_key || k == "error" || k == "state" + }); + if !has_answer { + let _ = stream + .write_all(b"HTTP/1.1 404 Not Found\r\nContent-Length: 0\r\nConnection: close\r\n\r\n") + .await; + return None; + } + Some((stream, query)) + } + let mut pending = futures::stream::FuturesUnordered::new(); + let (mut stream, query) = loop { + tokio::select! { + accepted = listener.accept() => { + let (conn, _) = accepted.map_err(|e| format!("Callback accept failed: {e}"))?; + pending.push(read_request(conn, value_key.to_string())); + } + Some(done) = futures::StreamExt::next(&mut pending), if !pending.is_empty() => { + if let Some(hit) = done { + break hit; + } + } + } + }; + let query = query.as_str(); let (mut code, mut state, mut error) = (None, None, None); for pair in query.split('&') { @@ -314,7 +456,7 @@ async fn await_oauth_callback(listener: TcpListener, expected_state: &str) -> Re .map(|v| urlencoding::decode(v).unwrap_or_default().into_owned()) .unwrap_or_default(); match key { - "code" => code = Some(val), + k if k == value_key => code = Some(val), "state" => state = Some(val), "error" => error = Some(val), "error_description" if error.is_none() => error = Some(val), @@ -323,22 +465,22 @@ async fn await_oauth_callback(listener: TcpListener, expected_state: &str) -> Re } if let Some(err) = error { - let _ = stream.write_all(send(ERR_HTML).as_bytes()).await; + let _ = stream.write_all(send(&err_page).as_bytes()).await; return Err(format!("Provider denied authorization: {err}")); } let code = match code { Some(c) if !c.is_empty() => c, _ => { - let _ = stream.write_all(send(ERR_HTML).as_bytes()).await; + let _ = stream.write_all(send(&err_page).as_bytes()).await; return Err("No authorization code in callback".into()); } }; if state.as_deref() != Some(expected_state) { - let _ = stream.write_all(send(ERR_HTML).as_bytes()).await; + let _ = stream.write_all(send(&err_page).as_bytes()).await; return Err("OAuth state mismatch - possible CSRF".into()); } - let _ = stream.write_all(send(OK_HTML).as_bytes()).await; + let _ = stream.write_all(send(&ok_page).as_bytes()).await; let _ = stream.flush().await; Ok(code) } @@ -351,21 +493,48 @@ struct TokenResponse { expires_in: Option, } +/// Why a token request failed. Only `Rejected` - the provider answered +/// `invalid_grant`, RFC 6749's "this refresh token is expired, revoked or +/// already used" - ends a session. Anything else (network, proxy down, a +/// misconfigured client) is not the user's sign-in going bad and must not +/// sign them out. +enum TokenError { + Rejected(String), + Other(String), +} + +impl From for String { + fn from(e: TokenError) -> String { + match e { + TokenError::Rejected(m) | TokenError::Other(m) => m, + } + } +} + +/// Canonical error an adapter returns for an HTTP 401 from the provider API. +/// The command layer matches it to refresh once and retry. +pub(crate) const UNAUTHORIZED: &str = "provider API returned 401 Unauthorized"; + +/// Shown when the session can't be renewed. Starts with "Not signed in" so the +/// UI recognises it and offers "Sign in again". +const SESSION_EXPIRED: &str = + "Not signed in to this provider: the session expired or was revoked. Sign in again to continue."; + /// POST a token request to `url` - either the real provider token endpoint /// (public clients) or the stroke.click proxy (confidential clients, where the /// proxy injects the secret and `params` includes `provider`). -async fn post_token(url: &str, params: &[(&str, &str)]) -> Result { +async fn post_token(url: &str, params: &[(&str, &str)]) -> Result { let resp = http() .post(url) .form(params) .send() .await - .map_err(|e| format!("Token request failed: {e}"))?; + .map_err(|e| TokenError::Other(format!("Token request failed: {e}")))?; let status = resp.status().as_u16(); let text = resp .text() .await - .map_err(|e| format!("Token request failed: {e}"))?; + .map_err(|e| TokenError::Other(format!("Token request failed: {e}")))?; // A non-JSON body from the proxy usually means it isn't deployed (the request // hit the marketing site's HTML). Surface that clearly. let body: serde_json::Value = serde_json::from_str(&text).map_err(|_| { @@ -375,19 +544,26 @@ async fn post_token(url: &str, params: &[(&str, &str)]) -> Result, redirect_uri: &str, -) -> Result { +) -> Result { let mut params = vec![ ("grant_type", "authorization_code"), ("code", code), @@ -426,7 +602,7 @@ async fn refresh_token( provider_key: &str, is_public: bool, refresh: &str, -) -> Result { +) -> Result { let mut params = vec![ ("grant_type", "refresh_token"), ("refresh_token", refresh), @@ -453,65 +629,99 @@ async fn store_tokens( email: Option<&str>, ) -> Result<(), String> { let k = p.key(); - let mut map = crate::secrets::read_all_async(app).await; - map.insert(format!("__{k}_access__"), access.to_string()); - if let Some(r) = refresh { - map.insert(format!("__{k}_refresh__"), r.to_string()); - } - if let Some(exp) = expires_in { - map.insert(format!("__{k}_expires__"), (now_secs() + exp - 30).to_string()); - } - if let Some(e) = email { - map.insert(format!("__{k}_email__"), e.to_string()); - } - crate::secrets::write_all_async(app, map).await + let access = access.to_string(); + let refresh = refresh.map(str::to_string); + let email = email.map(str::to_string); + crate::secrets::update_async(app, move |map| { + map.insert(format!("__{k}_access__"), access); + if let Some(r) = refresh { + map.insert(format!("__{k}_refresh__"), r); + } + match expires_in { + // saturating: a provider answering expires_in < 30 used to underflow. + Some(exp) => { + map.insert(format!("__{k}_expires__"), (now_secs() + exp.saturating_sub(30)).to_string()); + } + // A new token without an expiry must not inherit the old one's. + None => { + map.remove(&format!("__{k}_expires__")); + } + } + if let Some(e) = email { + map.insert(format!("__{k}_email__"), e); + } + }) + .await } async fn clear_tokens(app: &tauri::AppHandle, p: Provider) -> Result<(), String> { let k = p.key(); - let mut map = crate::secrets::read_all_async(app).await; - for suffix in ["access", "refresh", "expires", "email"] { - map.remove(&format!("__{k}_{suffix}__")); - } - crate::secrets::write_all_async(app, map).await + crate::secrets::update_async(app, move |map| { + for suffix in ["access", "refresh", "expires", "email"] { + map.remove(&format!("__{k}_{suffix}__")); + } + }) + .await } +// One refresh at a time. Prisma and Supabase rotate refresh tokens: the first +// refresh consumes the stored one, so a second concurrent refresh with the same +// token is rejected. Opening the panel fires the status check and the database +// list together, which is exactly that race. +static REFRESH_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(()); + /// A valid access token, refreshing transparently if the stored one expired. /// For token-based providers, the "access token" is the pasted API key. -async fn valid_token(app: &tauri::AppHandle, p: Provider) -> Result { +/// +/// `rejected` is an access token the provider API just answered 401 for: it is +/// refreshed even if its stored expiry says it is still good. +async fn valid_token( + app: &tauri::AppHandle, + p: Provider, + rejected: Option<&str>, +) -> Result { let k = p.key(); - let map = crate::secrets::read_all_async(app).await; - let access = map - .get(&format!("__{k}_access__")) - .cloned() - .ok_or("Not signed in to this provider")?; + let fresh = |map: &std::collections::HashMap| -> Option { + let access = map.get(&format!("__{k}_access__"))?; + // Missing expiry => assume the token is long-lived and valid, mirroring + // the Cloudflare flow. A `0` default treated every such token as already + // expired and forced a needless re-login on every call. + let expires = map + .get(&format!("__{k}_expires__")) + .and_then(|s| s.parse::().ok()) + .unwrap_or(u64::MAX); + let usable = now_secs() < expires && rejected != Some(access.as_str()); + usable.then(|| access.clone()) + }; + let map = crate::secrets::read_all_async(app).await; + if !map.contains_key(&format!("__{k}_access__")) { + return Err(SESSION_EXPIRED.into()); + } if p.is_token_based() { - return Ok(access); + return Ok(map[&format!("__{k}_access__")].clone()); } - - // Missing expiry => assume the token is long-lived and valid, mirroring the - // Cloudflare flow (`unwrap_or(u64::MAX)`). A `0` default treated every such - // token as already expired and forced a needless re-login on every call - - // the main cause of "it asks me to sign in again each day". - let expires = map - .get(&format!("__{k}_expires__")) - .and_then(|s| s.parse::().ok()) - .unwrap_or(u64::MAX); - if now_secs() < expires { - return Ok(access); + if let Some(t) = fresh(&map) { + return Ok(t); } - // Past the stored expiry: try to renew silently with the refresh token. If we - // can't - no refresh token, or the refresh call fails (e.g. a transient proxy - // hiccup) - fall back to the existing access token instead of forcing a - // re-login. The downstream API call is the real arbiter: a genuinely dead - // token surfaces a clear auth error there (and the UI offers reconnect), while - // a still-valid or barely-past-buffer token keeps working. This favours long - // session persistence over eager sign-out. - let refresh = match map.get(&format!("__{k}_refresh__")).cloned() { - Some(r) => r, - None => return Ok(access), + let _guard = REFRESH_LOCK.lock().await; + // Whoever held the lock may already have refreshed. + let map = crate::secrets::read_all_async(app).await; + let Some(access) = map.get(&format!("__{k}_access__")).cloned() else { + return Err(SESSION_EXPIRED.into()); + }; + if let Some(t) = fresh(&map) { + return Ok(t); + } + let Some(refresh) = map.get(&format!("__{k}_refresh__")).cloned() else { + // Nothing to renew with. An expired-by-the-clock token may still work + // (the API is the real arbiter); one the API already refused won't. + if rejected.is_some() { + clear_tokens(app, p).await?; + return Err(SESSION_EXPIRED.into()); + } + return Ok(access); }; let cfg = p.oauth(); match refresh_token(&cfg, p.key(), p.is_public_client(), &refresh).await { @@ -527,7 +737,44 @@ async fn valid_token(app: &tauri::AppHandle, p: Provider) -> Result Ok(access), + // The provider says this sign-in is over: clear it so the UI shows + // "Sign in again" instead of failing on every call. + Err(TokenError::Rejected(_)) => { + clear_tokens(app, p).await?; + Err(SESSION_EXPIRED.into()) + } + // Network or proxy trouble. Keep the session: hand back the current + // token unless the API already refused it, in which case say why. + Err(TokenError::Other(e)) => { + if rejected.is_some() { + Err(format!("Could not renew the sign-in: {e}")) + } else { + Ok(access) + } + } + } +} + +/// Run a provider API call with a valid token. On a 401 the token is refreshed +/// once and the call retried; a second 401 ends the session. +async fn with_token(app: &tauri::AppHandle, p: Provider, call: F) -> Result +where + F: Fn(String) -> Fut, + Fut: std::future::Future>, +{ + let token = valid_token(app, p, None).await?; + match call(token.clone()).await { + Err(e) if e == UNAUTHORIZED => { + let renewed = valid_token(app, p, Some(&token)).await?; + match call(renewed).await { + Err(e) if e == UNAUTHORIZED => { + clear_tokens(app, p).await?; + Err(SESSION_EXPIRED.into()) + } + r => r, + } + } + r => r, } } @@ -555,6 +802,30 @@ pub async fn provider_start_oauth( oauth_cancel().notify_waiters(); tokio::time::sleep(std::time::Duration::from_millis(150)).await; + match p.sign_in() { + SignIn::AuthCode => {} + SignIn::DeviceCode => { + let t = device_code_sign_in(&app, &cfg).await?; + // No refresh token is stored even if one comes back: the refresh + // path posts form-encoded to the stroke.click proxy, which this + // provider isn't behind. Without one, an expired token ends the + // session cleanly ("Sign in again") instead of failing to renew. + store_tokens(&app, p, &t.access_token, None, t.expires_in, None).await?; + return Ok(ProviderOAuthStatus { connected: true, email: None }); + } + SignIn::TokenRedirect => { + let token = token_redirect_sign_in(&app, &cfg, &state).await?; + store_tokens(&app, p, &token, None, None, None).await?; + return Ok(ProviderOAuthStatus { connected: true, email: None }); + } + } + + if cfg.client_id.is_empty() { + return Err(format!( + "{} sign-in isn't set up in this build yet: its OAuth app client id is missing.", + p.label() + )); + } let (listener, port) = bind_callback_listener(p.callback_ports()).await?; let redirect_uri = p.redirect_uri(port); @@ -570,14 +841,23 @@ pub async fn provider_start_oauth( } else { String::new() }; + // OIDC only issues a refresh token for `offline_access` together with + // `prompt=consent`; Railway's access tokens last an hour without one. Only + // Railway: Neon and Prisma already return refresh tokens without it. + let prompt_param = if p == Provider::Railway { + "&prompt=consent" + } else { + "" + }; let auth_url = format!( - "{}?response_type=code&client_id={}&redirect_uri={}{}&state={}{}", + "{}?response_type=code&client_id={}&redirect_uri={}{}&state={}{}{}", cfg.auth_url, urlencoding::encode(cfg.client_id), urlencoding::encode(&redirect_uri), scope_param, urlencoding::encode(&state), pkce_param, + prompt_param, ); eprintln!("[provider oauth] {} authorize URL: {auth_url}", p.key()); @@ -591,14 +871,16 @@ pub async fn provider_start_oauth( let code = tokio::select! { r = tokio::time::timeout( std::time::Duration::from_secs(AUTH_TIMEOUT_SECS), - await_oauth_callback(listener, &state), + await_oauth_callback(listener, &state, "code", p.label()), ) => r.map_err(|_| "Authorization timed out".to_string())??, // Dropping the other branch's future here drops `listener` → port freed. _ = cancelled => return Err("cancelled".to_string()), }; let verifier_opt = if p.uses_pkce() { Some(verifier.as_str()) } else { None }; - let t = exchange_code(&cfg, p.key(), p.is_public_client(), &code, verifier_opt, &redirect_uri).await?; + let t = exchange_code(&cfg, p.key(), p.is_public_client(), &code, verifier_opt, &redirect_uri) + .await + .map_err(String::from)?; store_tokens( &app, p, @@ -615,6 +897,118 @@ pub async fn provider_start_oauth( }) } +/// OAuth 2.0 device code grant (RFC 8628). Opens the verification page with the +/// code prefilled, then polls the token endpoint until the user approves it, +/// declines it, the code expires, or they hit Cancel. +async fn device_code_sign_in(app: &tauri::AppHandle, cfg: &OAuthConfig) -> Result { + let resp = http() + .post(cfg.auth_url) + .json(&serde_json::json!({ "client_id": cfg.client_id })) + .send() + .await + .map_err(|e| format!("Device authorization failed: {e}"))?; + let status = resp.status().as_u16(); + let auth: serde_json::Value = resp + .json() + .await + .map_err(|e| format!("Device authorization returned bad JSON: {e}"))?; + if !(200..300).contains(&status) { + let msg = auth["error_description"].as_str().or_else(|| auth["error"].as_str()).unwrap_or("request failed"); + return Err(format!("Device authorization failed ({status}): {msg}")); + } + let device_code = auth["device_code"].as_str().ok_or("Device authorization: missing device_code")?.to_string(); + let verify_url = auth["verification_uri_complete"] + .as_str() + .or_else(|| auth["verification_uri"].as_str()) + .ok_or("Device authorization: missing verification URL")? + .to_string(); + let mut interval = auth["interval"].as_u64().unwrap_or(5).max(1); + let expires = auth["expires_in"].as_u64().unwrap_or(AUTH_TIMEOUT_SECS).min(AUTH_TIMEOUT_SECS); + + // The confirmation code goes to the UI, not just a log line: the user checks + // that the code on TiDB's page matches the one Stroke shows, which is the + // whole point of the device grant, and it's the way back if the tab closes. + let _ = tauri::Emitter::emit( + app, + "provider-device-code", + serde_json::json!({ + "userCode": auth["user_code"].as_str().unwrap_or_default(), + "verificationUri": auth["verification_uri"].as_str().unwrap_or(&verify_url), + "verificationUriComplete": verify_url, + "expiresIn": expires, + }), + ); + tauri_plugin_opener::OpenerExt::opener(app) + .open_url(verify_url, None::<&str>) + .map_err(|e| format!("Could not open browser: {e}"))?; + + let started = std::time::Instant::now(); + let cancelled = oauth_cancel().notified(); + tokio::pin!(cancelled); + // Counts a Cancel that lands between polls, not only one during a sleep. + cancelled.as_mut().enable(); + loop { + tokio::select! { + _ = tokio::time::sleep(std::time::Duration::from_secs(interval)) => {} + _ = &mut cancelled => return Err("cancelled".into()), + } + if started.elapsed().as_secs() >= expires { + return Err("Authorization timed out".into()); + } + let resp = http() + .post(cfg.token_url) + .json(&serde_json::json!({ + "device_code": device_code, + "grant_type": "urn:ietf:params:oauth:grant-type:device_code", + "client_id": cfg.client_id, + })) + .send() + .await + .map_err(|e| format!("Token request failed: {e}"))?; + let body: serde_json::Value = resp.json().await.unwrap_or_default(); + if let Some(access) = body["access_token"].as_str().filter(|t| !t.is_empty()) { + return Ok(TokenResponse { + access_token: access.to_string(), + refresh_token: None, + expires_in: body["expires_in"].as_u64(), + }); + } + match body["error"].as_str().unwrap_or("") { + "authorization_pending" | "" => {} + "slow_down" => interval += 5, + "expired_token" => return Err("Authorization timed out".into()), + "access_denied" => return Err("Provider denied authorization: access_denied".into()), + other => return Err(format!("OAuth token error: {other}")), + } + } +} + +/// Turso-style CLI login: the login page redirects to our localhost listener +/// with the API token in `?jwt=`. +async fn token_redirect_sign_in( + app: &tauri::AppHandle, + cfg: &OAuthConfig, + state: &str, +) -> Result { + let (listener, port) = bind_callback_listener(CALLBACK_PORTS).await?; + let url = format!( + "{}/?port={port}&redirect=true&type=cli&state={}", + cfg.auth_url.trim_end_matches('/'), + urlencoding::encode(state), + ); + tauri_plugin_opener::OpenerExt::opener(app) + .open_url(url, None::<&str>) + .map_err(|e| format!("Could not open browser: {e}"))?; + let cancelled = oauth_cancel().notified(); + tokio::select! { + r = tokio::time::timeout( + std::time::Duration::from_secs(AUTH_TIMEOUT_SECS), + await_oauth_callback(listener, state, "jwt", "Turso"), + ) => r.map_err(|_| "Authorization timed out".to_string())?, + _ = cancelled => Err("cancelled".to_string()), + } +} + /// Abort an in-flight OAuth wait (Cancel button) - frees the callback port. #[tauri::command] pub fn provider_cancel_oauth() { @@ -660,8 +1054,7 @@ pub async fn provider_list_databases( provider: String, ) -> Result, String> { let p = Provider::parse(&provider)?; - let token = valid_token(&app, p).await?; - p.list_databases(&token).await + with_token(&app, p, |token| async move { p.list_databases(&token).await }).await } #[tauri::command] @@ -671,6 +1064,62 @@ pub async fn provider_build_connection( db_ref: String, ) -> Result { let p = Provider::parse(&provider)?; - let token = valid_token(&app, p).await?; - p.build_connection(&token, &db_ref).await + with_token(&app, p, |token| { + let db_ref = db_ref.clone(); + async move { p.build_connection(&token, &db_ref).await } + }) + .await +} + +#[cfg(test)] +mod callback_tests { + use super::await_oauth_callback; + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + + async fn hit(port: u16, path: &str) -> String { + let mut s = tokio::net::TcpStream::connect(("127.0.0.1", port)).await.unwrap(); + s.write_all(format!("GET {path} HTTP/1.1\r\nHost: localhost\r\n\r\n").as_bytes()).await.unwrap(); + let mut out = String::new(); + let _ = s.read_to_string(&mut out).await; + out + } + + /// A stray request (favicon, preconnect) used to spend the only accept and + /// fail the sign-in. It must be answered 404 and the wait must go on. + #[tokio::test] + async fn stray_requests_do_not_consume_the_callback() { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let port = listener.local_addr().unwrap().port(); + let waiter = tokio::spawn(async move { await_oauth_callback(listener, "s1", "code", "Neon").await }); + // A socket that connects and never sends, then a favicon fetch. + let _silent = tokio::net::TcpStream::connect(("127.0.0.1", port)).await.unwrap(); + let favicon = hit(port, "/favicon.ico").await; + assert!(favicon.starts_with("HTTP/1.1 404"), "{favicon}"); + let ok = hit(port, "/oauth/callback?code=abc&state=s1").await; + assert!(ok.contains("200 OK")); + assert_eq!(waiter.await.unwrap().unwrap(), "abc"); + } +} + +#[cfg(test)] +mod merge_tests { + use super::{merge_partial, UNAUTHORIZED}; + + #[test] + fn one_ungranted_org_does_not_fail_the_list() { + let pages = vec![Ok(1), Err("PlanetScale API error (403)".to_string()), Ok(3)]; + assert_eq!(merge_partial(pages, "hint").unwrap(), vec![1, 3]); + } + + #[test] + fn all_failing_reports_the_first_error_with_the_hint() { + let pages: Vec> = vec![Err("403 a".into()), Err("403 b".into())]; + assert_eq!(merge_partial(pages, "Sign in again.").unwrap_err(), "403 a. Sign in again."); + } + + #[test] + fn an_ended_session_always_wins() { + let pages = vec![Ok(1), Err(UNAUTHORIZED.to_string())]; + assert_eq!(merge_partial(pages, "").unwrap_err(), UNAUTHORIZED); + } } diff --git a/src-tauri/src/providers/neon.rs b/src-tauri/src/providers/neon.rs index d0b62576..5e2d9bc2 100644 --- a/src-tauri/src/providers/neon.rs +++ b/src-tauri/src/providers/neon.rs @@ -40,6 +40,10 @@ async fn get(token: &str, path: &str) -> Result { .await .map_err(|e| format!("Neon request failed: {e}"))?; let status = resp.status().as_u16(); + // Before reading the body: a 401 may not be JSON, and the caller refreshes on it. + if status == 401 { + return Err(super::UNAUTHORIZED.into()); + } let text = resp .text() .await @@ -61,10 +65,25 @@ async fn get(token: &str, path: &str) -> Result { pub async fn list_databases(token: &str) -> Result, String> { let orgs_body = get(token, "/users/me/organizations").await?; + // One request per org, all in flight at once: with several orgs the serial + // loop cost a full round-trip each before the picker could show anything. + // try_join_all keeps the orgs' order, so the list reads the same as before. + let org_ids: Vec<&str> = orgs_body["organizations"] + .as_array() + .into_iter() + .flatten() + .filter_map(|org| org["id"].as_str()) + .collect(); + let pages = futures::future::join_all(org_ids.iter().map(|org_id| { + let path = format!("/projects?org_id={}", urlencoding::encode(org_id)); + async move { get(token, &path).await } + })) + .await; + // An org the token can't read is skipped rather than failing the list. + let pages = super::merge_partial(pages, "")?; + let mut out = Vec::new(); - for org in orgs_body["organizations"].as_array().into_iter().flatten() { - let Some(org_id) = org["id"].as_str() else { continue }; - let body = get(token, &format!("/projects?org_id={}", urlencoding::encode(org_id))).await?; + for body in &pages { for p in body["projects"].as_array().into_iter().flatten() { if let Some(id) = p["id"].as_str() { out.push(ProviderDatabase { @@ -85,21 +104,12 @@ pub async fn list_databases(token: &str) -> Result, String pub async fn build_connection(token: &str, project_id: &str) -> Result { // Find the project's default branch (fall back to the first). let branches = get(token, &format!("/projects/{project_id}/branches")).await?; - let branch_list = branches["branches"].as_array().cloned().unwrap_or_default(); - let branch = branch_list - .iter() - .find(|b| b["default"].as_bool() == Some(true) || b["primary"].as_bool() == Some(true)) - .or_else(|| branch_list.first()) - .ok_or("This Neon project has no branches")?; - let branch_id = branch["id"].as_str().ok_or("Neon: missing branch id")?; + let branch_id = default_branch(&branches)?; + let branch_id = branch_id.as_str(); let dbs = get(token, &format!("/projects/{project_id}/branches/{branch_id}/databases")).await?; - let first = dbs["databases"] - .as_array() - .and_then(|a| a.first()) - .ok_or("This Neon branch has no databases yet")?; - let db_name = first["name"].as_str().ok_or("Neon: missing database name")?; - let role = first["owner_name"].as_str().ok_or("Neon: missing owner role")?; + let (db_name, role) = first_database(&dbs)?; + let (db_name, role) = (db_name.as_str(), role.as_str()); let uri_body = get( token, @@ -129,6 +139,28 @@ pub async fn build_connection(token: &str, project_id: &str) -> Result Result { + let list = branches["branches"].as_array().cloned().unwrap_or_default(); + let branch = list + .iter() + .find(|b| b["default"].as_bool() == Some(true) || b["primary"].as_bool() == Some(true)) + .or_else(|| list.first()) + .ok_or("This Neon project has no branches")?; + branch["id"].as_str().map(String::from).ok_or_else(|| "Neon: missing branch id".into()) +} + +/// The branch's first database and the role that owns it. +fn first_database(dbs: &Value) -> Result<(String, String), String> { + let first = dbs["databases"] + .as_array() + .and_then(|a| a.first()) + .ok_or("This Neon branch has no databases yet")?; + let name = first["name"].as_str().ok_or("Neon: missing database name")?; + let role = first["owner_name"].as_str().ok_or("Neon: missing owner role")?; + Ok((name.to_string(), role.to_string())) +} + /// Parse `postgres://user:pass@host[:port]/db?query` into its parts (URL-decoding /// the userinfo, which Neon percent-encodes). pub(crate) fn parse_pg_uri(uri: &str) -> Result<(String, u16, String, String, String), String> { @@ -154,3 +186,38 @@ pub(crate) fn parse_pg_uri(uri: &str) -> Result<(String, u16, String, String, St let dec = |s: &str| urlencoding::decode(s).map(|c| c.into_owned()).unwrap_or_else(|_| s.to_string()); Ok((host, port, dec(user_raw), dec(pass_raw), dec(&database))) } + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn default_branch_is_preferred_over_the_first() { + let b = json!({ "branches": [ { "id": "br-dev" }, { "id": "br-main", "default": true } ] }); + assert_eq!(default_branch(&b).unwrap(), "br-main"); + let legacy = json!({ "branches": [ { "id": "br-a" }, { "id": "br-b", "primary": true } ] }); + assert_eq!(default_branch(&legacy).unwrap(), "br-b"); + assert_eq!(default_branch(&json!({ "branches": [ { "id": "only" } ] })).unwrap(), "only"); + assert!(default_branch(&json!({ "branches": [] })).is_err()); + } + + #[test] + fn first_database_carries_its_owner_role() { + let dbs = json!({ "databases": [ { "name": "neondb", "owner_name": "neondb_owner" } ] }); + assert_eq!(first_database(&dbs).unwrap(), ("neondb".into(), "neondb_owner".into())); + assert!(first_database(&json!({ "databases": [] })).unwrap_err().contains("no databases")); + } + + #[test] + fn connection_uri_decodes_credentials_and_drops_the_query() { + let (h, p, u, pw, d) = parse_pg_uri( + "postgresql://neondb_owner:p%40ss%2Fw@ep-x-pooler.c-4.us-east-1.aws.neon.tech/neondb?sslmode=require&channel_binding=require", + ) + .unwrap(); + assert_eq!((h.as_str(), p), ("ep-x-pooler.c-4.us-east-1.aws.neon.tech", 5432)); + assert_eq!((u.as_str(), pw.as_str(), d.as_str()), ("neondb_owner", "p@ss/w", "neondb")); + let (_, port, ..) = parse_pg_uri("postgres://u:p@h:6543/db").unwrap(); + assert_eq!(port, 6543); + } +} diff --git a/src-tauri/src/providers/nile.rs b/src-tauri/src/providers/nile.rs new file mode 100644 index 00000000..ffb940ce --- /dev/null +++ b/src-tauri/src/providers/nile.rs @@ -0,0 +1,161 @@ +/*! + * Nile adapter - Postgres. Sign in through the browser the way `nile connect` + * does, list every database across the user's workspaces, and create database + * credentials on connect (Nile hands out a credential id + password pair; the id + * is the Postgres user). + * + * Reuses the official CLI's public PKCE client (`nilecli`), like Neon's + * neonctl: no secret, token exchange straight to Nile, redirect to + * `http://localhost:{port}/callback` on any free port. + */ + +use super::{http, OAuthConfig, ProviderConnection, ProviderDatabase}; +use serde_json::Value; + +pub const OAUTH: OAuthConfig = OAuthConfig { + client_id: "nilecli", + auth_url: "https://console.thenile.dev/authorize", + token_url: "https://global.thenile.dev/oauth2/token", + // The CLI sends no scope; the token carries the developer's own access. + scopes: "", +}; + +const API: &str = "https://global.thenile.dev"; + +async fn send(req: reqwest::RequestBuilder, what: &str) -> Result { + let resp = req + .send() + .await + .map_err(|e| format!("Nile request failed ({what}): {}", super::describe(&e)))?; + let status = resp.status().as_u16(); + if status == 401 { + return Err(super::UNAUTHORIZED.into()); + } + let text = resp.text().await.map_err(|e| format!("Nile read failed: {e}"))?; + let body: Value = serde_json::from_str(&text).unwrap_or(Value::Null); + if !(200..300).contains(&status) { + let msg = body["message"] + .as_str() + .or_else(|| body["errors"][0].as_str()) + .map(String::from) + .unwrap_or_else(|| text.chars().take(160).collect()); + return Err(format!("Nile API error ({status}) on {what}: {msg}")); + } + Ok(body) +} + +async fn get(token: &str, path: &str, what: &str) -> Result { + send(http().get(format!("{API}{path}")).bearer_auth(token), what).await +} + +/// `AWS_US_WEST_2` → `us-west-2.db.thenile.dev`; Azure regions live under +/// `db.{region}.azure.thenile.dev` (the same mapping `nile connect` uses). +fn host_for(region: &str) -> String { + let lower = region.to_ascii_lowercase(); + let mut parts = lower.split('_'); + let cloud = parts.next().unwrap_or("aws"); + let id = parts.collect::>().join("-"); + if cloud == "azure" { + format!("db.{id}.azure.thenile.dev") + } else { + format!("{id}.db.thenile.dev") + } +} + +pub async fn list_databases(token: &str) -> Result, String> { + let workspaces = get(token, "/workspaces", "listing workspaces").await?; + let slugs: Vec<&str> = workspaces + .as_array() + .into_iter() + .flatten() + .filter_map(|w| w["slug"].as_str()) + .collect(); + // A workspace the token can't read is skipped rather than failing the list. + let pages = futures::future::join_all(slugs.iter().map(|slug| { + let path = format!("/workspaces/{}/databases", urlencoding::encode(slug)); + async move { get(token, &path, "listing databases").await.map(|b| (*slug, b)) } + })) + .await; + let pages = super::merge_partial(pages, "")?; + + Ok(pages.iter().flat_map(|(slug, dbs)| parse_databases(slug, dbs)).collect()) +} + +/// One workspace's databases → picker rows. `db_ref` is "{workspace}/{database}". +fn parse_databases(slug: &str, dbs: &Value) -> Vec { + dbs.as_array() + .into_iter() + .flatten() + .filter_map(|db| { + let name = db["name"].as_str()?; + let region = db["region"].as_str().unwrap_or_default(); + Some(ProviderDatabase { + db_ref: format!("{slug}/{name}"), + name: name.to_string(), + region: (!region.is_empty()).then(|| region.to_ascii_lowercase().replace('_', "-")), + kind: Some("Database".into()), + host: (!region.is_empty()).then(|| host_for(region)), + }) + }) + .collect() +} + +/// The credentials-create response → a connection. The credential id is the +/// Postgres user; the host comes from the database's region. +fn connection_from_credentials(creds: &Value, database: &str) -> Result { + let user = creds["id"].as_str().ok_or("Nile returned no credential id")?; + let password = creds["password"].as_str().ok_or("Nile returned no password")?; + let region = creds["database"]["region"] + .as_str() + .ok_or("Nile returned no region for this database")?; + Ok(ProviderConnection { + db_type: "postgres".into(), + host: host_for(region), + port: 5432, + username: user.to_string(), + password: password.to_string(), + database: database.to_string(), + ssl: true, + needs_password: false, + name: format!("Nile · {database}"), + }) +} + +pub async fn build_connection(token: &str, db_ref: &str) -> Result { + let (workspace, database) = db_ref.split_once('/').ok_or("Invalid Nile database reference")?; + let creds = send( + http() + .post(format!( + "{API}/workspaces/{}/databases/{}/credentials", + urlencoding::encode(workspace), + urlencoding::encode(database) + )) + .bearer_auth(token), + "creating database credentials", + ) + .await?; + connection_from_credentials(&creds, database) +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn databases_and_credentials_map_to_the_region_host() { + let rows = parse_databases("ws", &json!([{ "id": "d1", "name": "default", "region": "AWS_US_WEST_2", "status": "READY" }])); + assert_eq!(rows[0].db_ref, "ws/default"); + assert_eq!(rows[0].host.as_deref(), Some("us-west-2.db.thenile.dev")); + let creds = json!({ "id": "0190-cred", "password": "pw", "database": { "region": "AWS_US_WEST_2" } }); + let c = connection_from_credentials(&creds, "default").unwrap(); + assert_eq!((c.username.as_str(), c.host.as_str(), c.port), ("0190-cred", "us-west-2.db.thenile.dev", 5432)); + assert!(connection_from_credentials(&json!({ "id": "x" }), "d").is_err()); + } + + #[test] + fn regions_map_to_the_cli_hosts() { + assert_eq!(host_for("AWS_US_WEST_2"), "us-west-2.db.thenile.dev"); + assert_eq!(host_for("AZURE_EASTUS"), "db.eastus.azure.thenile.dev"); + } +} diff --git a/src-tauri/src/providers/planetscale.rs b/src-tauri/src/providers/planetscale.rs index 3a7fe8b2..d7cc8248 100644 --- a/src-tauri/src/providers/planetscale.rs +++ b/src-tauri/src/providers/planetscale.rs @@ -14,13 +14,18 @@ pub const OAUTH: OAuthConfig = OAuthConfig { auth_url: "https://auth.planetscale.com/oauth/authorize", // token_url is used by the proxy, not the app (the app posts to TOKEN_PROXY). token_url: "https://auth.planetscale.com/oauth/token", - // Exact scope set PlanetScale's own CLI uses (proven valid). NOTE the real - // scope strings differ from the dashboard checkbox labels: it's - // `read_organization` (singular), and password creation is covered by - // `write_databases` (there is no `manage_passwords` OAuth scope). Sent - // space-separated, WITHOUT PKCE (see Provider::uses_pkce). Each must be - // enabled on the OAuth app. - scopes: "read_databases write_databases read_user read_organization", + // Two organization scopes, and they are different permissions: + // `read_organizations` (user access) lists the orgs a user belongs to, which + // is the first call the picker makes - without it `/organizations` answers + // 403 "User does not have permission". `read_organization` (org access) + // reads a single org. Connecting creates a branch password, + // and PlanetScale's API reference requires `manage_passwords` for that, plus + // `manage_production_branch_passwords` when the branch is production (the + // default branch usually is). `write_databases` alone signs in and lists + // fine, then fails on the first connect. Sent space-separated, WITHOUT PKCE + // (see Provider::uses_pkce). Every scope here must also be ticked on the + // OAuth app, or the authorize page rejects the request. + scopes: "read_user read_organizations read_organization read_databases write_databases manage_passwords manage_production_branch_passwords", }; const API: &str = "https://api.planetscale.com/v1"; @@ -33,12 +38,27 @@ async fn get(token: &str, path: &str) -> Result { .await .map_err(|e| format!("PlanetScale request failed: {e}"))?; let status = resp.status().as_u16(); - let body: Value = resp - .json() + // Before reading the body: a 401 may not be JSON, and the caller refreshes on it. + if status == 401 { + return Err(super::UNAUTHORIZED.into()); + } + let text = resp + .text() .await - .map_err(|e| format!("PlanetScale: bad JSON: {e}"))?; + .map_err(|e| format!("PlanetScale read failed: {e}"))?; + let body: Value = serde_json::from_str(&text).unwrap_or(Value::Null); if status != 200 { - return Err(format!("PlanetScale API error ({status})")); + // Say which call and what PlanetScale said. "PlanetScale API error (403)" + // alone gave no way to tell a missing scope from an org that was never + // granted to this app. + let msg = body["message"] + .as_str() + .map(String::from) + .unwrap_or_else(|| text.chars().take(160).collect()); + return Err(format!("PlanetScale API error ({status}) on {path}: {msg}")); + } + if body.is_null() { + return Err(format!("PlanetScale returned a non-JSON response for {path}")); } Ok(body) } @@ -47,10 +67,32 @@ async fn get(token: &str, path: &str) -> Result { /// second lookup. pub async fn list_databases(token: &str) -> Result, String> { let orgs = get(token, "/organizations").await?; + // Every org's database list in flight at once rather than one after another. + let org_names: Vec<&str> = orgs["data"] + .as_array() + .into_iter() + .flatten() + .filter_map(|org| org["name"].as_str()) + .collect(); + let pages = futures::future::join_all(org_names.iter().map(|org_name| { + let path = format!("/organizations/{org_name}/databases"); + async move { get(token, &path).await } + })) + .await; + // `/organizations` lists every org the user belongs to, but the token only + // covers the ones picked on the consent screen: the rest answer 403. + let pages = super::merge_partial( + org_names.iter().zip(pages).map(|(org, r)| r.map(|b| (*org, b))).collect(), + "Stroke may not have been granted this organization: sign out, sign in again, and select it on PlanetScale's consent screen.", + )?; + Ok(parse_databases(&pages)) +} + +/// `/organizations/{org}/databases` pages → picker rows. `db_ref` is +/// "{org}/{database}" so build_connection can act without a second lookup. +fn parse_databases(pages: &[(&str, Value)]) -> Vec { let mut out = Vec::new(); - for org in orgs["data"].as_array().into_iter().flatten() { - let Some(org_name) = org["name"].as_str() else { continue }; - let dbs = get(token, &format!("/organizations/{org_name}/databases")).await?; + for (org_name, dbs) in pages { for db in dbs["data"].as_array().into_iter().flatten() { let Some(name) = db["name"].as_str() else { continue }; out.push(ProviderDatabase { @@ -62,7 +104,7 @@ pub async fn list_databases(token: &str) -> Result, String }); } } - Ok(out) + out } pub async fn build_connection(token: &str, db_ref: &str) -> Result { @@ -84,6 +126,10 @@ pub async fn build_connection(token: &str, db_ref: &str) -> Result Result Result { let username = body["username"].as_str().ok_or("PlanetScale: missing username")?; let password = body["plain_text"].as_str().ok_or("PlanetScale: missing password")?; let host = body["access_host_url"] @@ -111,3 +164,32 @@ pub async fn build_connection(token: &str, db_ref: &str) -> Result Result { .await .map_err(|e| format!("Prisma request failed: {e}"))?; let status = resp.status().as_u16(); + // Before reading the body: a 401 may not be JSON, and the caller refreshes on it. + if status == 401 { + return Err(super::UNAUTHORIZED.into()); + } let body: Value = resp .json() .await @@ -42,18 +46,13 @@ async fn get(token: &str, path: &str) -> Result { Ok(body) } -/// Format a Management-API error. A 401 almost always means the stored token -/// expired - tell the user to reconnect rather than showing "request failed". +/// Format a Management-API error. A 401 never gets here: `get`/`post` return +/// `UNAUTHORIZED` for it so the command layer can refresh and retry. fn api_error(status: u16, body: &Value) -> String { let msg = body["message"] .as_str() .or_else(|| body["error"].as_str()) .unwrap_or("request failed"); - if status == 401 { - return "Your Prisma session has expired. Click \"Sign out\" and sign in \ - again to reconnect." - .to_string(); - } format!("Prisma API error ({status}): {msg}") } @@ -66,6 +65,10 @@ async fn post(token: &str, path: &str, body: serde_json::Value) -> Result Option { } pub async fn build_connection(token: &str, project_id: &str) -> Result { - let proj = get(token, &format!("/projects/{project_id}")).await?; + // The project (for its name) and its databases are independent reads: both + // at once, so a connect is two round trips to Prisma's API instead of three. + let (proj_path, dbs_path) = (format!("/projects/{project_id}"), format!("/projects/{project_id}/databases")); + let (proj, dbs) = tokio::join!(get(token, &proj_path), get(token, &dbs_path)); + let (proj, dbs) = (proj?, dbs?); let name = proj["data"]["name"] .as_str() .or_else(|| proj["name"].as_str()) @@ -173,7 +180,6 @@ pub async fn build_connection(token: &str, project_id: &str) -> Result Result Result { + let conn = |host, port, username, password, database: String| ProviderConnection { + db_type: "postgres".into(), + host, + port, + username, + password, + database: if database.is_empty() { "postgres".into() } else { database }, + ssl: true, + needs_password: false, + name: format!("Prisma · {name}"), + }; + if let Some((host, user, pass, dbn)) = find_direct(created) { + return Ok(conn(host, 5432, user, pass, dbn)); + } + if let Some(uri) = find_pg_uri(created) { let (host, port, username, password, database) = super::neon::parse_pg_uri(&uri)?; - return Ok(ProviderConnection { - db_type: "postgres".into(), - host, - port, - username, - password, - database: if database.is_empty() { "postgres".into() } else { database }, - ssl: true, - needs_password: false, - name: format!("Prisma · {name}"), - }); + return Ok(conn(host, port, username, password, database)); } - - let shape: String = serde_json::to_string(&created).unwrap_or_default().chars().take(1600).collect(); + let shape: String = serde_json::to_string(created).unwrap_or_default().chars().take(1600).collect(); Err(format!( "Prisma created a connection but returned no direct credentials we could parse. \ Shape: {shape}" )) } + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn credentials_object_is_read_wherever_prisma_nests_it() { + let created = json!({ "data": { "endpoints": { "direct": { + "ppgDirectConnection": { "host": "db.prisma.io", "user": "u1", "pass": "p1" } + } } } }); + let c = connection_from_created(&created, "pulse").unwrap(); + assert_eq!((c.host.as_str(), c.username.as_str(), c.password.as_str()), ("db.prisma.io", "u1", "p1")); + assert_eq!(c.name, "Prisma · pulse"); + } + + #[test] + fn falls_back_to_a_postgres_uri_anywhere_in_the_response() { + let created = json!({ "data": { "connectionString": "postgres://u2:p%212@db.prisma.io:5432/postgres?sslmode=require" } }); + let c = connection_from_created(&created, "x").unwrap(); + assert_eq!((c.username.as_str(), c.password.as_str()), ("u2", "p!2")); + assert!(connection_from_created(&json!({ "data": {} }), "x").unwrap_err().contains("no direct credentials")); + } +} diff --git a/src-tauri/src/providers/railway.rs b/src-tauri/src/providers/railway.rs new file mode 100644 index 00000000..802a725b --- /dev/null +++ b/src-tauri/src/providers/railway.rs @@ -0,0 +1,244 @@ +/*! + * Railway adapter - Postgres, MySQL and Redis services. Sign in with "Login with + * Railway" (OAuth 2.0 + OIDC), list the database services across the projects + * the user shared on the consent screen, and connect through the service's public + * TCP proxy URL. + * + * Stroke registers as a Railway *native* app: a public client with PKCE and no + * secret (the discovery document lists `none` among the token endpoint auth + * methods), so the token exchange goes straight to Railway, like Neon. Access + * tokens last an hour; `offline_access` (with `prompt=consent`) returns a refresh + * token. + */ + +use super::{http, OAuthConfig, ProviderConnection, ProviderDatabase}; +use serde_json::{json, Value}; + +pub const OAUTH: OAuthConfig = OAuthConfig { + // The native OAuth app registered in Railway (workspace Developer settings), + // redirect URI http://localhost:8989/oauth/callback. Empty until it exists: + // sign-in then says so instead of opening a page Railway would reject. + client_id: "", + auth_url: "https://backboard.railway.com/oauth/auth", + token_url: "https://backboard.railway.com/oauth/token", + // workspace:viewer lists workspaces; project:member is what lets the token + // read a service's variables, which is where the connection URL lives. + scopes: "openid email profile offline_access workspace:viewer project:member", +}; + +const API: &str = "https://backboard.railway.com/graphql/v2"; + +async fn gql(token: &str, query: &str, variables: Value) -> Result { + let resp = http() + .post(API) + .bearer_auth(token) + .json(&json!({ "query": query, "variables": variables })) + .send() + .await + .map_err(|e| format!("Railway request failed: {e}"))?; + let status = resp.status().as_u16(); + if status == 401 { + return Err(super::UNAUTHORIZED.into()); + } + let body: Value = resp + .json() + .await + .map_err(|e| format!("Railway returned bad JSON: {e}"))?; + if let Some(err) = body["errors"].as_array().and_then(|e| e.first()) { + let msg = err["message"].as_str().unwrap_or("request failed"); + // GraphQL reports an expired or revoked token as an error, not a 401. + if msg.to_ascii_lowercase().contains("not authorized") && body["data"].is_null() { + return Err(super::UNAUTHORIZED.into()); + } + return Err(format!("Railway API error: {msg}")); + } + if !(200..300).contains(&status) { + return Err(format!("Railway API error ({status})")); + } + Ok(body["data"].clone()) +} + +/// Which engine a service runs, from its deploy image. Railway's database +/// templates are images (`ghcr.io/railwayapp-templates/postgres-ssl:17`, +/// `mysql:9`, `redis:8`); anything else (an app built from a repo) is skipped. +fn engine_of(image: &str) -> Option<&'static str> { + let name = image.rsplit('/').next().unwrap_or(image).to_ascii_lowercase(); + let name = name.split(':').next().unwrap_or(""); + if name.starts_with("postgres") || name.starts_with("timescale") || name.starts_with("pgvector") { + Some("postgres") + } else if name.starts_with("mysql") || name.starts_with("mariadb") { + Some("mysql") + } else if name.starts_with("redis") || name.starts_with("valkey") { + Some("redis") + } else { + None + } +} + +const PROJECTS_QUERY: &str = r#" +query StrokeProjects($workspaceId: String) { + projects(first: 100, workspaceId: $workspaceId) { + edges { node { + id name + environments { edges { node { + id name + serviceInstances { edges { node { serviceId serviceName source { image } } } } + } } } + } } + } +}"#; + +pub async fn list_databases(token: &str) -> Result, String> { + let me = gql(token, "query { me { workspaces { id name } } }", json!({})).await?; + let ids: Vec = me["me"]["workspaces"] + .as_array() + .into_iter() + .flatten() + .filter_map(|w| w["id"].as_str().map(String::from)) + .collect(); + // Every workspace's projects at once. A workspace the token wasn't granted + // errors on its own and is skipped rather than failing the whole list. + let pages = futures::future::join_all( + ids.iter().map(|id| gql(token, PROJECTS_QUERY, json!({ "workspaceId": id }))), + ) + .await; + let pages = super::merge_partial(pages, "")?; + Ok(parse_projects(&pages)) +} + +/// `projects` pages → one row per database service instance. Services whose +/// image isn't a database template are skipped; a service deployed in several +/// environments gets the environment in its name. +fn parse_projects(pages: &[Value]) -> Vec { + let mut out = Vec::new(); + let mut seen = std::collections::HashSet::new(); + for page in pages { + for p in page["projects"]["edges"].as_array().into_iter().flatten() { + let p = &p["node"]; + let (Some(pid), Some(pname)) = (p["id"].as_str(), p["name"].as_str()) else { continue }; + let envs = p["environments"]["edges"].as_array().cloned().unwrap_or_default(); + let multi_env = envs.len() > 1; + for e in &envs { + let e = &e["node"]; + let (Some(eid), ename) = (e["id"].as_str(), e["name"].as_str().unwrap_or("")) else { continue }; + for si in e["serviceInstances"]["edges"].as_array().into_iter().flatten() { + let si = &si["node"]; + let Some(sid) = si["serviceId"].as_str() else { continue }; + let Some(engine) = si["source"]["image"].as_str().and_then(engine_of) else { continue }; + if !seen.insert(format!("{eid}/{sid}")) { + continue; + } + let sname = si["serviceName"].as_str().unwrap_or(sid); + out.push(ProviderDatabase { + // Everything build_connection needs, so it can skip a lookup. + db_ref: json!({ "p": pid, "e": eid, "s": sid, "k": engine, "n": format!("{pname} / {sname}") }) + .to_string(), + name: if multi_env { format!("{pname} / {sname} ({ename})") } else { format!("{pname} / {sname}") }, + region: None, + kind: Some(match engine { "postgres" => "Postgres", "mysql" => "MySQL", _ => "Redis" }.into()), + host: None, + }); + } + } + } + } + out +} + +pub async fn build_connection(token: &str, db_ref: &str) -> Result { + let r: Value = serde_json::from_str(db_ref).map_err(|_| "Invalid Railway service reference")?; + let field = |k: &str| r[k].as_str().unwrap_or_default().to_string(); + let (engine, name) = (field("k"), field("n")); + let vars = gql( + token, + "query StrokeVars($p: String!, $e: String!, $s: String!) { variables(projectId: $p, environmentId: $e, serviceId: $s) }", + json!({ "p": field("p"), "e": field("e"), "s": field("s") }), + ) + .await?; + connection_from_vars(&engine, &name, &vars["variables"]) +} + +/// A service's variables → a connection, through its PUBLIC url: the plain one +/// points at `*.railway.internal`, which only resolves inside Railway's network. +fn connection_from_vars(engine: &str, name: &str, vars: &Value) -> Result { + let keys: &[&str] = match engine { + "postgres" => &["DATABASE_PUBLIC_URL"], + "mysql" => &["MYSQL_PUBLIC_URL", "DATABASE_PUBLIC_URL"], + _ => &["REDIS_PUBLIC_URL"], + }; + let raw = keys + .iter() + .find_map(|k| vars[*k].as_str().filter(|v| !v.is_empty())) + .ok_or_else(|| { + format!( + "{name} has no public URL. Turn on its TCP proxy in the service's Railway settings (Networking), then try again." + ) + })?; + let url = reqwest::Url::parse(raw).map_err(|e| format!("Railway returned an unreadable URL for {name}: {e}"))?; + let decode = |s: &str| urlencoding::decode(s).map(|c| c.into_owned()).unwrap_or_else(|_| s.to_string()); + + let default_port = match engine { "postgres" => 5432, "mysql" => 3306, _ => 6379 }; + Ok(ProviderConnection { + db_type: engine.to_string(), + host: url.host_str().ok_or_else(|| format!("No host in {name}'s URL"))?.to_string(), + port: url.port().unwrap_or(default_port), + username: decode(url.username()), + password: decode(url.password().unwrap_or_default()), + database: decode(url.path().trim_start_matches('/')), + // Railway's TCP proxy is plain TCP; the Postgres template serves TLS + // itself (postgres-ssl) but does not require it, and MySQL/Redis there + // are not TLS. + ssl: false, + needs_password: false, + name: format!("Railway · {name}"), + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn lists_database_services_only_and_names_environments() { + let page = json!({ "projects": { "edges": [ { "node": { + "id": "p1", "name": "shop", + "environments": { "edges": [ + { "node": { "id": "e1", "name": "production", "serviceInstances": { "edges": [ + { "node": { "serviceId": "s1", "serviceName": "Postgres", "source": { "image": "ghcr.io/railwayapp-templates/postgres-ssl:17" } } }, + { "node": { "serviceId": "s2", "serviceName": "web", "source": { "repo": "me/web" } } } + ] } } }, + { "node": { "id": "e2", "name": "staging", "serviceInstances": { "edges": [ + { "node": { "serviceId": "s1", "serviceName": "Postgres", "source": { "image": "postgres-ssl:17" } } } + ] } } } + ] } + } } ] } }); + let rows = parse_projects(&[page]); + assert_eq!(rows.len(), 2, "the web app is not a database"); + assert_eq!(rows[0].name, "shop / Postgres (production)"); + let r: Value = serde_json::from_str(&rows[1].db_ref).unwrap(); + assert_eq!((r["e"].as_str(), r["k"].as_str()), (Some("e2"), Some("postgres"))); + } + + #[test] + fn connects_through_the_public_url_and_decodes_credentials() { + let vars = json!({ + "DATABASE_URL": "postgresql://postgres:x@postgres.railway.internal:5432/railway", + "DATABASE_PUBLIC_URL": "postgresql://postgres:p%40ss@shortline.proxy.rlwy.net:41234/railway" + }); + let c = connection_from_vars("postgres", "shop / Postgres", &vars).unwrap(); + assert_eq!((c.host.as_str(), c.port), ("shortline.proxy.rlwy.net", 41234)); + assert_eq!((c.username.as_str(), c.password.as_str(), c.database.as_str()), ("postgres", "p@ss", "railway")); + let redis = json!({ "REDIS_PUBLIC_URL": "redis://default:pw@x.proxy.rlwy.net:6380" }); + assert_eq!(connection_from_vars("redis", "cache", &redis).unwrap().port, 6380); + let private_only = json!({ "DATABASE_URL": "postgresql://u:p@postgres.railway.internal:5432/db" }); + assert!(connection_from_vars("postgres", "x", &private_only).unwrap_err().contains("TCP proxy")); + } + + #[test] + fn engines_come_from_the_template_image() { + assert_eq!(engine_of("ghcr.io/railwayapp-templates/postgres-ssl:17"), Some("postgres")); + assert_eq!(engine_of("mysql:9"), Some("mysql")); + assert_eq!(engine_of("bitnami/redis:7.2"), Some("redis")); + assert_eq!(engine_of("ghcr.io/me/web-app:latest"), None); + } +} diff --git a/src-tauri/src/providers/supabase.rs b/src-tauri/src/providers/supabase.rs index 0f78c059..21ae6fcc 100644 --- a/src-tauri/src/providers/supabase.rs +++ b/src-tauri/src/providers/supabase.rs @@ -28,6 +28,10 @@ async fn get(token: &str, path: &str) -> Result { .await .map_err(|e| format!("Supabase request failed: {e}"))?; let status = resp.status().as_u16(); + // Before reading the body: a 401 may not be JSON, and the caller refreshes on it. + if status == 401 { + return Err(super::UNAUTHORIZED.into()); + } let body: Value = resp .json() .await @@ -60,77 +64,95 @@ pub async fn list_databases(token: &str) -> Result, String } pub async fn build_connection(token: &str, project_ref: &str) -> Result { - let proj = get(token, &format!("/projects/{project_ref}")).await?; - let name = proj["name"].as_str().unwrap_or(project_ref); + // The project and its pooler config are independent reads: fetch both at + // once instead of one after the other (one round trip to Supabase's API + // instead of two). A paused project answers the pooler call with `200 []`, + // so the status check below still decides what to say. + let (proj_path, pooler_path) = ( + format!("/projects/{project_ref}"), + format!("/projects/{project_ref}/config/database/pooler"), + ); + let (proj, pooler) = tokio::join!(get(token, &proj_path), get(token, &pooler_path)); + let proj = proj?; + let name = proj["name"].as_str().unwrap_or(project_ref).to_string(); + status_error(&proj, &name)?; - // Say what is actually wrong before asking for a pooler that cannot exist. - // - // A project has to be running to have a Supavisor tenant, and a free-tier - // project pauses itself after a week of inactivity. The pooler endpoint then - // answers `200 []` rather than an error, so the failure surfaced as - // "Supabase didn't return a pooler host for this project. Pooler config: []" - // - a dump of an empty array, which tells the user nothing they can act on - // when the real answer is "your project is asleep, go and wake it". - if let Some(status) = proj["status"].as_str() { - match status { - "ACTIVE_HEALTHY" | "ACTIVE_UNHEALTHY" | "UNKNOWN" => {} - "INACTIVE" | "PAUSING" | "PAUSE_FAILED" => { - return Err(format!( - "{name} is paused, so it has no pooler to connect to. Restore it from your \ - Supabase dashboard, wait for it to come up, then try again." - )) - } - "COMING_UP" | "RESTORING" | "RESTARTING" | "RESIZING" | "UPGRADING" => { - return Err(format!( - "{name} is still starting up ({status}). Give it a moment and try again." - )) - } - "INIT_FAILED" | "RESTORE_FAILED" => { - return Err(format!( - "{name} failed to start ({status}). Check the project in your Supabase \ - dashboard - there is nothing to connect to until it comes up." - )) - } - "GOING_DOWN" | "REMOVED" => { - return Err(format!("{name} is being removed ({status}).")) - } - other => { - return Err(format!( - "{name} is not running ({other}). Check the project in your Supabase dashboard." - )) - } + let pooler = pooler.map_err(|e| { + if e.contains("scope") || e.contains("403") { + "Your Supabase OAuth app is missing the \"Database Pooling Config: Read\" scope, \ + which Stroke needs to find the correct pooler host. Fix: in the Supabase \ + dashboard open your OAuth app → Scopes → enable Database Pooling Config (Read) \ + (enabling all Read scopes is fine), save, then Disconnect here and sign in again \ + to re-authorize." + .to_string() + } else { + format!("Couldn't fetch Supabase pooler config (needed for a working host): {e}") } - } + })?; + let (host, user) = pooler_target(&proj, &pooler, project_ref, &name)?; - // Use the Supavisor pooler, NOT the direct host. The direct connection - // `db..supabase.co:5432` is IPv6-only, so it fails with "network - // unreachable" on the many networks without IPv6. The shared pooler is - // IPv4-compatible on every tier. Its username embeds the project ref - // (`postgres.`), and session mode (port 5432) behaves like a normal - // Postgres connection - the right choice for a GUI client. - // - // Ask the Management API for the exact pooler host - never guess the region - // prefix (aws-0 vs aws-1), since a wrong host connects to a pooler node that - // doesn't host this project's tenant → "tenant/user … not found". - let pooler = get(token, &format!("/projects/{project_ref}/config/database/pooler")) - .await - .map_err(|e| { - if e.contains("scope") || e.contains("403") { - "Your Supabase OAuth app is missing the \"Database Pooling Config: Read\" scope, \ - which Stroke needs to find the correct pooler host. Fix: in the Supabase \ - dashboard open your OAuth app → Scopes → enable Database Pooling Config (Read) \ - (enabling all Read scopes is fine), save, then Disconnect here and sign in again \ - to re-authorize." - .to_string() - } else { - format!("Couldn't fetch Supabase pooler config (needed for a working host): {e}") - } - })?; - let obj = pooler.as_array().and_then(|a| a.first()).unwrap_or(&pooler); + Ok(ProviderConnection { + db_type: "postgres".into(), + host, + port: 5432, // session mode - normal Postgres semantics, best for a GUI + username: user, + password: String::new(), + database: "postgres".into(), + ssl: true, + needs_password: true, // Supabase never returns the DB password + name: format!("Supabase · {name}"), + }) +} +/// Say what is actually wrong before asking for a pooler that cannot exist. +/// +/// A project has to be running to have a Supavisor tenant, and a free-tier +/// project pauses itself after a week of inactivity. The pooler endpoint then +/// answers `200 []` rather than an error, so the failure surfaced as +/// "Supabase didn't return a pooler host for this project. Pooler config: []" +/// - a dump of an empty array, which tells the user nothing they can act on +/// when the real answer is "your project is asleep, go and wake it". +fn status_error(proj: &Value, name: &str) -> Result<(), String> { + let Some(status) = proj["status"].as_str() else { return Ok(()) }; + match status { + "ACTIVE_HEALTHY" | "ACTIVE_UNHEALTHY" | "UNKNOWN" => Ok(()), + "INACTIVE" | "PAUSING" | "PAUSE_FAILED" => Err(format!( + "{name} is paused, so it has no pooler to connect to. Restore it from your \ + Supabase dashboard, wait for it to come up, then try again." + )), + "COMING_UP" | "RESTORING" | "RESTARTING" | "RESIZING" | "UPGRADING" => Err(format!( + "{name} is still starting up ({status}). Give it a moment and try again." + )), + "INIT_FAILED" | "RESTORE_FAILED" => Err(format!( + "{name} failed to start ({status}). Check the project in your Supabase \ + dashboard - there is nothing to connect to until it comes up." + )), + "GOING_DOWN" | "REMOVED" => Err(format!("{name} is being removed ({status}).")), + other => Err(format!( + "{name} is not running ({other}). Check the project in your Supabase dashboard." + )), + } +} + +/// Which host and user to connect with. +/// +/// The Supavisor pooler, NOT the direct host: `db..supabase.co:5432` is +/// IPv6-only, so it fails with "network unreachable" on the many networks +/// without IPv6, while the shared pooler is IPv4-compatible on every tier. Its +/// username embeds the project ref (`postgres.`), and session mode (5432) +/// behaves like a normal Postgres connection. The host comes from the pooler +/// config, never guessed: a wrong region prefix (aws-0 vs aws-1) reaches a +/// pooler node without this tenant ("tenant/user … not found"). +/// +/// Running but with no Supavisor config: fall back to the direct host rather +/// than refuse outright. It is IPv6-only unless the project has the IPv4 +/// add-on, so it can still fail on an IPv4-only network, but a connection that +/// might work beats an error that certainly doesn't. +fn pooler_target(proj: &Value, pooler: &Value, project_ref: &str, name: &str) -> Result<(String, String), String> { + let obj = pooler.as_array().and_then(|a| a.first()).unwrap_or(pooler); let mut host = String::new(); let mut user = format!("postgres.{project_ref}"); - // The connection_string carries the exact pooler host + tenant user - parse it. + // The connection_string carries the exact pooler host + tenant user. if let Some(cs) = obj["connection_string"].as_str().or_else(|| obj["connectionString"].as_str()) { if let Ok((h, _p, u, _pw, _d)) = super::neon::parse_pg_uri(cs) { if !h.is_empty() { @@ -146,19 +168,9 @@ pub async fn build_connection(token: &str, project_ref: &str) -> Result.supabase.co` is IPv6-only unless the project has the IPv4 - // add-on, so this can still fail on an IPv4-only network; the message below - // names that, because "network unreachable" from a Postgres driver does not. if host.is_empty() { if let Some(direct) = proj["database"]["host"].as_str().filter(|h| !h.is_empty()) { host = direct.to_string(); @@ -171,16 +183,37 @@ pub async fn build_connection(token: &str, project_ref: &str) -> Result Result { + // One retry when the request never got a response (a socket the far end had + // already closed). Safe for the POST too: a request that failed to send + // created nothing. + let retry = req.try_clone(); + let resp = match req.send().await { + Ok(r) => r, + Err(e) if e.is_connect() || e.is_request() => match retry { + Some(r) => r + .send() + .await + .map_err(|e| format!("TiDB Cloud request failed ({what}): {}", super::describe(&e)))?, + None => return Err(format!("TiDB Cloud request failed ({what}): {}", super::describe(&e))), + }, + Err(e) => return Err(format!("TiDB Cloud request failed ({what}): {}", super::describe(&e))), + }; + let status = resp.status().as_u16(); + // Before reading the body: a 401 may not be JSON, and the caller refreshes on it. + if status == 401 { + return Err(super::UNAUTHORIZED.into()); + } + let text = resp + .text() + .await + .map_err(|e| format!("TiDB Cloud read failed: {e}"))?; + let body: Value = serde_json::from_str(&text).map_err(|_| { + let snippet: String = text.chars().take(200).collect(); + format!("TiDB Cloud non-JSON response for {what} (HTTP {status}): {snippet}") + })?; + if !(200..300).contains(&status) { + let msg = body["message"] + .as_str() + .or_else(|| body["error"]["message"].as_str()) + .unwrap_or("request failed"); + return Err(format!("TiDB Cloud API error ({status}) on {what}: {msg}")); + } + Ok(body) +} + +async fn get(token: &str, url: String, what: &str) -> Result { + send(http().get(url).bearer_auth(token), what).await +} + +/// Each Starter/Essential cluster is one connectable database. +pub async fn list_databases(token: &str) -> Result, String> { + let mut out = Vec::new(); + let mut page_token = String::new(); + for _ in 0..MAX_PAGES { + let mut url = format!("{SERVERLESS_API}/clusters?pageSize=100"); + if !page_token.is_empty() { + url.push_str(&format!("&pageToken={}", urlencoding::encode(&page_token))); + } + let body = get(token, url, "listing clusters").await?; + out.extend(parse_clusters(&body)); + match body["nextPageToken"].as_str() { + Some(next) if !next.is_empty() => page_token = next.to_string(), + _ => break, + } + } + Ok(out) +} + +/// One `clusters` page → picker rows. +fn parse_clusters(body: &Value) -> Vec { + body["clusters"] + .as_array() + .into_iter() + .flatten() + .filter_map(|c| { + let id = c["clusterId"].as_str()?; + Some(ProviderDatabase { + db_ref: id.to_string(), + name: c["displayName"].as_str().unwrap_or(id).to_string(), + region: c["region"]["displayName"] + .as_str() + .or_else(|| c["region"]["regionId"].as_str()) + .map(String::from), + kind: Some("Cluster".into()), + host: c["endpoints"]["public"]["host"].as_str().map(String::from), + }) + }) + .collect() +} + +/// Whether a cluster can be connected to right now, and its public endpoint. +/// Errors are worded for the picker's toast ("is paused", "still starting"). +fn endpoint_of(c: &Value, name: &str) -> Result<(String, u16), String> { + match c["state"].as_str().unwrap_or("ACTIVE") { + "ACTIVE" => {} + "PAUSED" | "PAUSING" | "INACTIVE" => { + return Err(format!("{name} is paused. Resume it from the TiDB Cloud console, then try again.")) + } + s @ ("CREATING" | "RESUMING" | "RESTORING" | "UPGRADING" | "MODIFYING" | "MAINTENANCE" | "IMPORTING") => { + return Err(format!("{name} is still starting up ({s}). Give it a moment and try again.")) + } + s => return Err(format!("{name} is not available to connect to ({s}).")), + } + let public = &c["endpoints"]["public"]; + if public["disabled"].as_bool() == Some(true) { + return Err(format!( + "{name} has its public endpoint turned off. Enable it in the TiDB Cloud console to connect from Stroke." + )); + } + let host = public["host"] + .as_str() + .ok_or_else(|| format!("TiDB Cloud returned no public host for {name}"))? + .to_string(); + let port = public["port"].as_u64().and_then(|p| u16::try_from(p).ok()).unwrap_or(4000); + Ok((host, port)) +} + +/// The SQL user name the server created: with `autoPrefix` it is +/// `{userPrefix}.{userName}`, echoed back in the response when present. +fn created_user(created: &Value, prefix: &str, user: &str) -> String { + created["userName"] + .as_str() + .filter(|u| u.contains('.')) + .map(String::from) + .unwrap_or_else(|| if prefix.is_empty() { user.to_string() } else { format!("{prefix}.{user}") }) +} + +/// Alphanumeric only: TiDB accepts symbols in passwords, but a URL-shaped +/// character in one has to be escaped everywhere the credential travels. +fn alnum(n: usize) -> String { + super::random_base64url(n * 2) + .chars() + .filter(char::is_ascii_alphanumeric) + .take(n) + .collect() +} + +pub async fn build_connection(token: &str, cluster_id: &str) -> Result { + let c = get( + token, + format!("{SERVERLESS_API}/clusters/{}", urlencoding::encode(cluster_id)), + "reading the cluster", + ) + .await?; + let name = c["displayName"].as_str().unwrap_or(cluster_id).to_string(); + + let (host, port) = endpoint_of(&c, &name)?; + let prefix = c["userPrefix"].as_str().unwrap_or_default(); + + // A fresh admin user per connect, like PlanetScale's minted password. The + // frontend reuses a saved connection's credentials when it has them, so + // this only runs for a cluster Stroke hasn't connected to before. + let user = format!("stroke_{}", alnum(6).to_lowercase()); + let password = alnum(24); + let created = send( + http() + .post(format!("{IAM_API}/clusters/{}/sqlUsers", urlencoding::encode(cluster_id))) + .bearer_auth(token) + .json(&json!({ + "userName": user, + "password": password, + "builtinRole": "role_admin", + "authMethod": "mysql_native_password", + "autoPrefix": true, + })), + "creating a SQL user", + ) + .await?; + let username = created_user(&created, prefix, &user); + + Ok(ProviderConnection { + db_type: "mysql".into(), + host, + port, + username, + password, + // No default schema: the sidebar lists every database on the cluster. + database: String::new(), + ssl: true, + needs_password: false, + name: format!("TiDB · {name}"), + }) +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn clusters_list_with_region_and_host() { + let body = json!({ "clusters": [ + { "clusterId": "1094", "displayName": "test", "region": { "regionId": "aws-ap-northeast-1", "displayName": "Tokyo (ap-northeast-1)" }, + "endpoints": { "public": { "host": "gateway01.ap-northeast-1.prod.aws.tidbcloud.com", "port": 4000 } } }, + { "displayName": "no id, skipped" } + ]}); + let rows = parse_clusters(&body); + assert_eq!(rows.len(), 1); + assert_eq!(rows[0].region.as_deref(), Some("Tokyo (ap-northeast-1)")); + } + + #[test] + fn cluster_state_and_endpoint_gate_the_connect() { + let ok = json!({ "state": "ACTIVE", "endpoints": { "public": { "host": "h", "port": 4000 } } }); + assert_eq!(endpoint_of(&ok, "t").unwrap(), ("h".into(), 4000)); + assert!(endpoint_of(&json!({ "state": "PAUSED" }), "t").unwrap_err().contains("is paused")); + assert!(endpoint_of(&json!({ "state": "RESUMING" }), "t").unwrap_err().contains("still starting")); + let off = json!({ "state": "ACTIVE", "endpoints": { "public": { "disabled": true } } }); + assert!(endpoint_of(&off, "t").unwrap_err().contains("public endpoint")); + } + + #[test] + fn sql_user_gets_the_cluster_prefix() { + assert_eq!(created_user(&json!({ "userName": "3Xq.stroke_ab" }), "3Xq", "stroke_ab"), "3Xq.stroke_ab"); + assert_eq!(created_user(&json!({}), "3Xq", "stroke_ab"), "3Xq.stroke_ab"); + assert_eq!(created_user(&json!({}), "", "stroke_ab"), "stroke_ab"); + } +} diff --git a/src-tauri/src/providers/turso.rs b/src-tauri/src/providers/turso.rs new file mode 100644 index 00000000..eee49afe --- /dev/null +++ b/src-tauri/src/providers/turso.rs @@ -0,0 +1,183 @@ +/*! + * Turso adapter - libSQL. Sign in through the browser the way `turso auth login` + * does, list every database across the user's organizations, and create a + * database token on connect. + * + * Turso has no OAuth app registration. Its CLI opens + * `https://api.turso.tech/?port={port}&redirect=true&type=cli&state={state}`, + * and after sign-in the browser is sent to `http://localhost:{port}/?jwt=…&state=…` + * with a platform API token (valid about a week, no refresh token). This reuses + * that page, the same way Neon reuses neonctl's client. + */ + +use super::{http, OAuthConfig, ProviderConnection, ProviderDatabase}; +use serde_json::Value; + +pub const API: &str = "https://api.turso.tech"; + +pub const OAUTH: OAuthConfig = OAuthConfig { + // No client: the CLI login page identifies callers with `type=cli`. + client_id: "", + auth_url: API, + token_url: "", + scopes: "", +}; + +/// Cap on database list pages per organization. +const MAX_PAGES: usize = 20; + +async fn send(req: reqwest::RequestBuilder, what: &str) -> Result { + let resp = req + .send() + .await + .map_err(|e| format!("Turso request failed ({what}): {}", super::describe(&e)))?; + let status = resp.status().as_u16(); + // Before reading the body: a 401 may not be JSON, and the caller refreshes on it. + if status == 401 { + return Err(super::UNAUTHORIZED.into()); + } + let text = resp + .text() + .await + .map_err(|e| format!("Turso read failed: {e}"))?; + let body: Value = serde_json::from_str(&text).map_err(|_| { + let snippet: String = text.chars().take(200).collect(); + format!("Turso non-JSON response for {what} (HTTP {status}): {snippet}") + })?; + if !(200..300).contains(&status) { + let msg = body["error"].as_str().unwrap_or("request failed"); + return Err(format!("Turso API error ({status}) on {what}: {msg}")); + } + Ok(body) +} + +async fn get(token: &str, path: &str, what: &str) -> Result { + send(http().get(format!("{API}{path}")).bearer_auth(token), what).await +} + +/// Turso's JSON mixes `Name`/`Hostname` with camelCase keys; read either spelling. +fn field<'a>(v: &'a Value, upper: &str, lower: &str) -> Option<&'a str> { + v[upper].as_str().or_else(|| v[lower].as_str()) +} + +/// One `databases` page of an org → picker rows. `db_ref` is "{org}/{name}". +fn parse_page(org: &str, body: &Value) -> Vec { + body["databases"] + .as_array() + .into_iter() + .flatten() + .filter_map(|db| { + let name = field(db, "Name", "name")?; + Some(ProviderDatabase { + db_ref: format!("{org}/{name}"), + name: name.to_string(), + region: field(db, "primaryRegion", "primary_region").map(String::from), + kind: Some("Database".into()), + host: field(db, "Hostname", "hostname").map(String::from), + }) + }) + .collect() +} + +async fn list_org(token: &str, org: &str) -> Result, String> { + let base = format!("/v1/organizations/{}/databases", urlencoding::encode(org)); + let mut out = Vec::new(); + let mut cursor = String::new(); + for _ in 0..MAX_PAGES { + let path = if cursor.is_empty() { + base.clone() + } else { + format!("{base}?cursor={}", urlencoding::encode(&cursor)) + }; + let body = get(token, &path, "listing databases").await?; + out.extend(parse_page(org, &body)); + match body["pagination"]["next"].as_str() { + Some(next) if !next.is_empty() => cursor = next.to_string(), + _ => break, + } + } + Ok(out) +} + +pub async fn list_databases(token: &str) -> Result, String> { + let orgs = get(token, "/v2/organizations", "listing organizations").await?; + let slugs: Vec<&str> = orgs["organizations"] + .as_array() + .into_iter() + .flatten() + .filter_map(|o| o["slug"].as_str()) + .collect(); + // Every organization's list in flight at once, in the orgs' own order. An + // org the token can't read is skipped rather than failing the whole list. + let pages = futures::future::join_all(slugs.iter().map(|slug| list_org(token, slug))).await; + Ok(super::merge_partial(pages, "")?.into_iter().flatten().collect()) +} + +pub async fn build_connection(token: &str, db_ref: &str) -> Result { + let (org, database) = db_ref + .split_once('/') + .ok_or("Invalid Turso database reference")?; + let (org_enc, db_enc) = (urlencoding::encode(org), urlencoding::encode(database)); + + // The hostname lookup and the token mint don't depend on each other, so + // both go out at once: one round trip instead of two. + let db_path = format!("/v1/organizations/{org_enc}/databases/{db_enc}"); + let (db, minted) = tokio::join!( + get(token, &db_path, "reading the database"), + // A full-access token with no expiry: the connection is saved and + // reused, and a token that lapsed in a week would break it with no way + // to renew in place. + send( + http() + .post(format!( + "{API}/v1/organizations/{org_enc}/databases/{db_enc}/auth/tokens?expiration=never&authorization=full-access" + )) + .bearer_auth(token) + .json(&serde_json::json!({})), + "creating a database token", + ), + ); + let (db, minted) = (db?, minted?); + let host = field(&db["database"], "Hostname", "hostname") + .ok_or_else(|| format!("Turso returned no hostname for {database}"))? + .to_string(); + let jwt = minted["jwt"] + .as_str() + .ok_or("Turso did not return a database token")? + .to_string(); + + // libSQL has no host/port/user split. The adapter contract carries the URL in + // `host` and the token in `password`; the frontend maps them to `url` and + // `authToken` for a libsql connection. + Ok(ProviderConnection { + db_type: "libsql".into(), + host: format!("libsql://{host}"), + port: 443, + username: String::new(), + password: jwt, + database: database.to_string(), + ssl: true, + needs_password: false, + name: format!("Turso · {database}"), + }) +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn reads_either_casing_of_turso_fields() { + let body = json!({ "databases": [ + { "Name": "cool", "Hostname": "cool-me.aws-ap-south-1.turso.io", "primaryRegion": "aws-ap-south-1" }, + { "name": "lower", "hostname": "lower-me.turso.io" }, + { "Hostname": "no-name-skipped" } + ]}); + let rows = parse_page("me", &body); + assert_eq!(rows.len(), 2); + assert_eq!(rows[0].db_ref, "me/cool"); + assert_eq!(rows[0].region.as_deref(), Some("aws-ap-south-1")); + assert_eq!(rows[1].host.as_deref(), Some("lower-me.turso.io")); + } +} diff --git a/src-tauri/src/providers/upstash.rs b/src-tauri/src/providers/upstash.rs new file mode 100644 index 00000000..029c9e6e --- /dev/null +++ b/src-tauri/src/providers/upstash.rs @@ -0,0 +1,133 @@ +/*! + * Upstash adapter - Redis. Upstash has no OAuth for third-party apps, so this is + * the first paste-a-token provider: the account email and a Developer API key + * (Console → Account → Management API), stored together as `email:key` and sent + * as HTTP basic auth. List every Redis database, and connect with the + * database's own endpoint and password, which the API returns directly. + */ + +use super::{http, OAuthConfig, ProviderConnection, ProviderDatabase}; +use serde_json::Value; + +/// Unused for a token provider; present so `Provider::oauth` stays total. +pub const OAUTH: OAuthConfig = OAuthConfig { + client_id: "", + auth_url: "", + token_url: "", + scopes: "", +}; + +const API: &str = "https://api.upstash.com/v2"; + +/// The stored credential is `email:api_key`. An email never contains a colon, +/// so the first one is the split. +fn credentials(token: &str) -> Result<(&str, &str), String> { + token + .split_once(':') + .filter(|(e, k)| !e.is_empty() && !k.is_empty()) + .ok_or_else(|| "Upstash needs the account email and a Developer API key.".to_string()) +} + +async fn get(token: &str, path: &str, what: &str) -> Result { + let (email, key) = credentials(token)?; + let resp = http() + .get(format!("{API}{path}")) + .basic_auth(email, Some(key)) + .send() + .await + .map_err(|e| format!("Upstash request failed ({what}): {}", super::describe(&e)))?; + let status = resp.status().as_u16(); + // A wrong email/key pair is a 401: the command layer turns a second one into + // "Not signed in", which puts the token form back. + if status == 401 || status == 403 { + return Err(super::UNAUTHORIZED.into()); + } + let text = resp.text().await.map_err(|e| format!("Upstash read failed: {e}"))?; + if !(200..300).contains(&status) { + let snippet: String = text.chars().take(160).collect(); + return Err(format!("Upstash API error ({status}) on {what}: {snippet}")); + } + serde_json::from_str(&text).map_err(|_| format!("Upstash returned a non-JSON response for {what}")) +} + +pub async fn list_databases(token: &str) -> Result, String> { + let body = get(token, "/redis/databases", "listing databases").await?; + Ok(parse_list(&body)) +} + +fn parse_list(body: &Value) -> Vec { + body.as_array() + .into_iter() + .flatten() + .filter_map(|db| { + let id = db["database_id"].as_str()?; + Some(ProviderDatabase { + db_ref: id.to_string(), + name: db["database_name"].as_str().unwrap_or(id).to_string(), + region: db["primary_region"] + .as_str() + .or_else(|| db["region"].as_str()) + .filter(|r| !r.is_empty() && *r != "global") + .map(String::from), + kind: Some("Redis".into()), + host: db["endpoint"].as_str().map(String::from), + }) + }) + .collect() +} + +/// A database's detail → a connection with its own endpoint and password. +fn connection_from_detail(db: &Value, database_id: &str) -> Result { + let name = db["database_name"].as_str().unwrap_or(database_id).to_string(); + if let Some(state) = db["state"].as_str().filter(|s| *s != "active") { + return Err(format!("{name} is not active ({state}). Check it in the Upstash console.")); + } + let host = db["endpoint"].as_str().ok_or_else(|| format!("Upstash returned no endpoint for {name}"))?; + let password = db["password"].as_str().ok_or_else(|| format!("Upstash returned no password for {name}"))?; + let port = db["port"].as_u64().and_then(|p| u16::try_from(p).ok()).unwrap_or(6379); + Ok(ProviderConnection { + db_type: "redis".into(), + host: host.to_string(), + port, + username: "default".into(), + password: password.to_string(), + database: "0".into(), + // Upstash only accepts TLS on the Redis port. + ssl: db["tls"].as_bool().unwrap_or(true), + needs_password: false, + name: format!("Upstash · {name}"), + }) +} + +pub async fn build_connection(token: &str, database_id: &str) -> Result { + let db = get( + token, + &format!("/redis/database/{}", urlencoding::encode(database_id)), + "reading the database", + ) + .await?; + connection_from_detail(&db, database_id) +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn list_and_detail_map_to_a_tls_redis_connection() { + let rows = parse_list(&json!([{ "database_id": "abc", "database_name": "test", "region": "global", "primary_region": "ap-south-1", "endpoint": "fine-cat-1.upstash.io" }])); + assert_eq!((rows[0].db_ref.as_str(), rows[0].region.as_deref()), ("abc", Some("ap-south-1"))); + let detail = json!({ "database_name": "test", "state": "active", "endpoint": "fine-cat-1.upstash.io", "port": 6379, "password": "pw", "tls": true }); + let c = connection_from_detail(&detail, "abc").unwrap(); + assert_eq!((c.db_type.as_str(), c.host.as_str(), c.port, c.ssl), ("redis", "fine-cat-1.upstash.io", 6379, true)); + assert!(connection_from_detail(&json!({ "database_name": "t", "state": "suspended" }), "abc").unwrap_err().contains("not active")); + } + + #[test] + fn stored_token_splits_into_email_and_key() { + assert_eq!(credentials("me@x.dev:abc:def").unwrap(), ("me@x.dev", "abc:def")); + assert!(credentials("no-colon").is_err()); + assert!(credentials(":key").is_err()); + } +} diff --git a/src-tauri/src/secrets.rs b/src-tauri/src/secrets.rs index c9ec4f03..2bf7ba92 100644 --- a/src-tauri/src/secrets.rs +++ b/src-tauri/src/secrets.rs @@ -1,4 +1,5 @@ use std::collections::HashMap; +use std::path::Path; use std::sync::{Mutex, OnceLock}; use tauri::Manager; @@ -12,15 +13,31 @@ fn cache() -> &'static Mutex>> { CACHE.get_or_init(|| Mutex::new(None)) } -// All secrets (AI keys, provider OAuth tokens, Cloudflare tokens) live in a -// single JSON blob stored in the OS keychain - macOS Keychain, Windows -// Credential Manager, or Linux Secret Service. A legacy plaintext file -// (`ai-keys.json`) is migrated in on first read and then deleted; it also -// remains the fallback store if the keychain is unavailable, so credentials are -// never lost or unreadable. +// Serialises read-modify-write of the vault. Every caller rewrites the whole +// map, so two unguarded updates (a provider token refresh and an AI key save, +// say) each wrote back the map they read and one silently undid the other. +static VAULT_LOCK: Mutex<()> = Mutex::new(()); + +// All secrets (AI keys, provider OAuth tokens, Cloudflare tokens) live in one +// JSON map stored in the OS keychain - macOS Keychain, Windows Credential +// Manager, or Linux Secret Service. A plaintext file (`ai-keys.json`) is the +// fallback when the keychain can't hold it, and the pre-keychain format. const KEYCHAIN_SERVICE: &str = "app.stroke.desktop"; const KEYCHAIN_ACCOUNT: &str = "secrets-vault"; +// Windows caps one credential at 2560 bytes of UTF-16 - 1280 characters. A +// single OAuth access + refresh token pair is past that, so the map is split +// across several credentials there. Before this, every write after the first +// sign-in failed, the tokens went to the fallback file, and the next launch +// read the older, smaller credential first: signed out on every restart. +#[cfg(windows)] +const KEYCHAIN_CHUNK_UTF16: usize = 1200; +// macOS and Linux have no practical limit; one chunk keeps a single item. +#[cfg(not(windows))] +const KEYCHAIN_CHUNK_UTF16: usize = 1 << 20; + +type Map = HashMap; + fn legacy_path(app: &tauri::AppHandle) -> std::path::PathBuf { app.path() .app_data_dir() @@ -28,57 +45,234 @@ fn legacy_path(app: &tauri::AppHandle) -> std::path::PathBuf { .join("ai-keys.json") } -fn keychain_entry() -> Option { - keyring::Entry::new(KEYCHAIN_SERVICE, KEYCHAIN_ACCOUNT).ok() +// ── Credential store ───────────────────────────────────────────────────────── + +/// Named secret slots. The OS keychain in the app; an in-memory map in tests. +trait CredStore { + fn get(&self, account: &str) -> Result, String>; + fn set(&self, account: &str, value: &str) -> Result<(), String>; + fn delete(&self, account: &str) -> Result<(), String>; + /// Largest value one slot takes, in UTF-16 code units. + fn max_utf16(&self) -> usize; } -fn parse_map(s: &str) -> HashMap { - serde_json::from_str::>(s).unwrap_or_default() +struct Keychain; + +impl Keychain { + fn entry(account: &str) -> Result { + keyring::Entry::new(KEYCHAIN_SERVICE, account).map_err(|e| e.to_string()) + } } -fn write_keychain(map: &HashMap) -> Result<(), String> { - let entry = keychain_entry().ok_or_else(|| "keychain unavailable".to_string())?; - let json = serde_json::to_string(map).map_err(|e| e.to_string())?; - entry.set_password(&json).map_err(|e| e.to_string()) -} - -/// Read the keychain back and confirm it holds exactly what we intended to store. -/// This is what makes a silently-non-persisting backend (e.g. keyring's mock -/// store) detectable: a write can "succeed" yet not round-trip, in which case we -/// must keep the plaintext file rather than delete it and lose the data. -fn keychain_holds(expected: &HashMap) -> bool { - keychain_entry() - .and_then(|e| e.get_password().ok()) - .map(|json| parse_map(&json) == *expected) - .unwrap_or(false) -} - -/// Load the durable store: keychain first (encrypted at rest), then the legacy -/// plaintext file, migrating the file into the keychain when that round-trips. -fn load_durable(app: &tauri::AppHandle) -> HashMap { - if let Some(entry) = keychain_entry() { - if let Ok(json) = entry.get_password() { - let map = parse_map(&json); - if !map.is_empty() { - return map; - } +impl CredStore for Keychain { + fn get(&self, account: &str) -> Result, String> { + match Self::entry(account)?.get_password() { + Ok(s) => Ok(Some(s)), + Err(keyring::Error::NoEntry) => Ok(None), + Err(e) => Err(e.to_string()), + } + } + fn set(&self, account: &str, value: &str) -> Result<(), String> { + Self::entry(account)?.set_password(value).map_err(|e| e.to_string()) + } + fn delete(&self, account: &str) -> Result<(), String> { + match Self::entry(account)?.delete_credential() { + Ok(()) | Err(keyring::Error::NoEntry) => Ok(()), + Err(e) => Err(e.to_string()), } } - // Legacy plaintext file: the fallback store, and the pre-keychain format. - let path = legacy_path(app); - let map = std::fs::read_to_string(&path) + fn max_utf16(&self) -> usize { + KEYCHAIN_CHUNK_UTF16 + } +} + +// ── Chunked vault format ───────────────────────────────────────────────────── +// +// `secrets-vault` holds a small header, `{"v":2,"gen":"…","chunks":N}`, and the +// map's JSON is split over `secrets-vault..`. Each write uses a fresh +// `gen`, so the header flips from the old chunk set to the new one in a single +// write and a failure part-way never leaves a mix of both. A header that is not +// a v2 object is the old format: the whole map in one credential. + +#[derive(serde::Serialize, serde::Deserialize)] +struct Header { + v: u8, + gen: String, + chunks: usize, +} + +fn chunk_account(gen: &str, i: usize) -> String { + format!("{KEYCHAIN_ACCOUNT}.{gen}.{i}") +} + +/// Split on char boundaries into pieces of at most `max` UTF-16 units. +fn split_utf16(s: &str, max: usize) -> Vec { + let mut out = vec![String::new()]; + let mut units = 0; + for c in s.chars() { + let n = c.len_utf16(); + if units + n > max { + out.push(String::new()); + units = 0; + } + out.last_mut().expect("non-empty").push(c); + units += n; + } + out +} + +fn parse_map(s: &str) -> Option { + serde_json::from_str::(s).ok() +} + +fn read_header(store: &dyn CredStore) -> Result)>, String> { + let Some(raw) = store.get(KEYCHAIN_ACCOUNT)? else { + return Ok(None); + }; + let header = serde_json::from_str::(&raw) .ok() - .map(|s| parse_map(&s)) - .unwrap_or_default(); - // Migrate into the keychain only if it verifiably persisted; otherwise leave - // the file exactly where it is. - if !map.is_empty() && write_keychain(&map).is_ok() && keychain_holds(&map) { - let _ = std::fs::remove_file(&path); + .filter(|v| v.get("v").and_then(|v| v.as_u64()) == Some(2)) + .and_then(|v| serde_json::from_value::
(v).ok()); + Ok(Some((raw, header))) +} + +/// The stored map, `None` when the keychain holds nothing. +fn read_vault(store: &dyn CredStore) -> Result, String> { + let Some((raw, header)) = read_header(store)? else { + return Ok(None); + }; + let Some(h) = header else { + // Old single-credential format. + return parse_map(&raw).map(Some).ok_or_else(|| "keychain vault is not valid JSON".into()); + }; + let mut json = String::new(); + for i in 0..h.chunks { + json.push_str( + &store + .get(&chunk_account(&h.gen, i))? + .ok_or_else(|| format!("keychain vault is missing part {i}"))?, + ); } - map + parse_map(&json).map(Some).ok_or_else(|| "keychain vault is not valid JSON".into()) +} + +/// Write `map` and prove it round-trips. On failure nothing the old header +/// points at has been touched. +fn write_vault(store: &dyn CredStore, map: &Map) -> Result<(), String> { + let json = serde_json::to_string(map).map_err(|e| e.to_string())?; + let old = read_header(store).ok().flatten().and_then(|(_, h)| h); + let gen = { + let mut b = [0u8; 6]; + getrandom::getrandom(&mut b).map_err(|e| e.to_string())?; + b.iter().map(|x| format!("{x:02x}")).collect::() + }; + let chunks = split_utf16(&json, store.max_utf16()); + let cleanup_new = |upto: usize| { + for i in 0..upto { + let _ = store.delete(&chunk_account(&gen, i)); + } + }; + for (i, c) in chunks.iter().enumerate() { + if let Err(e) = store.set(&chunk_account(&gen, i), c) { + cleanup_new(i); + return Err(e); + } + } + let header = serde_json::to_string(&Header { v: 2, gen: gen.clone(), chunks: chunks.len() }) + .map_err(|e| e.to_string())?; + if let Err(e) = store.set(KEYCHAIN_ACCOUNT, &header) { + cleanup_new(chunks.len()); + return Err(e); + } + // A write can "succeed" yet not persist (keyring's mock store, a locked + // keyring): only a read-back that matches counts. + if read_vault(store).ok().flatten().as_ref() != Some(map) { + return Err("keychain did not keep the vault".into()); + } + if let Some(h) = old { + for i in 0..h.chunks { + let _ = store.delete(&chunk_account(&h.gen, i)); + } + } + Ok(()) +} + +/// Remove the vault from the keychain entirely (header first, so a half-done +/// delete reads as empty rather than as a broken vault). +fn clear_vault(store: &dyn CredStore) -> Result<(), String> { + let old = read_header(store)?.and_then(|(_, h)| h); + store.delete(KEYCHAIN_ACCOUNT)?; + if let Some(h) = old { + for i in 0..h.chunks { + let _ = store.delete(&chunk_account(&h.gen, i)); + } + } + Ok(()) +} + +// ── Fallback file ──────────────────────────────────────────────────────────── + +fn write_file(path: &Path, map: &Map) -> Result<(), String> { + use std::io::Write; + if let Some(parent) = path.parent() { + std::fs::create_dir_all(parent).map_err(|e| e.to_string())?; + } + let json = serde_json::to_string(map).map_err(|e| e.to_string())?; + let tmp = path.with_extension("json.tmp"); + let mut opts = std::fs::OpenOptions::new(); + opts.write(true).create(true).truncate(true); + // Plaintext secrets: private to this user. + #[cfg(unix)] + { + use std::os::unix::fs::OpenOptionsExt; + opts.mode(0o600); + } + let mut f = opts.open(&tmp).map_err(|e| e.to_string())?; + f.write_all(json.as_bytes()) + .and_then(|_| f.sync_all()) + .map_err(|e| e.to_string())?; + drop(f); + std::fs::rename(&tmp, path).map_err(|e| e.to_string()) +} + +// ── Durable load / persist ─────────────────────────────────────────────────── +// +// Invariant: the fallback file exists only while it is the newest copy. A +// keychain write that round-trips deletes it; a keychain write that fails +// writes it and then removes the keychain copy. So on load the file, when +// present, always wins - which also rescues installs where the old code left a +// stale credential in front of a newer file. + +fn load_from(path: &Path, store: &dyn CredStore) -> Map { + if let Some(map) = std::fs::read_to_string(path).ok().and_then(|s| parse_map(&s)) { + // Move it into the keychain when that verifiably works. + let _ = persist_to(path, store, &map); + return map; + } + read_vault(store).ok().flatten().unwrap_or_default() +} + +fn persist_to(path: &Path, store: &dyn CredStore, map: &Map) -> Result<(), String> { + if write_vault(store, map).is_ok() { + if path.exists() && std::fs::remove_file(path).is_err() { + // Could not delete it: keep it current instead, or it would win the + // next load with old contents. + write_file(path, map)?; + } + return Ok(()); + } + // Keychain missing, too small, or not round-tripping: the file becomes the + // store, and the keychain copy goes so it can never be read instead of it. + write_file(path, map)?; + let _ = clear_vault(store); + Ok(()) +} + +fn load_durable(app: &tauri::AppHandle) -> Map { + load_from(&legacy_path(app), &Keychain) } -pub(crate) fn read_all(app: &tauri::AppHandle) -> HashMap { +pub(crate) fn read_all(app: &tauri::AppHandle) -> Map { let mut guard = cache().lock().unwrap_or_else(|e| e.into_inner()); if let Some(map) = guard.as_ref() { return map.clone(); @@ -88,25 +282,21 @@ pub(crate) fn read_all(app: &tauri::AppHandle) -> HashMap { map } -pub(crate) fn write_all(app: &tauri::AppHandle, map: &HashMap) -> Result<(), String> { +fn write_unlocked(app: &tauri::AppHandle, map: &Map) -> Result<(), String> { // Session cache is authoritative first, so reads right after this always see // the new value regardless of what the durable backend does. *cache().lock().unwrap_or_else(|e| e.into_inner()) = Some(map.clone()); + persist_to(&legacy_path(app), &Keychain, map) +} - // Prefer the keychain, but only trust it - and drop the plaintext copy - once - // a read-back proves the data actually persisted. - if write_keychain(map).is_ok() && keychain_holds(map) { - let _ = std::fs::remove_file(legacy_path(app)); - return Ok(()); - } - // Keychain missing, mock, or not round-tripping - persist to the file so - // secrets survive a restart. - let path = legacy_path(app); - if let Some(parent) = path.parent() { - std::fs::create_dir_all(parent).map_err(|e| e.to_string())?; - } - let json = serde_json::to_string(map).map_err(|e| e.to_string())?; - std::fs::write(&path, json).map_err(|e| e.to_string()) +/// Read, change and write the vault as one step. Use this for every change: +/// a separate read and write can lose a concurrent update. +pub(crate) fn update(app: &tauri::AppHandle, f: impl FnOnce(&mut Map) -> R) -> Result { + let _guard = VAULT_LOCK.lock().unwrap_or_else(|e| e.into_inner()); + let mut map = read_all(app); + let r = f(&mut map); + write_unlocked(app, &map)?; + Ok(r) } // ── Keychain access must never run on the caller's thread ──────────────────── @@ -121,8 +311,8 @@ pub(crate) fn write_all(app: &tauri::AppHandle, map: &HashMap) - // queries. // // So every path that can reach the keychain goes through `read_all_async` / -// `write_all_async`, which move the blocking work to the dedicated blocking -// pool. The sync `read_all`/`write_all` stay for use *inside* those closures. +// `update_async`, which move the blocking work to the dedicated blocking pool. +// The sync `read_all`/`update` stay for use *inside* those closures. async fn off_thread(f: F) -> Result where F: FnOnce() -> T + Send + 'static, @@ -133,19 +323,20 @@ where .map_err(|e| e.to_string()) } -pub(crate) async fn read_all_async(app: &tauri::AppHandle) -> HashMap { +pub(crate) async fn read_all_async(app: &tauri::AppHandle) -> Map { let app = app.clone(); // A join error here means the closure panicked; an empty vault is the same // answer callers already handle for "nothing stored yet". off_thread(move || read_all(&app)).await.unwrap_or_default() } -pub(crate) async fn write_all_async( - app: &tauri::AppHandle, - map: HashMap, -) -> Result<(), String> { +pub(crate) async fn update_async(app: &tauri::AppHandle, f: F) -> Result +where + F: FnOnce(&mut Map) -> R + Send + 'static, + R: Send + 'static, +{ let app = app.clone(); - off_thread(move || write_all(&app, &map)).await? + off_thread(move || update(&app, f)).await? } #[tauri::command] @@ -154,17 +345,14 @@ pub async fn ai_store_key( profile_id: String, api_key: String, ) -> Result<(), String> { - // Read and write in one hop so the pair shares a single keychain unlock. - off_thread(move || { - let mut map = read_all(&app); + update_async(&app, move |map| { if api_key.is_empty() { map.remove(&profile_id); } else { map.insert(profile_id, api_key); } - write_all(&app, &map) }) - .await? + .await } #[tauri::command] @@ -178,10 +366,135 @@ pub async fn ai_load_key(app: tauri::AppHandle, profile_id: String) -> Result Result<(), String> { - off_thread(move || { - let mut map = read_all(&app); + update_async(&app, move |map| { map.remove(&profile_id); - write_all(&app, &map) }) - .await? + .await +} + +#[cfg(test)] +mod tests { + use super::*; + use std::cell::RefCell; + + /// In-memory keychain with Windows' per-credential limit, and switches to + /// make it refuse writes. + struct MemStore { + slots: RefCell>, + max: usize, + fail_writes: RefCell, + } + + impl MemStore { + fn new(max: usize) -> Self { + Self { slots: RefCell::new(HashMap::new()), max, fail_writes: RefCell::new(false) } + } + } + + impl CredStore for MemStore { + fn get(&self, account: &str) -> Result, String> { + Ok(self.slots.borrow().get(account).cloned()) + } + fn set(&self, account: &str, value: &str) -> Result<(), String> { + if *self.fail_writes.borrow() { + return Err("keychain unavailable".into()); + } + // What Windows Credential Manager enforces (keyring's TooLong). + if value.encode_utf16().count() * 2 > 2560 { + return Err("Attribute 'password' is longer than platform limit of 2560 chars".into()); + } + self.slots.borrow_mut().insert(account.into(), value.into()); + Ok(()) + } + fn delete(&self, account: &str) -> Result<(), String> { + self.slots.borrow_mut().remove(account); + Ok(()) + } + fn max_utf16(&self) -> usize { + self.max + } + } + + fn big_map() -> Map { + // Roughly an access + refresh token pair per provider, plus an AI key. + let token = "x".repeat(1800); + HashMap::from([ + ("__prisma_access__".into(), token.clone()), + ("__prisma_refresh__".into(), token.clone()), + ("__neon_access__".into(), token), + ("openai".into(), "sk-test".into()), + ]) + } + + fn tmp_path(name: &str) -> std::path::PathBuf { + let dir = std::env::temp_dir().join(format!("stroke-secrets-{name}-{}", std::process::id())); + let _ = std::fs::remove_dir_all(&dir); + std::fs::create_dir_all(&dir).unwrap(); + dir.join("ai-keys.json") + } + + #[test] + fn a_vault_past_the_windows_limit_round_trips_in_chunks() { + let store = MemStore::new(1200); + let map = big_map(); + write_vault(&store, &map).unwrap(); + assert_eq!(read_vault(&store).unwrap(), Some(map)); + assert!(store.slots.borrow().len() > 2, "expected several chunks"); + } + + #[test] + fn rewriting_drops_the_previous_chunks() { + let store = MemStore::new(1200); + write_vault(&store, &big_map()).unwrap(); + let small = HashMap::from([("k".to_string(), "v".to_string())]); + write_vault(&store, &small).unwrap(); + assert_eq!(read_vault(&store).unwrap(), Some(small)); + // Header + one chunk. + assert_eq!(store.slots.borrow().len(), 2); + } + + #[test] + fn the_old_single_credential_format_still_reads() { + let store = MemStore::new(1200); + store.slots.borrow_mut().insert(KEYCHAIN_ACCOUNT.into(), r#"{"openai":"sk-old"}"#.into()); + assert_eq!(read_vault(&store).unwrap().unwrap()["openai"], "sk-old"); + } + + #[test] + fn a_newer_fallback_file_beats_a_stale_credential() { + // What the old code left behind on Windows: a small credential from + // before sign-in, and the tokens only in the file. + let store = MemStore::new(1200); + store.slots.borrow_mut().insert(KEYCHAIN_ACCOUNT.into(), r#"{"openai":"sk-old"}"#.into()); + let path = tmp_path("stale"); + write_file(&path, &big_map()).unwrap(); + + let loaded = load_from(&path, &store); + assert_eq!(loaded, big_map()); + // ...and it was moved into the (chunked) keychain, file gone. + assert_eq!(read_vault(&store).unwrap(), Some(big_map())); + assert!(!path.exists()); + } + + #[test] + fn a_failing_keychain_falls_back_to_the_file_and_stays_signed_in() { + let store = MemStore::new(1200); + store.slots.borrow_mut().insert(KEYCHAIN_ACCOUNT.into(), r#"{"openai":"sk-old"}"#.into()); + *store.fail_writes.borrow_mut() = true; + let path = tmp_path("fallback"); + + persist_to(&path, &store, &big_map()).unwrap(); + // The stale credential can't be read in place of the file any more. + assert_eq!(read_vault(&store).unwrap(), None); + // Next launch. + assert_eq!(load_from(&path, &store), big_map()); + } + + #[test] + fn split_respects_the_limit_and_char_boundaries() { + let s = "aé😀".repeat(500); + let parts = split_utf16(&s, 7); + assert!(parts.iter().all(|p| p.encode_utf16().count() <= 7)); + assert_eq!(parts.concat(), s); + } } diff --git a/src/app.css b/src/app.css index 0aca2303..26a850a9 100644 --- a/src/app.css +++ b/src/app.css @@ -7,6 +7,7 @@ @import "@fontsource-variable/fira-code"; @import "@fontsource-variable/source-code-pro"; @import "@fontsource-variable/space-grotesk"; +@import "@fontsource-variable/source-serif-4"; @import "@fontsource/ibm-plex-sans/400.css"; @import "@fontsource/ibm-plex-sans/500.css"; @import "@fontsource/ibm-plex-sans/600.css"; @@ -445,6 +446,10 @@ html[data-os="linux"] { --window-radius: 8px; } --color-link-hover: var(--link-hover); --font-sans: "Geist Variable", ui-sans-serif, system-ui, sans-serif; --font-mono: "Geist Mono Variable", ui-monospace, monospace; + /* Dialog titles and headings. `--heading-font` is set by the font preset + (the Claude preset pairs Inter with a serif); this block is inlined, so the + utility has to point at a variable the preset can override. */ + --font-heading: var(--heading-font, var(--font-sans)); } @layer base { diff --git a/src/lib/components/ChartView.svelte b/src/lib/components/ChartView.svelte index cd8d8d5e..0b978ae5 100644 --- a/src/lib/components/ChartView.svelte +++ b/src/lib/components/ChartView.svelte @@ -104,18 +104,29 @@ // Aggregating charts (bar/pie/line) still represent the whole set closely; the // toolbar flags when data was sampled. const MAX_CHART_ROWS = 50_000 - const sampled = $derived(rows.length > MAX_CHART_ROWS) + // A windowed table browse hands over a sparse array as long as the whole + // table, with holes where pages aren't loaded yet. Indexing into it sampled + // holes (`undefined`) and the coercion pass threw on `row.map`. Chart the rows + // that exist: forEach skips holes, and a dense array passes through as is. + const loadedRows = $derived.by(() => { + /** @type {any[]} */ + const out = [] + rows.forEach((r) => { if (r) out.push(r) }) + return out.length === rows.length ? rows : out + }) + const partial = $derived(loadedRows.length < rows.length) + const sampled = $derived(loadedRows.length > MAX_CHART_ROWS) const chartRows = $derived.by(() => { - if (rows.length <= MAX_CHART_ROWS) return rows - const step = rows.length / MAX_CHART_ROWS + if (loadedRows.length <= MAX_CHART_ROWS) return loadedRows + const step = loadedRows.length / MAX_CHART_ROWS const out = new Array(MAX_CHART_ROWS) - for (let i = 0; i < MAX_CHART_ROWS; i++) out[i] = rows[Math.floor(i * step)] + for (let i = 0; i < MAX_CHART_ROWS; i++) out[i] = loadedRows[Math.floor(i * step)] return out }) // Sniff row data to detect numeric columns that the DB reported with no/wrong type const effectiveColumns = $derived.by(() => { - if (!rows.length) return columns + if (!loadedRows.length) return columns return columns.map((col, i) => { if (colType(col) === 'number') return col const samples = chartRows.slice(0, 20).map(r => /** @type {any} */ (r)[i]).filter(v => v != null && v !== '') @@ -129,7 +140,7 @@ // Coerce string-encoded numerics to actual numbers in rows const effectiveRows = $derived.by(() => { - if (!rows.length) return rows + if (!loadedRows.length) return loadedRows return chartRows.map(row => /** @type {any[]} */ (row).map((v, i) => { if (effectiveColumns[i] && colType(effectiveColumns[i]) === 'number' && typeof v === 'string') { @@ -478,7 +489,7 @@
- {rows.length.toLocaleString()} rows + {loadedRows.length.toLocaleString()}{partial ? ` of ${rows.length.toLocaleString()} loaded` : ' rows'} @@ -552,7 +563,7 @@ - {#if rows.length === 0} + {#if loadedRows.length === 0}

No data to display

@@ -563,7 +574,7 @@ {/if} {#if sampled}
- sampled {MAX_CHART_ROWS.toLocaleString()} of {rows.length.toLocaleString()} rows + sampled {MAX_CHART_ROWS.toLocaleString()} of {loadedRows.length.toLocaleString()} {partial ? 'loaded ' : ''}rows
{/if} {#snippet failed(error, reset)} diff --git a/src/lib/components/CloudflareLogin.svelte b/src/lib/components/CloudflareLogin.svelte index 573de0ae..f4b557ab 100644 --- a/src/lib/components/CloudflareLogin.svelte +++ b/src/lib/components/CloudflareLogin.svelte @@ -11,7 +11,9 @@ import SearchableMenu from './SearchableMenu.svelte' import ProviderAuthPanel from './ProviderAuthPanel.svelte' import { Button } from '$lib/components/ui/button/index.js' - import { cfStartOAuth, cfOAuthStatus, cfLogout } from '$lib/cloudflare.js' + import ConfirmDialog from './ConfirmDialog.svelte' + import { cfStartOAuth, cfOAuthStatus, cfLogout, cfGetValidToken } from '$lib/cloudflare.js' + import { readProviderList, writeProviderList, clearProviderLists } from '$lib/provider-list-cache.js' import { cloudflareListAccounts, cloudflareListD1Databases } from '$lib/api.js' import { cn } from '$lib/utils.js' @@ -113,7 +115,35 @@ } const shownError = $derived(friendlyError(errorMsg)) + const CACHE = 'cloudflare:' + /** @param {string} accountId */ + const dbsKey = (accountId) => `${CACHE}dbs:${accountId}` + + /** The account to land on: the saved connection's, else the last one used, else the first. */ + function pickAccount(/** @type {Array<{id: string}>} */ list) { + /** @type {string | undefined} */ + const last = readProviderList(`${CACHE}lastAccount`) + for (const want of [selectedAccountId, initialAccountId, last]) { + if (want && list.some((a) => a.id === want)) return want + } + return list[0]?.id ?? '' + } + onMount(async () => { + // Open on the accounts and databases from last time, then refresh both + // underneath. The token read fails with "not signed in" by itself when the + // session is gone, so the status round-trip only runs with nothing to show. + /** @type {{ email?: string, accounts?: typeof accounts } | undefined} */ + const cached = readProviderList(`${CACHE}accounts`) + if (cached?.accounts?.length) { + email = cached.email ?? '' + accounts = cached.accounts + selectedAccountId = pickAccount(accounts) + databases = readProviderList(dbsKey(selectedAccountId)) ?? [] + phase = 'selecting' + void loadAccounts({ quiet: true }) + return + } try { const status = await cfOAuthStatus() if (status.connected) { @@ -138,53 +168,81 @@ } } - async function loadAccounts() { - phase = 'fetching' + /** + * A failure while refreshing a cached view keeps that view: a network blip + * shouldn't swap a usable list for an error card. An ended session still + * surfaces, because nothing on the cached list would connect. + * @param {unknown} e @param {boolean} quiet + */ + function fail(e, quiet) { + const msg = String(e) + const ended = /not signed in|no.*token|unauthor/i.test(msg) + if (ended) clearProviderLists(CACHE) + if (quiet && !ended) return + phase = 'error' + errorMsg = msg + } + + /** @param {{ quiet?: boolean }} [opts] */ + async function loadAccounts({ quiet = false } = {}) { + if (!quiet) phase = 'fetching' errorMsg = '' try { - const { cfGetValidToken } = await import('$lib/cloudflare.js') const token = await withTimeout(cfGetValidToken(), 20_000, 'reading your Cloudflare session') + // Start the D1 list for the account we expect to land on alongside the + // account list, instead of waiting for one to finish before the other. + /** @type {string} */ + const guess = accounts.length + ? pickAccount(accounts) + : initialAccountId || readProviderList(`${CACHE}lastAccount`) || '' + const early = guess ? cloudflareListD1Databases(token, guess) : null + early?.catch(() => {}) // settled below or abandoned; never an unhandled rejection accounts = await withTimeout( cloudflareListAccounts(token), 20_000, 'listing your Cloudflare accounts', ) + writeProviderList(`${CACHE}accounts`, { email, accounts: $state.snapshot(accounts) }) phase = 'selecting' // Auto-select an account so the D1 database list loads immediately; the // user can still switch accounts via the dropdown when there are several. - // A saved connection's own account wins, otherwise the first one. - if (accounts.length && !selectedAccountId) { - const seeded = accounts.some((a) => a.id === initialAccountId) - ? initialAccountId - : accounts[0].id - await selectAccount(seeded) - } + const id = pickAccount(accounts) + if (id) await selectAccount(id, { token, pending: id === guess ? early : null, quiet }) } catch (e) { - phase = 'error' - errorMsg = String(e) + fail(e, quiet) } } - async function selectAccount(id) { - selectedAccountId = id - selectedDbUuid = '' - databases = [] - loadingDbs = true + /** + * @param {string} id + * @param {{ token?: string, pending?: Promise | null, quiet?: boolean }} [opts] + */ + async function selectAccount(id, { token, pending = null, quiet = false } = {}) { + if (id !== selectedAccountId) { + selectedAccountId = id + selectedDbUuid = '' + // Another account's last-known list, if we have one, while it refreshes. + databases = readProviderList(dbsKey(id)) ?? [] + } + writeProviderList(`${CACHE}lastAccount`, id) + loadingDbs = databases.length === 0 try { - const { cfGetValidToken } = await import('$lib/cloudflare.js') - const token = await withTimeout(cfGetValidToken(), 20_000, 'reading your Cloudflare session') - databases = await withTimeout( - cloudflareListD1Databases(token, id), + const tok = token ?? (await withTimeout(cfGetValidToken(), 20_000, 'reading your Cloudflare session')) + const list = await withTimeout( + pending ?? cloudflareListD1Databases(tok, id), 20_000, 'listing D1 databases for this account', ) + // The user may have switched accounts while this was in flight. + if (id !== selectedAccountId) return + databases = list + writeProviderList(dbsKey(id), $state.snapshot(list)) } catch (e) { // Show the error card - staying in 'selecting' rendered a misleading // "No D1 databases in this account" empty state over a real failure. - phase = 'error' - errorMsg = String(e) + if (id === selectedAccountId) fail(e, quiet) } finally { - loadingDbs = false + if (id === selectedAccountId) loadingDbs = false } } @@ -193,7 +251,6 @@ const db = databases.find(d => d.uuid === uuid) if (!db) return try { - const { cfGetValidToken } = await import('$lib/cloudflare.js') const token = await cfGetValidToken() onselect({ accountId: selectedAccountId, @@ -207,8 +264,12 @@ } } + /** Sign-out asks first: saved D1 connections depend on this sign-in. */ + let confirmSignOut = $state(false) + async function handleLogout() { await cfLogout() + clearProviderLists(CACHE) phase = 'idle' email = '' accounts = [] @@ -278,8 +339,8 @@ Try again {#if signedIn} - {/if} @@ -302,20 +363,16 @@

{email}

{/if}
- + {#if accounts.length > 1} -
- Account +
+ Account ({ value: a.id, label: a.name }))} placeholder="Search accounts…" @@ -330,7 +387,7 @@ class="field-surface flex h-9 w-full items-center gap-2 bg-muted/25 pl-3 pr-2.5 text-left text-ui-xs transition-[border-color,box-shadow] hover:border-border focus:outline-none data-[state=open]:border-ring" > - {selectedAccountName || '- select account -'} + {selectedAccountName || 'Select an account'} @@ -348,8 +405,8 @@ {#if selectedAccountId} -
- D1 Database +
+ D1 database{#if databases.length}{databases.length}{/if} {#if databases.length > 0} - {selectedDbName || '- select database -'} + {selectedDbName || 'Select a database'} @@ -387,16 +444,21 @@ Loading databases…
{:else} -
- -

No D1 databases in this account.

- + +
+
+ +
+
+

No databases yet

+

This account has no D1 databases. Create one with wrangler d1 create or in the Cloudflare dashboard, then refresh.

+
+
{/if}
@@ -406,3 +468,16 @@
+ + void handleLogout()} +/> diff --git a/src/lib/components/ConnectionModal.svelte b/src/lib/components/ConnectionModal.svelte index a1b2809b..173be8f0 100644 --- a/src/lib/components/ConnectionModal.svelte +++ b/src/lib/components/ConnectionModal.svelte @@ -41,6 +41,7 @@ import SearchableMenu from "./SearchableMenu.svelte"; import { Popover, PopoverTrigger, PopoverContent } from "$lib/components/ui/popover/index.js"; import PasswordInput from "./PasswordInput.svelte"; + import Kbd from "./Kbd.svelte"; import { requireUnlock } from "$lib/stores/app-lock.js"; import { readClipboardText } from "$lib/clipboard.js"; import { Checkbox } from "$lib/components/ui/checkbox/index.js"; @@ -54,6 +55,8 @@ import { toast } from "$lib/components/ui/sonner/toast.svelte.js"; import { parseConnectionUri, detectConnectionUri } from "$lib/connection-uri.js"; import { PROVIDERS, providerBuildConnection } from "$lib/providers.js"; + import { providerOf, engineLabel } from "$lib/connection-provider.js"; + import ConfirmDialog from "./ConfirmDialog.svelte"; let { open = $bindable(false), @@ -67,6 +70,8 @@ /** Name of the live session, '' when nothing is connected. Drives Disconnect. */ activeConnectionName = "", ondisconnect = () => {}, + /** A saved connection was deleted here; the shell ends its session if open. */ + onremoved = (/** @type {string} */ id) => {}, } = $props(); const CATEGORIES = [ @@ -158,6 +163,31 @@ label: "Prisma Postgres", desc: "Paste a Prisma Postgres connection string", }, + { + id: "tidb", + label: "TiDB Cloud", + desc: "Serverless MySQL, sign in & pick a cluster", + }, + { + id: "turso", + label: "Turso", + desc: "Edge SQLite, sign in & pick a database", + }, + { + id: "railway", + label: "Railway", + desc: "Postgres, MySQL & Redis, sign in & pick a service", + }, + { + id: "nile", + label: "Nile", + desc: "Multi-tenant Postgres, sign in & pick a database", + }, + { + id: "upstash", + label: "Upstash", + desc: "Serverless Redis, connect with an API key", + }, ], }, ]; @@ -185,6 +215,11 @@ "supabase", "planetscale", "prisma", + "tidb", + "turso", + "railway", + "nile", + "upstash", "d1", "redis", ]; @@ -201,9 +236,12 @@ // Provider (sign-in) ids are surfaced as cards on their own tab, so keep them // out of the manual Type dropdown. - const PROVIDER_IDS = ["neon", "supabase", "planetscale", "prisma"]; + const PROVIDER_IDS = ["neon", "supabase", "planetscale", "prisma", "tidb", "turso", "railway", "nile", "upstash"]; // Providers temporarily turned off (shown as a disabled tab, not connectable). - const DISABLED_TABS = new Set(["planetscale"]); + // Railway: the adapter is done, but its OAuth app isn't registered yet, so + // there is no client id to sign in with. + /** @type {Set} */ + const DISABLED_TABS = new Set(["railway"]); // Subtle per-engine icon tint (color-500/600), theme-aware via Tailwind tokens. const ENGINE_TINT = { @@ -219,10 +257,16 @@ "duckdb-memory": "text-yellow-500/80", d1: "text-orange-500/80", libsql: "text-emerald-500/80", + docker: "text-sky-500/80", neon: "text-emerald-500/80", supabase: "text-emerald-500/80", planetscale: "text-foreground/70", prisma: "text-indigo-500/80", + tidb: "text-red-500/80", + turso: "text-teal-500/80", + railway: "text-foreground/80", + nile: "text-violet-500/80", + upstash: "text-emerald-500/80", drizzle: "text-lime-500/80", redis: "text-red-500/80", }; @@ -306,7 +350,6 @@ // Top-level entry mode: connect manually vs sign in with a hosting provider. let entryMode = $state(/** @type {'manual'|'provider'} */ ("manual")); // Advanced (SSL / SSH / read-only) disclosure - collapsed by default. - let advancedOpen = $state(false); let name = $state(""); let host = $state("127.0.0.1"); let port = $state("5432"); @@ -465,10 +508,10 @@ /** Providers with an account flow, in the order they are offered. */ - const PROVIDER_CARDS = ["neon", "supabase", "prisma", "planetscale", "d1"]; + const PROVIDER_CARDS = ["neon", "supabase", "prisma", "planetscale", "tidb", "turso", "railway", "nile", "upstash", "d1"]; /** Names for providers a URI can identify but the catalog has no card for. */ - const PROVIDER_LABELS = { turso: "Turso", "prisma-postgres": "Prisma Postgres" }; + const PROVIDER_LABELS = { "prisma-postgres": "Prisma Postgres" }; /** The front page's paste bar. */ let quickUri = $state(""); @@ -987,7 +1030,6 @@ // An existing connection already answered "what are you connecting to", so it // opens on its details. A new one starts at the choice. step = conn ? "form" : "pick"; - advancedOpen = false; flashedFields = new Set(); error = ""; testOk = false; @@ -1006,30 +1048,30 @@ async function connectProviderConnection(conn) { error = ""; // Credentials reused from a saved connection can have been revoked in the - // provider's console since. Probe them first - connectWith reports failures - // itself, so letting it fail would toast a scary auth error a moment before - // the retry silently succeeded. - if (conn.reusedSaved) { - const probe = { - name: conn.name, - host: conn.host, - port: conn.port, - database: conn.database, - user: conn.username, - password: conn.password, - ssl: conn.ssl, - }; - let usable = true; - try { - if (conn.db_type === "mysql") await testMysqlConnection(probe); - else await testPostgresConnection(probe); - } catch { - usable = false; - } - const spec = usable - ? conn - : await providerBuildConnection(dbType, conn.reusedSaved); - await connectProviderResolved(spec); + // provider's console since. They used to be probed first with a full test + // connect, then connected again for real: two complete handshakes on every + // reuse, which against a far region (TiDB in Tokyo, Nile in us-west-2) was + // most of the wait. Now the real connect goes first, and only a rejected + // login falls back to minting fresh credentials - silently, with no auth + // toast ahead of the retry. + if (conn.reuse) { + const ref = conn.providerRef; + await connectWith( + // The saved entry as it is, with this panel's read-only choice. + { ...conn.reuse, providerRef: ref, readOnly: readOnly || conn.reuse.readOnly || undefined }, + { + onAuthFailure: async () => { + const fresh = await providerBuildConnection(dbType, ref); + // Supabase never returns a password: a rejected saved one can't be + // replaced from here, so say so rather than connect with none. + if (fresh.needs_password) { + failWith(`The saved password for ${conn.reuse.name} was rejected. Pick the database again and enter the current password.`); + return; + } + await connectProviderResolved({ ...fresh, providerRef: ref }); + }, + }, + ); return; } await connectProviderResolved(conn); @@ -1038,13 +1080,51 @@ /** * Build a SavedConnection from a resolved provider spec and connect. * @param {import('$lib/providers.js').ProviderConnection} conn + * @param {{ onAuthFailure?: () => Promise }} [opts] */ - async function connectProviderResolved(conn) { + async function connectProviderResolved(conn, opts = {}) { error = ""; // dbType is the provider id while the provider flow is showing - tag the // connection with it so the status bar can offer switching to the account's // other databases later. const providerId = PROVIDER_IDS.includes(dbType) ? dbType : undefined; + // libsql has no host/port/user: the adapter hands over the `libsql://` URL + // in `host` and the database token in `password`. + if (conn.db_type === "libsql") { + const existing = saved.find((s) => s.type === "libsql" && s.url === conn.host); + await connectWith({ + id: existing?.id ?? newConnectionId(), + type: "libsql", + name: conn.name, + url: conn.host, + authToken: conn.password || undefined, + provider: providerId, + providerRef: conn.providerRef, + readOnly: readOnly || undefined, + }, opts); + return; + } + // Redis (Upstash, Railway): the saved shape has `db` and `tls`, not a + // database name and `ssl`. + if (conn.db_type === "redis") { + const existing = saved.find( + (s) => s.type === "redis" && s.host === conn.host && s.port === conn.port, + ); + await connectWith({ + id: existing?.id ?? newConnectionId(), + type: "redis", + name: conn.name, + host: conn.host, + port: conn.port, + password: conn.password, + db: Number(conn.database) || 0, + tls: conn.ssl, + provider: providerId, + providerRef: conn.providerRef, + readOnly: readOnly || undefined, + }, opts); + return; + } const type = conn.db_type === "mysql" ? "mysql" : "postgres"; // Reuse an existing saved entry for this exact database (host + user) instead // of piling up duplicates - connectWith upserts it, keeping the saved password. @@ -1063,8 +1143,9 @@ password: conn.password, ssl: conn.ssl, provider: providerId, + providerRef: conn.providerRef, readOnly: readOnly || undefined, - }); + }, opts); } /** @@ -1376,7 +1457,6 @@ actionLabel: "Enable SSL & retry", action: () => { ssl = true; - advancedOpen = true; void handleTest(); }, }; @@ -1396,12 +1476,20 @@ e, ) ) - return { - title: "Nothing answered at that address", - hint: `Is the server running, and is ${host}:${port} the right host and port?`, - actionLabel: "Edit host", - action: () => focusField("cn-host"), - }; + // A provider connect has no host field on screen: name the address it + // actually dialled, and say what the driver said, instead of pointing at + // the empty manual form ("is 127.0.0.1: the right host"). + return failTarget + ? { + title: `Couldn't reach ${failTarget.host}`, + hint: `${failTarget.host}:${failTarget.port} did not answer. ${String(e).slice(0, 200)}`, + } + : { + title: "Nothing answered at that address", + hint: `Is the server running, and is ${host}:${port} the right host and port?`, + actionLabel: "Edit host", + action: () => focusField("cn-host"), + }; if ( /name or service not known|nodename nor servname|getaddrinfo|failed to lookup|dns/.test( e, @@ -1433,6 +1521,8 @@ catch { return conn.url || "—"; } } if (conn.type === "d1") return conn.database || conn.name || "—"; + // Redis has no database name worth showing; where it lives is the useful bit. + if (conn.type === "redis") return conn.host ? `${conn.host}${conn.port ? `:${conn.port}` : ""}` : "—"; return conn.database || "—"; } @@ -1468,13 +1558,38 @@ }); }); - function handleDelete(id) { + // Deleting also clears the connection's history, saved queries, charts and + // chats (purgeConnectionData), so it is asked first rather than done on click. + /** @type {{ id: string, name: string, fromKeyboard: boolean } | null} */ + let pendingDelete = $state(null); + let confirmDeleteOpen = $state(false); + + /** @param {string} id @param {boolean} [fromKeyboard] */ + function handleDelete(id, fromKeyboard = false) { + const conn = saved.find((c) => c.id === id); + pendingDelete = { id, name: conn?.name || conn?.database || conn?.host || conn?.filePath || "this connection", fromKeyboard }; + confirmDeleteOpen = true; + } + + function confirmDelete() { + if (!pendingDelete) return; + const { id, fromKeyboard } = pendingDelete; + pendingDelete = null; saved = removeConnection(id).sort(byLastConnected); if (id === lastId) { lastId = null; setLastConnectionId(null); } if (editingId === id) resetForm(null); + onremoved(id); + // The row that had focus no longer exists - hand it to whatever took its + // place rather than letting it fall back to the document. + if (fromKeyboard) { + void tick().then(() => { + if (savedMatches.length) focusFirstSavedRow(); + else savedSearchEl?.focus(); + }); + } } /** Open `conn` with the driver its `type` calls for. */ @@ -1496,10 +1611,25 @@ return connectPostgres(conn); } - async function connectWith(conn) { + /** The address a non-form connect dialled, for the unreachable hint. @type {{ host: string, port: number } | null} */ + let failTarget = $state(null); + + /** A driver error that means "these credentials are no good", on any engine. */ + const AUTH_FAILURE = + /password authentication failed|authentication failed|access denied|invalid password|28P01|\b1045\b|role ".*" does not exist/i; + + /** + * @param {any} conn + * @param {{ onAuthFailure?: () => Promise }} [opts] run instead of the + * error toast when the credentials are rejected (reused provider credentials + * that were revoked: mint fresh ones and try once more). + */ + async function connectWith(conn, opts = {}) { const myOp = ++opId; connecting = conn.id; error = ""; + // Provider and saved connections dial an address that isn't in the form. + failTarget = conn.provider && conn.host ? { host: conn.host, port: conn.port } : null; try { await openConnection(conn); if (myOp !== opId) return; // cancelled by the user @@ -1509,7 +1639,12 @@ open = false; await onconnected(updated, conn.id); } catch (e) { - if (myOp === opId) failWith(friendlyError(e)); + if (myOp !== opId) return; + if (opts.onAuthFailure && AUTH_FAILURE.test(String(e))) { + await opts.onAuthFailure(); + return; + } + failWith(friendlyError(e)); } finally { if (myOp === opId) connecting = null; } @@ -1599,12 +1734,48 @@ */ function onRefreshKey(e) { if (!open) return; - if (!(e.ctrlKey || e.metaKey) || e.altKey || e.shiftKey) return; - if (e.key.toLowerCase() !== "r") return; - e.preventDefault(); - e.stopPropagation(); - saved = loadSavedConnections().sort(byLastConnected); - void refreshLocal(); + if (!(e.ctrlKey || e.metaKey) || e.altKey) return; + const key = e.key.toLowerCase(); + /** @param {() => void} run */ + const take = (run) => { e.preventDefault(); e.stopPropagation(); run(); }; + // Mod+Shift+Enter: resume the last connection, the footer's "Resume". + if (e.shiftKey) { + if (key === "enter") { + const last = saved.find((c) => c.id === lastId); + if (last && !isBusy) take(() => void connectWith(last)); + } + return; + } + // The dialog's own chords, caught in the capture phase for the same reason + // as ⌘R (see below): the fields inside stop propagation, and + // Mod+N / Mod+F / Mod+L would otherwise reach the app or the webview. + if (key === "r") { + take(() => { + saved = loadSavedConnections().sort(byLastConnected); + void refreshLocal(); + }); + } else if (key === "n") { + take(() => newConnectionForm()); + } else if (key === "f") { + take(() => { + if (!railOpen) { railOpen = true; saveRail(); } + void tick().then(() => { savedSearchEl?.focus(); savedSearchEl?.select?.(); }); + }); + } else if (key === "l") { + take(() => { + // The paste field on the picker; on the form, the "Paste a URL" bar. + if (step === "pick") { + quickUriEl?.focus(); + quickUriEl?.select?.(); + } else if (entryMode === "manual" && hasFieldToggle) { + importOpen = true; + void tick().then(() => document.getElementById("cn-import-uri")?.focus()); + } else { + backToPick(); + void tick().then(() => quickUriEl?.focus()); + } + }); + } } async function refreshLocal() { @@ -2220,28 +2391,37 @@ 10px off the card's own border. h-14 with `leading-tight` on both lines puts 12px above and below the pair and 2px between them, which is what separates a title from its caption rather than stacking them. --> -
+ +
+