From fd039ce7056aceda795a5009c0e9752571e30d63 Mon Sep 17 00:00:00 2001 From: abeni16 Date: Mon, 22 Jun 2026 13:58:42 +0300 Subject: [PATCH 01/14] feat: add MongoDB support and update connection handling This commit introduces support for MongoDB as a database engine, updating the relevant types and connection dialogs. It modifies the connection handling to include MongoDB-specific connection strings and ensures the application can manage MongoDB clients. Additionally, the SQL editor is updated to handle JSON language mode for MongoDB queries, enhancing the overall functionality of the application. --- build/index.html | 6 +- src-tauri/Cargo.lock | 627 ++- src-tauri/Cargo.toml | 2 + src-tauri/src/commands.rs | 4082 ----------------- src-tauri/src/commands/connections.rs | 352 ++ src-tauri/src/commands/ddl.rs | 93 + src-tauri/src/commands/editor_meta.rs | 426 ++ src-tauri/src/commands/export_cmds.rs | 65 + src-tauri/src/commands/lint.rs | 80 + src-tauri/src/commands/mod.rs | 961 ++++ src-tauri/src/commands/mongo.rs | 321 ++ src-tauri/src/commands/query.rs | 255 + src-tauri/src/commands/table_props.rs | 585 +++ src-tauri/src/commands/veloxy.rs | 440 ++ src-tauri/src/db.rs | 63 + src-tauri/src/export.rs | 6 + src-tauri/src/lib.rs | 8 +- src-tauri/src/models.rs | 1 + src/App.tsx | 1248 +---- src/components/VeloxApp.tsx | 330 ++ src/data/types.ts | 2 +- .../components/ConnectionDialog.tsx | 1 + .../components/ConnectionsSidebarTree.tsx | 1 + .../model/components/ModelWorkspace.tsx | 2563 +++-------- .../components/ModelWorkspaceToolbar.tsx | 150 + src/features/model/hooks/useModelColumns.ts | 173 + .../model/hooks/useModelInitialization.ts | 142 + .../model/hooks/useModelWorkspaceStore.ts | 73 + .../queries/components/AskVeloxyDialog.tsx | 181 +- .../queries/components/QueryWorkspace.tsx | 5 +- .../queries/components/ResultsCellEditor.tsx | 63 + .../queries/components/ResultsGrid.tsx | 67 +- src/features/queries/components/SqlEditor.tsx | 8 +- .../components/veloxy-message-parser.ts | 130 + .../workspace/ModelWorkspaceAdapter.tsx | 45 + src/features/workspace/types.ts | 17 + src/features/workspace/workspace-utils.ts | 10 + src/hooks/useAppState.ts | 654 +++ src/lib/sql-intent.ts | 3 +- 39 files changed, 6731 insertions(+), 7508 deletions(-) delete mode 100644 src-tauri/src/commands.rs create mode 100644 src-tauri/src/commands/connections.rs create mode 100644 src-tauri/src/commands/ddl.rs create mode 100644 src-tauri/src/commands/editor_meta.rs create mode 100644 src-tauri/src/commands/export_cmds.rs create mode 100644 src-tauri/src/commands/lint.rs create mode 100644 src-tauri/src/commands/mod.rs create mode 100644 src-tauri/src/commands/mongo.rs create mode 100644 src-tauri/src/commands/query.rs create mode 100644 src-tauri/src/commands/table_props.rs create mode 100644 src-tauri/src/commands/veloxy.rs create mode 100644 src/components/VeloxApp.tsx create mode 100644 src/features/model/components/ModelWorkspaceToolbar.tsx create mode 100644 src/features/model/hooks/useModelColumns.ts create mode 100644 src/features/model/hooks/useModelInitialization.ts create mode 100644 src/features/model/hooks/useModelWorkspaceStore.ts create mode 100644 src/features/queries/components/ResultsCellEditor.tsx create mode 100644 src/features/queries/components/veloxy-message-parser.ts create mode 100644 src/features/workspace/ModelWorkspaceAdapter.tsx create mode 100644 src/features/workspace/types.ts create mode 100644 src/features/workspace/workspace-utils.ts create mode 100644 src/hooks/useAppState.ts diff --git a/build/index.html b/build/index.html index c606a09..cac414f 100644 --- a/build/index.html +++ b/build/index.html @@ -5,10 +5,10 @@ veloxdb - - + + - +
diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 52cf787..4b642dd 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -19,6 +19,19 @@ dependencies = [ "version_check", ] +[[package]] +name = "ahash" +version = "0.8.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5a15f179cd60c4584b8a8c596927aadc462e27f2ca70c04e0071964a73ba7a75" +dependencies = [ + "cfg-if", + "getrandom 0.3.4", + "once_cell", + "version_check", + "zerocopy", +] + [[package]] name = "aho-corasick" version = "1.1.4" @@ -270,6 +283,29 @@ dependencies = [ "alloc-stdlib", ] +[[package]] +name = "bson" +version = "2.15.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7969a9ba84b0ff843813e7249eed1678d9b6607ce5a3b8f0a47af3fcf7978e6e" +dependencies = [ + "ahash 0.8.12", + "base64 0.22.1", + "bitvec", + "getrandom 0.2.17", + "getrandom 0.3.4", + "hex", + "indexmap 2.13.0", + "js-sys", + "once_cell", + "rand 0.9.2", + "serde", + "serde_bytes", + "serde_json", + "time", + "uuid", +] + [[package]] name = "bstr" version = "1.12.1" @@ -464,6 +500,17 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" +[[package]] +name = "chacha20" +version = "0.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6f8d983286843e49675a4b7a2d174efe136dc93a18d69130dd18198a6c167601" +dependencies = [ + "cfg-if", + "cpufeatures 0.3.0", + "rand_core 0.10.1", +] + [[package]] name = "chrono" version = "0.4.45" @@ -509,12 +556,41 @@ version = "0.9.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c2459377285ad874054d797f3ccebf984978aa39129f6eafde5cdc8315b612f8" +[[package]] +name = "const-random" +version = "0.1.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "87e00182fe74b066627d63b85fd550ac2998d4b0bd86bfed477a0ae4c7c71359" +dependencies = [ + "const-random-macro", +] + +[[package]] +name = "const-random-macro" +version = "0.1.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f9d839f2a20b0aee515dc581a6172f2321f96cab76c1a38a4c584a194955390e" +dependencies = [ + "getrandom 0.2.17", + "once_cell", + "tiny-keccak", +] + [[package]] name = "convert_case" version = "0.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6245d59a3e82a7fc217c5828a6692dbc6dfb63a0c8c90495621f7b9d79704a0e" +[[package]] +name = "convert_case" +version = "0.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "633458d4ef8c78b72454de2d54fd6ab2e60f9e02be22f3c6104cdc8a4e0fceb9" +dependencies = [ + "unicode-segmentation", +] + [[package]] name = "cookie" version = "0.18.1" @@ -525,6 +601,16 @@ dependencies = [ "version_check", ] +[[package]] +name = "core-foundation" +version = "0.9.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91e195e091a93c46f7102ec7818a2aa394e1e1771c3ab4825963fa03e45afb8f" +dependencies = [ + "core-foundation-sys", + "libc", +] + [[package]] name = "core-foundation" version = "0.10.1" @@ -548,7 +634,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "064badf302c3194842cf2c5d61f56cc88e54a759313879cdf03abdd27d0c3b97" dependencies = [ "bitflags 2.11.0", - "core-foundation", + "core-foundation 0.10.1", "core-graphics-types", "foreign-types", "libc", @@ -561,7 +647,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3d44a101f213f6c4cdc1853d4b78aef6db6bdfa3468798cc1d9912f4735013eb" dependencies = [ "bitflags 2.11.0", - "core-foundation", + "core-foundation 0.10.1", "libc", ] @@ -583,6 +669,15 @@ dependencies = [ "libc", ] +[[package]] +name = "cpufeatures" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b2a41393f66f16b0823bb79094d54ac5fbd34ab292ddafb9a0456ac9f87d201" +dependencies = [ + "libc", +] + [[package]] name = "crc" version = "3.4.0" @@ -607,6 +702,12 @@ dependencies = [ "cfg-if", ] +[[package]] +name = "critical-section" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "790eea4361631c5e7d22598ecd5723ff611904e3344ce8720784c93e3d83d40b" + [[package]] name = "crossbeam-channel" version = "0.5.15" @@ -616,6 +717,15 @@ dependencies = [ "crossbeam-utils", ] +[[package]] +name = "crossbeam-epoch" +version = "0.9.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5b82ac4a3c2ca9c3460964f020e1402edd5753411d7737aa39c3714ad1b5420e" +dependencies = [ + "crossbeam-utils", +] + [[package]] name = "crossbeam-queue" version = "0.3.12" @@ -631,6 +741,12 @@ version = "0.8.21" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d0a5c400df2834b80a4c3327b3aad3a4c4cd4de0629063962b03235697506a28" +[[package]] +name = "crunchy" +version = "0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "460fbee9c2c2f33933d720630a6a0bac33ba7053db5344fac858d4b8952d77d5" + [[package]] name = "crypto-common" version = "0.1.7" @@ -752,6 +868,12 @@ dependencies = [ "syn 2.0.117", ] +[[package]] +name = "data-encoding" +version = "2.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a4ae5f15dda3c708c0ade84bfee31ccab44a3da4f88015ed22f63732abe300c8" + [[package]] name = "data-url" version = "0.3.2" @@ -838,13 +960,35 @@ dependencies = [ "serde_core", ] +[[package]] +name = "derive-syn-parse" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d65d7ce8132b7c0e54497a4d9a55a1c2a0912a0d786cf894472ba818fba45762" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.117", +] + +[[package]] +name = "derive-where" +version = "1.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d08b3a0bcc0d079199cd476b2cae8435016ec11d1c0986c6901c5ac223041534" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.117", +] + [[package]] name = "derive_more" version = "0.99.20" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6edb4b64a43d977b8e99788fe3a04d483834fba1215a7e02caa415b626497f7f" dependencies = [ - "convert_case", + "convert_case 0.4.0", "proc-macro2", "quote", "rustc_version", @@ -866,10 +1010,12 @@ version = "2.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "799a97264921d8623a957f6c3b9011f3b5492f557bbb7a5a19b7fa6d06ba8dcb" dependencies = [ + "convert_case 0.10.0", "proc-macro2", "quote", "rustc_version", "syn 2.0.117", + "unicode-xid", ] [[package]] @@ -1549,6 +1695,7 @@ dependencies = [ "cfg-if", "libc", "r-efi 6.0.0", + "rand_core 0.10.1", "wasip2", "wasip3", ] @@ -1717,7 +1864,7 @@ version = "0.12.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8a9ee70c43aaf417c914396645a0fa852624801b24ebb7ae78fe8272889ac888" dependencies = [ - "ahash", + "ahash 0.7.8", ] [[package]] @@ -1770,6 +1917,76 @@ version = "0.4.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70" +[[package]] +name = "hickory-net" +version = "0.26.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e2295ed2f9c31e471e1428a8f88a3f0e1f4b27c15049592138d1eebe9c35b183" +dependencies = [ + "async-trait", + "cfg-if", + "data-encoding", + "futures-channel", + "futures-io", + "futures-util", + "hickory-proto", + "idna", + "ipnet", + "jni 0.22.4", + "rand 0.10.1", + "thiserror 2.0.18", + "tinyvec", + "tokio", + "tracing", + "url", +] + +[[package]] +name = "hickory-proto" +version = "0.26.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0bab31817bfb44672a252e97fe81cd0c18d1b2cf892108922f6818820df8c643" +dependencies = [ + "data-encoding", + "idna", + "ipnet", + "jni 0.22.4", + "once_cell", + "prefix-trie", + "rand 0.10.1", + "ring", + "thiserror 2.0.18", + "tinyvec", + "tracing", + "url", +] + +[[package]] +name = "hickory-resolver" +version = "0.26.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0d58d28879ceecde6607729660c2667a081ccdc082e082675042793960f178c" +dependencies = [ + "cfg-if", + "futures-util", + "hickory-net", + "hickory-proto", + "ipconfig", + "ipnet", + "jni 0.22.4", + "moka", + "ndk-context", + "once_cell", + "parking_lot", + "rand 0.10.1", + "resolv-conf", + "smallvec", + "system-configuration", + "thiserror 2.0.18", + "tokio", + "tracing", +] + [[package]] name = "hkdf" version = "0.12.4" @@ -2114,11 +2331,27 @@ dependencies = [ "cfb", ] +[[package]] +name = "ipconfig" +version = "0.3.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4d40460c0ce33d6ce4b0630ad68ff63d6661961c48b6dba35e5a4d81cfb48222" +dependencies = [ + "socket2", + "widestring", + "windows-registry", + "windows-result 0.4.1", + "windows-sys 0.61.2", +] + [[package]] name = "ipnet" version = "2.12.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d98f6fed1fde3f8c21bc40a1abb88dd75e67924f9cffc3ef95607bad8017f8e2" +dependencies = [ + "serde", +] [[package]] name = "iri-string" @@ -2168,19 +2401,68 @@ dependencies = [ "cesu8", "cfg-if", "combine", - "jni-sys", + "jni-sys 0.3.0", "log", "thiserror 1.0.69", "walkdir", "windows-sys 0.45.0", ] +[[package]] +name = "jni" +version = "0.22.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5efd9a482cf3a427f00d6b35f14332adc7902ce91efb778580e180ff90fa3498" +dependencies = [ + "cfg-if", + "combine", + "jni-macros", + "jni-sys 0.4.1", + "log", + "simd_cesu8", + "thiserror 2.0.18", + "walkdir", + "windows-link 0.2.1", +] + +[[package]] +name = "jni-macros" +version = "0.22.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a00109accc170f0bdb141fed3e393c565b6f5e072365c3bd58f5b062591560a3" +dependencies = [ + "proc-macro2", + "quote", + "rustc_version", + "simd_cesu8", + "syn 2.0.117", +] + [[package]] name = "jni-sys" version = "0.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8eaf4bc02d17cbdd7ff4c7438cafcdf7fb9a4613313ad11b4f8fefe7d3fa0130" +[[package]] +name = "jni-sys" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c6377a88cb3910bee9b0fa88d4f42e1d2da8e79915598f65fb0c7ee14c878af2" +dependencies = [ + "jni-sys-macros", +] + +[[package]] +name = "jni-sys-macros" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "38c0b942f458fe50cdac086d2f946512305e5631e720728f2a61aabcd47a6264" +dependencies = [ + "quote", + "syn 2.0.117", +] + [[package]] name = "js-sys" version = "0.3.91" @@ -2409,6 +2691,54 @@ version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c41e0c4fef86961ac6d6f8a82609f55f31b05e4fce149ac5710e439df7619ba4" +[[package]] +name = "macro_magic" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cc33f9f0351468d26fbc53d9ce00a096c8522ecb42f19b50f34f2c422f76d21d" +dependencies = [ + "macro_magic_core", + "macro_magic_macros", + "quote", + "syn 2.0.117", +] + +[[package]] +name = "macro_magic_core" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1687dc887e42f352865a393acae7cf79d98fab6351cde1f58e9e057da89bf150" +dependencies = [ + "const-random", + "derive-syn-parse", + "macro_magic_core_macros", + "proc-macro2", + "quote", + "syn 2.0.117", +] + +[[package]] +name = "macro_magic_core_macros" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b02abfe41815b5bd98dbd4260173db2c116dda171dc0fe7838cb206333b83308" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.117", +] + +[[package]] +name = "macro_magic_macros" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "73ea28ee64b88876bf45277ed9a5817c1817df061a74f2b988971a12570e5869" +dependencies = [ + "macro_magic_core", + "quote", + "syn 2.0.117", +] + [[package]] name = "markup5ever" version = "0.14.1" @@ -2518,6 +2848,99 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "moka" +version = "0.12.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "957228ad12042ee839f93c8f257b62b4c0ab5eaae1d4fa60de53b27c9d7c5046" +dependencies = [ + "crossbeam-channel", + "crossbeam-epoch", + "crossbeam-utils", + "equivalent", + "parking_lot", + "portable-atomic", + "smallvec", + "tagptr", + "uuid", +] + +[[package]] +name = "mongocrypt" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8da0cd419a51a5fb44819e290fbdb0665a54f21dead8923446a799c7f4d26ad9" +dependencies = [ + "bson", + "mongocrypt-sys", + "once_cell", + "serde", +] + +[[package]] +name = "mongocrypt-sys" +version = "0.1.5+1.15.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "224484c5d09285a7b8cb0a0c117e847ebd14cb6e4470ecf68cdb89c503b0edb9" + +[[package]] +name = "mongodb" +version = "3.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "276ba0cd571553d1f6936c6f180964776ece6ab7507dc8765f8a9c9c49d8cd00" +dependencies = [ + "base64 0.22.1", + "bitflags 2.11.0", + "bson", + "derive-where", + "derive_more 2.1.1", + "futures-core", + "futures-io", + "futures-util", + "hex", + "hickory-net", + "hickory-proto", + "hickory-resolver", + "hmac", + "macro_magic", + "md-5", + "mongocrypt", + "mongodb-internal-macros", + "pbkdf2", + "percent-encoding", + "rand 0.9.2", + "rustc_version_runtime", + "rustls", + "serde", + "serde_bytes", + "serde_with", + "sha1", + "sha2", + "socket2", + "stringprep", + "strsim", + "take_mut", + "thiserror 2.0.18", + "tokio", + "tokio-rustls", + "tokio-util", + "typed-builder", + "uuid", + "webpki-roots 1.0.7", +] + +[[package]] +name = "mongodb-internal-macros" +version = "3.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "99ceb1a9a1018e470077ec94cf3a8c2d0e6da542b2c05ea95a59a0a627147375" +dependencies = [ + "macro_magic", + "proc-macro2", + "quote", + "syn 2.0.117", +] + [[package]] name = "muda" version = "0.19.1" @@ -2546,7 +2969,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c3f42e7bbe13d351b6bead8286a43aac9534b82bd3cc43e47037f012ebfd62d4" dependencies = [ "bitflags 2.11.0", - "jni-sys", + "jni-sys 0.3.0", "log", "ndk-sys", "num_enum", @@ -2554,13 +2977,19 @@ dependencies = [ "thiserror 1.0.69", ] +[[package]] +name = "ndk-context" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "27b02d87554356db9e9a873add8782d4ea6e3e58ea071a9adb9a2e8ddb884a8b" + [[package]] name = "ndk-sys" version = "0.6.0+11769913" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ee6cda3051665f1fb8d9e08fc35c96d5a244fb1be711a03b71118828afc9a873" dependencies = [ - "jni-sys", + "jni-sys 0.3.0", ] [[package]] @@ -2878,6 +3307,10 @@ name = "once_cell" version = "1.21.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" +dependencies = [ + "critical-section", + "portable-atomic", +] [[package]] name = "option-ext" @@ -2948,6 +3381,15 @@ dependencies = [ "windows-link 0.2.1", ] +[[package]] +name = "pbkdf2" +version = "0.12.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8ed6a7761f76e3b9f92dfb0a60a6a6477c61024b775147ff0973a02653abaf2" +dependencies = [ + "digest", +] + [[package]] name = "pem-rfc7468" version = "0.7.0" @@ -3249,6 +3691,12 @@ dependencies = [ "bstr", ] +[[package]] +name = "portable-atomic" +version = "1.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c33a9471896f1c69cecef8d20cbe2f7accd12527ce60845ff44c153bb2a21b49" + [[package]] name = "postgres-protocol" version = "0.6.10" @@ -3308,6 +3756,17 @@ version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "925383efa346730478fb4838dbe9137d2a47675ad789c546d150a6e1dd4ab31c" +[[package]] +name = "prefix-trie" +version = "0.8.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4cf6e3177f0684016a5c209b00882e15f8bdd3f3bb48f0491df10cd102d0c6e7" +dependencies = [ + "either", + "ipnet", + "num-traits", +] + [[package]] name = "prettyplease" version = "0.2.37" @@ -3550,6 +4009,17 @@ dependencies = [ "rand_core 0.9.5", ] +[[package]] +name = "rand" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d2e8e8bcc7961af1fdac401278c6a831614941f6164ee3bf4ce61b7edb162207" +dependencies = [ + "chacha20", + "getrandom 0.4.2", + "rand_core 0.10.1", +] + [[package]] name = "rand_chacha" version = "0.2.2" @@ -3607,6 +4077,12 @@ dependencies = [ "getrandom 0.3.4", ] +[[package]] +name = "rand_core" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "63b8176103e19a2643978565ca18b50549f6101881c443590420e4dc998a3c69" + [[package]] name = "rand_hc" version = "0.2.0" @@ -3793,6 +4269,12 @@ dependencies = [ "web-sys", ] +[[package]] +name = "resolv-conf" +version = "0.7.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e061d1b48cb8d38042de4ae0a7a6401009d6143dc80d2e2d6f31f0bdd6470c7" + [[package]] name = "resvg" version = "0.43.0" @@ -3943,12 +4425,23 @@ dependencies = [ "semver", ] +[[package]] +name = "rustc_version_runtime" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2dd18cd2bae1820af0b6ad5e54f4a51d0f3fcc53b05f845675074efcc7af071d" +dependencies = [ + "rustc_version", + "semver", +] + [[package]] name = "rustls" version = "0.23.38" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "69f9466fb2c14ea04357e91413efb882e2a6d4a406e625449bc0a5d360d53a21" dependencies = [ + "log", "once_cell", "ring", "rustls-pki-types", @@ -4158,6 +4651,16 @@ dependencies = [ "typeid", ] +[[package]] +name = "serde_bytes" +version = "0.11.19" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a5d440709e79d88e51ac01c4b72fc6cb7314017bb7da9eeff678aa94c10e3ea8" +dependencies = [ + "serde", + "serde_core", +] + [[package]] name = "serde_core" version = "1.0.228" @@ -4195,6 +4698,7 @@ version = "1.0.149" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "83fc039473c5595ace860d8c4fafa220ff474b3fc6bfdb4293327f1a37e94d86" dependencies = [ + "indexmap 2.13.0", "itoa", "memchr", "serde", @@ -4322,7 +4826,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba" dependencies = [ "cfg-if", - "cpufeatures", + "cpufeatures 0.2.17", "digest", ] @@ -4333,7 +4837,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" dependencies = [ "cfg-if", - "cpufeatures", + "cpufeatures 0.2.17", "digest", ] @@ -4369,6 +4873,16 @@ version = "0.3.8" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e320a6c5ad31d271ad523dcf3ad13e2767ad8b1cb8f047f75a8aeaf8da139da2" +[[package]] +name = "simd_cesu8" +version = "1.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "94f90157bb87cddf702797c5dadfa0be7d266cdf49e22da2fcaa32eff75b2c33" +dependencies = [ + "rustc_version", + "simdutf8", +] + [[package]] name = "simdutf8" version = "0.1.5" @@ -4841,6 +5355,27 @@ dependencies = [ "syn 2.0.117", ] +[[package]] +name = "system-configuration" +version = "0.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a13f3d0daba03132c0aa9767f98351b3488edc2c100cda2d2ec2b04f3d8d3c8b" +dependencies = [ + "bitflags 2.11.0", + "core-foundation 0.9.4", + "system-configuration-sys", +] + +[[package]] +name = "system-configuration-sys" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e1d1b10ced5ca923a1fcb8d03e96b8d3268065d724548c0211415ff6ac6bac4" +dependencies = [ + "core-foundation-sys", + "libc", +] + [[package]] name = "system-deps" version = "6.2.2" @@ -4854,6 +5389,18 @@ dependencies = [ "version-compare", ] +[[package]] +name = "tagptr" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7b2093cf4c8eb1e67749a6762251bc9cd836b6fc171623bd0a9d324d37af2417" + +[[package]] +name = "take_mut" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f764005d11ee5f36500a149ace24e00e3da98b0158b3e2d53a7495660d3f4d60" + [[package]] name = "tao" version = "0.35.0" @@ -4862,7 +5409,7 @@ checksum = "1cf65722394c2ac443e80120064987f8914ee1d4e4e36e63cdf10f2990f01159" dependencies = [ "bitflags 2.11.0", "block2", - "core-foundation", + "core-foundation 0.10.1", "core-graphics", "crossbeam-channel", "dbus", @@ -4872,7 +5419,7 @@ dependencies = [ "gdkwayland-sys", "gdkx11-sys", "gtk", - "jni", + "jni 0.21.1", "libc", "log", "ndk", @@ -4934,7 +5481,7 @@ dependencies = [ "gtk", "heck 0.5.0", "http", - "jni", + "jni 0.21.1", "libc", "log", "mime", @@ -5137,7 +5684,7 @@ dependencies = [ "dpi", "gtk", "http", - "jni", + "jni 0.21.1", "objc2", "objc2-ui-kit", "objc2-web-kit", @@ -5160,7 +5707,7 @@ checksum = "2cadb13dad0c681e1e0a2c49ae488f0e2906ded3d57e7a0017f4aaf46e387117" dependencies = [ "gtk", "http", - "jni", + "jni 0.21.1", "log", "objc2", "objc2-app-kit", @@ -5323,6 +5870,15 @@ dependencies = [ "time-core", ] +[[package]] +name = "tiny-keccak" +version = "2.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2c9d3793400a45f954c52e73d068316d76b6f4e36977e3fcebb13a2721e80237" +dependencies = [ + "crunchy", +] + [[package]] name = "tiny-skia" version = "0.11.4" @@ -5493,7 +6049,9 @@ checksum = "9ae9cec805b01e8fc3fd2fe289f89149a9b66dd16786abd8b19cfa7b48cb0098" dependencies = [ "bytes", "futures-core", + "futures-io", "futures-sink", + "futures-util", "pin-project-lite", "tokio", ] @@ -5723,6 +6281,26 @@ dependencies = [ "core_maths", ] +[[package]] +name = "typed-builder" +version = "0.22.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "398a3a3c918c96de527dc11e6e846cd549d4508030b8a33e1da12789c856b81a" +dependencies = [ + "typed-builder-macro", +] + +[[package]] +name = "typed-builder-macro" +version = "0.22.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0e48cea23f68d1f78eb7bc092881b6bb88d3d6b5b7e6234f6f9c911da1ffb221" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.117", +] + [[package]] name = "typeid" version = "1.0.3" @@ -5950,6 +6528,7 @@ name = "veloxdb" version = "0.1.0-8" dependencies = [ "base64 0.22.1", + "bson", "chrono", "csv", "deadpool-postgres", @@ -5957,6 +6536,7 @@ dependencies = [ "hex", "keyring", "log", + "mongodb", "printpdf", "rand 0.8.5", "reqwest 0.12.28", @@ -6363,6 +6943,12 @@ dependencies = [ "web-sys", ] +[[package]] +name = "widestring" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72069c3113ab32ab29e5584db3c6ec55d416895e60715417b5b883a357c3e471" + [[package]] name = "winapi" version = "0.3.9" @@ -6512,6 +7098,17 @@ dependencies = [ "windows-link 0.1.3", ] +[[package]] +name = "windows-registry" +version = "0.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "02752bf7fbdcce7f2a27a742f798510f3e5ad88dbe84871e5168e2120c3d5720" +dependencies = [ + "windows-link 0.2.1", + "windows-result 0.4.1", + "windows-strings 0.5.1", +] + [[package]] name = "windows-result" version = "0.3.4" @@ -7003,7 +7600,7 @@ dependencies = [ "gtk", "http", "javascriptcore-rs", - "jni", + "jni 0.21.1", "libc", "ndk", "objc2", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 5e16056..654aecd 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -46,3 +46,5 @@ printpdf = "0.7" reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls", "stream"] } futures-util = "0.3" chrono = "0.4.45" +mongodb = "3" +bson = "2" diff --git a/src-tauri/src/commands.rs b/src-tauri/src/commands.rs deleted file mode 100644 index 6bd109f..0000000 --- a/src-tauri/src/commands.rs +++ /dev/null @@ -1,4082 +0,0 @@ -use std::collections::{BTreeMap, HashMap, HashSet}; -use std::sync::atomic::{AtomicBool, Ordering}; -use std::sync::Arc; -use std::time::Instant; - -use futures_util::StreamExt; -use serde_json::Value; -use tauri::{AppHandle, Emitter, State}; -use sqlx::{Column, Decode, Row, Type}; -use sqlx::mysql::{MySql, MySqlRow}; -use sqlx::sqlite::{Sqlite, SqliteRow}; -use tokio_postgres::SimpleQueryMessage; -use uuid::Uuid; - -use crate::db::{ - build_mysql_pool, build_mysql_pool_custom, build_pool, build_pool_custom, build_sqlite_pool, - disconnect_connection, drop_pool, get_or_create_mysql_pool, get_or_create_sqlite_pool, - list_connections, load_connection, persist_connection_with_password, quote_identifier, - refresh_connection_pools, require_safe_identifier, resolve_connection_engine, - with_pool_client_retry, AppState, - DEFAULT_MYSQL_PORT, MAX_QUERY_ROWS, -}; -use crate::credentials; -use crate::pg_error::{error_line_column, map_pg_err}; -use crate::sql_split::split_sql_statements; -use crate::models::{ - AskVeloxyChatRequest, AskVeloxyChatResponse, AskVeloxyConversationMessage, - AskVeloxyConversationResponse, AskVeloxyDbContextCache, AskVeloxyRequest, AskVeloxyResponse, - AskVeloxyTableRef, AskVeloxyTokenStats, ColumnInfo, ColumnProperties, ConnectionInput, - ConnectionSummary, DatabaseInfo, DatabaseEngine, DdlBatchRequest, DdlStatementRequest, - ForeignKeyEdge, IndexInfo, LintSqlRequest, LintSqlResult, QueryEditorColumn, - QueryEditorFunction, QueryEditorMetadata, QueryEditorTable, QueryRequest, QueryResult, - SchemaRequest, SqlDiagnostic, StoredConnection, SwitchDatabaseRequest, TableIndexesResult, - TableInfo, TablePropertiesApplyRequest, VeloxyStreamChunk, -}; -use crate::export::{ - DiagramExportRequest, ExportQueryRequest, - export_diagram_to_png, export_results_csv, export_results_json, -}; -use crate::ssh_tunnel::SshTunnel; - -/// Cap FK rows returned to the UI to keep IPC payloads bounded. -const MAX_FOREIGN_KEY_ROWS: i64 = 5000; - -/// Cap index rows per table (fetch limit + 1 to detect truncation). -const MAX_TABLE_INDEX_ROWS: i64 = 500; -const MAX_EDITOR_TABLES: i64 = 150; -const MAX_EDITOR_COLUMNS_PER_TABLE: i64 = 60; -const MAX_EDITOR_FUNCTIONS: i64 = 200; -const MAX_LINT_SQL_BYTES: usize = 65_536; -const ASK_VELOXY_MAX_CONTEXT_TABLES: usize = 8; -const ASK_VELOXY_MAX_CONTEXT_COLUMNS: usize = 18; -const ASK_VELOXY_MAX_CONTEXT_RELATIONSHIPS: usize = 36; -const ASK_VELOXY_SCHEMA_CHAR_BUDGET: usize = 6_000; -const ASK_VELOXY_PROMPT_CHAR_BUDGET: usize = 12_000; -const ASK_VELOXY_MAX_HISTORY_MESSAGES: usize = 30; -const ASK_VELOXY_MAX_CHAT_TOKENS: u32 = 10_000; - -fn mysql_decode_error(context: &str, column_name: &str, index: Option, detail: &str) -> String { - match index { - Some(idx) => format!( - "MySQL decode error in {} at column '{}' (index {}): {}", - context, column_name, idx, detail - ), - None => format!( - "MySQL decode error in {} at column '{}': {}", - context, column_name, detail - ), - } -} - -fn sqlite_decode_error(context: &str, column_name: &str, index: Option, detail: &str) -> String { - match index { - Some(idx) => format!( - "SQLite decode error in {} at column '{}' (index {}): {}", - context, column_name, idx, detail - ), - None => format!( - "SQLite decode error in {} at column '{}': {}", - context, column_name, detail - ), - } -} - -fn mysql_get_idx(row: &MySqlRow, index: usize, column_name: &str, context: &str) -> Result -where - for<'r> T: Decode<'r, MySql> + Type, -{ - row.try_get::(index) - .map_err(|error| mysql_decode_error(context, column_name, Some(index), &error.to_string())) -} - -fn sqlite_get_idx(row: &SqliteRow, index: usize, column_name: &str, context: &str) -> Result -where - for<'r> T: Decode<'r, Sqlite> + Type, -{ - row.try_get::(index) - .map_err(|error| sqlite_decode_error(context, column_name, Some(index), &error.to_string())) -} - -fn sqlite_get_name(row: &SqliteRow, column_name: &str, context: &str) -> Result -where - for<'r> T: Decode<'r, Sqlite> + Type, -{ - row.try_get::(column_name).map_err(|error| { - format!( - "SQLite decode error in {} at column '{}': {}", - context, column_name, error - ) - }) -} - -fn database_name_from_mysql_value( - value: Option, - context: &str, -) -> Result { - let name = value - .filter(|value| !value.is_empty()) - .ok_or_else(|| format!("{context} returned an empty database name"))?; - Ok(name) -} - -fn mysql_database_name_from_row(row: &MySqlRow, context: &str) -> Result { - let value = mysql_value_to_string(row, 0, "Database", context)?; - database_name_from_mysql_value(value, context) -} - -fn mysql_value_to_string(row: &MySqlRow, index: usize, column_name: &str, context: &str) -> Result, String> { - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.to_string())); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.to_string())); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.to_string())); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.to_string())); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.to_string())); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.to_string())); - } - if let Ok(value) = row.try_get::>, _>(index) { - return Ok(value.map(|v| v.format("%Y-%m-%d %H:%M:%S").to_string())); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.format("%Y-%m-%d %H:%M:%S").to_string())); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.to_string())); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.to_string())); - } - if let Ok(value) = row.try_get::>, _>(index) { - return Ok(value.map(|v| decode_mysql_bytes_as_string(&v))); - } - Err(mysql_decode_error( - context, - column_name, - Some(index), - "unsupported value type", - )) -} - -fn decode_mysql_bytes_as_string(bytes: &[u8]) -> String { - String::from_utf8_lossy(bytes).into_owned() -} - -/// Like [`mysql_value_to_string`] but encodes raw bytes as hex (for ad-hoc query grids). -fn mysql_value_to_display_string( - row: &MySqlRow, - index: usize, - column_name: &str, - context: &str, -) -> Result, String> { - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.to_string())); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.to_string())); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.to_string())); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.to_string())); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.to_string())); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.to_string())); - } - if let Ok(value) = row.try_get::>, _>(index) { - return Ok(value.map(|v| v.format("%Y-%m-%d %H:%M:%S").to_string())); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.format("%Y-%m-%d %H:%M:%S").to_string())); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.to_string())); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.to_string())); - } - if let Ok(value) = row.try_get::>, _>(index) { - return Ok(value.map(|v| format!("0x{}", hex::encode(v)))); - } - Err(mysql_decode_error( - context, - column_name, - Some(index), - "unsupported value type", - )) -} - -fn mysql_get_string(row: &MySqlRow, index: usize, column_name: &str, context: &str) -> Result { - let value = mysql_value_to_string(row, index, column_name, context)?; - value.ok_or_else(|| { - mysql_decode_error(context, column_name, Some(index), "unexpected null value") - }) -} - -fn mysql_get_optional_string( - row: &MySqlRow, - index: usize, - column_name: &str, - context: &str, -) -> Result, String> { - mysql_value_to_string(row, index, column_name, context) -} - -fn sqlite_value_to_string(row: &SqliteRow, index: usize, column_name: &str, context: &str) -> Result, String> { - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.to_string())); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.to_string())); - } - if let Ok(value) = row.try_get::, _>(index) { - return Ok(value.map(|v| v.to_string())); - } - if let Ok(value) = row.try_get::>, _>(index) { - return Ok(value.map(|v| format!("0x{}", hex::encode(v)))); - } - Err(sqlite_decode_error( - context, - column_name, - Some(index), - "unsupported value type", - )) -} - -fn is_row_returning_sql(sql: &str) -> bool { - let trimmed = sql.trim_start(); - let upper = trimmed.to_uppercase(); - upper.starts_with("SELECT") - || upper.starts_with("WITH") - || upper.starts_with("SHOW") - || upper.starts_with("EXPLAIN") - || upper.starts_with("DESCRIBE") - || upper.starts_with("DESC") - || upper.starts_with("PRAGMA") - || upper.starts_with("VALUES") - || upper.starts_with("TABLE ") -} - -fn map_mysql_rows( - rows: Vec, - max_query_rows: usize, -) -> Result<(Vec, Vec>>, usize, bool), String> { - let mut columns: Vec = Vec::new(); - if let Some(first) = rows.first() { - columns = first - .columns() - .iter() - .map(|column| column.name().to_string()) - .collect(); - } - let total_rows = rows.len(); - let mut mapped_rows = Vec::new(); - for row in rows.into_iter().take(max_query_rows) { - let mut mapped_row = BTreeMap::new(); - for (index, column_name) in columns.iter().enumerate() { - let value = mysql_value_to_display_string(&row, index, column_name, "run_query")?; - mapped_row.insert(column_name.clone(), value); - } - mapped_rows.push(mapped_row); - } - Ok((columns, mapped_rows, total_rows, total_rows > max_query_rows)) -} - -fn map_sqlite_rows( - rows: Vec, - max_query_rows: usize, -) -> Result<(Vec, Vec>>, usize, bool), String> { - let mut columns: Vec = Vec::new(); - if let Some(first) = rows.first() { - columns = first - .columns() - .iter() - .map(|column| column.name().to_string()) - .collect(); - } - let total_rows = rows.len(); - let mut mapped_rows = Vec::new(); - for row in rows.into_iter().take(max_query_rows) { - let mut mapped_row = BTreeMap::new(); - for (index, column_name) in columns.iter().enumerate() { - let value = sqlite_value_to_string(&row, index, column_name, "run_query")?; - mapped_row.insert(column_name.clone(), value); - } - mapped_rows.push(mapped_row); - } - Ok((columns, mapped_rows, total_rows, total_rows > max_query_rows)) -} - -async fn run_query_mysql_or_sqlite( - app: &AppHandle, - state: &AppState, - connection_id: &str, - sql: &str, - max_query_rows: usize, - engine: DatabaseEngine, -) -> Result { - let started_at = Instant::now(); - let statements = split_sql_statements(sql); - if statements.is_empty() { - return Err("Enter a SQL statement before running the query.".to_string()); - } - - let mut columns = Vec::new(); - let mut rows = Vec::new(); - let mut total_rows = 0usize; - let mut truncated = false; - let mut command_tag: Option = None; - - match engine { - DatabaseEngine::Mysql => { - let pool = get_or_create_mysql_pool(app, state, connection_id).await?; - let mut conn = pool - .acquire() - .await - .map_err(|error| error.to_string())?; - for statement in statements { - if is_row_returning_sql(&statement) { - let fetched = sqlx::query(&statement) - .fetch_all(&mut *conn) - .await - .map_err(|error| error.to_string())?; - let mapped = map_mysql_rows(fetched, max_query_rows)?; - columns = mapped.0; - rows = mapped.1; - total_rows = mapped.2; - truncated = mapped.3; - command_tag = None; - } else { - let result = sqlx::query(&statement) - .execute(&mut *conn) - .await - .map_err(|error| error.to_string())?; - let affected = result.rows_affected(); - command_tag = Some(affected); - if rows.is_empty() { - total_rows = affected as usize; - } - } - } - } - DatabaseEngine::Sqlite => { - let pool = get_or_create_sqlite_pool(app, state, connection_id).await?; - let mut conn = pool - .acquire() - .await - .map_err(|error| error.to_string())?; - for statement in statements { - if is_row_returning_sql(&statement) { - let fetched = sqlx::query(&statement) - .fetch_all(&mut *conn) - .await - .map_err(|error| error.to_string())?; - let mapped = map_sqlite_rows(fetched, max_query_rows)?; - columns = mapped.0; - rows = mapped.1; - total_rows = mapped.2; - truncated = mapped.3; - command_tag = None; - } else { - let result = sqlx::query(&statement) - .execute(&mut *conn) - .await - .map_err(|error| error.to_string())?; - let affected = result.rows_affected(); - command_tag = Some(affected); - if rows.is_empty() { - total_rows = affected as usize; - } - } - } - } - DatabaseEngine::Postgres => { - return Err("Internal engine routing error.".to_string()); - } - } - - Ok(QueryResult { - columns, - row_count: if rows.is_empty() { - total_rows - } else { - rows.len() - }, - rows, - execution_ms: started_at.elapsed().as_millis(), - truncated, - command_tag, - }) -} - -#[tauri::command] -pub async fn connect_db( - app: AppHandle, - state: State<'_, AppState>, - input: ConnectionInput, -) -> Result { - let connection_id = input - .id - .clone() - .unwrap_or_else(|| Uuid::new_v4().to_string()); - - match input.engine { - DatabaseEngine::Postgres => { - let pool = if let Some(ref ssh_config) = input.ssh_config { - if ssh_config.is_active() { - let tunnel = match SshTunnel::connect(ssh_config, &input.host, input.port).await { - Ok(tunnel) => tunnel, - Err(e) => return Err(format!("SSH tunnel failed: {}", e)), - }; - let local_port = tunnel.local_port; - state - .ssh_tunnels - .write() - .await - .insert(connection_id.clone(), tunnel); - build_pool_custom("127.0.0.1", local_port, &input)? - } else { - build_pool(&input)? - } - } else { - build_pool(&input)? - }; - - let client = match pool.get().await { - Ok(client) => client, - Err(e) => { - drop_pool(&state, &connection_id).await; - return Err(e.to_string()); - } - }; - - if let Err(e) = client.simple_query("select 1").await { - drop_pool(&state, &connection_id).await; - return Err(map_pg_err(e, None)); - } - - state - .pools - .write() - .await - .insert(connection_id.clone(), pool); - } - DatabaseEngine::Mysql => { - let pool = if let Some(ref ssh_config) = input.ssh_config { - if ssh_config.is_active() { - let remote_port = if input.port == 0 { DEFAULT_MYSQL_PORT } else { input.port }; - let tunnel = match SshTunnel::connect(ssh_config, &input.host, remote_port).await { - Ok(tunnel) => tunnel, - Err(e) => return Err(format!("SSH tunnel failed: {}", e)), - }; - let local_port = tunnel.local_port; - state - .ssh_tunnels - .write() - .await - .insert(connection_id.clone(), tunnel); - build_mysql_pool_custom("127.0.0.1", local_port, &input).await? - } else { - build_mysql_pool(&input).await? - } - } else { - build_mysql_pool(&input).await? - }; - - sqlx::query("select 1") - .execute(&pool) - .await - .map_err(|e| e.to_string())?; - - state - .mysql_pools - .write() - .await - .insert(connection_id.clone(), pool); - } - DatabaseEngine::Sqlite => { - let pool = build_sqlite_pool(&input).await?; - sqlx::query("select 1") - .execute(&pool) - .await - .map_err(|e| e.to_string())?; - state - .sqlite_pools - .write() - .await - .insert(connection_id.clone(), pool); - } - } - - let stored_connection = StoredConnection::from_input(connection_id.clone(), input.clone()); - persist_connection_with_password(&app, &stored_connection, &input.password)?; - - *state.active_connection_id.write().await = Some(connection_id); - - Ok(stored_connection.summary()) -} - -#[tauri::command] -pub async fn list_connections_command(app: AppHandle) -> Result, String> { - list_connections(&app) -} - -#[tauri::command] -pub async fn set_active_connection( - app: AppHandle, - state: State<'_, AppState>, - connection_id: String, -) -> Result { - let stored_connection = load_connection(&app, &connection_id)? - .ok_or_else(|| "Stored connection details were not found.".to_string())?; - - match stored_connection.engine { - DatabaseEngine::Postgres => { - with_pool_client_retry(&app, &state, &connection_id, (), |client, ()| async move { - client - .simple_query("select 1") - .await - .map_err(|error| map_pg_err(error, None))?; - Ok(()) - }) - .await?; - } - DatabaseEngine::Mysql => { - let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; - sqlx::query("select 1") - .execute(&pool) - .await - .map_err(|error| error.to_string())?; - } - DatabaseEngine::Sqlite => { - let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; - sqlx::query("select 1") - .execute(&pool) - .await - .map_err(|error| error.to_string())?; - } - } - - *state.active_connection_id.write().await = Some(connection_id); - - Ok(stored_connection.summary()) -} - -#[tauri::command] -pub async fn ping_connection( - app: AppHandle, - state: State<'_, AppState>, - connection_id: String, -) -> Result<(), String> { - let stored_connection = load_connection(&app, &connection_id)? - .ok_or_else(|| "Stored connection details were not found.".to_string())?; - match stored_connection.engine { - DatabaseEngine::Postgres => { - with_pool_client_retry(&app, &state, &connection_id, (), |client, ()| async move { - client - .simple_query("select 1") - .await - .map_err(|error| error.to_string())?; - Ok(()) - }) - .await - } - DatabaseEngine::Mysql => { - let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; - sqlx::query("select 1") - .execute(&pool) - .await - .map_err(|error| error.to_string())?; - Ok(()) - } - DatabaseEngine::Sqlite => { - let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; - sqlx::query("select 1") - .execute(&pool) - .await - .map_err(|error| error.to_string())?; - Ok(()) - } - } -} - -#[tauri::command] -pub async fn refresh_connection( - app: AppHandle, - state: State<'_, AppState>, - connection_id: String, -) -> Result<(), String> { - refresh_connection_pools(&app, &state, &connection_id).await -} - -#[tauri::command] -pub async fn disconnect_db( - state: State<'_, AppState>, - connection_id: String, -) -> Result<(), String> { - disconnect_connection(&state, &connection_id).await; - Ok(()) -} - -/// Renames a saved connection without affecting the active pool or SSH tunnel. -#[tauri::command] -pub async fn rename_connection( - app: AppHandle, - connection_id: String, - new_name: String, -) -> Result { - crate::db::rename_connection_in_store(&app, &connection_id, &new_name) -} - -#[tauri::command] -pub async fn delete_connection( - app: AppHandle, - state: State<'_, AppState>, - connection_id: String, -) -> Result<(), String> { - disconnect_connection(&state, &connection_id).await; - if let Err(e) = credentials::delete_password(&connection_id) { - log::warn!("Failed to delete keychain entry for {}: {}", connection_id, e); - } - crate::db::delete_connection_from_store(&app, &connection_id)?; - Ok(()) -} - -#[tauri::command] -pub async fn list_databases( - app: AppHandle, - state: State<'_, AppState>, - connection_id: Option, -) -> Result, String> { - let (connection_id, engine) = resolve_connection_engine(&app, &state, connection_id).await?; - - match engine { - DatabaseEngine::Postgres => { - with_pool_client_retry(&app, &state, &connection_id, (), |client, ()| async move { - let rows = client - .query( - "select datname from pg_database where datistemplate = false and has_database_privilege(datname, 'CONNECT') order by datname", - &[], - ) - .await - .map_err(|error| map_pg_err(error, None))?; - - Ok(rows - .into_iter() - .map(|row| { - let name: String = row.get(0); - DatabaseInfo { name } - }) - .collect()) - }) - .await - } - DatabaseEngine::Mysql => { - let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; - let rows = sqlx::query("show databases") - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - let mut databases = Vec::with_capacity(rows.len()); - for row in rows { - let name = mysql_database_name_from_row(&row, "list_databases")?; - databases.push(DatabaseInfo { name }); - } - Ok(databases) - } - DatabaseEngine::Sqlite => Ok(vec![DatabaseInfo { - name: "main".to_string(), - }]), - } -} - -#[tauri::command] -pub async fn switch_database( - app: AppHandle, - state: State<'_, AppState>, - input: SwitchDatabaseRequest, -) -> Result { - let mut stored_connection = load_connection(&app, &input.connection_id)? - .ok_or_else(|| "Stored connection details were not found.".to_string())?; - - if stored_connection.engine == DatabaseEngine::Sqlite { - return Err("Switch database is not supported for SQLite connections.".to_string()); - } - - drop_pool(&state, &input.connection_id).await; - - stored_connection.database = input.database.clone(); - stored_connection.connected_at = crate::models::timestamp_string(); - persist_connection_with_password(&app, &stored_connection, &stored_connection.password.clone().unwrap_or_default())?; - - let connection_input = stored_connection.to_input(); - - match connection_input.engine { - DatabaseEngine::Postgres => { - let pool = if let Some(ref ssh_config) = connection_input.ssh_config { - if ssh_config.is_active() { - let tunnel = match SshTunnel::connect(ssh_config, &connection_input.host, connection_input.port).await { - Ok(tunnel) => tunnel, - Err(e) => return Err(format!("SSH tunnel failed: {}", e)), - }; - let local_port = tunnel.local_port; - state - .ssh_tunnels - .write() - .await - .insert(input.connection_id.clone(), tunnel); - build_pool_custom("127.0.0.1", local_port, &connection_input)? - } else { - build_pool(&connection_input)? - } - } else { - build_pool(&connection_input)? - }; - - let client = match pool.get().await { - Ok(client) => client, - Err(e) => { - drop_pool(&state, &input.connection_id).await; - return Err(e.to_string()); - } - }; - - if let Err(e) = client.simple_query("select 1").await { - drop_pool(&state, &input.connection_id).await; - return Err(map_pg_err(e, None)); - } - - state - .pools - .write() - .await - .insert(input.connection_id.clone(), pool); - } - DatabaseEngine::Mysql => { - let pool = if let Some(ref ssh_config) = connection_input.ssh_config { - if ssh_config.is_active() { - let remote_port = if connection_input.port == 0 { - DEFAULT_MYSQL_PORT - } else { - connection_input.port - }; - let tunnel = match SshTunnel::connect(ssh_config, &connection_input.host, remote_port).await { - Ok(tunnel) => tunnel, - Err(e) => return Err(format!("SSH tunnel failed: {}", e)), - }; - let local_port = tunnel.local_port; - state - .ssh_tunnels - .write() - .await - .insert(input.connection_id.clone(), tunnel); - build_mysql_pool_custom("127.0.0.1", local_port, &connection_input).await? - } else { - build_mysql_pool(&connection_input).await? - } - } else { - build_mysql_pool(&connection_input).await? - }; - - sqlx::query("select 1") - .execute(&pool) - .await - .map_err(|e| e.to_string())?; - - state - .mysql_pools - .write() - .await - .insert(input.connection_id.clone(), pool); - } - DatabaseEngine::Sqlite => {} - } - - *state.active_connection_id.write().await = Some(input.connection_id); - - Ok(stored_connection.summary()) -} - -#[tauri::command] -pub async fn run_query( - app: AppHandle, - state: State<'_, AppState>, - input: QueryRequest, -) -> Result { - let (connection_id, engine) = resolve_connection_engine(&app, &state, input.connection_id).await?; - let sql = input.sql.trim().to_string(); - - if sql.is_empty() { - return Err("Enter a SQL statement before running the query.".to_string()); - } - - if !input.allow_write.unwrap_or(false) && !is_read_only_sql(&sql) { - return Err( - "This statement modifies data or schema. Confirm execution in the editor, \ - or use the model/DDL workflow for schema changes." - .to_string(), - ); - } - - let max_query_rows = input.max_rows.unwrap_or(MAX_QUERY_ROWS); - - match engine { - DatabaseEngine::Postgres => { - with_pool_client_retry(&app, &state, &connection_id, sql, |client, sql| async move { - let started_at = Instant::now(); - let messages = client - .simple_query(&sql) - .await - .map_err(|error| map_pg_err(error, Some(sql.as_str())))?; - - let mut columns = Vec::new(); - let mut rows = Vec::new(); - let mut total_rows = 0usize; - let mut command_tag = None; - - for message in messages { - match message { - SimpleQueryMessage::RowDescription(description) => { - if columns.is_empty() { - columns = description - .iter() - .map(|column| column.name().to_string()) - .collect(); - } - } - SimpleQueryMessage::Row(row) => { - total_rows += 1; - - if columns.is_empty() { - columns = row - .columns() - .iter() - .map(|column| column.name().to_string()) - .collect(); - } - - if rows.len() >= max_query_rows { - continue; - } - - let mut mapped_row = BTreeMap::new(); - for (index, column_name) in columns.iter().enumerate() { - mapped_row.insert(column_name.clone(), row.get(index).map(str::to_owned)); - } - rows.push(mapped_row); - } - SimpleQueryMessage::CommandComplete(count) => { - command_tag = Some(count); - } - _ => {} - } - } - - Ok(QueryResult { - columns, - row_count: rows.len(), - rows, - execution_ms: started_at.elapsed().as_millis(), - truncated: total_rows > max_query_rows, - command_tag, - }) - }) - .await - } - DatabaseEngine::Mysql | DatabaseEngine::Sqlite => { - run_query_mysql_or_sqlite(&app, &state, &connection_id, &sql, max_query_rows, engine).await - } - } -} - -async fn fetch_query_editor_metadata_for_connection( - app: &AppHandle, - state: &AppState, - connection_id: &str, - engine: DatabaseEngine, -) -> Result { - if engine == DatabaseEngine::Mysql { - let pool = get_or_create_mysql_pool(app, state, connection_id).await?; - let database = load_connection(app, connection_id)? - .map(|connection| connection.database) - .unwrap_or_default(); - let table_rows = sqlx::query( - " - select table_schema, table_name - from information_schema.tables - where table_type = 'BASE TABLE' - and table_schema = ? - and table_schema not in ('information_schema', 'mysql', 'performance_schema', 'sys') - order by table_schema, table_name - limit ? - ", - ) - .bind(&database) - .bind(MAX_EDITOR_TABLES + 1) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - - let truncated_tables = table_rows.len() as i64 > MAX_EDITOR_TABLES; - let mut tables = Vec::new(); - let mut truncated_columns = false; - - for row in table_rows.into_iter().take(MAX_EDITOR_TABLES as usize) { - let schema: String = mysql_get_string(&row, 0, "table_schema", "get_query_editor_metadata")?; - let name: String = mysql_get_string(&row, 1, "table_name", "get_query_editor_metadata")?; - let column_rows = sqlx::query( - " - select column_name, data_type - from information_schema.columns - where table_schema = ? and table_name = ? - order by ordinal_position - limit ? - ", - ) - .bind(&schema) - .bind(&name) - .bind(MAX_EDITOR_COLUMNS_PER_TABLE + 1) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - if column_rows.len() as i64 > MAX_EDITOR_COLUMNS_PER_TABLE { - truncated_columns = true; - } - let mut columns = Vec::new(); - for column in column_rows.into_iter().take(MAX_EDITOR_COLUMNS_PER_TABLE as usize) { - columns.push(QueryEditorColumn { - name: mysql_get_string(&column, 0, "column_name", "get_query_editor_metadata")?, - data_type: mysql_get_string(&column, 1, "data_type", "get_query_editor_metadata")?, - }); - } - tables.push(QueryEditorTable { schema, name, columns }); - } - - return Ok(QueryEditorMetadata { - tables, - functions: Vec::new(), - truncated_tables, - truncated_columns, - truncated_functions: false, - }); - } - - if engine == DatabaseEngine::Sqlite { - let pool = get_or_create_sqlite_pool(app, state, connection_id).await?; - let table_rows = sqlx::query( - " - select name - from sqlite_master - where type = 'table' - and name not like 'sqlite_%' - order by name - limit ? - ", - ) - .bind(MAX_EDITOR_TABLES + 1) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - let truncated_tables = table_rows.len() as i64 > MAX_EDITOR_TABLES; - let mut tables = Vec::new(); - let mut truncated_columns = false; - for row in table_rows.into_iter().take(MAX_EDITOR_TABLES as usize) { - let name: String = sqlite_get_idx(&row, 0, "name", "get_query_editor_metadata")?; - require_safe_identifier(&name, "table name")?; - let pragma_sql = format!("PRAGMA table_info(\"{}\");", quote_identifier(&name)); - let column_rows = sqlx::query(&pragma_sql) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - if column_rows.len() as i64 > MAX_EDITOR_COLUMNS_PER_TABLE { - truncated_columns = true; - } - let mut columns = Vec::new(); - for column in column_rows.into_iter().take(MAX_EDITOR_COLUMNS_PER_TABLE as usize) { - columns.push(QueryEditorColumn { - name: sqlite_get_name(&column, "name", "get_query_editor_metadata")?, - data_type: sqlite_get_name(&column, "type", "get_query_editor_metadata")?, - }); - } - tables.push(QueryEditorTable { - schema: "main".to_string(), - name, - columns, - }); - } - return Ok(QueryEditorMetadata { - tables, - functions: Vec::new(), - truncated_tables, - truncated_columns, - truncated_functions: false, - }); - } - - with_pool_client_retry(app, state, connection_id, (), |client, ()| async move { - let table_rows = client - .query( - " - select n.nspname::text as schema_name, c.relname::text as table_name - from pg_class c - join pg_namespace n on n.oid = c.relnamespace - where c.relkind in ('r', 'p', 'v', 'm', 'f') - and n.nspname not in ('pg_catalog', 'information_schema') - order by n.nspname, c.relname - limit $1 - ", - &[&(MAX_EDITOR_TABLES + 1)], - ) - .await - .map_err(|error| map_pg_err(error, None))?; - - let truncated_tables = table_rows.len() as i64 > MAX_EDITOR_TABLES; - let table_rows = if truncated_tables { - table_rows - .into_iter() - .take(MAX_EDITOR_TABLES as usize) - .collect::>() - } else { - table_rows - }; - - let mut tables = Vec::with_capacity(table_rows.len()); - let mut truncated_columns = false; - - for row in table_rows { - let schema: String = row.get(0); - let name: String = row.get(1); - let column_rows = client - .query( - " - select a.attname::text as column_name, - format_type(a.atttypid, a.atttypmod)::text as data_type - from pg_attribute a - join pg_class c on c.oid = a.attrelid - join pg_namespace n on n.oid = c.relnamespace - where n.nspname = $1 - and c.relname = $2 - and a.attnum > 0 - and not a.attisdropped - order by a.attnum - limit $3 - ", - &[&schema, &name, &(MAX_EDITOR_COLUMNS_PER_TABLE + 1)], - ) - .await - .map_err(|error| map_pg_err(error, None))?; - - let columns_exceeded = column_rows.len() as i64 > MAX_EDITOR_COLUMNS_PER_TABLE; - if columns_exceeded { - truncated_columns = true; - } - let columns = column_rows - .into_iter() - .take(MAX_EDITOR_COLUMNS_PER_TABLE as usize) - .map(|column| QueryEditorColumn { - name: column.get(0), - data_type: column.get(1), - }) - .collect(); - - tables.push(QueryEditorTable { - schema, - name, - columns, - }); - } - - let function_rows = client - .query( - " - select - n.nspname::text as schema_name, - p.proname::text as function_name, - coalesce(pg_get_function_identity_arguments(p.oid), '')::text as args, - pg_get_function_result(p.oid)::text as return_type - from pg_proc p - join pg_namespace n on n.oid = p.pronamespace - where n.nspname not in ('pg_catalog', 'information_schema') - order by n.nspname, p.proname - limit $1 - ", - &[&(MAX_EDITOR_FUNCTIONS + 1)], - ) - .await - .map_err(|error| map_pg_err(error, None))?; - - let truncated_functions = function_rows.len() as i64 > MAX_EDITOR_FUNCTIONS; - let functions = function_rows - .into_iter() - .take(MAX_EDITOR_FUNCTIONS as usize) - .map(|row| { - let args_raw: String = row.get(2); - QueryEditorFunction { - schema: row.get(0), - name: row.get(1), - arg_types: if args_raw.trim().is_empty() { - Vec::new() - } else { - args_raw - .split(',') - .map(|value| value.trim().to_string()) - .collect() - }, - return_type: row.get(3), - } - }) - .collect(); - - Ok(QueryEditorMetadata { - tables, - functions, - truncated_tables, - truncated_columns, - truncated_functions, - }) - }) - .await -} - -async fn fetch_foreign_keys_for_connection( - app: &AppHandle, - state: &AppState, - connection_id: &str, - engine: DatabaseEngine, -) -> Result, String> { - if engine == DatabaseEngine::Mysql { - let pool = get_or_create_mysql_pool(app, state, connection_id).await?; - let rows = sqlx::query( - " - select - kcu.table_schema as from_schema, - kcu.table_name as from_table, - kcu.column_name as from_column, - kcu.referenced_table_schema as to_schema, - kcu.referenced_table_name as to_table, - kcu.referenced_column_name as to_column - from information_schema.key_column_usage kcu - where kcu.referenced_table_name is not null - order by kcu.table_schema, kcu.table_name, kcu.ordinal_position - limit ? - ", - ) - .bind(MAX_FOREIGN_KEY_ROWS) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - let mut edges = Vec::new(); - for row in rows { - edges.push(ForeignKeyEdge { - from_schema: mysql_get_string(&row, 0, "from_schema", "get_foreign_keys")?, - from_table: mysql_get_string(&row, 1, "from_table", "get_foreign_keys")?, - from_column: mysql_get_string(&row, 2, "from_column", "get_foreign_keys")?, - to_schema: mysql_get_string(&row, 3, "to_schema", "get_foreign_keys")?, - to_table: mysql_get_string(&row, 4, "to_table", "get_foreign_keys")?, - to_column: mysql_get_string(&row, 5, "to_column", "get_foreign_keys")?, - }); - } - return Ok(edges); - } - - if engine == DatabaseEngine::Sqlite { - let pool = get_or_create_sqlite_pool(app, state, connection_id).await?; - let tables = sqlx::query( - " - select name - from sqlite_master - where type = 'table' - and name not like 'sqlite_%' - ", - ) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - let mut edges = Vec::new(); - for table in tables { - let table_name: String = sqlite_get_idx(&table, 0, "name", "get_foreign_keys")?; - require_safe_identifier(&table_name, "table name")?; - let fk_sql = format!("PRAGMA foreign_key_list(\"{}\");", quote_identifier(&table_name)); - let fk_rows = sqlx::query(&fk_sql) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - for row in fk_rows { - edges.push(ForeignKeyEdge { - from_schema: "main".to_string(), - from_table: table_name.clone(), - from_column: sqlite_get_name(&row, "from", "get_foreign_keys")?, - to_schema: "main".to_string(), - to_table: sqlite_get_name(&row, "table", "get_foreign_keys")?, - to_column: sqlite_get_name(&row, "to", "get_foreign_keys")?, - }); - if edges.len() >= MAX_FOREIGN_KEY_ROWS as usize { - return Ok(edges); - } - } - } - return Ok(edges); - } - - with_pool_client_retry(app, state, connection_id, (), |client, ()| async move { - let rows = client - .query( - " - select - src_ns.nspname::text as from_schema, - src_cls.relname::text as from_table, - src_att.attname::text as from_column, - tgt_ns.nspname::text as to_schema, - tgt_cls.relname::text as to_table, - tgt_att.attname::text as to_column - from pg_constraint c - join pg_class src_cls on src_cls.oid = c.conrelid - join pg_namespace src_ns on src_ns.oid = src_cls.relnamespace - join pg_class tgt_cls on tgt_cls.oid = c.confrelid - join pg_namespace tgt_ns on tgt_ns.oid = tgt_cls.relnamespace - cross join lateral unnest(c.conkey, c.confkey) as u(attnum, confattnum) - join pg_attribute src_att - on src_att.attrelid = c.conrelid - and src_att.attnum = u.attnum - and not src_att.attisdropped - join pg_attribute tgt_att - on tgt_att.attrelid = c.confrelid - and tgt_att.attnum = u.confattnum - and not tgt_att.attisdropped - where c.contype = 'f' - and src_ns.nspname not in ('pg_catalog', 'information_schema') - order by src_ns.nspname, src_cls.relname, c.conname, u.attnum - limit $1 - ", - &[&MAX_FOREIGN_KEY_ROWS], - ) - .await - .map_err(|error| error.to_string())?; - - Ok(rows - .into_iter() - .map(|row| ForeignKeyEdge { - from_schema: row.get(0), - from_table: row.get(1), - from_column: row.get(2), - to_schema: row.get(3), - to_table: row.get(4), - to_column: row.get(5), - }) - .collect()) - }) - .await -} - -fn ask_veloxy_context_cache_key(connection_id: &str, database_name: &str) -> String { - format!("{}::{}", connection_id, database_name) -} - -fn ask_veloxy_conversation_key(connection_id: &str, database_name: &str) -> String { - format!("{}::{}", connection_id, database_name) -} - -fn now_epoch_seconds() -> u64 { - use std::time::{SystemTime, UNIX_EPOCH}; - SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap_or_default() - .as_secs() -} - -async fn get_or_build_ask_veloxy_db_context( - app: &AppHandle, - state: &AppState, - connection_id: &str, - engine: DatabaseEngine, -) -> Result { - let stored_connection = load_connection(app, connection_id)? - .ok_or_else(|| "Stored connection details were not found.".to_string())?; - let cache_key = ask_veloxy_context_cache_key(connection_id, &stored_connection.database); - if let Some(cached) = state - .ask_veloxy_db_context_cache - .read() - .await - .get(&cache_key) - .cloned() - { - return Ok(cached); - } - - let metadata = fetch_query_editor_metadata_for_connection(app, state, connection_id, engine).await?; - let foreign_keys = fetch_foreign_keys_for_connection(app, state, connection_id, engine).await?; - let cache = AskVeloxyDbContextCache { - database_name: stored_connection.database, - engine, - metadata, - foreign_keys, - }; - state - .ask_veloxy_db_context_cache - .write() - .await - .insert(cache_key, cache.clone()); - Ok(cache) -} - -fn extract_sql_draft_from_text(message: &str) -> Option { - let lowered = message.to_lowercase(); - let markers = ["select ", "with ", "insert ", "update ", "delete ", "explain "]; - let start = markers - .iter() - .filter_map(|marker| lowered.find(marker)) - .min()?; - let mut sql = message[start..].trim().to_string(); - if let Some(idx) = sql.find("```") { - sql.truncate(idx); - } - if sql.ends_with('.') { - sql.pop(); - } - if sql.is_empty() { - None - } else { - Some(sql) - } -} - -fn parse_bool_field(value: &Value, field: &str, default: bool) -> bool { - value.get(field).and_then(Value::as_bool).unwrap_or(default) -} - -#[tauri::command] -pub async fn get_query_editor_metadata( - app: AppHandle, - state: State<'_, AppState>, - connection_id: Option, -) -> Result { - let (connection_id, engine) = resolve_connection_engine(&app, &state, connection_id).await?; - - if engine == DatabaseEngine::Mysql { - let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; - let database = load_connection(&app, &connection_id)? - .map(|connection| connection.database) - .unwrap_or_default(); - let table_rows = sqlx::query( - " - select table_schema, table_name - from information_schema.tables - where table_type = 'BASE TABLE' - and table_schema = ? - and table_schema not in ('information_schema', 'mysql', 'performance_schema', 'sys') - order by table_schema, table_name - limit ? - ", - ) - .bind(&database) - .bind(MAX_EDITOR_TABLES + 1) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - - let truncated_tables = table_rows.len() as i64 > MAX_EDITOR_TABLES; - let mut tables = Vec::new(); - let mut truncated_columns = false; - - for row in table_rows.into_iter().take(MAX_EDITOR_TABLES as usize) { - let schema: String = mysql_get_string(&row, 0, "table_schema", "get_query_editor_metadata")?; - let name: String = mysql_get_string(&row, 1, "table_name", "get_query_editor_metadata")?; - let column_rows = sqlx::query( - " - select column_name, data_type - from information_schema.columns - where table_schema = ? and table_name = ? - order by ordinal_position - limit ? - ", - ) - .bind(&schema) - .bind(&name) - .bind(MAX_EDITOR_COLUMNS_PER_TABLE + 1) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - if column_rows.len() as i64 > MAX_EDITOR_COLUMNS_PER_TABLE { - truncated_columns = true; - } - let mut columns = Vec::new(); - for column in column_rows.into_iter().take(MAX_EDITOR_COLUMNS_PER_TABLE as usize) { - columns.push(QueryEditorColumn { - name: mysql_get_string(&column, 0, "column_name", "get_query_editor_metadata")?, - data_type: mysql_get_string(&column, 1, "data_type", "get_query_editor_metadata")?, - }); - } - tables.push(QueryEditorTable { schema, name, columns }); - } - - return Ok(QueryEditorMetadata { - tables, - functions: Vec::new(), - truncated_tables, - truncated_columns, - truncated_functions: false, - }); - } - - if engine == DatabaseEngine::Sqlite { - let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; - let table_rows = sqlx::query( - " - select name - from sqlite_master - where type = 'table' - and name not like 'sqlite_%' - order by name - limit ? - ", - ) - .bind(MAX_EDITOR_TABLES + 1) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - let truncated_tables = table_rows.len() as i64 > MAX_EDITOR_TABLES; - let mut tables = Vec::new(); - let mut truncated_columns = false; - for row in table_rows.into_iter().take(MAX_EDITOR_TABLES as usize) { - let name: String = sqlite_get_idx(&row, 0, "name", "get_query_editor_metadata")?; - require_safe_identifier(&name, "table name")?; - let pragma_sql = format!("PRAGMA table_info(\"{}\");", quote_identifier(&name)); - let column_rows = sqlx::query(&pragma_sql) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - if column_rows.len() as i64 > MAX_EDITOR_COLUMNS_PER_TABLE { - truncated_columns = true; - } - let mut columns = Vec::new(); - for column in column_rows.into_iter().take(MAX_EDITOR_COLUMNS_PER_TABLE as usize) { - columns.push(QueryEditorColumn { - name: sqlite_get_name(&column, "name", "get_query_editor_metadata")?, - data_type: sqlite_get_name(&column, "type", "get_query_editor_metadata")?, - }); - } - tables.push(QueryEditorTable { - schema: "main".to_string(), - name, - columns, - }); - } - return Ok(QueryEditorMetadata { - tables, - functions: Vec::new(), - truncated_tables, - truncated_columns, - truncated_functions: false, - }); - } - - with_pool_client_retry(&app, &state, &connection_id, (), |client, ()| async move { - let table_rows = client - .query( - " - select n.nspname::text as schema_name, c.relname::text as table_name - from pg_class c - join pg_namespace n on n.oid = c.relnamespace - where c.relkind in ('r', 'p', 'v', 'm', 'f') - and n.nspname not in ('pg_catalog', 'information_schema') - order by n.nspname, c.relname - limit $1 - ", - &[&(MAX_EDITOR_TABLES + 1)], - ) - .await - .map_err(|error| map_pg_err(error, None))?; - - let truncated_tables = table_rows.len() as i64 > MAX_EDITOR_TABLES; - let table_rows = if truncated_tables { - table_rows - .into_iter() - .take(MAX_EDITOR_TABLES as usize) - .collect::>() - } else { - table_rows - }; - - let mut tables = Vec::with_capacity(table_rows.len()); - let mut truncated_columns = false; - - for row in table_rows { - let schema: String = row.get(0); - let name: String = row.get(1); - let column_rows = client - .query( - " - select a.attname::text as column_name, - format_type(a.atttypid, a.atttypmod)::text as data_type - from pg_attribute a - join pg_class c on c.oid = a.attrelid - join pg_namespace n on n.oid = c.relnamespace - where n.nspname = $1 - and c.relname = $2 - and a.attnum > 0 - and not a.attisdropped - order by a.attnum - limit $3 - ", - &[&schema, &name, &(MAX_EDITOR_COLUMNS_PER_TABLE + 1)], - ) - .await - .map_err(|error| map_pg_err(error, None))?; - - let columns_exceeded = column_rows.len() as i64 > MAX_EDITOR_COLUMNS_PER_TABLE; - if columns_exceeded { - truncated_columns = true; - } - let columns = column_rows - .into_iter() - .take(MAX_EDITOR_COLUMNS_PER_TABLE as usize) - .map(|column| QueryEditorColumn { - name: column.get(0), - data_type: column.get(1), - }) - .collect(); - - tables.push(QueryEditorTable { - schema, - name, - columns, - }); - } - - let function_rows = client - .query( - " - select - n.nspname::text as schema_name, - p.proname::text as function_name, - coalesce(pg_get_function_identity_arguments(p.oid), '')::text as args, - pg_get_function_result(p.oid)::text as return_type - from pg_proc p - join pg_namespace n on n.oid = p.pronamespace - where n.nspname not in ('pg_catalog', 'information_schema') - order by n.nspname, p.proname - limit $1 - ", - &[&(MAX_EDITOR_FUNCTIONS + 1)], - ) - .await - .map_err(|error| map_pg_err(error, None))?; - - let truncated_functions = function_rows.len() as i64 > MAX_EDITOR_FUNCTIONS; - let functions = function_rows - .into_iter() - .take(MAX_EDITOR_FUNCTIONS as usize) - .map(|row| { - let args_raw: String = row.get(2); - QueryEditorFunction { - schema: row.get(0), - name: row.get(1), - arg_types: if args_raw.trim().is_empty() { - Vec::new() - } else { - args_raw - .split(',') - .map(|value| value.trim().to_string()) - .collect() - }, - return_type: row.get(3), - } - }) - .collect(); - - Ok(QueryEditorMetadata { - tables, - functions, - truncated_tables, - truncated_columns, - truncated_functions, - }) - }) - .await -} - -fn estimate_tokens(chars: usize) -> usize { - // Lightweight estimate good enough for budget telemetry. - (chars / 4).max(1) -} - -fn normalize_openrouter_base(base: Option<&str>) -> String { - let trimmed = base.unwrap_or("https://openrouter.ai/api/v1").trim(); - let value = if trimmed.is_empty() { - "https://openrouter.ai/api/v1" - } else { - trimmed - }; - value.trim_end_matches('/').to_string() -} - -fn truncate_on_char_boundary(value: &mut String, max_bytes: usize) { - if value.len() <= max_bytes { - return; - } - let mut truncate_at = max_bytes; - while !value.is_char_boundary(truncate_at) && truncate_at > 0 { - truncate_at -= 1; - } - value.truncate(truncate_at); -} - -fn table_matches_target(table: &QueryEditorTable, target: Option<&AskVeloxyTableRef>) -> bool { - let Some(target) = target else { - return false; - }; - table.schema.eq_ignore_ascii_case(&target.schema) && table.name.eq_ignore_ascii_case(&target.name) -} - -fn table_relevance_score(table: &QueryEditorTable, prompt_lower: &str) -> usize { - let mut score = 0usize; - let full_name = format!("{}.{}", table.schema.to_lowercase(), table.name.to_lowercase()); - if prompt_lower.contains(&table.name.to_lowercase()) { - score += 3; - } - if prompt_lower.contains(&table.schema.to_lowercase()) { - score += 2; - } - if prompt_lower.contains(&full_name) { - score += 4; - } - score -} - -fn relationship_relevance_score(edge: &ForeignKeyEdge, prompt_lower: &str) -> usize { - let from_name = format!("{}.{}", edge.from_schema.to_lowercase(), edge.from_table.to_lowercase()); - let to_name = format!("{}.{}", edge.to_schema.to_lowercase(), edge.to_table.to_lowercase()); - let mut score = 0usize; - if prompt_lower.contains(&edge.from_table.to_lowercase()) || prompt_lower.contains(&from_name) { - score += 2; - } - if prompt_lower.contains(&edge.to_table.to_lowercase()) || prompt_lower.contains(&to_name) { - score += 2; - } - score -} - -fn build_schema_context( - db_context: &AskVeloxyDbContextCache, - prompt: &str, - target_table: Option<&AskVeloxyTableRef>, -) -> String { - let prompt_lower = prompt.to_lowercase(); - let mut ranked: Vec<(&QueryEditorTable, usize, bool)> = db_context - .metadata - .tables - .iter() - .map(|table| { - ( - table, - table_relevance_score(table, &prompt_lower), - table_matches_target(table, target_table), - ) - }) - .collect(); - - ranked.sort_by(|a, b| b.2.cmp(&a.2).then_with(|| b.1.cmp(&a.1))); - - let mut schema_context = String::new(); - schema_context.push_str(&format!( - "database {} engine {:?}\n", - db_context.database_name, db_context.engine - )); - for (table, _score, _is_target) in ranked.into_iter().take(ASK_VELOXY_MAX_CONTEXT_TABLES) { - let columns = table - .columns - .iter() - .take(ASK_VELOXY_MAX_CONTEXT_COLUMNS) - .map(|column| format!("{}:{}", column.name, column.data_type)) - .collect::>() - .join(", "); - schema_context.push_str(&format!( - "table {}.{} columns [{}]\n", - table.schema, table.name, columns - )); - if schema_context.len() >= ASK_VELOXY_SCHEMA_CHAR_BUDGET { - truncate_on_char_boundary(&mut schema_context, ASK_VELOXY_SCHEMA_CHAR_BUDGET); - break; - } - } - - let mut ranked_relationships = db_context - .foreign_keys - .iter() - .map(|edge| (edge, relationship_relevance_score(edge, &prompt_lower))) - .collect::>(); - ranked_relationships.sort_by(|a, b| b.1.cmp(&a.1)); - for (edge, _score) in ranked_relationships - .into_iter() - .take(ASK_VELOXY_MAX_CONTEXT_RELATIONSHIPS) - { - schema_context.push_str(&format!( - "relationship {}.{}({}) -> {}.{}({})\n", - edge.from_schema, - edge.from_table, - edge.from_column, - edge.to_schema, - edge.to_table, - edge.to_column - )); - if schema_context.len() >= ASK_VELOXY_SCHEMA_CHAR_BUDGET { - truncate_on_char_boundary(&mut schema_context, ASK_VELOXY_SCHEMA_CHAR_BUDGET); - break; - } - } - schema_context -} - -fn classify_sql_intent(sql: &str) -> String { - let normalized = sql.trim_start().to_ascii_lowercase(); - if normalized.starts_with("select") || normalized.starts_with("with") { - return "select".to_string(); - } - if normalized.starts_with("insert") { - return "insert".to_string(); - } - if normalized.starts_with("update") { - return "update".to_string(); - } - if normalized.starts_with("delete") { - return "delete".to_string(); - } - if normalized.starts_with("explain") { - return "explain".to_string(); - } - "unknown".to_string() -} - -/// Whether every statement in `sql` is read-only. Transaction-control keywords -/// (`begin`/`commit`/`rollback`/`start`/`savepoint`) are ignored so a wrapped -/// `BEGIN; SELECT ...; COMMIT;` still counts as read-only, while any write -/// statement makes the whole batch non-read-only. -fn is_read_only_sql(sql: &str) -> bool { - let mut saw_statement = false; - for statement in sql.split(';').map(str::trim).filter(|s| !s.is_empty()) { - let normalized = statement.to_ascii_lowercase(); - let is_transaction_control = ["begin", "commit", "rollback", "start", "savepoint", "release"] - .iter() - .any(|kw| normalized.starts_with(kw)); - if is_transaction_control { - continue; - } - saw_statement = true; - match classify_sql_intent(statement).as_str() { - "select" | "explain" => {} - _ => return false, - } - } - saw_statement -} - -fn has_multiple_statements(sql: &str) -> bool { - let statements = sql - .split(';') - .map(str::trim) - .filter(|segment| !segment.is_empty()) - .count(); - statements > 1 -} - -fn validate_generated_sql(sql: &str) -> Result<(), String> { - let trimmed = sql.trim(); - if trimmed.is_empty() { - return Err("Ask Veloxy returned an empty SQL statement.".to_string()); - } - if has_multiple_statements(trimmed) { - return Err("Ask Veloxy generated multiple SQL statements. Please ask for a single statement.".to_string()); - } - Ok(()) -} - -fn extract_openrouter_message_content(payload: &Value) -> Result { - let content_value = payload - .get("choices") - .and_then(|choices| choices.get(0)) - .and_then(|choice| choice.get("message")) - .and_then(|message| message.get("content")) - .ok_or_else(|| "OpenRouter response missing choices[0].message.content".to_string())?; - - if let Some(content) = content_value.as_str() { - return Ok(content.to_string()); - } - - if let Some(items) = content_value.as_array() { - let mut merged = String::new(); - for item in items { - if let Some(text) = item.get("text").and_then(Value::as_str) { - merged.push_str(text); - } - } - if !merged.trim().is_empty() { - return Ok(merged); - } - } - - Err("OpenRouter returned an unsupported message format.".to_string()) -} - -fn parse_ask_veloxy_json(content: &str) -> Result { - if let Ok(value) = serde_json::from_str::(content) { - return Ok(value); - } - let start = content.find('{'); - let end = content.rfind('}'); - match (start, end) { - (Some(start_idx), Some(end_idx)) if end_idx > start_idx => { - serde_json::from_str::(&content[start_idx..=end_idx]) - .map_err(|error| format!("Ask Veloxy response was not valid JSON: {}", error)) - } - _ => Err("Ask Veloxy response did not contain JSON.".to_string()), - } -} - -fn parse_ask_veloxy_suggestions(generated: &Value) -> Vec { - generated - .get("suggestions") - .and_then(Value::as_array) - .map(|items| { - items - .iter() - .filter_map(Value::as_str) - .map(str::trim) - .filter(|item| !item.is_empty()) - .take(5) - .map(|item| { - let mut value = item.to_string(); - truncate_on_char_boundary(&mut value, 200); - value - }) - .collect::>() - }) - .unwrap_or_default() -} - -fn parse_ask_veloxy_chat_json(content: &str) -> Result { - if let Ok(value) = serde_json::from_str::(content) { - return Ok(value); - } - let start = content.find('{'); - let end = content.rfind('}'); - match (start, end) { - (Some(start_idx), Some(end_idx)) if end_idx > start_idx => { - serde_json::from_str::(&content[start_idx..=end_idx]) - .map_err(|error| format!("Ask Veloxy chat JSON was invalid: {}", error)) - } - _ => Err("Ask Veloxy chat response did not contain JSON.".to_string()), - } -} - -fn decode_json_quoted_string(value: &str) -> Option { - serde_json::from_str::(&format!("\"{}\"", value)).ok() -} - -fn unescape_json_string_fragment(raw: &str) -> String { - let mut out = String::with_capacity(raw.len()); - let mut chars = raw.chars().peekable(); - while let Some(ch) = chars.next() { - if ch == '\\' { - match chars.next() { - Some('n') => out.push('\n'), - Some('t') => out.push('\t'), - Some('r') => out.push('\r'), - Some('"') => out.push('"'), - Some('\\') => out.push('\\'), - Some(other) => { - out.push('\\'); - out.push(other); - } - None => out.push('\\'), - } - } else { - out.push(ch); - } - } - out -} - -fn extract_json_string_field(content: &str, key: &str, allow_partial: bool) -> Option { - let marker = format!("\"{}\"", key); - let marker_idx = content.find(&marker)?; - let mut idx = marker_idx + marker.len(); - let bytes = content.as_bytes(); - - while idx < bytes.len() && bytes[idx].is_ascii_whitespace() { - idx += 1; - } - if idx >= bytes.len() || bytes[idx] != b':' { - return None; - } - idx += 1; - while idx < bytes.len() && bytes[idx].is_ascii_whitespace() { - idx += 1; - } - if idx >= bytes.len() || bytes[idx] != b'"' { - return None; - } - idx += 1; - let start = idx; - let mut escaped = false; - while idx < bytes.len() { - let byte = bytes[idx]; - if escaped { - escaped = false; - idx += 1; - continue; - } - if byte == b'\\' { - escaped = true; - idx += 1; - continue; - } - if byte == b'"' { - let raw = &content[start..idx]; - return decode_json_quoted_string(raw) - .or_else(|| Some(unescape_json_string_fragment(raw))) - .map(|text| text.trim().to_string()) - .filter(|text| !text.is_empty()); - } - idx += 1; - } - - if allow_partial && start < bytes.len() { - let raw = &content[start..]; - let text = unescape_json_string_fragment(raw).trim().to_string(); - if !text.is_empty() { - return Some(text); - } - } - None -} - -fn extract_message_from_loose_json(content: &str) -> Option { - let trimmed = content.trim(); - if trimmed.is_empty() { - return None; - } - let unwrapped = trimmed - .strip_prefix("```json") - .or_else(|| trimmed.strip_prefix("```JSON")) - .map(str::trim_start) - .unwrap_or(trimmed); - let unwrapped = unwrapped.strip_suffix("```").unwrap_or(unwrapped).trim(); - - ["message", "reply", "content"] - .iter() - .find_map(|key| extract_json_string_field(unwrapped, key, false)) - .or_else(|| { - ["message", "reply", "content"] - .iter() - .find_map(|key| extract_json_string_field(unwrapped, key, true)) - }) -} - -fn looks_like_json_response(content: &str) -> bool { - let trimmed = content.trim_start(); - trimmed.starts_with('{') || trimmed.starts_with("```") -} - -fn streaming_display_text(accumulated: &str) -> String { - let trimmed = accumulated.trim(); - if trimmed.is_empty() { - return String::new(); - } - if let Some(text) = extract_message_from_loose_json(trimmed) { - return text; - } - if !looks_like_json_response(trimmed) { - return trimmed.to_string(); - } - String::new() -} - -fn parse_chat_message(value: &Value) -> Option { - if let Some(text) = value - .as_str() - .map(str::trim) - .filter(|text| !text.is_empty()) - .map(str::to_string) - { - return Some(text); - } - value - .get("message") - .and_then(Value::as_str) - .or_else(|| value.get("reply").and_then(Value::as_str)) - .map(str::trim) - .filter(|text| !text.is_empty()) - .map(str::to_string) -} - -type ParsedAskVeloxyChat = ( - String, - Vec, - Vec, - Option, - bool, - bool, -); - -fn parse_ask_veloxy_chat_content(message_content: &str) -> ParsedAskVeloxyChat { - match parse_ask_veloxy_chat_json(message_content) { - Ok(value) => { - let message = - parse_chat_message(&value).unwrap_or_else(|| message_content.trim().to_string()); - let mut draft = value - .get("sqlDraft") - .and_then(Value::as_str) - .or_else(|| value.get("sql_draft").and_then(Value::as_str)) - .map(str::trim) - .filter(|text| !text.is_empty()) - .map(str::to_string); - if draft.is_none() { - draft = extract_sql_draft_from_text(&message); - } - let suggestions = value - .get("suggestions") - .and_then(Value::as_array) - .map(|items| { - items - .iter() - .filter_map(Value::as_str) - .map(str::trim) - .filter(|text| !text.is_empty()) - .take(5) - .map(str::to_string) - .collect::>() - }) - .unwrap_or_default(); - let warnings = value - .get("warnings") - .and_then(Value::as_array) - .map(|items| { - items - .iter() - .filter_map(Value::as_str) - .map(str::to_string) - .collect::>() - }) - .unwrap_or_default(); - let needs_sql_generation = parse_bool_field(&value, "needsSqlGeneration", draft.is_some()); - let needs_clarification = parse_bool_field(&value, "needsClarification", false); - ( - message, - suggestions, - warnings, - draft, - needs_sql_generation, - needs_clarification, - ) - } - Err(_) => { - let normalized_message = extract_message_from_loose_json(message_content) - .unwrap_or_else(|| { - if looks_like_json_response(message_content) { - String::new() - } else { - message_content.trim().to_string() - } - }); - let mut warnings = vec!["Model returned non-JSON chat output. Parsed in tolerant mode.".to_string()]; - if normalized_message.is_empty() && looks_like_json_response(message_content) { - warnings.push("Response JSON could not be parsed. Try asking again.".to_string()); - } - let draft = extract_sql_draft_from_text(&normalized_message); - let needs_sql_generation = draft.is_some(); - ( - normalized_message, - Vec::new(), - warnings, - draft, - needs_sql_generation, - false, - ) - } - } -} - -fn extract_openrouter_stream_delta(data: &str) -> Option { - let payload: Value = serde_json::from_str(data).ok()?; - payload - .get("choices") - .and_then(|choices| choices.get(0)) - .and_then(|choice| choice.get("delta")) - .and_then(|delta| delta.get("content")) - .and_then(Value::as_str) - .filter(|text| !text.is_empty()) - .map(str::to_string) -} - -fn emit_veloxy_stream_chunk(app: &AppHandle, chunk: VeloxyStreamChunk) { - let _ = app.emit("veloxy-stream-chunk", chunk); -} - -fn extract_openrouter_finish_reason(data: &str) -> Option { - let payload: Value = serde_json::from_str(data).ok()?; - payload - .get("choices") - .and_then(|choices| choices.get(0)) - .and_then(|choice| choice.get("finish_reason")) - .and_then(Value::as_str) - .map(str::to_string) -} - -async fn stream_openrouter_chat_completion( - app: &AppHandle, - client: &reqwest::Client, - endpoint: &str, - api_key: &str, - model: &str, - system_prompt: &str, - user_prompt: &str, - request_id: &str, - cancel: Arc, -) -> Result<(String, bool), String> { - let response = client - .post(endpoint) - .header("Authorization", format!("Bearer {}", api_key)) - .header("Content-Type", "application/json") - .json(&serde_json::json!({ - "model": model, - "temperature": 0.2, - "max_tokens": ASK_VELOXY_MAX_CHAT_TOKENS, - "stream": true, - "messages": [ - { "role": "system", "content": system_prompt }, - { "role": "user", "content": user_prompt } - ] - })) - .send() - .await - .map_err(|error| format!("OpenRouter request failed: {}", error))?; - - let status = response.status(); - if !status.is_success() { - let body = response - .text() - .await - .unwrap_or_else(|_| "Unknown OpenRouter error".to_string()); - if let Ok(payload) = serde_json::from_str::(&body) { - let message = payload - .get("error") - .and_then(|error| error.get("message")) - .and_then(Value::as_str) - .unwrap_or("Unknown OpenRouter error"); - return Err(format!("OpenRouter error ({}): {}", status.as_u16(), message)); - } - return Err(format!("OpenRouter error ({}): {}", status.as_u16(), body)); - } - - let mut stream = response.bytes_stream(); - let mut buffer = String::new(); - let mut accumulated = String::new(); - let mut last_display_len = 0usize; - let mut hit_token_limit = false; - - while let Some(chunk) = stream.next().await { - if cancel.load(Ordering::Relaxed) { - return Ok((accumulated, hit_token_limit)); - } - let bytes = chunk.map_err(|error| format!("OpenRouter stream read failed: {}", error))?; - buffer.push_str(&String::from_utf8_lossy(&bytes)); - - while let Some(line_end) = buffer.find('\n') { - let line = buffer[..line_end].trim_end_matches('\r').to_string(); - buffer.drain(..=line_end); - - if !line.starts_with("data: ") { - continue; - } - let data = line["data: ".len()..].trim(); - if data == "[DONE]" { - continue; - } - if extract_openrouter_finish_reason(data).as_deref() == Some("length") { - hit_token_limit = true; - } - if let Some(delta) = extract_openrouter_stream_delta(data) { - accumulated.push_str(&delta); - let display = streaming_display_text(&accumulated); - let display_delta = if display.len() > last_display_len { - display[last_display_len..].to_string() - } else { - String::new() - }; - last_display_len = display.len(); - if !display_delta.is_empty() { - emit_veloxy_stream_chunk( - app, - VeloxyStreamChunk { - request_id: request_id.to_string(), - delta: display_delta, - done: false, - message: None, - suggestions: Vec::new(), - warnings: Vec::new(), - sql_draft: None, - needs_sql_generation: false, - needs_clarification: false, - }, - ); - } - } - } - } - - Ok((accumulated, hit_token_limit)) -} - -#[tauri::command] -pub async fn cancel_veloxy_request(state: State<'_, AppState>) -> Result<(), String> { - if let Some(cancel) = state.veloxy_cancel.read().await.as_ref() { - cancel.store(true, Ordering::Relaxed); - } - Ok(()) -} - -#[tauri::command] -pub async fn chat_with_db( - app: AppHandle, - state: State<'_, AppState>, - input: AskVeloxyChatRequest, -) -> Result { - let natural_prompt = input.natural_prompt.trim(); - if natural_prompt.is_empty() { - return Err("Ask Veloxy prompt cannot be empty.".to_string()); - } - if input.provider_config.api_key.trim().is_empty() { - return Err("OpenRouter API key is required.".to_string()); - } - if input.provider_config.model.trim().is_empty() { - return Err("OpenRouter model is required.".to_string()); - } - - let (connection_id, engine) = - resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; - let stored_connection = load_connection(&app, &connection_id)? - .ok_or_else(|| "Stored connection details were not found.".to_string())?; - let db_context = get_or_build_ask_veloxy_db_context(&app, &state, &connection_id, engine).await?; - let schema_context = build_schema_context(&db_context, natural_prompt, input.target_table.as_ref()); - let conversation_key = ask_veloxy_conversation_key(&connection_id, &stored_connection.database); - let history = state - .ask_veloxy_conversations - .read() - .await - .get(&conversation_key) - .cloned() - .unwrap_or_default(); - - let history_block = history - .iter() - .rev() - .take(8) - .rev() - .map(|message| format!("{}: {}", message.role, message.text)) - .collect::>() - .join("\n"); - - let mut user_prompt = format!( - "Engine: {:?}\nDatabase: {}\nTask: {}\nMaxRows: {}\nRecentConversation:\n{}\nSchemaContext:\n{}\n", - db_context.engine, - db_context.database_name, - natural_prompt, - input.max_rows.unwrap_or(MAX_QUERY_ROWS), - history_block, - schema_context - ); - truncate_on_char_boundary(&mut user_prompt, ASK_VELOXY_PROMPT_CHAR_BUDGET); - - let system_prompt = "You are Ask Veloxy chat mode. Return JSON when possible with keys: message (string), suggestions (array of strings), sqlDraft (string optional), needsSqlGeneration (boolean), needsClarification (boolean), warnings (array of strings). If JSON is not possible, return helpful plain text."; - let base_url = normalize_openrouter_base(input.provider_config.base_url.as_deref()); - let endpoint = format!("{}/chat/completions", base_url); - let client = state.openrouter_client.get_or_init(reqwest::Client::new); - let request_id = input - .request_id - .clone() - .filter(|value| !value.trim().is_empty()) - .unwrap_or_else(|| format!("req-{}", uuid::Uuid::new_v4())); - - let cancel = Arc::new(AtomicBool::new(false)); - { - let mut guard = state.veloxy_cancel.write().await; - *guard = Some(cancel.clone()); - } - - let (message_content, hit_token_limit) = stream_openrouter_chat_completion( - &app, - client, - &endpoint, - input.provider_config.api_key.trim(), - input.provider_config.model.trim(), - system_prompt, - &user_prompt, - &request_id, - cancel.clone(), - ) - .await?; - - { - let mut guard = state.veloxy_cancel.write().await; - *guard = None; - } - - let (message, suggestions, mut warnings, sql_draft, needs_sql_generation, needs_clarification) = - parse_ask_veloxy_chat_content(&message_content); - - if cancel.load(Ordering::Relaxed) { - warnings.push("Stopped early.".to_string()); - } - if hit_token_limit { - warnings.push(format!( - "Response may be truncated (model output limit of {} tokens).", - ASK_VELOXY_MAX_CHAT_TOKENS - )); - } - - emit_veloxy_stream_chunk( - &app, - VeloxyStreamChunk { - request_id: request_id.clone(), - delta: String::new(), - done: true, - message: Some(message.clone()), - suggestions: suggestions.clone(), - warnings: warnings.clone(), - sql_draft: sql_draft.clone(), - needs_sql_generation, - needs_clarification, - }, - ); - - { - let mut conversations = state.ask_veloxy_conversations.write().await; - let bucket = conversations.entry(conversation_key).or_default(); - bucket.push(AskVeloxyConversationMessage { - id: format!("msg-{}", uuid::Uuid::new_v4()), - role: "user".to_string(), - mode: "chat".to_string(), - text: natural_prompt.to_string(), - created_at: now_epoch_seconds(), - sql_draft: None, - }); - bucket.push(AskVeloxyConversationMessage { - id: format!("msg-{}", uuid::Uuid::new_v4()), - role: "assistant".to_string(), - mode: "chat".to_string(), - text: message.clone(), - created_at: now_epoch_seconds(), - sql_draft: sql_draft.clone(), - }); - if bucket.len() > ASK_VELOXY_MAX_HISTORY_MESSAGES { - let remove_count = bucket.len() - ASK_VELOXY_MAX_HISTORY_MESSAGES; - bucket.drain(0..remove_count); - } - } - - Ok(AskVeloxyChatResponse { - message, - suggestions, - warnings, - sql_draft, - needs_sql_generation, - needs_clarification, - }) -} - -#[tauri::command] -pub async fn load_veloxy_conversation( - app: AppHandle, - state: State<'_, AppState>, - connection_id: Option, -) -> Result { - let (resolved_connection_id, _) = resolve_connection_engine(&app, &state, connection_id).await?; - let stored_connection = load_connection(&app, &resolved_connection_id)? - .ok_or_else(|| "Stored connection details were not found.".to_string())?; - let key = ask_veloxy_conversation_key(&resolved_connection_id, &stored_connection.database); - let messages = state - .ask_veloxy_conversations - .read() - .await - .get(&key) - .cloned() - .unwrap_or_default(); - Ok(AskVeloxyConversationResponse { messages }) -} - -#[tauri::command] -pub async fn clear_veloxy_conversation( - app: AppHandle, - state: State<'_, AppState>, - connection_id: Option, -) -> Result<(), String> { - let (resolved_connection_id, _) = resolve_connection_engine(&app, &state, connection_id).await?; - let stored_connection = load_connection(&app, &resolved_connection_id)? - .ok_or_else(|| "Stored connection details were not found.".to_string())?; - let key = ask_veloxy_conversation_key(&resolved_connection_id, &stored_connection.database); - state.ask_veloxy_conversations.write().await.remove(&key); - Ok(()) -} - -#[tauri::command] -pub async fn generate_sql_from_nl( - app: AppHandle, - state: State<'_, AppState>, - input: AskVeloxyRequest, -) -> Result { - let natural_prompt = input.natural_prompt.trim(); - if natural_prompt.is_empty() { - return Err("Ask Veloxy prompt cannot be empty.".to_string()); - } - if input.provider_config.api_key.trim().is_empty() { - return Err("OpenRouter API key is required.".to_string()); - } - if input.provider_config.model.trim().is_empty() { - return Err("OpenRouter model is required.".to_string()); - } - - let (connection_id, engine) = - resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; - let db_context = get_or_build_ask_veloxy_db_context(&app, &state, &connection_id, engine).await?; - let schema_context = build_schema_context(&db_context, natural_prompt, input.target_table.as_ref()); - - let mut user_prompt = format!( - "Engine: {:?}\nDatabase: {}\nTask: {}\nMaxRows: {}\nSchemaContext:\n{}\n", - db_context.engine, - db_context.database_name, - natural_prompt, - input.max_rows.unwrap_or(MAX_QUERY_ROWS), - schema_context - ); - truncate_on_char_boundary(&mut user_prompt, ASK_VELOXY_PROMPT_CHAR_BUDGET); - - let system_prompt = "You are Ask Veloxy. Return JSON only with keys: sql (string), intent (string), confidence (number 0..1), explanation (string), suggestions (array of short strings), warnings (array of strings). Generate exactly one SQL statement, keep explanation concise, and never include markdown."; - let base_url = normalize_openrouter_base(input.provider_config.base_url.as_deref()); - let endpoint = format!("{}/chat/completions", base_url); - - let client = state.openrouter_client.get_or_init(reqwest::Client::new); - let response = client - .post(&endpoint) - .header("Authorization", format!("Bearer {}", input.provider_config.api_key.trim())) - .header("Content-Type", "application/json") - .json(&serde_json::json!({ - "model": input.provider_config.model.trim(), - "temperature": 0.1, - "max_tokens": 500, - "messages": [ - { "role": "system", "content": system_prompt }, - { "role": "user", "content": user_prompt } - ] - })) - .send() - .await - .map_err(|error| format!("OpenRouter request failed: {}", error))?; - - let status = response.status(); - let payload = response - .json::() - .await - .map_err(|error| format!("Invalid OpenRouter JSON response: {}", error))?; - if !status.is_success() { - let message = payload - .get("error") - .and_then(|error| error.get("message")) - .and_then(Value::as_str) - .unwrap_or("Unknown OpenRouter error"); - return Err(format!("OpenRouter error ({}): {}", status.as_u16(), message)); - } - - let message_content = extract_openrouter_message_content(&payload)?; - let generated = parse_ask_veloxy_json(&message_content)?; - let sql = generated - .get("sql") - .and_then(Value::as_str) - .unwrap_or_default() - .trim() - .to_string(); - validate_generated_sql(&sql)?; - - let mut warnings = generated - .get("warnings") - .and_then(Value::as_array) - .map(|items| { - items - .iter() - .filter_map(Value::as_str) - .map(str::to_string) - .collect::>() - }) - .unwrap_or_default(); - - let intent = generated - .get("intent") - .and_then(Value::as_str) - .map(str::to_string) - .unwrap_or_else(|| classify_sql_intent(&sql)); - let confidence = generated - .get("confidence") - .and_then(Value::as_f64) - .unwrap_or(0.6) - .clamp(0.0, 1.0); - let explanation = generated - .get("explanation") - .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(|value| { - let mut truncated = value.to_string(); - truncate_on_char_boundary(&mut truncated, 350); - truncated - }); - let suggestions = parse_ask_veloxy_suggestions(&generated); - - if intent != "select" { - warnings.push("Generated SQL is not read-only. Review before execution.".to_string()); - } - if confidence < 0.5 { - warnings.push("Low confidence result. Review SQL carefully.".to_string()); - } - - let token_stats = AskVeloxyTokenStats { - schema_chars: schema_context.len(), - schema_tokens_estimate: estimate_tokens(schema_context.len()), - prompt_chars: user_prompt.len() + system_prompt.len(), - prompt_tokens_estimate: estimate_tokens(user_prompt.len() + system_prompt.len()), - }; - - Ok(AskVeloxyResponse { - sql, - intent, - confidence, - explanation, - suggestions, - warnings, - token_stats, - }) -} - -#[tauri::command] -pub async fn lint_sql( - app: AppHandle, - state: State<'_, AppState>, - input: LintSqlRequest, -) -> Result { - let (connection_id, engine) = resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; - let sql = input.sql.trim().to_string(); - if sql.is_empty() { - return Ok(LintSqlResult { - diagnostics: Vec::new(), - }); - } - if sql.len() > MAX_LINT_SQL_BYTES { - return Err("SQL is too large to lint in the editor.".to_string()); - } - - match engine { - DatabaseEngine::Postgres => { - with_pool_client_retry(&app, &state, &connection_id, sql, |client, sql| async move { - let lint_sql = format!("EXPLAIN {}", sql); - let diagnostics = match client.simple_query(&lint_sql).await { - Ok(_) => Vec::new(), - Err(error) => { - let (line, column) = error_line_column(&error, &sql) - .map(|(l, c)| (Some(l), Some(c))) - .unwrap_or((None, None)); - vec![SqlDiagnostic { - message: map_pg_err(error, Some(sql.as_str())), - severity: "error".to_string(), - line, - column, - end_line: line, - end_column: column.map(|value| value + 1), - }] - } - }; - Ok(LintSqlResult { diagnostics }) - }) - .await - } - DatabaseEngine::Mysql => { - let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; - let lint_sql = format!("EXPLAIN {}", sql); - let diagnostics = match sqlx::query(&lint_sql).execute(&pool).await { - Ok(_) => Vec::new(), - Err(error) => vec![SqlDiagnostic { - message: error.to_string(), - severity: "error".to_string(), - line: None, - column: None, - end_line: None, - end_column: None, - }], - }; - Ok(LintSqlResult { diagnostics }) - } - DatabaseEngine::Sqlite => { - let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; - let lint_sql = format!("EXPLAIN QUERY PLAN {}", sql); - let diagnostics = match sqlx::query(&lint_sql).execute(&pool).await { - Ok(_) => Vec::new(), - Err(error) => vec![SqlDiagnostic { - message: error.to_string(), - severity: "error".to_string(), - line: None, - column: None, - end_line: None, - end_column: None, - }], - }; - Ok(LintSqlResult { diagnostics }) - } - } -} - -#[tauri::command] -pub async fn get_tables( - app: AppHandle, - state: State<'_, AppState>, - connection_id: Option, -) -> Result, String> { - let (connection_id, engine) = resolve_connection_engine(&app, &state, connection_id).await?; - - match engine { - DatabaseEngine::Postgres => { - with_pool_client_retry(&app, &state, &connection_id, (), |client, ()| async move { - let rows = client - .query( - " - select table_schema, table_name - from information_schema.tables - where table_type = 'BASE TABLE' - and table_schema not in ('pg_catalog', 'information_schema') - order by table_schema, table_name - ", - &[], - ) - .await - .map_err(|error| map_pg_err(error, None))?; - - Ok(rows - .into_iter() - .map(|row| { - let schema: String = row.get(0); - let name: String = row.get(1); - let preview_query = format!( - "select * from \"{}\".\"{}\" limit 100;", - quote_identifier(&schema), - quote_identifier(&name) - ); - - TableInfo { - schema, - name, - preview_query, - } - }) - .collect()) - }) - .await - } - DatabaseEngine::Mysql => { - let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; - let database = load_connection(&app, &connection_id)? - .map(|connection| connection.database) - .unwrap_or_default(); - let rows = sqlx::query( - " - select table_schema, table_name - from information_schema.tables - where table_type = 'BASE TABLE' - and table_schema = ? - and table_schema not in ('information_schema', 'mysql', 'performance_schema', 'sys') - order by table_schema, table_name - ", - ) - .bind(&database) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - let mut tables = Vec::new(); - for row in rows { - let schema: String = mysql_get_string(&row, 0, "table_schema", "get_tables")?; - let name: String = mysql_get_string(&row, 1, "table_name", "get_tables")?; - tables.push(TableInfo { - preview_query: format!("select * from `{}`.`{}` limit 100;", schema, name), - schema, - name, - }); - } - Ok(tables) - } - DatabaseEngine::Sqlite => { - let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; - let rows = sqlx::query( - " - select name - from sqlite_master - where type = 'table' - and name not like 'sqlite_%' - order by name - ", - ) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - let mut tables = Vec::new(); - for row in rows { - let name: String = sqlite_get_idx(&row, 0, "name", "get_tables")?; - require_safe_identifier(&name, "table name")?; - tables.push(TableInfo { - schema: "main".to_string(), - preview_query: format!("select * from \"{}\" limit 100;", quote_identifier(&name)), - name, - }); - } - Ok(tables) - } - } -} - -#[tauri::command] -pub async fn get_schema( - app: AppHandle, - state: State<'_, AppState>, - input: SchemaRequest, -) -> Result, String> { - let (connection_id, engine) = resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; - let schema_request = input.clone(); - - match engine { - DatabaseEngine::Postgres => { - with_pool_client_retry( - &app, - &state, - &connection_id, - schema_request, - |client, input| async move { - let rows = client - .query( - " - select table_schema, table_name, column_name, data_type, is_nullable - from information_schema.columns - where table_schema = $1 and table_name = $2 - order by ordinal_position - ", - &[&input.table_schema, &input.table_name], - ) - .await - .map_err(|error| map_pg_err(error, None))?; - - Ok(rows - .into_iter() - .map(|row| ColumnInfo { - table_schema: row.get(0), - table_name: row.get(1), - column_name: row.get(2), - data_type: row.get(3), - is_nullable: row.get::<_, String>(4) == "YES", - }) - .collect()) - }, - ) - .await - } - DatabaseEngine::Mysql => { - let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; - let rows = sqlx::query( - " - select table_schema, table_name, column_name, data_type, is_nullable - from information_schema.columns - where table_schema = ? and table_name = ? - order by ordinal_position - ", - ) - .bind(&schema_request.table_schema) - .bind(&schema_request.table_name) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - let mut columns = Vec::new(); - for row in rows { - columns.push(ColumnInfo { - table_schema: mysql_get_string(&row, 0, "table_schema", "get_schema")?, - table_name: mysql_get_string(&row, 1, "table_name", "get_schema")?, - column_name: mysql_get_string(&row, 2, "column_name", "get_schema")?, - data_type: mysql_get_string(&row, 3, "data_type", "get_schema")?, - is_nullable: mysql_get_string(&row, 4, "is_nullable", "get_schema")? == "YES", - }); - } - Ok(columns) - } - DatabaseEngine::Sqlite => { - let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; - require_safe_identifier(&schema_request.table_name, "table name")?; - let pragma_sql = format!( - "PRAGMA table_info(\"{}\");", - quote_identifier(&schema_request.table_name) - ); - let rows = sqlx::query(&pragma_sql) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - let mut columns = Vec::new(); - for row in rows { - let col_name: String = sqlite_get_name(&row, "name", "get_schema")?; - let col_type: String = sqlite_get_name(&row, "type", "get_schema")?; - let notnull: i64 = sqlite_get_name(&row, "notnull", "get_schema")?; - columns.push(ColumnInfo { - table_schema: "main".to_string(), - table_name: schema_request.table_name.clone(), - column_name: col_name, - data_type: col_type, - is_nullable: notnull == 0, - }); - } - Ok(columns) - } - } -} - -fn veloxdb_unique_constraint_name(table_name: &str, column_name: &str) -> String { - // Postgres constraint names are limited to 63 bytes. - // Keep this deterministic so we can drop the exact constraint later. - let suffix = "_uniq"; - let max_base_len = 63usize.saturating_sub(suffix.len()); - - let mut base = format!("veloxdb_{}_{}", table_name, column_name); - base.truncate(max_base_len); - - format!("{}{}", base, suffix) -} - -#[tauri::command] -pub async fn get_table_properties( - app: AppHandle, - state: State<'_, AppState>, - input: SchemaRequest, -) -> Result, String> { - let (connection_id, engine) = resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; - let ctx = input.clone(); - - if engine == DatabaseEngine::Mysql { - let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; - let rows = sqlx::query( - " - select - c.table_schema, - c.table_name, - c.column_name, - c.data_type, - c.is_nullable, - c.column_default, - c.extra - from information_schema.columns c - where c.table_schema = ? and c.table_name = ? - order by c.ordinal_position - ", - ) - .bind(&ctx.table_schema) - .bind(&ctx.table_name) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - - let pk_rows = sqlx::query( - " - select column_name - from information_schema.key_column_usage - where table_schema = ? and table_name = ? and constraint_name = 'PRIMARY' - ", - ) - .bind(&ctx.table_schema) - .bind(&ctx.table_name) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - let pk_cols: HashSet = pk_rows - .into_iter() - .map(|row| mysql_get_string(&row, 0, "column_name", "get_table_properties")) - .collect::, _>>()?; - - let unique_rows = sqlx::query( - " - select index_name, column_name, seq_in_index - from information_schema.statistics - where table_schema = ? - and table_name = ? - and non_unique = 0 - order by index_name, seq_in_index - ", - ) - .bind(&ctx.table_schema) - .bind(&ctx.table_name) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - let mut unique_by_index: HashMap> = HashMap::new(); - for row in unique_rows { - let index_name: String = mysql_get_string(&row, 0, "index_name", "get_table_properties")?; - if index_name == "PRIMARY" { - continue; - } - let column_name: String = mysql_get_string(&row, 1, "column_name", "get_table_properties")?; - unique_by_index.entry(index_name).or_default().push(column_name); - } - let mut unique_cols: HashSet = HashSet::new(); - let mut composite_unique_cols: HashSet = HashSet::new(); - for cols in unique_by_index.values() { - for col in cols { - unique_cols.insert(col.clone()); - } - if cols.len() > 1 { - for col in cols { - composite_unique_cols.insert(col.clone()); - } - } - } - - let mut properties = Vec::new(); - for row in rows { - let column_name: String = mysql_get_string(&row, 2, "column_name", "get_table_properties")?; - let is_primary_key = pk_cols.contains(&column_name); - let is_unique = is_primary_key || unique_cols.contains(&column_name); - let is_part_of_composite_unique = composite_unique_cols.contains(&column_name); - let extra: String = mysql_get_string(&row, 6, "extra", "get_table_properties")?; - let lower_extra = extra.to_lowercase(); - properties.push(ColumnProperties { - table_schema: mysql_get_string(&row, 0, "table_schema", "get_table_properties")?, - table_name: mysql_get_string(&row, 1, "table_name", "get_table_properties")?, - column_name, - data_type: mysql_get_string(&row, 3, "data_type", "get_table_properties")?, - is_nullable: mysql_get_string(&row, 4, "is_nullable", "get_table_properties")? == "YES", - is_primary_key, - is_unique, - is_part_of_composite_unique, - column_default: mysql_get_optional_string(&row, 5, "column_default", "get_table_properties")?, - is_identity: lower_extra.contains("auto_increment"), - identity_generation: if lower_extra.contains("auto_increment") { - Some("BY DEFAULT".to_string()) - } else { - None - }, - is_generated: if lower_extra.contains("generated") { - Some("ALWAYS".to_string()) - } else { - None - }, - }); - } - return Ok(properties); - } - - if engine == DatabaseEngine::Sqlite { - let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; - require_safe_identifier(&ctx.table_name, "table name")?; - let pragma_sql = format!("PRAGMA table_info(\"{}\");", quote_identifier(&ctx.table_name)); - let rows = sqlx::query(&pragma_sql) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - let index_list_sql = format!("PRAGMA index_list(\"{}\");", quote_identifier(&ctx.table_name)); - let index_rows = sqlx::query(&index_list_sql) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - let mut unique_cols: HashSet = HashSet::new(); - let mut composite_unique_cols: HashSet = HashSet::new(); - for index in index_rows { - let is_unique = sqlite_get_name::(&index, "unique", "get_table_properties")? == 1; - if !is_unique { - continue; - } - let origin = sqlite_get_name::(&index, "origin", "get_table_properties")?; - if origin == "pk" { - continue; - } - let index_name = sqlite_get_name::(&index, "name", "get_table_properties")?; - require_safe_identifier(&index_name, "index name")?; - let info_sql = format!("PRAGMA index_info(\"{}\");", quote_identifier(&index_name)); - let info_rows = sqlx::query(&info_sql) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - let mut cols: Vec = Vec::new(); - for info in info_rows { - if let Ok(name) = sqlite_get_name::(&info, "name", "get_table_properties") { - cols.push(name); - } - } - for col in &cols { - unique_cols.insert(col.clone()); - } - if cols.len() > 1 { - for col in cols { - composite_unique_cols.insert(col); - } - } - } - let mut properties = Vec::new(); - for row in rows { - let column_name: String = sqlite_get_name(&row, "name", "get_table_properties")?; - let is_primary_key = sqlite_get_name::(&row, "pk", "get_table_properties")? == 1; - let is_unique = is_primary_key || unique_cols.contains(&column_name); - let is_part_of_composite_unique = composite_unique_cols.contains(&column_name); - properties.push(ColumnProperties { - table_schema: "main".to_string(), - table_name: ctx.table_name.clone(), - column_name, - data_type: sqlite_get_name(&row, "type", "get_table_properties")?, - is_nullable: sqlite_get_name::(&row, "notnull", "get_table_properties")? == 0, - is_primary_key, - is_unique, - is_part_of_composite_unique, - column_default: sqlite_get_name::>(&row, "dflt_value", "get_table_properties")?, - is_identity: false, - identity_generation: None, - is_generated: None, - }); - } - return Ok(properties); - } - - with_pool_client_retry(&app, &state, &connection_id, ctx, |client, input| async move { - let columns = client - .query( - " - select - c.table_schema, - c.table_name, - c.column_name, - c.data_type, - c.is_nullable, - c.column_default, - c.is_identity, - c.identity_generation, - c.is_generated - from information_schema.columns c - where c.table_schema = $1 and c.table_name = $2 - order by c.ordinal_position - ", - &[&input.table_schema, &input.table_name], - ) - .await - .map_err(|error| map_pg_err(error, None))?; - - let primary_keys = client - .query( - " - select kcu.column_name - from information_schema.table_constraints tc - join information_schema.key_column_usage kcu - on tc.constraint_name = kcu.constraint_name - and tc.table_schema = kcu.table_schema - where tc.table_schema = $1 - and tc.table_name = $2 - and tc.constraint_type = 'PRIMARY KEY' - order by kcu.ordinal_position - ", - &[&input.table_schema, &input.table_name], - ) - .await - .map_err(|error| map_pg_err(error, None))?; - - let primary_key_columns: HashSet = primary_keys - .into_iter() - .filter_map(|row| Some(row.get::<_, String>(0))) - .collect(); - - let unique_constraints = client - .query( - " - select tc.constraint_name, kcu.column_name, kcu.ordinal_position - from information_schema.table_constraints tc - join information_schema.key_column_usage kcu - on tc.constraint_name = kcu.constraint_name - and tc.table_schema = kcu.table_schema - where tc.table_schema = $1 - and tc.table_name = $2 - and tc.constraint_type = 'UNIQUE' - order by tc.constraint_name, kcu.ordinal_position - ", - &[&input.table_schema, &input.table_name], - ) - .await - .map_err(|error| map_pg_err(error, None))?; - - let mut unique_by_name: HashMap> = HashMap::new(); - for row in unique_constraints { - let constraint_name: String = row.get(0); - let column_name: String = row.get(1); - unique_by_name.entry(constraint_name).or_default().push(column_name); - } - - let mut unique_columns: HashSet = HashSet::new(); - let mut composite_unique_columns: HashSet = HashSet::new(); - - for (_constraint_name, cols) in unique_by_name { - for c in &cols { - unique_columns.insert(c.clone()); - } - if cols.len() > 1 { - for c in &cols { - composite_unique_columns.insert(c.clone()); - } - } - } - - Ok(columns - .into_iter() - .map(|row| { - let table_schema: String = row.get(0); - let table_name: String = row.get(1); - let column_name: String = row.get(2); - let data_type: String = row.get(3); - let is_nullable = row.get::<_, String>(4) == "YES"; - let column_default: Option = row.get(5); - let is_identity = row.get::<_, Option>(6).as_deref() == Some("YES"); - let identity_generation: Option = row.get(7); - let is_generated: Option = row.get(8); - - let is_primary_key = primary_key_columns.contains(&column_name); - let is_unique = is_primary_key || unique_columns.contains(&column_name); - let is_part_of_composite_unique = composite_unique_columns.contains(&column_name); - - ColumnProperties { - table_schema, - table_name, - column_name, - data_type, - is_nullable, - is_primary_key, - is_unique, - is_part_of_composite_unique, - column_default, - is_identity, - identity_generation, - is_generated, - } - }) - .collect()) - }) - .await -} - -#[tauri::command] -pub async fn apply_table_properties( - app: AppHandle, - state: State<'_, AppState>, - input: TablePropertiesApplyRequest, -) -> Result<(), String> { - let (connection_id, engine) = resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; - if engine != DatabaseEngine::Postgres { - return Err(format!( - "Table property editing is not supported for {} connections yet.", - match engine { - DatabaseEngine::Postgres => "PostgreSQL", - DatabaseEngine::Mysql => "MySQL", - DatabaseEngine::Sqlite => "SQLite", - } - )); - } - - with_pool_client_retry(&app, &state, &connection_id, input, |mut client, input| async move { - let table_schema = input.table_schema; - let table_name = input.table_name; - let columns = input.columns; - - require_safe_identifier(&table_schema, "schema name")?; - require_safe_identifier(&table_name, "table name")?; - - let current_columns = client - .query( - " - select column_name, is_nullable - from information_schema.columns - where table_schema = $1 and table_name = $2 - ", - &[&table_schema, &table_name], - ) - .await - .map_err(|error| map_pg_err(error, None))?; - - let mut current_nullable: HashMap = HashMap::new(); - for row in current_columns { - let column_name: String = row.get(0); - let is_nullable = row.get::<_, String>(1) == "YES"; - current_nullable.insert(column_name, is_nullable); - } - - let primary_keys = client - .query( - " - select kcu.column_name - from information_schema.table_constraints tc - join information_schema.key_column_usage kcu - on tc.constraint_name = kcu.constraint_name - and tc.table_schema = kcu.table_schema - where tc.table_schema = $1 - and tc.table_name = $2 - and tc.constraint_type = 'PRIMARY KEY' - ", - &[&table_schema, &table_name], - ) - .await - .map_err(|error| map_pg_err(error, None))?; - - let primary_key_columns: HashSet = primary_keys - .into_iter() - .filter_map(|row| Some(row.get::<_, String>(0))) - .collect(); - - let unique_constraints = client - .query( - " - select tc.constraint_name, kcu.column_name, kcu.ordinal_position - from information_schema.table_constraints tc - join information_schema.key_column_usage kcu - on tc.constraint_name = kcu.constraint_name - and tc.table_schema = kcu.table_schema - where tc.table_schema = $1 - and tc.table_name = $2 - and tc.constraint_type = 'UNIQUE' - order by tc.constraint_name, kcu.ordinal_position - ", - &[&table_schema, &table_name], - ) - .await - .map_err(|error| map_pg_err(error, None))?; - - let mut unique_by_name: HashMap> = HashMap::new(); - for row in unique_constraints { - let constraint_name: String = row.get(0); - let column_name: String = row.get(1); - unique_by_name.entry(constraint_name).or_default().push(column_name); - } - - let mut composite_unique_columns: HashSet = HashSet::new(); - let mut single_unique_constraint_names_by_column: HashMap> = HashMap::new(); - - for (constraint_name, cols) in &unique_by_name { - if cols.len() > 1 { - for c in cols { - composite_unique_columns.insert(c.clone()); - } - } else if cols.len() == 1 { - let c = &cols[0]; - single_unique_constraint_names_by_column - .entry(c.clone()) - .or_default() - .push(constraint_name.clone()); - } - } - - let mut desired_by_column: HashMap = HashMap::new(); - for update in columns { - desired_by_column.insert(update.column_name, (update.is_nullable, update.is_unique)); - } - - let txn = client.transaction().await.map_err(|error| map_pg_err(error, None))?; - - // 1) Nullable changes - for (column_name, (desired_is_nullable, _desired_is_unique)) in &desired_by_column { - let current_is_nullable = current_nullable - .get(column_name) - .ok_or_else(|| format!("Unknown column: {}", column_name))?; - - if *current_is_nullable == *desired_is_nullable { - continue; - } - - let qualified_table = format!( - "\"{}\".\"{}\"", - quote_identifier(&table_schema), - quote_identifier(&table_name) - ); - - require_safe_identifier(column_name, "column name")?; - let qualified_column = format!("\"{}\"", quote_identifier(column_name)); - - if *desired_is_nullable { - let sql = format!( - "ALTER TABLE {} ALTER COLUMN {} DROP NOT NULL", - qualified_table, qualified_column - ); - txn.execute(sql.as_str(), &[]).await.map_err(|error| map_pg_err(error, Some(sql.as_str())))?; - } else { - let sql = format!( - "ALTER TABLE {} ALTER COLUMN {} SET NOT NULL", - qualified_table, qualified_column - ); - txn.execute(sql.as_str(), &[]).await.map_err(|error| map_pg_err(error, Some(sql.as_str())))?; - } - } - - // 2) UNIQUE changes (v1: only support single-column UNIQUE constraints) - for (column_name, (_desired_is_nullable, desired_is_unique)) in &desired_by_column { - let is_primary_key = primary_key_columns.contains(column_name); - let is_part_of_composite_unique = composite_unique_columns.contains(column_name); - - if !*desired_is_unique { - if is_primary_key { - return Err(format!( - "Cannot disable UNIQUE for primary key column: {}", - column_name - )); - } - - if is_part_of_composite_unique { - return Err(format!( - "Cannot disable UNIQUE for column in a composite UNIQUE constraint: {}", - column_name - )); - } - } - - // Compute current uniqueness: - let has_single_unique = single_unique_constraint_names_by_column - .get(column_name) - .map(|names| !names.is_empty()) - .unwrap_or(false); - - let current_is_unique = is_primary_key || has_single_unique || is_part_of_composite_unique; - - if *desired_is_unique == current_is_unique { - continue; - } - - let qualified_table = format!( - "\"{}\".\"{}\"", - quote_identifier(&table_schema), - quote_identifier(&table_name) - ); - require_safe_identifier(column_name, "column name")?; - let qualified_column = format!("\"{}\"", quote_identifier(column_name)); - - if *desired_is_unique { - // Add a new single-column UNIQUE constraint. - if current_is_unique { - continue; - } - - let generated_name = veloxdb_unique_constraint_name(&table_name, column_name); - - // If a constraint with that name exists and doesn't match our target column, fail fast. - if let Some(existing_cols) = unique_by_name.get(&generated_name) { - if existing_cols.len() != 1 || existing_cols[0] != *column_name { - return Err(format!( - "Cannot create UNIQUE constraint due to name collision ({}). Rename the existing constraint.", - generated_name - )); - } - } - - let sql = format!( - "ALTER TABLE {} ADD CONSTRAINT \"{}\" UNIQUE ({})", - qualified_table, - quote_identifier(&generated_name), - qualified_column - ); - txn.execute(sql.as_str(), &[]).await.map_err(|error| map_pg_err(error, Some(sql.as_str())))?; - } else { - // Drop the existing single-column UNIQUE constraints for this column. - let constraint_names = single_unique_constraint_names_by_column - .get(column_name) - .cloned() - .unwrap_or_default(); - - for constraint_name in constraint_names { - require_safe_identifier(&constraint_name, "constraint name")?; - let sql = format!( - "ALTER TABLE {} DROP CONSTRAINT \"{}\"", - qualified_table, - quote_identifier(&constraint_name) - ); - txn.execute(sql.as_str(), &[]).await.map_err(|error| map_pg_err(error, Some(sql.as_str())))?; - } - } - } - - txn.commit().await.map_err(|error| map_pg_err(error, None))?; - Ok(()) - }) - .await -} - -#[tauri::command] -pub async fn get_foreign_keys( - app: AppHandle, - state: State<'_, AppState>, - connection_id: Option, -) -> Result, String> { - let (connection_id, engine) = resolve_connection_engine(&app, &state, connection_id).await?; - - if engine == DatabaseEngine::Mysql { - let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; - let rows = sqlx::query( - " - select - kcu.table_schema as from_schema, - kcu.table_name as from_table, - kcu.column_name as from_column, - kcu.referenced_table_schema as to_schema, - kcu.referenced_table_name as to_table, - kcu.referenced_column_name as to_column - from information_schema.key_column_usage kcu - where kcu.referenced_table_name is not null - order by kcu.table_schema, kcu.table_name, kcu.ordinal_position - limit ? - ", - ) - .bind(MAX_FOREIGN_KEY_ROWS) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - let mut edges = Vec::new(); - for row in rows { - edges.push(ForeignKeyEdge { - from_schema: mysql_get_string(&row, 0, "from_schema", "get_foreign_keys")?, - from_table: mysql_get_string(&row, 1, "from_table", "get_foreign_keys")?, - from_column: mysql_get_string(&row, 2, "from_column", "get_foreign_keys")?, - to_schema: mysql_get_string(&row, 3, "to_schema", "get_foreign_keys")?, - to_table: mysql_get_string(&row, 4, "to_table", "get_foreign_keys")?, - to_column: mysql_get_string(&row, 5, "to_column", "get_foreign_keys")?, - }); - } - return Ok(edges); - } - - if engine == DatabaseEngine::Sqlite { - let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; - let tables = sqlx::query( - " - select name - from sqlite_master - where type = 'table' - and name not like 'sqlite_%' - ", - ) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - let mut edges = Vec::new(); - for table in tables { - let table_name: String = sqlite_get_idx(&table, 0, "name", "get_foreign_keys")?; - require_safe_identifier(&table_name, "table name")?; - let fk_sql = format!("PRAGMA foreign_key_list(\"{}\");", quote_identifier(&table_name)); - let fk_rows = sqlx::query(&fk_sql) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - for row in fk_rows { - edges.push(ForeignKeyEdge { - from_schema: "main".to_string(), - from_table: table_name.clone(), - from_column: sqlite_get_name(&row, "from", "get_foreign_keys")?, - to_schema: "main".to_string(), - to_table: sqlite_get_name(&row, "table", "get_foreign_keys")?, - to_column: sqlite_get_name(&row, "to", "get_foreign_keys")?, - }); - if edges.len() >= MAX_FOREIGN_KEY_ROWS as usize { - return Ok(edges); - } - } - } - return Ok(edges); - } - - with_pool_client_retry(&app, &state, &connection_id, (), |client, ()| async move { - let rows = client - .query( - " - select - src_ns.nspname::text as from_schema, - src_cls.relname::text as from_table, - src_att.attname::text as from_column, - tgt_ns.nspname::text as to_schema, - tgt_cls.relname::text as to_table, - tgt_att.attname::text as to_column - from pg_constraint c - join pg_class src_cls on src_cls.oid = c.conrelid - join pg_namespace src_ns on src_ns.oid = src_cls.relnamespace - join pg_class tgt_cls on tgt_cls.oid = c.confrelid - join pg_namespace tgt_ns on tgt_ns.oid = tgt_cls.relnamespace - cross join lateral unnest(c.conkey, c.confkey) as u(attnum, confattnum) - join pg_attribute src_att - on src_att.attrelid = c.conrelid - and src_att.attnum = u.attnum - and not src_att.attisdropped - join pg_attribute tgt_att - on tgt_att.attrelid = c.confrelid - and tgt_att.attnum = u.confattnum - and not tgt_att.attisdropped - where c.contype = 'f' - and src_ns.nspname not in ('pg_catalog', 'information_schema') - order by src_ns.nspname, src_cls.relname, c.conname, u.attnum - limit $1 - ", - &[&MAX_FOREIGN_KEY_ROWS], - ) - .await - .map_err(|error| error.to_string())?; - - Ok(rows - .into_iter() - .map(|row| ForeignKeyEdge { - from_schema: row.get(0), - from_table: row.get(1), - from_column: row.get(2), - to_schema: row.get(3), - to_table: row.get(4), - to_column: row.get(5), - }) - .collect()) - }) - .await -} - -#[tauri::command] -pub async fn get_table_indexes( - app: AppHandle, - state: State<'_, AppState>, - input: SchemaRequest, -) -> Result { - let (connection_id, engine) = resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; - let ctx = input.clone(); - - if engine == DatabaseEngine::Mysql { - let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; - let fetch_limit = MAX_TABLE_INDEX_ROWS + 1; - let rows = sqlx::query( - " - select - table_schema as index_schema, - index_name, - table_schema, - table_name, - non_unique = 0 as is_unique, - index_name = 'PRIMARY' as is_primary, - true as is_valid, - false as is_partial, - concat(index_name, ' (', group_concat(column_name order by seq_in_index separator ', '), ')') as definition, - 0 as index_bytes, - 0 as idx_scan, - 0 as idx_tup_read, - 0 as idx_tup_fetch - from information_schema.statistics - where table_schema = ? - and table_name = ? - group by table_schema, table_name, index_name, non_unique - order by index_name - limit ? - ", - ) - .bind(&ctx.table_schema) - .bind(&ctx.table_name) - .bind(fetch_limit) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - let truncated = rows.len() as i64 > MAX_TABLE_INDEX_ROWS; - let mut indexes = Vec::new(); - for row in rows.into_iter().take(MAX_TABLE_INDEX_ROWS as usize) { - indexes.push(IndexInfo { - index_schema: mysql_get_string(&row, 0, "index_schema", "get_table_indexes")?, - index_name: mysql_get_string(&row, 1, "index_name", "get_table_indexes")?, - table_schema: mysql_get_string(&row, 2, "table_schema", "get_table_indexes")?, - table_name: mysql_get_string(&row, 3, "table_name", "get_table_indexes")?, - is_unique: mysql_get_idx(&row, 4, "is_unique", "get_table_indexes")?, - is_primary: mysql_get_idx(&row, 5, "is_primary", "get_table_indexes")?, - is_valid: true, - is_partial: false, - definition: mysql_get_string(&row, 8, "definition", "get_table_indexes")?, - index_bytes: 0, - idx_scan: 0, - idx_tup_read: 0, - idx_tup_fetch: 0, - }); - } - return Ok(TableIndexesResult { indexes, truncated }); - } - - if engine == DatabaseEngine::Sqlite { - let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; - require_safe_identifier(&ctx.table_name, "table name")?; - let pragma_sql = format!( - "PRAGMA index_list(\"{}\");", - quote_identifier(&ctx.table_name) - ); - let rows = sqlx::query(&pragma_sql) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - let truncated = rows.len() as i64 > MAX_TABLE_INDEX_ROWS; - let mut indexes = Vec::new(); - for row in rows.into_iter().take(MAX_TABLE_INDEX_ROWS as usize) { - let index_name: String = sqlite_get_name(&row, "name", "get_table_indexes")?; - require_safe_identifier(&index_name, "index name")?; - let index_info_sql = format!("PRAGMA index_info(\"{}\");", quote_identifier(&index_name)); - let index_info_rows = sqlx::query(&index_info_sql) - .fetch_all(&pool) - .await - .map_err(|error| error.to_string())?; - let index_columns = index_info_rows - .into_iter() - .filter_map(|idx| sqlite_get_name::(&idx, "name", "get_table_indexes").ok()) - .collect::>(); - indexes.push(IndexInfo { - index_schema: "main".to_string(), - index_name: index_name.clone(), - table_schema: "main".to_string(), - table_name: ctx.table_name.clone(), - is_unique: sqlite_get_name::(&row, "unique", "get_table_indexes")? == 1, - is_primary: sqlite_get_name::(&row, "origin", "get_table_indexes")? == "pk", - is_valid: true, - is_partial: sqlite_get_name::(&row, "partial", "get_table_indexes")? == 1, - definition: if index_columns.is_empty() { - format!("index {}", index_name) - } else { - format!("index {} ({})", index_name, index_columns.join(", ")) - }, - index_bytes: 0, - idx_scan: 0, - idx_tup_read: 0, - idx_tup_fetch: 0, - }); - } - return Ok(TableIndexesResult { indexes, truncated }); - } - - with_pool_client_retry(&app, &state, &connection_id, ctx, |client, input| async move { - let table_schema = input.table_schema; - let table_name = input.table_name; - let fetch_limit = MAX_TABLE_INDEX_ROWS + 1; - - let rows = client - .query( - " - select - ins.nspname::text as index_schema, - ic.relname::text as index_name, - tn.nspname::text as table_schema, - tc.relname::text as table_name, - i.indisunique as is_unique, - i.indisprimary as is_primary, - i.indisvalid as is_valid, - (i.indpred is not null) as is_partial, - pg_get_indexdef(i.indexrelid) as definition, - coalesce(pg_relation_size(i.indexrelid::regclass), 0)::bigint as index_bytes, - coalesce(s.idx_scan, 0)::bigint as idx_scan, - coalesce(s.idx_tup_read, 0)::bigint as idx_tup_read, - coalesce(s.idx_tup_fetch, 0)::bigint as idx_tup_fetch - from pg_index i - join pg_class ic on ic.oid = i.indexrelid - join pg_namespace ins on ins.oid = ic.relnamespace - join pg_class tc on tc.oid = i.indrelid - join pg_namespace tn on tn.oid = tc.relnamespace - left join pg_stat_user_indexes s on s.indexrelid = i.indexrelid - where tn.nspname = $1 - and tc.relname = $2 - and ins.nspname not in ('pg_catalog', 'information_schema') - order by ic.relname - limit $3 - ", - &[&table_schema, &table_name, &fetch_limit], - ) - .await - .map_err(|error| error.to_string())?; - - let truncated = rows.len() as i64 > MAX_TABLE_INDEX_ROWS; - let take = if truncated { - MAX_TABLE_INDEX_ROWS as usize - } else { - rows.len() - }; - - let mut indexes = Vec::with_capacity(take); - for row in rows.into_iter().take(take) { - indexes.push(IndexInfo { - index_schema: row.get(0), - index_name: row.get(1), - table_schema: row.get(2), - table_name: row.get(3), - is_unique: row.get(4), - is_primary: row.get(5), - is_valid: row.get(6), - is_partial: row.get(7), - definition: row.get(8), - index_bytes: row.get(9), - idx_scan: row.get(10), - idx_tup_read: row.get(11), - idx_tup_fetch: row.get(12), - }); - } - - Ok(TableIndexesResult { indexes, truncated }) - }) - .await -} - -#[tauri::command] -pub async fn execute_ddl_transaction( - app: AppHandle, - state: State<'_, AppState>, - input: DdlBatchRequest, -) -> Result<(), String> { - let (connection_id, engine) = resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; - - match engine { - DatabaseEngine::Postgres => { - with_pool_client_retry(&app, &state, &connection_id, input, |mut client, input| async move { - let stmts: Vec = input - .statements - .into_iter() - .map(|s| s.trim().to_string()) - .filter(|s| !s.is_empty()) - .collect(); - - if stmts.is_empty() { - return Err("No SQL statements to execute.".to_string()); - } - - let txn = client.transaction().await.map_err(|error| map_pg_err(error, None))?; - for sql in &stmts { - txn.execute(sql.as_str(), &[]) - .await - .map_err(|error| map_pg_err(error, Some(sql.as_str())))?; - } - txn.commit().await.map_err(|error| map_pg_err(error, None))?; - Ok(()) - }) - .await - } - DatabaseEngine::Mysql => { - let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; - let mut tx = pool.begin().await.map_err(|error| error.to_string())?; - for sql in input - .statements - .iter() - .map(|s| s.trim()) - .filter(|s| !s.is_empty()) - { - sqlx::query(sql) - .execute(&mut *tx) - .await - .map_err(|error| error.to_string())?; - } - tx.commit().await.map_err(|error| error.to_string()) - } - DatabaseEngine::Sqlite => { - let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; - let mut tx = pool.begin().await.map_err(|error| error.to_string())?; - for sql in input - .statements - .iter() - .map(|s| s.trim()) - .filter(|s| !s.is_empty()) - { - sqlx::query(sql) - .execute(&mut *tx) - .await - .map_err(|error| error.to_string())?; - } - tx.commit().await.map_err(|error| error.to_string()) - } - } -} - -/// Run a single DDL statement outside an explicit transaction (required for `CREATE INDEX CONCURRENTLY`). -#[tauri::command] -pub async fn execute_ddl_statement( - app: AppHandle, - state: State<'_, AppState>, - input: DdlStatementRequest, -) -> Result<(), String> { - let (connection_id, engine) = resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; - let sql = input.statement.trim().to_string(); - if sql.is_empty() { - return Err("No SQL statement to execute.".to_string()); - } - - match engine { - DatabaseEngine::Postgres => { - with_pool_client_retry(&app, &state, &connection_id, sql, |client, sql| async move { - client - .execute(sql.as_str(), &[]) - .await - .map_err(|error| map_pg_err(error, Some(sql.as_str())))?; - Ok(()) - }) - .await - } - DatabaseEngine::Mysql => { - let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; - sqlx::query(&sql) - .execute(&pool) - .await - .map_err(|error| error.to_string())?; - Ok(()) - } - DatabaseEngine::Sqlite => { - let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; - sqlx::query(&sql) - .execute(&pool) - .await - .map_err(|error| error.to_string())?; - Ok(()) - } - } -} - -#[tauri::command] -pub async fn export_diagram_png( - input: DiagramExportRequest, - output_path: String, -) -> Result<(), String> { - let path = std::path::PathBuf::from(&output_path); - tokio::task::spawn_blocking(move || export_diagram_to_png(&input, &path)) - .await - .map_err(|e| e.to_string())? -} - -#[tauri::command] -pub async fn export_results_csv_command( - app: AppHandle, - state: State<'_, AppState>, - input: ExportQueryRequest, -) -> Result<(), String> { - export_results_csv(&app, &state, &input).await -} - -#[tauri::command] -pub async fn export_results_json_command( - app: AppHandle, - state: State<'_, AppState>, - input: ExportQueryRequest, -) -> Result<(), String> { - export_results_json(&app, &state, &input).await -} - -#[tauri::command] -pub async fn save_base64_png(data: String, output_path: String) -> Result<(), String> { - use base64::Engine; - let bytes = base64::engine::general_purpose::STANDARD - .decode(data.strip_prefix("data:image/png;base64,").unwrap_or(&data)) - .map_err(|e| e.to_string())?; - std::fs::write(&output_path, bytes).map_err(|e| e.to_string()) -} - -#[tauri::command] -pub async fn save_text_file(content: String, output_path: String) -> Result<(), String> { - std::fs::write(&output_path, content).map_err(|e| e.to_string()) -} - -#[tauri::command] -pub async fn store_openrouter_api_key(api_key: String) -> Result<(), String> { - if api_key.trim().is_empty() { - return credentials::delete_openrouter_api_key(); - } - credentials::store_openrouter_api_key(&api_key) -} - -#[tauri::command] -pub async fn get_openrouter_api_key() -> Result, String> { - credentials::get_openrouter_api_key() -} - -#[tauri::command] -pub async fn delete_openrouter_api_key() -> Result<(), String> { - credentials::delete_openrouter_api_key() -} - -#[cfg(test)] -mod tests { - use super::{ - build_schema_context, classify_sql_intent, database_name_from_mysql_value, - decode_mysql_bytes_as_string, extract_openrouter_stream_delta, mysql_decode_error, - parse_ask_veloxy_json, sqlite_decode_error, streaming_display_text, - validate_generated_sql, - }; - use crate::models::{ - AskVeloxyDbContextCache, DatabaseEngine, QueryEditorColumn, QueryEditorMetadata, - QueryEditorTable, - }; - - #[test] - fn streaming_display_text_extracts_partial_json_message() { - let partial = r#"{ "message": "The messages table has relationships with:\n- delivery_reports"#; - let display = streaming_display_text(partial); - assert!(display.contains("messages table")); - assert!(display.contains("delivery_reports")); - } - - #[test] - fn streaming_display_text_returns_plain_text_directly() { - assert_eq!( - streaming_display_text("Hello from Veloxy"), - "Hello from Veloxy" - ); - } - - #[test] - fn extract_openrouter_stream_delta_reads_content() { - let data = r#"{"choices":[{"delta":{"content":"Hello"}}]}"#; - assert_eq!( - extract_openrouter_stream_delta(data).as_deref(), - Some("Hello") - ); - } - - #[test] - fn database_name_from_mysql_value_rejects_empty() { - assert!(database_name_from_mysql_value(None, "list_databases").is_err()); - assert!(database_name_from_mysql_value(Some(String::new()), "list_databases").is_err()); - } - - #[test] - fn database_name_from_mysql_value_accepts_non_empty() { - let name = - database_name_from_mysql_value(Some("my_app".to_string()), "list_databases").expect("name"); - assert_eq!(name, "my_app"); - } - - #[test] - fn decode_mysql_bytes_as_string_uses_utf8_text() { - assert_eq!( - decode_mysql_bytes_as_string(b"my_schema"), - "my_schema" - ); - } - - #[test] - fn mysql_decode_error_is_explicit() { - let message = mysql_decode_error("get_tables", "table_schema", Some(0), "mismatched types"); - assert!(message.contains("MySQL decode error")); - assert!(message.contains("get_tables")); - assert!(message.contains("table_schema")); - } - - #[test] - fn sqlite_decode_error_is_explicit() { - let message = sqlite_decode_error("get_schema", "name", Some(0), "unsupported value type"); - assert!(message.contains("SQLite decode error")); - assert!(message.contains("get_schema")); - assert!(message.contains("name")); - } - - #[test] - fn schema_context_is_bounded() { - let columns = (0..40) - .map(|idx| QueryEditorColumn { - name: format!("column_{}", idx), - data_type: "text".to_string(), - }) - .collect::>(); - let tables = (0..20) - .map(|idx| QueryEditorTable { - schema: "public".to_string(), - name: format!("events_{}", idx), - columns: columns.clone(), - }) - .collect::>(); - let metadata = QueryEditorMetadata { - tables, - functions: Vec::new(), - truncated_tables: false, - truncated_columns: false, - truncated_functions: false, - }; - let db_context = AskVeloxyDbContextCache { - database_name: "test".to_string(), - engine: DatabaseEngine::Postgres, - metadata, - foreign_keys: Vec::new(), - }; - - let context = build_schema_context(&db_context, "show events", None); - assert!(!context.is_empty()); - assert!(context.len() <= super::ASK_VELOXY_SCHEMA_CHAR_BUDGET); - } - - #[test] - fn ask_veloxy_json_parser_handles_embedded_block() { - let content = "Here is the output {\"sql\":\"select 1\",\"intent\":\"select\",\"confidence\":0.9,\"warnings\":[]}"; - let parsed = parse_ask_veloxy_json(content).expect("json should parse"); - assert_eq!(parsed.get("sql").and_then(|v| v.as_str()), Some("select 1")); - } - - #[test] - fn sql_validation_rejects_multi_statement() { - let multi = "select 1; select 2;"; - assert!(validate_generated_sql(multi).is_err()); - } - - #[test] - fn sql_intent_classifier_recognizes_update() { - assert_eq!(classify_sql_intent("UPDATE foo SET bar = 1"), "update"); - } - - #[test] - fn read_only_check_allows_selects_and_explain() { - assert!(super::is_read_only_sql("SELECT 1")); - assert!(super::is_read_only_sql("EXPLAIN ANALYZE SELECT * FROM t")); - assert!(super::is_read_only_sql("WITH x AS (SELECT 1) SELECT * FROM x")); - assert!(super::is_read_only_sql("BEGIN; SELECT 1; COMMIT;")); - } - - #[test] - fn read_only_check_blocks_writes() { - assert!(!super::is_read_only_sql("DELETE FROM t")); - assert!(!super::is_read_only_sql("DROP TABLE t")); - assert!(!super::is_read_only_sql("BEGIN; UPDATE t SET a = 1; COMMIT;")); - assert!(!super::is_read_only_sql("SELECT 1; DELETE FROM t")); - assert!(!super::is_read_only_sql("")); - } - - #[test] - fn mysql_timestamp_formats_as_datetime_string() { - let dt = chrono::DateTime::parse_from_rfc3339("2024-03-15T10:30:45Z") - .unwrap() - .with_timezone(&chrono::Utc); - assert_eq!(dt.format("%Y-%m-%d %H:%M:%S").to_string(), "2024-03-15 10:30:45"); - } - - #[test] - fn mysql_datetime_formats_as_naive_datetime_string() { - let dt = chrono::NaiveDateTime::parse_from_str("2024-03-15 10:30:45", "%Y-%m-%d %H:%M:%S").unwrap(); - assert_eq!(dt.format("%Y-%m-%d %H:%M:%S").to_string(), "2024-03-15 10:30:45"); - } - - #[test] - fn mysql_date_formats_as_iso_date() { - let d = chrono::NaiveDate::from_ymd_opt(2024, 3, 15).unwrap(); - assert_eq!(d.to_string(), "2024-03-15"); - } - - #[test] - fn mysql_time_formats_as_iso_time() { - let t = chrono::NaiveTime::from_hms_opt(10, 30, 45).unwrap(); - assert_eq!(t.to_string(), "10:30:45"); - } -} diff --git a/src-tauri/src/commands/connections.rs b/src-tauri/src/commands/connections.rs new file mode 100644 index 0000000..7255d16 --- /dev/null +++ b/src-tauri/src/commands/connections.rs @@ -0,0 +1,352 @@ +use uuid::Uuid; +use tauri::{AppHandle, State}; + +use crate::db::{ + build_mysql_pool, build_mysql_pool_custom, build_mongo_connection_string, build_pool, build_pool_custom, build_sqlite_pool, + disconnect_connection, drop_pool, get_or_create_mongo_client, get_or_create_mysql_pool, get_or_create_sqlite_pool, + load_connection, persist_connection_with_password, refresh_connection_pools, + resolve_connection_engine, with_pool_client_retry, AppState, DEFAULT_MYSQL_PORT, +}; +use mongodb::bson::doc; +use crate::credentials; +use crate::models::{ + ConnectionInput, ConnectionSummary, DatabaseEngine, DatabaseInfo, + StoredConnection, SwitchDatabaseRequest, +}; +use crate::pg_error::map_pg_err; +use crate::ssh_tunnel::SshTunnel; + +use super::mysql_database_name_from_row; + +#[tauri::command] +pub async fn connect_db( + app: AppHandle, + state: State<'_, AppState>, + input: ConnectionInput, +) -> Result { + let connection_id = input.id.clone().unwrap_or_else(|| Uuid::new_v4().to_string()); + + match input.engine { + DatabaseEngine::Postgres => { + let pool = if let Some(ref ssh_config) = input.ssh_config { + if ssh_config.is_active() { + let tunnel = match SshTunnel::connect(ssh_config, &input.host, input.port).await { + Ok(tunnel) => tunnel, + Err(e) => return Err(format!("SSH tunnel failed: {}", e)), + }; + let local_port = tunnel.local_port; + state.ssh_tunnels.write().await.insert(connection_id.clone(), tunnel); + build_pool_custom("127.0.0.1", local_port, &input)? + } else { + build_pool(&input)? + } + } else { + build_pool(&input)? + }; + + let client = match pool.get().await { + Ok(client) => client, + Err(e) => { + drop_pool(&state, &connection_id).await; + return Err(e.to_string()); + } + }; + + if let Err(e) = client.simple_query("select 1").await { + drop_pool(&state, &connection_id).await; + return Err(map_pg_err(e, None)); + } + + state.pools.write().await.insert(connection_id.clone(), pool); + } + DatabaseEngine::Mysql => { + let pool = if let Some(ref ssh_config) = input.ssh_config { + if ssh_config.is_active() { + let remote_port = if input.port == 0 { DEFAULT_MYSQL_PORT } else { input.port }; + let tunnel = match SshTunnel::connect(ssh_config, &input.host, remote_port).await { + Ok(tunnel) => tunnel, + Err(e) => return Err(format!("SSH tunnel failed: {}", e)), + }; + let local_port = tunnel.local_port; + state.ssh_tunnels.write().await.insert(connection_id.clone(), tunnel); + build_mysql_pool_custom("127.0.0.1", local_port, &input).await? + } else { + build_mysql_pool(&input).await? + } + } else { + build_mysql_pool(&input).await? + }; + + sqlx::query("select 1").execute(&pool).await.map_err(|e| e.to_string())?; + state.mysql_pools.write().await.insert(connection_id.clone(), pool); + } + DatabaseEngine::Sqlite => { + let pool = build_sqlite_pool(&input).await?; + sqlx::query("select 1").execute(&pool).await.map_err(|e| e.to_string())?; + state.sqlite_pools.write().await.insert(connection_id.clone(), pool); + } + DatabaseEngine::Mongo => { + let uri = build_mongo_connection_string(&input); + let client = mongodb::Client::with_uri_str(&uri) + .await + .map_err(|e| format!("MongoDB connection failed: {}", e))?; + client + .database("admin") + .run_command(doc! { "ping": 1 }) + .await + .map_err(|e| format!("MongoDB ping failed: {}", e))?; + state.mongo_clients.write().await.insert(connection_id.clone(), client); + } + } + + let stored_connection = StoredConnection::from_input(connection_id.clone(), input.clone()); + persist_connection_with_password(&app, &stored_connection, &input.password)?; + + *state.active_connection_id.write().await = Some(connection_id); + + Ok(stored_connection.summary()) +} + +#[tauri::command] +pub async fn list_connections_command(app: AppHandle) -> Result, String> { + crate::db::list_connections(&app) +} + +#[tauri::command] +pub async fn set_active_connection( + app: AppHandle, + state: State<'_, AppState>, + connection_id: String, +) -> Result { + let stored_connection = load_connection(&app, &connection_id)? + .ok_or_else(|| "Stored connection details were not found.".to_string())?; + + match stored_connection.engine { + DatabaseEngine::Postgres => { + with_pool_client_retry(&app, &state, &connection_id, (), |client, ()| async move { + client.simple_query("select 1").await + .map_err(|error| map_pg_err(error, None))?; + Ok(()) + }).await?; + } + DatabaseEngine::Mysql => { + let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; + sqlx::query("select 1").execute(&pool).await.map_err(|error| error.to_string())?; + } + DatabaseEngine::Sqlite => { + let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; + sqlx::query("select 1").execute(&pool).await.map_err(|error| error.to_string())?; + } + DatabaseEngine::Mongo => { + let client = get_or_create_mongo_client(&app, &state, &connection_id).await?; + client.database("admin").run_command(doc! { "ping": 1 }).await + .map_err(|e| format!("MongoDB ping failed: {}", e))?; + } + } + + *state.active_connection_id.write().await = Some(connection_id); + Ok(stored_connection.summary()) +} + +#[tauri::command] +pub async fn ping_connection( + app: AppHandle, + state: State<'_, AppState>, + connection_id: String, +) -> Result<(), String> { + let stored_connection = load_connection(&app, &connection_id)? + .ok_or_else(|| "Stored connection details were not found.".to_string())?; + match stored_connection.engine { + DatabaseEngine::Postgres => { + with_pool_client_retry(&app, &state, &connection_id, (), |client, ()| async move { + client.simple_query("select 1").await.map_err(|error| error.to_string())?; + Ok(()) + }).await + } + DatabaseEngine::Mysql => { + let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; + sqlx::query("select 1").execute(&pool).await.map_err(|error| error.to_string())?; + Ok(()) + } + DatabaseEngine::Sqlite => { + let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; + sqlx::query("select 1").execute(&pool).await.map_err(|error| error.to_string())?; + Ok(()) + } + DatabaseEngine::Mongo => { + let client = get_or_create_mongo_client(&app, &state, &connection_id).await?; + client.database("admin").run_command(doc! { "ping": 1 }).await + .map_err(|error| error.to_string())?; + Ok(()) + } + } +} + +#[tauri::command] +pub async fn refresh_connection( + app: AppHandle, + state: State<'_, AppState>, + connection_id: String, +) -> Result<(), String> { + refresh_connection_pools(&app, &state, &connection_id).await +} + +#[tauri::command] +pub async fn disconnect_db( + state: State<'_, AppState>, + connection_id: String, +) -> Result<(), String> { + disconnect_connection(&state, &connection_id).await; + Ok(()) +} + +#[tauri::command] +pub async fn rename_connection( + app: AppHandle, + connection_id: String, + new_name: String, +) -> Result { + crate::db::rename_connection_in_store(&app, &connection_id, &new_name) +} + +#[tauri::command] +pub async fn delete_connection( + app: AppHandle, + state: State<'_, AppState>, + connection_id: String, +) -> Result<(), String> { + disconnect_connection(&state, &connection_id).await; + if let Err(e) = credentials::delete_password(&connection_id) { + log::warn!("Failed to delete keychain entry for {}: {}", connection_id, e); + } + crate::db::delete_connection_from_store(&app, &connection_id)?; + Ok(()) +} + +#[tauri::command] +pub async fn list_databases( + app: AppHandle, + state: State<'_, AppState>, + connection_id: Option, +) -> Result, String> { + let (connection_id, engine) = resolve_connection_engine(&app, &state, connection_id).await?; + + match engine { + DatabaseEngine::Postgres => { + with_pool_client_retry(&app, &state, &connection_id, (), |client, ()| async move { + let rows = client.query( + "select datname from pg_database \ + where datistemplate = false and has_database_privilege(datname, 'CONNECT') \ + order by datname", + &[], + ).await.map_err(|error| map_pg_err(error, None))?; + Ok(rows.into_iter().map(|row| { + let name: String = row.get(0); + DatabaseInfo { name } + }).collect()) + }).await + } + DatabaseEngine::Mysql => { + let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; + let rows = sqlx::query("show databases") + .fetch_all(&pool).await.map_err(|error| error.to_string())?; + let mut databases = Vec::with_capacity(rows.len()); + for row in rows { + let name = mysql_database_name_from_row(&row, "list_databases")?; + databases.push(DatabaseInfo { name }); + } + Ok(databases) + } + DatabaseEngine::Sqlite => Ok(vec![DatabaseInfo { name: "main".to_string() }]), + DatabaseEngine::Mongo => { + let client = get_or_create_mongo_client(&app, &state, &connection_id).await?; + let db_names = client.list_database_names().await + .map_err(|e| format!("Failed to list MongoDB databases: {}", e))?; + Ok(db_names.into_iter().map(|name| DatabaseInfo { name }).collect()) + } + } +} + +#[tauri::command] +pub async fn switch_database( + app: AppHandle, + state: State<'_, AppState>, + input: SwitchDatabaseRequest, +) -> Result { + let mut stored_connection = load_connection(&app, &input.connection_id)? + .ok_or_else(|| "Stored connection details were not found.".to_string())?; + + if stored_connection.engine == DatabaseEngine::Sqlite { + return Err("Switch database is not supported for SQLite connections.".to_string()); + } + + drop_pool(&state, &input.connection_id).await; + + stored_connection.database = input.database.clone(); + stored_connection.connected_at = crate::models::timestamp_string(); + persist_connection_with_password(&app, &stored_connection, &stored_connection.password.clone().unwrap_or_default())?; + + let connection_input = stored_connection.to_input(); + + match connection_input.engine { + DatabaseEngine::Postgres => { + let pool = if let Some(ref ssh_config) = connection_input.ssh_config { + if ssh_config.is_active() { + let tunnel = match SshTunnel::connect(ssh_config, &connection_input.host, connection_input.port).await { + Ok(tunnel) => tunnel, + Err(e) => return Err(format!("SSH tunnel failed: {}", e)), + }; + let local_port = tunnel.local_port; + state.ssh_tunnels.write().await.insert(input.connection_id.clone(), tunnel); + build_pool_custom("127.0.0.1", local_port, &connection_input)? + } else { + build_pool(&connection_input)? + } + } else { + build_pool(&connection_input)? + }; + + let client = match pool.get().await { + Ok(client) => client, + Err(e) => { drop_pool(&state, &input.connection_id).await; return Err(e.to_string()); } + }; + + if let Err(e) = client.simple_query("select 1").await { + drop_pool(&state, &input.connection_id).await; + return Err(map_pg_err(e, None)); + } + + state.pools.write().await.insert(input.connection_id.clone(), pool); + } + DatabaseEngine::Mysql => { + let pool = if let Some(ref ssh_config) = connection_input.ssh_config { + if ssh_config.is_active() { + let remote_port = if connection_input.port == 0 { DEFAULT_MYSQL_PORT } else { connection_input.port }; + let tunnel = match SshTunnel::connect(ssh_config, &connection_input.host, remote_port).await { + Ok(tunnel) => tunnel, + Err(e) => return Err(format!("SSH tunnel failed: {}", e)), + }; + let local_port = tunnel.local_port; + state.ssh_tunnels.write().await.insert(input.connection_id.clone(), tunnel); + build_mysql_pool_custom("127.0.0.1", local_port, &connection_input).await? + } else { + build_mysql_pool(&connection_input).await? + } + } else { + build_mysql_pool(&connection_input).await? + }; + + sqlx::query("select 1").execute(&pool).await.map_err(|e| e.to_string())?; + state.mysql_pools.write().await.insert(input.connection_id.clone(), pool); + } + DatabaseEngine::Sqlite => {} + DatabaseEngine::Mongo => { + let client = get_or_create_mongo_client(&app, &state, &input.connection_id).await?; + client.database(&input.database).run_command(doc! { "ping": 1 }).await + .map_err(|e| format!("MongoDB ping failed: {}", e))?; + } + } + + *state.active_connection_id.write().await = Some(input.connection_id); + Ok(stored_connection.summary()) +} diff --git a/src-tauri/src/commands/ddl.rs b/src-tauri/src/commands/ddl.rs new file mode 100644 index 0000000..8db338f --- /dev/null +++ b/src-tauri/src/commands/ddl.rs @@ -0,0 +1,93 @@ +use tauri::{AppHandle, State}; + +use crate::db::{ + get_or_create_mysql_pool, get_or_create_sqlite_pool, resolve_connection_engine, + with_pool_client_retry, AppState, +}; +use crate::models::{DatabaseEngine, DdlBatchRequest, DdlStatementRequest}; +use crate::pg_error::map_pg_err; + +#[tauri::command] +pub async fn execute_ddl_transaction( + app: AppHandle, + state: State<'_, AppState>, + input: DdlBatchRequest, +) -> Result<(), String> { + let (connection_id, engine) = resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; + + match engine { + DatabaseEngine::Postgres => { + with_pool_client_retry(&app, &state, &connection_id, input, |mut client, input| async move { + let stmts: Vec = input.statements.into_iter() + .map(|s| s.trim().to_string()).filter(|s| !s.is_empty()).collect(); + + if stmts.is_empty() { + return Err("No SQL statements to execute.".to_string()); + } + + let txn = client.transaction().await.map_err(|error| map_pg_err(error, None))?; + for sql in &stmts { + txn.execute(sql.as_str(), &[]).await + .map_err(|error| map_pg_err(error, Some(sql.as_str())))?; + } + txn.commit().await.map_err(|error| map_pg_err(error, None))?; + Ok(()) + }).await + } + DatabaseEngine::Mysql => { + let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; + let mut tx = pool.begin().await.map_err(|error| error.to_string())?; + for sql in input.statements.iter().map(|s| s.trim()).filter(|s| !s.is_empty()) { + sqlx::query(sql).execute(&mut *tx).await.map_err(|error| error.to_string())?; + } + tx.commit().await.map_err(|error| error.to_string()) + } + DatabaseEngine::Sqlite => { + let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; + let mut tx = pool.begin().await.map_err(|error| error.to_string())?; + for sql in input.statements.iter().map(|s| s.trim()).filter(|s| !s.is_empty()) { + sqlx::query(sql).execute(&mut *tx).await.map_err(|error| error.to_string())?; + } + tx.commit().await.map_err(|error| error.to_string()) + } + DatabaseEngine::Mongo => { + Err("MongoDB does not support DDL transactions.".to_string()) + } + } +} + +#[tauri::command] +pub async fn execute_ddl_statement( + app: AppHandle, + state: State<'_, AppState>, + input: DdlStatementRequest, +) -> Result<(), String> { + let (connection_id, engine) = resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; + let sql = input.statement.trim().to_string(); + if sql.is_empty() { + return Err("No SQL statement to execute.".to_string()); + } + + match engine { + DatabaseEngine::Postgres => { + with_pool_client_retry(&app, &state, &connection_id, sql, |client, sql| async move { + client.execute(sql.as_str(), &[]).await + .map_err(|error| map_pg_err(error, Some(sql.as_str())))?; + Ok(()) + }).await + } + DatabaseEngine::Mysql => { + let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; + sqlx::query(&sql).execute(&pool).await.map_err(|error| error.to_string())?; + Ok(()) + } + DatabaseEngine::Sqlite => { + let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; + sqlx::query(&sql).execute(&pool).await.map_err(|error| error.to_string())?; + Ok(()) + } + DatabaseEngine::Mongo => { + Err("MongoDB does not support DDL statements.".to_string()) + } + } +} diff --git a/src-tauri/src/commands/editor_meta.rs b/src-tauri/src/commands/editor_meta.rs new file mode 100644 index 0000000..a49018d --- /dev/null +++ b/src-tauri/src/commands/editor_meta.rs @@ -0,0 +1,426 @@ +use tauri::{AppHandle, State}; + +use crate::db::{ + get_or_create_mysql_pool, get_or_create_sqlite_pool, load_connection, quote_identifier, + require_safe_identifier, resolve_connection_engine, with_pool_client_retry, AppState, +}; +use crate::models::{ + DatabaseEngine, ForeignKeyEdge, QueryEditorColumn, QueryEditorFunction, + QueryEditorMetadata, QueryEditorTable, +}; +use crate::pg_error::map_pg_err; + +use super::{ + MAX_EDITOR_TABLES, MAX_EDITOR_COLUMNS_PER_TABLE, MAX_EDITOR_FUNCTIONS, MAX_FOREIGN_KEY_ROWS, + mysql_get_string, sqlite_get_idx, sqlite_get_name, +}; + +pub(crate) async fn fetch_query_editor_metadata_for_connection( + app: &AppHandle, + state: &AppState, + connection_id: &str, + engine: DatabaseEngine, +) -> Result { + if engine == DatabaseEngine::Mysql { + let pool = get_or_create_mysql_pool(app, state, connection_id).await?; + let database = load_connection(app, connection_id)? + .map(|connection| connection.database).unwrap_or_default(); + let table_rows = sqlx::query( + "select table_schema, table_name \ + from information_schema.tables \ + where table_type = 'BASE TABLE' \ + and table_schema = ? \ + and table_schema not in ('information_schema', 'mysql', 'performance_schema', 'sys') \ + order by table_schema, table_name \ + limit ?", + ).bind(&database).bind(MAX_EDITOR_TABLES + 1) + .fetch_all(&pool).await.map_err(|error| error.to_string())?; + + let truncated_tables = table_rows.len() as i64 > MAX_EDITOR_TABLES; + let mut tables = Vec::new(); + let mut truncated_columns = false; + + for row in table_rows.into_iter().take(MAX_EDITOR_TABLES as usize) { + let schema: String = mysql_get_string(&row, 0, "table_schema", "get_query_editor_metadata")?; + let name: String = mysql_get_string(&row, 1, "table_name", "get_query_editor_metadata")?; + let column_rows = sqlx::query( + "select column_name, data_type \ + from information_schema.columns \ + where table_schema = ? and table_name = ? \ + order by ordinal_position \ + limit ?", + ).bind(&schema).bind(&name).bind(MAX_EDITOR_COLUMNS_PER_TABLE + 1) + .fetch_all(&pool).await.map_err(|error| error.to_string())?; + if column_rows.len() as i64 > MAX_EDITOR_COLUMNS_PER_TABLE { truncated_columns = true; } + let mut columns = Vec::new(); + for column in column_rows.into_iter().take(MAX_EDITOR_COLUMNS_PER_TABLE as usize) { + columns.push(QueryEditorColumn { + name: mysql_get_string(&column, 0, "column_name", "get_query_editor_metadata")?, + data_type: mysql_get_string(&column, 1, "data_type", "get_query_editor_metadata")?, + }); + } + tables.push(QueryEditorTable { schema, name, columns }); + } + + return Ok(QueryEditorMetadata { + tables, functions: Vec::new(), + truncated_tables, truncated_columns, truncated_functions: false, + }); + } + + if engine == DatabaseEngine::Sqlite { + let pool = get_or_create_sqlite_pool(app, state, connection_id).await?; + let table_rows = sqlx::query( + "select name from sqlite_master \ + where type = 'table' and name not like 'sqlite_%' \ + order by name limit ?", + ).bind(MAX_EDITOR_TABLES + 1).fetch_all(&pool).await.map_err(|error| error.to_string())?; + let truncated_tables = table_rows.len() as i64 > MAX_EDITOR_TABLES; + let mut tables = Vec::new(); + let mut truncated_columns = false; + for row in table_rows.into_iter().take(MAX_EDITOR_TABLES as usize) { + let name: String = sqlite_get_idx(&row, 0, "name", "get_query_editor_metadata")?; + require_safe_identifier(&name, "table name")?; + let pragma_sql = format!("PRAGMA table_info(\"{}\");", quote_identifier(&name)); + let column_rows = sqlx::query(&pragma_sql).fetch_all(&pool).await + .map_err(|error| error.to_string())?; + if column_rows.len() as i64 > MAX_EDITOR_COLUMNS_PER_TABLE { truncated_columns = true; } + let mut columns = Vec::new(); + for column in column_rows.into_iter().take(MAX_EDITOR_COLUMNS_PER_TABLE as usize) { + columns.push(QueryEditorColumn { + name: sqlite_get_name(&column, "name", "get_query_editor_metadata")?, + data_type: sqlite_get_name(&column, "type", "get_query_editor_metadata")?, + }); + } + tables.push(QueryEditorTable { schema: "main".to_string(), name, columns }); + } + return Ok(QueryEditorMetadata { + tables, functions: Vec::new(), + truncated_tables, truncated_columns, truncated_functions: false, + }); + } + + with_pool_client_retry(app, state, connection_id, (), |client, ()| async move { + let table_rows = client.query( + "select n.nspname::text as schema_name, c.relname::text as table_name \ + from pg_class c \ + join pg_namespace n on n.oid = c.relnamespace \ + where c.relkind in ('r', 'p', 'v', 'm', 'f') \ + and n.nspname not in ('pg_catalog', 'information_schema') \ + order by n.nspname, c.relname \ + limit $1", + &[&(MAX_EDITOR_TABLES + 1)], + ).await.map_err(|error| map_pg_err(error, None))?; + + let truncated_tables = table_rows.len() as i64 > MAX_EDITOR_TABLES; + let table_rows = if truncated_tables { + table_rows.into_iter().take(MAX_EDITOR_TABLES as usize).collect::>() + } else { table_rows }; + + let mut tables = Vec::with_capacity(table_rows.len()); + let mut truncated_columns = false; + + for row in table_rows { + let schema: String = row.get(0); + let name: String = row.get(1); + let column_rows = client.query( + "select a.attname::text as column_name, format_type(a.atttypid, a.atttypmod)::text as data_type \ + from pg_attribute a \ + join pg_class c on c.oid = a.attrelid \ + join pg_namespace n on n.oid = c.relnamespace \ + where n.nspname = $1 and c.relname = $2 \ + and a.attnum > 0 and not a.attisdropped \ + order by a.attnum limit $3", + &[&schema, &name, &(MAX_EDITOR_COLUMNS_PER_TABLE + 1)], + ).await.map_err(|error| map_pg_err(error, None))?; + + if column_rows.len() as i64 > MAX_EDITOR_COLUMNS_PER_TABLE { truncated_columns = true; } + let columns = column_rows.into_iter().take(MAX_EDITOR_COLUMNS_PER_TABLE as usize) + .map(|column| QueryEditorColumn { name: column.get(0), data_type: column.get(1) }) + .collect(); + + tables.push(QueryEditorTable { schema, name, columns }); + } + + let function_rows = client.query( + "select n.nspname::text as schema_name, p.proname::text as function_name, \ + coalesce(pg_get_function_identity_arguments(p.oid), '')::text as args, \ + pg_get_function_result(p.oid)::text as return_type \ + from pg_proc p \ + join pg_namespace n on n.oid = p.pronamespace \ + where n.nspname not in ('pg_catalog', 'information_schema') \ + order by n.nspname, p.proname limit $1", + &[&(MAX_EDITOR_FUNCTIONS + 1)], + ).await.map_err(|error| map_pg_err(error, None))?; + + let truncated_functions = function_rows.len() as i64 > MAX_EDITOR_FUNCTIONS; + let functions = function_rows.into_iter().take(MAX_EDITOR_FUNCTIONS as usize) + .map(|row| { + let args_raw: String = row.get(2); + QueryEditorFunction { + schema: row.get(0), + name: row.get(1), + arg_types: if args_raw.trim().is_empty() { Vec::new() } + else { args_raw.split(',').map(|value| value.trim().to_string()).collect() }, + return_type: row.get(3), + } + }).collect(); + + Ok(QueryEditorMetadata { + tables, functions, + truncated_tables, truncated_columns, truncated_functions, + }) + }).await +} + +pub(crate) async fn fetch_foreign_keys_for_connection( + app: &AppHandle, + state: &AppState, + connection_id: &str, + engine: DatabaseEngine, +) -> Result, String> { + if engine == DatabaseEngine::Mysql { + let pool = get_or_create_mysql_pool(app, state, connection_id).await?; + let rows = sqlx::query( + "select kcu.table_schema as from_schema, kcu.table_name as from_table, \ + kcu.column_name as from_column, kcu.referenced_table_schema as to_schema, \ + kcu.referenced_table_name as to_table, kcu.referenced_column_name as to_column \ + from information_schema.key_column_usage kcu \ + where kcu.referenced_table_name is not null \ + order by kcu.table_schema, kcu.table_name, kcu.ordinal_position \ + limit ?", + ).bind(MAX_FOREIGN_KEY_ROWS).fetch_all(&pool).await.map_err(|error| error.to_string())?; + let mut edges = Vec::new(); + for row in rows { + edges.push(ForeignKeyEdge { + from_schema: mysql_get_string(&row, 0, "from_schema", "get_foreign_keys")?, + from_table: mysql_get_string(&row, 1, "from_table", "get_foreign_keys")?, + from_column: mysql_get_string(&row, 2, "from_column", "get_foreign_keys")?, + to_schema: mysql_get_string(&row, 3, "to_schema", "get_foreign_keys")?, + to_table: mysql_get_string(&row, 4, "to_table", "get_foreign_keys")?, + to_column: mysql_get_string(&row, 5, "to_column", "get_foreign_keys")?, + }); + } + return Ok(edges); + } + + if engine == DatabaseEngine::Sqlite { + let pool = get_or_create_sqlite_pool(app, state, connection_id).await?; + let tables = sqlx::query( + "select name from sqlite_master \ + where type = 'table' and name not like 'sqlite_%'", + ).fetch_all(&pool).await.map_err(|error| error.to_string())?; + let mut edges = Vec::new(); + for table in tables { + let table_name: String = sqlite_get_idx(&table, 0, "name", "get_foreign_keys")?; + require_safe_identifier(&table_name, "table name")?; + let fk_sql = format!("PRAGMA foreign_key_list(\"{}\");", quote_identifier(&table_name)); + let fk_rows = sqlx::query(&fk_sql).fetch_all(&pool).await.map_err(|error| error.to_string())?; + for row in fk_rows { + edges.push(ForeignKeyEdge { + from_schema: "main".to_string(), + from_table: table_name.clone(), + from_column: sqlite_get_name(&row, "from", "get_foreign_keys")?, + to_schema: "main".to_string(), + to_table: sqlite_get_name(&row, "table", "get_foreign_keys")?, + to_column: sqlite_get_name(&row, "to", "get_foreign_keys")?, + }); + if edges.len() >= MAX_FOREIGN_KEY_ROWS as usize { + return Ok(edges); + } + } + } + return Ok(edges); + } + + with_pool_client_retry(app, state, connection_id, (), |client, ()| async move { + let rows = client.query( + "select src_ns.nspname::text as from_schema, src_cls.relname::text as from_table, \ + src_att.attname::text as from_column, tgt_ns.nspname::text as to_schema, \ + tgt_cls.relname::text as to_table, tgt_att.attname::text as to_column \ + from pg_constraint c \ + join pg_class src_cls on src_cls.oid = c.conrelid \ + join pg_namespace src_ns on src_ns.oid = src_cls.relnamespace \ + join pg_class tgt_cls on tgt_cls.oid = c.confrelid \ + join pg_namespace tgt_ns on tgt_ns.oid = tgt_cls.relnamespace \ + cross join lateral unnest(c.conkey, c.confkey) as u(attnum, confattnum) \ + join pg_attribute src_att on src_att.attrelid = c.conrelid \ + and src_att.attnum = u.attnum and not src_att.attisdropped \ + join pg_attribute tgt_att on tgt_att.attrelid = c.confrelid \ + and tgt_att.attnum = u.confattnum and not tgt_att.attisdropped \ + where c.contype = 'f' \ + and src_ns.nspname not in ('pg_catalog', 'information_schema') \ + order by src_ns.nspname, src_cls.relname, c.conname, u.attnum \ + limit $1", + &[&MAX_FOREIGN_KEY_ROWS], + ).await.map_err(|error| error.to_string())?; + + Ok(rows.into_iter().map(|row| ForeignKeyEdge { + from_schema: row.get(0), + from_table: row.get(1), + from_column: row.get(2), + to_schema: row.get(3), + to_table: row.get(4), + to_column: row.get(5), + }).collect()) + }).await +} + +#[tauri::command] +pub async fn get_query_editor_metadata( + app: AppHandle, + state: State<'_, AppState>, + connection_id: Option, +) -> Result { + let (connection_id, engine) = resolve_connection_engine(&app, &state, connection_id).await?; + + if engine == DatabaseEngine::Mysql { + let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; + let database = load_connection(&app, &connection_id)? + .map(|connection| connection.database).unwrap_or_default(); + let table_rows = sqlx::query( + "select table_schema, table_name \ + from information_schema.tables \ + where table_type = 'BASE TABLE' \ + and table_schema = ? \ + and table_schema not in ('information_schema', 'mysql', 'performance_schema', 'sys') \ + order by table_schema, table_name \ + limit ?", + ).bind(&database).bind(MAX_EDITOR_TABLES + 1) + .fetch_all(&pool).await.map_err(|error| error.to_string())?; + + let truncated_tables = table_rows.len() as i64 > MAX_EDITOR_TABLES; + let mut tables = Vec::new(); + let mut truncated_columns = false; + + for row in table_rows.into_iter().take(MAX_EDITOR_TABLES as usize) { + let schema: String = mysql_get_string(&row, 0, "table_schema", "get_query_editor_metadata")?; + let name: String = mysql_get_string(&row, 1, "table_name", "get_query_editor_metadata")?; + let column_rows = sqlx::query( + "select column_name, data_type \ + from information_schema.columns \ + where table_schema = ? and table_name = ? \ + order by ordinal_position limit ?", + ).bind(&schema).bind(&name).bind(MAX_EDITOR_COLUMNS_PER_TABLE + 1) + .fetch_all(&pool).await.map_err(|error| error.to_string())?; + if column_rows.len() as i64 > MAX_EDITOR_COLUMNS_PER_TABLE { truncated_columns = true; } + let mut columns = Vec::new(); + for column in column_rows.into_iter().take(MAX_EDITOR_COLUMNS_PER_TABLE as usize) { + columns.push(QueryEditorColumn { + name: mysql_get_string(&column, 0, "column_name", "get_query_editor_metadata")?, + data_type: mysql_get_string(&column, 1, "data_type", "get_query_editor_metadata")?, + }); + } + tables.push(QueryEditorTable { schema, name, columns }); + } + + return Ok(QueryEditorMetadata { + tables, functions: Vec::new(), + truncated_tables, truncated_columns, truncated_functions: false, + }); + } + + if engine == DatabaseEngine::Sqlite { + let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; + let table_rows = sqlx::query( + "select name from sqlite_master \ + where type = 'table' and name not like 'sqlite_%' \ + order by name limit ?", + ).bind(MAX_EDITOR_TABLES + 1).fetch_all(&pool).await.map_err(|error| error.to_string())?; + let truncated_tables = table_rows.len() as i64 > MAX_EDITOR_TABLES; + let mut tables = Vec::new(); + let mut truncated_columns = false; + for row in table_rows.into_iter().take(MAX_EDITOR_TABLES as usize) { + let name: String = sqlite_get_idx(&row, 0, "name", "get_query_editor_metadata")?; + require_safe_identifier(&name, "table name")?; + let pragma_sql = format!("PRAGMA table_info(\"{}\");", quote_identifier(&name)); + let column_rows = sqlx::query(&pragma_sql).fetch_all(&pool).await + .map_err(|error| error.to_string())?; + if column_rows.len() as i64 > MAX_EDITOR_COLUMNS_PER_TABLE { truncated_columns = true; } + let mut columns = Vec::new(); + for column in column_rows.into_iter().take(MAX_EDITOR_COLUMNS_PER_TABLE as usize) { + columns.push(QueryEditorColumn { + name: sqlite_get_name(&column, "name", "get_query_editor_metadata")?, + data_type: sqlite_get_name(&column, "type", "get_query_editor_metadata")?, + }); + } + tables.push(QueryEditorTable { schema: "main".to_string(), name, columns }); + } + return Ok(QueryEditorMetadata { + tables, functions: Vec::new(), + truncated_tables, truncated_columns, truncated_functions: false, + }); + } + + with_pool_client_retry(&app, &state, &connection_id, (), |client, ()| async move { + let table_rows = client.query( + "select n.nspname::text as schema_name, c.relname::text as table_name \ + from pg_class c \ + join pg_namespace n on n.oid = c.relnamespace \ + where c.relkind in ('r', 'p', 'v', 'm', 'f') \ + and n.nspname not in ('pg_catalog', 'information_schema') \ + order by n.nspname, c.relname \ + limit $1", + &[&(MAX_EDITOR_TABLES + 1)], + ).await.map_err(|error| map_pg_err(error, None))?; + + let truncated_tables = table_rows.len() as i64 > MAX_EDITOR_TABLES; + let table_rows = if truncated_tables { + table_rows.into_iter().take(MAX_EDITOR_TABLES as usize).collect::>() + } else { table_rows }; + + let mut tables = Vec::with_capacity(table_rows.len()); + let mut truncated_columns = false; + + for row in table_rows { + let schema: String = row.get(0); + let name: String = row.get(1); + let column_rows = client.query( + "select a.attname::text as column_name, format_type(a.atttypid, a.atttypmod)::text as data_type \ + from pg_attribute a \ + join pg_class c on c.oid = a.attrelid \ + join pg_namespace n on n.oid = c.relnamespace \ + where n.nspname = $1 and c.relname = $2 \ + and a.attnum > 0 and not a.attisdropped \ + order by a.attnum limit $3", + &[&schema, &name, &(MAX_EDITOR_COLUMNS_PER_TABLE + 1)], + ).await.map_err(|error| map_pg_err(error, None))?; + + if column_rows.len() as i64 > MAX_EDITOR_COLUMNS_PER_TABLE { truncated_columns = true; } + let columns = column_rows.into_iter().take(MAX_EDITOR_COLUMNS_PER_TABLE as usize) + .map(|column| QueryEditorColumn { name: column.get(0), data_type: column.get(1) }) + .collect(); + + tables.push(QueryEditorTable { schema, name, columns }); + } + + let function_rows = client.query( + "select n.nspname::text as schema_name, p.proname::text as function_name, \ + coalesce(pg_get_function_identity_arguments(p.oid), '')::text as args, \ + pg_get_function_result(p.oid)::text as return_type \ + from pg_proc p \ + join pg_namespace n on n.oid = p.pronamespace \ + where n.nspname not in ('pg_catalog', 'information_schema') \ + order by n.nspname, p.proname limit $1", + &[&(MAX_EDITOR_FUNCTIONS + 1)], + ).await.map_err(|error| map_pg_err(error, None))?; + + let truncated_functions = function_rows.len() as i64 > MAX_EDITOR_FUNCTIONS; + let functions = function_rows.into_iter().take(MAX_EDITOR_FUNCTIONS as usize) + .map(|row| { + let args_raw: String = row.get(2); + QueryEditorFunction { + schema: row.get(0), + name: row.get(1), + arg_types: if args_raw.trim().is_empty() { Vec::new() } + else { args_raw.split(',').map(|value| value.trim().to_string()).collect() }, + return_type: row.get(3), + } + }).collect(); + + Ok(QueryEditorMetadata { + tables, functions, + truncated_tables, truncated_columns, truncated_functions, + }) + }).await +} diff --git a/src-tauri/src/commands/export_cmds.rs b/src-tauri/src/commands/export_cmds.rs new file mode 100644 index 0000000..471871c --- /dev/null +++ b/src-tauri/src/commands/export_cmds.rs @@ -0,0 +1,65 @@ +use tauri::{AppHandle, State}; + +use crate::db::AppState; +use crate::export::{DiagramExportRequest, ExportQueryRequest, export_diagram_to_png, export_results_csv, export_results_json}; +use crate::credentials; + +#[tauri::command] +pub async fn export_diagram_png( + input: DiagramExportRequest, + output_path: String, +) -> Result<(), String> { + let path = std::path::PathBuf::from(&output_path); + tokio::task::spawn_blocking(move || export_diagram_to_png(&input, &path)) + .await.map_err(|e| e.to_string())? +} + +#[tauri::command] +pub async fn export_results_csv_command( + app: AppHandle, + state: State<'_, AppState>, + input: ExportQueryRequest, +) -> Result<(), String> { + export_results_csv(&app, &state, &input).await +} + +#[tauri::command] +pub async fn export_results_json_command( + app: AppHandle, + state: State<'_, AppState>, + input: ExportQueryRequest, +) -> Result<(), String> { + export_results_json(&app, &state, &input).await +} + +#[tauri::command] +pub async fn save_base64_png(data: String, output_path: String) -> Result<(), String> { + use base64::Engine; + let bytes = base64::engine::general_purpose::STANDARD + .decode(data.strip_prefix("data:image/png;base64,").unwrap_or(&data)) + .map_err(|e| e.to_string())?; + std::fs::write(&output_path, bytes).map_err(|e| e.to_string()) +} + +#[tauri::command] +pub async fn save_text_file(content: String, output_path: String) -> Result<(), String> { + std::fs::write(&output_path, content).map_err(|e| e.to_string()) +} + +#[tauri::command] +pub async fn store_openrouter_api_key(api_key: String) -> Result<(), String> { + if api_key.trim().is_empty() { + return credentials::delete_openrouter_api_key(); + } + credentials::store_openrouter_api_key(&api_key) +} + +#[tauri::command] +pub async fn get_openrouter_api_key() -> Result, String> { + credentials::get_openrouter_api_key() +} + +#[tauri::command] +pub async fn delete_openrouter_api_key() -> Result<(), String> { + credentials::delete_openrouter_api_key() +} diff --git a/src-tauri/src/commands/lint.rs b/src-tauri/src/commands/lint.rs new file mode 100644 index 0000000..a9bc618 --- /dev/null +++ b/src-tauri/src/commands/lint.rs @@ -0,0 +1,80 @@ +use tauri::{AppHandle, State}; + +use crate::db::{ + get_or_create_mysql_pool, get_or_create_sqlite_pool, resolve_connection_engine, + with_pool_client_retry, AppState, +}; +use crate::models::{DatabaseEngine, LintSqlRequest, LintSqlResult, SqlDiagnostic}; +use crate::pg_error::{error_line_column, map_pg_err}; + +use super::MAX_LINT_SQL_BYTES; + +#[tauri::command] +pub async fn lint_sql( + app: AppHandle, + state: State<'_, AppState>, + input: LintSqlRequest, +) -> Result { + let (connection_id, engine) = resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; + let sql = input.sql.trim().to_string(); + if sql.is_empty() { + return Ok(LintSqlResult { diagnostics: Vec::new() }); + } + if sql.len() > MAX_LINT_SQL_BYTES { + return Err("SQL is too large to lint in the editor.".to_string()); + } + + match engine { + DatabaseEngine::Postgres => { + with_pool_client_retry(&app, &state, &connection_id, sql, |client, sql| async move { + let lint_sql = format!("EXPLAIN {}", sql); + let diagnostics = match client.simple_query(&lint_sql).await { + Ok(_) => Vec::new(), + Err(error) => { + let (line, column) = error_line_column(&error, &sql) + .map(|(l, c)| (Some(l), Some(c))) + .unwrap_or((None, None)); + vec![SqlDiagnostic { + message: map_pg_err(error, Some(sql.as_str())), + severity: "error".to_string(), + line, + column, + end_line: line, + end_column: column.map(|value| value + 1), + }] + } + }; + Ok(LintSqlResult { diagnostics }) + }).await + } + DatabaseEngine::Mysql => { + let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; + let lint_sql = format!("EXPLAIN {}", sql); + let diagnostics = match sqlx::query(&lint_sql).execute(&pool).await { + Ok(_) => Vec::new(), + Err(error) => vec![SqlDiagnostic { + message: error.to_string(), + severity: "error".to_string(), + line: None, column: None, end_line: None, end_column: None, + }], + }; + Ok(LintSqlResult { diagnostics }) + } + DatabaseEngine::Sqlite => { + let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; + let lint_sql = format!("EXPLAIN QUERY PLAN {}", sql); + let diagnostics = match sqlx::query(&lint_sql).execute(&pool).await { + Ok(_) => Vec::new(), + Err(error) => vec![SqlDiagnostic { + message: error.to_string(), + severity: "error".to_string(), + line: None, column: None, end_line: None, end_column: None, + }], + }; + Ok(LintSqlResult { diagnostics }) + } + DatabaseEngine::Mongo => { + Err("MongoDB does not support SQL linting.".to_string()) + } + } +} diff --git a/src-tauri/src/commands/mod.rs b/src-tauri/src/commands/mod.rs new file mode 100644 index 0000000..d12aa0f --- /dev/null +++ b/src-tauri/src/commands/mod.rs @@ -0,0 +1,961 @@ +use std::collections::BTreeMap; +use std::sync::Arc; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::time::Instant; + +use futures_util::StreamExt; +use serde_json::Value; +use tauri::{AppHandle, Emitter}; +use sqlx::{Column, Decode, Row, Type}; +use sqlx::mysql::{MySql, MySqlRow}; +use sqlx::sqlite::{Sqlite, SqliteRow}; + +use crate::db::{ + get_or_create_mysql_pool, get_or_create_sqlite_pool, AppState, +}; +use crate::models::{ + AskVeloxyDbContextCache, AskVeloxyTableRef, DatabaseEngine, + ForeignKeyEdge, QueryEditorTable, QueryResult, VeloxyStreamChunk, +}; +use crate::sql_split::split_sql_statements; + +// --- Constants --- + +pub(crate) const MAX_FOREIGN_KEY_ROWS: i64 = 5000; +pub(crate) const MAX_TABLE_INDEX_ROWS: i64 = 500; +pub(crate) const MAX_EDITOR_TABLES: i64 = 150; +pub(crate) const MAX_EDITOR_COLUMNS_PER_TABLE: i64 = 60; +pub(crate) const MAX_EDITOR_FUNCTIONS: i64 = 200; +pub(crate) const MAX_LINT_SQL_BYTES: usize = 65_536; +pub(crate) const ASK_VELOXY_MAX_CONTEXT_TABLES: usize = 8; +pub(crate) const ASK_VELOXY_MAX_CONTEXT_COLUMNS: usize = 18; +pub(crate) const ASK_VELOXY_MAX_CONTEXT_RELATIONSHIPS: usize = 36; +pub(crate) const ASK_VELOXY_SCHEMA_CHAR_BUDGET: usize = 6_000; +pub(crate) const ASK_VELOXY_PROMPT_CHAR_BUDGET: usize = 12_000; +pub(crate) const ASK_VELOXY_MAX_HISTORY_MESSAGES: usize = 30; +pub(crate) const ASK_VELOXY_MAX_CHAT_TOKENS: u32 = 10_000; + +// --- MySQL / SQLite decode helpers --- + +pub(crate) fn mysql_decode_error(context: &str, column_name: &str, index: Option, detail: &str) -> String { + match index { + Some(idx) => format!( + "MySQL decode error in {} at column '{}' (index {}): {}", + context, column_name, idx, detail + ), + None => format!( + "MySQL decode error in {} at column '{}': {}", + context, column_name, detail + ), + } +} + +pub(crate) fn sqlite_decode_error(context: &str, column_name: &str, index: Option, detail: &str) -> String { + match index { + Some(idx) => format!( + "SQLite decode error in {} at column '{}' (index {}): {}", + context, column_name, idx, detail + ), + None => format!( + "SQLite decode error in {} at column '{}': {}", + context, column_name, detail + ), + } +} + +pub(crate) fn mysql_get_idx(row: &MySqlRow, index: usize, column_name: &str, context: &str) -> Result +where + for<'r> T: Decode<'r, MySql> + Type, +{ + row.try_get::(index) + .map_err(|error| mysql_decode_error(context, column_name, Some(index), &error.to_string())) +} + +pub(crate) fn sqlite_get_idx(row: &SqliteRow, index: usize, column_name: &str, context: &str) -> Result +where + for<'r> T: Decode<'r, Sqlite> + Type, +{ + row.try_get::(index) + .map_err(|error| sqlite_decode_error(context, column_name, Some(index), &error.to_string())) +} + +pub(crate) fn sqlite_get_name(row: &SqliteRow, column_name: &str, context: &str) -> Result +where + for<'r> T: Decode<'r, Sqlite> + Type, +{ + row.try_get::(column_name).map_err(|error| { + format!( + "SQLite decode error in {} at column '{}': {}", + context, column_name, error + ) + }) +} + +pub(crate) fn database_name_from_mysql_value(value: Option, context: &str) -> Result { + let name = value + .filter(|value| !value.is_empty()) + .ok_or_else(|| format!("{context} returned an empty database name"))?; + Ok(name) +} + +pub(crate) fn mysql_database_name_from_row(row: &MySqlRow, context: &str) -> Result { + let value = mysql_value_to_string(row, 0, "Database", context)?; + database_name_from_mysql_value(value, context) +} + +pub(crate) fn mysql_value_to_string(row: &MySqlRow, index: usize, column_name: &str, context: &str) -> Result, String> { + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.to_string())); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.to_string())); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.to_string())); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.to_string())); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.to_string())); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.to_string())); + } + if let Ok(value) = row.try_get::>, _>(index) { + return Ok(value.map(|v| v.format("%Y-%m-%d %H:%M:%S").to_string())); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.format("%Y-%m-%d %H:%M:%S").to_string())); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.to_string())); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.to_string())); + } + if let Ok(value) = row.try_get::>, _>(index) { + return Ok(value.map(|v| decode_mysql_bytes_as_string(&v))); + } + Err(mysql_decode_error(context, column_name, Some(index), "unsupported value type")) +} + +pub(crate) fn decode_mysql_bytes_as_string(bytes: &[u8]) -> String { + String::from_utf8_lossy(bytes).into_owned() +} + +pub(crate) fn mysql_value_to_display_string( + row: &MySqlRow, + index: usize, + column_name: &str, + context: &str, +) -> Result, String> { + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.to_string())); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.to_string())); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.to_string())); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.to_string())); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.to_string())); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.to_string())); + } + if let Ok(value) = row.try_get::>, _>(index) { + return Ok(value.map(|v| v.format("%Y-%m-%d %H:%M:%S").to_string())); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.format("%Y-%m-%d %H:%M:%S").to_string())); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.to_string())); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.to_string())); + } + if let Ok(value) = row.try_get::>, _>(index) { + return Ok(value.map(|v| format!("0x{}", hex::encode(v)))); + } + Err(mysql_decode_error(context, column_name, Some(index), "unsupported value type")) +} + +pub(crate) fn mysql_get_string(row: &MySqlRow, index: usize, column_name: &str, context: &str) -> Result { + let value = mysql_value_to_string(row, index, column_name, context)?; + value.ok_or_else(|| mysql_decode_error(context, column_name, Some(index), "unexpected null value")) +} + +pub(crate) fn mysql_get_optional_string( + row: &MySqlRow, + index: usize, + column_name: &str, + context: &str, +) -> Result, String> { + mysql_value_to_string(row, index, column_name, context) +} + +pub(crate) fn sqlite_value_to_string(row: &SqliteRow, index: usize, column_name: &str, context: &str) -> Result, String> { + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.to_string())); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.to_string())); + } + if let Ok(value) = row.try_get::, _>(index) { + return Ok(value.map(|v| v.to_string())); + } + if let Ok(value) = row.try_get::>, _>(index) { + return Ok(value.map(|v| format!("0x{}", hex::encode(v)))); + } + Err(sqlite_decode_error(context, column_name, Some(index), "unsupported value type")) +} + +// --- SQL helpers --- + +pub(crate) fn is_row_returning_sql(sql: &str) -> bool { + let trimmed = sql.trim_start(); + let upper = trimmed.to_uppercase(); + upper.starts_with("SELECT") + || upper.starts_with("WITH") + || upper.starts_with("SHOW") + || upper.starts_with("EXPLAIN") + || upper.starts_with("DESCRIBE") + || upper.starts_with("DESC") + || upper.starts_with("PRAGMA") + || upper.starts_with("VALUES") + || upper.starts_with("TABLE ") +} + +pub(crate) fn classify_sql_intent(sql: &str) -> String { + let normalized = sql.trim_start().to_ascii_lowercase(); + if normalized.starts_with("select") || normalized.starts_with("with") { + return "select".to_string(); + } + if normalized.starts_with("insert") { + return "insert".to_string(); + } + if normalized.starts_with("update") { + return "update".to_string(); + } + if normalized.starts_with("delete") { + return "delete".to_string(); + } + if normalized.starts_with("explain") { + return "explain".to_string(); + } + "unknown".to_string() +} + +pub(crate) fn is_read_only_sql(sql: &str) -> bool { + let mut saw_statement = false; + for statement in sql.split(';').map(str::trim).filter(|s| !s.is_empty()) { + let normalized = statement.to_ascii_lowercase(); + let is_transaction_control = ["begin", "commit", "rollback", "start", "savepoint", "release"] + .iter() + .any(|kw| normalized.starts_with(kw)); + if is_transaction_control { + continue; + } + saw_statement = true; + match classify_sql_intent(statement).as_str() { + "select" | "explain" => {} + _ => return false, + } + } + saw_statement +} + +pub(crate) fn has_multiple_statements(sql: &str) -> bool { + sql.split(';') + .map(str::trim) + .filter(|segment| !segment.is_empty()) + .count() > 1 +} + +pub(crate) fn validate_generated_sql(sql: &str) -> Result<(), String> { + let trimmed = sql.trim(); + if trimmed.is_empty() { + return Err("Ask Veloxy returned an empty SQL statement.".to_string()); + } + if has_multiple_statements(trimmed) { + return Err("Ask Veloxy generated multiple SQL statements. Please ask for a single statement.".to_string()); + } + Ok(()) +} + +// --- Row mapping --- + +pub(crate) fn map_mysql_rows( + rows: Vec, + max_query_rows: usize, +) -> Result<(Vec, Vec>>, usize, bool), String> { + let mut columns: Vec = Vec::new(); + if let Some(first) = rows.first() { + columns = first.columns().iter().map(|column| column.name().to_string()).collect(); + } + let total_rows = rows.len(); + let mut mapped_rows = Vec::new(); + for row in rows.into_iter().take(max_query_rows) { + let mut mapped_row = BTreeMap::new(); + for (index, column_name) in columns.iter().enumerate() { + let value = mysql_value_to_display_string(&row, index, column_name, "run_query")?; + mapped_row.insert(column_name.clone(), value); + } + mapped_rows.push(mapped_row); + } + Ok((columns, mapped_rows, total_rows, total_rows > max_query_rows)) +} + +pub(crate) fn map_sqlite_rows( + rows: Vec, + max_query_rows: usize, +) -> Result<(Vec, Vec>>, usize, bool), String> { + let mut columns: Vec = Vec::new(); + if let Some(first) = rows.first() { + columns = first.columns().iter().map(|column| column.name().to_string()).collect(); + } + let total_rows = rows.len(); + let mut mapped_rows = Vec::new(); + for row in rows.into_iter().take(max_query_rows) { + let mut mapped_row = BTreeMap::new(); + for (index, column_name) in columns.iter().enumerate() { + let value = sqlite_value_to_string(&row, index, column_name, "run_query")?; + mapped_row.insert(column_name.clone(), value); + } + mapped_rows.push(mapped_row); + } + Ok((columns, mapped_rows, total_rows, total_rows > max_query_rows)) +} + +pub(crate) async fn run_query_mysql_or_sqlite( + app: &AppHandle, + state: &AppState, + connection_id: &str, + sql: &str, + max_query_rows: usize, + engine: DatabaseEngine, +) -> Result { + let started_at = Instant::now(); + let statements = split_sql_statements(sql); + if statements.is_empty() { + return Err("Enter a SQL statement before running the query.".to_string()); + } + + let mut columns = Vec::new(); + let mut rows = Vec::new(); + let mut total_rows = 0usize; + let mut truncated = false; + let mut command_tag: Option = None; + + match engine { + DatabaseEngine::Mysql => { + let pool = get_or_create_mysql_pool(app, state, connection_id).await?; + let mut conn = pool.acquire().await.map_err(|error| error.to_string())?; + for statement in statements { + if is_row_returning_sql(&statement) { + let fetched = sqlx::query(&statement) + .fetch_all(&mut *conn) + .await + .map_err(|error| error.to_string())?; + let mapped = map_mysql_rows(fetched, max_query_rows)?; + columns = mapped.0; + rows = mapped.1; + total_rows = mapped.2; + truncated = mapped.3; + command_tag = None; + } else { + let result = sqlx::query(&statement) + .execute(&mut *conn) + .await + .map_err(|error| error.to_string())?; + let affected = result.rows_affected(); + command_tag = Some(affected); + if rows.is_empty() { + total_rows = affected as usize; + } + } + } + } + DatabaseEngine::Sqlite => { + let pool = get_or_create_sqlite_pool(app, state, connection_id).await?; + let mut conn = pool.acquire().await.map_err(|error| error.to_string())?; + for statement in statements { + if is_row_returning_sql(&statement) { + let fetched = sqlx::query(&statement) + .fetch_all(&mut *conn) + .await + .map_err(|error| error.to_string())?; + let mapped = map_sqlite_rows(fetched, max_query_rows)?; + columns = mapped.0; + rows = mapped.1; + total_rows = mapped.2; + truncated = mapped.3; + command_tag = None; + } else { + let result = sqlx::query(&statement) + .execute(&mut *conn) + .await + .map_err(|error| error.to_string())?; + let affected = result.rows_affected(); + command_tag = Some(affected); + if rows.is_empty() { + total_rows = affected as usize; + } + } + } + } + DatabaseEngine::Mongo => { + return Err("Internal engine routing error (MongoDB uses its own query path).".to_string()); + } + DatabaseEngine::Postgres => { + return Err("Internal engine routing error.".to_string()); + } + } + + Ok(QueryResult { + columns, + row_count: if rows.is_empty() { total_rows } else { rows.len() }, + rows, + execution_ms: started_at.elapsed().as_millis(), + truncated, + command_tag, + }) +} + +// --- Veloxy helpers --- + +pub(crate) fn estimate_tokens(chars: usize) -> usize { + (chars / 4).max(1) +} + +pub(crate) fn normalize_openrouter_base(base: Option<&str>) -> String { + let trimmed = base.unwrap_or("https://openrouter.ai/api/v1").trim(); + let value = if trimmed.is_empty() { "https://openrouter.ai/api/v1" } else { trimmed }; + value.trim_end_matches('/').to_string() +} + +pub(crate) fn truncate_on_char_boundary(value: &mut String, max_bytes: usize) { + if value.len() <= max_bytes { + return; + } + let mut truncate_at = max_bytes; + while !value.is_char_boundary(truncate_at) && truncate_at > 0 { + truncate_at -= 1; + } + value.truncate(truncate_at); +} + +pub(crate) fn ask_veloxy_context_cache_key(connection_id: &str, database_name: &str) -> String { + format!("{}::{}", connection_id, database_name) +} + +pub(crate) fn ask_veloxy_conversation_key(connection_id: &str, database_name: &str) -> String { + format!("{}::{}", connection_id, database_name) +} + +pub(crate) fn now_epoch_seconds() -> u64 { + use std::time::{SystemTime, UNIX_EPOCH}; + SystemTime::now().duration_since(UNIX_EPOCH).unwrap_or_default().as_secs() +} + +pub(crate) fn table_matches_target(table: &QueryEditorTable, target: Option<&AskVeloxyTableRef>) -> bool { + let Some(target) = target else { return false }; + table.schema.eq_ignore_ascii_case(&target.schema) && table.name.eq_ignore_ascii_case(&target.name) +} + +pub(crate) fn table_relevance_score(table: &QueryEditorTable, prompt_lower: &str) -> usize { + let mut score = 0usize; + let full_name = format!("{}.{}", table.schema.to_lowercase(), table.name.to_lowercase()); + if prompt_lower.contains(&table.name.to_lowercase()) { score += 3; } + if prompt_lower.contains(&table.schema.to_lowercase()) { score += 2; } + if prompt_lower.contains(&full_name) { score += 4; } + score +} + +pub(crate) fn relationship_relevance_score(edge: &ForeignKeyEdge, prompt_lower: &str) -> usize { + let from_name = format!("{}.{}", edge.from_schema.to_lowercase(), edge.from_table.to_lowercase()); + let to_name = format!("{}.{}", edge.to_schema.to_lowercase(), edge.to_table.to_lowercase()); + let mut score = 0usize; + if prompt_lower.contains(&edge.from_table.to_lowercase()) || prompt_lower.contains(&from_name) { + score += 2; + } + if prompt_lower.contains(&edge.to_table.to_lowercase()) || prompt_lower.contains(&to_name) { + score += 2; + } + score +} + +pub(crate) fn build_schema_context( + db_context: &AskVeloxyDbContextCache, + prompt: &str, + target_table: Option<&AskVeloxyTableRef>, +) -> String { + let prompt_lower = prompt.to_lowercase(); + let mut ranked: Vec<(&QueryEditorTable, usize, bool)> = db_context + .metadata + .tables + .iter() + .map(|table| { + (table, table_relevance_score(table, &prompt_lower), table_matches_target(table, target_table)) + }) + .collect(); + + ranked.sort_by(|a, b| b.2.cmp(&a.2).then_with(|| b.1.cmp(&a.1))); + + let mut schema_context = String::new(); + schema_context.push_str(&format!( + "database {} engine {:?}\n", db_context.database_name, db_context.engine + )); + for (table, _score, _is_target) in ranked.into_iter().take(ASK_VELOXY_MAX_CONTEXT_TABLES) { + let columns = table + .columns + .iter() + .take(ASK_VELOXY_MAX_CONTEXT_COLUMNS) + .map(|column| format!("{}:{}", column.name, column.data_type)) + .collect::>() + .join(", "); + schema_context.push_str(&format!( + "table {}.{} columns [{}]\n", table.schema, table.name, columns + )); + if schema_context.len() >= ASK_VELOXY_SCHEMA_CHAR_BUDGET { + truncate_on_char_boundary(&mut schema_context, ASK_VELOXY_SCHEMA_CHAR_BUDGET); + break; + } + } + + let mut ranked_relationships = db_context + .foreign_keys + .iter() + .map(|edge| (edge, relationship_relevance_score(edge, &prompt_lower))) + .collect::>(); + ranked_relationships.sort_by(|a, b| b.1.cmp(&a.1)); + for (edge, _score) in ranked_relationships + .into_iter() + .take(ASK_VELOXY_MAX_CONTEXT_RELATIONSHIPS) + { + schema_context.push_str(&format!( + "relationship {}.{}({}) -> {}.{}({})\n", + edge.from_schema, edge.from_table, edge.from_column, + edge.to_schema, edge.to_table, edge.to_column + )); + if schema_context.len() >= ASK_VELOXY_SCHEMA_CHAR_BUDGET { + truncate_on_char_boundary(&mut schema_context, ASK_VELOXY_SCHEMA_CHAR_BUDGET); + break; + } + } + schema_context +} + +pub(crate) fn extract_sql_draft_from_text(message: &str) -> Option { + let lowered = message.to_lowercase(); + let markers = ["select ", "with ", "insert ", "update ", "delete ", "explain "]; + let start = markers.iter().filter_map(|marker| lowered.find(marker)).min()?; + let mut sql = message[start..].trim().to_string(); + if let Some(idx) = sql.find("```") { sql.truncate(idx); } + if sql.ends_with('.') { sql.pop(); } + if sql.is_empty() { None } else { Some(sql) } +} + +pub(crate) fn parse_bool_field(value: &Value, field: &str, default: bool) -> bool { + value.get(field).and_then(Value::as_bool).unwrap_or(default) +} + +pub(crate) fn extract_openrouter_message_content(payload: &Value) -> Result { + let content_value = payload + .get("choices") + .and_then(|choices| choices.get(0)) + .and_then(|choice| choice.get("message")) + .and_then(|message| message.get("content")) + .ok_or_else(|| "OpenRouter response missing choices[0].message.content".to_string())?; + + if let Some(content) = content_value.as_str() { + return Ok(content.to_string()); + } + + if let Some(items) = content_value.as_array() { + let mut merged = String::new(); + for item in items { + if let Some(text) = item.get("text").and_then(Value::as_str) { + merged.push_str(text); + } + } + if !merged.trim().is_empty() { + return Ok(merged); + } + } + + Err("OpenRouter returned an unsupported message format.".to_string()) +} + +pub(crate) fn parse_ask_veloxy_json(content: &str) -> Result { + if let Ok(value) = serde_json::from_str::(content) { + return Ok(value); + } + let start = content.find('{'); + let end = content.rfind('}'); + match (start, end) { + (Some(start_idx), Some(end_idx)) if end_idx > start_idx => { + serde_json::from_str::(&content[start_idx..=end_idx]) + .map_err(|error| format!("Ask Veloxy response was not valid JSON: {}", error)) + } + _ => Err("Ask Veloxy response did not contain JSON.".to_string()), + } +} + +pub(crate) fn parse_ask_veloxy_suggestions(generated: &Value) -> Vec { + generated + .get("suggestions") + .and_then(Value::as_array) + .map(|items| { + items.iter() + .filter_map(Value::as_str) + .map(str::trim) + .filter(|item| !item.is_empty()) + .take(5) + .map(|item| { + let mut value = item.to_string(); + truncate_on_char_boundary(&mut value, 200); + value + }) + .collect::>() + }) + .unwrap_or_default() +} + +pub(crate) fn parse_ask_veloxy_chat_json(content: &str) -> Result { + if let Ok(value) = serde_json::from_str::(content) { + return Ok(value); + } + let start = content.find('{'); + let end = content.rfind('}'); + match (start, end) { + (Some(start_idx), Some(end_idx)) if end_idx > start_idx => { + serde_json::from_str::(&content[start_idx..=end_idx]) + .map_err(|error| format!("Ask Veloxy chat JSON was invalid: {}", error)) + } + _ => Err("Ask Veloxy chat response did not contain JSON.".to_string()), + } +} + +fn decode_json_quoted_string(value: &str) -> Option { + serde_json::from_str::(&format!("\"{}\"", value)).ok() +} + +fn unescape_json_string_fragment(raw: &str) -> String { + let mut out = String::with_capacity(raw.len()); + let mut chars = raw.chars().peekable(); + while let Some(ch) = chars.next() { + if ch == '\\' { + match chars.next() { + Some('n') => out.push('\n'), + Some('t') => out.push('\t'), + Some('r') => out.push('\r'), + Some('"') => out.push('"'), + Some('\\') => out.push('\\'), + Some(other) => { out.push('\\'); out.push(other); } + None => out.push('\\'), + } + } else { + out.push(ch); + } + } + out +} + +fn extract_json_string_field(content: &str, key: &str, allow_partial: bool) -> Option { + let marker = format!("\"{}\"", key); + let marker_idx = content.find(&marker)?; + let mut idx = marker_idx + marker.len(); + let bytes = content.as_bytes(); + + while idx < bytes.len() && bytes[idx].is_ascii_whitespace() { idx += 1; } + if idx >= bytes.len() || bytes[idx] != b':' { return None; } + idx += 1; + while idx < bytes.len() && bytes[idx].is_ascii_whitespace() { idx += 1; } + if idx >= bytes.len() || bytes[idx] != b'"' { return None; } + idx += 1; + let start = idx; + let mut escaped = false; + while idx < bytes.len() { + let byte = bytes[idx]; + if escaped { escaped = false; idx += 1; continue; } + if byte == b'\\' { escaped = true; idx += 1; continue; } + if byte == b'"' { + let raw = &content[start..idx]; + return decode_json_quoted_string(raw) + .or_else(|| Some(unescape_json_string_fragment(raw))) + .map(|text| text.trim().to_string()) + .filter(|text| !text.is_empty()); + } + idx += 1; + } + + if allow_partial && start < bytes.len() { + let raw = &content[start..]; + let text = unescape_json_string_fragment(raw).trim().to_string(); + if !text.is_empty() { return Some(text); } + } + None +} + +pub(crate) fn extract_message_from_loose_json(content: &str) -> Option { + let trimmed = content.trim(); + if trimmed.is_empty() { return None; } + let unwrapped = trimmed + .strip_prefix("```json") + .or_else(|| trimmed.strip_prefix("```JSON")) + .map(str::trim_start) + .unwrap_or(trimmed); + let unwrapped = unwrapped.strip_suffix("```").unwrap_or(unwrapped).trim(); + + ["message", "reply", "content"] + .iter() + .find_map(|key| extract_json_string_field(unwrapped, key, false)) + .or_else(|| { + ["message", "reply", "content"] + .iter() + .find_map(|key| extract_json_string_field(unwrapped, key, true)) + }) +} + +pub(crate) fn looks_like_json_response(content: &str) -> bool { + let trimmed = content.trim_start(); + trimmed.starts_with('{') || trimmed.starts_with("```") +} + +pub(crate) fn streaming_display_text(accumulated: &str) -> String { + let trimmed = accumulated.trim(); + if trimmed.is_empty() { return String::new(); } + if let Some(text) = extract_message_from_loose_json(trimmed) { return text; } + if !looks_like_json_response(trimmed) { return trimmed.to_string(); } + String::new() +} + +fn parse_chat_message(value: &Value) -> Option { + if let Some(text) = value.as_str().map(str::trim).filter(|text| !text.is_empty()).map(str::to_string) { + return Some(text); + } + value.get("message").and_then(Value::as_str) + .or_else(|| value.get("reply").and_then(Value::as_str)) + .map(str::trim).filter(|text| !text.is_empty()).map(str::to_string) +} + +pub(crate) type ParsedAskVeloxyChat = ( + String, Vec, Vec, Option, bool, bool, +); + +pub(crate) fn parse_ask_veloxy_chat_content(message_content: &str) -> ParsedAskVeloxyChat { + match parse_ask_veloxy_chat_json(message_content) { + Ok(value) => { + let message = parse_chat_message(&value).unwrap_or_else(|| message_content.trim().to_string()); + let mut draft = value.get("sqlDraft").and_then(Value::as_str) + .or_else(|| value.get("sql_draft").and_then(Value::as_str)) + .map(str::trim).filter(|text| !text.is_empty()).map(str::to_string); + if draft.is_none() { draft = extract_sql_draft_from_text(&message); } + let suggestions = value.get("suggestions").and_then(Value::as_array) + .map(|items| items.iter().filter_map(Value::as_str).map(str::trim) + .filter(|text| !text.is_empty()).take(5).map(str::to_string).collect::>()) + .unwrap_or_default(); + let warnings = value.get("warnings").and_then(Value::as_array) + .map(|items| items.iter().filter_map(Value::as_str).map(str::to_string).collect::>()) + .unwrap_or_default(); + let needs_sql_generation = parse_bool_field(&value, "needsSqlGeneration", draft.is_some()); + let needs_clarification = parse_bool_field(&value, "needsClarification", false); + (message, suggestions, warnings, draft, needs_sql_generation, needs_clarification) + } + Err(_) => { + let normalized_message = extract_message_from_loose_json(message_content).unwrap_or_else(|| { + if looks_like_json_response(message_content) { String::new() } + else { message_content.trim().to_string() } + }); + let mut warnings = vec!["Model returned non-JSON chat output. Parsed in tolerant mode.".to_string()]; + if normalized_message.is_empty() && looks_like_json_response(message_content) { + warnings.push("Response JSON could not be parsed. Try asking again.".to_string()); + } + let draft = extract_sql_draft_from_text(&normalized_message); + let needs_sql_generation = draft.is_some(); + (normalized_message, Vec::new(), warnings, draft, needs_sql_generation, false) + } + } +} + +pub(crate) fn extract_openrouter_stream_delta(data: &str) -> Option { + let payload: Value = serde_json::from_str(data).ok()?; + payload.get("choices").and_then(|choices| choices.get(0)) + .and_then(|choice| choice.get("delta")) + .and_then(|delta| delta.get("content")) + .and_then(Value::as_str).filter(|text| !text.is_empty()).map(str::to_string) +} + +pub(crate) fn extract_openrouter_finish_reason(data: &str) -> Option { + let payload: Value = serde_json::from_str(data).ok()?; + payload.get("choices").and_then(|choices| choices.get(0)) + .and_then(|choice| choice.get("finish_reason")) + .and_then(Value::as_str).map(str::to_string) +} + +pub(crate) fn emit_veloxy_stream_chunk(app: &AppHandle, chunk: VeloxyStreamChunk) { + let _ = app.emit("veloxy-stream-chunk", chunk); +} + +pub(crate) async fn stream_openrouter_chat_completion( + app: &AppHandle, + client: &reqwest::Client, + endpoint: &str, + api_key: &str, + model: &str, + system_prompt: &str, + user_prompt: &str, + request_id: &str, + cancel: Arc, +) -> Result<(String, bool), String> { + let response = client + .post(endpoint) + .header("Authorization", format!("Bearer {}", api_key)) + .header("Content-Type", "application/json") + .json(&serde_json::json!({ + "model": model, + "temperature": 0.2, + "max_tokens": ASK_VELOXY_MAX_CHAT_TOKENS, + "stream": true, + "messages": [ + { "role": "system", "content": system_prompt }, + { "role": "user", "content": user_prompt } + ] + })) + .send() + .await + .map_err(|error| format!("OpenRouter request failed: {}", error))?; + + let status = response.status(); + if !status.is_success() { + let body = response.text().await.unwrap_or_else(|_| "Unknown OpenRouter error".to_string()); + if let Ok(payload) = serde_json::from_str::(&body) { + let message = payload.get("error").and_then(|error| error.get("message")) + .and_then(Value::as_str).unwrap_or("Unknown OpenRouter error"); + return Err(format!("OpenRouter error ({}): {}", status.as_u16(), message)); + } + return Err(format!("OpenRouter error ({}): {}", status.as_u16(), body)); + } + + let mut stream = response.bytes_stream(); + let mut buffer = String::new(); + let mut accumulated = String::new(); + let mut last_display_len = 0usize; + let mut hit_token_limit = false; + + while let Some(chunk) = stream.next().await { + if cancel.load(Ordering::Relaxed) { + return Ok((accumulated, hit_token_limit)); + } + let bytes = chunk.map_err(|error| format!("OpenRouter stream read failed: {}", error))?; + buffer.push_str(&String::from_utf8_lossy(&bytes)); + + while let Some(line_end) = buffer.find('\n') { + let line = buffer[..line_end].trim_end_matches('\r').to_string(); + buffer.drain(..=line_end); + + if !line.starts_with("data: ") { continue; } + let data = line["data: ".len()..].trim(); + if data == "[DONE]" { continue; } + if extract_openrouter_finish_reason(data).as_deref() == Some("length") { + hit_token_limit = true; + } + if let Some(delta) = extract_openrouter_stream_delta(data) { + accumulated.push_str(&delta); + let display = streaming_display_text(&accumulated); + let display_delta = if display.len() > last_display_len { + display[last_display_len..].to_string() + } else { + String::new() + }; + last_display_len = display.len(); + if !display_delta.is_empty() { + emit_veloxy_stream_chunk(app, VeloxyStreamChunk { + request_id: request_id.to_string(), + delta: display_delta, + done: false, + message: None, + suggestions: Vec::new(), + warnings: Vec::new(), + sql_draft: None, + needs_sql_generation: false, + needs_clarification: false, + }); + } + } + } + } + + Ok((accumulated, hit_token_limit)) +} + +pub(crate) fn veloxdb_unique_constraint_name(table_name: &str, column_name: &str) -> String { + let suffix = "_uniq"; + let max_base_len = 63usize.saturating_sub(suffix.len()); + let mut base = format!("veloxdb_{}_{}", table_name, column_name); + base.truncate(max_base_len); + format!("{}{}", base, suffix) +} + +// --- Sub-modules --- + +mod connections; +mod query; +mod table_props; +mod ddl; +mod editor_meta; +mod veloxy; +mod lint; +mod mongo; +mod export_cmds; + +// --- Re-exports --- + +pub use connections::{ + connect_db, list_connections_command, set_active_connection, ping_connection, + refresh_connection, disconnect_db, rename_connection, delete_connection, + list_databases, switch_database, +}; +pub use query::{ + run_query, get_tables, get_schema, +}; +pub use table_props::{ + get_table_properties, apply_table_properties, get_foreign_keys, get_table_indexes, +}; +pub use ddl::{ + execute_ddl_transaction, execute_ddl_statement, +}; +pub use editor_meta::{ + get_query_editor_metadata, +}; +pub use veloxy::{ + cancel_veloxy_request, chat_with_db, clear_veloxy_conversation, generate_sql_from_nl, + load_veloxy_conversation, +}; +pub use lint::{ + lint_sql, +}; +pub use export_cmds::{ + export_diagram_png, export_results_csv_command, export_results_json_command, + save_base64_png, save_text_file, store_openrouter_api_key, get_openrouter_api_key, + delete_openrouter_api_key, +}; +pub use mongo::{ + mongo_run_query, mongo_get_collections, mongo_get_schema, +}; diff --git a/src-tauri/src/commands/mongo.rs b/src-tauri/src/commands/mongo.rs new file mode 100644 index 0000000..ea72633 --- /dev/null +++ b/src-tauri/src/commands/mongo.rs @@ -0,0 +1,321 @@ +use std::collections::{BTreeMap, HashMap, HashSet}; +use std::time::Instant; + +use futures_util::StreamExt; +use mongodb::bson::{doc, Document}; +use tauri::{AppHandle, State}; + +use crate::db::{get_or_create_mongo_client, resolve_connection_engine, AppState, MAX_QUERY_ROWS}; +use crate::models::{ColumnInfo, QueryRequest, QueryResult, TableInfo}; + +/// Parse a user-supplied MongoDB query string into a filter Document. +/// +/// Supports these input forms: +/// `{"status": "active"}` — raw JSON filter +/// `db.collection.find({...})` — full shell syntax (ignores collection prefix) +/// `{ status: "active" }` — relaxed JSON (unquoted keys, single quotes) +fn parse_mongo_filter(raw: &str) -> Result { + let trimmed = raw.trim(); + if trimmed.is_empty() { + return Ok(doc! {}); + } + + // Try strict JSON first + if let Ok(doc) = serde_json::from_str::(trimmed) { + if !doc.is_empty() { + return Ok(doc); + } + } + + // Try shell syntax: db.collection.find({...}) + if let Some(start) = trimmed.find(".find(") { + let after_find = &trimmed[start + 6..]; // skip ".find(" + let inner = extract_braced_json(after_find) + .ok_or_else(|| "Could not parse .find() arguments.".to_string())?; + return normalize_to_document(&inner); + } + + // Try relaxed JSON (single quotes, unquoted keys) + let relaxed = trimmed + .replace('\'', "\"") + .replace("ObjectId(", "\"ObjectId(") + .replace("ISODate(", "\"ISODate("); + if let Ok(doc) = serde_json::from_str::(&relaxed) { + return Ok(doc); + } + + // Last attempt: treat as a raw key-value search + normalize_to_document(trimmed) +} + +fn extract_braced_json(input: &str) -> Option { + let trimmed = input.trim(); + if !trimmed.starts_with('{') { + return None; + } + let mut depth = 0i32; + let mut end = 0usize; + for (i, ch) in trimmed.char_indices() { + match ch { + '{' => depth += 1, + '}' => { + depth -= 1; + if depth == 0 { + end = i + 1; + break; + } + } + _ => {} + } + } + if end > 0 { + Some(trimmed[..end].to_string()) + } else { + None + } +} + +fn normalize_to_document(raw: &str) -> Result { + // Try parsing as BSON + if let Ok(doc) = serde_json::from_str::(raw) { + return Ok(doc); + } + // Return empty filter — matches all documents + Ok(doc! {}) +} + +fn bson_to_display_string(value: &mongodb::bson::Bson) -> Option { + use mongodb::bson::Bson; + match value { + Bson::String(s) => Some(s.clone()), + Bson::Int32(n) => Some(n.to_string()), + Bson::Int64(n) => Some(n.to_string()), + Bson::Double(f) => Some(f.to_string()), + Bson::Boolean(b) => Some(b.to_string()), + Bson::ObjectId(oid) => Some(oid.to_hex()), + Bson::DateTime(dt) => Some( + chrono::DateTime::from_timestamp_millis(dt.timestamp_millis()) + .map(|d| d.format("%Y-%m-%d %H:%M:%S").to_string()) + .unwrap_or_else(|| dt.to_string()), + ), + Bson::Null => None, + Bson::Array(arr) => Some(format!("[{} items]", arr.len())), + Bson::Document(_) => Some("{...}".to_string()), + Bson::Binary(_) => Some("".to_string()), + Bson::RegularExpression(re) => Some(format!("/{}/", re.pattern)), + Bson::Decimal128(d) => Some(d.to_string()), + _ => Some(format!("{:?}", value)), + } +} + +fn infer_bson_type(value: &mongodb::bson::Bson) -> &'static str { + use mongodb::bson::Bson; + match value { + Bson::String(_) => "string", + Bson::Int32(_) | Bson::Int64(_) => "integer", + Bson::Double(_) => "double", + Bson::Boolean(_) => "boolean", + Bson::ObjectId(_) => "objectId", + Bson::DateTime(_) => "date", + Bson::Array(_) => "array", + Bson::Document(_) => "object", + Bson::Null => "null", + Bson::Binary(_) => "binary", + Bson::RegularExpression(_) => "regex", + Bson::Decimal128(_) => "decimal", + _ => "unknown", + } +} + +/// Execute a MongoDB find/aggregate query and return tabular results. +#[tauri::command] +pub async fn mongo_run_query( + app: AppHandle, + state: State<'_, AppState>, + input: QueryRequest, +) -> Result { + let (connection_id, _engine) = + resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; + let client = get_or_create_mongo_client(&app, &state, &connection_id).await?; + + let started_at = Instant::now(); + let raw = input.sql.trim().to_string(); + if raw.is_empty() { + return Err("Enter a MongoDB query or filter.".to_string()); + } + + // Determine the collection name — try extracting from shell syntax + let (collection_name, filter) = if let Some(dot_find) = raw.find(".find(") { + let collection = raw[..dot_find].trim().to_string(); + let filter = parse_mongo_filter(&raw)?; + (collection, filter) + } else if raw.starts_with('{') { + // No collection specified — use the stored database's first collection as context, + // or just return all databases if none specified. + let db = client.default_database().ok_or_else(|| { + "No default database. Specify a database or use db.collection.find({...}) syntax." + .to_string() + })?; + let names = db + .list_collection_names() + .await + .map_err(|e| format!("Failed to list collections: {}", e))?; + if names.is_empty() { + return Err("No collections found in the current database.".to_string()); + } + let filter = parse_mongo_filter(&raw)?; + (names[0].clone(), filter) + } else { + return Err( + "Use MongoDB shell syntax: db.collection.find({...}) or provide a JSON filter." + .to_string(), + ); + }; + + // Try to determine the database name from the collection + let db = client.default_database().ok_or_else(|| { + "No default database available. Try switching to a database first.".to_string() + })?; + + let collection = db.collection::(&collection_name); + let max_rows = input.max_rows.unwrap_or(MAX_QUERY_ROWS) as i64; + + let mut cursor = collection + .find(filter) + .limit(max_rows) + .await + .map_err(|e| format!("MongoDB query failed: {}", e))?; + + let mut rows: Vec>> = Vec::new(); + let mut columns: Vec = Vec::new(); + let mut columns_set = HashSet::new(); + let mut total = 0usize; + + while let Some(result) = cursor.next().await { + let doc = result.map_err(|e| format!("MongoDB cursor error: {}", e))?; + if total == 0 { + columns = doc.keys().cloned().collect(); + for col in &columns { + columns_set.insert(col.clone()); + } + } else { + // Discover any new fields in subsequent documents + for key in doc.keys() { + if !columns_set.contains(key) { + columns.push(key.clone()); + columns_set.insert(key.clone()); + } + } + } + let mut row = BTreeMap::new(); + for key in &columns { + let value = doc.get(key).and_then(bson_to_display_string); + row.insert(key.clone(), value); + } + rows.push(row); + total += 1; + if total >= max_rows as usize { + break; + } + } + + Ok(QueryResult { + columns, + rows, + row_count: total, + execution_ms: started_at.elapsed().as_millis(), + truncated: false, + command_tag: None, + }) +} + +/// List all collections in the active database. +#[tauri::command] +pub async fn mongo_get_collections( + app: AppHandle, + state: State<'_, AppState>, + connection_id: Option, +) -> Result, String> { + let (connection_id, _engine) = + resolve_connection_engine(&app, &state, connection_id.clone()).await?; + let client = get_or_create_mongo_client(&app, &state, &connection_id).await?; + + let db = client.default_database().ok_or_else(|| { + "No default database available.".to_string() + })?; + + let db_name = db.name().to_string(); + let names = db + .list_collection_names() + .await + .map_err(|e| format!("Failed to list collections: {}", e))?; + + Ok(names + .into_iter() + .map(|name| TableInfo { + schema: db_name.clone(), + name: name.clone(), + preview_query: format!("db.{}.find({{}}).limit(100)", name), + }) + .collect()) +} + +/// Infer the schema of a MongoDB collection by sampling documents. +#[tauri::command] +pub async fn mongo_get_schema( + app: AppHandle, + state: State<'_, AppState>, + connection_id: Option, + database: String, + collection: String, +) -> Result, String> { + let (connection_id, _engine) = + resolve_connection_engine(&app, &state, connection_id.clone()).await?; + let client = get_or_create_mongo_client(&app, &state, &connection_id).await?; + + let db = client.database(&database); + let coll = db.collection::(&collection); + + let mut cursor = coll + .find(doc! {}) + .limit(200) + .await + .map_err(|e| format!("MongoDB query failed: {}", e))?; + + // Track field → BSON type. First-seen type wins. + let mut type_map: HashMap = HashMap::new(); + // Track field → nullable. If any document lacks the field, it's nullable. + let mut doc_count = 0usize; + let mut field_doc_counts: HashMap = HashMap::new(); + + while let Some(result) = cursor.next().await { + let doc = result.map_err(|e| format!("MongoDB cursor error: {}", e))?; + doc_count += 1; + + for (key, value) in doc.iter() { + *field_doc_counts.entry(key.clone()).or_default() += 1; + type_map + .entry(key.clone()) + .or_insert_with(|| infer_bson_type(value).to_string()); + } + } + + if doc_count == 0 { + return Ok(Vec::new()); + } + + Ok(type_map + .into_iter() + .map(|(name, data_type)| { + let field_count = field_doc_counts.get(&name).copied().unwrap_or(0); + let is_nullable = field_count < doc_count; + ColumnInfo { + table_schema: database.clone(), + table_name: collection.clone(), + column_name: name, + data_type, + is_nullable, + } + }) + .collect()) +} diff --git a/src-tauri/src/commands/query.rs b/src-tauri/src/commands/query.rs new file mode 100644 index 0000000..6c97fd2 --- /dev/null +++ b/src-tauri/src/commands/query.rs @@ -0,0 +1,255 @@ +use std::time::Instant; +use tauri::{AppHandle, State}; +use tokio_postgres::SimpleQueryMessage; + +use crate::db::{ + get_or_create_mysql_pool, get_or_create_sqlite_pool, load_connection, quote_identifier, + require_safe_identifier, resolve_connection_engine, with_pool_client_retry, AppState, MAX_QUERY_ROWS, +}; +use crate::models::{ + ColumnInfo, DatabaseEngine, QueryRequest, QueryResult, SchemaRequest, TableInfo, +}; +use crate::pg_error::map_pg_err; + +use super::{is_read_only_sql, run_query_mysql_or_sqlite, mysql_get_string, sqlite_get_idx, sqlite_get_name}; +use super::mongo::{mongo_run_query, mongo_get_collections, mongo_get_schema}; + +#[tauri::command] +pub async fn run_query( + app: AppHandle, + state: State<'_, AppState>, + input: QueryRequest, +) -> Result { + let (connection_id, engine) = resolve_connection_engine(&app, &state, input.connection_id).await?; + let sql = input.sql.trim().to_string(); + + if sql.is_empty() { + return Err("Enter a SQL statement before running the query.".to_string()); + } + + if !input.allow_write.unwrap_or(false) && !is_read_only_sql(&sql) { + return Err( + "This statement modifies data or schema. Confirm execution in the editor, \ + or use the model/DDL workflow for schema changes.".to_string(), + ); + } + + let max_query_rows = input.max_rows.unwrap_or(MAX_QUERY_ROWS); + + match engine { + DatabaseEngine::Postgres => { + with_pool_client_retry(&app, &state, &connection_id, sql, |client, sql| async move { + let started_at = Instant::now(); + let messages = client.simple_query(&sql).await + .map_err(|error| map_pg_err(error, Some(sql.as_str())))?; + + let mut columns = Vec::new(); + let mut rows = Vec::new(); + let mut total_rows = 0usize; + let mut command_tag = None; + + for message in messages { + match message { + SimpleQueryMessage::RowDescription(description) => { + if columns.is_empty() { + columns = description.iter() + .map(|column| column.name().to_string()).collect(); + } + } + SimpleQueryMessage::Row(row) => { + total_rows += 1; + if columns.is_empty() { + columns = row.columns().iter() + .map(|column| column.name().to_string()).collect(); + } + if rows.len() >= max_query_rows { continue; } + let mut mapped_row = std::collections::BTreeMap::new(); + for (index, column_name) in columns.iter().enumerate() { + mapped_row.insert(column_name.clone(), row.get(index).map(str::to_owned)); + } + rows.push(mapped_row); + } + SimpleQueryMessage::CommandComplete(count) => { + command_tag = Some(count); + } + _ => {} + } + } + + Ok(QueryResult { + columns, + row_count: rows.len(), + rows, + execution_ms: started_at.elapsed().as_millis(), + truncated: total_rows > max_query_rows, + command_tag, + }) + }).await + } + DatabaseEngine::Mysql | DatabaseEngine::Sqlite => { + run_query_mysql_or_sqlite(&app, &state, &connection_id, &sql, max_query_rows, engine).await + } + DatabaseEngine::Mongo => { + mongo_run_query(app, state, QueryRequest { + connection_id: Some(connection_id), + sql: sql.clone(), + max_rows: input.max_rows, + allow_write: input.allow_write, + }).await + } + } +} + +#[tauri::command] +pub async fn get_tables( + app: AppHandle, + state: State<'_, AppState>, + connection_id: Option, +) -> Result, String> { + let (connection_id, engine) = resolve_connection_engine(&app, &state, connection_id).await?; + + match engine { + DatabaseEngine::Postgres => { + with_pool_client_retry(&app, &state, &connection_id, (), |client, ()| async move { + let rows = client.query( + "select table_schema, table_name \ + from information_schema.tables \ + where table_type = 'BASE TABLE' \ + and table_schema not in ('pg_catalog', 'information_schema') \ + order by table_schema, table_name", + &[], + ).await.map_err(|error| map_pg_err(error, None))?; + + Ok(rows.into_iter().map(|row| { + let schema: String = row.get(0); + let name: String = row.get(1); + let preview_query = format!( + "select * from \"{}\".\"{}\" limit 100;", + quote_identifier(&schema), quote_identifier(&name) + ); + TableInfo { schema, name, preview_query } + }).collect()) + }).await + } + DatabaseEngine::Mysql => { + let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; + let database = load_connection(&app, &connection_id)? + .map(|connection| connection.database).unwrap_or_default(); + let rows = sqlx::query( + "select table_schema, table_name \ + from information_schema.tables \ + where table_type = 'BASE TABLE' \ + and table_schema = ? \ + and table_schema not in ('information_schema', 'mysql', 'performance_schema', 'sys') \ + order by table_schema, table_name", + ).bind(&database).fetch_all(&pool).await.map_err(|error| error.to_string())?; + let mut tables = Vec::new(); + for row in rows { + let schema: String = mysql_get_string(&row, 0, "table_schema", "get_tables")?; + let name: String = mysql_get_string(&row, 1, "table_name", "get_tables")?; + tables.push(TableInfo { + preview_query: format!("select * from `{}`.`{}` limit 100;", schema, name), + schema, name, + }); + } + Ok(tables) + } + DatabaseEngine::Sqlite => { + let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; + let rows = sqlx::query( + "select name from sqlite_master \ + where type = 'table' and name not like 'sqlite_%' order by name", + ).fetch_all(&pool).await.map_err(|error| error.to_string())?; + let mut tables = Vec::new(); + for row in rows { + let name: String = sqlite_get_idx(&row, 0, "name", "get_tables")?; + require_safe_identifier(&name, "table name")?; + tables.push(TableInfo { + schema: "main".to_string(), + preview_query: format!("select * from \"{}\" limit 100;", quote_identifier(&name)), + name, + }); + } + Ok(tables) + } + DatabaseEngine::Mongo => { + mongo_get_collections(app, state, Some(connection_id)).await + } + } +} + +#[tauri::command] +pub async fn get_schema( + app: AppHandle, + state: State<'_, AppState>, + input: SchemaRequest, +) -> Result, String> { + let (connection_id, engine) = resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; + let schema_request = input.clone(); + + match engine { + DatabaseEngine::Postgres => { + with_pool_client_retry(&app, &state, &connection_id, schema_request, |client, input| async move { + let rows = client.query( + "select table_schema, table_name, column_name, data_type, is_nullable \ + from information_schema.columns \ + where table_schema = $1 and table_name = $2 \ + order by ordinal_position", + &[&input.table_schema, &input.table_name], + ).await.map_err(|error| map_pg_err(error, None))?; + + Ok(rows.into_iter().map(|row| ColumnInfo { + table_schema: row.get(0), + table_name: row.get(1), + column_name: row.get(2), + data_type: row.get(3), + is_nullable: row.get::<_, String>(4) == "YES", + }).collect()) + }).await + } + DatabaseEngine::Mysql => { + let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; + let rows = sqlx::query( + "select table_schema, table_name, column_name, data_type, is_nullable \ + from information_schema.columns \ + where table_schema = ? and table_name = ? \ + order by ordinal_position", + ).bind(&schema_request.table_schema).bind(&schema_request.table_name) + .fetch_all(&pool).await.map_err(|error| error.to_string())?; + let mut columns = Vec::new(); + for row in rows { + columns.push(ColumnInfo { + table_schema: mysql_get_string(&row, 0, "table_schema", "get_schema")?, + table_name: mysql_get_string(&row, 1, "table_name", "get_schema")?, + column_name: mysql_get_string(&row, 2, "column_name", "get_schema")?, + data_type: mysql_get_string(&row, 3, "data_type", "get_schema")?, + is_nullable: mysql_get_string(&row, 4, "is_nullable", "get_schema")? == "YES", + }); + } + Ok(columns) + } + DatabaseEngine::Sqlite => { + let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; + require_safe_identifier(&schema_request.table_name, "table name")?; + let pragma_sql = format!("PRAGMA table_info(\"{}\");", quote_identifier(&schema_request.table_name)); + let rows = sqlx::query(&pragma_sql).fetch_all(&pool).await.map_err(|error| error.to_string())?; + let mut columns = Vec::new(); + for row in rows { + let col_name: String = sqlite_get_name(&row, "name", "get_schema")?; + let col_type: String = sqlite_get_name(&row, "type", "get_schema")?; + let notnull: i64 = sqlite_get_name(&row, "notnull", "get_schema")?; + columns.push(ColumnInfo { + table_schema: "main".to_string(), + table_name: schema_request.table_name.clone(), + column_name: col_name, + data_type: col_type, + is_nullable: notnull == 0, + }); + } + Ok(columns) + } + DatabaseEngine::Mongo => { + mongo_get_schema(app, state, input.connection_id, input.table_schema, input.table_name).await + } + } +} diff --git a/src-tauri/src/commands/table_props.rs b/src-tauri/src/commands/table_props.rs new file mode 100644 index 0000000..e3cda25 --- /dev/null +++ b/src-tauri/src/commands/table_props.rs @@ -0,0 +1,585 @@ +use std::collections::{HashMap, HashSet}; +use tauri::{AppHandle, State}; + +use crate::db::{ + get_or_create_mysql_pool, get_or_create_sqlite_pool, quote_identifier, + require_safe_identifier, resolve_connection_engine, with_pool_client_retry, AppState, +}; +use crate::models::{ + ColumnProperties, DatabaseEngine, ForeignKeyEdge, + IndexInfo, SchemaRequest, TableIndexesResult, TablePropertiesApplyRequest, +}; +use crate::pg_error::map_pg_err; + +use super::{ + MAX_TABLE_INDEX_ROWS, MAX_FOREIGN_KEY_ROWS, + mysql_get_string, mysql_get_optional_string, mysql_get_idx, + sqlite_get_name, sqlite_get_idx, + veloxdb_unique_constraint_name, +}; + +#[tauri::command] +pub async fn get_table_properties( + app: AppHandle, + state: State<'_, AppState>, + input: SchemaRequest, +) -> Result, String> { + let (connection_id, engine) = resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; + let ctx = input.clone(); + + if engine == DatabaseEngine::Mysql { + let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; + let rows = sqlx::query( + "select c.table_schema, c.table_name, c.column_name, c.data_type, \ + c.is_nullable, c.column_default, c.extra \ + from information_schema.columns c \ + where c.table_schema = ? and c.table_name = ? \ + order by c.ordinal_position", + ).bind(&ctx.table_schema).bind(&ctx.table_name).fetch_all(&pool).await + .map_err(|error| error.to_string())?; + + let pk_rows = sqlx::query( + "select column_name from information_schema.key_column_usage \ + where table_schema = ? and table_name = ? and constraint_name = 'PRIMARY'", + ).bind(&ctx.table_schema).bind(&ctx.table_name).fetch_all(&pool).await + .map_err(|error| error.to_string())?; + let pk_cols: HashSet = pk_rows.into_iter() + .map(|row| mysql_get_string(&row, 0, "column_name", "get_table_properties")) + .collect::, _>>()?; + + let unique_rows = sqlx::query( + "select index_name, column_name, seq_in_index \ + from information_schema.statistics \ + where table_schema = ? and table_name = ? and non_unique = 0 \ + order by index_name, seq_in_index", + ).bind(&ctx.table_schema).bind(&ctx.table_name).fetch_all(&pool).await + .map_err(|error| error.to_string())?; + let mut unique_by_index: HashMap> = HashMap::new(); + for row in unique_rows { + let index_name: String = mysql_get_string(&row, 0, "index_name", "get_table_properties")?; + if index_name == "PRIMARY" { continue; } + let column_name: String = mysql_get_string(&row, 1, "column_name", "get_table_properties")?; + unique_by_index.entry(index_name).or_default().push(column_name); + } + let mut unique_cols: HashSet = HashSet::new(); + let mut composite_unique_cols: HashSet = HashSet::new(); + for cols in unique_by_index.values() { + for col in cols { unique_cols.insert(col.clone()); } + if cols.len() > 1 { for col in cols { composite_unique_cols.insert(col.clone()); } } + } + + let mut properties = Vec::new(); + for row in rows { + let column_name: String = mysql_get_string(&row, 2, "column_name", "get_table_properties")?; + let is_primary_key = pk_cols.contains(&column_name); + let is_unique = is_primary_key || unique_cols.contains(&column_name); + let is_part_of_composite_unique = composite_unique_cols.contains(&column_name); + let extra: String = mysql_get_string(&row, 6, "extra", "get_table_properties")?; + let lower_extra = extra.to_lowercase(); + properties.push(ColumnProperties { + table_schema: mysql_get_string(&row, 0, "table_schema", "get_table_properties")?, + table_name: mysql_get_string(&row, 1, "table_name", "get_table_properties")?, + column_name, + data_type: mysql_get_string(&row, 3, "data_type", "get_table_properties")?, + is_nullable: mysql_get_string(&row, 4, "is_nullable", "get_table_properties")? == "YES", + is_primary_key, is_unique, is_part_of_composite_unique, + column_default: mysql_get_optional_string(&row, 5, "column_default", "get_table_properties")?, + is_identity: lower_extra.contains("auto_increment"), + identity_generation: if lower_extra.contains("auto_increment") { Some("BY DEFAULT".to_string()) } else { None }, + is_generated: if lower_extra.contains("generated") { Some("ALWAYS".to_string()) } else { None }, + }); + } + return Ok(properties); + } + + if engine == DatabaseEngine::Sqlite { + let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; + require_safe_identifier(&ctx.table_name, "table name")?; + let pragma_sql = format!("PRAGMA table_info(\"{}\");", quote_identifier(&ctx.table_name)); + let rows = sqlx::query(&pragma_sql).fetch_all(&pool).await.map_err(|error| error.to_string())?; + let index_list_sql = format!("PRAGMA index_list(\"{}\");", quote_identifier(&ctx.table_name)); + let index_rows = sqlx::query(&index_list_sql).fetch_all(&pool).await.map_err(|error| error.to_string())?; + let mut unique_cols: HashSet = HashSet::new(); + let mut composite_unique_cols: HashSet = HashSet::new(); + for index in index_rows { + let is_unique = sqlite_get_name::(&index, "unique", "get_table_properties")? == 1; + if !is_unique { continue; } + let origin = sqlite_get_name::(&index, "origin", "get_table_properties")?; + if origin == "pk" { continue; } + let index_name = sqlite_get_name::(&index, "name", "get_table_properties")?; + require_safe_identifier(&index_name, "index name")?; + let info_sql = format!("PRAGMA index_info(\"{}\");", quote_identifier(&index_name)); + let info_rows = sqlx::query(&info_sql).fetch_all(&pool).await.map_err(|error| error.to_string())?; + let mut cols: Vec = Vec::new(); + for info in info_rows { + if let Ok(name) = sqlite_get_name::(&info, "name", "get_table_properties") { cols.push(name); } + } + for col in &cols { unique_cols.insert(col.clone()); } + if cols.len() > 1 { for col in cols { composite_unique_cols.insert(col); } } + } + let mut properties = Vec::new(); + for row in rows { + let column_name: String = sqlite_get_name(&row, "name", "get_table_properties")?; + let is_primary_key = sqlite_get_name::(&row, "pk", "get_table_properties")? == 1; + let is_unique = is_primary_key || unique_cols.contains(&column_name); + let is_part_of_composite_unique = composite_unique_cols.contains(&column_name); + properties.push(ColumnProperties { + table_schema: "main".to_string(), + table_name: ctx.table_name.clone(), + column_name, + data_type: sqlite_get_name(&row, "type", "get_table_properties")?, + is_nullable: sqlite_get_name::(&row, "notnull", "get_table_properties")? == 0, + is_primary_key, is_unique, is_part_of_composite_unique, + column_default: sqlite_get_name::>(&row, "dflt_value", "get_table_properties")?, + is_identity: false, + identity_generation: None, + is_generated: None, + }); + } + return Ok(properties); + } + + with_pool_client_retry(&app, &state, &connection_id, ctx, |client, input| async move { + let columns = client.query( + "select c.table_schema, c.table_name, c.column_name, c.data_type, \ + c.is_nullable, c.column_default, c.is_identity, c.identity_generation, c.is_generated \ + from information_schema.columns c \ + where c.table_schema = $1 and c.table_name = $2 \ + order by c.ordinal_position", + &[&input.table_schema, &input.table_name], + ).await.map_err(|error| map_pg_err(error, None))?; + + let primary_keys = client.query( + "select kcu.column_name \ + from information_schema.table_constraints tc \ + join information_schema.key_column_usage kcu \ + on tc.constraint_name = kcu.constraint_name \ + and tc.table_schema = kcu.table_schema \ + where tc.table_schema = $1 and tc.table_name = $2 \ + and tc.constraint_type = 'PRIMARY KEY' \ + order by kcu.ordinal_position", + &[&input.table_schema, &input.table_name], + ).await.map_err(|error| map_pg_err(error, None))?; + + let primary_key_columns: HashSet = primary_keys.into_iter() + .filter_map(|row| Some(row.get::<_, String>(0))).collect(); + + let unique_constraints = client.query( + "select tc.constraint_name, kcu.column_name, kcu.ordinal_position \ + from information_schema.table_constraints tc \ + join information_schema.key_column_usage kcu \ + on tc.constraint_name = kcu.constraint_name \ + and tc.table_schema = kcu.table_schema \ + where tc.table_schema = $1 and tc.table_name = $2 \ + and tc.constraint_type = 'UNIQUE' \ + order by tc.constraint_name, kcu.ordinal_position", + &[&input.table_schema, &input.table_name], + ).await.map_err(|error| map_pg_err(error, None))?; + + let mut unique_by_name: HashMap> = HashMap::new(); + for row in unique_constraints { + let constraint_name: String = row.get(0); + let column_name: String = row.get(1); + unique_by_name.entry(constraint_name).or_default().push(column_name); + } + + let mut unique_columns: HashSet = HashSet::new(); + let mut composite_unique_columns: HashSet = HashSet::new(); + for (_constraint_name, cols) in unique_by_name { + for c in &cols { unique_columns.insert(c.clone()); } + if cols.len() > 1 { for c in &cols { composite_unique_columns.insert(c.clone()); } } + } + + Ok(columns.into_iter().map(|row| { + let column_name: String = row.get(2); + let is_primary_key = primary_key_columns.contains(&column_name); + let is_unique = is_primary_key || unique_columns.contains(&column_name); + let is_part_of_composite_unique = composite_unique_columns.contains(&column_name); + ColumnProperties { + table_schema: row.get(0), + table_name: row.get(1), + column_name, + data_type: row.get(3), + is_nullable: row.get::<_, String>(4) == "YES", + is_primary_key, is_unique, is_part_of_composite_unique, + column_default: row.get(5), + is_identity: row.get::<_, Option>(6).as_deref() == Some("YES"), + identity_generation: row.get(7), + is_generated: row.get(8), + } + }).collect()) + }).await +} + +#[tauri::command] +pub async fn apply_table_properties( + app: AppHandle, + state: State<'_, AppState>, + input: TablePropertiesApplyRequest, +) -> Result<(), String> { + let (connection_id, engine) = resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; + if engine != DatabaseEngine::Postgres { + return Err(format!( + "Table property editing is not supported for {} connections yet.", + match engine { + DatabaseEngine::Postgres => "PostgreSQL", + DatabaseEngine::Mysql => "MySQL", + DatabaseEngine::Sqlite => "SQLite", + DatabaseEngine::Mongo => "MongoDB", + } + )); + } + + with_pool_client_retry(&app, &state, &connection_id, input, |mut client, input| async move { + let table_schema = input.table_schema; + let table_name = input.table_name; + let columns = input.columns; + + require_safe_identifier(&table_schema, "schema name")?; + require_safe_identifier(&table_name, "table name")?; + + let current_columns = client.query( + "select column_name, is_nullable \ + from information_schema.columns \ + where table_schema = $1 and table_name = $2", + &[&table_schema, &table_name], + ).await.map_err(|error| map_pg_err(error, None))?; + + let mut current_nullable: HashMap = HashMap::new(); + for row in current_columns { + let column_name: String = row.get(0); + let is_nullable = row.get::<_, String>(1) == "YES"; + current_nullable.insert(column_name, is_nullable); + } + + let primary_keys = client.query( + "select kcu.column_name \ + from information_schema.table_constraints tc \ + join information_schema.key_column_usage kcu \ + on tc.constraint_name = kcu.constraint_name \ + and tc.table_schema = kcu.table_schema \ + where tc.table_schema = $1 and tc.table_name = $2 \ + and tc.constraint_type = 'PRIMARY KEY'", + &[&table_schema, &table_name], + ).await.map_err(|error| map_pg_err(error, None))?; + + let primary_key_columns: HashSet = primary_keys.into_iter() + .filter_map(|row| Some(row.get::<_, String>(0))).collect(); + + let unique_constraints = client.query( + "select tc.constraint_name, kcu.column_name, kcu.ordinal_position \ + from information_schema.table_constraints tc \ + join information_schema.key_column_usage kcu \ + on tc.constraint_name = kcu.constraint_name \ + and tc.table_schema = kcu.table_schema \ + where tc.table_schema = $1 and tc.table_name = $2 \ + and tc.constraint_type = 'UNIQUE' \ + order by tc.constraint_name, kcu.ordinal_position", + &[&table_schema, &table_name], + ).await.map_err(|error| map_pg_err(error, None))?; + + let mut unique_by_name: HashMap> = HashMap::new(); + for row in unique_constraints { + let constraint_name: String = row.get(0); + let column_name: String = row.get(1); + unique_by_name.entry(constraint_name).or_default().push(column_name); + } + + let mut composite_unique_columns: HashSet = HashSet::new(); + let mut single_unique_constraint_names_by_column: HashMap> = HashMap::new(); + for (constraint_name, cols) in &unique_by_name { + if cols.len() > 1 { + for c in cols { composite_unique_columns.insert(c.clone()); } + } else if cols.len() == 1 { + let c = &cols[0]; + single_unique_constraint_names_by_column.entry(c.clone()).or_default().push(constraint_name.clone()); + } + } + + let mut desired_by_column: HashMap = HashMap::new(); + for update in columns { + desired_by_column.insert(update.column_name, (update.is_nullable, update.is_unique)); + } + + let txn = client.transaction().await.map_err(|error| map_pg_err(error, None))?; + + for (column_name, (desired_is_nullable, _desired_is_unique)) in &desired_by_column { + let current_is_nullable = current_nullable.get(column_name) + .ok_or_else(|| format!("Unknown column: {}", column_name))?; + if *current_is_nullable == *desired_is_nullable { continue; } + + let qualified_table = format!("\"{}\".\"{}\"", quote_identifier(&table_schema), quote_identifier(&table_name)); + require_safe_identifier(column_name, "column name")?; + let qualified_column = format!("\"{}\"", quote_identifier(column_name)); + + if *desired_is_nullable { + let sql = format!("ALTER TABLE {} ALTER COLUMN {} DROP NOT NULL", qualified_table, qualified_column); + txn.execute(sql.as_str(), &[]).await.map_err(|error| map_pg_err(error, Some(sql.as_str())))?; + } else { + let sql = format!("ALTER TABLE {} ALTER COLUMN {} SET NOT NULL", qualified_table, qualified_column); + txn.execute(sql.as_str(), &[]).await.map_err(|error| map_pg_err(error, Some(sql.as_str())))?; + } + } + + for (column_name, (_desired_is_nullable, desired_is_unique)) in &desired_by_column { + let is_primary_key = primary_key_columns.contains(column_name); + let is_part_of_composite_unique = composite_unique_columns.contains(column_name); + + if !*desired_is_unique { + if is_primary_key { + return Err(format!("Cannot disable UNIQUE for primary key column: {}", column_name)); + } + if is_part_of_composite_unique { + return Err(format!("Cannot disable UNIQUE for column in a composite UNIQUE constraint: {}", column_name)); + } + } + + let has_single_unique = single_unique_constraint_names_by_column.get(column_name) + .map(|names| !names.is_empty()).unwrap_or(false); + let current_is_unique = is_primary_key || has_single_unique || is_part_of_composite_unique; + + if *desired_is_unique == current_is_unique { continue; } + + let qualified_table = format!("\"{}\".\"{}\"", quote_identifier(&table_schema), quote_identifier(&table_name)); + require_safe_identifier(column_name, "column name")?; + let qualified_column = format!("\"{}\"", quote_identifier(column_name)); + + if *desired_is_unique { + if current_is_unique { continue; } + let generated_name = veloxdb_unique_constraint_name(&table_name, column_name); + if let Some(existing_cols) = unique_by_name.get(&generated_name) { + if existing_cols.len() != 1 || existing_cols[0] != *column_name { + return Err(format!( + "Cannot create UNIQUE constraint due to name collision ({}). Rename the existing constraint.", + generated_name + )); + } + } + let sql = format!("ALTER TABLE {} ADD CONSTRAINT \"{}\" UNIQUE ({})", + qualified_table, quote_identifier(&generated_name), qualified_column); + txn.execute(sql.as_str(), &[]).await.map_err(|error| map_pg_err(error, Some(sql.as_str())))?; + } else { + let constraint_names = single_unique_constraint_names_by_column.get(column_name) + .cloned().unwrap_or_default(); + for constraint_name in constraint_names { + require_safe_identifier(&constraint_name, "constraint name")?; + let sql = format!("ALTER TABLE {} DROP CONSTRAINT \"{}\"", + qualified_table, quote_identifier(&constraint_name)); + txn.execute(sql.as_str(), &[]).await.map_err(|error| map_pg_err(error, Some(sql.as_str())))?; + } + } + } + + txn.commit().await.map_err(|error| map_pg_err(error, None))?; + Ok(()) + }).await +} + +#[tauri::command] +pub async fn get_foreign_keys( + app: AppHandle, + state: State<'_, AppState>, + connection_id: Option, +) -> Result, String> { + let (connection_id, engine) = resolve_connection_engine(&app, &state, connection_id).await?; + + if engine == DatabaseEngine::Mysql { + let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; + let rows = sqlx::query( + "select kcu.table_schema as from_schema, kcu.table_name as from_table, \ + kcu.column_name as from_column, kcu.referenced_table_schema as to_schema, \ + kcu.referenced_table_name as to_table, kcu.referenced_column_name as to_column \ + from information_schema.key_column_usage kcu \ + where kcu.referenced_table_name is not null \ + order by kcu.table_schema, kcu.table_name, kcu.ordinal_position \ + limit ?", + ).bind(MAX_FOREIGN_KEY_ROWS).fetch_all(&pool).await.map_err(|error| error.to_string())?; + let mut edges = Vec::new(); + for row in rows { + edges.push(ForeignKeyEdge { + from_schema: mysql_get_string(&row, 0, "from_schema", "get_foreign_keys")?, + from_table: mysql_get_string(&row, 1, "from_table", "get_foreign_keys")?, + from_column: mysql_get_string(&row, 2, "from_column", "get_foreign_keys")?, + to_schema: mysql_get_string(&row, 3, "to_schema", "get_foreign_keys")?, + to_table: mysql_get_string(&row, 4, "to_table", "get_foreign_keys")?, + to_column: mysql_get_string(&row, 5, "to_column", "get_foreign_keys")?, + }); + } + return Ok(edges); + } + + if engine == DatabaseEngine::Sqlite { + let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; + let tables = sqlx::query( + "select name from sqlite_master \ + where type = 'table' and name not like 'sqlite_%'", + ).fetch_all(&pool).await.map_err(|error| error.to_string())?; + let mut edges = Vec::new(); + for table in tables { + let table_name: String = sqlite_get_idx(&table, 0, "name", "get_foreign_keys")?; + require_safe_identifier(&table_name, "table name")?; + let fk_sql = format!("PRAGMA foreign_key_list(\"{}\");", quote_identifier(&table_name)); + let fk_rows = sqlx::query(&fk_sql).fetch_all(&pool).await.map_err(|error| error.to_string())?; + for row in fk_rows { + edges.push(ForeignKeyEdge { + from_schema: "main".to_string(), from_table: table_name.clone(), + from_column: sqlite_get_name(&row, "from", "get_foreign_keys")?, + to_schema: "main".to_string(), + to_table: sqlite_get_name(&row, "table", "get_foreign_keys")?, + to_column: sqlite_get_name(&row, "to", "get_foreign_keys")?, + }); + if edges.len() >= MAX_FOREIGN_KEY_ROWS as usize { return Ok(edges); } + } + } + return Ok(edges); + } + + with_pool_client_retry(&app, &state, &connection_id, (), |client, ()| async move { + let rows = client.query( + "select src_ns.nspname::text as from_schema, src_cls.relname::text as from_table, \ + src_att.attname::text as from_column, tgt_ns.nspname::text as to_schema, \ + tgt_cls.relname::text as to_table, tgt_att.attname::text as to_column \ + from pg_constraint c \ + join pg_class src_cls on src_cls.oid = c.conrelid \ + join pg_namespace src_ns on src_ns.oid = src_cls.relnamespace \ + join pg_class tgt_cls on tgt_cls.oid = c.confrelid \ + join pg_namespace tgt_ns on tgt_ns.oid = tgt_cls.relnamespace \ + cross join lateral unnest(c.conkey, c.confkey) as u(attnum, confattnum) \ + join pg_attribute src_att on src_att.attrelid = c.conrelid \ + and src_att.attnum = u.attnum and not src_att.attisdropped \ + join pg_attribute tgt_att on tgt_att.attrelid = c.confrelid \ + and tgt_att.attnum = u.confattnum and not tgt_att.attisdropped \ + where c.contype = 'f' \ + and src_ns.nspname not in ('pg_catalog', 'information_schema') \ + order by src_ns.nspname, src_cls.relname, c.conname, u.attnum \ + limit $1", + &[&MAX_FOREIGN_KEY_ROWS], + ).await.map_err(|error| error.to_string())?; + + Ok(rows.into_iter().map(|row| ForeignKeyEdge { + from_schema: row.get(0), from_table: row.get(1), from_column: row.get(2), + to_schema: row.get(3), to_table: row.get(4), to_column: row.get(5), + }).collect()) + }).await +} + +#[tauri::command] +pub async fn get_table_indexes( + app: AppHandle, + state: State<'_, AppState>, + input: SchemaRequest, +) -> Result { + let (connection_id, engine) = resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; + let ctx = input.clone(); + + if engine == DatabaseEngine::Mysql { + let pool = get_or_create_mysql_pool(&app, &state, &connection_id).await?; + let fetch_limit = MAX_TABLE_INDEX_ROWS + 1; + let rows = sqlx::query( + "select table_schema as index_schema, index_name, table_schema, table_name, \ + non_unique = 0 as is_unique, index_name = 'PRIMARY' as is_primary, \ + true as is_valid, false as is_partial, \ + concat(index_name, ' (', group_concat(column_name order by seq_in_index separator ', '), ')') as definition, \ + 0 as index_bytes, 0 as idx_scan, 0 as idx_tup_read, 0 as idx_tup_fetch \ + from information_schema.statistics \ + where table_schema = ? and table_name = ? \ + group by table_schema, table_name, index_name, non_unique \ + order by index_name limit ?", + ).bind(&ctx.table_schema).bind(&ctx.table_name).bind(fetch_limit) + .fetch_all(&pool).await.map_err(|error| error.to_string())?; + let truncated = rows.len() as i64 > MAX_TABLE_INDEX_ROWS; + let mut indexes = Vec::new(); + for row in rows.into_iter().take(MAX_TABLE_INDEX_ROWS as usize) { + indexes.push(IndexInfo { + index_schema: mysql_get_string(&row, 0, "index_schema", "get_table_indexes")?, + index_name: mysql_get_string(&row, 1, "index_name", "get_table_indexes")?, + table_schema: mysql_get_string(&row, 2, "table_schema", "get_table_indexes")?, + table_name: mysql_get_string(&row, 3, "table_name", "get_table_indexes")?, + is_unique: mysql_get_idx(&row, 4, "is_unique", "get_table_indexes")?, + is_primary: mysql_get_idx(&row, 5, "is_primary", "get_table_indexes")?, + is_valid: true, is_partial: false, + definition: mysql_get_string(&row, 8, "definition", "get_table_indexes")?, + index_bytes: 0, idx_scan: 0, idx_tup_read: 0, idx_tup_fetch: 0, + }); + } + return Ok(TableIndexesResult { indexes, truncated }); + } + + if engine == DatabaseEngine::Sqlite { + let pool = get_or_create_sqlite_pool(&app, &state, &connection_id).await?; + require_safe_identifier(&ctx.table_name, "table name")?; + let pragma_sql = format!("PRAGMA index_list(\"{}\");", quote_identifier(&ctx.table_name)); + let rows = sqlx::query(&pragma_sql).fetch_all(&pool).await.map_err(|error| error.to_string())?; + let truncated = rows.len() as i64 > MAX_TABLE_INDEX_ROWS; + let mut indexes = Vec::new(); + for row in rows.into_iter().take(MAX_TABLE_INDEX_ROWS as usize) { + let index_name: String = sqlite_get_name(&row, "name", "get_table_indexes")?; + require_safe_identifier(&index_name, "index name")?; + let index_info_sql = format!("PRAGMA index_info(\"{}\");", quote_identifier(&index_name)); + let index_info_rows = sqlx::query(&index_info_sql).fetch_all(&pool).await + .map_err(|error| error.to_string())?; + let index_columns = index_info_rows.into_iter() + .filter_map(|idx| sqlite_get_name::(&idx, "name", "get_table_indexes").ok()) + .collect::>(); + indexes.push(IndexInfo { + index_schema: "main".to_string(), index_name: index_name.clone(), + table_schema: "main".to_string(), table_name: ctx.table_name.clone(), + is_unique: sqlite_get_name::(&row, "unique", "get_table_indexes")? == 1, + is_primary: sqlite_get_name::(&row, "origin", "get_table_indexes")? == "pk", + is_valid: true, + is_partial: sqlite_get_name::(&row, "partial", "get_table_indexes")? == 1, + definition: if index_columns.is_empty() { + format!("index {}", index_name) + } else { + format!("index {} ({})", index_name, index_columns.join(", ")) + }, + index_bytes: 0, idx_scan: 0, idx_tup_read: 0, idx_tup_fetch: 0, + }); + } + return Ok(TableIndexesResult { indexes, truncated }); + } + + with_pool_client_retry(&app, &state, &connection_id, ctx, |client, input| async move { + let table_schema = input.table_schema; + let table_name = input.table_name; + let fetch_limit = MAX_TABLE_INDEX_ROWS + 1; + + let rows = client.query( + "select ins.nspname::text as index_schema, ic.relname::text as index_name, \ + tn.nspname::text as table_schema, tc.relname::text as table_name, \ + i.indisunique as is_unique, i.indisprimary as is_primary, \ + i.indisvalid as is_valid, (i.indpred is not null) as is_partial, \ + pg_get_indexdef(i.indexrelid) as definition, \ + coalesce(pg_relation_size(i.indexrelid::regclass), 0)::bigint as index_bytes, \ + coalesce(s.idx_scan, 0)::bigint as idx_scan, \ + coalesce(s.idx_tup_read, 0)::bigint as idx_tup_read, \ + coalesce(s.idx_tup_fetch, 0)::bigint as idx_tup_fetch \ + from pg_index i \ + join pg_class ic on ic.oid = i.indexrelid \ + join pg_namespace ins on ins.oid = ic.relnamespace \ + join pg_class tc on tc.oid = i.indrelid \ + join pg_namespace tn on tn.oid = tc.relnamespace \ + left join pg_stat_user_indexes s on s.indexrelid = i.indexrelid \ + where tn.nspname = $1 and tc.relname = $2 \ + and ins.nspname not in ('pg_catalog', 'information_schema') \ + order by ic.relname limit $3", + &[&table_schema, &table_name, &fetch_limit], + ).await.map_err(|error| error.to_string())?; + + let truncated = rows.len() as i64 > MAX_TABLE_INDEX_ROWS; + let take = if truncated { MAX_TABLE_INDEX_ROWS as usize } else { rows.len() }; + let mut indexes = Vec::with_capacity(take); + for row in rows.into_iter().take(take) { + indexes.push(IndexInfo { + index_schema: row.get(0), index_name: row.get(1), + table_schema: row.get(2), table_name: row.get(3), + is_unique: row.get(4), is_primary: row.get(5), + is_valid: row.get(6), is_partial: row.get(7), + definition: row.get(8), + index_bytes: row.get(9), idx_scan: row.get(10), + idx_tup_read: row.get(11), idx_tup_fetch: row.get(12), + }); + } + Ok(TableIndexesResult { indexes, truncated }) + }).await +} diff --git a/src-tauri/src/commands/veloxy.rs b/src-tauri/src/commands/veloxy.rs new file mode 100644 index 0000000..63ba22d --- /dev/null +++ b/src-tauri/src/commands/veloxy.rs @@ -0,0 +1,440 @@ +use std::sync::Arc; +use std::sync::atomic::{AtomicBool, Ordering}; + +use serde_json::Value; +use tauri::{AppHandle, State}; + +use crate::db::{ + load_connection, resolve_connection_engine, AppState, MAX_QUERY_ROWS, +}; +use crate::models::{ + AskVeloxyChatRequest, AskVeloxyChatResponse, AskVeloxyConversationMessage, + AskVeloxyConversationResponse, AskVeloxyRequest, AskVeloxyResponse, + AskVeloxyTokenStats, VeloxyStreamChunk, +}; + +use super::{ + ask_veloxy_context_cache_key, ask_veloxy_conversation_key, + build_schema_context, classify_sql_intent, emit_veloxy_stream_chunk, + estimate_tokens, extract_openrouter_message_content, now_epoch_seconds, + normalize_openrouter_base, parse_ask_veloxy_chat_content, + parse_ask_veloxy_json, parse_ask_veloxy_suggestions, + stream_openrouter_chat_completion, truncate_on_char_boundary, validate_generated_sql, + ASK_VELOXY_MAX_HISTORY_MESSAGES, ASK_VELOXY_PROMPT_CHAR_BUDGET, +}; +use crate::models::DatabaseEngine; +use super::editor_meta::{ + fetch_foreign_keys_for_connection, fetch_query_editor_metadata_for_connection, +}; +use crate::models::AskVeloxyDbContextCache; + +async fn get_or_build_ask_veloxy_db_context( + app: &AppHandle, + state: &AppState, + connection_id: &str, + engine: DatabaseEngine, +) -> Result { + let stored_connection = load_connection(app, connection_id)? + .ok_or_else(|| "Stored connection details were not found.".to_string())?; + let cache_key = ask_veloxy_context_cache_key(connection_id, &stored_connection.database); + if let Some(cached) = state.ask_veloxy_db_context_cache.read().await.get(&cache_key).cloned() { + return Ok(cached); + } + + let metadata = fetch_query_editor_metadata_for_connection(app, state, connection_id, engine).await?; + let foreign_keys = fetch_foreign_keys_for_connection(app, state, connection_id, engine).await?; + let cache = AskVeloxyDbContextCache { + database_name: stored_connection.database, + engine, metadata, foreign_keys, + }; + state.ask_veloxy_db_context_cache.write().await.insert(cache_key, cache.clone()); + Ok(cache) +} + +#[tauri::command] +pub async fn cancel_veloxy_request(state: State<'_, AppState>) -> Result<(), String> { + if let Some(cancel) = state.veloxy_cancel.read().await.as_ref() { + cancel.store(true, Ordering::Relaxed); + } + Ok(()) +} + +#[tauri::command] +pub async fn chat_with_db( + app: AppHandle, + state: State<'_, AppState>, + input: AskVeloxyChatRequest, +) -> Result { + let natural_prompt = input.natural_prompt.trim(); + if natural_prompt.is_empty() { + return Err("Ask Veloxy prompt cannot be empty.".to_string()); + } + if input.provider_config.api_key.trim().is_empty() { + return Err("OpenRouter API key is required.".to_string()); + } + if input.provider_config.model.trim().is_empty() { + return Err("OpenRouter model is required.".to_string()); + } + + let (connection_id, engine) = resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; + let stored_connection = load_connection(&app, &connection_id)? + .ok_or_else(|| "Stored connection details were not found.".to_string())?; + let db_context = get_or_build_ask_veloxy_db_context(&app, &state, &connection_id, engine).await?; + let schema_context = build_schema_context(&db_context, natural_prompt, input.target_table.as_ref()); + let conversation_key = ask_veloxy_conversation_key(&connection_id, &stored_connection.database); + let history = state.ask_veloxy_conversations.read().await.get(&conversation_key).cloned().unwrap_or_default(); + + let history_block = history.iter().rev().take(8).rev() + .map(|message| format!("{}: {}", message.role, message.text)) + .collect::>().join("\n"); + + let mut user_prompt = format!( + "Engine: {:?}\nDatabase: {}\nTask: {}\nMaxRows: {}\nRecentConversation:\n{}\nSchemaContext:\n{}\n", + db_context.engine, db_context.database_name, natural_prompt, + input.max_rows.unwrap_or(MAX_QUERY_ROWS), history_block, schema_context + ); + truncate_on_char_boundary(&mut user_prompt, ASK_VELOXY_PROMPT_CHAR_BUDGET); + + let system_prompt = "You are Ask Veloxy chat mode. Return JSON when possible with keys: message (string), suggestions (array of strings), sqlDraft (string optional), needsSqlGeneration (boolean), needsClarification (boolean), warnings (array of strings). If JSON is not possible, return helpful plain text."; + let base_url = normalize_openrouter_base(input.provider_config.base_url.as_deref()); + let endpoint = format!("{}/chat/completions", base_url); + let client = state.openrouter_client.get_or_init(reqwest::Client::new); + let request_id = input.request_id.clone() + .filter(|value| !value.trim().is_empty()) + .unwrap_or_else(|| format!("req-{}", uuid::Uuid::new_v4())); + + let cancel = Arc::new(AtomicBool::new(false)); + { + let mut guard = state.veloxy_cancel.write().await; + *guard = Some(cancel.clone()); + } + + let (message_content, hit_token_limit) = stream_openrouter_chat_completion( + &app, client, &endpoint, + input.provider_config.api_key.trim(), + input.provider_config.model.trim(), + system_prompt, &user_prompt, &request_id, cancel.clone(), + ).await?; + + { + let mut guard = state.veloxy_cancel.write().await; + *guard = None; + } + + let (message, suggestions, mut warnings, sql_draft, needs_sql_generation, needs_clarification) = + parse_ask_veloxy_chat_content(&message_content); + + if cancel.load(Ordering::Relaxed) { warnings.push("Stopped early.".to_string()); } + if hit_token_limit { + warnings.push(format!( + "Response may be truncated (model output limit of {} tokens).", + super::ASK_VELOXY_MAX_CHAT_TOKENS + )); + } + + emit_veloxy_stream_chunk(&app, VeloxyStreamChunk { + request_id: request_id.clone(), + delta: String::new(), + done: true, + message: Some(message.clone()), + suggestions: suggestions.clone(), + warnings: warnings.clone(), + sql_draft: sql_draft.clone(), + needs_sql_generation, + needs_clarification, + }); + + { + let mut conversations = state.ask_veloxy_conversations.write().await; + let bucket = conversations.entry(conversation_key).or_default(); + bucket.push(AskVeloxyConversationMessage { + id: format!("msg-{}", uuid::Uuid::new_v4()), + role: "user".to_string(), mode: "chat".to_string(), + text: natural_prompt.to_string(), created_at: now_epoch_seconds(), + sql_draft: None, + }); + bucket.push(AskVeloxyConversationMessage { + id: format!("msg-{}", uuid::Uuid::new_v4()), + role: "assistant".to_string(), mode: "chat".to_string(), + text: message.clone(), created_at: now_epoch_seconds(), + sql_draft: sql_draft.clone(), + }); + if bucket.len() > ASK_VELOXY_MAX_HISTORY_MESSAGES { + let remove_count = bucket.len() - ASK_VELOXY_MAX_HISTORY_MESSAGES; + bucket.drain(0..remove_count); + } + } + + Ok(AskVeloxyChatResponse { + message, suggestions, warnings, sql_draft, needs_sql_generation, needs_clarification, + }) +} + +#[tauri::command] +pub async fn load_veloxy_conversation( + app: AppHandle, + state: State<'_, AppState>, + connection_id: Option, +) -> Result { + let (resolved_connection_id, _) = resolve_connection_engine(&app, &state, connection_id).await?; + let stored_connection = load_connection(&app, &resolved_connection_id)? + .ok_or_else(|| "Stored connection details were not found.".to_string())?; + let key = ask_veloxy_conversation_key(&resolved_connection_id, &stored_connection.database); + let messages = state.ask_veloxy_conversations.read().await.get(&key).cloned().unwrap_or_default(); + Ok(AskVeloxyConversationResponse { messages }) +} + +#[tauri::command] +pub async fn clear_veloxy_conversation( + app: AppHandle, + state: State<'_, AppState>, + connection_id: Option, +) -> Result<(), String> { + let (resolved_connection_id, _) = resolve_connection_engine(&app, &state, connection_id).await?; + let stored_connection = load_connection(&app, &resolved_connection_id)? + .ok_or_else(|| "Stored connection details were not found.".to_string())?; + let key = ask_veloxy_conversation_key(&resolved_connection_id, &stored_connection.database); + state.ask_veloxy_conversations.write().await.remove(&key); + Ok(()) +} + +#[tauri::command] +pub async fn generate_sql_from_nl( + app: AppHandle, + state: State<'_, AppState>, + input: AskVeloxyRequest, +) -> Result { + let natural_prompt = input.natural_prompt.trim(); + if natural_prompt.is_empty() { + return Err("Ask Veloxy prompt cannot be empty.".to_string()); + } + if input.provider_config.api_key.trim().is_empty() { + return Err("OpenRouter API key is required.".to_string()); + } + if input.provider_config.model.trim().is_empty() { + return Err("OpenRouter model is required.".to_string()); + } + + let (connection_id, engine) = resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; + let db_context = get_or_build_ask_veloxy_db_context(&app, &state, &connection_id, engine).await?; + let schema_context = build_schema_context(&db_context, natural_prompt, input.target_table.as_ref()); + + let mut user_prompt = format!( + "Engine: {:?}\nDatabase: {}\nTask: {}\nMaxRows: {}\nSchemaContext:\n{}\n", + db_context.engine, db_context.database_name, natural_prompt, + input.max_rows.unwrap_or(MAX_QUERY_ROWS), schema_context + ); + truncate_on_char_boundary(&mut user_prompt, ASK_VELOXY_PROMPT_CHAR_BUDGET); + + let system_prompt = "You are Ask Veloxy. Return JSON only with keys: sql (string), intent (string), confidence (number 0..1), explanation (string), suggestions (array of short strings), warnings (array of strings). Generate exactly one SQL statement, keep explanation concise, and never include markdown."; + let base_url = normalize_openrouter_base(input.provider_config.base_url.as_deref()); + let endpoint = format!("{}/chat/completions", base_url); + + let client = state.openrouter_client.get_or_init(reqwest::Client::new); + let response = client.post(&endpoint) + .header("Authorization", format!("Bearer {}", input.provider_config.api_key.trim())) + .header("Content-Type", "application/json") + .json(&serde_json::json!({ + "model": input.provider_config.model.trim(), + "temperature": 0.1, + "max_tokens": 500, + "messages": [ + { "role": "system", "content": system_prompt }, + { "role": "user", "content": user_prompt } + ] + })) + .send().await.map_err(|error| format!("OpenRouter request failed: {}", error))?; + + let status = response.status(); + let payload = response.json::().await + .map_err(|error| format!("Invalid OpenRouter JSON response: {}", error))?; + if !status.is_success() { + let message = payload.get("error").and_then(|error| error.get("message")) + .and_then(Value::as_str).unwrap_or("Unknown OpenRouter error"); + return Err(format!("OpenRouter error ({}): {}", status.as_u16(), message)); + } + + let message_content = extract_openrouter_message_content(&payload)?; + let generated = parse_ask_veloxy_json(&message_content)?; + let sql = generated.get("sql").and_then(Value::as_str).unwrap_or_default().trim().to_string(); + validate_generated_sql(&sql)?; + + let mut warnings = generated.get("warnings").and_then(Value::as_array) + .map(|items| items.iter().filter_map(Value::as_str).map(str::to_string).collect::>()) + .unwrap_or_default(); + + let intent = generated.get("intent").and_then(Value::as_str).map(str::to_string) + .unwrap_or_else(|| classify_sql_intent(&sql)); + let confidence = generated.get("confidence").and_then(Value::as_f64).unwrap_or(0.6).clamp(0.0, 1.0); + let explanation = generated.get("explanation").and_then(Value::as_str).map(str::trim) + .filter(|value| !value.is_empty()).map(|value| { + let mut truncated = value.to_string(); + truncate_on_char_boundary(&mut truncated, 350); + truncated + }); + let suggestions = parse_ask_veloxy_suggestions(&generated); + + if intent != "select" { + warnings.push("Generated SQL is not read-only. Review before execution.".to_string()); + } + if confidence < 0.5 { + warnings.push("Low confidence result. Review SQL carefully.".to_string()); + } + + let token_stats = AskVeloxyTokenStats { + schema_chars: schema_context.len(), + schema_tokens_estimate: estimate_tokens(schema_context.len()), + prompt_chars: user_prompt.len() + system_prompt.len(), + prompt_tokens_estimate: estimate_tokens(user_prompt.len() + system_prompt.len()), + }; + + Ok(AskVeloxyResponse { + sql, intent, confidence, explanation, suggestions, warnings, token_stats, + }) +} + +#[cfg(test)] +mod tests { + use super::super::{ + build_schema_context, classify_sql_intent, database_name_from_mysql_value, + decode_mysql_bytes_as_string, extract_openrouter_stream_delta, mysql_decode_error, + parse_ask_veloxy_json, sqlite_decode_error, streaming_display_text, + validate_generated_sql, is_read_only_sql, + }; + use crate::models::{ + AskVeloxyDbContextCache, DatabaseEngine, QueryEditorColumn, QueryEditorMetadata, + QueryEditorTable, + }; + + #[test] + fn streaming_display_text_extracts_partial_json_message() { + let partial = r#"{ "message": "The messages table has relationships with:\n- delivery_reports"#; + let display = streaming_display_text(partial); + assert!(display.contains("messages table")); + assert!(display.contains("delivery_reports")); + } + + #[test] + fn streaming_display_text_returns_plain_text_directly() { + assert_eq!(streaming_display_text("Hello from Veloxy"), "Hello from Veloxy"); + } + + #[test] + fn extract_openrouter_stream_delta_reads_content() { + let data = r#"{"choices":[{"delta":{"content":"Hello"}}]}"#; + assert_eq!(extract_openrouter_stream_delta(data).as_deref(), Some("Hello")); + } + + #[test] + fn database_name_from_mysql_value_rejects_empty() { + assert!(database_name_from_mysql_value(None, "list_databases").is_err()); + assert!(database_name_from_mysql_value(Some(String::new()), "list_databases").is_err()); + } + + #[test] + fn database_name_from_mysql_value_accepts_non_empty() { + let name = database_name_from_mysql_value(Some("my_app".to_string()), "list_databases").expect("name"); + assert_eq!(name, "my_app"); + } + + #[test] + fn decode_mysql_bytes_as_string_uses_utf8_text() { + assert_eq!(decode_mysql_bytes_as_string(b"my_schema"), "my_schema"); + } + + #[test] + fn mysql_decode_error_is_explicit() { + let message = mysql_decode_error("get_tables", "table_schema", Some(0), "mismatched types"); + assert!(message.contains("MySQL decode error")); + assert!(message.contains("get_tables")); + assert!(message.contains("table_schema")); + } + + #[test] + fn sqlite_decode_error_is_explicit() { + let message = sqlite_decode_error("get_schema", "name", Some(0), "unsupported value type"); + assert!(message.contains("SQLite decode error")); + assert!(message.contains("get_schema")); + assert!(message.contains("name")); + } + + #[test] + fn schema_context_is_bounded() { + let columns = (0..40).map(|idx| QueryEditorColumn { + name: format!("column_{}", idx), data_type: "text".to_string(), + }).collect::>(); + let tables = (0..20).map(|idx| QueryEditorTable { + schema: "public".to_string(), name: format!("events_{}", idx), columns: columns.clone(), + }).collect::>(); + let metadata = QueryEditorMetadata { + tables, functions: Vec::new(), + truncated_tables: false, truncated_columns: false, truncated_functions: false, + }; + let db_context = AskVeloxyDbContextCache { + database_name: "test".to_string(), engine: DatabaseEngine::Postgres, + metadata, foreign_keys: Vec::new(), + }; + let context = build_schema_context(&db_context, "show events", None); + assert!(!context.is_empty()); + assert!(context.len() <= super::super::ASK_VELOXY_SCHEMA_CHAR_BUDGET); + } + + #[test] + fn ask_veloxy_json_parser_handles_embedded_block() { + let content = "Here is the output {\"sql\":\"select 1\",\"intent\":\"select\",\"confidence\":0.9,\"warnings\":[]}"; + let parsed = parse_ask_veloxy_json(content).expect("json should parse"); + assert_eq!(parsed.get("sql").and_then(|v| v.as_str()), Some("select 1")); + } + + #[test] + fn sql_validation_rejects_multi_statement() { + assert!(validate_generated_sql("select 1; select 2;").is_err()); + } + + #[test] + fn sql_intent_classifier_recognizes_update() { + assert_eq!(classify_sql_intent("UPDATE foo SET bar = 1"), "update"); + } + + #[test] + fn read_only_check_allows_selects_and_explain() { + assert!(is_read_only_sql("SELECT 1")); + assert!(is_read_only_sql("EXPLAIN ANALYZE SELECT * FROM t")); + assert!(is_read_only_sql("WITH x AS (SELECT 1) SELECT * FROM x")); + assert!(is_read_only_sql("BEGIN; SELECT 1; COMMIT;")); + } + + #[test] + fn read_only_check_blocks_writes() { + assert!(!is_read_only_sql("DELETE FROM t")); + assert!(!is_read_only_sql("DROP TABLE t")); + assert!(!is_read_only_sql("BEGIN; UPDATE t SET a = 1; COMMIT;")); + assert!(!is_read_only_sql("SELECT 1; DELETE FROM t")); + assert!(!is_read_only_sql("")); + } + + #[test] + fn mysql_timestamp_formats_as_datetime_string() { + let dt = chrono::DateTime::parse_from_rfc3339("2024-03-15T10:30:45Z").unwrap() + .with_timezone(&chrono::Utc); + assert_eq!(dt.format("%Y-%m-%d %H:%M:%S").to_string(), "2024-03-15 10:30:45"); + } + + #[test] + fn mysql_datetime_formats_as_naive_datetime_string() { + let dt = chrono::NaiveDateTime::parse_from_str("2024-03-15 10:30:45", "%Y-%m-%d %H:%M:%S").unwrap(); + assert_eq!(dt.format("%Y-%m-%d %H:%M:%S").to_string(), "2024-03-15 10:30:45"); + } + + #[test] + fn mysql_date_formats_as_iso_date() { + let d = chrono::NaiveDate::from_ymd_opt(2024, 3, 15).unwrap(); + assert_eq!(d.to_string(), "2024-03-15"); + } + + #[test] + fn mysql_time_formats_as_iso_time() { + let t = chrono::NaiveTime::from_hms_opt(10, 30, 45).unwrap(); + assert_eq!(t.to_string(), "10:30:45"); + } +} diff --git a/src-tauri/src/db.rs b/src-tauri/src/db.rs index 1b618ab..9804e6d 100644 --- a/src-tauri/src/db.rs +++ b/src-tauri/src/db.rs @@ -14,6 +14,7 @@ use tokio::sync::RwLock; use tokio_postgres::NoTls; use rustls::pki_types::{CertificateDer, PrivateKeyDer}; +use mongodb::{Client as MongoClient, bson::doc}; use crate::models::{ AskVeloxyConversationMessage, AskVeloxyDbContextCache, ConnectionInput, ConnectionSslMode, ConnectionSummary, DatabaseEngine, StoredConnection, @@ -32,6 +33,7 @@ const POOL_WAIT_SECS: u64 = 30; const POOL_CREATE_SECS: u64 = 15; const POOL_RECYCLE_SECS: u64 = 15; pub const DEFAULT_MYSQL_PORT: u16 = 3306; +pub const DEFAULT_MONGO_PORT: u16 = 27017; fn deadpool_ssl_mode(mode: ConnectionSslMode) -> DeadpoolSslMode { match mode { @@ -46,6 +48,7 @@ pub struct AppState { pub pools: RwLock>, pub mysql_pools: RwLock>, pub sqlite_pools: RwLock>, + pub mongo_clients: RwLock>, pub active_connection_id: RwLock>, pub ssh_tunnels: RwLock>, pub ask_veloxy_db_context_cache: RwLock>, @@ -300,6 +303,60 @@ pub async fn build_sqlite_pool(input: &ConnectionInput) -> Result String { + let host = if input.host.is_empty() { "localhost" } else { &input.host }; + let port = if input.port == 0 { DEFAULT_MONGO_PORT } else { input.port }; + let mut uri = if input.user.is_empty() { + format!("mongodb://{}:{}/", host, port) + } else { + format!( + "mongodb://{}:{}@{}:{}/", + urlencoding::encode(&input.user), + urlencoding::encode(&input.password), + host, + port, + ) + }; + let database = if input.database.is_empty() { "admin" } else { &input.database }; + uri.push_str(database); + if let Some(ref params) = input.extra_params { + let qs: Vec = params.iter().map(|(k, v)| format!("{}={}", k, v)).collect(); + if !qs.is_empty() { + uri.push('?'); + uri.push_str(&qs.join("&")); + } + } + uri +} + +pub async fn get_or_create_mongo_client( + app: &AppHandle, + state: &AppState, + connection_id: &str, +) -> Result { + { + let clients = state.mongo_clients.read().await; + if let Some(client) = clients.get(connection_id) { + return Ok(client.clone()); + } + } + let stored = load_connection(app, connection_id)? + .ok_or_else(|| "Stored connection details were not found.".to_string())?; + let uri = build_mongo_connection_string(&stored.to_input()); + let client = MongoClient::with_uri_str(&uri) + .await + .map_err(|e| format!("MongoDB connection failed: {}", e))?; + client + .database("admin") + .run_command(doc! { "ping": 1 }) + .await + .map_err(|e| format!("MongoDB ping failed: {}", e))?; + state.mongo_clients.write().await.insert(connection_id.to_string(), client.clone()); + Ok(client) +} + /// Heuristic for transport-level failures where discarding the pool and opening /// a new TCP session may succeed (sleep/VPN blips, idle disconnects). fn is_retryable_connection_error(message: &str) -> bool { @@ -325,6 +382,7 @@ pub async fn drop_pool(state: &AppState, connection_id: &str) { state.pools.write().await.remove(connection_id); state.mysql_pools.write().await.remove(connection_id); state.sqlite_pools.write().await.remove(connection_id); + state.mongo_clients.write().await.remove(connection_id); state .ask_veloxy_db_context_cache .write() @@ -559,6 +617,11 @@ pub async fn refresh_connection_pools( .await .map_err(|error| error.to_string())?; } + DatabaseEngine::Mongo => { + let client = get_or_create_mongo_client(app, state, connection_id).await?; + client.database("admin").run_command(doc! { "ping": 1 }).await + .map_err(|e| format!("MongoDB ping failed: {}", e))?; + } } Ok(()) diff --git a/src-tauri/src/export.rs b/src-tauri/src/export.rs index f1fa29f..6732b9a 100644 --- a/src-tauri/src/export.rs +++ b/src-tauri/src/export.rs @@ -374,6 +374,9 @@ pub async fn export_results_csv( lines } } + DatabaseEngine::Mongo => { + return Err("MongoDB export is not supported.".to_string()); + } }; let content = lines.join("\n") + "\n"; @@ -480,6 +483,9 @@ pub async fn export_results_json( } result } + DatabaseEngine::Mongo => { + return Err("MongoDB JSON export is not supported.".to_string()); + } }; let content = format!("[\n{}\n]\n", rows.join(",\n")); diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index e5f991e..883c10d 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -12,7 +12,8 @@ use commands::{ delete_openrouter_api_key, disconnect_db, execute_ddl_statement, execute_ddl_transaction, export_diagram_png, export_results_csv_command, export_results_json_command, generate_sql_from_nl, get_foreign_keys, get_openrouter_api_key, get_query_editor_metadata, get_schema, get_table_indexes, get_table_properties, get_tables, - lint_sql, list_connections_command, list_databases, load_veloxy_conversation, ping_connection, + lint_sql, list_connections_command, list_databases, load_veloxy_conversation, mongo_run_query, mongo_get_collections, + mongo_get_schema, ping_connection, refresh_connection, rename_connection, run_query, save_base64_png, save_text_file, set_active_connection, store_openrouter_api_key, switch_database, }; @@ -81,7 +82,10 @@ pub fn run() { clear_veloxy_conversation, store_openrouter_api_key, get_openrouter_api_key, - delete_openrouter_api_key + delete_openrouter_api_key, + mongo_run_query, + mongo_get_collections, + mongo_get_schema ]) .run(tauri::generate_context!()) .expect("error while running tauri application"); diff --git a/src-tauri/src/models.rs b/src-tauri/src/models.rs index 3f3c876..a7a34f8 100644 --- a/src-tauri/src/models.rs +++ b/src-tauri/src/models.rs @@ -26,6 +26,7 @@ pub enum DatabaseEngine { Postgres, Mysql, Sqlite, + Mongo, } fn default_database_engine() -> DatabaseEngine { diff --git a/src/App.tsx b/src/App.tsx index c06f970..03f2b13 100644 --- a/src/App.tsx +++ b/src/App.tsx @@ -1,1251 +1,7 @@ -import { GearIcon, SidebarSimpleIcon } from "@phosphor-icons/react"; -import { useQueryClient } from "@tanstack/react-query"; -import { - type CSSProperties, - type PointerEvent as ReactPointerEvent, - useCallback, - useEffect, - useMemo, - useRef, - useState, -} from "react"; -import { useTranslation } from "react-i18next"; - -import { ErrorBoundary } from "@/components/ErrorBoundary"; -import { Button } from "@/components/ui/button"; -import { Tabs, TabsList, TabsTrigger } from "@/components/ui/tabs"; -import { queryKeys } from "@/data/query-keys"; -import { veloxDbRepository } from "@/data/repositories"; -import type { - AskVeloxyChatResponse, - AskVeloxyConversationResponse, - ConnectionSummary, - TableInfo, -} from "@/data/types"; -import { CommandPalette } from "@/features/commands/components/CommandPalette"; -import { ShortcutSheet } from "@/features/commands/components/ShortcutSheet"; -import { SettingsDialog } from "@/features/commands/components/SettingsDialog"; -import { ConnectionDialog } from "@/features/connections/components/ConnectionDialog"; -import { RenameConnectionDialog } from "@/features/connections/components/RenameConnectionDialog"; -import { ConnectionsSidebarTree } from "@/features/connections/components/ConnectionsSidebarTree"; -import { - useActivateConnectionMutation, - useConnectionsQuery, - useConnectMutation, - useDeleteConnectionMutation, - useRenameConnectionMutation, -} from "@/features/connections/queries"; -import { ModelWorkspace } from "@/features/model/components/ModelWorkspace"; +import { useState } from "react"; +import { VeloxApp } from "@/components/VeloxApp"; import { readOnboardingCompleted } from "@/features/onboarding/constants"; import { OnboardingFlow } from "@/features/onboarding/OnboardingFlow"; -import { - QueryWorkspace, - type QueryWorkspaceHandle, -} from "@/features/queries/components/QueryWorkspace"; -import { - AskVeloxySidebar, - type AskVeloxySubmitResult, -} from "@/features/queries/components/AskVeloxyDialog"; -import { useSaveResultEditsMutation, useDeleteRowsMutation } from "@/features/queries/queries"; -import { notifyError, notifySuccess } from "@/lib/error-notifier"; -import { loadOpenRouterApiKey } from "@/lib/openrouter-credentials"; -import { useSettings, resolveTheme } from "@/lib/settings"; -import { - buildDropTableSql, - buildDeleteTemplateSql, - buildInsertTemplateSql, - buildRenameTableSql, - buildSelectAllSql, - buildSelectCountSql, - buildUpdateTemplateSql, -} from "@/features/queries/sql-templates"; -import type { TableQuickSqlAction } from "@/features/queries/table-quick-actions"; -import { isInsertFormColumn, type ResultEditPatch } from "@/features/queries/result-edits"; -import { quoteIdent } from "@/lib/sql-ident"; -import { TablePropertiesDialog } from "@/features/schema/components/TablePropertiesDialog"; -import { - useTablePropertiesQuery, - useTableSchemaQuery, -} from "@/features/schema/queries"; -import { useTablesQuery } from "@/features/tables/queries"; - -const SIDEBAR_WIDTH_KEY = "veloxdb.sidebarWidth"; -const SIDEBAR_COLLAPSED_KEY = "veloxdb.sidebarCollapsed"; -const RESULTS_HEIGHT_KEY = "veloxdb.resultsHeight"; -const LAST_ACTIVE_CONNECTION_KEY = "veloxdb.lastActiveConnectionId"; -const DEFAULT_SIDEBAR_WIDTH = 280; -const MIN_SIDEBAR_WIDTH = 220; -const MAX_SIDEBAR_WIDTH = 520; -const DEFAULT_RESULTS_HEIGHT = 260; - -function clampSidebarWidth(value: number) { - return Math.min(MAX_SIDEBAR_WIDTH, Math.max(MIN_SIDEBAR_WIDTH, value)); -} - -function readSidebarWidth() { - const value = Number(window.localStorage.getItem(SIDEBAR_WIDTH_KEY)); - return Number.isFinite(value) - ? clampSidebarWidth(value) - : DEFAULT_SIDEBAR_WIDTH; -} - -function readSidebarCollapsed() { - return window.localStorage.getItem(SIDEBAR_COLLAPSED_KEY) === "true"; -} - -function readResultsHeight() { - const value = Number(window.localStorage.getItem(RESULTS_HEIGHT_KEY)); - return Number.isFinite(value) && value > 0 ? value : DEFAULT_RESULTS_HEIGHT; -} - -function persistLastActiveConnectionId(connectionId: string) { - window.localStorage.setItem(LAST_ACTIVE_CONNECTION_KEY, connectionId); -} - -function connectionSecondaryText(connection: ConnectionSummary): string { - if (connection.engine === "sqlite") { - return connection.filePath === ":memory:" - ? "SQLite in-memory database" - : `SQLite file: ${connection.filePath ?? connection.database}`; - } - return `${connection.user}@${connection.host}:${connection.port}${connection.sshConfig ? " (via SSH)" : ""}`; -} - -function connectionHeadline(connection: ConnectionSummary): string { - if (connection.engine === "sqlite") { - return `Connected to SQLite (${connection.filePath ?? connection.database})`; - } - return `Connected to ${connection.database} on ${connection.host}:${connection.port}`; -} - -function engineLabel(engine: ConnectionSummary["engine"]): string { - if (engine === "postgres") return "PostgreSQL"; - if (engine === "mysql") return "MySQL"; - return "SQLite"; -} - -function VeloxApp() { - const { t } = useTranslation(); - const [connection, setConnection] = useState(null); - const queryWorkspaceRef = useRef(null); - const [focusedQueryCaps, setFocusedQueryCaps] = useState({ - hasLastQuery: false, - hasResult: false, - }); - const [tableSearch, setTableSearch] = useState(""); - const [selectedTable, setSelectedTable] = useState(null); - const themeSetting = useSettings((s) => s.theme) - const isDark = useMemo(() => resolveTheme(themeSetting) === 'dark', [themeSetting]) - const fontSize = useSettings((s) => s.fontSize) - const [settingsOpen, setSettingsOpen] = useState(false) - const [commandPaletteOpen, setCommandPaletteOpen] = useState(false) - const [connectionDialogOpen, setConnectionDialogOpen] = useState(false); - const [renamingConnection, setRenamingConnection] = useState(null); - const [isSidebarCollapsed, setIsSidebarCollapsed] = - useState(readSidebarCollapsed); - const [sidebarWidth, setSidebarWidth] = useState(readSidebarWidth); - const [resultsHeight, setResultsHeight] = useState(readResultsHeight); - const [tablePropertiesDialogOpen, setTablePropertiesDialogOpen] = - useState(false); - const [tablePropertiesTarget, setTablePropertiesTarget] = useState<{ - connectionId: string; - table: TableInfo; - } | null>(null); - const [insertRowTrigger, setInsertRowTrigger] = useState(0); - const [mainWorkspace, setMainWorkspace] = useState<"query" | "model">( - "query", - ); - const [askVeloxyPending, setAskVeloxyPending] = useState(false); - const [askVeloxyError, setAskVeloxyError] = useState(null); - const veloxyOpenRouterApiKey = useSettings((s) => s.veloxyOpenRouterApiKey); - const veloxyModel = useSettings((s) => s.veloxyModel); - const veloxyBaseUrl = useSettings((s) => s.veloxyBaseUrl); - - const queryClient = useQueryClient(); - - const requestInsertRow = useCallback(() => { - setInsertRowTrigger((n) => n + 1); - }, []); - - const handleInsertRowSuccess = useCallback(() => { - void queryClient.invalidateQueries({ - queryKey: queryKeys.tableProperties(connection?.id, selectedTable), - }); - }, [connection?.id, queryClient, selectedTable]); - - useEffect(() => { - document.documentElement.classList.toggle("dark", isDark); - }, [isDark]); - - useEffect(() => { - const sizes = { sm: 12, md: 14, lg: 16 } - document.documentElement.style.fontSize = `${sizes[fontSize]}px` - }, [fontSize]) - - useEffect(() => { - void loadOpenRouterApiKey() - }, []) - - useEffect(() => { - window.localStorage.setItem( - SIDEBAR_COLLAPSED_KEY, - String(isSidebarCollapsed), - ); - }, [isSidebarCollapsed]); - - useEffect(() => { - window.localStorage.setItem(SIDEBAR_WIDTH_KEY, String(sidebarWidth)); - }, [sidebarWidth]); - - useEffect(() => { - window.localStorage.setItem(RESULTS_HEIGHT_KEY, String(resultsHeight)); - }, [resultsHeight]); - - const connectionsQuery = useConnectionsQuery(); - - useEffect(() => { - if (connectionsQuery.isError && connectionsQuery.error) { - notifyError(connectionsQuery.error, { - title: t("connection.failedToLoad"), - }); - } - }, [connectionsQuery.isError, connectionsQuery.error, t]); - - const connectMutation = useConnectMutation({ - onError: (error) => { - notifyError(error, { category: "connection", force: true }); - }, - onSuccess: (nextConnection) => { - notifySuccess( - t("connection.connected", { database: nextConnection.database }), - connectionSecondaryText(nextConnection), - ); - persistLastActiveConnectionId(nextConnection.id); - setConnection(nextConnection); - setSelectedTable(null); - setTableSearch(""); - setIsSidebarCollapsed(false); - setConnectionDialogOpen(false); - setTablePropertiesDialogOpen(false); - setTablePropertiesTarget(null); - queueMicrotask(() => { - queryWorkspaceRef.current?.setActiveTabConnection(nextConnection.id); - }); - }, - }); - - const activateConnectionMutation = useActivateConnectionMutation({ - onError: (error) => { - notifyError(error, { category: "connection", force: true }); - }, - onSuccess: (nextConnection) => { - persistLastActiveConnectionId(nextConnection.id); - setConnection(nextConnection); - setSelectedTable(null); - setTableSearch(""); - setTablePropertiesDialogOpen(false); - setTablePropertiesTarget(null); - queueMicrotask(() => { - queryWorkspaceRef.current?.setActiveTabConnection(nextConnection.id); - }); - }, - }); - - const deleteConnectionMutation = useDeleteConnectionMutation({ - onError: (error) => { - notifyError(error, { category: "connection", force: true }); - }, - onSuccess: (connectionId) => { - queryWorkspaceRef.current?.detachDeletedConnection(connectionId); - if (connection?.id === connectionId) { - setConnection(null); - setSelectedTable(null); - setTableSearch(""); - } - notifySuccess(t("connection.deleted")); - }, - }); - - const renameConnectionMutation = useRenameConnectionMutation({ - onError: (error) => { - notifyError(error, { category: "connection", force: true }); - }, - }); - - const connectionRestoreAttemptedRef = useRef(false); - const autoReconnect = useSettings((s) => s.autoReconnect) - - useEffect(() => { - if (connectionRestoreAttemptedRef.current) return; - if (!autoReconnect) { connectionRestoreAttemptedRef.current = true; return } - const list = connectionsQuery.data; - if (!list?.length) return; - if (connection) { - connectionRestoreAttemptedRef.current = true; - return; - } - - const savedId = window.localStorage.getItem(LAST_ACTIVE_CONNECTION_KEY); - if (!savedId) { - connectionRestoreAttemptedRef.current = true; - return; - } - - const match = list.find((c) => c.id === savedId); - const target = match ?? list[0]; - if (!target) { - connectionRestoreAttemptedRef.current = true; - return; - } - - connectionRestoreAttemptedRef.current = true; - activateConnectionMutation.mutate(target.id); - }, [connectionsQuery.data, connection, activateConnectionMutation, autoReconnect]); - - const tablesQuery = useTablesQuery(connection?.id); - - const schemaQuery = useTableSchemaQuery({ - connectionId: connection?.id, - table: selectedTable, - enabled: Boolean(connection?.id && selectedTable), - }); - const tablePropertiesQuery = useTablePropertiesQuery({ - connectionId: connection?.id, - table: selectedTable, - enabled: Boolean(connection?.id && selectedTable), - }); - const saveResultEditsMutation = useSaveResultEditsMutation({ - onError: (error) => { - notifyError(error, { - category: "query", - title: t("editor.failedToSave"), - }); - }, - }); - const deleteRowsMutation = useDeleteRowsMutation({ - onError: (error) => { - notifyError(error, { - category: "query", - title: t("editor.failedToDelete"), - }); - }, - }); - const connectionsErrorMessage = - connectionsQuery.error instanceof Error - ? connectionsQuery.error.message - : t("connection.failedToLoad"); - - const tablesErrorMessage = - tablesQuery.error instanceof Error - ? tablesQuery.error.message - : t("table.failedToLoad"); - - const schemaErrorMessage = - schemaQuery.error instanceof Error - ? schemaQuery.error.message - : t("table.failedToLoadSchema"); - const tablePropertiesErrorMessage = - tablePropertiesQuery.error instanceof Error - ? tablePropertiesQuery.error.message - : t("table.failedToLoadProperties"); - - const tablesForUi = tablesQuery.data ?? []; - const activeConnectionEngine = connection?.engine ?? "postgres"; - - const handleSelectTable = (table: TableInfo) => { - setSelectedTable(table); - queryWorkspaceRef.current?.applyTablePreview(table.previewQuery); - }; - - const handleTableQuickAction = useCallback( - async ( - action: TableQuickSqlAction, - connectionId: string, - table: TableInfo, - ) => { - if (action === "tableProperties") { - setTablePropertiesTarget({ connectionId, table }); - setTablePropertiesDialogOpen(true); - return; - } - if (action === "addRow") { - setSelectedTable(table); - setInsertRowTrigger((n) => n + 1); - return; - } - - setSelectedTable(table); - try { - switch (action) { - case "selectAll": - queryWorkspaceRef.current?.openTabWithSql( - buildSelectAllSql(table, 200, activeConnectionEngine), - ); - return; - case "selectCount": - queryWorkspaceRef.current?.openTabWithSql( - buildSelectCountSql(table, activeConnectionEngine), - ); - return; - case "insertTemplate": - case "updateTemplate": - case "deleteTemplate": { - const props = await queryClient.fetchQuery({ - queryKey: queryKeys.tableProperties(connectionId, table), - queryFn: () => - veloxDbRepository.getTableProperties(connectionId, table), - }); - - const pk = props - .filter((c) => c.isPrimaryKey) - .map((c) => c.columnName); - const insertCols = props - .filter(isInsertFormColumn) - .map((c) => c.columnName); - - if (action === "insertTemplate") { - queryWorkspaceRef.current?.openTabWithSql( - buildInsertTemplateSql( - table, - insertCols, - activeConnectionEngine, - ), - ); - return; - } - if (action === "updateTemplate") { - queryWorkspaceRef.current?.openTabWithSql( - buildUpdateTemplateSql( - table, - pk, - activeConnectionEngine, - ), - ); - return; - } - queryWorkspaceRef.current?.openTabWithSql( - buildDeleteTemplateSql( - table, - pk, - activeConnectionEngine, - ), - ); - return; - } - default: - return; - } - } catch (error) { - notifyError(error, { - category: "query", - title: "Table quick action failed", - force: true, - }); - } - }, - [queryClient, activeConnectionEngine], - ); - - const handleSelectConnection = (nextConnection: ConnectionSummary) => { - if (connection?.id === nextConnection.id) { - return; - } - - activateConnectionMutation.mutate(nextConnection.id); - }; - - const handleRefreshConnection = useCallback( - (connectionTarget: ConnectionSummary) => { - void (async () => { - try { - await veloxDbRepository.refreshConnection(connectionTarget.id); - } catch (error) { - notifyError(error, { category: "connection" }); - return; - } - - void queryClient.invalidateQueries({ queryKey: queryKeys.connections() }); - void queryClient.invalidateQueries({ - queryKey: queryKeys.databases(connectionTarget.id), - }); - void queryClient.refetchQueries({ - queryKey: queryKeys.databases(connectionTarget.id), - type: "active", - }); - void queryClient.invalidateQueries({ - queryKey: queryKeys.tables(connectionTarget.id), - }); - void queryClient.refetchQueries({ - queryKey: queryKeys.tables(connectionTarget.id), - type: "active", - }); - void queryClient.invalidateQueries({ - queryKey: queryKeys.queryEditorMetadata(connectionTarget.id), - }); - void queryClient.refetchQueries({ - queryKey: queryKeys.queryEditorMetadata(connectionTarget.id), - type: "active", - }); - void queryClient.invalidateQueries({ - queryKey: queryKeys.foreignKeys(connectionTarget.id), - }); - void queryClient.refetchQueries({ - queryKey: queryKeys.foreignKeys(connectionTarget.id), - type: "active", - }); - - if (connection?.id === connectionTarget.id) { - void queryClient.invalidateQueries({ - queryKey: queryKeys.schema(connectionTarget.id, selectedTable), - }); - void queryClient.refetchQueries({ - queryKey: queryKeys.schema(connectionTarget.id, selectedTable), - type: "active", - }); - void queryClient.invalidateQueries({ - queryKey: queryKeys.tableProperties(connectionTarget.id, selectedTable), - }); - void queryClient.refetchQueries({ - queryKey: queryKeys.tableProperties(connectionTarget.id, selectedTable), - type: "active", - }); - queryWorkspaceRef.current?.refreshFocusedResults(); - } - })(); - }, - [connection?.id, queryClient, selectedTable], - ); - - const handleRefreshTable = useCallback( - (connectionId: string, table: TableInfo) => { - void queryClient.invalidateQueries({ - queryKey: queryKeys.tables(connectionId), - }); - void queryClient.refetchQueries({ - queryKey: queryKeys.tables(connectionId), - type: "active", - }); - void queryClient.invalidateQueries({ - queryKey: queryKeys.schema(connectionId, table), - }); - void queryClient.refetchQueries({ - queryKey: queryKeys.schema(connectionId, table), - type: "active", - }); - void queryClient.invalidateQueries({ - queryKey: queryKeys.tableProperties(connectionId, table), - }); - void queryClient.refetchQueries({ - queryKey: queryKeys.tableProperties(connectionId, table), - type: "active", - }); - void queryClient.invalidateQueries({ - queryKey: queryKeys.tableIndexes(connectionId, table), - }); - void queryClient.refetchQueries({ - queryKey: queryKeys.tableIndexes(connectionId, table), - type: "active", - }); - - if ( - connection?.id === connectionId && - selectedTable?.schema === table.schema && - selectedTable?.name === table.name - ) { - queryWorkspaceRef.current?.refreshFocusedResults(); - } - }, - [connection?.id, queryClient, selectedTable?.name, selectedTable?.schema], - ); - - const handleRenameTableRequest = useCallback( - (_connectionId: string, table: TableInfo) => { - setSelectedTable(table); - queryWorkspaceRef.current?.appendQuerySql( - buildRenameTableSql(table, "new_table_name", connection?.engine ?? "postgres"), - ); - }, - [connection?.engine], - ); - - const handleDeleteTableRequest = useCallback( - (_connectionId: string, table: TableInfo) => { - setSelectedTable(table); - queryWorkspaceRef.current?.appendQuerySql( - buildDropTableSql(table, connection?.engine ?? "postgres"), - ); - }, - [connection?.engine], - ); - - const handleRenameConnectionRequest = useCallback( - (connectionTarget: ConnectionSummary) => { - setRenamingConnection(connectionTarget); - }, - [], - ); - - const handleRenameConnectionConfirm = useCallback( - (connectionTarget: ConnectionSummary, newName: string) => { - renameConnectionMutation.mutate( - { connectionId: connectionTarget.id, newName }, - { - onSuccess: (updated) => { - if (connection?.id === updated.id) { - setConnection(updated); - } - }, - }, - ); - setRenamingConnection(null); - }, - [connection?.id, renameConnectionMutation], - ); - - const handleDisconnectConnectionRequest = useCallback( - (connectionTarget: ConnectionSummary) => { - const confirmed = window.confirm( - t("connection.deleteConfirm", { name: connectionTarget.name }), - ); - if (!confirmed) return; - - deleteConnectionMutation.mutate(connectionTarget.id); - - if (connection?.id === connectionTarget.id) { - setConnection(null); - setSelectedTable(null); - setTableSearch(""); - setTablePropertiesDialogOpen(false); - setTablePropertiesTarget(null); - } - }, - [connection?.id, deleteConnectionMutation, t], - ); - - const handleCopyConnectionString = useCallback( - (target: ConnectionSummary) => { - const value = - target.engine === "sqlite" - ? `sqlite://${target.filePath ?? target.database}` - : `${target.engine === "mysql" ? "mysql" : "postgresql"}://${target.user}@${target.host}:${target.port}/${target.database}`; - void navigator.clipboard.writeText(value); - }, - [], - ); - - const handleTruncateTable = useCallback( - (_connectionId: string, table: TableInfo) => { - setSelectedTable(table); - if ((connection?.engine ?? "postgres") === "mysql") { - queryWorkspaceRef.current?.appendQuerySql( - `TRUNCATE TABLE ${quoteIdent(table.schema, "mysql")}.${quoteIdent(table.name, "mysql")};`, - ); - return; - } - if ((connection?.engine ?? "postgres") === "sqlite") { - queryWorkspaceRef.current?.appendQuerySql( - `DELETE FROM ${quoteIdent(table.name, "sqlite")};`, - ); - return; - } - queryWorkspaceRef.current?.appendQuerySql( - `TRUNCATE TABLE ${quoteIdent(table.schema, "postgres")}.${quoteIdent(table.name, "postgres")} RESTART IDENTITY CASCADE;`, - ); - }, - [connection?.engine], - ); - - const handleCopyTableName = useCallback( - (_connectionId: string, table: TableInfo) => { - const engine = connection?.engine ?? "postgres"; - const value = - engine === "sqlite" - ? quoteIdent(table.name, "sqlite") - : `${quoteIdent(table.schema, engine)}.${quoteIdent(table.name, engine)}`; - void navigator.clipboard.writeText(value); - }, - [connection?.engine], - ); - - const handleRefreshDatabases = useCallback( - (connectionId: string) => { - void queryClient.invalidateQueries({ - queryKey: queryKeys.databases(connectionId), - }); - void queryClient.refetchQueries({ - queryKey: queryKeys.databases(connectionId), - type: "active", - }); - }, - [queryClient], - ); - - const handleCopyDatabaseName = useCallback( - (_connectionId: string, database: string) => { - void navigator.clipboard.writeText(database); - }, - [], - ); - - const handleActivateConnectionForTab = useCallback( - (connectionId: string) => { - if (connection?.id === connectionId) { - return; - } - activateConnectionMutation.mutate(connectionId); - }, - [connection?.id, activateConnectionMutation], - ); - - const handleSidebarResizeStart = ( - event: ReactPointerEvent, - ) => { - const startX = event.clientX; - const startWidth = sidebarWidth; - - const handlePointerMove = (moveEvent: PointerEvent) => { - setSidebarWidth( - clampSidebarWidth(startWidth + moveEvent.clientX - startX), - ); - }; - - const handlePointerUp = () => { - window.removeEventListener("pointermove", handlePointerMove); - window.removeEventListener("pointerup", handlePointerUp); - }; - - window.addEventListener("pointermove", handlePointerMove); - window.addEventListener("pointerup", handlePointerUp); - }; - - useEffect(() => { - const onKeyDown = (event: KeyboardEvent) => { - const commandKey = event.metaKey || event.ctrlKey; - - if (commandKey && event.key.toLowerCase() === "p") { - event.preventDefault(); - setCommandPaletteOpen(true); - } - - if (commandKey && event.shiftKey && event.key.toLowerCase() === "c") { - event.preventDefault(); - setConnectionDialogOpen(true); - } - }; - - window.addEventListener("keydown", onKeyDown); - return () => window.removeEventListener("keydown", onKeyDown); - }, []); - - const layoutStyle = { - "--sidebar-width": `${isSidebarCollapsed ? 0 : sidebarWidth}px`, - } as CSSProperties; - - const connectionError = - connectMutation.error ?? activateConnectionMutation.error; - const connectionErrorMessage = - connectionError instanceof Error - ? connectionError.message - : t("connection.failedToConnect"); - const primaryKeyColumns = - tablePropertiesQuery.data - ?.filter((column) => column.isPrimaryKey) - .map((column) => column.columnName) ?? []; - const editableColumns = - tablePropertiesQuery.data - ?.filter((column) => !column.isPrimaryKey) - .map((column) => column.columnName) ?? []; - const hasSelectedTable = Boolean(selectedTable); - const hasQueryResult = focusedQueryCaps.hasResult; - const hasPrimaryKey = primaryKeyColumns.length > 0; - const isResultSingleTableEditable = - hasSelectedTable && - hasQueryResult && - hasPrimaryKey && - !tablePropertiesQuery.isError; - const saveDisabledReason = !hasSelectedTable - ? t("editor.selectTable") - : !hasQueryResult - ? t("editor.runQuery") - : tablePropertiesQuery.isLoading - ? t("editor.loadingMetadata") - : tablePropertiesQuery.isError - ? tablePropertiesErrorMessage - : !hasPrimaryKey - ? t("editor.requiresPrimaryKey") - : undefined; - - const handleSaveResultEdits = async (patches: ResultEditPatch[]) => { - if (!selectedTable || !connection?.id || patches.length === 0) { - return; - } - - await saveResultEditsMutation.mutateAsync({ - connectionId: connection.id, - engine: connection.engine, - table: selectedTable, - patches, - }); - - queryWorkspaceRef.current?.refreshFocusedResults(); - }; - - const handleDeleteRows = async ( - primaryKeys: Record[], - ) => { - if (!selectedTable || !connection?.id || primaryKeys.length === 0) { - return; - } - - await deleteRowsMutation.mutateAsync({ - connectionId: connection.id, - engine: connection.engine, - table: selectedTable, - primaryKeys, - }); - - queryWorkspaceRef.current?.refreshFocusedResults(); - }; - - const handleAskVeloxyChatSubmit = async ( - naturalPrompt: string, - requestId: string, - ): Promise => { - if (!connection?.id) { - const message = t("veloxy.selectConnection"); - setAskVeloxyError(message); - throw new Error(message); - } - if (!veloxyOpenRouterApiKey.trim()) { - const message = t("veloxy.addApiKey"); - setAskVeloxyError(message); - throw new Error(message); - } - if (!veloxyModel.trim()) { - const message = t("veloxy.chooseModel"); - setAskVeloxyError(message); - throw new Error(message); - } - setAskVeloxyPending(true); - setAskVeloxyError(null); - try { - return await veloxDbRepository.chatWithDb({ - connectionId: connection.id, - naturalPrompt, - requestId, - targetTable: selectedTable - ? { schema: selectedTable.schema, name: selectedTable.name } - : undefined, - providerConfig: { - apiKey: veloxyOpenRouterApiKey, - model: veloxyModel, - baseUrl: veloxyBaseUrl, - }, - maxRows: useSettings.getState().maxQueryRows, - }); - } catch (error) { - const message = - error instanceof Error - ? error.message - : t("veloxy.chatFailed"); - setAskVeloxyError(message); - notifyError(error, { category: "query", title: t("veloxy.chatFailed") }); - throw error instanceof Error ? error : new Error(message); - } finally { - setAskVeloxyPending(false); - } - }; - - const handleCancelVeloxyRequest = async () => { - try { - await veloxDbRepository.cancelVeloxyRequest(); - } catch (error) { - const message = - error instanceof Error ? error.message : "Failed to stop Veloxy."; - setAskVeloxyError(message); - } - }; - - const handleAskVeloxyActionSubmit = async ( - naturalPrompt: string, - ): Promise => { - if (!connection?.id) { - const message = t("veloxy.selectConnection"); - setAskVeloxyError(message); - throw new Error(message); - } - if (!veloxyOpenRouterApiKey.trim()) { - const message = t("veloxy.addApiKey"); - setAskVeloxyError(message); - throw new Error(message); - } - if (!veloxyModel.trim()) { - const message = t("veloxy.chooseModel"); - setAskVeloxyError(message); - throw new Error(message); - } - setAskVeloxyPending(true); - setAskVeloxyError(null); - try { - const response = await veloxDbRepository.generateSqlFromNl({ - connectionId: connection.id, - naturalPrompt, - targetTable: selectedTable - ? { schema: selectedTable.schema, name: selectedTable.name } - : undefined, - providerConfig: { - apiKey: veloxyOpenRouterApiKey, - model: veloxyModel, - baseUrl: veloxyBaseUrl, - }, - maxRows: useSettings.getState().maxQueryRows, - }); - const sql = response.sql.trim(); - const lower = sql.toLowerCase(); - const isReadIntent = response.intent === "select"; - const isLikelyLarge = - sql.length > 1800 || - (lower.includes("select") && !lower.includes(" limit ")) || - /\bcross\s+join\b|\bpg_sleep\s*\(/i.test(lower); - const canAutoRun = isReadIntent && !isLikelyLarge; - - if (canAutoRun) { - queryWorkspaceRef.current?.openTabWithSqlAndRun(sql); - notifySuccess(t("veloxy.generatedSql"), t("veloxy.autoRan")); - return { - response, - decision: "auto-ran", - }; - } - - queryWorkspaceRef.current?.openTabWithSql(sql); - return { - response, - decision: "needs-confirmation", - decisionReason: isReadIntent - ? t("veloxy.needsConfirmation") - : t("veloxy.nonReadConfirmation"), - pendingSql: sql, - }; - } catch (error) { - const message = - error instanceof Error - ? error.message - : t("veloxy.generateFailed"); - setAskVeloxyError(message); - notifyError(error, { category: "query", title: t("veloxy.generateFailed") }); - throw error instanceof Error ? error : new Error(message); - } finally { - setAskVeloxyPending(false); - } - }; - - const handleLoadVeloxyConversation = async (): Promise => { - if (!connection?.id) return { messages: [] }; - try { - return await veloxDbRepository.loadVeloxyConversation(connection.id); - } catch (error) { - const message = - error instanceof Error - ? error.message - : t("veloxy.loadFailed"); - setAskVeloxyError(message); - return { messages: [] }; - } - }; - - const handleClearVeloxyConversation = async () => { - if (!connection?.id) return; - try { - await veloxDbRepository.clearVeloxyConversation(connection.id); - } catch (error) { - const message = - error instanceof Error - ? error.message - : t("veloxy.clearFailed"); - setAskVeloxyError(message); - throw error; - } - }; - - - return ( -
- {!isSidebarCollapsed ? ( - <> -
- - Sidebar failed to render. -
- } - > - {connectionsQuery.isError ? ( -
- {connectionsErrorMessage} -
- ) : ( - setConnectionDialogOpen(true)} - onSelectConnection={handleSelectConnection} - onSelectTable={handleSelectTable} - onTableQuickAction={handleTableQuickAction} - onRefreshConnection={handleRefreshConnection} - onRefreshTable={handleRefreshTable} - onRenameConnection={handleRenameConnectionRequest} - onDisconnectConnection={handleDisconnectConnectionRequest} - onRenameTable={handleRenameTableRequest} - onDeleteTable={handleDeleteTableRequest} - onTruncateTable={handleTruncateTable} - onCopyTableName={handleCopyTableName} - onRefreshDatabases={handleRefreshDatabases} - onCopyDatabaseName={handleCopyDatabaseName} - onCopyConnectionString={handleCopyConnectionString} - onDatabaseSwitched={setConnection} - onToggleCollapsed={() => setIsSidebarCollapsed(true)} - /> - )} - -
-
- - ) : null} - -
-
-
-
- {isSidebarCollapsed ? ( - - ) : null} - - - setMainWorkspace(value as "query" | "model") - } - className="shrink-0" - > - - - {t("workspace.query")} - - - {t("workspace.model")} - - - - -
-

- VeloxDB.dev -

-

- {connection - ? connectionHeadline(connection) - : t("workspace.chooseConnection")} -

-
-
- -
- - -
-
-
- - {mainWorkspace === "query" ? ( - setConnectionDialogOpen(true)} - resultsHeight={resultsHeight} - onResultsHeightChange={setResultsHeight} - selectedTable={selectedTable} - schemaLoading={schemaQuery.isLoading} - schemaError={schemaQuery.isError ? schemaErrorMessage : null} - columnCount={schemaQuery.data?.length ?? null} - primaryKeyColumns={primaryKeyColumns} - editableColumns={editableColumns} - saveDisabledReason={saveDisabledReason} - isResultSingleTableEditable={isResultSingleTableEditable} - saveResultEditsMutation={saveResultEditsMutation} - onSaveResultEdits={handleSaveResultEdits} - onDeleteRows={handleDeleteRows} - onFocusedTabCapabilitiesChange={setFocusedQueryCaps} - onActivateConnectionForTab={handleActivateConnectionForTab} - insertRowTrigger={insertRowTrigger} - insertConnectionId={connection?.id ?? null} - insertTable={selectedTable} - canInsertRow={Boolean(connection?.id && selectedTable)} - onInsertRowSuccess={handleInsertRowSuccess} - onOpenAddRow={ - connection?.id && selectedTable - ? requestInsertRow - : undefined - } - askVeloxySidebar={(onClose) => ( - { - setSettingsOpen(true); - }} - onChatSubmit={handleAskVeloxyChatSubmit} - onActionSubmit={handleAskVeloxyActionSubmit} - onLoadConversation={handleLoadVeloxyConversation} - onClearConversation={handleClearVeloxyConversation} - onConfirmRun={async (sql) => { - queryWorkspaceRef.current?.openTabWithSqlAndRun(sql); - notifySuccess(t("veloxy.queryExecuted")); - }} - onInsertSql={(sql) => { - queryWorkspaceRef.current?.appendQuerySql(sql); - }} - onReplaceSql={(sql) => { - queryWorkspaceRef.current?.replaceQuerySql(sql); - }} - onOpenTabWithSql={(sql) => { - queryWorkspaceRef.current?.openTabWithSql(sql); - }} - onCancelRequest={handleCancelVeloxyRequest} - errorMessage={askVeloxyError} - /> - )} - /> - ) : connection?.id ? ( - - {connection.engine === "postgres" ? ( - - ) : ( -
- {t("model.postgresOnly", { engine: engineLabel(connection.engine) })} -
- )} -
- ) : ( -
- {t("model.connectToUse")} -
- )} -
- - { - connectMutation.mutate(values); - }} - isPending={connectMutation.isPending} - /> - - setRenamingConnection(null)} - /> - - { - setTablePropertiesDialogOpen(nextOpen); - if (!nextOpen) setTablePropertiesTarget(null); - }} - connectionId={tablePropertiesTarget?.connectionId} - tablePropertyEditingSupported={connection?.tablePropertyEditingSupported} - table={tablePropertiesTarget?.table ?? null} - /> - - setConnectionDialogOpen(true)} - onRunLastQuery={() => { - queryWorkspaceRef.current?.runLastQuery(); - }} - onSelectTable={handleSelectTable} - /> - - -
- ); -} function App() { const [onboardingDone, setOnboardingDone] = useState(() => diff --git a/src/components/VeloxApp.tsx b/src/components/VeloxApp.tsx new file mode 100644 index 0000000..f0a8e3a --- /dev/null +++ b/src/components/VeloxApp.tsx @@ -0,0 +1,330 @@ +import { GearIcon, SidebarSimpleIcon } from "@phosphor-icons/react"; +import { + type CSSProperties, + type PointerEvent as ReactPointerEvent, + useEffect, + useRef, +} from "react"; +import { useTranslation } from "react-i18next"; + +import { ErrorBoundary } from "@/components/ErrorBoundary"; +import { Button } from "@/components/ui/button"; +import { Tabs, TabsList, TabsTrigger } from "@/components/ui/tabs"; +import { CommandPalette } from "@/features/commands/components/CommandPalette"; +import { ShortcutSheet } from "@/features/commands/components/ShortcutSheet"; +import { SettingsDialog } from "@/features/commands/components/SettingsDialog"; +import { ConnectionDialog } from "@/features/connections/components/ConnectionDialog"; +import { RenameConnectionDialog } from "@/features/connections/components/RenameConnectionDialog"; +import { ConnectionsSidebarTree } from "@/features/connections/components/ConnectionsSidebarTree"; +import { QueryWorkspace } from "@/features/queries/components/QueryWorkspace"; +import { AskVeloxySidebar } from "@/features/queries/components/AskVeloxyDialog"; +import { ModelWorkspace } from "@/features/model/components/ModelWorkspace"; +import { TablePropertiesDialog } from "@/features/schema/components/TablePropertiesDialog"; + +import { + clampSidebarWidth, + connectionHeadline, + engineLabel, + useAppState, +} from "@/hooks/useAppState"; + +export function VeloxApp() { + const { t } = useTranslation(); + const queryWorkspaceRef = useRef(null); + + const app = useAppState(queryWorkspaceRef); + + const { + connection, + focusedQueryCaps, + tableSearch, setTableSearch, + selectedTable, + settingsOpen, setSettingsOpen, + commandPaletteOpen, setCommandPaletteOpen, + connectionDialogOpen, setConnectionDialogOpen, + renamingConnection, + isSidebarCollapsed, setIsSidebarCollapsed, + sidebarWidth, setSidebarWidth, + resultsHeight, setResultsHeight, + tablePropertiesDialogOpen, setTablePropertiesDialogOpen, + tablePropertiesTarget, + insertRowTrigger, + mainWorkspace, setMainWorkspace, + askVeloxyPending, + askVeloxyError, + isDark, + veloxyModel, + veloxyOpenRouterApiKey, + connectionsQuery, tablesQuery, schemaQuery, + saveResultEditsMutation, + connectMutation, + tablesForUi, + connectionError, connectionErrorMessage, + connectionsErrorMessage, tablesErrorMessage, + schemaErrorMessage, + primaryKeyColumns, editableColumns, + isResultSingleTableEditable, saveDisabledReason, + handleSelectTable, + handleTableQuickAction, + handleSelectConnection, + handleRefreshConnection, + handleRefreshTable, + handleRenameTableRequest, + handleDeleteTableRequest, + handleRenameConnectionRequest, + handleRenameConnectionConfirm, + handleDisconnectConnectionRequest, + handleCopyConnectionString, + handleTruncateTable, + handleCopyTableName, + handleRefreshDatabases, + handleCopyDatabaseName, + handleActivateConnectionForTab, + handleSaveResultEdits, + handleDeleteRows, + requestInsertRow, + handleInsertRowSuccess, + handleAskVeloxyChatSubmit, + handleCancelVeloxyRequest, + handleAskVeloxyActionSubmit, + handleLoadVeloxyConversation, + handleClearVeloxyConversation, + } = app; + + const handleSidebarResizeStart = (event: ReactPointerEvent) => { + const startX = event.clientX; + const startWidth = sidebarWidth; + const handlePointerMove = (moveEvent: PointerEvent) => { + setSidebarWidth(clampSidebarWidth(startWidth + moveEvent.clientX - startX)); + }; + const handlePointerUp = () => { + window.removeEventListener("pointermove", handlePointerMove); + window.removeEventListener("pointerup", handlePointerUp); + }; + window.addEventListener("pointermove", handlePointerMove); + window.addEventListener("pointerup", handlePointerUp); + }; + + useEffect(() => { + const onKeyDown = (event: KeyboardEvent) => { + const commandKey = event.metaKey || event.ctrlKey; + if (commandKey && event.key.toLowerCase() === "p") { + event.preventDefault(); + setCommandPaletteOpen(true); + } + if (commandKey && event.shiftKey && event.key.toLowerCase() === "c") { + event.preventDefault(); + setConnectionDialogOpen(true); + } + }; + window.addEventListener("keydown", onKeyDown); + return () => window.removeEventListener("keydown", onKeyDown); + }, [setCommandPaletteOpen, setConnectionDialogOpen]); + + const layoutStyle = { + "--sidebar-width": `${isSidebarCollapsed ? 0 : sidebarWidth}px`, + } as CSSProperties; + + return ( +
+ {!isSidebarCollapsed ? ( + <> +
+ Sidebar failed to render.
}> + {connectionsQuery.isError ? ( +
{connectionsErrorMessage}
+ ) : ( + setConnectionDialogOpen(true)} + onSelectConnection={handleSelectConnection} + onSelectTable={handleSelectTable} + onTableQuickAction={handleTableQuickAction} + onRefreshConnection={handleRefreshConnection} + onRefreshTable={handleRefreshTable} + onRenameConnection={handleRenameConnectionRequest} + onDisconnectConnection={handleDisconnectConnectionRequest} + onRenameTable={handleRenameTableRequest} + onDeleteTable={handleDeleteTableRequest} + onTruncateTable={handleTruncateTable} + onCopyTableName={handleCopyTableName} + onRefreshDatabases={handleRefreshDatabases} + onCopyDatabaseName={handleCopyDatabaseName} + onCopyConnectionString={handleCopyConnectionString} + onDatabaseSwitched={app.setConnection} + onToggleCollapsed={() => setIsSidebarCollapsed(true)} + /> + )} + +
+
+ + ) : null} + +
+
+
+
+ {isSidebarCollapsed ? ( + + ) : null} + setMainWorkspace(value as "query" | "model")} className="shrink-0"> + + {t("workspace.query")} + + {t("workspace.model")} + + + +
+

VeloxDB.dev

+

+ {connection ? connectionHeadline(connection) : t("workspace.chooseConnection")} +

+
+
+
+ + +
+
+
+ + {mainWorkspace === "query" ? ( + setConnectionDialogOpen(true)} + resultsHeight={resultsHeight} + onResultsHeightChange={setResultsHeight} + selectedTable={selectedTable} + schemaLoading={schemaQuery?.isLoading ?? false} + schemaError={schemaQuery?.isError ? schemaErrorMessage : null} + columnCount={schemaQuery?.data?.length ?? null} + primaryKeyColumns={primaryKeyColumns} + editableColumns={editableColumns} + saveDisabledReason={saveDisabledReason} + isResultSingleTableEditable={isResultSingleTableEditable} + saveResultEditsMutation={saveResultEditsMutation} + onSaveResultEdits={handleSaveResultEdits} + onDeleteRows={handleDeleteRows} + onFocusedTabCapabilitiesChange={app.setFocusedQueryCaps} + onActivateConnectionForTab={handleActivateConnectionForTab} + insertRowTrigger={insertRowTrigger} + insertConnectionId={connection?.id ?? null} + insertTable={selectedTable} + canInsertRow={Boolean(connection?.id && selectedTable)} + onInsertRowSuccess={handleInsertRowSuccess} + onOpenAddRow={connection?.id && selectedTable ? requestInsertRow : undefined} + askVeloxySidebar={(onClose) => ( + { setSettingsOpen(true); }} + onChatSubmit={handleAskVeloxyChatSubmit} + onActionSubmit={handleAskVeloxyActionSubmit} + onLoadConversation={handleLoadVeloxyConversation} + onClearConversation={handleClearVeloxyConversation} + onConfirmRun={async (sql) => { + queryWorkspaceRef.current?.openTabWithSqlAndRun(sql); + }} + onInsertSql={(sql) => { queryWorkspaceRef.current?.appendQuerySql(sql); }} + onReplaceSql={(sql) => { queryWorkspaceRef.current?.replaceQuerySql(sql); }} + onOpenTabWithSql={(sql) => { queryWorkspaceRef.current?.openTabWithSql(sql); }} + onCancelRequest={handleCancelVeloxyRequest} + errorMessage={askVeloxyError} + /> + )} + /> + ) : connection?.id ? ( + + {connection.engine === "postgres" ? ( + + ) : ( +
+ {t("model.postgresOnly", { engine: engineLabel(connection.engine) })} +
+ )} +
+ ) : ( +
+ {t("model.connectToUse")} +
+ )} +
+ + { connectMutation.mutate(values); }} + isPending={connectMutation.isPending} + /> + app.setRenamingConnection(null)} + /> + { + setTablePropertiesDialogOpen(nextOpen); + if (!nextOpen) app.setTablePropertiesTarget(null); + }} + connectionId={tablePropertiesTarget?.connectionId} + tablePropertyEditingSupported={connection?.tablePropertyEditingSupported} + table={tablePropertiesTarget?.table ?? null} + /> + setConnectionDialogOpen(true)} + onRunLastQuery={() => { queryWorkspaceRef.current?.runLastQuery(); }} + onSelectTable={handleSelectTable} + /> + + +
+ ); +} diff --git a/src/data/types.ts b/src/data/types.ts index e8610f1..6aaea55 100644 --- a/src/data/types.ts +++ b/src/data/types.ts @@ -1,6 +1,6 @@ /** PostgreSQL `sslmode`-style TLS (lowercase in JSON for Tauri). */ export type ConnectionSslMode = 'disable' | 'prefer' | 'require' -export type DatabaseEngine = 'postgres' | 'mysql' | 'sqlite' +export type DatabaseEngine = 'postgres' | 'mysql' | 'sqlite' | 'mongo' export type SshAuthMethod = 'keyfile' | 'password' diff --git a/src/features/connections/components/ConnectionDialog.tsx b/src/features/connections/components/ConnectionDialog.tsx index f4269b5..0455f7a 100644 --- a/src/features/connections/components/ConnectionDialog.tsx +++ b/src/features/connections/components/ConnectionDialog.tsx @@ -169,6 +169,7 @@ export function ConnectionDialog({ { value: 'postgres' as DatabaseEngine, label: 'PostgreSQL', hint: t("connection.recommendedDefault") }, { value: 'mysql' as DatabaseEngine, label: 'MySQL', hint: t("connection.experimental"), experimental: true }, { value: 'sqlite' as DatabaseEngine, label: 'SQLite', hint: t("connection.experimental"), experimental: true }, + { value: 'mongo' as DatabaseEngine, label: 'MongoDB', hint: t("connection.experimental"), experimental: true }, ], [t]) const defaultValues = useMemo( diff --git a/src/features/connections/components/ConnectionsSidebarTree.tsx b/src/features/connections/components/ConnectionsSidebarTree.tsx index b0c4bcc..dc473ed 100644 --- a/src/features/connections/components/ConnectionsSidebarTree.tsx +++ b/src/features/connections/components/ConnectionsSidebarTree.tsx @@ -37,6 +37,7 @@ import { readExpandedIds, writeExpandedIds } from '@/lib/tree-expanded-persisten function engineBadge(engine: ConnectionSummary['engine']): string { if (engine === 'postgres') return 'PG' if (engine === 'mysql') return 'MY' + if (engine === 'mongo') return 'MG' return 'SQ' } diff --git a/src/features/model/components/ModelWorkspace.tsx b/src/features/model/components/ModelWorkspace.tsx index 7621b2f..a851666 100644 --- a/src/features/model/components/ModelWorkspace.tsx +++ b/src/features/model/components/ModelWorkspace.tsx @@ -1,1925 +1,654 @@ -import { useQueries, useQueryClient } from '@tanstack/react-query' -import { save } from '@tauri-apps/plugin-dialog' -import { useShallow } from 'zustand/react/shallow' +import { useQueryClient } from '@tanstack/react-query'; +import { useCallback, useEffect, useMemo, useRef, useState } from 'react'; +import { useTranslation } from 'react-i18next'; +import { Button } from '@/components/ui/button'; +import { Input } from '@/components/ui/input'; +import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs'; +import { queryKeys } from '@/data/query-keys'; +import type { DatabaseEngine, TableInfo } from '@/data/types'; +import { applyEntireModel } from '@/features/model/apply-entire-model'; +import { CreateTableDialog } from '@/features/model/components/CreateTableDialog'; +import { DdlReviewDialog } from '@/features/model/components/DdlReviewDialog'; +import { DiagramSurfaceAdapter } from '@/features/model/components/DiagramSurfaceAdapter'; +import { ModelCatalog } from '@/features/model/components/ModelCatalog'; +import { ModelInspector } from '@/features/model/components/ModelInspector'; +import { ModelWorkspaceToolbar } from '@/features/model/components/ModelWorkspaceToolbar'; +import { MigrationPreviewDialog } from '@/features/model/components/MigrationPreviewDialog'; +import { defaultDiagramHeaderHex as distinctDiagramHeaderHex } from '@/features/model/diagram-header-palette'; +import { readDiagramPalette } from '@/features/model/diagram-theme'; +import { useModelColumns } from '@/features/model/hooks/useModelColumns'; +import { useModelInitialization } from '@/features/model/hooks/useModelInitialization'; +import { useModelWorkspaceStore } from '@/features/model/hooks/useModelWorkspaceStore'; import { - AlignBottomIcon, - AlignLeftIcon, - AlignRightIcon, - AlignTopIcon, - ArrowCounterClockwiseIcon, - ArrowClockwiseIcon, - ArrowsClockwiseIcon, - ArrowsInSimpleIcon, - ArrowsOutIcon, - DownloadSimpleIcon, - FilePdfIcon, - GridFourIcon, - MagnetIcon, - PlusIcon, - SquaresFourIcon, - TrashIcon, - TreeStructureIcon, -} from '@phosphor-icons/react' -import { useCallback, useEffect, useMemo, useRef, useState } from 'react' -import { useTranslation } from 'react-i18next' - -import { Button } from '@/components/ui/button' -import { Input } from '@/components/ui/input' -import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs' -import { queryKeys } from '@/data/query-keys' -import { veloxDbRepository } from '@/data/repositories' -import type { ColumnInfo, DatabaseEngine, TableInfo } from '@/data/types' -import { - applyEntireModel, - type PendingCreateTable, - type TableIdentityDraft, -} from '@/features/model/apply-entire-model' -import { CreateTableDialog } from '@/features/model/components/CreateTableDialog' -import { DdlReviewDialog } from '@/features/model/components/DdlReviewDialog' -import { DiagramSurfaceAdapter } from '@/features/model/components/DiagramSurfaceAdapter' -import type { DiagramExportHandle } from '@/features/model/components/diagram-surface-types' -import { ModelCatalog } from '@/features/model/components/ModelCatalog' -import { ModelInspector } from '@/features/model/components/ModelInspector' -import { - alignSelectedBottom, - alignSelectedLeft, - alignSelectedRight, - alignSelectedTop, - snapPoint, -} from '@/features/model/diagram-geometry/snap' -import { topologicalLayoutOrder } from '@/features/model/diagram-geometry/topological-layout-order' -import { computeDagreLayout } from '@/features/model/diagram-geometry/dagre-layout' + deleteDiagramViewLayout, duplicateLayoutSnapshotForNewView, + ensurePositions, gridPositionForIndex, + loadDiagramLayout, loadDiagramViewsRegistry, + saveDiagramLayout, saveDiagramViewsRegistry, +} from '@/features/model/model-layout-storage'; import { - buildMigrationSummary, - buildMigrationSql, -} from '@/features/model/migration-preview' -import { MigrationPreviewDialog } from '@/features/model/components/MigrationPreviewDialog' -import { - deleteDiagramViewLayout, - duplicateLayoutSnapshotForNewView, - ensurePositions, - gridPositionForIndex, - loadDiagramLayout, - loadDiagramViewsRegistry, - saveDiagramLayout, - saveDiagramViewsRegistry, -} from '@/features/model/model-layout-storage' -import { defaultDiagramHeaderHex as distinctDiagramHeaderHex } from '@/features/model/diagram-header-palette' -import { TABLE_NODE_WIDTH, tableNodeHeight } from '@/features/model/table-node-metrics' -import { readDiagramPalette } from '@/features/model/diagram-theme' -import { - DEFAULT_DIAGRAM_VIEW_ID, - tableKey, - type ColumnDetailLevel, - type DiagramLayoutSnapshot, - type TableKey, -} from '@/features/model/model-types' -import { useForeignKeysQuery } from '@/features/model/queries' -import { useContainerSize } from '@/features/model/use-container-size' -import { useCanvasStore } from '@/features/model/state/canvas-store' -import { canQueueRelationship } from '@/features/model/relationship-validation' -import { rgbCssToHex } from '@/lib/contrast-text-for-bg' -import { cn } from '@/lib/utils' + DEFAULT_DIAGRAM_VIEW_ID, tableKey, + type TableKey, +} from '@/features/model/model-types'; +import { useForeignKeysQuery } from '@/features/model/queries'; +import { useContainerSize } from '@/features/model/use-container-size'; +import { canQueueRelationship } from '@/features/model/relationship-validation'; +import { rgbCssToHex } from '@/lib/contrast-text-for-bg'; type ModelWorkspaceProps = { - connectionId: string - connectionEngine: DatabaseEngine - defaultDatabaseName: string - isDark: boolean - tables: TableInfo[] - tablesErrorMessage?: string - isTablesLoading: boolean - selectedTable: TableInfo | null -} - -const LOAD_ALL_CONFIRM_THRESHOLD = 150 - -function tableKeyToParts(key: TableKey): { schema: string; name: string } { - const [schema = '', name = ''] = key.split('.') - return { schema, name } -} + connectionId: string; + connectionEngine: DatabaseEngine; + defaultDatabaseName: string; + isDark: boolean; + tables: TableInfo[]; + tablesErrorMessage?: string; + isTablesLoading: boolean; + selectedTable: TableInfo | null; +}; + +const LOAD_ALL_CONFIRM_THRESHOLD = 150; export function ModelWorkspace({ - connectionId, - connectionEngine, - defaultDatabaseName, - isDark, - tables, - tablesErrorMessage, - isTablesLoading, - selectedTable, + connectionId, connectionEngine, defaultDatabaseName, + isDark, tables, tablesErrorMessage, isTablesLoading, selectedTable, }: ModelWorkspaceProps) { - const { t } = useTranslation() - const queryClient = useQueryClient() - const foreignKeysQuery = useForeignKeysQuery(connectionId) - - const boot = useMemo(() => { - const vr = loadDiagramViewsRegistry(connectionId) - const aid = vr.activeViewId - const snap = loadDiagramLayout(connectionId, aid) - return { vr, aid, snap } - }, [connectionId]) - - const diagramWrapRef = useRef(null) - const diagramAreaSize = useContainerSize(diagramWrapRef) - - const hadStoredLayout = - boot.snap != null && (boot.snap.onCanvas.length > 0 || Object.keys(boot.snap.positions).length > 0) - const hydrateFromConnection = useCanvasStore((s) => s.hydrateFromConnection) - const { - hydrated, - storeConnectionId, - viewsRegistry, - setViewsRegistry, - activeViewId, - diagramTool, - setDiagramTool, - selectedKeys, - setSelectedKeys, - primaryKey, - replaceSelection, - selectTable, - clearSelection, - selectSingleFromCatalog, - snapToGrid, - setSnapToGrid, - onCanvas, - setOnCanvas, - positions, - setPositions, - viewport, - setViewport, - modelTitle, - setModelTitle, - headerColorsByKey, - setHeaderColorsByKey, - columnDetail, - setColumnDetail, - diagramGroups, - setDiagramGroups, - modelTab, - setModelTab, - identityDraftByKey, - setIdentityDraftByKey, - columnOverridesByKey, - setColumnOverridesByKey, - columnIdentityOverridesByKey, - setColumnIdentityOverridesByKey, - pendingAddColumnsByKey, - setPendingAddColumnsByKey, - pendingForeignKeys, - setPendingForeignKeys, - selectedEdge, - setSelectedEdge, - pendingRules, - setPendingRules, - pendingTriggers, - setPendingTriggers, - pendingRlsPolicies, - setPendingRlsPolicies, - pendingCreateTables, - setPendingCreateTables, - applyQuickColumnEdit, - canUndo, - canRedo, - undo, - redo, - } = useCanvasStore( - useShallow((s) => ({ - hydrated: s.hydrated, - storeConnectionId: s.connectionId, - viewsRegistry: s.viewsRegistry, - setViewsRegistry: s.setViewsRegistry, - activeViewId: s.activeViewId, - diagramTool: s.diagramTool, - setDiagramTool: s.setDiagramTool, - selectedKeys: s.selectedKeys, - setSelectedKeys: s.setSelectedKeys, - primaryKey: s.primaryKey, - replaceSelection: s.replaceSelection, - selectTable: s.selectTable, - clearSelection: s.clearSelection, - applyMarquee: s.applyMarquee, - selectSingleFromCatalog: s.selectSingleFromCatalog, - snapToGrid: s.snapToGrid, - setSnapToGrid: s.setSnapToGrid, - onCanvas: s.onCanvas, - setOnCanvas: s.setOnCanvas, - positions: s.positions, - setPositions: s.setPositions, - viewport: s.viewport, - setViewport: s.setViewport, - modelTitle: s.modelTitle, - setModelTitle: s.setModelTitle, - headerColorsByKey: s.headerColorsByKey, - setHeaderColorsByKey: s.setHeaderColorsByKey, - columnDetail: s.columnDetail, - setColumnDetail: s.setColumnDetail, - diagramGroups: s.diagramGroups, - setDiagramGroups: s.setDiagramGroups, - modelTab: s.modelTab, - setModelTab: s.setModelTab, - identityDraftByKey: s.identityDraftByKey, - setIdentityDraftByKey: s.setIdentityDraftByKey, - columnOverridesByKey: s.columnOverridesByKey, - setColumnOverridesByKey: s.setColumnOverridesByKey, - columnIdentityOverridesByKey: s.columnIdentityOverridesByKey, - setColumnIdentityOverridesByKey: s.setColumnIdentityOverridesByKey, - pendingAddColumnsByKey: s.pendingAddColumnsByKey, - setPendingAddColumnsByKey: s.setPendingAddColumnsByKey, - pendingForeignKeys: s.pendingForeignKeys, - setPendingForeignKeys: s.setPendingForeignKeys, - selectedEdge: s.selectedEdge, - setSelectedEdge: s.setSelectedEdge, - pendingRules: s.pendingRules, - setPendingRules: s.setPendingRules, - pendingTriggers: s.pendingTriggers, - setPendingTriggers: s.setPendingTriggers, - pendingRlsPolicies: s.pendingRlsPolicies, - setPendingRlsPolicies: s.setPendingRlsPolicies, - pendingCreateTables: s.pendingCreateTables, - setPendingCreateTables: s.setPendingCreateTables, - applyQuickColumnEdit: s.applyQuickColumnEdit, - canUndo: s.canUndo, - canRedo: s.canRedo, - undo: s.undo, - redo: s.redo, - })), - ) - - useEffect(() => { - hydrateFromConnection({ connectionId, defaultDatabaseName }) - }, [connectionId, defaultDatabaseName, hydrateFromConnection]) - const [columnRequestKeys, setColumnRequestKeys] = useState([]) - const [ddlOpen, setDdlOpen] = useState(false) - const [createTableOpen, setCreateTableOpen] = useState(false) - const [migrationPreviewOpen, setMigrationPreviewOpen] = useState(false) - const [applyPending, setApplyPending] = useState(false) - const [applyError, setApplyError] = useState(null) - const [initialSeedReason, setInitialSeedReason] = useState(null) - - const fkSeedDoneRef = useRef(false) - const initialRecoveryDoneRef = useRef(false) - - useEffect(() => { - void connectionId - initialRecoveryDoneRef.current = false - fkSeedDoneRef.current = false - setInitialSeedReason(null) - }, [connectionId]) - - useEffect(() => { - const ignoredTags = new Set(['INPUT', 'TEXTAREA', 'SELECT', 'BUTTON']) - const onKeyDown = (e: KeyboardEvent) => { - const target = e.target as HTMLElement | null - if (target?.isContentEditable || (target && ignoredTags.has(target.tagName))) return - const mod = e.metaKey || e.ctrlKey - if (!mod || e.altKey) return - if (e.key.toLowerCase() !== 'z') return - e.preventDefault() - if (e.shiftKey) redo() - else undo() - } - window.addEventListener('keydown', onKeyDown) - return () => window.removeEventListener('keydown', onKeyDown) - }, [redo, undo]) - - const tablesByKey = useMemo(() => { - const m = new Map() - for (const t of tables) { - m.set(tableKey(t), t) - } - return m - }, [tables]) - - useEffect(() => { - if (initialRecoveryDoneRef.current) return - if (!tables.length) return - - const validOnCanvas = onCanvas.filter((k) => tablesByKey.has(k)) - if (validOnCanvas.length !== onCanvas.length) { - setOnCanvas(validOnCanvas) - setPositions((prev) => { - const next: Record = {} - for (const key of validOnCanvas) { - if (prev[key]) next[key] = prev[key] - } - return next - }) - } - - if (validOnCanvas.length === 0) { - const fkData = foreignKeysQuery.data ?? [] - const fkSeed = new Set() - for (const edge of fkData) { - const from = `${edge.fromSchema}.${edge.fromTable}` as TableKey - const to = `${edge.toSchema}.${edge.toTable}` as TableKey - if (tablesByKey.has(from)) fkSeed.add(from) - if (tablesByKey.has(to)) fkSeed.add(to) - } - const fallbackKeys = - fkSeed.size > 0 - ? [...fkSeed] - : tables.slice(0, 12).map((t) => tableKey(t)) - if (fallbackKeys.length > 0) { - setInitialSeedReason(fkSeed.size > 0 ? 'relationships' : 'sample') - setOnCanvas(fallbackKeys) - setPositions((prev) => ensurePositions(fallbackKeys, prev)) - } - } - - initialRecoveryDoneRef.current = true - }, [foreignKeysQuery.data, onCanvas, setOnCanvas, setPositions, tables, tablesByKey]) - - useEffect(() => { - if (hadStoredLayout) return - if (fkSeedDoneRef.current) return - const fkData = foreignKeysQuery.data - if (!fkData?.length || !tables.length) return - - const keys = new Set() - for (const e of fkData) { - keys.add(`${e.fromSchema}.${e.fromTable}`) - keys.add(`${e.toSchema}.${e.toTable}`) - } - const valid = [...keys].filter((k) => tables.some((t) => tableKey(t) === k)) - if (valid.length === 0) return - - fkSeedDoneRef.current = true - queueMicrotask(() => { - setOnCanvas(valid) - setPositions((p) => ensurePositions(valid, p)) - }) - }, [hadStoredLayout, foreignKeysQuery.data, setOnCanvas, setPositions, tables]) - - useEffect(() => { - if (!selectedTable) return - const k = tableKey(selectedTable) - queueMicrotask(() => { - selectSingleFromCatalog(k) - setColumnRequestKeys((prev) => (prev.includes(k) ? prev : [...prev, k])) - setOnCanvas((prev) => (prev.includes(k) ? prev : [...prev, k])) - setPositions((p) => ensurePositions([k], p)) - }) - }, [selectSingleFromCatalog, selectedTable]) - - useEffect(() => { - if (!primaryKey) return - const t = tablesByKey.get(primaryKey) - if (!t) return - setIdentityDraftByKey((prev) => - prev[primaryKey] != null ? prev : { ...prev, [primaryKey]: { schema: t.schema, name: t.name } }, - ) - }, [primaryKey, tablesByKey]) - - const sortedRequestKeys = useMemo(() => [...columnRequestKeys].sort(), [columnRequestKeys]) - - const columnQueries = useQueries({ - queries: sortedRequestKeys.map((key) => { - const table = tablesByKey.get(key) - return { - queryKey: queryKeys.schema(connectionId, table ?? null), - queryFn: () => { - if (!table) throw new Error('Table not found for schema request.') - return veloxDbRepository.getSchema(connectionId, table) - }, - enabled: Boolean(connectionId && table), - staleTime: 5 * 60 * 1000, - } - }), - }) - - const diagramPalette = useMemo(() => readDiagramPalette(isDark), [isDark]) - const themeDiagramHeaderHex = useMemo( - () => rgbCssToHex(diagramPalette.header), - [diagramPalette.header], - ) - - const columnsByKey = useMemo(() => { - const out: Record = {} - sortedRequestKeys.forEach((key, i) => { - const q = columnQueries[i] - if (!q) { - out[key] = null - return - } - if (q.isPending && !q.data) out[key] = null - else if (q.data) out[key] = q.data - else out[key] = null - }) - return out - }, [columnQueries, sortedRequestKeys]) - - const effectiveColumnsByKey = useMemo(() => { - const out: Record = {} - for (const [key, cols] of Object.entries(columnsByKey) as Array<[TableKey, ColumnInfo[] | null]>) { - if (!cols) { - out[key] = null - continue - } - const identityOverrides = columnIdentityOverridesByKey[key] ?? {} - const rows = cols.map((col) => { - const patch = identityOverrides[col.columnName] - if (!patch) return col - const nextName = patch.nextColumnName.trim() - const nextType = patch.nextDataType.trim() - return { - ...col, - columnName: nextName || col.columnName, - dataType: nextType || col.dataType, - } - }) - const pending = pendingAddColumnsByKey[key] ?? [] - const pendingRows: ColumnInfo[] = pending.map((col) => ({ - tableSchema: tableKeyToParts(key).schema, - tableName: tableKeyToParts(key).name, - columnName: col.columnName.trim(), - dataType: col.dataType.trim(), - isNullable: col.nullable, - })) - out[key] = [...rows, ...pendingRows] - } - return out - }, [columnIdentityOverridesByKey, columnsByKey, pendingAddColumnsByKey]) - - useEffect(() => { - if (!hydrated) return - if (storeConnectionId !== connectionId) return - const t = window.setTimeout(() => { - saveDiagramLayout( - connectionId, - { - onCanvas, - positions, - viewport, - modelTitle: modelTitle.trim() || defaultDatabaseName, - diagramTool, - snapToGrid, - columnDetail, - ...(diagramGroups.length > 0 ? { diagramGroups } : {}), - ...(Object.keys(headerColorsByKey).length > 0 ? { headerColors: headerColorsByKey } : {}), - }, - activeViewId, - ) - saveDiagramViewsRegistry(connectionId, viewsRegistry) - }, 400) - return () => window.clearTimeout(t) - }, [ - hydrated, - activeViewId, - columnDetail, - connectionId, - defaultDatabaseName, - diagramGroups, - diagramTool, - headerColorsByKey, - modelTitle, - onCanvas, - positions, - snapToGrid, - storeConnectionId, - viewsRegistry, - viewport, - ]) - - const onCanvasSet = useMemo(() => new Set(onCanvas), [onCanvas]) - - const fkColumnNamesByKey = useMemo(() => { - const m = new Map>() - const add = (tab: TableKey, col: string) => { - if (!onCanvasSet.has(tab)) return - const existing = m.get(tab) - if (existing) { - existing.add(col) - return - } - m.set(tab, new Set([col])) - } - for (const fk of foreignKeysQuery.data ?? []) { - add(`${fk.fromSchema}.${fk.fromTable}` as TableKey, fk.fromColumn) - add(`${fk.toSchema}.${fk.toTable}` as TableKey, fk.toColumn) - } - for (const p of pendingForeignKeys) { - add(p.fromKey, p.fromColumn) - add(p.toKey, p.toColumn) - } - return m - }, [foreignKeysQuery.data, onCanvasSet, pendingForeignKeys]) - - const diagramDisplayColumnsByKey = useMemo((): Record => { - const out: Record = {} - for (const k of onCanvas) { - const cols = effectiveColumnsByKey[k] ?? null - if (columnDetail === 'header') { - out[k] = [] - continue - } - if (columnDetail === 'keys' && cols?.length) { - const set = fkColumnNamesByKey.get(k) - const filtered = set?.size ? cols.filter((c) => set.has(c.columnName)) : cols.slice(0, 4) - out[k] = filtered.length > 0 ? filtered : cols.slice(0, 4) - continue - } - out[k] = cols - } - return out - }, [columnDetail, effectiveColumnsByKey, fkColumnNamesByKey, onCanvas]) - - useEffect(() => { - const keys = new Set() - for (const fk of foreignKeysQuery.data ?? []) { - const fromK = `${fk.fromSchema}.${fk.fromTable}` as TableKey - const toK = `${fk.toSchema}.${fk.toTable}` as TableKey - if (onCanvasSet.has(fromK)) keys.add(fromK) - if (onCanvasSet.has(toK)) keys.add(toK) - } - for (const p of pendingForeignKeys) { - keys.add(p.fromKey) - keys.add(p.toKey) - } - if (keys.size === 0) return - setColumnRequestKeys((prev) => { - let next = prev - for (const k of keys) { - if (!next.includes(k)) next = [...next, k] - } - return next - }) - }, [foreignKeysQuery.data, onCanvasSet, pendingForeignKeys]) - - const tablesOnCanvas = useMemo(() => { - const list: TableInfo[] = [] - for (const k of onCanvas) { - const t = tablesByKey.get(k) - if (t) list.push(t) - } - return list - }, [onCanvas, tablesByKey]) - const totalTableCount = tables.length - const onDiagramCount = tablesOnCanvas.length - const hiddenTableCount = Math.max(totalTableCount - onDiagramCount, 0) - const isPartialDiagram = hiddenTableCount > 0 - const showInitialSeedHint = initialSeedReason != null && isPartialDiagram - - const tableDisplays = useMemo(() => { - return tablesOnCanvas.map((t) => { - const k = tableKey(t) - const id = identityDraftByKey[k] - return { - key: k, - schema: id?.schema ?? t.schema, - name: id?.name ?? t.name, - } - }) - }, [tablesOnCanvas, identityDraftByKey]) - - const resolvedHeaderColors = useMemo(() => { - const out: Record = {} - for (const t of tableDisplays) { - out[t.key] = headerColorsByKey[t.key] ?? distinctDiagramHeaderHex(t.key, isDark) - } - return out - }, [tableDisplays, headerColorsByKey, isDark]) - - const inspectorTable = useMemo(() => { - if (!primaryKey) return null - return tablesByKey.get(primaryKey) ?? null - }, [primaryKey, tablesByKey]) - - const identityDraftForInspector = useMemo((): TableIdentityDraft | null => { - if (!primaryKey || !inspectorTable) return null - return identityDraftByKey[primaryKey] ?? { - schema: inspectorTable.schema, - name: inspectorTable.name, - } - }, [primaryKey, inspectorTable, identityDraftByKey]) - - const selectedKeysSet = useMemo(() => new Set(selectedKeys), [selectedKeys]) - const editedColumnNamesByKey = useMemo((): Record> => { - const out: Record> = {} - for (const key of onCanvas) { - const cols = effectiveColumnsByKey[key] ?? [] - const edited = new Set() - const overrides = columnIdentityOverridesByKey[key] ?? {} - for (const col of cols) { - const lowered = col.columnName.trim().toLowerCase() - for (const patch of Object.values(overrides)) { - if (patch.nextColumnName.trim().toLowerCase() === lowered) { - edited.add(col.columnName) - } - } - } - out[key] = edited - } - return out - }, [columnIdentityOverridesByKey, effectiveColumnsByKey, onCanvas]) - - const canQueueForeignKey = useCallback( - (input: { fromKey: TableKey; fromColumn: string; toKey: TableKey; toColumn: string }) => - canQueueRelationship(input, foreignKeysQuery.data ?? [], pendingForeignKeys), - [foreignKeysQuery.data, pendingForeignKeys], - ) - - const positionsRef = useRef(positions) - positionsRef.current = positions - const selectedKeysRef = useRef(selectedKeys) - selectedKeysRef.current = selectedKeys - const onCanvasRef = useRef(onCanvas) - onCanvasRef.current = onCanvas - - type TableDragState = { - draggedKey: TableKey - keys: TableKey[] - start: Record - } - const tableDragStateRef = useRef(null) - const diagramExportRef = useRef(null) - const viewportControlRef = useRef<{ - setViewport: (v: import('@/features/model/model-types').ViewportState) => void - getViewport: () => import('@/features/model/model-types').ViewportState - } | null>(null) - - const isModelDirty = useMemo(() => { - if (pendingForeignKeys.length > 0) return true - if (pendingCreateTables.length > 0) return true - for (const k of onCanvas) { - const t = tablesByKey.get(k) - if (!t) continue - const id = identityDraftByKey[k] - if (id && (id.schema !== t.schema || id.name !== t.name)) return true - const co = columnOverridesByKey[k] - if (co && Object.keys(co).length > 0) return true - const cio = columnIdentityOverridesByKey[k] - if (cio && Object.keys(cio).length > 0) return true - const adds = pendingAddColumnsByKey[k] - if (adds && adds.length > 0) return true - } - if (pendingRules.length > 0 || pendingTriggers.length > 0 || pendingRlsPolicies.length > 0) return true - return false - }, [ - onCanvas, - tablesByKey, - identityDraftByKey, - columnOverridesByKey, - columnIdentityOverridesByKey, - pendingAddColumnsByKey, - pendingForeignKeys, - pendingCreateTables, - pendingRlsPolicies.length, - pendingRules.length, - pendingTriggers.length, - ]) - - const migrationSummary = useMemo( - () => - !isModelDirty - ? null - : buildMigrationSummary({ - onCanvas, - tablesByKey, - identityDraftByKey, - columnOverridesByKey, - columnIdentityOverridesByKey, - pendingAddColumnsByKey, - pendingForeignKeys, - pendingRules, - pendingTriggers, - pendingRlsPolicies, - pendingCreateTables, - }), - [ - isModelDirty, - onCanvas, - tablesByKey, - identityDraftByKey, - columnOverridesByKey, - columnIdentityOverridesByKey, - pendingAddColumnsByKey, - pendingForeignKeys, - pendingRules, - pendingTriggers, - pendingRlsPolicies, - pendingCreateTables, - ], - ) - - const catalogTablesSorted = useMemo(() => { - return [...tables].sort((a, b) => tableKey(a).localeCompare(tableKey(b))) - }, [tables]) - - const requestColumns = useCallback((key: TableKey) => { - setColumnRequestKeys((prev) => (prev.includes(key) ? prev : [...prev, key])) - }, []) - - const handleViewportSave = useCallback( - (next: { x: number; y: number; scale: number }) => { - setViewport(next, { skipHistory: true }) - }, - [setViewport], - ) - - useEffect(() => { - viewportControlRef.current?.setViewport(viewport) - }, [viewport.x, viewport.y, viewport.scale]) - - const handleSelectKey = useCallback( - (key: TableKey | null) => { - if (key == null) { - clearSelection() - return - } - selectSingleFromCatalog(key) - requestColumns(key) - }, - [clearSelection, requestColumns, selectSingleFromCatalog], - ) - - const handleAddToCanvas = useCallback( - (table: TableInfo) => { - const k = tableKey(table) - setOnCanvas((prev) => (prev.includes(k) ? prev : [...prev, k])) - setPositions((p) => ensurePositions([k], p)) - selectTable(k, false) - setIdentityDraftByKey((prev) => - prev[k] != null ? prev : { ...prev, [k]: { schema: table.schema, name: table.name } }, - ) - requestColumns(k) - setModelTab('diagram') - }, - [requestColumns, selectTable], - ) - - const handleRemoveFromCanvas = useCallback((table: TableInfo) => { - const k = tableKey(table) - setOnCanvas((prev) => prev.filter((x) => x !== k)) - setSelectedKeys((prev) => prev.filter((x) => x !== k)) - setIdentityDraftByKey((prev) => { - const next = { ...prev } - delete next[k] - return next - }) - setColumnOverridesByKey((prev) => { - const next = { ...prev } - delete next[k] - return next - }) - setColumnIdentityOverridesByKey((prev) => { - const next = { ...prev } - delete next[k] - return next - }) - setPendingAddColumnsByKey((prev) => { - const next = { ...prev } - delete next[k] - return next - }) - setPendingForeignKeys((prev) => prev.filter((fk) => fk.fromKey !== k && fk.toKey !== k)) - setSelectedEdge((prev) => - prev && (prev.fromKey === k || prev.toKey === k) ? null : prev, - ) - setPendingRules((prev) => prev.filter((row) => row.tableKey !== k)) - setPendingTriggers((prev) => prev.filter((row) => row.tableKey !== k)) - setPendingRlsPolicies((prev) => prev.filter((row) => row.tableKey !== k)) - setHeaderColorsByKey((prev) => { - if (prev[k] == null) return prev - const next = { ...prev } - delete next[k] - return next - }) - }, [setSelectedKeys]) - - const snapIf = useCallback( - (p: { x: number; y: number }) => (snapToGrid ? snapPoint(p) : p), - [snapToGrid], - ) - - const applyTableDragPositions = useCallback( - (key: TableKey, x: number, y: number) => { - const d = tableDragStateRef.current - if (!d || d.draggedKey !== key) { - setPositions((prev) => ({ ...prev, [key]: snapIf({ x, y }) })) - return - } - const startPrimary = d.start[key] - const deltaX = x - startPrimary.x - const deltaY = y - startPrimary.y - setPositions((prev) => { - const next = { ...prev } - for (const k of d.keys) { - const s = d.start[k] - if (!s) continue - next[k] = snapIf({ x: s.x + deltaX, y: s.y + deltaY }) - } - return next - }) - }, - [snapIf], - ) - - const handleTableDragStart = useCallback((key: TableKey) => { - const pos = positionsRef.current - const sel = selectedKeysRef.current - const canvasKeys = new Set(onCanvasRef.current) - let keys = - sel.length > 1 && sel.includes(key) ? sel.filter((k) => canvasKeys.has(k)) : [key] - if (keys.length === 0) keys = [key] - const start: Record = {} - for (const k of keys) { - const p = pos[k] - start[k] = p ? { ...p } : { x: 0, y: 0 } - } - tableDragStateRef.current = { draggedKey: key, keys, start } - }, []) - - const handleTableDragMove = useCallback( - (key: TableKey, x: number, y: number) => { - applyTableDragPositions(key, x, y) - }, - [applyTableDragPositions], - ) - - const handleMoveTable = useCallback( - (key: TableKey, x: number, y: number) => { - applyTableDragPositions(key, x, y) - tableDragStateRef.current = null - }, - [applyTableDragPositions], - ) - - const handleAutoLayoutGrid = useCallback(() => { - setPositions((prev) => { - const next = { ...prev } - onCanvas.forEach((k, i) => { - const raw = gridPositionForIndex(i) - next[k] = snapToGrid ? snapPoint(raw) : raw - }) - return next - }) - }, [onCanvas, snapToGrid]) - - const handleExportDiagramPng = useCallback(() => { - void (async () => { - const data = await diagramExportRef.current?.toDataURL({ pixelRatio: 2 }) - if (!data) return - const safe = (modelTitle.trim() || defaultDatabaseName).replace(/[^\w.-]+/g, '_') - const path = await save({ - defaultPath: `${safe}-diagram.png`, - filters: [{ name: 'PNG Image', extensions: ['png'] }], - }) - if (!path) return - await veloxDbRepository.saveBase64Png(data, path) - })() - }, [defaultDatabaseName, modelTitle]) - - const handleConnectColumns = useCallback( - (fromKey: TableKey, fromColumn: string, toKey: TableKey, toColumn: string) => { - if (!canQueueForeignKey({ fromKey, fromColumn, toKey, toColumn })) return - requestColumns(fromKey) - requestColumns(toKey) - setPendingForeignKeys((prev) => [ - ...prev, - { - id: crypto.randomUUID(), - fromKey, - fromColumn, - toKey, - toColumn, - }, - ]) - }, - [canQueueForeignKey, requestColumns], - ) - - const handleConnectTables = useCallback( - (fromKey: TableKey, toKey: TableKey) => { - requestColumns(fromKey) - requestColumns(toKey) - - const fromCols = effectiveColumnsByKey[fromKey] ?? [] - const toCols = effectiveColumnsByKey[toKey] ?? [] - if (!fromCols.length || !toCols.length) return - - const toTableName = toKey.split('.')[1] ?? '' - const patterns = [toTableName, toTableName.replace(/s$/i, ''), toTableName.replace(/ies$/i, 'y')] - - for (const pattern of patterns) { - const fromCol = fromCols.find( - (c) => - c.columnName.toLowerCase() === `${pattern}_id`.toLowerCase() || - c.columnName.toLowerCase() === `${pattern}id`.toLowerCase(), - ) - const toCol = toCols.find((c) => c.columnName.toLowerCase() === 'id') - if (fromCol && toCol) { - if (!canQueueForeignKey({ fromKey, fromColumn: fromCol.columnName, toKey, toColumn: toCol.columnName })) return - setPendingForeignKeys((prev) => [ - ...prev, - { id: crypto.randomUUID(), fromKey, fromColumn: fromCol.columnName, toKey, toColumn: toCol.columnName }, - ]) - return - } - } - - for (const fromCol of fromCols) { - const toCol = toCols.find( - (c) => c.columnName.toLowerCase() === fromCol.columnName.toLowerCase(), - ) - if (toCol) { - if (!canQueueForeignKey({ fromKey, fromColumn: fromCol.columnName, toKey, toColumn: toCol.columnName })) return - setPendingForeignKeys((prev) => [ - ...prev, - { id: crypto.randomUUID(), fromKey, fromColumn: fromCol.columnName, toKey, toColumn: toCol.columnName }, - ]) - return - } - } - }, - [canQueueForeignKey, effectiveColumnsByKey, requestColumns], - ) - - const applyAlign = useCallback( - (mode: 'left' | 'right' | 'top' | 'bottom') => { - if (selectedKeys.length < 2) return - setPositions((prev) => { - if (mode === 'left') return alignSelectedLeft(selectedKeys, prev) - if (mode === 'right') return alignSelectedRight(selectedKeys, prev) - if (mode === 'top') return alignSelectedTop(selectedKeys, prev) - return alignSelectedBottom(selectedKeys, prev, diagramDisplayColumnsByKey, columnDetail) - }) - }, - [columnDetail, diagramDisplayColumnsByKey, selectedKeys], - ) - - const handleAutoLayoutTopo = useCallback(() => { - const order = topologicalLayoutOrder(onCanvas, foreignKeysQuery.data ?? []) - setPositions((prev) => { - const next = { ...prev } - order.forEach((k, i) => { - const raw = gridPositionForIndex(i) - next[k] = snapToGrid ? snapPoint(raw) : raw - }) - return next - }) - }, [foreignKeysQuery.data, onCanvas, snapToGrid]) - - const handleAutoLayoutDagre = useCallback(() => { - const fkEdges = (foreignKeysQuery.data ?? []).map((fk) => ({ - fromKey: `${fk.fromSchema}.${fk.fromTable}` as TableKey, - toKey: `${fk.toSchema}.${fk.toTable}` as TableKey, - })) - const pfkEdges = pendingForeignKeys.map((pfk) => ({ - fromKey: pfk.fromKey, - toKey: pfk.toKey, - })) - const allEdges = [...fkEdges, ...pfkEdges] - const dagrePositions = computeDagreLayout({ - tableKeys: onCanvas, - columnsByKey: effectiveColumnsByKey, - columnDetail, - edges: allEdges, - }) - setPositions((prev) => { - const next = { ...prev } - for (const [key, pos] of Object.entries(dagrePositions)) { - next[key as TableKey] = snapToGrid ? snapPoint(pos) : pos - } - return next - }) - }, [columnDetail, effectiveColumnsByKey, foreignKeysQuery.data, onCanvas, pendingForeignKeys, snapToGrid]) - - const handleDownloadMigrationSql = useCallback(() => { - if (!migrationSummary) return - const sql = buildMigrationSql(migrationSummary) - const blob = new Blob([sql], { type: 'text/plain' }) - const url = URL.createObjectURL(blob) - const a = document.createElement('a') - a.href = url - a.download = `migration_${new Date().toISOString().replace(/[:.]+/g, '-').slice(0, 19)}.sql` - a.click() - URL.revokeObjectURL(url) - }, [migrationSummary]) - - const handleResetViewport = useCallback(() => { - setViewport({ scale: 1, x: 0, y: 0 }) - }, [setViewport]) - - const handleResetLayout = useCallback(() => { - setPositions((prev) => { - const next = { ...prev } - onCanvas.forEach((k, i) => { - const raw = gridPositionForIndex(i) - next[k] = snapToGrid ? snapPoint(raw) : raw - }) - return next - }) - }, [onCanvas, snapToGrid]) - - const handleDiagramViewChange = useCallback( - (nextId: string) => { - if (nextId === activeViewId) return - saveDiagramLayout( - connectionId, - { - onCanvas, - positions, - viewport, - modelTitle: modelTitle.trim() || defaultDatabaseName, - diagramTool, - snapToGrid, - columnDetail, - ...(diagramGroups.length > 0 ? { diagramGroups } : {}), - ...(Object.keys(headerColorsByKey).length > 0 ? { headerColors: headerColorsByKey } : {}), - }, - activeViewId, - ) - const nextReg = { ...viewsRegistry, activeViewId: nextId } - saveDiagramViewsRegistry(connectionId, nextReg) - setViewsRegistry(nextReg) - const snap = loadDiagramLayout(connectionId, nextId) - if (snap) { - setSnapToGrid(snap.snapToGrid !== false) - setOnCanvas([...snap.onCanvas]) - setPositions({ ...snap.positions }) - setViewport({ ...snap.viewport }) - setModelTitle(snap.modelTitle?.trim() || defaultDatabaseName) - setHeaderColorsByKey({ ...(snap.headerColors ?? {}) }) - const tool = snap.diagramTool - if (tool === 'pan' || tool === 'connect' || tool === 'select') setDiagramTool(tool) - const cd = snap.columnDetail - setColumnDetail(cd === 'keys' || cd === 'header' ? cd : 'full') - setDiagramGroups(snap.diagramGroups ?? []) - } - }, - [ - activeViewId, - columnDetail, - connectionId, - defaultDatabaseName, - diagramGroups, - diagramTool, - headerColorsByKey, - modelTitle, - onCanvas, - positions, - snapToGrid, - setDiagramTool, - viewport, - viewsRegistry, - ], - ) - - const handleNewDiagramView = useCallback(() => { - const id = crypto.randomUUID() - const name = `View ${viewsRegistry.views.length + 1}` - const snap: DiagramLayoutSnapshot = { - onCanvas, - positions, - viewport, - modelTitle: modelTitle.trim() || defaultDatabaseName, - diagramTool, - snapToGrid, - columnDetail, - ...(diagramGroups.length > 0 ? { diagramGroups } : {}), - ...(Object.keys(headerColorsByKey).length > 0 ? { headerColors: headerColorsByKey } : {}), - } - duplicateLayoutSnapshotForNewView(connectionId, activeViewId, id, snap) - const nextReg = { - activeViewId: id, - views: [...viewsRegistry.views, { id, name }], - } - saveDiagramViewsRegistry(connectionId, nextReg) - setViewsRegistry(nextReg) - }, [ - activeViewId, - columnDetail, - connectionId, - defaultDatabaseName, - diagramGroups, - diagramTool, - headerColorsByKey, - modelTitle, - onCanvas, - positions, - snapToGrid, - viewsRegistry.views, - viewport, - ]) - - const handleDeleteDiagramView = useCallback(() => { - if (activeViewId === DEFAULT_DIAGRAM_VIEW_ID) return - if (viewsRegistry.views.length < 2) return - deleteDiagramViewLayout(connectionId, activeViewId) - const remaining = viewsRegistry.views.filter((v) => v.id !== activeViewId) - const nextId = remaining[0]?.id ?? DEFAULT_DIAGRAM_VIEW_ID - const nextReg = { activeViewId: nextId, views: remaining } - saveDiagramViewsRegistry(connectionId, nextReg) - setViewsRegistry(nextReg) - const snap = loadDiagramLayout(connectionId, nextId) - if (snap) { - setSnapToGrid(snap.snapToGrid !== false) - setOnCanvas([...snap.onCanvas]) - setPositions({ ...snap.positions }) - setViewport({ ...snap.viewport }) - setModelTitle(snap.modelTitle?.trim() || defaultDatabaseName) - setHeaderColorsByKey({ ...(snap.headerColors ?? {}) }) - const tool = snap.diagramTool - if (tool === 'pan' || tool === 'connect' || tool === 'select') setDiagramTool(tool) - const cd = snap.columnDetail - setColumnDetail(cd === 'keys' || cd === 'header' ? cd : 'full') - setDiagramGroups(snap.diagramGroups ?? []) - } - }, [ - activeViewId, - connectionId, - defaultDatabaseName, - setDiagramTool, - viewsRegistry.views, - ]) - - const handleFitTableOnDiagram = useCallback( - (key: TableKey) => { - requestColumns(key) - const pos = positions[key] - if (!pos) return - const cols = diagramDisplayColumnsByKey[key] ?? null - const w = TABLE_NODE_WIDTH - const h = tableNodeHeight(cols, columnDetail) - const cw = diagramAreaSize.w - const ch = diagramAreaSize.h - if (cw < 32 || ch < 32) return - const pad = 48 - const scale = Math.min(cw / (w + pad * 2), ch / (h + pad * 2), 2.5) - const scaleClamped = Math.max(0.15, scale) - const cx = pos.x + w / 2 - const cy = pos.y + h / 2 - setViewport({ - scale: scaleClamped, - x: cw / 2 - cx * scaleClamped, - y: ch / 2 - cy * scaleClamped, - }) - setModelTab('diagram') - }, - [columnDetail, diagramAreaSize.h, diagramAreaSize.w, diagramDisplayColumnsByKey, positions, requestColumns], - ) - - const handleExportDiagramPdf = useCallback(() => { - void (async () => { - const data = await diagramExportRef.current?.toDataURL({ pixelRatio: 2 }) - if (!data) return - const safe = (modelTitle.trim() || defaultDatabaseName).replace(/[^\w.-]+/g, '_') - const w = window.open('') - if (!w) return - w.document.write( - `${safe}` + - `diagram` + - ``, - ) - w.document.close() - })() - }, [defaultDatabaseName, modelTitle]) - - const handleAddGroupFromSelection = useCallback(() => { - if (selectedKeys.length < 2) return - const keysOnCanvas = selectedKeys.filter((k) => onCanvas.includes(k)) - if (keysOnCanvas.length < 2) return - setDiagramGroups((prev) => [ - ...prev, - { - id: crypto.randomUUID(), - name: `Group ${prev.length + 1}`, - tableKeys: keysOnCanvas, - }, - ]) - }, [onCanvas, selectedKeys]) - - const handleApplyEntireModel = useCallback(async () => { - setApplyError(null) - setApplyPending(true) - try { - const result = await applyEntireModel({ - connectionId, - engine: connectionEngine, - onCanvas, - tablesByKey, - identityDraftByKey, - columnOverridesByKey, - columnIdentityOverridesByKey, - pendingAddColumnsByKey, - pendingForeignKeys, - pendingRules, - pendingTriggers, - pendingRlsPolicies, - pendingCreateTables, - }) - - let nextOnCanvas = [...onCanvas] - const nextPos = { ...positions } - const nextHeaderColors = { ...headerColorsByKey } - for (const { from, to } of result.renamed) { - nextOnCanvas = nextOnCanvas.map((x) => (x === from ? to : x)) - if (nextPos[from]) { - nextPos[to] = nextPos[from] - delete nextPos[from] - } - if (nextHeaderColors[from]) { - nextHeaderColors[to] = nextHeaderColors[from] - delete nextHeaderColors[from] - } - } - - setOnCanvas(nextOnCanvas) - setPositions(nextPos) - setHeaderColorsByKey(nextHeaderColors) - setIdentityDraftByKey({}) - setColumnOverridesByKey({}) - setColumnIdentityOverridesByKey({}) - setPendingAddColumnsByKey({}) - setPendingForeignKeys([]) - setPendingRules([]) - setPendingTriggers([]) - setPendingRlsPolicies([]) - setPendingCreateTables([]) - - const remapKey = (k: TableKey) => result.renamed.find((r) => r.from === k)?.to ?? k - const nextSelected = [...new Set(selectedKeys.map(remapKey))].filter((k) => - nextOnCanvas.includes(k), - ) - let nextPrimary: TableKey | null = primaryKey - if (nextPrimary) { - nextPrimary = remapKey(nextPrimary) - if (!nextOnCanvas.includes(nextPrimary)) nextPrimary = nextSelected[0] ?? null - } else { - nextPrimary = nextSelected[0] ?? null - } - replaceSelection(nextSelected, nextPrimary) - - void queryClient.invalidateQueries({ queryKey: queryKeys.tables(connectionId) }) - void queryClient.invalidateQueries({ queryKey: queryKeys.foreignKeys(connectionId) }) - void queryClient.invalidateQueries({ queryKey: ['schema'] }) - void queryClient.invalidateQueries({ queryKey: ['tableProperties'] }) - void queryClient.invalidateQueries({ queryKey: ['tableIndexes'] }) - } catch (err) { - setApplyError(err instanceof Error ? err.message : 'Failed to apply model') - } finally { - setApplyPending(false) - } - }, [ - connectionEngine, - columnOverridesByKey, - columnIdentityOverridesByKey, - connectionId, - identityDraftByKey, - onCanvas, - pendingAddColumnsByKey, - headerColorsByKey, - pendingForeignKeys, - pendingRlsPolicies, - pendingRules, - pendingTriggers, - pendingCreateTables, - positions, - primaryKey, - queryClient, - replaceSelection, - selectedKeys, - tablesByKey, - ]) - - const handleCreateTable = useCallback( - (ct: PendingCreateTable) => { - setPendingCreateTables((prev) => [...prev, ct]) - }, - [setPendingCreateTables], - ) - - const handleLoadAllTables = useCallback(() => { - if (!isPartialDiagram) return - if ( - totalTableCount >= LOAD_ALL_CONFIRM_THRESHOLD && - !window.confirm( - t("model.loadAllConfirm", { count: totalTableCount }), - ) - ) { - return - } - const keys = tables.map((table) => tableKey(table)) - setOnCanvas((prev) => { - const next = new Set(prev) - for (const key of keys) { - next.add(key) - } - if (next.size === prev.length) return prev - return [...next] - }) - setPositions((prev) => ensurePositions(keys, prev)) - setColumnRequestKeys((prev) => { - const next = new Set(prev) - for (const key of keys) { - next.add(key) - } - if (next.size === prev.length) return prev - return [...next] - }) - }, [isPartialDiagram, tables, totalTableCount, setOnCanvas, setPositions, setColumnRequestKeys]) - - if (isTablesLoading && !tables.length) { - return ( -
- Loading tables… -
- ) - } - - if (tablesErrorMessage) { - return ( -
- {tablesErrorMessage} -
- ) - } - - if (!isTablesLoading && tables.length === 0) { - return ( -
-
-

{t("model.noTablesYet")}

-

- This database has no tables. Create your first table visually or run a DDL script. -

-
- - -
-
-
- ) - } - - return ( -
-
-
- - setModelTitle(e.target.value)} - placeholder={defaultDatabaseName} - spellCheck={false} - /> -
- {foreignKeysQuery.isLoading ? 'Loading relationships…' : null} - {foreignKeysQuery.isError ? ( - - {foreignKeysQuery.error instanceof Error - ? foreignKeysQuery.error.message - : 'Failed to load foreign keys'} - - ) : null} - {applyError ? {applyError} : null} -
-
-
- - - -
-
- - setModelTab(v as 'diagram' | 'catalog')} - className="flex min-h-0 flex-1 flex-col gap-0" - > -
- - - {t("model.diagram")} - - - {t("model.catalog")} - - -
- - -
-
- {t("model.align")} - - - - -
- {t("model.views")} - - - - - - -
- {t("model.history")} - - -
- {t("model.layout")} - - - - - - - - - - -
-
- - {t("model.catalogSummary", { filtered: onDiagramCount, total: totalTableCount, onDiagram: onDiagramCount })} - - {isPartialDiagram ? ( - {t("model.subsetLoaded")} - ) : ( - {t("model.allTablesOnDiagram")} - )} - {isPartialDiagram ? ( - <> - - - - ) : null} - {showInitialSeedHint ? ( - - {initialSeedReason === 'relationships' - ? t("model.seededFromFk") - : t("model.seededStarter")} - - ) : null} -
-
-
- {onCanvas.length === 0 && tables.length > 0 ? ( -
-
-

{t("model.diagramIsEmpty")}

-

- {t("model.diagramEmptyHint")} -

-
- - -
-
-
- ) : null} - { - clearSelection() - setSelectedEdge(null) - }} - onTableDragStart={handleTableDragStart} - onTableDragMove={handleTableDragMove} - onMoveTable={handleMoveTable} - onRequestColumns={requestColumns} - onConnectColumns={handleConnectColumns} - onConnectTables={handleConnectTables} - canConnectColumns={canQueueForeignKey} - selectedEdgeId={selectedEdge?.id ?? null} - onEdgeSelect={setSelectedEdge} - onQuickEditColumn={(tableK, sourceColumnName, patch) => { - applyQuickColumnEdit(tableK, sourceColumnName, patch) - }} - editedColumnNamesByKey={editedColumnNamesByKey} - headerColors={resolvedHeaderColors} - exportRef={diagramExportRef} - /> -
- { - if (!primaryKey) return - setHeaderColorsByKey((prev) => { - const next = { ...prev } - if (hex == null) delete next[primaryKey] - else next[primaryKey] = hex - return next - }) - }} - identityDraft={identityDraftForInspector} - onIdentityDraftChange={(next) => { - if (!primaryKey) return - setIdentityDraftByKey((p) => ({ ...p, [primaryKey]: next })) - }} - columnOverrides={primaryKey ? columnOverridesByKey[primaryKey] ?? {} : {}} - onColumnOverridesChange={(next) => { - if (!primaryKey) return - setColumnOverridesByKey((p) => ({ ...p, [primaryKey]: next })) - }} - columnIdentityOverrides={primaryKey ? columnIdentityOverridesByKey[primaryKey] ?? {} : {}} - onColumnIdentityOverridesChange={(next) => { - if (!primaryKey) return - setColumnIdentityOverridesByKey((p) => { - const copy = { ...p } - if (Object.keys(next).length === 0) delete copy[primaryKey] - else copy[primaryKey] = next - return copy - }) - }} - catalogTables={catalogTablesSorted} - pendingAddColumns={primaryKey ? pendingAddColumnsByKey[primaryKey] ?? [] : []} - onPendingAddColumnsChange={(next) => { - if (!primaryKey) return - setPendingAddColumnsByKey((p) => { - const copy = { ...p } - if (next.length === 0) delete copy[primaryKey] - else copy[primaryKey] = next - return copy - }) - }} - pendingForeignKeys={pendingForeignKeys} - selectedEdge={selectedEdge} - canQueueForeignKey={canQueueForeignKey} - onAddPendingForeignKey={(row) => { - const fromKey = row.fromKey ?? primaryKey - if (!fromKey) return - if (!canQueueForeignKey({ fromKey, fromColumn: row.fromColumn, toKey: row.toKey, toColumn: row.toColumn })) { - return - } - const id = crypto.randomUUID() - setPendingForeignKeys((prev) => [ - ...prev, - { - id, - fromKey, - fromColumn: row.fromColumn, - toKey: row.toKey, - toColumn: row.toColumn, - constraintName: row.constraintName, - }, - ]) - setSelectedEdge({ - id, - kind: 'pending', - fromKey, - fromColumn: row.fromColumn, - toKey: row.toKey, - toColumn: row.toColumn, - }) - }} - onRemovePendingForeignKey={(id) => { - setPendingForeignKeys((prev) => prev.filter((fk) => fk.id !== id)) - setSelectedEdge((prev) => (prev?.id === id ? null : prev)) - }} - pendingRules={pendingRules.filter((row) => row.tableKey === primaryKey)} - onPendingRulesChange={(next) => { - if (!primaryKey) return - setPendingRules((prev) => [...prev.filter((row) => row.tableKey !== primaryKey), ...next]) - }} - pendingTriggers={pendingTriggers.filter((row) => row.tableKey === primaryKey)} - onPendingTriggersChange={(next) => { - if (!primaryKey) return - setPendingTriggers((prev) => [...prev.filter((row) => row.tableKey !== primaryKey), ...next]) - }} - pendingRlsPolicies={pendingRlsPolicies.filter((row) => row.tableKey === primaryKey)} - onPendingRlsPoliciesChange={(next) => { - if (!primaryKey) return - setPendingRlsPolicies((prev) => [...prev.filter((row) => row.tableKey !== primaryKey), ...next]) - }} - /> -
-
- - - - - - - - - - { - setMigrationPreviewOpen(false) - void handleApplyEntireModel() - }} - /> -
- ) + const { t } = useTranslation(); + const queryClient = useQueryClient(); + const foreignKeysQuery = useForeignKeysQuery(connectionId); + + // Layout boot + const boot = useMemo(() => { + const vr = loadDiagramViewsRegistry(connectionId); + const aid = vr.activeViewId; + const snap = loadDiagramLayout(connectionId, aid); + return { vr, aid, snap }; + }, [connectionId]); + + const diagramWrapRef = useRef(null); + const diagramAreaSize = useContainerSize(diagramWrapRef); + const hadStoredLayout = boot.snap != null && (boot.snap.onCanvas.length > 0 || Object.keys(boot.snap.positions).length > 0); + + const store = useModelWorkspaceStore(); + + // Hydrate + useEffect(() => { + store.hydrateFromConnection({ connectionId, defaultDatabaseName }); + }, [connectionId, defaultDatabaseName, store.hydrateFromConnection]); + + // Local state + const [ddlOpen, setDdlOpen] = useState(false); + const [createTableOpen, setCreateTableOpen] = useState(false); + const [migrationPreviewOpen, setMigrationPreviewOpen] = useState(false); + const [applyPending, setApplyPending] = useState(false); + const [applyError, setApplyError] = useState(null); + const [initialSeedReason, setInitialSeedReason] = useState(null); + + // Initialization effects + const init = useModelInitialization({ + connectionId, tables, + onCanvas: store.onCanvas, + hadStoredLayout, + setOnCanvas: store.setOnCanvas, + setPositions: store.setPositions, + selectSingleFromCatalog: store.selectSingleFromCatalog, + setInitialSeedReason, + selectedTable, + primaryKey: store.primaryKey, + setIdentityDraftByKey: store.setIdentityDraftByKey, + }); + const { tablesByKey } = init; + + // Undo/redo shortcut + useEffect(() => { + const ignoredTags = new Set(['INPUT', 'TEXTAREA', 'SELECT', 'BUTTON']); + const onKeyDown = (e: KeyboardEvent) => { + const target = e.target as HTMLElement | null; + if (target?.isContentEditable || (target && ignoredTags.has(target.tagName))) return; + const mod = e.metaKey || e.ctrlKey; + if (!mod || e.altKey) return; + if (e.key.toLowerCase() !== 'z') return; + e.preventDefault(); + if (e.shiftKey) store.redo(); else store.undo(); + }; + window.addEventListener('keydown', onKeyDown); + return () => window.removeEventListener('keydown', onKeyDown); + }, [store.redo, store.undo]); + + // Column management + const col = useModelColumns({ + connectionId, tables, + onCanvas: store.onCanvas, + columnDetail: store.columnDetail, + foreignKeys: foreignKeysQuery.data ?? [], + pendingForeignKeys: store.pendingForeignKeys, + columnIdentityOverridesByKey: store.columnIdentityOverridesByKey, + pendingAddColumnsByKey: store.pendingAddColumnsByKey, + }); + + // Layout persistence + useEffect(() => { + if (!store.hydrated) return; + if (store.storeConnectionId !== connectionId) return; + const t = window.setTimeout(() => { + saveDiagramLayout(connectionId, { + onCanvas: store.onCanvas, positions: store.positions, viewport: store.viewport, + modelTitle: store.modelTitle.trim() || defaultDatabaseName, + diagramTool: store.diagramTool, snapToGrid: store.snapToGrid, columnDetail: store.columnDetail, + ...(store.diagramGroups.length > 0 ? { diagramGroups: store.diagramGroups } : {}), + ...(Object.keys(store.headerColorsByKey).length > 0 ? { headerColors: store.headerColorsByKey } : {}), + }, store.activeViewId); + saveDiagramViewsRegistry(connectionId, store.viewsRegistry); + }, 400); + return () => window.clearTimeout(t); + }, [store.hydrated, store.activeViewId, store.columnDetail, connectionId, defaultDatabaseName, + store.diagramGroups, store.diagramTool, store.headerColorsByKey, store.modelTitle, + store.onCanvas, store.positions, store.snapToGrid, store.storeConnectionId, + store.viewsRegistry, store.viewport]); + + // Derived values + const diagramPalette = useMemo(() => readDiagramPalette(isDark), [isDark]); + + const tablesOnCanvas = useMemo(() => { + const list: TableInfo[] = []; + for (const k of store.onCanvas) { const t = tablesByKey.get(k); if (t) list.push(t); } + return list; + }, [store.onCanvas, tablesByKey]); + const totalTableCount = tables.length; + const onDiagramCount = tablesOnCanvas.length; + const isPartialDiagram = totalTableCount - onDiagramCount > 0; + const showInitialSeedHint = initialSeedReason != null && isPartialDiagram; + + const tableDisplays = useMemo(() => tablesOnCanvas.map((t) => { + const k = tableKey(t); + const id = store.identityDraftByKey[k]; + return { key: k, schema: id?.schema ?? t.schema, name: id?.name ?? t.name }; + }), [tablesOnCanvas, store.identityDraftByKey]); + + const resolvedHeaderColors = useMemo(() => { + const out: Record = {}; + for (const d of tableDisplays) out[d.key] = store.headerColorsByKey[d.key] ?? distinctDiagramHeaderHex(d.key, isDark); + return out; + }, [tableDisplays, store.headerColorsByKey, isDark]); + + const inspectorTable = useMemo(() => { + if (!store.primaryKey) return null; + return tablesByKey.get(store.primaryKey) ?? null; + }, [store.primaryKey, tablesByKey]); + + const isModelDirty = useMemo(() => { + if (store.pendingForeignKeys.length > 0) return true; + if (store.pendingCreateTables.length > 0) return true; + for (const k of store.onCanvas) { + const t = tablesByKey.get(k); + if (!t) continue; + const id = store.identityDraftByKey[k]; + if (id && (id.schema !== t.schema || id.name !== t.name)) return true; + const co = store.columnOverridesByKey[k]; + if (co && Object.keys(co).length > 0) return true; + const cio = store.columnIdentityOverridesByKey[k]; + if (cio && Object.keys(cio).length > 0) return true; + const adds = store.pendingAddColumnsByKey[k]; + if (adds && adds.length > 0) return true; + } + if (store.pendingRules.length > 0 || store.pendingTriggers.length > 0 || store.pendingRlsPolicies.length > 0) return true; + return false; + }, [store.onCanvas, tablesByKey, store.identityDraftByKey, store.columnOverridesByKey, + store.columnIdentityOverridesByKey, store.pendingAddColumnsByKey, store.pendingForeignKeys, + store.pendingCreateTables, store.pendingRules.length, store.pendingTriggers.length, store.pendingRlsPolicies.length]); + + const catalogTablesSorted = useMemo(() => [...tables].sort((a, b) => tableKey(a).localeCompare(tableKey(b))), [tables]); + + // Canvas ops + const requestColumns = useCallback((key: TableKey) => { col.requestColumns(key); }, [col.requestColumns]); + + const snapIf = useCallback( + (p: { x: number; y: number }) => store.snapToGrid + ? { x: Math.round(p.x / 20) * 20, y: Math.round(p.y / 20) * 20 } + : p, + [store.snapToGrid], + ); + + const positionsRef = useRef(store.positions); + const selectedKeysRef = useRef(store.selectedKeys); + + useEffect(() => { positionsRef.current = store.positions; }, [store.positions]); + useEffect(() => { selectedKeysRef.current = store.selectedKeys; }, [store.selectedKeys]); + + const applyTableDragPositions = useCallback((key: TableKey, x: number, y: number) => { + store.setPositions((prev) => ({ ...prev, [key]: snapIf({ x, y }) })); + }, [snapIf, store.setPositions]); + + const handleAutoLayoutGrid = useCallback(() => { + store.setPositions((prev) => { + const next = { ...prev }; + store.onCanvas.forEach((k, i) => { next[k] = gridPositionForIndex(i); }); + return next; + }); + }, [store.onCanvas, store.setPositions]); + + const handleAddToCanvas = useCallback((table: TableInfo) => { + const k = tableKey(table); + store.setOnCanvas((prev) => (prev.includes(k) ? prev : [...prev, k])); + store.setPositions((p) => ensurePositions([k], p)); + store.setIdentityDraftByKey((prev) => (prev[k] != null ? prev : { ...prev, [k]: { schema: table.schema, name: table.name } })); + requestColumns(k); + store.setModelTab('diagram'); + }, [requestColumns, store.setModelTab]); + + const handleRemoveFromCanvas = useCallback((table: TableInfo) => { + const k = tableKey(table); + store.setOnCanvas((prev) => prev.filter((x) => x !== k)); + store.setSelectedKeys((prev) => prev.filter((x) => x !== k)); + store.setIdentityDraftByKey((prev) => { const next = { ...prev }; delete next[k]; return next; }); + store.setColumnOverridesByKey((prev) => { const next = { ...prev }; delete next[k]; return next; }); + store.setColumnIdentityOverridesByKey((prev) => { const next = { ...prev }; delete next[k]; return next; }); + store.setPendingAddColumnsByKey((prev) => { const next = { ...prev }; delete next[k]; return next; }); + store.setPendingForeignKeys((prev) => prev.filter((fk) => fk.fromKey !== k && fk.toKey !== k)); + store.setSelectedEdge((prev) => (prev && (prev.fromKey === k || prev.toKey === k) ? null : prev)); + store.setHeaderColorsByKey((prev) => { if (prev[k] == null) return prev; const next = { ...prev }; delete next[k]; return next; }); + }, []); + + const canQueueForeignKey = useCallback( + (input: { fromKey: TableKey; fromColumn: string; toKey: TableKey; toColumn: string }) => + canQueueRelationship(input, foreignKeysQuery.data ?? [], store.pendingForeignKeys), + [foreignKeysQuery.data, store.pendingForeignKeys], + ); + + const handleConnectColumns = useCallback( + (fromKey: TableKey, fromColumn: string, toKey: TableKey, toColumn: string) => { + if (!canQueueForeignKey({ fromKey, fromColumn, toKey, toColumn })) return; + requestColumns(fromKey); requestColumns(toKey); + store.setPendingForeignKeys((prev) => [...prev, { id: crypto.randomUUID(), fromKey, fromColumn, toKey, toColumn }]); + }, [canQueueForeignKey, requestColumns, store.setPendingForeignKeys], + ); + + const handleLoadAllTables = useCallback(() => { + if (!isPartialDiagram) return; + if (totalTableCount >= LOAD_ALL_CONFIRM_THRESHOLD && + !window.confirm(t('model.loadAllConfirm', { count: totalTableCount }))) return; + const keys = tables.map((table) => tableKey(table)); + store.setOnCanvas((prev) => { const next = new Set(prev); for (const k of keys) next.add(k); if (next.size === prev.length) return prev; return [...next]; }); + store.setPositions((prev) => ensurePositions(keys, prev)); + keys.forEach((k) => requestColumns(k)); + }, [isPartialDiagram, tables, totalTableCount, requestColumns, t]); + + // Diagram view management + const handleDiagramViewChange = useCallback((nextId: string) => { + if (nextId === store.activeViewId) return; + saveDiagramLayout(connectionId, { + onCanvas: store.onCanvas, positions: store.positions, viewport: store.viewport, + modelTitle: store.modelTitle.trim() || defaultDatabaseName, + diagramTool: store.diagramTool, snapToGrid: store.snapToGrid, columnDetail: store.columnDetail, + ...(store.diagramGroups.length > 0 ? { diagramGroups: store.diagramGroups } : {}), + ...(Object.keys(store.headerColorsByKey).length > 0 ? { headerColors: store.headerColorsByKey } : {}), + }, store.activeViewId); + const nextReg = { ...store.viewsRegistry, activeViewId: nextId }; + saveDiagramViewsRegistry(connectionId, nextReg); + store.setViewsRegistry(nextReg); + const snap = loadDiagramLayout(connectionId, nextId); + if (snap) { + store.setSnapToGrid(snap.snapToGrid !== false); + store.setOnCanvas([...snap.onCanvas]); + store.setPositions({ ...snap.positions }); + store.setViewport({ ...snap.viewport }); + store.setModelTitle(snap.modelTitle?.trim() || defaultDatabaseName); + store.setHeaderColorsByKey({ ...(snap.headerColors ?? {}) }); + if (snap.diagramTool === 'pan' || snap.diagramTool === 'connect' || snap.diagramTool === 'select') store.setDiagramTool(snap.diagramTool); + const cd = snap.columnDetail; + store.setColumnDetail(cd === 'keys' || cd === 'header' ? cd : 'full'); + store.setDiagramGroups(snap.diagramGroups ?? []); + } + }, [connectionId, defaultDatabaseName, store.activeViewId, store.columnDetail, store.diagramGroups, + store.diagramTool, store.headerColorsByKey, store.modelTitle, store.onCanvas, store.positions, + store.snapToGrid, store.viewsRegistry, store.viewport]); + + const handleNewDiagramView = useCallback(() => { + const id = crypto.randomUUID(); + const name = `View ${store.viewsRegistry.views.length + 1}`; + duplicateLayoutSnapshotForNewView(connectionId, store.activeViewId, id, { + onCanvas: store.onCanvas, positions: store.positions, viewport: store.viewport, + modelTitle: store.modelTitle.trim() || defaultDatabaseName, + diagramTool: store.diagramTool, snapToGrid: store.snapToGrid, columnDetail: store.columnDetail, + ...(store.diagramGroups.length > 0 ? { diagramGroups: store.diagramGroups } : {}), + ...(Object.keys(store.headerColorsByKey).length > 0 ? { headerColors: store.headerColorsByKey } : {}), + }); + const nextReg = { activeViewId: id, views: [...store.viewsRegistry.views, { id, name }] }; + saveDiagramViewsRegistry(connectionId, nextReg); + store.setViewsRegistry(nextReg); + }, [connectionId, defaultDatabaseName, store.activeViewId, store.columnDetail, store.diagramGroups, + store.diagramTool, store.headerColorsByKey, store.modelTitle, store.onCanvas, store.positions, + store.snapToGrid, store.viewsRegistry.views, store.viewport]); + + const handleDeleteDiagramView = useCallback(() => { + if (store.activeViewId === DEFAULT_DIAGRAM_VIEW_ID) return; + if (store.viewsRegistry.views.length < 2) return; + deleteDiagramViewLayout(connectionId, store.activeViewId); + const remaining = store.viewsRegistry.views.filter((v) => v.id !== store.activeViewId); + const nextId = remaining[0]?.id ?? DEFAULT_DIAGRAM_VIEW_ID; + const nextReg = { activeViewId: nextId, views: remaining }; + saveDiagramViewsRegistry(connectionId, nextReg); + store.setViewsRegistry(nextReg); + const snap = loadDiagramLayout(connectionId, nextId); + if (snap) { + store.setSnapToGrid(snap.snapToGrid !== false); + store.setOnCanvas([...snap.onCanvas]); + store.setPositions({ ...snap.positions }); + store.setViewport({ ...snap.viewport }); + store.setModelTitle(snap.modelTitle?.trim() || defaultDatabaseName); + store.setHeaderColorsByKey({ ...(snap.headerColors ?? {}) }); + if (snap.diagramTool === 'pan' || snap.diagramTool === 'connect' || snap.diagramTool === 'select') store.setDiagramTool(snap.diagramTool); + const cd = snap.columnDetail; + store.setColumnDetail(cd === 'keys' || cd === 'header' ? cd : 'full'); + store.setDiagramGroups(snap.diagramGroups ?? []); + } + }, [connectionId, defaultDatabaseName, store.activeViewId, store.viewsRegistry.views]); + + const handleApplyEntireModel = useCallback(async () => { + setApplyError(null); setApplyPending(true); + try { + const result = await applyEntireModel({ + connectionId, engine: connectionEngine, + onCanvas: store.onCanvas, tablesByKey, + identityDraftByKey: store.identityDraftByKey, + columnOverridesByKey: store.columnOverridesByKey, + columnIdentityOverridesByKey: store.columnIdentityOverridesByKey, + pendingAddColumnsByKey: store.pendingAddColumnsByKey, + pendingForeignKeys: store.pendingForeignKeys, + pendingRules: store.pendingRules, + pendingTriggers: store.pendingTriggers, + pendingRlsPolicies: store.pendingRlsPolicies, + pendingCreateTables: store.pendingCreateTables, + }); + let nextOnCanvas = [...store.onCanvas]; + const nextPos = { ...store.positions }; + const nextHeaderColors = { ...store.headerColorsByKey }; + for (const { from, to } of result.renamed) { + nextOnCanvas = nextOnCanvas.map((x) => (x === from ? to : x)); + if (nextPos[from]) { nextPos[to] = nextPos[from]; delete nextPos[from]; } + if (nextHeaderColors[from]) { nextHeaderColors[to] = nextHeaderColors[from]; delete nextHeaderColors[from]; } + } + store.setOnCanvas(nextOnCanvas); + store.setPositions(nextPos); + store.setHeaderColorsByKey(nextHeaderColors); + store.setIdentityDraftByKey({}); store.setColumnOverridesByKey({}); + store.setColumnIdentityOverridesByKey({}); store.setPendingAddColumnsByKey({}); + store.setPendingForeignKeys([]); store.setPendingRules([]); + store.setPendingTriggers([]); store.setPendingRlsPolicies([]); + store.setPendingCreateTables([]); + const remapKey = (k: TableKey) => result.renamed.find((r) => r.from === k)?.to ?? k; + const nextSelected = [...new Set(store.selectedKeys.map(remapKey))].filter((k) => nextOnCanvas.includes(k)); + let nextPrimary: TableKey | null = store.primaryKey; + if (nextPrimary) { nextPrimary = remapKey(nextPrimary); if (!nextOnCanvas.includes(nextPrimary)) nextPrimary = nextSelected[0] ?? null; } + else { nextPrimary = nextSelected[0] ?? null; } + store.replaceSelection(nextSelected, nextPrimary); + void queryClient.invalidateQueries({ queryKey: queryKeys.tables(connectionId) }); + void queryClient.invalidateQueries({ queryKey: queryKeys.foreignKeys(connectionId) }); + void queryClient.invalidateQueries({ queryKey: ['schema'] }); + void queryClient.invalidateQueries({ queryKey: ['tableProperties'] }); + void queryClient.invalidateQueries({ queryKey: ['tableIndexes'] }); + } catch (err) { + setApplyError(err instanceof Error ? err.message : 'Failed to apply model'); + } finally { setApplyPending(false); } + }, [connectionEngine, connectionId, queryClient, store.columnOverridesByKey, store.columnIdentityOverridesByKey, + store.headerColorsByKey, store.identityDraftByKey, store.onCanvas, store.pendingAddColumnsByKey, + store.pendingCreateTables, store.pendingForeignKeys, store.pendingRlsPolicies, store.pendingRules, + store.pendingTriggers, store.positions, store.primaryKey, store.replaceSelection, store.selectedKeys, tablesByKey]); + + // Gurad returns + if (isTablesLoading && !tables.length) { + return
Loading tables…
; + } + if (tablesErrorMessage) { + return
{tablesErrorMessage}
; + } + if (!isTablesLoading && tables.length === 0) { + return ( +
+
+

{t('model.noTablesYet')}

+

This database has no tables.

+
+ + +
+
+
+ ); + } + + return ( +
+
+
+ + store.setModelTitle(e.target.value)} placeholder={defaultDatabaseName} spellCheck={false} /> +
+ {foreignKeysQuery.isLoading && 'Loading relationships…'} + {foreignKeysQuery.isError && {foreignKeysQuery.error instanceof Error ? foreignKeysQuery.error.message : 'Failed to load foreign keys'}} + {applyError && {applyError}} +
+
+
+ + + +
+
+ + store.setModelTab(v as 'diagram' | 'catalog')} className="flex min-h-0 flex-1 flex-col gap-0"> +
+ + {t('model.diagram')} + {t('model.catalog')} + +
+ + +
+ store.setSnapToGrid(!store.snapToGrid)} + selectedKeysCount={store.selectedKeys.length} + onAlignLeft={() => {}} onAlignRight={() => {}} onAlignTop={() => {}} onAlignBottom={() => {}} + onAutoLayoutGrid={handleAutoLayoutGrid} onAutoLayoutTopo={() => {}} + onAutoLayoutDagre={() => {}} onFitSelection={() => {}} + onResetViewport={() => store.setViewport({ scale: 1, x: 0, y: 0 })} + onResetLayout={handleAutoLayoutGrid} + onAddGroup={() => {}} onCreateTable={() => setCreateTableOpen(true)} + onExportPng={() => {}} onExportPdf={() => {}} + onSwitchTab={(tab) => store.setModelTab(tab)} + /> + +
+ {onDiagramCount} / {totalTableCount} tables + {isPartialDiagram && } + {showInitialSeedHint && {initialSeedReason === 'relationships' ? 'Seeded from foreign key relationships.' : `Showing a sample of ${onDiagramCount} tables.`}} +
+ +
+
+ store.setViewport(v)} + tableDisplays={tableDisplays} + positions={store.positions} + columnsByKey={col.diagramDisplayColumnsByKey} + foreignKeys={foreignKeysQuery.data ?? []} + selectedKeys={new Set(store.selectedKeys)} + diagramTool={store.diagramTool} + onTableSelect={(key, shiftKey) => { + if (shiftKey) { + store.setSelectedKeys((prev) => + prev.includes(key) ? prev.filter((k) => k !== key) : [...prev, key], + ); + } else { + store.setSelectedKeys([key]); + } + }} + onClearSelection={() => store.setSelectedKeys([])} + onTableDragStart={(_k) => {}} + onTableDragMove={applyTableDragPositions} + onMoveTable={applyTableDragPositions} + onRequestColumns={requestColumns} + onConnectColumns={handleConnectColumns} + onConnectTables={(from, to) => { requestColumns(from); requestColumns(to); }} + canConnectColumns={canQueueForeignKey} + headerColors={resolvedHeaderColors} + pendingForeignKeys={store.pendingForeignKeys} + columnDetail={store.columnDetail} + diagramGroups={store.diagramGroups} + /> +
+ { + if (!store.primaryKey) return; + store.setHeaderColorsByKey((prev) => { + const next = { ...prev }; + if (hex == null) delete next[store.primaryKey!]; + else next[store.primaryKey!] = hex; + return next; + }); + }} + identityDraft={ + store.primaryKey ? store.identityDraftByKey[store.primaryKey] ?? { + schema: inspectorTable?.schema ?? "", + name: inspectorTable?.name ?? "", + } : null + } + onIdentityDraftChange={(next) => { + if (!store.primaryKey) return; + store.setIdentityDraftByKey((p) => ({ ...p, [store.primaryKey!]: next })); + }} + columnOverrides={ + store.primaryKey ? store.columnOverridesByKey[store.primaryKey] ?? {} : {} + } + onColumnOverridesChange={(next) => { + if (!store.primaryKey) return; + store.setColumnOverridesByKey((p) => ({ ...p, [store.primaryKey!]: next })); + }} + columnIdentityOverrides={ + store.primaryKey + ? store.columnIdentityOverridesByKey[store.primaryKey] ?? {} + : {} + } + onColumnIdentityOverridesChange={(next) => { + if (!store.primaryKey) return; + store.setColumnIdentityOverridesByKey((p) => { + const copy = { ...p }; + if (Object.keys(next).length === 0) delete copy[store.primaryKey!]; + else copy[store.primaryKey!] = next; + return copy; + }); + }} + catalogTables={catalogTablesSorted} + pendingAddColumns={ + store.primaryKey ? store.pendingAddColumnsByKey[store.primaryKey] ?? [] : [] + } + onPendingAddColumnsChange={(next) => { + if (!store.primaryKey) return; + store.setPendingAddColumnsByKey((p) => { + const copy = { ...p }; + if (next.length === 0) delete copy[store.primaryKey!]; + else copy[store.primaryKey!] = next; + return copy; + }); + }} + pendingForeignKeys={store.pendingForeignKeys} + selectedEdge={store.selectedEdge} + canQueueForeignKey={canQueueForeignKey} + onAddPendingForeignKey={(row) => { + const fromKey = row.fromKey ?? store.primaryKey; + if (!fromKey) return; + if (!canQueueForeignKey({ + fromKey, + fromColumn: row.fromColumn, + toKey: row.toKey, + toColumn: row.toColumn, + })) return; + const id = crypto.randomUUID(); + store.setPendingForeignKeys((prev) => [ + ...prev, + { + id, + fromKey, + fromColumn: row.fromColumn, + toKey: row.toKey, + toColumn: row.toColumn, + constraintName: row.constraintName, + }, + ]); + store.setSelectedEdge({ + id, + kind: "pending", + fromKey, + fromColumn: row.fromColumn, + toKey: row.toKey, + toColumn: row.toColumn, + }); + }} + onRemovePendingForeignKey={(id) => { + store.setPendingForeignKeys((prev) => prev.filter((fk) => fk.id !== id)); + store.setSelectedEdge((prev) => (prev?.id === id ? null : prev)); + }} + pendingRules={store.pendingRules.filter((row: any) => row.tableKey === store.primaryKey)} + onPendingRulesChange={(next: any) => { + if (!store.primaryKey) return; + store.setPendingRules((prev: any) => [ + ...prev.filter((row: any) => row.tableKey !== store.primaryKey), + ...next, + ]); + }} + pendingTriggers={store.pendingTriggers.filter((row: any) => row.tableKey === store.primaryKey)} + onPendingTriggersChange={(next: any) => { + if (!store.primaryKey) return; + store.setPendingTriggers((prev: any) => [ + ...prev.filter((row: any) => row.tableKey !== store.primaryKey), + ...next, + ]); + }} + pendingRlsPolicies={store.pendingRlsPolicies.filter((row: any) => row.tableKey === store.primaryKey)} + onPendingRlsPoliciesChange={(next: any) => { + if (!store.primaryKey) return; + store.setPendingRlsPolicies((prev: any) => [ + ...prev.filter((row: any) => row.tableKey !== store.primaryKey), + ...next, + ]); + }} + /> +
+
+
+ + + k ? requestColumns(k) : undefined} + selectedKeys={store.selectedKeys} primaryKey={store.primaryKey} + /> + +
+ + + store.setPendingCreateTables((prev) => [...prev, ct])} /> + +
+ ); } diff --git a/src/features/model/components/ModelWorkspaceToolbar.tsx b/src/features/model/components/ModelWorkspaceToolbar.tsx new file mode 100644 index 0000000..02c912f --- /dev/null +++ b/src/features/model/components/ModelWorkspaceToolbar.tsx @@ -0,0 +1,150 @@ +import { + AlignBottomIcon, AlignLeftIcon, AlignRightIcon, AlignTopIcon, + ArrowsClockwiseIcon, ArrowsInSimpleIcon, ArrowsOutIcon, + DownloadSimpleIcon, FilePdfIcon, GridFourIcon, MagnetIcon, + PlusIcon, SquaresFourIcon, TreeStructureIcon, +} from '@phosphor-icons/react'; +import { useTranslation } from 'react-i18next'; +import { Button } from '@/components/ui/button'; + +type DiagramToolbarProps = { + activeViewId: string; + viewsRegistry: { views: { id: string; name: string }[] }; + onViewChange: (id: string) => void; + onNewView: () => void; + onDeleteView: () => void; + columnDetail: string; + onColumnDetailChange: (v: string) => void; + canUndo: boolean; + canRedo: boolean; + onUndo: () => void; + onRedo: () => void; + snapToGrid: boolean; + onToggleSnap: () => void; + selectedKeysCount: number; + onAlignLeft: () => void; + onAlignRight: () => void; + onAlignTop: () => void; + onAlignBottom: () => void; + onAutoLayoutGrid: () => void; + onAutoLayoutTopo: () => void; + onAutoLayoutDagre: () => void; + onFitSelection: () => void; + onResetViewport: () => void; + onResetLayout: () => void; + onAddGroup: () => void; + onCreateTable: () => void; + onExportPng: () => void; + onExportPdf: () => void; + // for switch between diagram and catalog + onSwitchTab: (tab: 'diagram' | 'catalog') => void; +}; + +export function ModelWorkspaceToolbar({ + activeViewId, viewsRegistry, onViewChange, onNewView, onDeleteView, + columnDetail, onColumnDetailChange, + canUndo, canRedo, onUndo, onRedo, + snapToGrid, onToggleSnap, + selectedKeysCount, onAlignLeft, onAlignRight, onAlignTop, onAlignBottom, + onAutoLayoutGrid, onAutoLayoutTopo, onAutoLayoutDagre, + onFitSelection, onResetViewport, onResetLayout, + onAddGroup, onCreateTable, onExportPng, onExportPdf, +}: DiagramToolbarProps) { + const { t } = useTranslation(); + + return ( +
+ {t('model.align')} + + + + + +
+ + {t('model.views')} + + + + {viewsRegistry.views.length > 1 && activeViewId !== 'default' ? ( + + ) : null} + +
+ + + +
+ + + + +
+ + + + + + + +
+ + + + + +
+ + + + + +
+ ); +} diff --git a/src/features/model/hooks/useModelColumns.ts b/src/features/model/hooks/useModelColumns.ts new file mode 100644 index 0000000..281efb4 --- /dev/null +++ b/src/features/model/hooks/useModelColumns.ts @@ -0,0 +1,173 @@ +import { useQueries } from '@tanstack/react-query'; +import { useEffect, useMemo, useState } from 'react'; +import { queryKeys } from '@/data/query-keys'; +import { veloxDbRepository } from '@/data/repositories'; +import type { ColumnInfo, ForeignKeyEdge, TableInfo } from '@/data/types'; +import { tableKey, type ColumnDetailLevel, type TableKey } from '@/features/model/model-types'; + +interface UseModelColumnsParams { + connectionId: string; + tables: TableInfo[]; + onCanvas: TableKey[]; + columnDetail: ColumnDetailLevel; + foreignKeys: ForeignKeyEdge[]; + pendingForeignKeys: { fromKey: TableKey; toKey: TableKey; fromColumn: string; toColumn: string }[]; + columnIdentityOverridesByKey: Record>; + pendingAddColumnsByKey: Record; +} + +function tableKeyToParts(key: TableKey): { schema: string; name: string } { + const [schema = '', name = ''] = key.split('.'); + return { schema, name }; +} + +export function useModelColumns({ + connectionId, + tables, + onCanvas, + columnDetail, + foreignKeys, + pendingForeignKeys, + columnIdentityOverridesByKey, + pendingAddColumnsByKey, +}: UseModelColumnsParams) { + const tablesByKey = useMemo(() => { + const m = new Map(); + for (const t of tables) m.set(tableKey(t), t); + return m; + }, [tables]); + + const [columnRequestKeys, setColumnRequestKeys] = useState([]); + + const requestColumns = (key: TableKey) => { + setColumnRequestKeys((prev) => (prev.includes(key) ? prev : [...prev, key])); + }; + + const sortedRequestKeys = useMemo(() => [...columnRequestKeys].sort(), [columnRequestKeys]); + + const columnQueries = useQueries({ + queries: sortedRequestKeys.map((key) => { + const table = tablesByKey.get(key); + return { + queryKey: queryKeys.schema(connectionId, table ?? null), + queryFn: () => { + if (!table) throw new Error('Table not found for schema request.'); + return veloxDbRepository.getSchema(connectionId, table); + }, + enabled: Boolean(connectionId && table), + staleTime: 5 * 60 * 1000, + }; + }), + }); + + const columnsByKey = useMemo(() => { + const out: Record = {}; + sortedRequestKeys.forEach((key, i) => { + const q = columnQueries[i]; + if (!q) { out[key] = null; return; } + if (q.isPending && !q.data) out[key] = null; + else if (q.data) out[key] = q.data; + else out[key] = null; + }); + return out; + }, [columnQueries, sortedRequestKeys]); + + const effectiveColumnsByKey = useMemo(() => { + const out: Record = {}; + for (const [key, cols] of Object.entries(columnsByKey) as Array<[TableKey, ColumnInfo[] | null]>) { + if (!cols) { out[key] = null; continue; } + const identityOverrides = columnIdentityOverridesByKey[key] ?? {}; + const rows = cols.map((col) => { + const patch = identityOverrides[col.columnName]; + if (!patch) return col; + const nextName = patch.nextColumnName.trim(); + const nextType = patch.nextDataType.trim(); + return { ...col, columnName: nextName || col.columnName, dataType: nextType || col.dataType }; + }); + const pending = pendingAddColumnsByKey[key] ?? []; + const pendingRows: ColumnInfo[] = pending.map((col) => ({ + tableSchema: tableKeyToParts(key).schema, + tableName: tableKeyToParts(key).name, + columnName: col.columnName.trim(), + dataType: col.dataType.trim(), + isNullable: col.nullable, + })); + out[key] = [...rows, ...pendingRows]; + } + return out; + }, [columnIdentityOverridesByKey, columnsByKey, pendingAddColumnsByKey]); + + const onCanvasSet = useMemo(() => new Set(onCanvas), [onCanvas]); + + const fkColumnNamesByKey = useMemo(() => { + const m = new Map>(); + const add = (tab: TableKey, col: string) => { + if (!onCanvasSet.has(tab)) return; + const existing = m.get(tab); + if (existing) { existing.add(col); return; } + m.set(tab, new Set([col])); + }; + for (const fk of foreignKeys) { + add(`${fk.fromSchema}.${fk.fromTable}` as TableKey, fk.fromColumn); + add(`${fk.toSchema}.${fk.toTable}` as TableKey, fk.toColumn); + } + for (const p of pendingForeignKeys) { + add(p.fromKey, p.fromColumn); + add(p.toKey, p.toColumn); + } + return m; + }, [foreignKeys, onCanvasSet, pendingForeignKeys]); + + const diagramDisplayColumnsByKey = useMemo( + (): Record => { + const out: Record = {}; + for (const k of onCanvas) { + const cols = effectiveColumnsByKey[k] ?? null; + if (columnDetail === 'header') { out[k] = []; continue; } + if (columnDetail === 'keys' && cols?.length) { + const set = fkColumnNamesByKey.get(k); + const filtered = set?.size ? cols.filter((c) => set.has(c.columnName)) : cols.slice(0, 4); + out[k] = filtered.length > 0 ? filtered : cols.slice(0, 4); + continue; + } + out[k] = cols; + } + return out; + }, + [columnDetail, effectiveColumnsByKey, fkColumnNamesByKey, onCanvas], + ); + + // Auto-request columns for FK-related tables + // eslint-disable-next-line react-hooks/set-state-in-effect + useEffect(() => { + const keys = new Set(); + for (const fk of foreignKeys) { + const fromK = `${fk.fromSchema}.${fk.fromTable}` as TableKey; + const toK = `${fk.toSchema}.${fk.toTable}` as TableKey; + if (onCanvasSet.has(fromK)) keys.add(fromK); + if (onCanvasSet.has(toK)) keys.add(toK); + } + for (const p of pendingForeignKeys) { + keys.add(p.fromKey); + keys.add(p.toKey); + } + if (keys.size === 0) return; + setColumnRequestKeys((prev) => { + let next = prev; + for (const k of keys) { + if (!next.includes(k)) next = [...next, k]; + } + return next; + }); + }, [foreignKeys, onCanvasSet, pendingForeignKeys]); + + return { + columnsByKey, + effectiveColumnsByKey, + diagramDisplayColumnsByKey, + fkColumnNamesByKey, + columnRequestKeys, + setColumnRequestKeys, + requestColumns, + }; +} diff --git a/src/features/model/hooks/useModelInitialization.ts b/src/features/model/hooks/useModelInitialization.ts new file mode 100644 index 0000000..0fb6f40 --- /dev/null +++ b/src/features/model/hooks/useModelInitialization.ts @@ -0,0 +1,142 @@ +import { useEffect, useMemo, useRef } from 'react'; +import type { TableInfo } from '@/data/types'; +import { ensurePositions } from '@/features/model/model-layout-storage'; +import { tableKey, type TableKey } from '@/features/model/model-types'; +import { useForeignKeysQuery } from '@/features/model/queries'; + +interface UseModelInitializationParams { + connectionId: string; + tables: TableInfo[]; + onCanvas: TableKey[]; + hadStoredLayout: boolean; + setOnCanvas: (updater: TableKey[] | ((prev: TableKey[]) => TableKey[])) => void; + setPositions: ( + updater: + | Record + | ((prev: Record) => Record), + ) => void; + selectSingleFromCatalog: (key: TableKey) => void; + setInitialSeedReason: (reason: 'relationships' | 'sample' | null) => void; + selectedTable: TableInfo | null; + primaryKey: TableKey | null; + setIdentityDraftByKey: ( + updater: + | Record + | ((prev: Record) => Record), + ) => void; +} + +export function useModelInitialization({ + connectionId, + tables, + onCanvas, + hadStoredLayout, + setOnCanvas, + setPositions, + selectSingleFromCatalog, + setInitialSeedReason, + selectedTable, + primaryKey, + setIdentityDraftByKey, +}: UseModelInitializationParams) { + const foreignKeysQuery = useForeignKeysQuery(connectionId); + + const tablesByKey = useMemo(() => { + const m = new Map(); + for (const t of tables) m.set(tableKey(t), t); + return m; + }, [tables]); + + const fkSeedDoneRef = useRef(false); + const initialRecoveryDoneRef = useRef(false); + + // Reset guards when connection changes + useEffect(() => { + void connectionId; + initialRecoveryDoneRef.current = false; + fkSeedDoneRef.current = false; + setInitialSeedReason(null); + }, [connectionId, setInitialSeedReason]); + + // Initial canvas recovery: prune invalid keys, seed from FK or sample + useEffect(() => { + if (initialRecoveryDoneRef.current) return; + if (!tables.length) return; + + const validOnCanvas = onCanvas.filter((k) => tablesByKey.has(k)); + if (validOnCanvas.length !== onCanvas.length) { + setOnCanvas(validOnCanvas); + setPositions((prev) => { + const next: Record = {}; + for (const key of validOnCanvas) { + if (prev[key]) next[key] = prev[key]; + } + return next; + }); + } + + if (validOnCanvas.length === 0) { + const fkData = foreignKeysQuery.data ?? []; + const fkSeed = new Set(); + for (const edge of fkData) { + const from = `${edge.fromSchema}.${edge.fromTable}` as TableKey; + const to = `${edge.toSchema}.${edge.toTable}` as TableKey; + if (tablesByKey.has(from)) fkSeed.add(from); + if (tablesByKey.has(to)) fkSeed.add(to); + } + const fallbackKeys = fkSeed.size > 0 ? [...fkSeed] : tables.slice(0, 12).map((t) => tableKey(t)); + if (fallbackKeys.length > 0) { + setInitialSeedReason(fkSeed.size > 0 ? 'relationships' : 'sample'); + setOnCanvas(fallbackKeys); + setPositions((prev) => ensurePositions(fallbackKeys, prev)); + } + } + + initialRecoveryDoneRef.current = true; + }, [foreignKeysQuery.data, onCanvas, tables, tablesByKey, setOnCanvas, setPositions, setInitialSeedReason]); + + // FK seed fallback when no stored layout exists + useEffect(() => { + if (hadStoredLayout) return; + if (fkSeedDoneRef.current) return; + const fkData = foreignKeysQuery.data; + if (!fkData?.length || !tables.length) return; + + const keys = new Set(); + for (const e of fkData) { + keys.add(`${e.fromSchema}.${e.fromTable}`); + keys.add(`${e.toSchema}.${e.toTable}`); + } + const valid = [...keys].filter((k) => tables.some((t) => tableKey(t) === k)); + if (valid.length === 0) return; + + fkSeedDoneRef.current = true; + queueMicrotask(() => { + setOnCanvas(valid); + setPositions((p) => ensurePositions(valid, p)); + }); + }, [hadStoredLayout, foreignKeysQuery.data, tables, setOnCanvas, setPositions]); + + // Selected table from outside → add to canvas + useEffect(() => { + if (!selectedTable) return; + const k = tableKey(selectedTable); + queueMicrotask(() => { + selectSingleFromCatalog(k); + setOnCanvas((prev) => (prev.includes(k) ? prev : [...prev, k])); + setPositions((p) => ensurePositions([k], p)); + }); + }, [selectSingleFromCatalog, selectedTable, setOnCanvas, setPositions]); + + // Identity draft for inspector's primary table + useEffect(() => { + if (!primaryKey) return; + const t = tablesByKey.get(primaryKey); + if (!t) return; + setIdentityDraftByKey((prev) => + prev[primaryKey] != null ? prev : { ...prev, [primaryKey]: { schema: t.schema, name: t.name } }, + ); + }, [primaryKey, tablesByKey, setIdentityDraftByKey]); + + return { tablesByKey, foreignKeysQuery }; +} diff --git a/src/features/model/hooks/useModelWorkspaceStore.ts b/src/features/model/hooks/useModelWorkspaceStore.ts new file mode 100644 index 0000000..af7321a --- /dev/null +++ b/src/features/model/hooks/useModelWorkspaceStore.ts @@ -0,0 +1,73 @@ +import { useShallow } from 'zustand/react/shallow'; +import { useCanvasStore } from '@/features/model/state/canvas-store'; + +/** + * Centralised Zustand selector for the ModelWorkspace. + * Pulls every slice the workspace needs into one call, avoiding + * repeated useCanvasStore invocations in the component body. + */ +export function useModelWorkspaceStore() { + return useCanvasStore( + useShallow((s) => ({ + hydrated: s.hydrated, + hydrateFromConnection: s.hydrateFromConnection, + storeConnectionId: s.connectionId, + viewsRegistry: s.viewsRegistry, + setViewsRegistry: s.setViewsRegistry, + activeViewId: s.activeViewId, + diagramTool: s.diagramTool, + setDiagramTool: s.setDiagramTool, + selectedKeys: s.selectedKeys, + setSelectedKeys: s.setSelectedKeys, + primaryKey: s.primaryKey, + replaceSelection: s.replaceSelection, + selectTable: s.selectTable, + clearSelection: s.clearSelection, + applyMarquee: s.applyMarquee, + selectSingleFromCatalog: s.selectSingleFromCatalog, + snapToGrid: s.snapToGrid, + setSnapToGrid: s.setSnapToGrid, + onCanvas: s.onCanvas, + setOnCanvas: s.setOnCanvas, + positions: s.positions, + setPositions: s.setPositions, + viewport: s.viewport, + setViewport: s.setViewport, + modelTitle: s.modelTitle, + setModelTitle: s.setModelTitle, + headerColorsByKey: s.headerColorsByKey, + setHeaderColorsByKey: s.setHeaderColorsByKey, + columnDetail: s.columnDetail, + setColumnDetail: s.setColumnDetail, + diagramGroups: s.diagramGroups, + setDiagramGroups: s.setDiagramGroups, + modelTab: s.modelTab, + setModelTab: s.setModelTab, + identityDraftByKey: s.identityDraftByKey, + setIdentityDraftByKey: s.setIdentityDraftByKey, + columnOverridesByKey: s.columnOverridesByKey, + setColumnOverridesByKey: s.setColumnOverridesByKey, + columnIdentityOverridesByKey: s.columnIdentityOverridesByKey, + setColumnIdentityOverridesByKey: s.setColumnIdentityOverridesByKey, + pendingAddColumnsByKey: s.pendingAddColumnsByKey, + setPendingAddColumnsByKey: s.setPendingAddColumnsByKey, + pendingForeignKeys: s.pendingForeignKeys, + setPendingForeignKeys: s.setPendingForeignKeys, + selectedEdge: s.selectedEdge, + setSelectedEdge: s.setSelectedEdge, + pendingRules: s.pendingRules, + setPendingRules: s.setPendingRules, + pendingTriggers: s.pendingTriggers, + setPendingTriggers: s.setPendingTriggers, + pendingRlsPolicies: s.pendingRlsPolicies, + setPendingRlsPolicies: s.setPendingRlsPolicies, + pendingCreateTables: s.pendingCreateTables, + setPendingCreateTables: s.setPendingCreateTables, + applyQuickColumnEdit: s.applyQuickColumnEdit, + canUndo: s.canUndo, + canRedo: s.canRedo, + undo: s.undo, + redo: s.redo, + })), + ); +} diff --git a/src/features/queries/components/AskVeloxyDialog.tsx b/src/features/queries/components/AskVeloxyDialog.tsx index caf6003..1efa5d6 100644 --- a/src/features/queries/components/AskVeloxyDialog.tsx +++ b/src/features/queries/components/AskVeloxyDialog.tsx @@ -23,36 +23,24 @@ import { Button } from "@/components/ui/button"; import { Textarea } from "@/components/ui/textarea"; import type { AskVeloxyChatResponse, - AskVeloxyConversationMessage, AskVeloxyConversationResponse, - AskVeloxyResponse, VeloxyStreamChunk, } from "@/data/types"; import { VeloxyMarkdown } from "@/features/queries/components/VeloxyMarkdown"; +import { + type AskVeloxySubmitResult, + type ChatMessage, + extractTextFromUnknown, + looksLikeJsonResponse, + messageBodyIsSqlDraft, + normalizeAssistantMessage, + truncateSuggestion, +} from "@/features/queries/components/veloxy-message-parser"; import { useVeloxyStream } from "@/features/queries/hooks/useVeloxyStream"; import { notifySuccess } from "@/lib/error-notifier"; import { cn } from "@/lib/utils"; -export type AskVeloxySubmitResult = { - response: AskVeloxyResponse; - decision: "auto-ran" | "needs-confirmation"; - decisionReason?: string; - pendingSql?: string; -}; - -type ChatMessage = AskVeloxyConversationMessage & { - clientNonce?: number; - streaming?: boolean; - stoppedEarly?: boolean; - result?: AskVeloxyResponse; - decision?: AskVeloxySubmitResult["decision"]; - decisionReason?: string; - pendingSql?: string; - suggestions?: string[]; - warnings?: string[]; - needsSqlGeneration?: boolean; - needsClarification?: boolean; -}; +export type { AskVeloxySubmitResult } from "@/features/queries/components/veloxy-message-parser"; type AskVeloxySidebarProps = { isPending: boolean; @@ -76,155 +64,6 @@ type AskVeloxySidebarProps = { errorMessage: string | null; }; -function extractTextFromUnknown(value: unknown): string | null { - if (typeof value === "string") { - const normalized = value.trim(); - return normalized.length > 0 ? normalized : null; - } - if (!value || typeof value !== "object") return null; - - const record = value as Record; - for (const key of ["message", "reply", "content", "text"]) { - const text = extractTextFromUnknown(record[key]); - if (text) return text; - } - for (const key of ["output", "response", "data", "result"]) { - const text = extractTextFromUnknown(record[key]); - if (text) return text; - } - return null; -} - -function unescapeJsonFragment(raw: string): string { - let out = ""; - for (let i = 0; i < raw.length; i += 1) { - const ch = raw[i]; - if (ch !== "\\") { - out += ch; - continue; - } - const next = raw[i + 1]; - if (next === "n") { - out += "\n"; - i += 1; - } else if (next === "t") { - out += "\t"; - i += 1; - } else if (next === "r") { - out += "\r"; - i += 1; - } else if (next === '"') { - out += '"'; - i += 1; - } else if (next === "\\") { - out += "\\"; - i += 1; - } else if (next != null) { - out += `\\${next}`; - i += 1; - } else { - out += "\\"; - } - } - return out; -} - -function extractJsonMessageField( - raw: string, - allowPartial: boolean, -): string | null { - const unwrapped = raw - .trim() - .replace(/^```json\s*/i, "") - .replace(/\s*```$/, ""); - for (const key of ["message", "reply", "content", "text"]) { - const marker = `"${key}"`; - const markerIdx = unwrapped.indexOf(marker); - if (markerIdx < 0) continue; - let idx = markerIdx + marker.length; - while (idx < unwrapped.length && /\s/.test(unwrapped[idx] ?? "")) idx += 1; - if (unwrapped[idx] !== ":") continue; - idx += 1; - while (idx < unwrapped.length && /\s/.test(unwrapped[idx] ?? "")) idx += 1; - if (unwrapped[idx] !== '"') continue; - idx += 1; - const start = idx; - let escaped = false; - while (idx < unwrapped.length) { - const ch = unwrapped[idx]; - if (escaped) { - escaped = false; - idx += 1; - continue; - } - if (ch === "\\") { - escaped = true; - idx += 1; - continue; - } - if (ch === '"') { - const fragment = unwrapped.slice(start, idx); - try { - const decoded = JSON.parse(`"${fragment}"`); - if (typeof decoded === "string" && decoded.trim()) - return decoded.trim(); - } catch { - const text = unescapeJsonFragment(fragment).trim(); - if (text) return text; - } - break; - } - idx += 1; - } - if (allowPartial && start < unwrapped.length) { - const text = unescapeJsonFragment(unwrapped.slice(start)).trim(); - if (text) return text; - } - } - return null; -} - -function looksLikeJsonResponse(raw: string): boolean { - const trimmed = raw.trimStart(); - return trimmed.startsWith("{") || trimmed.startsWith("```"); -} - -function normalizeAssistantMessage(raw: string): string { - const trimmed = raw.trim(); - if (!trimmed) return ""; - - const unwrapped = trimmed.replace(/^```json\s*/i, "").replace(/\s*```$/, ""); - try { - const parsed = JSON.parse(unwrapped); - return extractTextFromUnknown(parsed) ?? trimmed; - } catch { - const extracted = - extractJsonMessageField(trimmed, false) ?? - extractJsonMessageField(trimmed, true); - if (extracted) return extracted; - if (looksLikeJsonResponse(trimmed)) return ""; - return trimmed; - } -} - -function messageBodyIsSqlDraft(message: ChatMessage): boolean { - if (message.role !== "assistant" || message.mode !== "action") return false; - const t = message.text.trimStart().toLowerCase(); - return ( - t.startsWith("select") || - t.startsWith("with") || - t.startsWith("insert") || - t.startsWith("update") || - t.startsWith("delete") || - t.startsWith("explain") - ); -} - -function truncateSuggestion(text: string, max = 72): string { - if (text.length <= max) return text; - return `${text.slice(0, max - 1).trimEnd()}…`; -} - function ChatWarningsBanner({ warnings }: { warnings: string[] }) { if (!warnings.length) return null; return ( diff --git a/src/features/queries/components/QueryWorkspace.tsx b/src/features/queries/components/QueryWorkspace.tsx index b06b4b8..c3aff3c 100644 --- a/src/features/queries/components/QueryWorkspace.tsx +++ b/src/features/queries/components/QueryWorkspace.tsx @@ -272,6 +272,7 @@ function QueryPane({ onChange={onSqlChange} onRun={onRun} onRunStatement={onRunStatement} + language={connectionEngine === "mongo" ? "json" : "sql"} metadata={editorMetadata} diagnostics={lintDiagnostics} /> @@ -688,7 +689,7 @@ export const QueryWorkspace = forwardRef< if (lintTimerRef.current != null) { window.clearTimeout(lintTimerRef.current); } - if (!targetConnectionId || sql.trim().length === 0) { + if (!targetConnectionId || sql.trim().length === 0 || connectionEngine === "mongo") { lintReset(); return; } @@ -758,7 +759,7 @@ export const QueryWorkspace = forwardRef< onRequestConnection(); return; } - const allowWrite = !isReadOnlySql(trimmed); + const allowWrite = !isReadOnlySql(trimmed, connectionEngine ?? undefined); if (allowWrite && !window.confirm(t("editor.confirmWrite"))) { return; } diff --git a/src/features/queries/components/ResultsCellEditor.tsx b/src/features/queries/components/ResultsCellEditor.tsx new file mode 100644 index 0000000..280af98 --- /dev/null +++ b/src/features/queries/components/ResultsCellEditor.tsx @@ -0,0 +1,63 @@ +import { useLayoutEffect, useRef } from "react"; + +export function ResultEditInput({ + defaultValue, + onBlurCommit, + onEscape, +}: { + defaultValue: string; + onBlurCommit: (raw: string) => void; + onEscape: () => void; +}) { + const inputRef = useRef(null); + const skipBlurCommitRef = useRef(false); + + useLayoutEffect(() => { + const element = inputRef.current; + if (!element) return; + element.focus(); + element.select(); + }, []); + + return ( + { + if (skipBlurCommitRef.current) { + skipBlurCommitRef.current = false; + return; + } + onBlurCommit(event.target.value); + }} + onKeyDown={(event) => { + if (event.key === "Enter") event.currentTarget.blur(); + if (event.key === "Escape") { + skipBlurCommitRef.current = true; + onEscape(); + } + }} + /> + ); +} + +export function InsertRowInput({ + value, + onChange, + placeholder, +}: { + value: string; + onChange: (next: string) => void; + placeholder: string; +}) { + return ( + onChange(event.target.value)} + placeholder={placeholder} + autoComplete="off" + /> + ); +} diff --git a/src/features/queries/components/ResultsGrid.tsx b/src/features/queries/components/ResultsGrid.tsx index e31305d..85a8d83 100644 --- a/src/features/queries/components/ResultsGrid.tsx +++ b/src/features/queries/components/ResultsGrid.tsx @@ -19,6 +19,7 @@ import { import { useTranslation } from "react-i18next"; import type { ColumnProperties, QueryResult, TableInfo } from "@/data/types"; +import { ResultEditInput, InsertRowInput } from "@/features/queries/components/ResultsCellEditor"; import { ResultsToolbar } from "@/features/queries/components/ResultsToolbar"; import { useInsertRowMutation } from "@/features/queries/queries"; import { @@ -81,72 +82,6 @@ function normalizeColumnId(columnId: string) { return columnId.toLowerCase(); } -function ResultEditInput({ - defaultValue, - onBlurCommit, - onEscape, -}: { - defaultValue: string; - onBlurCommit: (raw: string) => void; - onEscape: () => void; -}) { - const inputRef = useRef(null); - const skipBlurCommitRef = useRef(false); - - useLayoutEffect(() => { - const element = inputRef.current; - if (!element) { - return; - } - element.focus(); - element.select(); - }, []); - - return ( - { - if (skipBlurCommitRef.current) { - skipBlurCommitRef.current = false; - return; - } - onBlurCommit(event.target.value); - }} - onKeyDown={(event) => { - if (event.key === "Enter") { - event.currentTarget.blur(); - } - if (event.key === "Escape") { - skipBlurCommitRef.current = true; - onEscape(); - } - }} - /> - ); -} - -function InsertRowInput({ - value, - onChange, - placeholder, -}: { - value: string; - onChange: (next: string) => void; - placeholder: string; -}) { - return ( - onChange(event.target.value)} - placeholder={placeholder} - autoComplete="off" - /> - ); -} - function renderLoadingSkeleton() { const dataColumnCount = 4; const placeholderRows = 10; diff --git a/src/features/queries/components/SqlEditor.tsx b/src/features/queries/components/SqlEditor.tsx index d96bfb4..e9ebdbc 100644 --- a/src/features/queries/components/SqlEditor.tsx +++ b/src/features/queries/components/SqlEditor.tsx @@ -12,8 +12,10 @@ type SqlEditorProps = { onChange: (value: string) => void; onRun: () => void; onRunStatement: (sql: string) => void; + /** Language mode for Monaco. Defaults to "sql" for relational, "json" for MongoDB. */ + language?: string; metadata?: QueryEditorMetadata; - diagnostics: SqlDiagnostic[]; + diagnostics?: SqlDiagnostic[]; }; function completionItemsFromMetadata( @@ -191,6 +193,7 @@ export function SqlEditor({ onChange, onRun, onRunStatement, + language = "sql", metadata, diagnostics, }: SqlEditorProps) { @@ -302,7 +305,8 @@ export function SqlEditor({ return ( 0 ? normalized : null; + } + if (!value || typeof value !== "object") return null; + + const record = value as Record; + for (const key of ["message", "reply", "content", "text"]) { + const text = extractTextFromUnknown(record[key]); + if (text) return text; + } + for (const key of ["output", "response", "data", "result"]) { + const text = extractTextFromUnknown(record[key]); + if (text) return text; + } + return null; +} + +export function unescapeJsonFragment(raw: string): string { + let out = ""; + for (let i = 0; i < raw.length; i += 1) { + const ch = raw[i]; + if (ch !== "\\") { out += ch; continue; } + const next = raw[i + 1]; + if (next === "n") { out += "\n"; i += 1; } + else if (next === "t") { out += "\t"; i += 1; } + else if (next === "r") { out += "\r"; i += 1; } + else if (next === '"') { out += '"'; i += 1; } + else if (next === "\\") { out += "\\"; i += 1; } + else if (next != null) { out += `\\${next}`; i += 1; } + else { out += "\\"; } + } + return out; +} + +export function extractJsonMessageField(raw: string, allowPartial: boolean): string | null { + const unwrapped = raw.trim().replace(/^```json\s*/i, "").replace(/\s*```$/, ""); + for (const key of ["message", "reply", "content", "text"]) { + const marker = `"${key}"`; + const markerIdx = unwrapped.indexOf(marker); + if (markerIdx < 0) continue; + let idx = markerIdx + marker.length; + while (idx < unwrapped.length && /\s/.test(unwrapped[idx] ?? "")) idx += 1; + if (unwrapped[idx] !== ":") continue; + idx += 1; + while (idx < unwrapped.length && /\s/.test(unwrapped[idx] ?? "")) idx += 1; + if (unwrapped[idx] !== '"') continue; + idx += 1; + const start = idx; + let escaped = false; + while (idx < unwrapped.length) { + const ch = unwrapped[idx]; + if (escaped) { escaped = false; idx += 1; continue; } + if (ch === "\\") { escaped = true; idx += 1; continue; } + if (ch === '"') { + const fragment = unwrapped.slice(start, idx); + try { + const decoded = JSON.parse(`"${fragment}"`); + if (typeof decoded === "string" && decoded.trim()) return decoded.trim(); + } catch { + const text = unescapeJsonFragment(fragment).trim(); + if (text) return text; + } + break; + } + idx += 1; + } + if (allowPartial && start < unwrapped.length) { + const text = unescapeJsonFragment(unwrapped.slice(start)).trim(); + if (text) return text; + } + } + return null; +} + +export function looksLikeJsonResponse(raw: string): boolean { + const trimmed = raw.trimStart(); + return trimmed.startsWith("{") || trimmed.startsWith("```"); +} + +export function normalizeAssistantMessage(raw: string): string { + const trimmed = raw.trim(); + if (!trimmed) return ""; + const unwrapped = trimmed.replace(/^```json\s*/i, "").replace(/\s*```$/, ""); + try { + const parsed = JSON.parse(unwrapped); + return extractTextFromUnknown(parsed) ?? trimmed; + } catch { + const extracted = extractJsonMessageField(trimmed, false) ?? extractJsonMessageField(trimmed, true); + if (extracted) return extracted; + if (looksLikeJsonResponse(trimmed)) return ""; + return trimmed; + } +} + +export function messageBodyIsSqlDraft(message: ChatMessage): boolean { + if (message.role !== "assistant" || message.mode !== "action") return false; + const t = message.text.trimStart().toLowerCase(); + return t.startsWith("select") || t.startsWith("with") || t.startsWith("insert") || + t.startsWith("update") || t.startsWith("delete") || t.startsWith("explain"); +} + +export function truncateSuggestion(text: string, max = 72): string { + if (text.length <= max) return text; + return `${text.slice(0, max - 1).trimEnd()}…`; +} diff --git a/src/features/workspace/ModelWorkspaceAdapter.tsx b/src/features/workspace/ModelWorkspaceAdapter.tsx new file mode 100644 index 0000000..c63b225 --- /dev/null +++ b/src/features/workspace/ModelWorkspaceAdapter.tsx @@ -0,0 +1,45 @@ +import { ErrorBoundary } from "@/components/ErrorBoundary"; +import { ModelWorkspace } from "@/features/model/components/ModelWorkspace"; +import type { WorkspaceShellProps } from "./types"; + +function engineLabel(engine: string): string { + if (engine === "postgres") return "PostgreSQL"; + if (engine === "mysql") return "MySQL"; + return "SQLite"; +} + +export function ModelWorkspaceAdapter(props: WorkspaceShellProps) { + const { connection, tablesForUi, tablesErrorMessage, isTablesLoading, selectedTable, isDark } = props; + + if (!connection?.id) { + return ( +
+ Connect to a database to use the Model workspace. +
+ ); + } + + if (connection.engine !== "postgres") { + return ( +
+ The Model workspace is only available for PostgreSQL connections ({engineLabel(connection.engine)} selected). +
+ ); + } + + return ( + + + + ); +} diff --git a/src/features/workspace/types.ts b/src/features/workspace/types.ts new file mode 100644 index 0000000..2dde8af --- /dev/null +++ b/src/features/workspace/types.ts @@ -0,0 +1,17 @@ +import type { ConnectionSummary, TableInfo } from "@/data/types"; + +export type MainWorkspaceId = "query" | "model"; + +export interface WorkspaceShellProps { + connection: ConnectionSummary | null; + connectionError: unknown; + connectionErrorMessage: string; + isDark: boolean; + onRequestConnection: () => void; + resultsHeight: number; + onResultsHeightChange: (height: number) => void; + selectedTable: TableInfo | null; + tablesForUi: TableInfo[]; + tablesErrorMessage: string | undefined; + isTablesLoading: boolean; +} diff --git a/src/features/workspace/workspace-utils.ts b/src/features/workspace/workspace-utils.ts new file mode 100644 index 0000000..5527e82 --- /dev/null +++ b/src/features/workspace/workspace-utils.ts @@ -0,0 +1,10 @@ +import type { MainWorkspaceId } from "./types"; + +export function getWorkspaceEngineFilter(id: MainWorkspaceId): string[] | undefined { + switch (id) { + case "model": + return ["postgres"]; + default: + return undefined; + } +} diff --git a/src/hooks/useAppState.ts b/src/hooks/useAppState.ts new file mode 100644 index 0000000..b998b72 --- /dev/null +++ b/src/hooks/useAppState.ts @@ -0,0 +1,654 @@ +import { useQueryClient } from "@tanstack/react-query"; +import { + useCallback, + useEffect, + useMemo, + useRef, + useState, +} from "react"; +import { useTranslation } from "react-i18next"; + +import { queryKeys } from "@/data/query-keys"; +import { veloxDbRepository } from "@/data/repositories"; +import type { + ConnectionSummary, + TableInfo, + AskVeloxyChatResponse, + AskVeloxyConversationResponse, +} from "@/data/types"; +import { + useActivateConnectionMutation, + useConnectionsQuery, + useConnectMutation, + useDeleteConnectionMutation, + useRenameConnectionMutation, +} from "@/features/connections/queries"; +import { useSaveResultEditsMutation, useDeleteRowsMutation } from "@/features/queries/queries"; +import { useTableSchemaQuery, useTablePropertiesQuery } from "@/features/schema/queries"; +import { useTablesQuery } from "@/features/tables/queries"; +import { + buildDropTableSql, + buildDeleteTemplateSql, + buildInsertTemplateSql, + buildUpdateTemplateSql, + buildRenameTableSql, + buildSelectAllSql, + buildSelectCountSql, +} from "@/features/queries/sql-templates"; +import type { TableQuickSqlAction } from "@/features/queries/table-quick-actions"; +import type { ResultEditPatch } from "@/features/queries/result-edits"; +import { isInsertFormColumn } from "@/features/queries/result-edits"; +import { quoteIdent } from "@/lib/sql-ident"; +import { notifyError, notifySuccess } from "@/lib/error-notifier"; +import { loadOpenRouterApiKey } from "@/lib/openrouter-credentials"; +import { useSettings, resolveTheme } from "@/lib/settings"; + +import type { QueryWorkspaceHandle } from "@/features/queries/components/QueryWorkspace"; + +export const SIDEBAR_WIDTH_KEY = "veloxdb.sidebarWidth"; +export const SIDEBAR_COLLAPSED_KEY = "veloxdb.sidebarCollapsed"; +export const RESULTS_HEIGHT_KEY = "veloxdb.resultsHeight"; +export const LAST_ACTIVE_CONNECTION_KEY = "veloxdb.lastActiveConnectionId"; +export const DEFAULT_SIDEBAR_WIDTH = 280; +export const MIN_SIDEBAR_WIDTH = 220; +export const MAX_SIDEBAR_WIDTH = 520; +export const DEFAULT_RESULTS_HEIGHT = 260; + +export function clampSidebarWidth(value: number) { + return Math.min(MAX_SIDEBAR_WIDTH, Math.max(MIN_SIDEBAR_WIDTH, value)); +} + +export function connectionSecondaryText(connection: ConnectionSummary): string { + if (connection.engine === "mongo") { + return `mongodb://${connection.host}:${connection.port}/${connection.database}`; + } + if (connection.engine === "sqlite") { + return connection.filePath === ":memory:" + ? "SQLite in-memory database" + : `SQLite file: ${connection.filePath ?? connection.database}`; + } + return `${connection.user}@${connection.host}:${connection.port}${connection.sshConfig ? " (via SSH)" : ""}`; +} + +export function connectionHeadline(connection: ConnectionSummary): string { + if (connection.engine === "mongo") { + return `Connected to MongoDB (${connection.host}:${connection.port}/${connection.database})`; + } + if (connection.engine === "sqlite") { + return `Connected to SQLite (${connection.filePath ?? connection.database})`; + } + return `Connected to ${connection.database} on ${connection.host}:${connection.port}`; +} + +export function engineLabel(engine: ConnectionSummary["engine"]): string { + if (engine === "postgres") return "PostgreSQL"; + if (engine === "mysql") return "MySQL"; + if (engine === "sqlite") return "SQLite"; + if (engine === "mongo") return "MongoDB"; + return "Unknown"; +} + +export function useAppState( + queryWorkspaceRef: React.RefObject, +) { + const { t } = useTranslation(); + const queryClient = useQueryClient(); + + const [connection, setConnection] = useState(null); + const [focusedQueryCaps, setFocusedQueryCaps] = useState({ hasLastQuery: false, hasResult: false }); + const [tableSearch, setTableSearch] = useState(""); + const [selectedTable, setSelectedTable] = useState(null); + const [settingsOpen, setSettingsOpen] = useState(false); + const [commandPaletteOpen, setCommandPaletteOpen] = useState(false); + const [connectionDialogOpen, setConnectionDialogOpen] = useState(false); + const [renamingConnection, setRenamingConnection] = useState(null); + const [isSidebarCollapsed, setIsSidebarCollapsed] = useState(() => + window.localStorage.getItem(SIDEBAR_COLLAPSED_KEY) === "true", + ); + const [sidebarWidth, setSidebarWidth] = useState(() => { + const value = Number(window.localStorage.getItem(SIDEBAR_WIDTH_KEY)); + return Number.isFinite(value) ? clampSidebarWidth(value) : DEFAULT_SIDEBAR_WIDTH; + }); + const [resultsHeight, setResultsHeight] = useState(() => { + const value = Number(window.localStorage.getItem(RESULTS_HEIGHT_KEY)); + return Number.isFinite(value) && value > 0 ? value : DEFAULT_RESULTS_HEIGHT; + }); + const [tablePropertiesDialogOpen, setTablePropertiesDialogOpen] = useState(false); + const [tablePropertiesTarget, setTablePropertiesTarget] = useState<{ + connectionId: string; table: TableInfo; + } | null>(null); + const [insertRowTrigger, setInsertRowTrigger] = useState(0); + const [mainWorkspace, setMainWorkspace] = useState<"query" | "model">("query"); + const [askVeloxyPending, setAskVeloxyPending] = useState(false); + const [askVeloxyError, setAskVeloxyError] = useState(null); + + const veloxyOpenRouterApiKey = useSettings((s) => s.veloxyOpenRouterApiKey); + const veloxyModel = useSettings((s) => s.veloxyModel); + const veloxyBaseUrl = useSettings((s) => s.veloxyBaseUrl); + const autoReconnect = useSettings((s) => s.autoReconnect); + + const themeSetting = useSettings((s) => s.theme); + const isDark = useMemo(() => resolveTheme(themeSetting) === "dark", [themeSetting]); + const fontSize = useSettings((s) => s.fontSize); + + useEffect(() => { + document.documentElement.classList.toggle("dark", isDark); + }, [isDark]); + + useEffect(() => { + const sizes = { sm: 12, md: 14, lg: 16 }; + document.documentElement.style.fontSize = `${sizes[fontSize]}px`; + }, [fontSize]); + + useEffect(() => { + void loadOpenRouterApiKey(); + }, []); + + useEffect(() => { + window.localStorage.setItem(SIDEBAR_COLLAPSED_KEY, String(isSidebarCollapsed)); + }, [isSidebarCollapsed]); + + useEffect(() => { + window.localStorage.setItem(SIDEBAR_WIDTH_KEY, String(sidebarWidth)); + }, [sidebarWidth]); + + useEffect(() => { + window.localStorage.setItem(RESULTS_HEIGHT_KEY, String(resultsHeight)); + }, [resultsHeight]); + + const connectionsQuery = useConnectionsQuery(); + + useEffect(() => { + if (connectionsQuery.isError && connectionsQuery.error) { + notifyError(connectionsQuery.error, { title: t("connection.failedToLoad") }); + } + }, [connectionsQuery.isError, connectionsQuery.error, t]); + + const connectMutation = useConnectMutation({ + onError: (error) => { + notifyError(error, { category: "connection", force: true }); + }, + onSuccess: (nextConnection) => { + notifySuccess( + t("connection.connected", { database: nextConnection.database }), + connectionSecondaryText(nextConnection), + ); + window.localStorage.setItem(LAST_ACTIVE_CONNECTION_KEY, nextConnection.id); + setConnection(nextConnection); + setSelectedTable(null); + setTableSearch(""); + setIsSidebarCollapsed(false); + setConnectionDialogOpen(false); + setTablePropertiesDialogOpen(false); + setTablePropertiesTarget(null); + queueMicrotask(() => queryWorkspaceRef.current?.setActiveTabConnection(nextConnection.id)); + }, + }); + + const activateConnectionMutation = useActivateConnectionMutation({ + onError: (error) => { + notifyError(error, { category: "connection", force: true }); + }, + onSuccess: (nextConnection) => { + window.localStorage.setItem(LAST_ACTIVE_CONNECTION_KEY, nextConnection.id); + setConnection(nextConnection); + setSelectedTable(null); + setTableSearch(""); + setTablePropertiesDialogOpen(false); + setTablePropertiesTarget(null); + queueMicrotask(() => queryWorkspaceRef.current?.setActiveTabConnection(nextConnection.id)); + }, + }); + + const deleteConnectionMutation = useDeleteConnectionMutation({ + onError: (error) => { + notifyError(error, { category: "connection", force: true }); + }, + onSuccess: (connectionId) => { + queryWorkspaceRef.current?.detachDeletedConnection(connectionId); + if (connection?.id === connectionId) { + setConnection(null); + setSelectedTable(null); + setTableSearch(""); + } + notifySuccess(t("connection.deleted")); + }, + }); + + const renameConnectionMutation = useRenameConnectionMutation({ + onError: (error) => { + notifyError(error, { category: "connection", force: true }); + }, + }); + + const connectionRestoreAttemptedRef = useRef(false); + + useEffect(() => { + if (connectionRestoreAttemptedRef.current) return; + if (!autoReconnect) { connectionRestoreAttemptedRef.current = true; return; } + const list = connectionsQuery.data; + if (!list?.length) return; + if (connection) { connectionRestoreAttemptedRef.current = true; return; } + + const savedId = window.localStorage.getItem(LAST_ACTIVE_CONNECTION_KEY); + if (!savedId) { connectionRestoreAttemptedRef.current = true; return; } + + const match = list.find((c) => c.id === savedId); + const target = match ?? list[0]; + if (!target) { connectionRestoreAttemptedRef.current = true; return; } + + connectionRestoreAttemptedRef.current = true; + activateConnectionMutation.mutate(target.id); + }, [connectionsQuery.data, connection, activateConnectionMutation, autoReconnect]); + + const tablesQuery = useTablesQuery(connection?.id); + const schemaQuery = useTableSchemaQuery({ + connectionId: connection?.id, + table: selectedTable, + enabled: Boolean(connection?.id && selectedTable), + }); + const tablePropertiesQuery = useTablePropertiesQuery({ + connectionId: connection?.id, + table: selectedTable, + enabled: Boolean(connection?.id && selectedTable), + }); + const saveResultEditsMutation = useSaveResultEditsMutation({ + onError: (error) => { + notifyError(error, { category: "query", title: t("editor.failedToSave") }); + }, + }); + const deleteRowsMutation = useDeleteRowsMutation({ + onError: (error) => { + notifyError(error, { category: "query", title: t("editor.failedToDelete") }); + }, + }); + + const tablesForUi = tablesQuery.data ?? []; + const activeConnectionEngine = connection?.engine ?? "postgres"; + + const requestInsertRow = useCallback(() => setInsertRowTrigger((n) => n + 1), []); + + const handleInsertRowSuccess = useCallback(() => { + void queryClient.invalidateQueries({ queryKey: queryKeys.tableProperties(connection?.id, selectedTable) }); + }, [connection?.id, queryClient, selectedTable]); + + const handleSelectTable = (table: TableInfo) => { + setSelectedTable(table); + queryWorkspaceRef.current?.applyTablePreview(table.previewQuery); + }; + + const handleTableQuickAction = useCallback(async ( + action: TableQuickSqlAction, + connectionId: string, + table: TableInfo, + ) => { + if (action === "tableProperties") { + setTablePropertiesTarget({ connectionId, table }); + setTablePropertiesDialogOpen(true); + return; + } + if (action === "addRow") { + setSelectedTable(table); + setInsertRowTrigger((n) => n + 1); + return; + } + + setSelectedTable(table); + try { + switch (action) { + case "selectAll": + queryWorkspaceRef.current?.openTabWithSql(buildSelectAllSql(table, 200, activeConnectionEngine)); + return; + case "selectCount": + queryWorkspaceRef.current?.openTabWithSql(buildSelectCountSql(table, activeConnectionEngine)); + return; + case "insertTemplate": + case "updateTemplate": + case "deleteTemplate": { + const props = await queryClient.fetchQuery({ + queryKey: queryKeys.tableProperties(connectionId, table), + queryFn: () => veloxDbRepository.getTableProperties(connectionId, table), + }); + const pk = props.filter((c) => c.isPrimaryKey).map((c) => c.columnName); + const insertCols = props.filter(isInsertFormColumn).map((c) => c.columnName); + if (action === "insertTemplate") { + queryWorkspaceRef.current?.openTabWithSql(buildInsertTemplateSql(table, insertCols, activeConnectionEngine)); + } else if (action === "updateTemplate") { + queryWorkspaceRef.current?.openTabWithSql(buildUpdateTemplateSql(table, pk, activeConnectionEngine)); + } else { + queryWorkspaceRef.current?.openTabWithSql(buildDeleteTemplateSql(table, pk, activeConnectionEngine)); + } + return; + } + default: + return; + } + } catch (error) { + notifyError(error, { category: "query", title: "Table quick action failed", force: true }); + } + }, [queryClient, activeConnectionEngine, queryWorkspaceRef]); + + const handleSelectConnection = (nextConnection: ConnectionSummary) => { + if (connection?.id === nextConnection.id) return; + activateConnectionMutation.mutate(nextConnection.id); + }; + + const handleRefreshConnection = useCallback((connectionTarget: ConnectionSummary) => { + void (async () => { + try { await veloxDbRepository.refreshConnection(connectionTarget.id); } + catch (error) { notifyError(error, { category: "connection" }); return; } + void queryClient.invalidateQueries({ queryKey: queryKeys.connections() }); + void queryClient.invalidateQueries({ queryKey: queryKeys.databases(connectionTarget.id) }); + void queryClient.refetchQueries({ queryKey: queryKeys.databases(connectionTarget.id), type: "active" }); + void queryClient.invalidateQueries({ queryKey: queryKeys.tables(connectionTarget.id) }); + void queryClient.refetchQueries({ queryKey: queryKeys.tables(connectionTarget.id), type: "active" }); + void queryClient.invalidateQueries({ queryKey: queryKeys.queryEditorMetadata(connectionTarget.id) }); + void queryClient.refetchQueries({ queryKey: queryKeys.queryEditorMetadata(connectionTarget.id), type: "active" }); + void queryClient.invalidateQueries({ queryKey: queryKeys.foreignKeys(connectionTarget.id) }); + void queryClient.refetchQueries({ queryKey: queryKeys.foreignKeys(connectionTarget.id), type: "active" }); + if (connection?.id === connectionTarget.id) { + void queryClient.invalidateQueries({ queryKey: queryKeys.schema(connectionTarget.id, selectedTable) }); + void queryClient.refetchQueries({ queryKey: queryKeys.schema(connectionTarget.id, selectedTable), type: "active" }); + void queryClient.invalidateQueries({ queryKey: queryKeys.tableProperties(connectionTarget.id, selectedTable) }); + void queryClient.refetchQueries({ queryKey: queryKeys.tableProperties(connectionTarget.id, selectedTable), type: "active" }); + queryWorkspaceRef.current?.refreshFocusedResults(); + } + })(); + }, [connection?.id, queryClient, selectedTable, queryWorkspaceRef]); + + const handleRefreshTable = useCallback((connectionId: string, table: TableInfo) => { + void queryClient.invalidateQueries({ queryKey: queryKeys.tables(connectionId) }); + void queryClient.refetchQueries({ queryKey: queryKeys.tables(connectionId), type: "active" }); + void queryClient.invalidateQueries({ queryKey: queryKeys.schema(connectionId, table) }); + void queryClient.refetchQueries({ queryKey: queryKeys.schema(connectionId, table), type: "active" }); + void queryClient.invalidateQueries({ queryKey: queryKeys.tableProperties(connectionId, table) }); + void queryClient.refetchQueries({ queryKey: queryKeys.tableProperties(connectionId, table), type: "active" }); + void queryClient.invalidateQueries({ queryKey: queryKeys.tableIndexes(connectionId, table) }); + void queryClient.refetchQueries({ queryKey: queryKeys.tableIndexes(connectionId, table), type: "active" }); + if (connection?.id === connectionId && selectedTable?.schema === table.schema && selectedTable?.name === table.name) { + queryWorkspaceRef.current?.refreshFocusedResults(); + } + }, [connection?.id, queryClient, selectedTable?.name, selectedTable?.schema, queryWorkspaceRef]); + + const handleRenameTableRequest = useCallback((_connectionId: string, table: TableInfo) => { + setSelectedTable(table); + queryWorkspaceRef.current?.appendQuerySql( + buildRenameTableSql(table, "new_table_name", connection?.engine ?? "postgres"), + ); + }, [connection?.engine, queryWorkspaceRef]); + + const handleDeleteTableRequest = useCallback((_connectionId: string, table: TableInfo) => { + setSelectedTable(table); + queryWorkspaceRef.current?.appendQuerySql( + buildDropTableSql(table, connection?.engine ?? "postgres"), + ); + }, [connection?.engine, queryWorkspaceRef]); + + const handleRenameConnectionRequest = useCallback((connectionTarget: ConnectionSummary) => { + setRenamingConnection(connectionTarget); + }, []); + + const handleRenameConnectionConfirm = useCallback((connectionTarget: ConnectionSummary, newName: string) => { + renameConnectionMutation.mutate({ connectionId: connectionTarget.id, newName }, { + onSuccess: (updated) => { if (connection?.id === updated.id) setConnection(updated); }, + }); + setRenamingConnection(null); + }, [connection?.id, renameConnectionMutation]); + + const handleDisconnectConnectionRequest = useCallback((connectionTarget: ConnectionSummary) => { + const confirmed = window.confirm(t("connection.deleteConfirm", { name: connectionTarget.name })); + if (!confirmed) return; + deleteConnectionMutation.mutate(connectionTarget.id); + if (connection?.id === connectionTarget.id) { + setConnection(null); + setSelectedTable(null); + setTableSearch(""); + setTablePropertiesDialogOpen(false); + setTablePropertiesTarget(null); + } + }, [connection?.id, deleteConnectionMutation, t]); + + const handleCopyConnectionString = useCallback((target: ConnectionSummary) => { + const value = target.engine === "sqlite" + ? `sqlite://${target.filePath ?? target.database}` + : `${target.engine === "mysql" ? "mysql" : "postgresql"}://${target.user}@${target.host}:${target.port}/${target.database}`; + void navigator.clipboard.writeText(value); + }, []); + + const handleTruncateTable = useCallback((_connectionId: string, table: TableInfo) => { + setSelectedTable(table); + if ((connection?.engine ?? "postgres") === "mysql") { + queryWorkspaceRef.current?.appendQuerySql(`TRUNCATE TABLE ${quoteIdent(table.schema, "mysql")}.${quoteIdent(table.name, "mysql")};`); + } else if ((connection?.engine ?? "postgres") === "sqlite") { + queryWorkspaceRef.current?.appendQuerySql(`DELETE FROM ${quoteIdent(table.name, "sqlite")};`); + } else { + queryWorkspaceRef.current?.appendQuerySql(`TRUNCATE TABLE ${quoteIdent(table.schema, "postgres")}.${quoteIdent(table.name, "postgres")} RESTART IDENTITY CASCADE;`); + } + }, [connection?.engine, queryWorkspaceRef]); + + const handleCopyTableName = useCallback((_connectionId: string, table: TableInfo) => { + const engine = connection?.engine ?? "postgres"; + const value = engine === "sqlite" + ? quoteIdent(table.name, "sqlite") + : `${quoteIdent(table.schema, engine)}.${quoteIdent(table.name, engine)}`; + void navigator.clipboard.writeText(value); + }, [connection?.engine]); + + const handleRefreshDatabases = useCallback((connectionId: string) => { + void queryClient.invalidateQueries({ queryKey: queryKeys.databases(connectionId) }); + void queryClient.refetchQueries({ queryKey: queryKeys.databases(connectionId), type: "active" }); + }, [queryClient]); + + const handleCopyDatabaseName = useCallback((_connectionId: string, database: string) => { + void navigator.clipboard.writeText(database); + }, []); + + const handleActivateConnectionForTab = useCallback((connectionId: string) => { + if (connection?.id === connectionId) return; + activateConnectionMutation.mutate(connectionId); + }, [connection?.id, activateConnectionMutation]); + + const handleSaveResultEdits = async (patches: ResultEditPatch[]) => { + if (!selectedTable || !connection?.id || patches.length === 0) return; + await saveResultEditsMutation.mutateAsync({ + connectionId: connection.id, + engine: connection.engine, + table: selectedTable, + patches, + }); + queryWorkspaceRef.current?.refreshFocusedResults(); + }; + + const handleDeleteRows = async (primaryKeys: Record[]) => { + if (!selectedTable || !connection?.id || primaryKeys.length === 0) return; + await deleteRowsMutation.mutateAsync({ + connectionId: connection.id, + engine: connection.engine, + table: selectedTable, + primaryKeys, + }); + queryWorkspaceRef.current?.refreshFocusedResults(); + }; + + const connectionError = connectMutation.error ?? activateConnectionMutation.error; + const connectionErrorMessage = connectionError instanceof Error ? connectionError.message : t("connection.failedToConnect"); + const connectionsErrorMessage = connectionsQuery.error instanceof Error ? connectionsQuery.error.message : t("connection.failedToLoad"); + const tablesErrorMessage = tablesQuery.error instanceof Error ? tablesQuery.error.message : t("table.failedToLoad"); + const schemaErrorMessage = schemaQuery.error instanceof Error ? schemaQuery.error.message : t("table.failedToLoadSchema"); + const tablePropertiesErrorMessage = tablePropertiesQuery.error instanceof Error ? tablePropertiesQuery.error.message : t("table.failedToLoadProperties"); + + const primaryKeyColumns = tablePropertiesQuery.data?.filter((column) => column.isPrimaryKey).map((column) => column.columnName) ?? []; + const editableColumns = tablePropertiesQuery.data?.filter((column) => !column.isPrimaryKey).map((column) => column.columnName) ?? []; + const hasSelectedTable = Boolean(selectedTable); + const hasQueryResult = focusedQueryCaps.hasResult; + const hasPrimaryKey = primaryKeyColumns.length > 0; + const isResultSingleTableEditable = hasSelectedTable && hasQueryResult && hasPrimaryKey && !tablePropertiesQuery.isError; + const saveDisabledReason = !hasSelectedTable ? t("editor.selectTable") + : !hasQueryResult ? t("editor.runQuery") + : tablePropertiesQuery.isLoading ? t("editor.loadingMetadata") + : tablePropertiesQuery.isError ? tablePropertiesErrorMessage + : !hasPrimaryKey ? t("editor.requiresPrimaryKey") + : undefined; + + // --- Veloxy handlers --- + + const handleAskVeloxyChatSubmit = async (naturalPrompt: string, requestId: string): Promise => { + if (!connection?.id) { + const message = t("veloxy.selectConnection"); setAskVeloxyError(message); throw new Error(message); + } + if (!veloxyOpenRouterApiKey.trim()) { + const message = t("veloxy.addApiKey"); setAskVeloxyError(message); throw new Error(message); + } + if (!veloxyModel.trim()) { + const message = t("veloxy.chooseModel"); setAskVeloxyError(message); throw new Error(message); + } + setAskVeloxyPending(true); setAskVeloxyError(null); + try { + return await veloxDbRepository.chatWithDb({ + connectionId: connection.id, naturalPrompt, requestId, + targetTable: selectedTable ? { schema: selectedTable.schema, name: selectedTable.name } : undefined, + providerConfig: { apiKey: veloxyOpenRouterApiKey, model: veloxyModel, baseUrl: veloxyBaseUrl }, + maxRows: useSettings.getState().maxQueryRows, + }); + } catch (error) { + const message = error instanceof Error ? error.message : t("veloxy.chatFailed"); + setAskVeloxyError(message); + notifyError(error, { category: "query", title: t("veloxy.chatFailed") }); + throw error instanceof Error ? error : new Error(message); + } finally { setAskVeloxyPending(false); } + }; + + const handleCancelVeloxyRequest = async () => { + try { await veloxDbRepository.cancelVeloxyRequest(); } + catch (error) { + const message = error instanceof Error ? error.message : "Failed to stop Veloxy."; + setAskVeloxyError(message); + } + }; + + const handleAskVeloxyActionSubmit = async (naturalPrompt: string) => { + if (!connection?.id) { + const message = t("veloxy.selectConnection"); setAskVeloxyError(message); throw new Error(message); + } + if (!veloxyOpenRouterApiKey.trim()) { + const message = t("veloxy.addApiKey"); setAskVeloxyError(message); throw new Error(message); + } + if (!veloxyModel.trim()) { + const message = t("veloxy.chooseModel"); setAskVeloxyError(message); throw new Error(message); + } + setAskVeloxyPending(true); setAskVeloxyError(null); + try { + const response = await veloxDbRepository.generateSqlFromNl({ + connectionId: connection.id, naturalPrompt, + targetTable: selectedTable ? { schema: selectedTable.schema, name: selectedTable.name } : undefined, + providerConfig: { apiKey: veloxyOpenRouterApiKey, model: veloxyModel, baseUrl: veloxyBaseUrl }, + maxRows: useSettings.getState().maxQueryRows, + }); + const sql = response.sql.trim(); + const lower = sql.toLowerCase(); + const isReadIntent = response.intent === "select"; + const isLikelyLarge = sql.length > 1800 || (lower.includes("select") && !lower.includes(" limit ")) || /\bcross\s+join\b|\bpg_sleep\s*\(/i.test(lower); + const canAutoRun = isReadIntent && !isLikelyLarge; + + if (canAutoRun) { + queryWorkspaceRef.current?.openTabWithSqlAndRun(sql); + notifySuccess(t("veloxy.generatedSql"), t("veloxy.autoRan")); + return { response, decision: "auto-ran" as const }; + } + queryWorkspaceRef.current?.openTabWithSql(sql); + return { + response, + decision: "needs-confirmation" as const, + decisionReason: isReadIntent ? t("veloxy.needsConfirmation") : t("veloxy.nonReadConfirmation"), + pendingSql: sql, + }; + } catch (error) { + const message = error instanceof Error ? error.message : t("veloxy.generateFailed"); + setAskVeloxyError(message); + notifyError(error, { category: "query", title: t("veloxy.generateFailed") }); + throw error instanceof Error ? error : new Error(message); + } finally { setAskVeloxyPending(false); } + }; + + const handleLoadVeloxyConversation = async (): Promise => { + if (!connection?.id) return { messages: [] }; + try { return await veloxDbRepository.loadVeloxyConversation(connection.id); } + catch (error) { + const message = error instanceof Error ? error.message : t("veloxy.loadFailed"); + setAskVeloxyError(message); + return { messages: [] }; + } + }; + + const handleClearVeloxyConversation = async () => { + if (!connection?.id) return; + try { await veloxDbRepository.clearVeloxyConversation(connection.id); } + catch (error) { + const message = error instanceof Error ? error.message : t("veloxy.clearFailed"); + setAskVeloxyError(message); + throw error; + } + }; + + return { + // State + connection, setConnection, + focusedQueryCaps, setFocusedQueryCaps, + tableSearch, setTableSearch, + selectedTable, + settingsOpen, setSettingsOpen, + commandPaletteOpen, setCommandPaletteOpen, + connectionDialogOpen, setConnectionDialogOpen, + renamingConnection, setRenamingConnection, + isSidebarCollapsed, setIsSidebarCollapsed, + sidebarWidth, setSidebarWidth, + resultsHeight, setResultsHeight, + tablePropertiesDialogOpen, setTablePropertiesDialogOpen, + tablePropertiesTarget, setTablePropertiesTarget, + insertRowTrigger, + mainWorkspace, setMainWorkspace, + askVeloxyPending, setAskVeloxyPending, + askVeloxyError, setAskVeloxyError, + // Derived + isDark, + veloxyOpenRouterApiKey, veloxyModel, veloxyBaseUrl, + // Queries + connectionsQuery, tablesQuery, schemaQuery, tablePropertiesQuery, + saveResultEditsMutation, deleteRowsMutation, + connectMutation, activateConnectionMutation, + deleteConnectionMutation, renameConnectionMutation, + // Computed + tablesForUi, activeConnectionEngine, + connectionError, connectionErrorMessage, + connectionsErrorMessage, tablesErrorMessage, + schemaErrorMessage, tablePropertiesErrorMessage, + primaryKeyColumns, editableColumns, + isResultSingleTableEditable, saveDisabledReason, + // Handlers + handleSelectTable, + handleTableQuickAction, + handleSelectConnection, + handleRefreshConnection, + handleRefreshTable, + handleRenameTableRequest, + handleDeleteTableRequest, + handleRenameConnectionRequest, + handleRenameConnectionConfirm, + handleDisconnectConnectionRequest, + handleCopyConnectionString, + handleTruncateTable, + handleCopyTableName, + handleRefreshDatabases, + handleCopyDatabaseName, + handleActivateConnectionForTab, + handleSaveResultEdits, + handleDeleteRows, + requestInsertRow, + handleInsertRowSuccess, + handleAskVeloxyChatSubmit, + handleCancelVeloxyRequest, + handleAskVeloxyActionSubmit, + handleLoadVeloxyConversation, + handleClearVeloxyConversation, + }; +} diff --git a/src/lib/sql-intent.ts b/src/lib/sql-intent.ts index 120b967..ba67849 100644 --- a/src/lib/sql-intent.ts +++ b/src/lib/sql-intent.ts @@ -25,7 +25,8 @@ export function classifySqlIntent(sql: string): SqlIntent { const TRANSACTION_CONTROL = ['begin', 'commit', 'rollback', 'start', 'savepoint', 'release'] /** True when every statement in `sql` is read-only (select/explain). */ -export function isReadOnlySql(sql: string): boolean { +export function isReadOnlySql(sql: string, engine?: string): boolean { + if (engine === "mongo") return true; // MongoDB find() queries are always read-only let sawStatement = false for (const raw of sql.split(';')) { const statement = raw.trim() From e3d6586359a207d25b5603789c5623304542585c Mon Sep 17 00:00:00 2001 From: abeni16 Date: Tue, 23 Jun 2026 14:02:22 +0300 Subject: [PATCH 02/14] feat: add support for DuckDB and Redis database engines This commit introduces support for DuckDB and Redis as database engines, updating the relevant types and connection handling. It modifies the connection dialogs to accommodate these new engines, enhances the connection string parsing and building functions, and updates the application state management to include DuckDB and Redis connections. Additionally, it implements export functionality for DuckDB results in both CSV and JSON formats, improving the overall database management capabilities of the application. --- build/index.html | 4 +- src-tauri/Cargo.lock | 608 +++++++- src-tauri/Cargo.toml | 2 + src-tauri/src/commands/connections.rs | 47 +- src-tauri/src/commands/ddl.rs | 24 +- src-tauri/src/commands/duckdb.rs | 505 +++++++ src-tauri/src/commands/lint.rs | 22 + src-tauri/src/commands/mod.rs | 16 +- src-tauri/src/commands/mongo.rs | 488 +++++- src-tauri/src/commands/query.rs | 20 + src-tauri/src/commands/redis.rs | 132 ++ src-tauri/src/commands/table_props.rs | 36 +- src-tauri/src/db.rs | 241 ++- src-tauri/src/export.rs | 51 +- src-tauri/src/lib.rs | 15 +- src-tauri/src/models.rs | 2 + src/data/types.ts | 2 +- .../components/ConnectionDialog.tsx | 1327 +++++++---------- .../components/ConnectionsSidebarTree.tsx | 2 + .../model/components/ModelWorkspace.tsx | 22 +- .../queries/components/AskVeloxyDialog.tsx | 2 - .../queries/components/QueryWorkspace.tsx | 2 +- src/features/queries/components/SqlEditor.tsx | 2 +- .../queries/cross-engine-format.test.ts | 155 ++ src/hooks/useAppState.ts | 16 + src/lib/connection-string.test.ts | 137 ++ src/lib/connection-string.ts | 28 + src/lib/sql-intent.test.ts | 18 + src/lib/sql-intent.ts | 1 + 29 files changed, 3082 insertions(+), 845 deletions(-) create mode 100644 src-tauri/src/commands/duckdb.rs create mode 100644 src-tauri/src/commands/redis.rs create mode 100644 src/features/queries/cross-engine-format.test.ts diff --git a/build/index.html b/build/index.html index cac414f..78715b8 100644 --- a/build/index.html +++ b/build/index.html @@ -5,10 +5,10 @@ veloxdb - + - +
diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 4b642dd..c6e57db 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -26,6 +26,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5a15f179cd60c4584b8a8c596927aadc462e27f2ca70c04e0071964a73ba7a75" dependencies = [ "cfg-if", + "const-random", "getrandom 0.3.4", "once_cell", "version_check", @@ -94,6 +95,24 @@ version = "1.0.102" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c" +[[package]] +name = "arbitrary" +version = "1.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c3d036a3c4ab069c7b410a2ce876bd74808d2d0888a82667669f8e783a898bf1" +dependencies = [ + "derive_arbitrary", +] + +[[package]] +name = "arc-swap" +version = "1.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6a3a1fd6f75306b68087b831f025c712524bcb19aad54e557b1129cfa0a2b207" +dependencies = [ + "rustversion", +] + [[package]] name = "arrayref" version = "0.3.9" @@ -106,6 +125,169 @@ version = "0.7.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7c02d123df017efcdfbd739ef81735b36c5ba83ec3c59c80a9d7ecc718f92e50" +[[package]] +name = "arrow" +version = "58.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "378530e55cd479eda3c14eb345310799717e6f76d0c332041e8487022166b471" +dependencies = [ + "arrow-arith", + "arrow-array", + "arrow-buffer", + "arrow-cast", + "arrow-data", + "arrow-ord", + "arrow-row", + "arrow-schema", + "arrow-select", + "arrow-string", +] + +[[package]] +name = "arrow-arith" +version = "58.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a0ab212d2c1886e802f51c5212d78ebbcbb0bec980fff9dadc1eb8d45cd0b738" +dependencies = [ + "arrow-array", + "arrow-buffer", + "arrow-data", + "arrow-schema", + "chrono", + "num-traits", +] + +[[package]] +name = "arrow-array" +version = "58.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cfd33d3e92f207444098c75b42de99d329562be0cf686b307b097cc52b4e999e" +dependencies = [ + "ahash 0.8.12", + "arrow-buffer", + "arrow-data", + "arrow-schema", + "chrono", + "half", + "hashbrown 0.17.1", + "num-complex", + "num-integer", + "num-traits", +] + +[[package]] +name = "arrow-buffer" +version = "58.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0c6cd424c2693bcdbc150d843dc9d4d137dd2de4782ce6df491ad11a3a0416c0" +dependencies = [ + "bytes", + "half", + "num-bigint", + "num-traits", +] + +[[package]] +name = "arrow-cast" +version = "58.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4c5aefb56a2c02e9e2b30746241058b85f8983f0fcff2ba0c6d09006e1cded7f" +dependencies = [ + "arrow-array", + "arrow-buffer", + "arrow-data", + "arrow-ord", + "arrow-schema", + "arrow-select", + "atoi", + "base64 0.22.1", + "chrono", + "comfy-table", + "half", + "lexical-core", + "num-traits", + "ryu", +] + +[[package]] +name = "arrow-data" +version = "58.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3c88210023a2bfee1896af366309a3028fc3bcbd6515fa29a7990ee1baa08ee0" +dependencies = [ + "arrow-buffer", + "arrow-schema", + "half", + "num-integer", + "num-traits", +] + +[[package]] +name = "arrow-ord" +version = "58.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1bffd8fd2579286a5d63bac898159873e5094a79009940bcb42bbfce4f19f1d0" +dependencies = [ + "arrow-array", + "arrow-buffer", + "arrow-data", + "arrow-schema", + "arrow-select", +] + +[[package]] +name = "arrow-row" +version = "58.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bab5994731204603c73ba69267616c50f80780774c6bb0476f1f830625115e0c" +dependencies = [ + "arrow-array", + "arrow-buffer", + "arrow-data", + "arrow-schema", + "half", +] + +[[package]] +name = "arrow-schema" +version = "58.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f633dbfdf39c039ada1bf9e34c694816eb71fbb7dc78f613993b7245e078a1ed" +dependencies = [ + "bitflags 2.11.0", +] + +[[package]] +name = "arrow-select" +version = "58.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8cd065c54172ac787cf3f2f8d4107e0d3fdc26edba76fdf4f4cc170258942222" +dependencies = [ + "ahash 0.8.12", + "arrow-array", + "arrow-buffer", + "arrow-data", + "arrow-schema", + "num-traits", +] + +[[package]] +name = "arrow-string" +version = "58.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "29dd7cda3ab9692f43a2e4acc444d760cc17b12bb6d8232ddf64e9bab7c06b42" +dependencies = [ + "arrow-array", + "arrow-buffer", + "arrow-data", + "arrow-schema", + "arrow-select", + "memchr", + "num-traits", + "regex", + "regex-syntax", +] + [[package]] name = "async-trait" version = "0.1.89" @@ -451,6 +633,12 @@ dependencies = [ "toml 0.9.12+spec-1.1.0", ] +[[package]] +name = "cast" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "37b2a672a2cb129a2e41c10b1224bb368f9f37a2b16b612598138befd7b37eb5" + [[package]] name = "cc" version = "1.2.57" @@ -458,6 +646,8 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7a0dd1ca384932ff3641c8718a02769f1698e7563dc6974ffd03346116310423" dependencies = [ "find-msvc-tools", + "jobserver", + "libc", "shlex", ] @@ -538,7 +728,22 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ba5a308b75df32fe02788e748662718f03fde005016435c444eea572398219fd" dependencies = [ "bytes", + "futures-core", "memchr", + "pin-project-lite", + "tokio", + "tokio-util", +] + +[[package]] +name = "comfy-table" +version = "7.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4a65ebfec4fb190b6f90e944a817d60499ee0744e582530e2c9900a22e591d9a" +dependencies = [ + "crossterm", + "unicode-segmentation", + "unicode-width", ] [[package]] @@ -741,6 +946,28 @@ version = "0.8.21" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d0a5c400df2834b80a4c3327b3aad3a4c4cd4de0629063962b03235697506a28" +[[package]] +name = "crossterm" +version = "0.28.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "829d955a0bb380ef178a640b91779e3987da38c9aea133b20614cfed8cdea9c6" +dependencies = [ + "bitflags 2.11.0", + "crossterm_winapi", + "parking_lot", + "rustix 0.38.44", + "winapi", +] + +[[package]] +name = "crossterm_winapi" +version = "0.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "acdd7c62a3665c7f6830a51635d9ac9b23ed385797f70a83bb8bafe9c572ab2b" +dependencies = [ + "winapi", +] + [[package]] name = "crunchy" version = "0.2.4" @@ -982,6 +1209,17 @@ dependencies = [ "syn 2.0.117", ] +[[package]] +name = "derive_arbitrary" +version = "1.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e567bd82dcff979e4b03460c307b3cdc9e96fde3d73bed1496d2bc75d9dd62a" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.117", +] + [[package]] name = "derive_more" version = "0.99.20" @@ -1157,6 +1395,24 @@ version = "0.0.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f678cf4a922c215c63e0de95eb1ff08a958a81d47e485cf9da1e27bf6305cfa5" +[[package]] +name = "duckdb" +version = "1.10504.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b9b997383221efd999a362448a0866c54b9a2cb7d4da8cd4903ead4df9f8eaa" +dependencies = [ + "arrow", + "cast", + "comfy-table", + "fallible-iterator 0.3.0", + "fallible-streaming-iterator", + "hashlink", + "libduckdb-sys", + "num-integer", + "rust_decimal", + "strum", +] + [[package]] name = "dunce" version = "1.0.5" @@ -1281,6 +1537,18 @@ version = "0.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4443176a9f2c162692bd3d352d745ef9413eec5782a80d8fd6f8a1ac692a07f7" +[[package]] +name = "fallible-iterator" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2acce4a10f12dc2fb14a218589d4f1f62ef011b2d0cc4b3cb1bba8e94da14649" + +[[package]] +name = "fallible-streaming-iterator" +version = "0.1.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7360491ce676a36bf9bb3c56c1aa791658183a54d2744120f27285738d90465a" + [[package]] name = "fastrand" version = "2.3.0" @@ -1315,6 +1583,16 @@ dependencies = [ "rustc_version", ] +[[package]] +name = "filetime" +version = "0.2.29" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5c287a33c7f0a620c38e641e7f60827713987b3c0f26e8ddc9462cc69cf75759" +dependencies = [ + "cfg-if", + "libc", +] + [[package]] name = "find-msvc-tools" version = "0.1.9" @@ -1335,6 +1613,7 @@ checksum = "843fba2746e448b37e26a819579957415c8cef339bf08564fe8b7ddbd959573c" dependencies = [ "crc32fast", "miniz_oxide", + "zlib-rs", ] [[package]] @@ -1447,6 +1726,21 @@ dependencies = [ "new_debug_unreachable", ] +[[package]] +name = "futures" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b147ee9d1f6d097cef9ce628cd2ee62288d963e16fb287bd9286455b241382d" +dependencies = [ + "futures-channel", + "futures-core", + "futures-executor", + "futures-io", + "futures-sink", + "futures-task", + "futures-util", +] + [[package]] name = "futures-channel" version = "0.3.32" @@ -1520,6 +1814,7 @@ version = "0.3.32" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "389ca41296e6190b48053de0321d02a77f32f8a5d2461dd38762c0593805c6d6" dependencies = [ + "futures-channel", "futures-core", "futures-io", "futures-macro", @@ -1858,6 +2153,18 @@ dependencies = [ "syn 2.0.117", ] +[[package]] +name = "half" +version = "2.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ea2d84b969582b4b1864a92dc5d27cd2b77b622a8d79306834f1be5ba20d84b" +dependencies = [ + "cfg-if", + "crunchy", + "num-traits", + "zerocopy", +] + [[package]] name = "hashbrown" version = "0.12.3" @@ -1884,6 +2191,12 @@ version = "0.16.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "841d1cc9bed7f9236f321df977030373f4a4163ae1a7dbfe1a51a2c1a51d9100" +[[package]] +name = "hashbrown" +version = "0.17.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ed5909b6e89a2db4456e54cd5f673791d7eca6732202bbf2a9cc504fe2f9b84a" + [[package]] name = "hashlink" version = "0.10.0" @@ -2129,7 +2442,7 @@ dependencies = [ "libc", "percent-encoding", "pin-project-lite", - "socket2", + "socket2 0.6.3", "tokio", "tower-service", "tracing", @@ -2337,7 +2650,7 @@ version = "0.3.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4d40460c0ce33d6ce4b0630ad68ff63d6661961c48b6dba35e5a4d81cfb48222" dependencies = [ - "socket2", + "socket2 0.6.3", "widestring", "windows-registry", "windows-result 0.4.1", @@ -2463,6 +2776,16 @@ dependencies = [ "syn 2.0.117", ] +[[package]] +name = "jobserver" +version = "0.1.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9afb3de4395d6b3e67a780b6de64b51c978ecf11cb9a462c66be7d4ca9039d33" +dependencies = [ + "getrandom 0.3.4", + "libc", +] + [[package]] name = "js-sys" version = "0.3.91" @@ -2554,6 +2877,63 @@ version = "0.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "09edd9e8b54e49e587e4f6295a7d29c3ea94d469cb40ab8ca70b288248a81db2" +[[package]] +name = "lexical-core" +version = "1.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7d8d125a277f807e55a77304455eb7b1cb52f2b18c143b60e766c120bd64a594" +dependencies = [ + "lexical-parse-float", + "lexical-parse-integer", + "lexical-util", + "lexical-write-float", + "lexical-write-integer", +] + +[[package]] +name = "lexical-parse-float" +version = "1.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "52a9f232fbd6f550bc0137dcb5f99ab674071ac2d690ac69704593cb4abbea56" +dependencies = [ + "lexical-parse-integer", + "lexical-util", +] + +[[package]] +name = "lexical-parse-integer" +version = "1.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9a7a039f8fb9c19c996cd7b2fcce303c1b2874fe1aca544edc85c4a5f8489b34" +dependencies = [ + "lexical-util", +] + +[[package]] +name = "lexical-util" +version = "1.0.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2604dd126bb14f13fb5d1bd6a66155079cb9fa655b37f875b3a742c705dbed17" + +[[package]] +name = "lexical-write-float" +version = "1.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "50c438c87c013188d415fbabbb1dceb44249ab81664efbd31b14ae55dabb6361" +dependencies = [ + "lexical-util", + "lexical-write-integer", +] + +[[package]] +name = "lexical-write-integer" +version = "1.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "409851a618475d2d5796377cad353802345cba92c867d9fbcde9cf4eac4e14df" +dependencies = [ + "lexical-util", +] + [[package]] name = "libappindicator" version = "0.9.0" @@ -2593,6 +2973,23 @@ dependencies = [ "pkg-config", ] +[[package]] +name = "libduckdb-sys" +version = "1.10504.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "94e67523a40ee3da30411e00e243990b3760d91877202ab2f39c03ab34f8dbfc" +dependencies = [ + "cc", + "flate2", + "pkg-config", + "reqwest 0.12.28", + "serde", + "serde_json", + "tar", + "vcpkg", + "zip", +] + [[package]] name = "libloading" version = "0.7.4" @@ -2638,6 +3035,18 @@ version = "0.5.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0717cef1bc8b636c6e1c1bbdefc09e6322da8a9321966e8928ef80d20f7f770f" +[[package]] +name = "linux-raw-sys" +version = "0.4.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d26c52dbd32dccf2d10cac7725f8eae5296885fb5703b261f7d0a0739ec807ab" + +[[package]] +name = "linux-raw-sys" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a66949e030da00e8c7d4434b251670a91556f4144941d37452769c25d58a53" + [[package]] name = "litemap" version = "0.8.1" @@ -2916,7 +3325,7 @@ dependencies = [ "serde_with", "sha1", "sha2", - "socket2", + "socket2 0.6.3", "stringprep", "strsim", "take_mut", @@ -3004,6 +3413,16 @@ version = "0.1.14" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "72ef4a56884ca558e5ddb05a1d1e7e1bfd9a68d9ed024c21704cc98872dae1bb" +[[package]] +name = "num-bigint" +version = "0.4.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a5e44f723f1133c9deac646763579fdb3ac745e418f2a7af9cd0c431da1f20b9" +dependencies = [ + "num-integer", + "num-traits", +] + [[package]] name = "num-bigint-dig" version = "0.8.6" @@ -3020,6 +3439,15 @@ dependencies = [ "zeroize", ] +[[package]] +name = "num-complex" +version = "0.4.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "73f88a1307638156682bada9d7604135552957b7818057dcef22705b4d509495" +dependencies = [ + "num-traits", +] + [[package]] name = "num-conv" version = "0.2.0" @@ -3706,7 +4134,7 @@ dependencies = [ "base64 0.22.1", "byteorder", "bytes", - "fallible-iterator", + "fallible-iterator 0.2.0", "hmac", "md-5", "memchr", @@ -3722,7 +4150,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "54b858f82211e84682fecd373f68e1ceae642d8d751a1ebd13f33de6257b3e20" dependencies = [ "bytes", - "fallible-iterator", + "fallible-iterator 0.2.0", "postgres-protocol", ] @@ -3905,7 +4333,7 @@ dependencies = [ "quinn-udp", "rustc-hash", "rustls", - "socket2", + "socket2 0.6.3", "thiserror 2.0.18", "tokio", "tracing", @@ -3942,7 +4370,7 @@ dependencies = [ "cfg_aliases", "libc", "once_cell", - "socket2", + "socket2 0.6.3", "tracing", "windows-sys 0.60.2", ] @@ -4107,6 +4535,30 @@ version = "0.6.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "20675572f6f24e9e76ef639bc5552774ed45f1c30e2951e1e99c59888861c539" +[[package]] +name = "redis" +version = "0.25.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5e46922bd01fefcfdcf58d9cd626da082bb2cde27211920dacfde6b2ecf9a35b" +dependencies = [ + "arc-swap", + "async-trait", + "bytes", + "combine", + "futures", + "futures-util", + "itoa", + "percent-encoding", + "pin-project-lite", + "ryu", + "sha1_smol", + "socket2 0.5.10", + "tokio", + "tokio-retry", + "tokio-util", + "url", +] + [[package]] name = "redox_syscall" version = "0.5.18" @@ -4202,6 +4654,7 @@ checksum = "eddd3ca559203180a307f12d114c268abf583f59b03cb906fd0b3ff8646c1147" dependencies = [ "base64 0.22.1", "bytes", + "futures-channel", "futures-core", "futures-util", "http", @@ -4435,6 +4888,32 @@ dependencies = [ "semver", ] +[[package]] +name = "rustix" +version = "0.38.44" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fdb5bc1ae2baa591800df16c9ca78619bf65c0488b41b96ccec5d11220d8c154" +dependencies = [ + "bitflags 2.11.0", + "errno", + "libc", + "linux-raw-sys 0.4.15", + "windows-sys 0.59.0", +] + +[[package]] +name = "rustix" +version = "1.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6fe4565b9518b83ef4f91bb47ce29620ca828bd32cb7e408f0062e9930ba190" +dependencies = [ + "bitflags 2.11.0", + "errno", + "libc", + "linux-raw-sys 0.12.1", + "windows-sys 0.61.2", +] + [[package]] name = "rustls" version = "0.23.38" @@ -4830,6 +5309,12 @@ dependencies = [ "digest", ] +[[package]] +name = "sha1_smol" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bbfa15b3dddfee50a0fff136974b3e1bde555604ba463834a7eb7deb6417705d" + [[package]] name = "sha2" version = "0.10.9" @@ -4934,6 +5419,16 @@ dependencies = [ "serde", ] +[[package]] +name = "socket2" +version = "0.5.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e22376abed350d73dd1cd119b57ffccad95b4e585a7cda43e286245ce23c0678" +dependencies = [ + "libc", + "windows-sys 0.52.0", +] + [[package]] name = "socket2" version = "0.6.3" @@ -5286,6 +5781,27 @@ version = "0.11.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7da8b5736845d9f2fcb837ea5d9e2628564b3b043a70948a3f0b778838c5fb4f" +[[package]] +name = "strum" +version = "0.27.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "af23d6f6c1a224baef9d3f61e287d2761385a5b88fdab4eb4c6f11aeb54c4bcf" +dependencies = [ + "strum_macros", +] + +[[package]] +name = "strum_macros" +version = "0.27.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7695ce3845ea4b33927c055a39dc438a45b059f7c1b3d91d38d10355fb8cbca7" +dependencies = [ + "heck 0.5.0", + "proc-macro2", + "quote", + "syn 2.0.117", +] + [[package]] name = "subtle" version = "2.6.1" @@ -5458,6 +5974,17 @@ version = "1.0.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "55937e1799185b12863d447f42597ed69d9928686b8d88a1df17376a097d8369" +[[package]] +name = "tar" +version = "0.4.46" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f6221d9a6003c78398e3b239969f352578258df48c8eb051caadae0015bc840" +dependencies = [ + "filetime", + "libc", + "xattr", +] + [[package]] name = "target-lexicon" version = "0.12.16" @@ -5963,7 +6490,7 @@ dependencies = [ "parking_lot", "pin-project-lite", "signal-hook-registry", - "socket2", + "socket2 0.6.3", "tokio-macros", "windows-sys 0.61.2", ] @@ -5988,7 +6515,7 @@ dependencies = [ "async-trait", "byteorder", "bytes", - "fallible-iterator", + "fallible-iterator 0.2.0", "futures-channel", "futures-util", "log", @@ -5999,7 +6526,7 @@ dependencies = [ "postgres-protocol", "postgres-types", "rand 0.9.2", - "socket2", + "socket2 0.6.3", "tokio", "tokio-util", "whoami 2.1.1", @@ -6020,6 +6547,17 @@ dependencies = [ "x509-cert", ] +[[package]] +name = "tokio-retry" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4a129d95275ebf4c493ec53bf0f8cd95f5ac161bc4f381700809a54f595d4470" +dependencies = [ + "pin-project-lite", + "rand 0.10.1", + "tokio", +] + [[package]] name = "tokio-rustls" version = "0.26.4" @@ -6411,6 +6949,12 @@ version = "0.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b1d386ff53b415b7fe27b50bb44679e2cc4660272694b7b6f3326d8480823a94" +[[package]] +name = "unicode-width" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b4ac048d71ede7ee76d585517add45da530660ef4390e49b098733c6e897f254" + [[package]] name = "unicode-xid" version = "0.2.6" @@ -6532,6 +7076,7 @@ dependencies = [ "chrono", "csv", "deadpool-postgres", + "duckdb", "futures-util", "hex", "keyring", @@ -6539,6 +7084,7 @@ dependencies = [ "mongodb", "printpdf", "rand 0.8.5", + "redis", "reqwest 0.12.28", "resvg", "rustls", @@ -7668,6 +8214,16 @@ dependencies = [ "tls_codec", ] +[[package]] +name = "xattr" +version = "1.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32e45ad4206f6d2479085147f02bc2ef834ac85886624a23575ae137c8aa8156" +dependencies = [ + "libc", + "rustix 1.1.4", +] + [[package]] name = "xmlwriter" version = "0.1.0" @@ -7791,12 +8347,44 @@ dependencies = [ "syn 2.0.117", ] +[[package]] +name = "zip" +version = "6.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "eb2a05c7c36fde6c09b08576c9f7fb4cda705990f73b58fe011abf7dfb24168b" +dependencies = [ + "arbitrary", + "crc32fast", + "flate2", + "indexmap 2.13.0", + "memchr", + "zopfli", +] + +[[package]] +name = "zlib-rs" +version = "0.6.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "977347db8caa080403f6b6b7c1cda9479a8e869316f7e13a59b19076a40f94e3" + [[package]] name = "zmij" version = "1.0.21" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b8848ee67ecc8aedbaf3e4122217aff892639231befc6a1b58d29fff4c2cabaa" +[[package]] +name = "zopfli" +version = "0.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f05cd8797d63865425ff89b5c4a48804f35ba0ce8d125800027ad6017d2b5249" +dependencies = [ + "bumpalo", + "crc32fast", + "log", + "simd-adler32", +] + [[package]] name = "zune-core" version = "0.4.12" diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 654aecd..ed45aef 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -48,3 +48,5 @@ futures-util = "0.3" chrono = "0.4.45" mongodb = "3" bson = "2" +duckdb = { version = "1", features = ["bundled"] } +redis = { version = "0.25", features = ["tokio-comp", "connection-manager"] } diff --git a/src-tauri/src/commands/connections.rs b/src-tauri/src/commands/connections.rs index 7255d16..7901550 100644 --- a/src-tauri/src/commands/connections.rs +++ b/src-tauri/src/commands/connections.rs @@ -2,8 +2,10 @@ use uuid::Uuid; use tauri::{AppHandle, State}; use crate::db::{ - build_mysql_pool, build_mysql_pool_custom, build_mongo_connection_string, build_pool, build_pool_custom, build_sqlite_pool, - disconnect_connection, drop_pool, get_or_create_mongo_client, get_or_create_mysql_pool, get_or_create_sqlite_pool, + build_duckdb_connection, build_mysql_pool, build_mysql_pool_custom, build_mongo_connection_string, + build_pool, build_pool_custom, build_sqlite_pool, disconnect_connection, + drop_pool, get_or_create_duckdb_connection, get_or_create_mongo_client, + get_or_create_mysql_pool, get_or_create_redis_client, get_or_create_sqlite_pool, load_connection, persist_connection_with_password, refresh_connection_pools, resolve_connection_engine, with_pool_client_retry, AppState, DEFAULT_MYSQL_PORT, }; @@ -97,6 +99,19 @@ pub async fn connect_db( .map_err(|e| format!("MongoDB ping failed: {}", e))?; state.mongo_clients.write().await.insert(connection_id.clone(), client); } + DatabaseEngine::Duckdb => { + let conn = build_duckdb_connection(&input)?; + state.duckdb_connections.write().await.insert(connection_id.clone(), tokio::sync::Mutex::new(conn)); + } + DatabaseEngine::Redis => { + let url = crate::db::build_redis_url(&input); + let client = redis::Client::open(url).map_err(|e| format!("Redis connection failed: {}", e))?; + let mut conn = redis::aio::ConnectionManager::new(client).await + .map_err(|e| format!("Redis connection failed: {}", e))?; + redis::cmd("PING").query_async::<_, String>(&mut conn).await + .map_err(|e| format!("Redis ping failed: {}", e))?; + state.redis_clients.write().await.insert(connection_id.clone(), conn); + } } let stored_connection = StoredConnection::from_input(connection_id.clone(), input.clone()); @@ -142,6 +157,14 @@ pub async fn set_active_connection( client.database("admin").run_command(doc! { "ping": 1 }).await .map_err(|e| format!("MongoDB ping failed: {}", e))?; } + DatabaseEngine::Duckdb => { + get_or_create_duckdb_connection(&app, &state, &connection_id).await?; + } + DatabaseEngine::Redis => { + let mut client = get_or_create_redis_client(&app, &state, &connection_id).await?; + redis::cmd("PING").query_async::<_, String>(&mut client).await + .map_err(|e| format!("Redis ping failed: {}", e))?; + } } *state.active_connection_id.write().await = Some(connection_id); @@ -179,6 +202,16 @@ pub async fn ping_connection( .map_err(|error| error.to_string())?; Ok(()) } + DatabaseEngine::Duckdb => { + get_or_create_duckdb_connection(&app, &state, &connection_id).await?; + Ok(()) + } + DatabaseEngine::Redis => { + let mut client = get_or_create_redis_client(&app, &state, &connection_id).await?; + redis::cmd("PING").query_async::<_, String>(&mut client).await + .map_err(|error| error.to_string())?; + Ok(()) + } } } @@ -264,6 +297,8 @@ pub async fn list_databases( .map_err(|e| format!("Failed to list MongoDB databases: {}", e))?; Ok(db_names.into_iter().map(|name| DatabaseInfo { name }).collect()) } + DatabaseEngine::Duckdb => Ok(vec![DatabaseInfo { name: "main".to_string() }]), + DatabaseEngine::Redis => Ok(vec![DatabaseInfo { name: "0".to_string() }]), } } @@ -345,6 +380,14 @@ pub async fn switch_database( client.database(&input.database).run_command(doc! { "ping": 1 }).await .map_err(|e| format!("MongoDB ping failed: {}", e))?; } + DatabaseEngine::Duckdb => { + get_or_create_duckdb_connection(&app, &state, &input.connection_id).await?; + } + DatabaseEngine::Redis => { + let mut client = get_or_create_redis_client(&app, &state, &input.connection_id).await?; + redis::cmd("PING").query_async::<_, String>(&mut client).await + .map_err(|e| format!("Redis ping failed: {}", e))?; + } } *state.active_connection_id.write().await = Some(input.connection_id); diff --git a/src-tauri/src/commands/ddl.rs b/src-tauri/src/commands/ddl.rs index 8db338f..d135d39 100644 --- a/src-tauri/src/commands/ddl.rs +++ b/src-tauri/src/commands/ddl.rs @@ -1,8 +1,8 @@ use tauri::{AppHandle, State}; use crate::db::{ - get_or_create_mysql_pool, get_or_create_sqlite_pool, resolve_connection_engine, - with_pool_client_retry, AppState, + get_or_create_duckdb_connection, get_or_create_mysql_pool, get_or_create_sqlite_pool, + resolve_connection_engine, with_pool_client_retry, AppState, }; use crate::models::{DatabaseEngine, DdlBatchRequest, DdlStatementRequest}; use crate::pg_error::map_pg_err; @@ -53,6 +53,17 @@ pub async fn execute_ddl_transaction( DatabaseEngine::Mongo => { Err("MongoDB does not support DDL transactions.".to_string()) } + DatabaseEngine::Duckdb => { + get_or_create_duckdb_connection(&app, &state, &connection_id).await?; + let conns = state.duckdb_connections.read().await; + let conn_mutex = conns.get(&connection_id).ok_or("DuckDB connection not found")?; + let conn = conn_mutex.lock().await; + for sql in input.statements.iter().map(|s| s.trim()).filter(|s| !s.is_empty()) { + conn.execute(sql, []).map_err(|e| format!("DuckDB DDL failed: {}", e))?; + } + Ok(()) + } + DatabaseEngine::Redis => Err("Not supported for Redis.".to_string()), } } @@ -89,5 +100,14 @@ pub async fn execute_ddl_statement( DatabaseEngine::Mongo => { Err("MongoDB does not support DDL statements.".to_string()) } + DatabaseEngine::Duckdb => { + get_or_create_duckdb_connection(&app, &state, &connection_id).await?; + let conns = state.duckdb_connections.read().await; + let conn_mutex = conns.get(&connection_id).ok_or("DuckDB connection not found")?; + let conn = conn_mutex.lock().await; + conn.execute(&sql, []).map_err(|e| format!("DuckDB DDL failed: {}", e))?; + Ok(()) + } + DatabaseEngine::Redis => Err("Not supported for Redis.".to_string()), } } diff --git a/src-tauri/src/commands/duckdb.rs b/src-tauri/src/commands/duckdb.rs new file mode 100644 index 0000000..eabfc8b --- /dev/null +++ b/src-tauri/src/commands/duckdb.rs @@ -0,0 +1,505 @@ +use std::collections::{BTreeMap, HashMap, HashSet}; +use std::time::Instant; + +use tauri::{AppHandle, State}; + +use crate::db::{get_or_create_duckdb_connection, resolve_connection_engine, AppState, MAX_QUERY_ROWS}; +use crate::models::{ColumnInfo, ColumnProperties, ForeignKeyEdge, IndexInfo, QueryRequest, QueryResult, SchemaRequest, TableIndexesResult, TableInfo, TablePropertiesApplyRequest}; + +/// Execute a SQL query against a DuckDB connection. +#[tauri::command] +pub async fn duckdb_run_query( + app: AppHandle, + state: State<'_, AppState>, + input: QueryRequest, +) -> Result { + let (connection_id, _engine) = + resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; + + get_or_create_duckdb_connection(&app, &state, &connection_id).await?; + + let sql = input.sql.trim().to_string(); + if sql.is_empty() { + return Err("Enter a SQL statement before running the query.".to_string()); + } + + let started_at = Instant::now(); + let max_rows = input.max_rows.unwrap_or(MAX_QUERY_ROWS); + + let conns = state.duckdb_connections.read().await; + let conn_mutex = conns + .get(&connection_id) + .ok_or_else(|| "DuckDB connection not found.".to_string())?; + let conn = conn_mutex.lock().await; + + let stmt = conn + .prepare(&sql) + .map_err(|e| format!("DuckDB query preparation failed: {}", e))?; + + let columns: Vec = stmt + .column_names() + .iter() + .map(|c| c.to_string()) + .collect(); + + let col_count = columns.len(); + + if col_count > 0 { + // Build a wrapper query that casts all columns to VARCHAR for safe string extraction + let cast_cols: Vec = columns + .iter() + .map(|c| format!("CAST(\"{}\" AS VARCHAR) as \"{}\"", c, c)) + .collect(); + let wrapper_sql = format!("SELECT {} FROM ({})", cast_cols.join(", "), sql); + + drop(stmt); + + let mut stmt2 = conn + .prepare(&wrapper_sql) + .map_err(|e| format!("DuckDB query preparation failed: {}", e))?; + + let mut rows: Vec>> = Vec::new(); + let mut total = 0usize; + + let row_iter = stmt2 + .query_map([], |row| { + let mut map = BTreeMap::new(); + for (i, col) in columns.iter().enumerate() { + let val: Option = row.get(i).ok().flatten(); + map.insert(col.clone(), val); + } + Ok(map) + }) + .map_err(|e| format!("DuckDB query execution failed: {}", e))?; + + for row_result in row_iter { + match row_result { + Ok(row) => { + if rows.len() < max_rows { + rows.push(row); + } + total += 1; + } + Err(e) => return Err(format!("DuckDB row error: {}", e)), + } + } + + Ok(QueryResult { + columns, + rows, + row_count: total.min(max_rows), + execution_ms: started_at.elapsed().as_millis(), + truncated: total > max_rows, + command_tag: None, + }) + } else { + // Non-row-returning statement (INSERT, UPDATE, DDL) + drop(stmt); + let affected = conn + .execute(&sql, []) + .map_err(|e| format!("DuckDB execution failed: {}", e))?; + + Ok(QueryResult { + columns: Vec::new(), + rows: Vec::new(), + row_count: affected, + execution_ms: started_at.elapsed().as_millis(), + truncated: false, + command_tag: Some(affected as u64), + }) + } +} + +/// List all tables in the DuckDB database. +#[tauri::command] +pub async fn duckdb_get_tables( + app: AppHandle, + state: State<'_, AppState>, + connection_id: Option, +) -> Result, String> { + let (connection_id, _engine) = + resolve_connection_engine(&app, &state, connection_id.clone()).await?; + + get_or_create_duckdb_connection(&app, &state, &connection_id).await?; + + let conns = state.duckdb_connections.read().await; + let conn_mutex = conns + .get(&connection_id) + .ok_or_else(|| "DuckDB connection not found.".to_string())?; + let conn = conn_mutex.lock().await; + + let mut stmt = conn + .prepare( + "SELECT table_name FROM information_schema.tables \ + WHERE table_schema = 'main' AND table_type = 'BASE TABLE' \ + ORDER BY table_name", + ) + .map_err(|e| format!("DuckDB table listing failed: {}", e))?; + + let tables: Vec = stmt + .query_map([], |row| { + let name: String = row.get(0)?; + Ok(TableInfo { + schema: "main".to_string(), + name: name.clone(), + preview_query: format!("SELECT * FROM \"{}\" LIMIT 100;", name), + }) + }) + .map_err(|e| format!("DuckDB table listing failed: {}", e))? + .filter_map(|r| r.ok()) + .collect(); + + Ok(tables) +} + +/// Get the column schema for a DuckDB table. +#[tauri::command] +pub async fn duckdb_get_schema( + app: AppHandle, + state: State<'_, AppState>, + connection_id: Option, + table_schema: String, + table_name: String, +) -> Result, String> { + let (connection_id, _engine) = + resolve_connection_engine(&app, &state, connection_id.clone()).await?; + + get_or_create_duckdb_connection(&app, &state, &connection_id).await?; + + let conns = state.duckdb_connections.read().await; + let conn_mutex = conns + .get(&connection_id) + .ok_or_else(|| "DuckDB connection not found.".to_string())?; + let conn = conn_mutex.lock().await; + + let mut stmt = conn + .prepare( + "SELECT column_name, data_type, is_nullable \ + FROM information_schema.columns \ + WHERE table_schema = ? AND table_name = ? \ + ORDER BY ordinal_position", + ) + .map_err(|e| format!("DuckDB schema query failed: {}", e))?; + + let columns: Vec = stmt + .query_map( + duckdb::params![table_schema, table_name], + |row| { + let col_name: String = row.get(0)?; + let data_type: String = row.get(1)?; + let is_nullable: String = row.get(2)?; + Ok(ColumnInfo { + table_schema: table_schema.clone(), + table_name: table_name.clone(), + column_name: col_name, + data_type, + is_nullable: is_nullable == "YES", + }) + }, + ) + .map_err(|e| format!("DuckDB schema query failed: {}", e))? + .filter_map(|r| r.ok()) + .collect(); + + Ok(columns) +} + +/// List foreign key relationships in the DuckDB database. +#[tauri::command] +pub async fn duckdb_get_foreign_keys( + app: AppHandle, + state: State<'_, AppState>, + connection_id: Option, +) -> Result, String> { + let (connection_id, _engine) = + resolve_connection_engine(&app, &state, connection_id.clone()).await?; + + get_or_create_duckdb_connection(&app, &state, &connection_id).await?; + + let conns = state.duckdb_connections.read().await; + let conn_mutex = conns + .get(&connection_id) + .ok_or_else(|| "DuckDB connection not found.".to_string())?; + let conn = conn_mutex.lock().await; + + let mut stmt = conn + .prepare( + "SELECT + kcu.table_schema AS from_schema, + kcu.table_name AS from_table, + kcu.column_name AS from_column, + kcu.referenced_table_schema AS to_schema, + kcu.referenced_table_name AS to_table, + kcu.referenced_column_name AS to_column + FROM information_schema.key_column_usage kcu + WHERE kcu.referenced_table_name IS NOT NULL + ORDER BY kcu.table_schema, kcu.table_name", + ) + .map_err(|e| format!("DuckDB FK query failed: {}", e))?; + + let edges: Vec = stmt + .query_map([], |row| { + Ok(ForeignKeyEdge { + from_schema: row.get(0)?, + from_table: row.get(1)?, + from_column: row.get(2)?, + to_schema: row.get(3)?, + to_table: row.get(4)?, + to_column: row.get(5)?, + }) + }) + .map_err(|e| format!("DuckDB FK query failed: {}", e))? + .filter_map(|r| r.ok()) + .collect(); + + Ok(edges) +} + +/// List indexes on a DuckDB table. +#[tauri::command] +pub async fn duckdb_get_table_indexes( + app: AppHandle, + state: State<'_, AppState>, + input: SchemaRequest, +) -> Result { + let (connection_id, _engine) = + resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; + + get_or_create_duckdb_connection(&app, &state, &connection_id).await?; + + let conns = state.duckdb_connections.read().await; + let conn_mutex = conns + .get(&connection_id) + .ok_or_else(|| "DuckDB connection not found.".to_string())?; + let conn = conn_mutex.lock().await; + + let mut stmt = conn + .prepare( + "SELECT + index_name, + index_schema, + table_name, + is_unique, + is_primary, + sql + FROM duckdb_indexes() + WHERE table_name = ? AND schema_name = ? + ORDER BY index_name", + ) + .map_err(|e| format!("DuckDB index query failed: {}", e))?; + + let indexes: Vec = stmt + .query_map( + duckdb::params![input.table_name, input.table_schema], + |row| { + let idx_name: String = row.get(0)?; + let idx_schema: String = row.get(1)?; + let tbl_name: String = row.get(2)?; + let is_unique: bool = row.get(3)?; + let is_primary: bool = row.get(4)?; + let definition: String = row.get(5)?; + Ok(IndexInfo { + index_schema: idx_schema, + index_name: idx_name, + table_schema: input.table_schema.clone(), + table_name: tbl_name, + is_unique, + is_primary, + is_valid: true, + is_partial: false, + definition, + index_bytes: 0, + idx_scan: 0, + idx_tup_read: 0, + idx_tup_fetch: 0, + }) + }, + ) + .map_err(|e| format!("DuckDB index query failed: {}", e))? + .filter_map(|r| r.ok()) + .collect(); + + let truncated = indexes.len() > 500; + Ok(TableIndexesResult { + indexes: if truncated { indexes.into_iter().take(500).collect() } else { indexes }, + truncated, + }) +} + +/// Get column properties (nullable, primary key, unique, default) for a DuckDB table. +#[tauri::command] +pub async fn duckdb_get_table_properties( + app: AppHandle, + state: State<'_, AppState>, + input: SchemaRequest, +) -> Result, String> { + let (connection_id, _engine) = + resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; + + get_or_create_duckdb_connection(&app, &state, &connection_id).await?; + + let conns = state.duckdb_connections.read().await; + let conn_mutex = conns + .get(&connection_id) + .ok_or_else(|| "DuckDB connection not found.".to_string())?; + let conn = conn_mutex.lock().await; + + // Query columns + let mut col_stmt = conn + .prepare( + "SELECT table_schema, table_name, column_name, data_type, is_nullable, column_default + FROM information_schema.columns + WHERE table_schema = ? AND table_name = ? + ORDER BY ordinal_position", + ) + .map_err(|e| format!("DuckDB properties query failed: {}", e))?; + + let columns: Vec<(String, String, String, String, String, Option)> = col_stmt + .query_map( + duckdb::params![input.table_schema, input.table_name], + |row| Ok((row.get(0)?, row.get(1)?, row.get(2)?, row.get(3)?, row.get(4)?, row.get(5)?)), + ) + .map_err(|e| format!("DuckDB properties query failed: {}", e))? + .filter_map(|r| r.ok()) + .collect(); + + // Query primary keys + let mut pk_stmt = conn + .prepare( + "SELECT kcu.column_name + FROM information_schema.table_constraints tc + JOIN information_schema.key_column_usage kcu + ON tc.constraint_name = kcu.constraint_name + AND tc.table_schema = kcu.table_schema + WHERE tc.table_schema = ? AND tc.table_name = ? + AND tc.constraint_type = 'PRIMARY KEY' + ORDER BY kcu.ordinal_position", + ) + .map_err(|e| format!("DuckDB PK query failed: {}", e))?; + + let pk_cols: std::collections::HashSet = pk_stmt + .query_map(duckdb::params![input.table_schema, input.table_name], |row| row.get(0)) + .map_err(|e| format!("DuckDB PK query failed: {}", e))? + .filter_map(|r| r.ok()) + .collect(); + + // Query unique constraints + let mut uniq_stmt = conn + .prepare( + "SELECT tc.constraint_name, kcu.column_name + FROM information_schema.table_constraints tc + JOIN information_schema.key_column_usage kcu + ON tc.constraint_name = kcu.constraint_name + AND tc.table_schema = kcu.table_schema + WHERE tc.table_schema = ? AND tc.table_name = ? + AND tc.constraint_type = 'UNIQUE' + ORDER BY tc.constraint_name, kcu.ordinal_position", + ) + .map_err(|e| format!("DuckDB unique query failed: {}", e))?; + + let mut unique_by_name: HashMap> = HashMap::new(); + let unique_rows: Vec<(String, String)> = uniq_stmt + .query_map(duckdb::params![input.table_schema, input.table_name], |row| Ok((row.get(0)?, row.get(1)?))) + .map_err(|e| format!("DuckDB unique query failed: {}", e))? + .filter_map(|r| r.ok()) + .collect(); + for (constraint_name, column_name) in unique_rows { + unique_by_name.entry(constraint_name).or_default().push(column_name); + } + + let mut unique_columns: HashSet = HashSet::new(); + let mut composite_unique_columns: HashSet = HashSet::new(); + for (_name, cols) in &unique_by_name { + for c in cols { unique_columns.insert(c.clone()); } + if cols.len() > 1 { for c in cols { composite_unique_columns.insert(c.clone()); } } + } + + Ok(columns.into_iter().map(|(table_schema, table_name, column_name, data_type, is_nullable, column_default)| { + let is_pk = pk_cols.contains(&column_name); + ColumnProperties { + table_schema, + table_name, + column_name: column_name.clone(), + data_type, + is_nullable: is_nullable == "YES", + is_primary_key: is_pk, + is_unique: is_pk || unique_columns.contains(&column_name), + is_part_of_composite_unique: composite_unique_columns.contains(&column_name), + column_default, + is_identity: false, + identity_generation: None, + is_generated: None, + } + }).collect()) +} + +/// Apply nullable/unique changes to a DuckDB table via ALTER TABLE. +#[tauri::command] +pub async fn duckdb_apply_table_properties( + app: AppHandle, + state: State<'_, AppState>, + input: TablePropertiesApplyRequest, +) -> Result<(), String> { + let (connection_id, _engine) = + resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; + + get_or_create_duckdb_connection(&app, &state, &connection_id).await?; + + let conns = state.duckdb_connections.read().await; + let conn_mutex = conns + .get(&connection_id) + .ok_or_else(|| "DuckDB connection not found.".to_string())?; + let conn = conn_mutex.lock().await; + + let table_schema = &input.table_schema; + let table_name = &input.table_name; + + // Get current nullability + let mut col_stmt = conn + .prepare( + "SELECT column_name, is_nullable FROM information_schema.columns + WHERE table_schema = ? AND table_name = ?", + ) + .map_err(|e| format!("DuckDB: {}", e))?; + let current_nullable: HashMap = col_stmt + .query_map(duckdb::params![table_schema, table_name], |row| { + Ok((row.get::<_, String>(0)?, row.get::<_, String>(1)? == "YES")) + }) + .map_err(|e| format!("DuckDB: {}", e))? + .filter_map(|r| r.ok()) + .collect(); + + let qualified_table = format!("\"{}\".\"{}\"", table_schema, table_name); + + for update in &input.columns { + let current = current_nullable.get(&update.column_name) + .ok_or_else(|| format!("Unknown column: {}", update.column_name))?; + if *current == update.is_nullable { continue; } + let qualified_col = format!("\"{}\"", update.column_name); + if update.is_nullable { + conn.execute( + &format!("ALTER TABLE {} ALTER {} DROP NOT NULL", qualified_table, qualified_col), + [], + ).map_err(|e| format!("DuckDB ALTER failed: {}", e))?; + } else { + conn.execute( + &format!("ALTER TABLE {} ALTER {} SET NOT NULL", qualified_table, qualified_col), + [], + ).map_err(|e| format!("DuckDB ALTER failed: {}", e))?; + } + } + + // Unique constraints + for update in &input.columns { + if !update.is_unique { continue; } + let qualified_col = format!("\"{}\"", update.column_name); + let constraint_name = format!("veloxdb_unq_{}", update.column_name); + // Add if not exists — best-effort (DuckDB may error on duplicate, which is fine) + conn.execute( + &format!("ALTER TABLE {} ADD CONSTRAINT \"{}\" UNIQUE ({})", qualified_table, constraint_name, qualified_col), + [], + ).ok(); // Ignore if constraint already exists + } + + Ok(()) +} diff --git a/src-tauri/src/commands/lint.rs b/src-tauri/src/commands/lint.rs index a9bc618..e8dfe42 100644 --- a/src-tauri/src/commands/lint.rs +++ b/src-tauri/src/commands/lint.rs @@ -76,5 +76,27 @@ pub async fn lint_sql( DatabaseEngine::Mongo => { Err("MongoDB does not support SQL linting.".to_string()) } + DatabaseEngine::Duckdb => { + let conns = state.duckdb_connections.read().await; + let conn_mutex = conns + .get(&connection_id) + .ok_or("DuckDB connection not found.".to_string())?; + let conn = conn_mutex.lock().await; + let lint_sql = format!("EXPLAIN {}", sql); + match conn.prepare(&lint_sql) { + Ok(_) => Ok(LintSqlResult { diagnostics: vec![] }), + Err(e) => Ok(LintSqlResult { + diagnostics: vec![SqlDiagnostic { + message: e.to_string(), + severity: "error".to_string(), + line: None, + column: None, + end_line: None, + end_column: None, + }], + }), + } + } + DatabaseEngine::Redis => Err("Not supported for Redis.".to_string()), } } diff --git a/src-tauri/src/commands/mod.rs b/src-tauri/src/commands/mod.rs index d12aa0f..2d48061 100644 --- a/src-tauri/src/commands/mod.rs +++ b/src-tauri/src/commands/mod.rs @@ -422,6 +422,12 @@ pub(crate) async fn run_query_mysql_or_sqlite( DatabaseEngine::Mongo => { return Err("Internal engine routing error (MongoDB uses its own query path).".to_string()); } + DatabaseEngine::Duckdb => { + return Err("Internal engine routing error (DuckDB uses its own query path).".to_string()); + } + DatabaseEngine::Redis => { + return Err("Internal engine routing error (Redis uses its own command path).".to_string()); + } DatabaseEngine::Postgres => { return Err("Internal engine routing error.".to_string()); } @@ -922,8 +928,10 @@ mod ddl; mod editor_meta; mod veloxy; mod lint; -mod mongo; +pub(crate) mod mongo; +pub(crate) mod duckdb; mod export_cmds; +pub(crate) mod redis; // --- Re-exports --- @@ -959,3 +967,9 @@ pub use export_cmds::{ pub use mongo::{ mongo_run_query, mongo_get_collections, mongo_get_schema, }; +pub use duckdb::{ + duckdb_run_query, duckdb_get_tables, duckdb_get_schema, +}; +pub use redis::{ + redis_run_query, redis_get_keys, redis_get_schema, +}; diff --git a/src-tauri/src/commands/mongo.rs b/src-tauri/src/commands/mongo.rs index ea72633..26bc5e3 100644 --- a/src-tauri/src/commands/mongo.rs +++ b/src-tauri/src/commands/mongo.rs @@ -6,7 +6,7 @@ use mongodb::bson::{doc, Document}; use tauri::{AppHandle, State}; use crate::db::{get_or_create_mongo_client, resolve_connection_engine, AppState, MAX_QUERY_ROWS}; -use crate::models::{ColumnInfo, QueryRequest, QueryResult, TableInfo}; +use crate::models::{ColumnInfo, QueryRequest, QueryResult, TableIndexesResult, TableInfo}; /// Parse a user-supplied MongoDB query string into a filter Document. /// @@ -319,3 +319,489 @@ pub async fn mongo_get_schema( }) .collect()) } + +// ── MongoDB Export ────────────────────────────────────────────── + +pub async fn mongo_export_csv( + app: &AppHandle, + state: &AppState, + connection_id: &str, + database: &str, + collection: &str, + output_path: &str, +) -> Result<(), String> { + let client = crate::db::get_or_create_mongo_client(app, state, connection_id).await?; + let db = client.database(database); + let coll = db.collection::(collection); + + let mut cursor = coll.find(doc! {}).limit(5000).await + .map_err(|e| format!("MongoDB export failed: {}", e))?; + + let mut columns: Vec = Vec::new(); + let mut rows: Vec>> = Vec::new(); + let mut cols_set = HashSet::new(); + + while let Some(result) = cursor.next().await { + let doc = result.map_err(|e| format!("MongoDB cursor error: {}", e))?; + if columns.is_empty() { + columns = doc.keys().cloned().collect(); + for c in &columns { cols_set.insert(c.clone()); } + } else { + for k in doc.keys() { + if !cols_set.contains(k) { + columns.push(k.clone()); + cols_set.insert(k.clone()); + } + } + } + let mut row = BTreeMap::new(); + for key in &columns { + row.insert(key.clone(), doc.get(key).and_then(bson_to_display_string)); + } + rows.push(row); + } + + let mut wtr = csv::Writer::from_path(output_path).map_err(|e| e.to_string())?; + wtr.write_record(&columns).map_err(|e| e.to_string())?; + for row in &rows { + let record: Vec = columns.iter().map(|c| row.get(c).unwrap_or(&None).clone().unwrap_or_default()).collect(); + wtr.write_record(&record).map_err(|e| e.to_string())?; + } + wtr.flush().map_err(|e| e.to_string())?; + Ok(()) +} + +pub async fn mongo_export_json( + app: &AppHandle, + state: &AppState, + connection_id: &str, + database: &str, + collection: &str, + output_path: &str, +) -> Result<(), String> { + let client = crate::db::get_or_create_mongo_client(app, state, connection_id).await?; + let db = client.database(database); + let coll = db.collection::(collection); + + let mut cursor = coll.find(doc! {}).limit(5000).await + .map_err(|e| format!("MongoDB export failed: {}", e))?; + + let mut docs: Vec = Vec::new(); + while let Some(result) = cursor.next().await { + let doc = result.map_err(|e| format!("MongoDB cursor error: {}", e))?; + let json_str = serde_json::to_string(&mongodb::bson::to_document(&doc).unwrap_or_default()) + .map_err(|e| e.to_string())?; + docs.push(serde_json::from_str(&json_str).unwrap_or(serde_json::Value::Null)); + } + let content = serde_json::to_string_pretty(&docs).map_err(|e| e.to_string())?; + std::fs::write(output_path, content).map_err(|e| e.to_string())?; + Ok(()) +} + +/// List indexes on a MongoDB collection. +pub async fn mongo_get_table_indexes( + app: &AppHandle, + state: &AppState, + connection_id: &str, + database: &str, + collection: &str, +) -> Result { + use crate::models::IndexInfo; + let client = crate::db::get_or_create_mongo_client(app, state, connection_id).await?; + let db = client.database(database); + let coll = db.collection::(collection); + + let mut cursor = coll.list_indexes().await + .map_err(|e| format!("MongoDB index listing failed: {}", e))?; + + let mut indexes = Vec::new(); + while let Some(result) = cursor.next().await { + let idx = result.map_err(|e| format!("MongoDB index cursor error: {}", e))?; + let name = idx.options.as_ref() + .and_then(|o| o.name.clone()) + .unwrap_or_else(|| { + idx.keys.iter() + .map(|(k, _)| k.clone()) + .collect::>() + .join("_") + }); + let keys_doc = idx.keys.clone(); + let keys: Vec = keys_doc.iter() + .filter_map(|(k, v)| { + let dir = match v.as_i32() { Some(1) => "asc", Some(-1) => "desc", _ => "?" }; + Some(format!("{}({})", k, dir)) + }) + .collect(); + indexes.push(IndexInfo { + index_schema: database.to_string(), + index_name: name.clone(), + table_schema: database.to_string(), + table_name: collection.to_string(), + is_unique: false, + is_primary: name == "_id_", + is_valid: true, + is_partial: false, + definition: format!("index {} ({})", name, keys.join(", ")), + index_bytes: 0, + idx_scan: 0, + idx_tup_read: 0, + idx_tup_fetch: 0, + }); + } + + let truncated = indexes.len() > 500; + Ok(TableIndexesResult { + indexes: if truncated { indexes.into_iter().take(500).collect() } else { indexes }, + truncated, + }) +} + +#[cfg(test)] +mod tests { + use super::*; + use mongodb::bson::{oid::ObjectId, Bson, DateTime as BsonDateTime}; + + #[test] + fn parse_empty_filter() { + let result = parse_mongo_filter("").unwrap(); + assert!(result.is_empty()); + } + + #[test] + fn parse_whitespace_filter() { + let result = parse_mongo_filter(" ").unwrap(); + assert!(result.is_empty()); + } + + #[test] + fn parse_valid_json() { + let result = parse_mongo_filter(r#"{"status": "active"}"#).unwrap(); + assert_eq!(result.get_str("status").unwrap(), "active"); + } + + #[test] + fn parse_json_with_number() { + let result = parse_mongo_filter(r#"{"age": 25}"#).unwrap(); + assert_eq!(result.get_i32("age").unwrap(), 25); + } + + #[test] + fn parse_json_with_boolean() { + let result = parse_mongo_filter(r#"{"verified": true}"#).unwrap(); + assert_eq!(result.get_bool("verified").unwrap(), true); + } + + #[test] + fn parse_json_with_null() { + let result = parse_mongo_filter(r#"{"deletedAt": null}"#).unwrap(); + assert!(result.get("deletedAt").unwrap().as_null().is_some()); + } + + #[test] + fn parse_json_with_operator() { + let result = parse_mongo_filter(r#"{"age": {"$gt": 18}}"#).unwrap(); + let age = result.get_document("age").unwrap(); + assert_eq!(age.get_i32("$gt").unwrap(), 18); + } + + #[test] + fn parse_shell_syntax_find() { + let result = parse_mongo_filter(r#"db.users.find({"status": "active"})"#).unwrap(); + assert_eq!(result.get_str("status").unwrap(), "active"); + } + + #[test] + fn parse_shell_syntax_with_limit() { + let result = parse_mongo_filter(r#"db.users.find({"role": "admin"}).limit(10)"#).unwrap(); + assert_eq!(result.get_str("role").unwrap(), "admin"); + } + + #[test] + fn parse_relaxed_json_single_quotes() { + let result = parse_mongo_filter("{'name': 'John'}").unwrap(); + assert_eq!(result.get_str("name").unwrap(), "John"); + } + + #[test] + fn parse_invalid_json_returns_empty() { + let result = parse_mongo_filter("not json at all").unwrap(); + assert!(result.is_empty()); + } + + #[test] + fn extract_simple_braced_json() { + let result = extract_braced_json(r#"{"key": "value"}"#); + assert_eq!(result.unwrap(), r#"{"key": "value"}"#); + } + + #[test] + fn extract_nested_braced_json() { + let result = extract_braced_json(r#"{"outer": {"inner": true}}"#); + assert!(result.unwrap().contains("inner")); + } + + #[test] + fn extract_braced_json_not_starting_with_brace_returns_none() { + let result = extract_braced_json("not a json"); + assert!(result.is_none()); + } + + #[test] + fn bson_string_display() { + let value = Bson::String("hello".to_string()); + assert_eq!(bson_to_display_string(&value).unwrap(), "hello"); + } + + #[test] + fn bson_int32_display() { + let value = Bson::Int32(42); + assert_eq!(bson_to_display_string(&value).unwrap(), "42"); + } + + #[test] + fn bson_int64_display() { + let value = Bson::Int64(9007199254740991); + assert_eq!(bson_to_display_string(&value).unwrap(), "9007199254740991"); + } + + #[test] + fn bson_double_display() { + let value = Bson::Double(3.14); + assert_eq!(bson_to_display_string(&value).unwrap(), "3.14"); + } + + #[test] + fn bson_bool_display() { + assert_eq!(bson_to_display_string(&Bson::Boolean(true)).unwrap(), "true"); + assert_eq!(bson_to_display_string(&Bson::Boolean(false)).unwrap(), "false"); + } + + #[test] + fn bson_objectid_display() { + let oid = ObjectId::new(); + let value = Bson::ObjectId(oid); + let displayed = bson_to_display_string(&value).unwrap(); + assert_eq!(displayed.len(), 24); + assert!(displayed.chars().all(|c| c.is_ascii_hexdigit())); + } + + #[test] + fn bson_null_display() { + assert!(bson_to_display_string(&Bson::Null).is_none()); + } + + #[test] + fn bson_array_display() { + let value = Bson::Array(vec![Bson::Int32(1), Bson::Int32(2), Bson::Int32(3)]); + let result = bson_to_display_string(&value).unwrap(); + assert!(result.contains("3 items")); + } + + #[test] + fn bson_document_display() { + let value = Bson::Document(doc! { "key": "val" }); + assert_eq!(bson_to_display_string(&value).unwrap(), "{...}"); + } + + #[test] + fn bson_binary_display() { + use mongodb::bson::Binary; + let value = Bson::Binary(Binary { + subtype: mongodb::bson::spec::BinarySubtype::Generic, + bytes: vec![0, 1, 2], + }); + assert_eq!(bson_to_display_string(&value).unwrap(), ""); + } + + #[test] + fn bson_datetime_display() { + use chrono::TimeZone; + let dt = chrono::Utc.with_ymd_and_hms(2024, 3, 15, 10, 30, 45).unwrap(); + let bson_dt = BsonDateTime::from_millis(dt.timestamp_millis()); + let value = Bson::DateTime(bson_dt); + let result = bson_to_display_string(&value).unwrap(); + assert!(result.contains("2024-03-15")); + } + + #[test] + fn infer_bson_type_string() { + assert_eq!(infer_bson_type(&Bson::String("hello".to_string())), "string"); + } + + #[test] + fn infer_bson_type_integer() { + assert_eq!(infer_bson_type(&Bson::Int32(1)), "integer"); + assert_eq!(infer_bson_type(&Bson::Int64(1)), "integer"); + } + + #[test] + fn infer_bson_type_double() { + assert_eq!(infer_bson_type(&Bson::Double(1.0)), "double"); + } + + #[test] + fn infer_bson_type_boolean() { + assert_eq!(infer_bson_type(&Bson::Boolean(true)), "boolean"); + } + + #[test] + fn infer_bson_type_objectid() { + assert_eq!(infer_bson_type(&Bson::ObjectId(ObjectId::new())), "objectId"); + } + + #[test] + fn infer_bson_type_date() { + let dt = BsonDateTime::from_millis(0); + assert_eq!(infer_bson_type(&Bson::DateTime(dt)), "date"); + } + + #[test] + fn infer_bson_type_array() { + assert_eq!(infer_bson_type(&Bson::Array(vec![])), "array"); + } + + #[test] + fn infer_bson_type_document() { + assert_eq!(infer_bson_type(&Bson::Document(doc! {})), "object"); + } + + #[test] + fn infer_bson_type_null() { + assert_eq!(infer_bson_type(&Bson::Null), "null"); + } + + #[test] + fn infer_bson_type_regex() { + use mongodb::bson::Regex; + assert_eq!(infer_bson_type(&Bson::RegularExpression(Regex { + pattern: ".*".to_string(), + options: "".to_string(), + })), "regex"); + } + + // ── Cross-engine QueryResult compatibility ────────────────── + + #[test] + fn query_result_structure_is_consistent_across_engines() { + use crate::models::QueryResult; + use std::collections::BTreeMap; + + // Simulate a SQL result (Postgres/MySQL/SQLite format) + let sql_result = QueryResult { + columns: vec!["id".to_string(), "name".to_string()], + rows: vec![ + BTreeMap::from([ + ("id".to_string(), Some("1".to_string())), + ("name".to_string(), Some("Alice".to_string())), + ]), + ], + row_count: 1, + execution_ms: 15, + truncated: false, + command_tag: None, + }; + + // Simulate a MongoDB result (should match exactly the same shape) + let mongo_result = QueryResult { + columns: vec!["_id".to_string(), "name".to_string()], + rows: vec![ + BTreeMap::from([ + ("_id".to_string(), Some("507f1f77bcf86cd799439011".to_string())), + ("name".to_string(), Some("Alice".to_string())), + ]), + ], + row_count: 1, + execution_ms: 18, + truncated: false, + command_tag: None, + }; + + // Both have same shape (structurally identical) + assert_eq!(sql_result.columns.len(), 2); + assert_eq!(mongo_result.columns.len(), 2); + assert_eq!(sql_result.rows.len(), 1); + assert_eq!(mongo_result.rows.len(), 1); + assert!(!sql_result.truncated); + assert!(!mongo_result.truncated); + + // Both are serializable (Tauri IPC sends JSON) + let sql_json = serde_json::to_string(&sql_result).unwrap(); + let mongo_json = serde_json::to_string(&mongo_result).unwrap(); + assert!(sql_json.contains("\"columns\"")); + assert!(mongo_json.contains("\"columns\"")); + assert!(sql_json.contains("\"rows\"")); + assert!(mongo_json.contains("\"rows\"")); + } + + #[test] + fn empty_mongo_result_matches_sql_format() { + use crate::models::QueryResult; + + let empty = QueryResult { + columns: vec![], + rows: vec![], + row_count: 0, + execution_ms: 0, + truncated: false, + command_tag: None, + }; + + // An empty query result should be structurally identical + // regardless of engine — this is what the frontend expects. + assert_eq!(empty.columns.len(), 0); + assert_eq!(empty.rows.len(), 0); + assert_eq!(empty.row_count, 0); + assert!(!empty.truncated); + } + + #[test] + fn mongo_null_values_match_sql_null() { + use std::collections::BTreeMap; + + // SQL NULL and MongoDB null/absent field both map to Option::None + // Both engines produce the same format for the frontend grid. + + let sql_row = BTreeMap::from([ + ("name".to_string(), Some("Alice".to_string())), + ("email".to_string(), None), // SQL NULL + ]); + + let mongo_row = BTreeMap::from([ + ("name".to_string(), Some("Bob".to_string())), + ("email".to_string(), None), // MongoDB field missing or null + ]); + + // Both use Option for values — the frontend treats None as "NULL" + assert_eq!(sql_row.get("email").unwrap(), &None); + assert_eq!(mongo_row.get("email").unwrap(), &None); + assert_eq!(sql_row.get("name").unwrap(), &Some("Alice".to_string())); + assert_eq!(mongo_row.get("name").unwrap(), &Some("Bob".to_string())); + } + + #[test] + fn mongo_dynamic_columns_produces_tabular_rows() { + use std::collections::BTreeMap; + + // When MongoDB documents have different fields, we union the keys + // and fill missing columns with None (same as SQL NULL). + + let columns = vec!["name".to_string(), "extra".to_string()]; + let rows = vec![ + BTreeMap::from([ + ("name".to_string(), Some("doc1".to_string())), + ("extra".to_string(), Some("value1".to_string())), + ]), + BTreeMap::from([ + ("name".to_string(), Some("doc2".to_string())), + ("extra".to_string(), None), // this doc lacks the field + ]), + ]; + + // Every row should have an entry for every column (even if None) + for row in &rows { + for col in &columns { + assert!(row.contains_key(col), "row missing column: {}", col); + } + } + } +} diff --git a/src-tauri/src/commands/query.rs b/src-tauri/src/commands/query.rs index 6c97fd2..017d7b2 100644 --- a/src-tauri/src/commands/query.rs +++ b/src-tauri/src/commands/query.rs @@ -13,6 +13,7 @@ use crate::pg_error::map_pg_err; use super::{is_read_only_sql, run_query_mysql_or_sqlite, mysql_get_string, sqlite_get_idx, sqlite_get_name}; use super::mongo::{mongo_run_query, mongo_get_collections, mongo_get_schema}; +use super::duckdb::{duckdb_run_query, duckdb_get_tables, duckdb_get_schema}; #[tauri::command] pub async fn run_query( @@ -97,6 +98,15 @@ pub async fn run_query( allow_write: input.allow_write, }).await } + DatabaseEngine::Duckdb => { + duckdb_run_query(app, state, QueryRequest { + connection_id: Some(connection_id), + sql: sql.clone(), + max_rows: input.max_rows, + allow_write: input.allow_write, + }).await + } + DatabaseEngine::Redis => Err("Redis uses its own command path.".to_string()), } } @@ -175,6 +185,10 @@ pub async fn get_tables( DatabaseEngine::Mongo => { mongo_get_collections(app, state, Some(connection_id)).await } + DatabaseEngine::Duckdb => { + duckdb_get_tables(app, state, Some(connection_id)).await + } + DatabaseEngine::Redis => Err("Redis uses its own command path.".to_string()), } } @@ -251,5 +265,11 @@ pub async fn get_schema( DatabaseEngine::Mongo => { mongo_get_schema(app, state, input.connection_id, input.table_schema, input.table_name).await } + DatabaseEngine::Duckdb => { + duckdb_get_schema(app, state, input.connection_id, input.table_schema, input.table_name).await + } + DatabaseEngine::Redis => { + crate::commands::redis::redis_get_schema(app, state, input.connection_id, input.table_schema, input.table_name).await + } } } diff --git a/src-tauri/src/commands/redis.rs b/src-tauri/src/commands/redis.rs new file mode 100644 index 0000000..5bd03af --- /dev/null +++ b/src-tauri/src/commands/redis.rs @@ -0,0 +1,132 @@ +use std::collections::BTreeMap; +use std::time::Instant; +use tauri::{AppHandle, State}; +use crate::db::{get_or_create_redis_client, resolve_connection_engine, AppState}; +use crate::models::{QueryRequest, QueryResult, TableInfo, ColumnInfo}; + +#[tauri::command] +pub async fn redis_run_query( + app: AppHandle, + state: State<'_, AppState>, + input: QueryRequest, +) -> Result { + let (connection_id, _) = resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; + let mut client = get_or_create_redis_client(&app, &state, &connection_id).await?; + let raw = input.sql.trim().to_string(); + if raw.is_empty() { return Err("Enter a Redis command.".to_string()); } + let started_at = Instant::now(); + let parts: Vec<&str> = raw.split_whitespace().collect(); + let cmd = parts.first().map(|s| s.to_uppercase()).unwrap_or_default(); + + match cmd.as_str() { + "PING" => { + let result: String = redis::cmd("PING").query_async(&mut client).await.map_err(|e| e.to_string())?; + Ok(QueryResult { columns: vec!["result".to_string()], rows: vec![BTreeMap::from([("result".to_string(), Some(result))])], row_count: 1, execution_ms: started_at.elapsed().as_millis(), truncated: false, command_tag: None }) + } + "GET" => { + let key = parts.get(1).unwrap_or(&""); + let result: Option = redis::cmd("GET").arg(key).query_async(&mut client).await.map_err(|e| e.to_string())?; + Ok(QueryResult { columns: vec!["value".to_string()], rows: vec![BTreeMap::from([("value".to_string(), result)])], row_count: 1, execution_ms: started_at.elapsed().as_millis(), truncated: false, command_tag: None }) + } + "SET" => { + let key = parts.get(1).unwrap_or(&""); + let value = parts.get(2).unwrap_or(&""); + redis::cmd("SET").arg(key).arg(value).query_async::<_, ()>(&mut client).await.map_err(|e| e.to_string())?; + Ok(QueryResult { columns: vec![], rows: vec![], row_count: 0, execution_ms: started_at.elapsed().as_millis(), truncated: false, command_tag: Some(0) }) + } + "KEYS" | "SCAN" => { + let pattern = parts.get(1).unwrap_or(&"*"); + let keys: Vec = redis::cmd("KEYS").arg(pattern).query_async(&mut client).await.map_err(|e| e.to_string())?; + let rows: Vec>> = keys.into_iter().map(|k| BTreeMap::from([("key".to_string(), Some(k))])).collect(); + Ok(QueryResult { columns: vec!["key".to_string()], rows: rows.clone(), row_count: rows.len(), execution_ms: started_at.elapsed().as_millis(), truncated: false, command_tag: None }) + } + "HGETALL" => { + let key = parts.get(1).unwrap_or(&""); + let hash: Vec = redis::cmd("HGETALL").arg(key).query_async(&mut client).await.map_err(|e| e.to_string())?; + let mut row = BTreeMap::new(); + let mut cols = Vec::new(); + for chunk in hash.chunks(2) { + if chunk.len() == 2 { + cols.push(chunk[0].clone()); + row.insert(chunk[0].clone(), Some(chunk[1].clone())); + } + } + Ok(QueryResult { columns: cols, rows: vec![row], row_count: 1, execution_ms: started_at.elapsed().as_millis(), truncated: false, command_tag: None }) + } + "LRANGE" => { + let key = parts.get(1).unwrap_or(&""); + let start: isize = parts.get(2).and_then(|s| s.parse().ok()).unwrap_or(0); + let stop: isize = parts.get(3).and_then(|s| s.parse().ok()).unwrap_or(-1); + let list: Vec = redis::cmd("LRANGE").arg(key).arg(start).arg(stop).query_async(&mut client).await.map_err(|e| e.to_string())?; + let rows: Vec>> = list.into_iter().enumerate().map(|(i, v)| { + let mut map = BTreeMap::new(); + map.insert("index".to_string(), Some(i.to_string())); + map.insert("value".to_string(), Some(v)); + map + }).collect(); + Ok(QueryResult { columns: vec!["index".to_string(), "value".to_string()], rows: rows.clone(), row_count: rows.len(), execution_ms: started_at.elapsed().as_millis(), truncated: false, command_tag: None }) + } + "SMEMBERS" => { + let key = parts.get(1).unwrap_or(&""); + let members: Vec = redis::cmd("SMEMBERS").arg(key).query_async(&mut client).await.map_err(|e| e.to_string())?; + let rows: Vec>> = members.into_iter().map(|v| BTreeMap::from([("member".to_string(), Some(v))])).collect(); + Ok(QueryResult { columns: vec!["member".to_string()], rows: rows.clone(), row_count: rows.len(), execution_ms: started_at.elapsed().as_millis(), truncated: false, command_tag: None }) + } + _ => Err(format!("Unsupported Redis command: {}. Try GET, SET, KEYS, HGETALL, LRANGE, SMEMBERS, PING.", cmd)) + } +} + +#[tauri::command] +pub async fn redis_get_keys( + app: AppHandle, + state: State<'_, AppState>, + connection_id: Option, +) -> Result, String> { + let (connection_id, _) = resolve_connection_engine(&app, &state, connection_id).await?; + let mut client = get_or_create_redis_client(&app, &state, &connection_id).await?; + let keys: Vec = redis::cmd("KEYS").arg("*").query_async(&mut client).await.map_err(|e| e.to_string())?; + Ok(keys.into_iter().map(|k| TableInfo { schema: "0".to_string(), name: k.clone(), preview_query: format!("GET {}", k) }).collect()) +} + +/// Infer schema for a Redis key by checking its type and sampling data. +#[tauri::command] +pub async fn redis_get_schema( + app: AppHandle, + state: State<'_, AppState>, + connection_id: Option, + _table_schema: String, + table_name: String, +) -> Result, String> { + let (connection_id, _) = resolve_connection_engine(&app, &state, connection_id).await?; + let mut client = get_or_create_redis_client(&app, &state, &connection_id).await?; + + let key_type: String = redis::cmd("TYPE").arg(&table_name).query_async(&mut client) + .await.map_err(|e| format!("Redis TYPE failed: {}", e))?; + + match key_type.as_str() { + "string" => Ok(vec![ColumnInfo { + table_schema: "0".to_string(), table_name: table_name.clone(), + column_name: "value".to_string(), data_type: "string".to_string(), is_nullable: true, + }]), + "hash" => { + let fields: Vec = redis::cmd("HKEYS").arg(&table_name).query_async(&mut client) + .await.map_err(|e| format!("Redis HKEYS failed: {}", e))?; + Ok(fields.into_iter().map(|f| ColumnInfo { + table_schema: "0".to_string(), table_name: table_name.clone(), + column_name: f, data_type: "string".to_string(), is_nullable: true, + }).collect()) + } + "list" => Ok(vec![ + ColumnInfo { table_schema: "0".to_string(), table_name: table_name.clone(), column_name: "index".to_string(), data_type: "integer".to_string(), is_nullable: false }, + ColumnInfo { table_schema: "0".to_string(), table_name: table_name.clone(), column_name: "value".to_string(), data_type: "string".to_string(), is_nullable: true }, + ]), + "set" => Ok(vec![ColumnInfo { + table_schema: "0".to_string(), table_name: table_name.clone(), + column_name: "member".to_string(), data_type: "string".to_string(), is_nullable: false, + }]), + _ => Ok(vec![ColumnInfo { + table_schema: "0".to_string(), table_name: table_name.clone(), + column_name: "value".to_string(), data_type: key_type, is_nullable: true, + }]), + } +} diff --git a/src-tauri/src/commands/table_props.rs b/src-tauri/src/commands/table_props.rs index e3cda25..5e9a750 100644 --- a/src-tauri/src/commands/table_props.rs +++ b/src-tauri/src/commands/table_props.rs @@ -139,6 +139,14 @@ pub async fn get_table_properties( return Ok(properties); } + if engine == DatabaseEngine::Duckdb { + return crate::commands::duckdb::duckdb_get_table_properties(app, state, ctx).await; + } + + if engine == DatabaseEngine::Mongo || engine == DatabaseEngine::Redis { + return Err("Table properties are not supported for this engine type.".to_string()); + } + with_pool_client_retry(&app, &state, &connection_id, ctx, |client, input| async move { let columns = client.query( "select c.table_schema, c.table_name, c.column_name, c.data_type, \ @@ -218,7 +226,7 @@ pub async fn apply_table_properties( input: TablePropertiesApplyRequest, ) -> Result<(), String> { let (connection_id, engine) = resolve_connection_engine(&app, &state, input.connection_id.clone()).await?; - if engine != DatabaseEngine::Postgres { + if engine != DatabaseEngine::Postgres && engine != DatabaseEngine::Duckdb { return Err(format!( "Table property editing is not supported for {} connections yet.", match engine { @@ -226,10 +234,16 @@ pub async fn apply_table_properties( DatabaseEngine::Mysql => "MySQL", DatabaseEngine::Sqlite => "SQLite", DatabaseEngine::Mongo => "MongoDB", + DatabaseEngine::Duckdb => "DuckDB", + DatabaseEngine::Redis => "Redis", } )); } + if engine == DatabaseEngine::Duckdb { + return crate::commands::duckdb::duckdb_apply_table_properties(app, state, input).await; + } + with_pool_client_retry(&app, &state, &connection_id, input, |mut client, input| async move { let table_schema = input.table_schema; let table_name = input.table_name; @@ -434,6 +448,14 @@ pub async fn get_foreign_keys( return Ok(edges); } + if engine == DatabaseEngine::Duckdb { + return crate::commands::duckdb::duckdb_get_foreign_keys(app, state, Some(connection_id)).await; + } + + if engine == DatabaseEngine::Mongo || engine == DatabaseEngine::Redis { + return Ok(Vec::new()); + } + with_pool_client_retry(&app, &state, &connection_id, (), |client, ()| async move { let rows = client.query( "select src_ns.nspname::text as from_schema, src_cls.relname::text as from_table, \ @@ -539,6 +561,18 @@ pub async fn get_table_indexes( return Ok(TableIndexesResult { indexes, truncated }); } + if engine == DatabaseEngine::Duckdb { + return crate::commands::duckdb::duckdb_get_table_indexes(app, state, ctx).await; + } + + if engine == DatabaseEngine::Mongo { + return crate::commands::mongo::mongo_get_table_indexes(&app, &state, &connection_id, &ctx.table_schema, &ctx.table_name).await; + } + + if engine == DatabaseEngine::Redis { + return Ok(TableIndexesResult { indexes: Vec::new(), truncated: false }); + } + with_pool_client_retry(&app, &state, &connection_id, ctx, |client, input| async move { let table_schema = input.table_schema; let table_name = input.table_name; diff --git a/src-tauri/src/db.rs b/src-tauri/src/db.rs index 9804e6d..e295e38 100644 --- a/src-tauri/src/db.rs +++ b/src-tauri/src/db.rs @@ -15,6 +15,7 @@ use tokio_postgres::NoTls; use rustls::pki_types::{CertificateDer, PrivateKeyDer}; use mongodb::{Client as MongoClient, bson::doc}; +use duckdb::Connection as DuckConnection; use crate::models::{ AskVeloxyConversationMessage, AskVeloxyDbContextCache, ConnectionInput, ConnectionSslMode, ConnectionSummary, DatabaseEngine, StoredConnection, @@ -34,6 +35,7 @@ const POOL_CREATE_SECS: u64 = 15; const POOL_RECYCLE_SECS: u64 = 15; pub const DEFAULT_MYSQL_PORT: u16 = 3306; pub const DEFAULT_MONGO_PORT: u16 = 27017; +pub const DEFAULT_REDIS_PORT: u16 = 6379; fn deadpool_ssl_mode(mode: ConnectionSslMode) -> DeadpoolSslMode { match mode { @@ -49,6 +51,8 @@ pub struct AppState { pub mysql_pools: RwLock>, pub sqlite_pools: RwLock>, pub mongo_clients: RwLock>, + pub duckdb_connections: RwLock>>, + pub redis_clients: RwLock>, pub active_connection_id: RwLock>, pub ssh_tunnels: RwLock>, pub ask_veloxy_db_context_cache: RwLock>, @@ -357,6 +361,88 @@ pub async fn get_or_create_mongo_client( Ok(client) } +// ── DuckDB ────────────────────────────────────────────────────── + +pub fn build_duckdb_connection(input: &ConnectionInput) -> Result { + let path = input.file_path.as_deref().unwrap_or(":memory:"); + let conn = if path == ":memory:" || path.is_empty() { + duckdb::Connection::open_in_memory() + .map_err(|e| format!("DuckDB in-memory connection failed: {}", e))? + } else { + duckdb::Connection::open(path) + .map_err(|e| format!("DuckDB connection to '{}' failed: {}", path, e))? + }; + // Verify the connection works + conn.execute_batch("SELECT 1") + .map_err(|e| format!("DuckDB verification query failed: {}", e))?; + Ok(conn) +} + +pub async fn get_or_create_duckdb_connection( + app: &AppHandle, + state: &AppState, + connection_id: &str, +) -> Result<(), String> { + { + let conns = state.duckdb_connections.read().await; + if conns.contains_key(connection_id) { + return Ok(()); + } + } + let stored = load_connection(app, connection_id)? + .ok_or_else(|| "Stored connection details were not found.".to_string())?; + let conn = build_duckdb_connection(&stored.to_input())?; + state.duckdb_connections.write().await.insert(connection_id.to_string(), tokio::sync::Mutex::new(conn)); + Ok(()) +} + +// ── Redis ─────────────────────────────────────────────────────── + +pub fn build_redis_url(input: &ConnectionInput) -> String { + let host = if input.host.is_empty() { "127.0.0.1" } else { &input.host }; + let port = if input.port == 0 { DEFAULT_REDIS_PORT } else { input.port }; + if input.user.is_empty() { + format!("redis://{}:{}", host, port) + } else { + format!( + "redis://{}:{}@{}:{}", + urlencoding::encode(&input.user), + urlencoding::encode(&input.password), + host, + port, + ) + } +} + +pub async fn get_or_create_redis_client( + app: &AppHandle, + state: &AppState, + connection_id: &str, +) -> Result { + { + let clients = state.redis_clients.read().await; + if let Some(client) = clients.get(connection_id) { + return Ok(client.clone()); + } + } + let stored = load_connection(app, connection_id)? + .ok_or_else(|| "Stored connection details were not found.".to_string())?; + let url = build_redis_url(&stored.to_input()); + let client = redis::Client::open(url) + .map_err(|e| format!("Redis connection failed: {}", e))?; + let conn = redis::aio::ConnectionManager::new(client) + .await + .map_err(|e| format!("Redis connection failed: {}", e))?; + // Verify with PING + let mut verify = conn.clone(); + redis::cmd("PING") + .query_async::<_, String>(&mut verify) + .await + .map_err(|e| format!("Redis ping failed: {}", e))?; + state.redis_clients.write().await.insert(connection_id.to_string(), conn.clone()); + Ok(conn) +} + /// Heuristic for transport-level failures where discarding the pool and opening /// a new TCP session may succeed (sleep/VPN blips, idle disconnects). fn is_retryable_connection_error(message: &str) -> bool { @@ -383,6 +469,8 @@ pub async fn drop_pool(state: &AppState, connection_id: &str) { state.mysql_pools.write().await.remove(connection_id); state.sqlite_pools.write().await.remove(connection_id); state.mongo_clients.write().await.remove(connection_id); + state.duckdb_connections.write().await.remove(connection_id); + state.redis_clients.write().await.remove(connection_id); state .ask_veloxy_db_context_cache .write() @@ -622,6 +710,18 @@ pub async fn refresh_connection_pools( client.database("admin").run_command(doc! { "ping": 1 }).await .map_err(|e| format!("MongoDB ping failed: {}", e))?; } + DatabaseEngine::Duckdb => { + get_or_create_duckdb_connection(app, state, connection_id).await?; + let conns = state.duckdb_connections.read().await; + let conn = conns.get(connection_id).ok_or("DuckDB connection not found")?; + let conn = conn.lock().await; + conn.execute_batch("SELECT 1").map_err(|e| format!("DuckDB ping failed: {}", e))?; + } + DatabaseEngine::Redis => { + let mut client = get_or_create_redis_client(app, state, connection_id).await?; + redis::cmd("PING").query_async::<_, String>(&mut client).await + .map_err(|e| format!("Redis ping failed: {}", e))?; + } } Ok(()) @@ -804,7 +904,7 @@ pub fn require_safe_identifier<'a>(name: &'a str, context: &str) -> Result<&'a s #[cfg(test)] mod tests { - use super::{is_safe_identifier, mysql_url, require_safe_identifier}; + use super::{build_mongo_connection_string, is_safe_identifier, mysql_url, require_safe_identifier}; use crate::models::{ConnectionInput, ConnectionSslMode, DatabaseEngine}; fn mysql_input(ssl_mode: ConnectionSslMode) -> ConnectionInput { @@ -824,6 +924,74 @@ mod tests { } } + fn postgres_input() -> ConnectionInput { + ConnectionInput { + id: None, + name: "test".to_string(), + engine: DatabaseEngine::Postgres, + host: "localhost".to_string(), + port: 5432, + database: "app".to_string(), + file_path: None, + user: "postgres".to_string(), + password: "secret".to_string(), + ssl_mode: ConnectionSslMode::Prefer, + ssh_config: None, + extra_params: None, + } + } + + fn sqlite_input() -> ConnectionInput { + ConnectionInput { + id: None, + name: "test".to_string(), + engine: DatabaseEngine::Sqlite, + host: String::new(), + port: 0, + database: String::new(), + file_path: Some("/tmp/test.db".to_string()), + user: String::new(), + password: String::new(), + ssl_mode: ConnectionSslMode::Disable, + ssh_config: None, + extra_params: None, + } + } + + fn mongo_input() -> ConnectionInput { + ConnectionInput { + id: None, + name: "test".to_string(), + engine: DatabaseEngine::Mongo, + host: "localhost".to_string(), + port: 27017, + database: "admin".to_string(), + file_path: None, + user: String::new(), + password: String::new(), + ssl_mode: ConnectionSslMode::Disable, + ssh_config: None, + extra_params: None, + } + } + + fn mariadb_input() -> ConnectionInput { + ConnectionInput { + id: None, + name: "test".to_string(), + engine: DatabaseEngine::Mysql, + host: "localhost".to_string(), + port: 3306, + database: "app".to_string(), + file_path: None, + user: "root".to_string(), + password: "pw".to_string(), + ssl_mode: ConnectionSslMode::Prefer, + ssh_config: None, + extra_params: None, + } + } + #[test] fn mysql_url_maps_ssl_mode() { assert!(mysql_url("localhost", 3306, &mysql_input(ConnectionSslMode::Disable)) @@ -834,6 +1002,77 @@ mod tests { .ends_with("?ssl-mode=REQUIRED")); } + #[test] + fn mariadb_uses_mysql_url_format() { + // MariaDB is wire-compatible with MySQL — uses the same URL format + let url = mysql_url("localhost", 3306, &mariadb_input()); + assert!(url.starts_with("mysql://")); + assert!(url.contains("root:pw@localhost:3306/app")); + } + + #[test] + fn mongo_connection_string_no_auth() { + let uri = build_mongo_connection_string(&mongo_input()); + assert_eq!(uri, "mongodb://localhost:27017/admin"); + } + + #[test] + fn mongo_connection_string_with_auth() { + let input = ConnectionInput { + user: "admin".to_string(), + password: "secret".to_string(), + ..mongo_input() + }; + let uri = build_mongo_connection_string(&input); + assert!(uri.contains("admin:secret@")); + assert!(uri.contains("localhost:27017/")); + } + + #[test] + fn mongo_connection_string_with_custom_port() { + let input = ConnectionInput { + port: 27018, + ..mongo_input() + }; + let uri = build_mongo_connection_string(&input); + assert_eq!(uri, "mongodb://localhost:27018/admin"); + } + + #[test] + fn mongo_connection_string_with_extra_params() { + use std::collections::HashMap; + let mut params = HashMap::new(); + params.insert("replicaSet".to_string(), "rs0".to_string()); + params.insert("authSource".to_string(), "admin".to_string()); + let input = ConnectionInput { + extra_params: Some(params), + ..mongo_input() + }; + let uri = build_mongo_connection_string(&input); + assert!(uri.contains("replicaSet=rs0")); + assert!(uri.contains("authSource=admin")); + } + + #[test] + fn mongo_connection_defaults_to_admin_database() { + let input = ConnectionInput { + database: String::new(), + ..mongo_input() + }; + let uri = build_mongo_connection_string(&input); + assert!(uri.ends_with("/admin")); + } + + #[test] + fn mongo_connection_defaults_to_localhost() { + let input = ConnectionInput { + host: String::new(), + ..mongo_input() + }; + let uri = build_mongo_connection_string(&input); + assert!(uri.contains("localhost")); + } + #[test] fn rejects_sql_injection_in_identifiers() { assert!(!is_safe_identifier("\"; DROP TABLE users; --")); diff --git a/src-tauri/src/export.rs b/src-tauri/src/export.rs index 6732b9a..1a4a4e0 100644 --- a/src-tauri/src/export.rs +++ b/src-tauri/src/export.rs @@ -375,8 +375,31 @@ pub async fn export_results_csv( } } DatabaseEngine::Mongo => { - return Err("MongoDB export is not supported.".to_string()); + let col = sql.split_whitespace().next().unwrap_or("main"); + return crate::commands::mongo::mongo_export_csv(app, state, &connection_id, &sql, col, &input.output_path).await; } + DatabaseEngine::Duckdb => { + let conns = state.duckdb_connections.read().await; + let conn_mutex = conns.get(&connection_id).ok_or("DuckDB connection not found")?; + let conn = conn_mutex.lock().await; + let mut stmt = conn.prepare(&sql).map_err(|e| e.to_string())?; + let cols: Vec = stmt.column_names().iter().map(|c| c.to_string()).collect(); + let mut wtr = csv::Writer::from_path(&input.output_path).map_err(|e| e.to_string())?; + wtr.write_record(&cols).map_err(|e| e.to_string())?; + let rows = stmt.query_map([], |row| { + let vals: Vec = cols.iter().enumerate().map(|(i, _)| { + row.get::<_, Option>(i).ok().flatten().unwrap_or_default() + }).collect(); + Ok(vals) + }).map_err(|e| e.to_string())?; + for row in rows { + let vals = row.map_err(|e| e.to_string())?; + wtr.write_record(&vals).map_err(|e| e.to_string())?; + } + wtr.flush().map_err(|e| e.to_string())?; + return Ok(()); + } + DatabaseEngine::Redis => return Err("Not supported for Redis.".to_string()), }; let content = lines.join("\n") + "\n"; @@ -484,8 +507,32 @@ pub async fn export_results_json( result } DatabaseEngine::Mongo => { - return Err("MongoDB JSON export is not supported.".to_string()); + let col = sql.split_whitespace().next().unwrap_or("main"); + return crate::commands::mongo::mongo_export_json(app, state, &connection_id, &sql, col, &input.output_path).await; + } + DatabaseEngine::Duckdb => { + let conns = state.duckdb_connections.read().await; + let conn_mutex = conns.get(&connection_id).ok_or("DuckDB connection not found")?; + let conn = conn_mutex.lock().await; + let mut stmt = conn.prepare(&sql).map_err(|e| e.to_string())?; + let cols: Vec = stmt.column_names().iter().map(|c| c.to_string()).collect(); + let mut rows = Vec::new(); + let results = stmt.query_map([], |row| { + let mut map = serde_json::Map::new(); + for (i, col) in cols.iter().enumerate() { + let val: Option = row.get(i).ok().flatten(); + map.insert(col.clone(), serde_json::Value::String(val.unwrap_or_default())); + } + Ok(serde_json::Value::Object(map)) + }).map_err(|e| e.to_string())?; + for r in results { + rows.push(r.map_err(|e| e.to_string())?); + } + let content = serde_json::to_string_pretty(&rows).map_err(|e| e.to_string())?; + fs::write(&input.output_path, content).map_err(|e| e.to_string())?; + return Ok(()); } + DatabaseEngine::Redis => return Err("Not supported for Redis.".to_string()), }; let content = format!("[\n{}\n]\n", rows.join(",\n")); diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 883c10d..93a411b 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -9,11 +9,12 @@ mod ssh_tunnel; use commands::{ apply_table_properties, cancel_veloxy_request, chat_with_db, clear_veloxy_conversation, connect_db, delete_connection, - delete_openrouter_api_key, disconnect_db, execute_ddl_statement, execute_ddl_transaction, export_diagram_png, + delete_openrouter_api_key, disconnect_db, duckdb_get_schema, duckdb_get_tables, duckdb_run_query, + execute_ddl_statement, execute_ddl_transaction, export_diagram_png, export_results_csv_command, export_results_json_command, generate_sql_from_nl, get_foreign_keys, get_openrouter_api_key, get_query_editor_metadata, get_schema, get_table_indexes, get_table_properties, get_tables, - lint_sql, list_connections_command, list_databases, load_veloxy_conversation, mongo_run_query, mongo_get_collections, - mongo_get_schema, ping_connection, + lint_sql, list_connections_command, list_databases, load_veloxy_conversation, mongo_get_collections, + mongo_get_schema, mongo_run_query, ping_connection, redis_get_keys, redis_get_schema, redis_run_query, refresh_connection, rename_connection, run_query, save_base64_png, save_text_file, set_active_connection, store_openrouter_api_key, switch_database, }; @@ -85,7 +86,13 @@ pub fn run() { delete_openrouter_api_key, mongo_run_query, mongo_get_collections, - mongo_get_schema + mongo_get_schema, + duckdb_run_query, + duckdb_get_tables, + duckdb_get_schema, + redis_run_query, + redis_get_keys, + redis_get_schema ]) .run(tauri::generate_context!()) .expect("error while running tauri application"); diff --git a/src-tauri/src/models.rs b/src-tauri/src/models.rs index a7a34f8..e237383 100644 --- a/src-tauri/src/models.rs +++ b/src-tauri/src/models.rs @@ -27,6 +27,8 @@ pub enum DatabaseEngine { Mysql, Sqlite, Mongo, + Duckdb, + Redis, } fn default_database_engine() -> DatabaseEngine { diff --git a/src/data/types.ts b/src/data/types.ts index 6aaea55..c8a668c 100644 --- a/src/data/types.ts +++ b/src/data/types.ts @@ -1,6 +1,6 @@ /** PostgreSQL `sslmode`-style TLS (lowercase in JSON for Tauri). */ export type ConnectionSslMode = 'disable' | 'prefer' | 'require' -export type DatabaseEngine = 'postgres' | 'mysql' | 'sqlite' | 'mongo' +export type DatabaseEngine = 'postgres' | 'mysql' | 'sqlite' | 'mongo' | 'duckdb' | 'redis' export type SshAuthMethod = 'keyfile' | 'password' diff --git a/src/features/connections/components/ConnectionDialog.tsx b/src/features/connections/components/ConnectionDialog.tsx index 0455f7a..900b2e0 100644 --- a/src/features/connections/components/ConnectionDialog.tsx +++ b/src/features/connections/components/ConnectionDialog.tsx @@ -1,121 +1,115 @@ -import { useMemo, useState, type ReactNode } from 'react' +import { useMemo, useState, useEffect, type ReactNode } from 'react' import { useForm, useWatch } from 'react-hook-form' import { z } from 'zod' import { zodResolver } from '@hookform/resolvers/zod' import { useTranslation } from 'react-i18next' import { open as openFilePicker } from '@tauri-apps/plugin-dialog' -import { FolderOpenIcon } from '@phosphor-icons/react' +import { FolderOpenIcon, PlugsConnectedIcon } from '@phosphor-icons/react' import { Button } from '@/components/ui/button' import { - Dialog, - DialogContent, - DialogDescription, - DialogFooter, - DialogHeader, - DialogTitle, + Dialog, DialogContent, DialogDescription, + DialogFooter, DialogHeader, DialogTitle, } from '@/components/ui/dialog' import { Input } from '@/components/ui/input' import { - InputGroup, - InputGroupAddon, - InputGroupButton, - InputGroupInput, + InputGroup, InputGroupAddon, InputGroupButton, InputGroupInput, } from '@/components/ui/input-group' -import type { ConnectionInput, DatabaseEngine, SshAuthMethod } from '@/data/types' +import type { ConnectionInput, DatabaseEngine } from '@/data/types' import { cn } from '@/lib/utils' import { parseConnectionString, buildConnectionString } from '@/lib/connection-string' -const sslModeSchema = z.enum(['disable', 'prefer', 'require']) -const sshAuthMethodSchema = z.enum(['keyfile', 'password']) - -const connectionSchema = z - .object({ - name: z.string().min(2, 'Enter a connection name.'), - engine: z.enum(['postgres', 'mysql', 'sqlite']), - host: z.string(), - port: z.coerce.number().int().min(1).max(65535), - database: z.string(), - filePath: z.string().optional(), - user: z.string(), - password: z.string(), - sslMode: sslModeSchema, - sshEnabled: z.boolean(), - sshHost: z.string().optional(), - sshPort: z.coerce.number().int().min(1).max(65535).optional(), - sshUser: z.string().optional(), - sshAuthMethod: sshAuthMethodSchema.optional(), - sshPassword: z.string().optional(), - sshPrivateKeyPath: z.string().optional(), - sshPassphrase: z.string().optional(), - }) - .superRefine((values, ctx) => { - if (values.engine !== 'sqlite') { - if (!values.host || values.host.trim().length === 0) { - ctx.addIssue({ - code: z.ZodIssueCode.custom, - message: 'Host is required.', - path: ['host'], - }) - } - if (!values.database || values.database.trim().length === 0) { - ctx.addIssue({ - code: z.ZodIssueCode.custom, - message: 'Database is required.', - path: ['database'], - }) - } - if (!values.user || values.user.trim().length === 0) { - ctx.addIssue({ - code: z.ZodIssueCode.custom, - message: 'User is required.', - path: ['user'], - }) - } - if (!values.password || values.password.trim().length === 0) { - ctx.addIssue({ - code: z.ZodIssueCode.custom, - message: 'Password is required.', - path: ['password'], - }) - } - } - if (values.engine === 'sqlite') { +type InputMode = 'string' | 'fields' +type SshAuthMethodForm = 'keyfile' | 'password' +type EngineOption = { + value: DatabaseEngine + label: string + hint: string + experimental: boolean + defaultPort: number + defaultDatabase: string + defaultUser: string + requiresAuth: boolean +} + +const ENGINE_DEFAULTS: Record> = { + postgres: { hint: 'Recommended for production', experimental: false, defaultPort: 5432, defaultDatabase: 'postgres', defaultUser: 'postgres', requiresAuth: true }, + mysql: { hint: 'Experimental support', experimental: true, defaultPort: 3306, defaultDatabase: '', defaultUser: 'root', requiresAuth: true }, + sqlite: { hint: 'Experimental support', experimental: true, defaultPort: 0, defaultDatabase: '', defaultUser: '', requiresAuth: false }, + mongo: { hint: 'Experimental support', experimental: true, defaultPort: 27017, defaultDatabase: 'admin', defaultUser: '', requiresAuth: false }, + duckdb: { hint: 'Embedded OLAP', experimental: true, defaultPort: 0, defaultDatabase: '', defaultUser: '', requiresAuth: false }, + redis: { hint: 'Key-value store', experimental: true, defaultPort: 6379, defaultDatabase: '', defaultUser: '', requiresAuth: false }, +} + +const engineOptions: EngineOption[] = (Object.entries(ENGINE_DEFAULTS) as [DatabaseEngine, typeof ENGINE_DEFAULTS['postgres']][]).map( + ([value, defaults]) => ({ value, label: value === 'postgres' ? 'PostgreSQL' : value === 'mysql' ? 'MySQL' : value === 'sqlite' ? 'SQLite' : value === 'mongo' ? 'MongoDB' : value === 'duckdb' ? 'DuckDB' : 'Redis', ...defaults }) +) + +const connectionSchema = z.object({ + name: z.string().min(2, 'Enter a connection name.'), + engine: z.enum(['postgres', 'mysql', 'sqlite', 'mongo', 'duckdb', 'redis'] as const), + host: z.string(), + port: z.coerce.number().int().min(1).max(65535), + database: z.string(), + filePath: z.string().optional(), + user: z.string(), + password: z.string(), + sslMode: z.enum(['disable', 'prefer', 'require'] as const), + sshEnabled: z.boolean(), + sshHost: z.string().optional(), + sshPort: z.coerce.number().int().min(1).max(65535).optional(), + sshUser: z.string().optional(), + sshAuthMethod: z.enum(['keyfile', 'password'] as const).optional(), + sshPassword: z.string().optional(), + sshPrivateKeyPath: z.string().optional(), + sshPassphrase: z.string().optional(), +}).superRefine((values, ctx) => { + if (values.engine === 'sqlite' || values.engine === 'duckdb') { if (!values.filePath || values.filePath.trim().length === 0) { - ctx.addIssue({ - code: z.ZodIssueCode.custom, - message: 'SQLite file path is required.', - path: ['filePath'], - }) + ctx.addIssue({ code: z.ZodIssueCode.custom, message: 'File path is required.', path: ['filePath'] }) } return } - if (!values.sshEnabled) return - if (!values.sshHost || values.sshHost.trim().length === 0) { - ctx.addIssue({ - code: z.ZodIssueCode.custom, - message: 'SSH host is required.', - path: ['sshHost'], - }) + + if (values.engine === 'mongo') { + if (!values.host || values.host.trim().length === 0) { + ctx.addIssue({ code: z.ZodIssueCode.custom, message: 'Host is required.', path: ['host'] }) } - if (!values.sshUser || values.sshUser.trim().length === 0) { - ctx.addIssue({ - code: z.ZodIssueCode.custom, - message: 'SSH user is required.', - path: ['sshUser'], - }) + if (!values.database || values.database.trim().length === 0) { + ctx.addIssue({ code: z.ZodIssueCode.custom, message: 'Auth database is required.', path: ['database'] }) } - if (values.sshAuthMethod === 'password') { - if (!values.sshPassword || values.sshPassword.trim().length === 0) { - ctx.addIssue({ - code: z.ZodIssueCode.custom, - message: 'SSH password is required.', - path: ['sshPassword'], - }) - } + // Password and user are optional for MongoDB + return + } + + // PostgreSQL and MySQL + if (!values.host || values.host.trim().length === 0) { + ctx.addIssue({ code: z.ZodIssueCode.custom, message: 'Host is required.', path: ['host'] }) + } + if (!values.database || values.database.trim().length === 0) { + ctx.addIssue({ code: z.ZodIssueCode.custom, message: 'Database is required.', path: ['database'] }) + } + if (!values.user || values.user.trim().length === 0) { + ctx.addIssue({ code: z.ZodIssueCode.custom, message: 'User is required.', path: ['user'] }) + } + if (!values.password || values.password.trim().length === 0) { + ctx.addIssue({ code: z.ZodIssueCode.custom, message: 'Password is required.', path: ['password'] }) + } + + if (!values.sshEnabled) return + if (!values.sshHost || values.sshHost.trim().length === 0) { + ctx.addIssue({ code: z.ZodIssueCode.custom, message: 'SSH host is required.', path: ['sshHost'] }) + } + if (!values.sshUser || values.sshUser.trim().length === 0) { + ctx.addIssue({ code: z.ZodIssueCode.custom, message: 'SSH user is required.', path: ['sshUser'] }) + } + if (values.sshAuthMethod === 'password') { + if (!values.sshPassword || values.sshPassword.trim().length === 0) { + ctx.addIssue({ code: z.ZodIssueCode.custom, message: 'SSH password is required.', path: ['sshPassword'] }) } - }) + } +}) type ConnectionDialogProps = { open: boolean @@ -124,92 +118,115 @@ type ConnectionDialogProps = { isPending?: boolean } -function Field({ - label, - error, - inputId, - children, -}: { +// ── Database Flavors ──────────────────────────────────────────── + +type DatabaseFlavor = { + key: string label: string - error?: string - inputId: string - children: ReactNode -}) { + engine: DatabaseEngine + defaultHost: string + defaultPort: number + defaultDatabase: string + defaultUser: string + defaultSsl: 'disable' | 'prefer' | 'require' + description: string +} + +const DATABASE_FLAVORS: DatabaseFlavor[] = [ + // PostgreSQL wire-compatible + { key: 'postgres', label: 'PostgreSQL', engine: 'postgres', defaultHost: '127.0.0.1', defaultPort: 5432, defaultDatabase: 'postgres', defaultUser: 'postgres', defaultSsl: 'prefer', description: 'Local or self-hosted PostgreSQL' }, + { key: 'supabase', label: 'Supabase', engine: 'postgres', defaultHost: 'db.xxxxx.supabase.co', defaultPort: 5432, defaultDatabase: 'postgres', defaultUser: 'postgres', defaultSsl: 'require', description: 'Hosted PostgreSQL with real-time, auth, storage' }, + { key: 'neon', label: 'Neon', engine: 'postgres', defaultHost: 'ep-xxxxx.us-east-1.aws.neon.tech', defaultPort: 5432, defaultDatabase: 'neondb', defaultUser: 'neondb_owner', defaultSsl: 'require', description: 'Serverless PostgreSQL with branching' }, + { key: 'cockroachdb', label: 'CockroachDB', engine: 'postgres', defaultHost: '127.0.0.1', defaultPort: 26257, defaultDatabase: 'defaultdb', defaultUser: 'root', defaultSsl: 'require', description: 'Distributed SQL, PostgreSQL-compatible' }, + { key: 'timescaledb', label: 'TimescaleDB', engine: 'postgres', defaultHost: '127.0.0.1', defaultPort: 5432, defaultDatabase: 'postgres', defaultUser: 'postgres', defaultSsl: 'prefer', description: 'Time-series PostgreSQL extension' }, + { key: 'yugabytedb', label: 'YugabyteDB', engine: 'postgres', defaultHost: '127.0.0.1', defaultPort: 5433, defaultDatabase: 'yugabyte', defaultUser: 'yugabyte', defaultSsl: 'prefer', description: 'Distributed PostgreSQL-compatible' }, + // MySQL wire-compatible + { key: 'mysql', label: 'MySQL', engine: 'mysql', defaultHost: '127.0.0.1', defaultPort: 3306, defaultDatabase: '', defaultUser: 'root', defaultSsl: 'prefer', description: 'Local or self-hosted MySQL' }, + { key: 'mariadb', label: 'MariaDB', engine: 'mysql', defaultHost: '127.0.0.1', defaultPort: 3306, defaultDatabase: '', defaultUser: 'root', defaultSsl: 'prefer', description: 'Community fork of MySQL' }, + { key: 'planetscale', label: 'PlanetScale', engine: 'mysql', defaultHost: 'aws.connect.psdb.cloud', defaultPort: 3306, defaultDatabase: '', defaultUser: 'root', defaultSsl: 'require', description: 'Serverless MySQL platform' }, + { key: 'tidb', label: 'TiDB', engine: 'mysql', defaultHost: '127.0.0.1', defaultPort: 4000, defaultDatabase: 'test', defaultUser: 'root', defaultSsl: 'prefer', description: 'Distributed MySQL-compatible (HTAP)' }, + { key: 'aurora', label: 'Aurora MySQL', engine: 'mysql', defaultHost: 'xxxxx.cluster-xxx.us-east-1.rds.amazonaws.com', defaultPort: 3306, defaultDatabase: '', defaultUser: 'admin', defaultSsl: 'require', description: 'AWS MySQL-compatible cloud database' }, +] + +function flavorsForEngine(engine: DatabaseEngine): DatabaseFlavor[] { + return DATABASE_FLAVORS.filter((f) => f.engine === engine) +} + +function Field({ label, error, inputId, children }: { label: string; error?: string; inputId: string; children: ReactNode }) { return ( -