diff --git a/config/localization-coverage-allowlist.json b/config/localization-coverage-allowlist.json index fe51488c706..f496c3b7402 100644 --- a/config/localization-coverage-allowlist.json +++ b/config/localization-coverage-allowlist.json @@ -1 +1,38 @@ -[] +[ + { + "area": "src/renderer", + "filePath": "src/renderer/src/components/database/data-grid-filters.ts", + "kind": "object-property:label", + "text": "LIKE", + "dynamic": false, + "count": 1, + "reason": "SQL operator keyword shown verbatim in the filter dropdown; not translatable copy." + }, + { + "area": "src/renderer", + "filePath": "src/renderer/src/components/database/data-grid-filters.ts", + "kind": "object-property:label", + "text": "ILIKE", + "dynamic": false, + "count": 1, + "reason": "SQL operator keyword shown verbatim in the filter dropdown; not translatable copy." + }, + { + "area": "src/renderer", + "filePath": "src/renderer/src/components/database/data-grid-filters.ts", + "kind": "object-property:label", + "text": "IS NULL", + "dynamic": false, + "count": 1, + "reason": "SQL operator keyword shown verbatim in the filter dropdown; not translatable copy." + }, + { + "area": "src/renderer", + "filePath": "src/renderer/src/components/database/data-grid-filters.ts", + "kind": "object-property:label", + "text": "IS NOT NULL", + "dynamic": false, + "count": 1, + "reason": "SQL operator keyword shown verbatim in the filter dropdown; not translatable copy." + } +] diff --git a/config/packaged-runtime-node-modules.cjs b/config/packaged-runtime-node-modules.cjs index d3f81817cee..d139c956a70 100644 --- a/config/packaged-runtime-node-modules.cjs +++ b/config/packaged-runtime-node-modules.cjs @@ -20,7 +20,12 @@ const PACKAGED_RUNTIME_PACKAGE_ROOTS = [ 'electron-updater', 'i18next', 'jsonc-parser', + 'mysql2', 'node-pty', + // Database client drivers — bundled + lazy-imported (await import) from the main + // process; absent from this allowlist => MODULE_NOT_FOUND in packaged builds. The + // graph walk pulls their transitive deps (pg-*, denque, named-placeholders, …). + 'pg', 'posthog-node', // serve-sim (for CLI JS entry + closure + state/middleware + to make packaged require('serve-sim') + its internal relatives work; mirrors other runtime JS like ws/yaml/zod. Natives/dylibs still via extraResources + the node_modules/serve-sim copy in resources from builder. Client if added too. 'serve-sim', diff --git a/package.json b/package.json index 88f598a7d08..271ca384387 100644 --- a/package.json +++ b/package.json @@ -99,7 +99,9 @@ "electron-updater": "^6.8.3", "i18next": "^26.3.1", "jsonc-parser": "^3.3.1", + "mysql2": "^3.22.5", "node-pty": "^1.1.0", + "pg": "^8.22.0", "posthog-node": "^5.33.3", "qrcode": "^1.5.4", "react-i18next": "^17.0.8", @@ -141,6 +143,7 @@ "@tiptap/react": "^3.22.5", "@tiptap/starter-kit": "^3.22.5", "@types/node": "^25.6.0", + "@types/pg": "^8.20.0", "@types/qrcode": "^1.5.6", "@types/react": "^19.2.14", "@types/react-dom": "^19.2.3", diff --git a/pnpm-lock.yaml b/pnpm-lock.yaml index 12052fba2e0..c06c87fa0f9 100644 --- a/pnpm-lock.yaml +++ b/pnpm-lock.yaml @@ -52,9 +52,15 @@ importers: jsonc-parser: specifier: ^3.3.1 version: 3.3.1 + mysql2: + specifier: ^3.22.5 + version: 3.22.5(@types/node@25.6.0) node-pty: specifier: ^1.1.0 version: 1.1.0(patch_hash=407ae07e1e0e2ff2e8b58696449c54c31e51d87535bc6aa4a7a7b0b561407282) + pg: + specifier: ^8.22.0 + version: 8.22.0 posthog-node: specifier: ^5.33.3 version: 5.33.3 @@ -173,6 +179,9 @@ importers: '@types/node': specifier: ^25.6.0 version: 25.6.0 + '@types/pg': + specifier: ^8.20.0 + version: 8.20.0 '@types/qrcode': specifier: ^1.5.6 version: 1.5.6 @@ -3030,6 +3039,9 @@ packages: '@types/node@25.6.0': resolution: {integrity: sha512-+qIYRKdNYJwY3vRCZMdJbPLJAtGjQBudzZzdzwQYkEPQd+PJGixUL5QfvCLDaULoLv+RhT3LDkwEfKaAkgSmNQ==} + '@types/pg@8.20.0': + resolution: {integrity: sha512-bEPFOaMAHTEP1EzpvHTbmwR8UsFyHSKsRisLIHVMXnpNefSbGA1bD6CVy+qKjGSqmZqNqBDV2azOBo8TgkcVow==} + '@types/plist@3.0.5': resolution: {integrity: sha512-E6OCaRmAe4WDmWNsL/9RMqdkkzDCY1etutkflWk4c+AcjDU07Pcz1fQwTX0TQz+Pxqn9i4L1TU3UFpjnrcDgxA==} @@ -3337,6 +3349,10 @@ packages: resolution: {integrity: sha512-+q/t7Ekv1EDY2l6Gda6LLiX14rU9TV20Wa3ofeQmwPFZbOMo9DXrLbOjFaaclkXKWidIaopwAObQDqwWtGUjqg==} engines: {node: '>= 4.0.0'} + aws-ssl-profiles@1.1.2: + resolution: {integrity: sha512-NZKeq9AfyQvEeNlN0zSYAaWrmBffJh3IELMZfRpJVWgrpEbtEpnjvzqBPf+mxoI287JohRDoa+/nsfqqiZmF6g==} + engines: {node: '>= 6.0.0'} + bail@2.0.2: resolution: {integrity: sha512-0xO6mYd7JB2YesxDKplafRpsiOzPt9V02ddPCLbY1xYGPOX24NTyN50qnUxgCPcSoYMhKpAuBTjQoRZCAkUDRw==} @@ -3892,6 +3908,10 @@ packages: resolution: {integrity: sha512-ZySD7Nf91aLB0RxL4KGrKHBXl7Eds1DAmEdcoVawXnLD7SDhpNgtuII2aAkg7a7QS41jxPSZ17p4VdGnMHk3MQ==} engines: {node: '>=0.4.0'} + denque@2.1.0: + resolution: {integrity: sha512-HVQE3AAb/pxF8fQAoiqpvg9i3evqug3hoiwakOyZAwJm+6vZehbkYXZ0l4JxS+I3QxM97v5aaRNhj8v5oBhekw==} + engines: {node: '>=0.10'} + depd@2.0.0: resolution: {integrity: sha512-g7nH6P6dyDioJogAAGprGpCtVImJhpPk/roCzdb3fIh61/s/nPsfR6onyMwkCAR/OlC3yBC0lESvUoQEAssIrw==} engines: {node: '>= 0.8'} @@ -4320,6 +4340,9 @@ packages: fuzzysort@3.1.0: resolution: {integrity: sha512-sR9BNCjBg6LNgwvxlBd0sBABvQitkLzoVY9MYYROQVX/FvfJ4Mai9LsGhDgd8qYdds0bY77VzYd5iuB+v5rwQQ==} + generate-function@2.3.1: + resolution: {integrity: sha512-eeB5GfMNeevm/GRYq20ShmsaGcmI81kIX2K9XQx5miC8KdHaC6Jm0qQ8ZNeGOi7wYB8OsdxKs+Y2oVuTFuVwKQ==} + gensync@1.0.0-beta.2: resolution: {integrity: sha512-3hN7NaskYvMDLQY55gnW3NQ+mesEAepTqlg+VEbj7zzqEMBVNhzcGYYeqFo/TlYz6eQiFcp1HcsCZO+nGgS8zg==} engines: {node: '>=6.9.0'} @@ -4662,6 +4685,9 @@ packages: is-promise@4.0.0: resolution: {integrity: sha512-hvpoI6korhJMnej285dSg6nu1+e6uxs7zG3BYAm5byqDsgJNWwxzM6z6iZiAgQR4TJ30JmBTOwqZUw3WlyH3AQ==} + is-property@1.0.2: + resolution: {integrity: sha512-Ks/IoX00TtClbGQr4TWXemAnktAQvYB7HzcCxDGqEZU6oCmb2INHuOoKxbtR+HFkmYWBKv/dOZtGRiAjDhj92g==} + is-regexp@3.1.0: resolution: {integrity: sha512-rbku49cWloU5bSMI+zaRaXdQHXnthP6DZ/vLnfdSKyL4zUzuWnomtOEiZZOd+ioQ+avFo/qau3KPTc7Fjy1uPA==} engines: {node: '>=12'} @@ -4909,6 +4935,9 @@ packages: resolution: {integrity: sha512-9ie8ItPR6tjY5uYJh8K/Zrv/RMZ5VOlOWvtZdEHYSTFKZfIBPQa9tOAEeAWhd+AnIneLJ22w5fjOYtoutpWq5w==} engines: {node: '>=18'} + long@5.3.2: + resolution: {integrity: sha512-mNAgZ1GmyNhD7AuqnTG3/VQ26o760+ZYBPKjPvugO8+nLbYfX6TVpJPseBvopbdY+qpZ/lKUnmEc1LeZYS3QAA==} + longest-streak@3.1.0: resolution: {integrity: sha512-9Ri+o0JYgehTaVBBDoMqIl8GXtbWg711O3srftcHhZ0dqnETqLaoIK0x17fUw9rFSlK/0NlsKe0Ahhyl5pXE2g==} @@ -4930,6 +4959,10 @@ packages: resolution: {integrity: sha512-Jo6dJ04CmSjuznwJSS3pUeWmd/H0ffTlkXXgwZi+eq1UCmqQwCh+eLsYOYCwY991i2Fah4h1BEMCx4qThGbsiA==} engines: {node: '>=10'} + lru.min@1.1.4: + resolution: {integrity: sha512-DqC6n3QQ77zdFpCMASA1a3Jlb64Hv2N2DciFGkO/4L9+q/IpIAuRlKOvCXabtRW6cQf8usbmM6BE/TOPysCdIA==} + engines: {bun: '>=1.0.0', deno: '>=1.30.0', node: '>=8.0.0'} + lucide-react@0.577.0: resolution: {integrity: sha512-4LjoFv2eEPwYDPg/CUdBJQSDfPyzXCRrVW1X7jrx/trgxnxkHFjnVZINbzvzxjN70dxychOfg+FTYwBiS3pQ5A==} peerDependencies: @@ -5228,6 +5261,16 @@ packages: resolution: {integrity: sha512-dkEJPVvun4FryqBmZ5KhDo0K9iDXAwn08tMLDinNdRBNPcYEDiWYysLcc6k3mjTMlbP9KyylvRpd4wFtwrT9rw==} engines: {node: ^20.17.0 || >=22.9.0} + mysql2@3.22.5: + resolution: {integrity: sha512-95uZ2TrPWAZdwpB3vvvDbmEMcNG8yIeNCyu6GUcr/QnWEE/wXm7+mhOCsdQfWQDTV7qYT/PDUZ4U4UPP4AsXqQ==} + engines: {node: '>= 8.0'} + peerDependencies: + '@types/node': '>= 8' + + named-placeholders@1.1.6: + resolution: {integrity: sha512-Tz09sEL2EEuv5fFowm419c1+a/jSMiBjI9gHxVLrVdbUkkNUUfjsVYs9pVZu5oCon/kmRh9TfLEObFtkVxmY0w==} + engines: {node: '>=8.0.0'} + nan@2.26.2: resolution: {integrity: sha512-0tTvBTYkt3tdGw22nrAy50x7gpbGCCFH3AFcyS5WiUu7Eu4vWlri1woE6qHBSfy11vksDqkiwjOnlR7WV8G1Hw==} @@ -5471,6 +5514,40 @@ packages: pend@1.2.0: resolution: {integrity: sha512-F3asv42UuXchdzt+xXqfW1OGlVBe+mxa2mqI0pg5yAHZPvFmY3Y6drSf/GQ1A86WgWEN9Kzh/WrgKa6iGcHXLg==} + pg-cloudflare@1.4.0: + resolution: {integrity: sha512-Vo7z/6rrQYxpNRylp4Tlob2elzbh+N/MOQbxFVWCxS7oEx6jF53GTJFxK2WWpKuBRkmiin4Mt+xofFDjx09R0A==} + + pg-connection-string@2.14.0: + resolution: {integrity: sha512-XwWDGcLRGCXAR8F/AM5bG7Q+A3Wm2s6QeEjlOKZLlH3UYcguiqCWKyWXVag5TLTIjR7oOJUY8kcADaZgWPyLeg==} + + pg-int8@1.0.1: + resolution: {integrity: sha512-WCtabS6t3c8SkpDBUlb1kjOs7l66xsGdKpIPZsg4wR+B3+u9UAum2odSsF9tnvxg80h4ZxLWMy4pRjOsFIqQpw==} + engines: {node: '>=4.0.0'} + + pg-pool@3.14.0: + resolution: {integrity: sha512-gKtPkFdQPU3DksooVLi9LsjZxrsBUZIpa+7aVx+LV5pNh0KzP4Zleud2po+ConrxbuXGBJ6Hfer6hdgpIBpBaw==} + peerDependencies: + pg: '>=8.0' + + pg-protocol@1.15.0: + resolution: {integrity: sha512-cq9sECI5s0+uPUXjbz8ioyPJni6RzsRib0US67i5IoTZKw8fNeYlVE7u8F4dG7vEJJtc5wdD1K189lCCUwqWTQ==} + + pg-types@2.2.0: + resolution: {integrity: sha512-qTAAlrEsl8s4OiEQY69wDvcMIdQN6wdz5ojQiOy6YRMuynxenON0O5oCpJI6lshc6scgAY8qvJ2On/p+CXY0GA==} + engines: {node: '>=4'} + + pg@8.22.0: + resolution: {integrity: sha512-8wih1vVIBMxoUM2oB4soJsD9tDnDpLv4OXBJ+EJzFsvycD+lfyIreC2gGHq78f8jbLLt+bvlPTFdFZfJkOuzAA==} + engines: {node: '>= 16.0.0'} + peerDependencies: + pg-native: '>=3.0.1' + peerDependenciesMeta: + pg-native: + optional: true + + pgpass@1.0.5: + resolution: {integrity: sha512-FdW9r/jQZhSeohs1Z3sI1yxFQNFvMcnmfuj4WBMUTxOrAyLMaTcE1aAMBiTlbMNaXvBCQuVi0R7hd8udDSP7ug==} + picocolors@1.1.1: resolution: {integrity: sha512-xceH2snhtb5M9liqDsmEw56le376mTZkEX/jEb/RxNFyegNul7eNslCXP9FDj/Lcu0X8KEyMceP2ntpaHrDEVA==} @@ -5529,6 +5606,22 @@ packages: resolution: {integrity: sha512-SoSL4+OSEtR99LHFZQiJLkT59C5B1amGO1NzTwj7TT1qCUgUO6hxOvzkOYxD+vMrXBM3XJIKzokoERdqQq/Zmg==} engines: {node: ^10 || ^12 || >=14} + postgres-array@2.0.0: + resolution: {integrity: sha512-VpZrUqU5A69eQyW2c5CA1jtLecCsN2U/bD6VilrFDWq5+5UIEVO7nazS3TEcHf1zuPYO/sqGvUvW62g86RXZuA==} + engines: {node: '>=4'} + + postgres-bytea@1.0.1: + resolution: {integrity: sha512-5+5HqXnsZPE65IJZSMkZtURARZelel2oXUEO8rH83VS/hxH5vv1uHquPg5wZs8yMAfdv971IU+kcPUczi7NVBQ==} + engines: {node: '>=0.10.0'} + + postgres-date@1.0.7: + resolution: {integrity: sha512-suDmjLVQg78nMK2UZ454hAG+OAW+HQPZ6n++TNDUX+L0+uUlLywnoxJKDou51Zm+zTCjrCl0Nq6J9C5hP9vK/Q==} + engines: {node: '>=0.10.0'} + + postgres-interval@1.2.0: + resolution: {integrity: sha512-9ZhXKM/rw350N1ovuWHbGxnGh/SNJ4cnxHiM0rxE4VN41wsg8P8zWn9hv/buK00RP4WvlOyr/RBDiptyxVbkZQ==} + engines: {node: '>=0.10.0'} + posthog-node@5.33.3: resolution: {integrity: sha512-6BkdRtRKf/y/j6JGXLUDWpruL9Rebpe2EtbSU3EB+yaI9ukTwEGT8XJQDyM8wcsxu1HdOnvAwN7gnOOu5OgY2A==} engines: {node: ^20.20.0 || >=22.22.0} @@ -6049,9 +6142,17 @@ packages: space-separated-tokens@2.0.2: resolution: {integrity: sha512-PEGlAwrG8yXGXRjW32fGbg66JAlOAwbObuqVoJpv/mRgoWDQfgH1wDPvtzWyUSNAXBGSk8h755YDbbcEy3SH2Q==} + split2@4.2.0: + resolution: {integrity: sha512-UcjcJOWknrNkF6PLX83qcHM6KHgVKNkV62Y8a5uYDVv9ydGQVwAHMKqHdJje1VTWpljG0WYpCDhrCdAOYH4TWg==} + engines: {node: '>= 10.x'} + sprintf-js@1.1.3: resolution: {integrity: sha512-Oo+0REFV59/rz3gfJNKQiBlwfHaSESl1pcGyABQsnnIfWOFt6JNj5gCog2U6MLZ//IGYD+nA8nI+mTShREReaA==} + sql-escaper@1.3.3: + resolution: {integrity: sha512-BsTCV265VpTp8tm1wyIm1xqQCS+Q9NHx2Sr+WcnUrgLrQ6yiDIvHYJV5gHxsj1lMBy2zm5twLaZao8Jd+S8JJw==} + engines: {bun: '>=1.0.0', deno: '>=2.0.0', node: '>=12.0.0'} + ssh2@1.17.0: resolution: {integrity: sha512-wPldCk3asibAjQ/kziWQQt1Wh3PgDFpC0XpwclzKcdT1vql6KeYxf5LIt4nlFkUeR8WuphYMKqUA56X4rjbfgQ==} engines: {node: '>=10.16.0'} @@ -6572,6 +6673,10 @@ packages: resolution: {integrity: sha512-yMqGBqtXyeN1e3TGYvgNgDVZ3j84W4cwkOXQswghol6APgZWaff9lnbvN7MHYJOiXsvGPXtjTYJEiC9J2wv9Eg==} engines: {node: '>=8.0'} + xtend@4.0.2: + resolution: {integrity: sha512-LKYU1iAXJXUgAXn9URjiu+MWhyUXHsvfp7mcuYm9dSUKK0/CjtrUwFAxD82/mCWbtLsGjFIad0wIsod4zrTAEQ==} + engines: {node: '>=0.4'} + y18n@4.0.3: resolution: {integrity: sha512-JKhqTOwSrqNA1NY5lSztJ1GrBiUodLMmIZuLiDaMRJ+itFd+ABVE8XBjOvIWL+rSqNDC74LCSFmlb/U4UZ4hJQ==} @@ -9116,6 +9221,12 @@ snapshots: dependencies: undici-types: 7.19.2 + '@types/pg@8.20.0': + dependencies: + '@types/node': 25.6.0 + pg-protocol: 1.15.0 + pg-types: 2.2.0 + '@types/plist@3.0.5': dependencies: '@types/node': 25.6.0 @@ -9438,6 +9549,8 @@ snapshots: at-least-node@1.0.0: {} + aws-ssl-profiles@1.1.2: {} + bail@2.0.2: {} balanced-match@1.0.2: {} @@ -9994,6 +10107,8 @@ snapshots: delayed-stream@1.0.0: {} + denque@2.1.0: {} + depd@2.0.0: {} dequal@2.0.3: {} @@ -10553,6 +10668,10 @@ snapshots: fuzzysort@3.1.0: {} + generate-function@2.3.1: + dependencies: + is-property: 1.0.2 + gensync@1.0.0-beta.2: {} get-caller-file@2.0.5: {} @@ -10950,6 +11069,8 @@ snapshots: is-promise@4.0.0: {} + is-property@1.0.2: {} + is-regexp@3.1.0: {} is-stream@2.0.1: {} @@ -11142,6 +11263,8 @@ snapshots: strip-ansi: 7.2.0 wrap-ansi: 9.0.2 + long@5.3.2: {} + longest-streak@3.1.0: {} lowercase-keys@2.0.0: {} @@ -11162,6 +11285,8 @@ snapshots: dependencies: yallist: 4.0.0 + lru.min@1.1.4: {} + lucide-react@0.577.0(react@19.2.5): dependencies: react: 19.2.5 @@ -11706,6 +11831,22 @@ snapshots: mute-stream@3.0.0: {} + mysql2@3.22.5(@types/node@25.6.0): + dependencies: + '@types/node': 25.6.0 + aws-ssl-profiles: 1.1.2 + denque: 2.1.0 + generate-function: 2.3.1 + iconv-lite: 0.7.2 + long: 5.3.2 + lru.min: 1.1.4 + named-placeholders: 1.1.6 + sql-escaper: 1.3.3 + + named-placeholders@1.1.6: + dependencies: + lru.min: 1.1.4 + nan@2.26.2: optional: true @@ -11973,6 +12114,41 @@ snapshots: pend@1.2.0: {} + pg-cloudflare@1.4.0: + optional: true + + pg-connection-string@2.14.0: {} + + pg-int8@1.0.1: {} + + pg-pool@3.14.0(pg@8.22.0): + dependencies: + pg: 8.22.0 + + pg-protocol@1.15.0: {} + + pg-types@2.2.0: + dependencies: + pg-int8: 1.0.1 + postgres-array: 2.0.0 + postgres-bytea: 1.0.1 + postgres-date: 1.0.7 + postgres-interval: 1.2.0 + + pg@8.22.0: + dependencies: + pg-connection-string: 2.14.0 + pg-pool: 3.14.0(pg@8.22.0) + pg-protocol: 1.15.0 + pg-types: 2.2.0 + pgpass: 1.0.5 + optionalDependencies: + pg-cloudflare: 1.4.0 + + pgpass@1.0.5: + dependencies: + split2: 4.2.0 + picocolors@1.1.1: {} picomatch@2.3.2: {} @@ -12030,6 +12206,16 @@ snapshots: picocolors: 1.1.1 source-map-js: 1.2.1 + postgres-array@2.0.0: {} + + postgres-bytea@1.0.1: {} + + postgres-date@1.0.7: {} + + postgres-interval@1.2.0: + dependencies: + xtend: 4.0.2 + posthog-node@5.33.3: dependencies: '@posthog/core': 1.28.3 @@ -12747,9 +12933,13 @@ snapshots: space-separated-tokens@2.0.2: {} + split2@4.2.0: {} + sprintf-js@1.1.3: optional: true + sql-escaper@1.3.3: {} + ssh2@1.17.0: dependencies: asn1: 0.2.6 @@ -13191,6 +13381,8 @@ snapshots: xmlbuilder@15.1.1: {} + xtend@4.0.2: {} + y18n@4.0.3: {} y18n@5.0.8: {} diff --git a/src/main/database/db-connection-manager.test.ts b/src/main/database/db-connection-manager.test.ts new file mode 100644 index 00000000000..fea0bfb01b4 --- /dev/null +++ b/src/main/database/db-connection-manager.test.ts @@ -0,0 +1,377 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest' +import type { DbConnectionRuntimeState, QueryHandle } from '../../shared/database-types' +import type { LiveConnection, ResolvedDbConfig } from './db-driver' + +// Controllable fake drivers so the manager's lifecycle can be exercised without +// pg/mysql2. `capturedOnError` lets a test simulate a dropped connection. +const { pgDriver, mysqlDriver, state } = vi.hoisted(() => { + const shared = { capturedOnError: null as ((err: unknown) => void) | null } + const makeDriver = (engine: 'postgres' | 'mysql') => ({ + testConnection: vi.fn(async () => {}), + connect: vi.fn(async (cfg: ResolvedDbConfig, onError: (err: unknown) => void) => { + shared.capturedOnError = onError + return { id: cfg.id, engine, raw: {} } as LiveConnection + }), + introspectSchemas: vi.fn(async () => ({ schemas: ['public'], truncated: false })), + introspectTables: vi.fn(async () => ({ tables: [], truncated: false })), + introspectColumns: vi.fn(async () => []), + query: vi.fn( + async ( + conn: LiveConnection, + _sql: string, + _opts: unknown, + onStart: (h: { connectionId: string; backendPid: number | null }) => void + ) => { + onStart({ connectionId: conn.id, backendPid: 123 }) + return { columns: [], rows: [], rowCount: 0, truncated: false, durationMs: 1 } + } + ), + execute: vi.fn( + async ( + conn: LiveConnection, + _statement: unknown, + _opts: unknown, + onStart: (h: { connectionId: string; backendPid: number | null }) => void + ) => { + onStart({ connectionId: conn.id, backendPid: 123 }) + return { columns: [], rows: [], rowCount: 0, truncated: false, durationMs: 1 } + } + ), + executeBatch: vi.fn( + async ( + conn: LiveConnection, + statements: unknown[], + _opts: unknown, + onStart: (h: { connectionId: string; backendPid: number | null }) => void + ) => { + onStart({ connectionId: conn.id, backendPid: 123 }) + return statements.map(() => 1) + } + ), + cancel: vi.fn(async () => {}), + close: vi.fn(async () => {}) + }) + return { pgDriver: makeDriver('postgres'), mysqlDriver: makeDriver('mysql'), state: shared } +}) + +vi.mock('./postgres-driver', () => ({ postgresDriver: pgDriver })) +vi.mock('./mysql-driver', () => ({ mysqlDriver })) + +import { DbConnectionManager } from './db-connection-manager' + +function cfg(overrides: Partial = {}): ResolvedDbConfig { + return { + id: 'c1', + engine: 'postgres', + host: 'db.example.com', + port: 5432, + database: 'app', + user: 'admin', + password: 'pw', + ssl: 'verify-full', + readOnly: false, + ...overrides + } +} + +describe('DbConnectionManager', () => { + let manager: DbConnectionManager + let statuses: DbConnectionRuntimeState[] + + beforeEach(() => { + vi.clearAllMocks() + state.capturedOnError = null + manager = new DbConnectionManager() + statuses = [] + manager.setStatusListener((s) => statuses.push(s)) + }) + + it('connect transitions connecting → connected and holds the connection', async () => { + const result = await manager.connect(cfg()) + expect(result.status).toBe('connected') + expect(manager.getConnection('c1')).toBeDefined() + expect(statuses.map((s) => s.status)).toEqual(['connecting', 'connected']) + }) + + it('routes by engine to the matching driver', async () => { + await manager.connect(cfg({ id: 'm1', engine: 'mysql' })) + expect(mysqlDriver.connect).toHaveBeenCalledTimes(1) + expect(pgDriver.connect).not.toHaveBeenCalled() + }) + + it('rejects a concurrent connect for the same id (race guard)', async () => { + let release: (() => void) | undefined + pgDriver.connect.mockImplementationOnce( + (c: ResolvedDbConfig) => + new Promise((resolve) => { + release = () => resolve({ id: c.id, engine: 'postgres', raw: {} } as LiveConnection) + }) + ) + const first = manager.connect(cfg()) + await expect(manager.connect(cfg())).rejects.toThrow('db_connect_in_progress') + release?.() + await first + }) + + it('returns the existing state without re-dialing when already connected', async () => { + await manager.connect(cfg()) + await manager.connect(cfg()) + expect(pgDriver.connect).toHaveBeenCalledTimes(1) + }) + + it('marks error and rethrows when the driver fails to connect', async () => { + pgDriver.connect.mockRejectedValueOnce(Object.assign(new Error('x'), { code: '28P01' })) + await expect(manager.connect(cfg())).rejects.toThrow() + const last = statuses.at(-1) + expect(last?.status).toBe('error') + expect(last?.error?.code).toBe('auth_failed') + expect(manager.getConnection('c1')).toBeUndefined() + }) + + it('degrades a dropped connection to lost without crashing, and drops it', async () => { + await manager.connect(cfg()) + expect(state.capturedOnError).toBeTypeOf('function') + // Simulate the driver 'error' event on a live connection. + state.capturedOnError?.(Object.assign(new Error('reset'), { code: 'ECONNRESET' })) + expect(manager.getStatus('c1').status).toBe('lost') + expect(manager.getConnection('c1')).toBeUndefined() + expect(pgDriver.close).toHaveBeenCalledTimes(1) + }) + + it('ignores a late error from a superseded connection', async () => { + // First connect captures onError bound to connection #1. + await manager.connect(cfg()) + const staleOnError = state.capturedOnError + // Disconnect + reconnect under the same id → a fresh live connection #2. + await manager.disconnect('c1') + await manager.connect(cfg()) + expect(manager.getConnection('c1')).toBeDefined() + // The OLD pool now emits a late error; it must NOT tear down connection #2. + staleOnError?.(Object.assign(new Error('late'), { code: 'ECONNRESET' })) + expect(manager.getStatus('c1').status).toBe('connected') + expect(manager.getConnection('c1')).toBeDefined() + }) + + it('disconnect closes the pool and resets status to idle', async () => { + await manager.connect(cfg()) + await manager.disconnect('c1') + expect(pgDriver.close).toHaveBeenCalledTimes(1) + expect(manager.getStatus('c1').status).toBe('idle') + expect(manager.getConnection('c1')).toBeUndefined() + }) + + it('disconnectAll closes every held connection', async () => { + await manager.connect(cfg({ id: 'c1' })) + await manager.connect(cfg({ id: 'm1', engine: 'mysql' })) + await manager.disconnectAll() + expect(pgDriver.close).toHaveBeenCalledTimes(1) + expect(mysqlDriver.close).toHaveBeenCalledTimes(1) + expect(manager.getConnection('c1')).toBeUndefined() + expect(manager.getConnection('m1')).toBeUndefined() + }) + + it('test performs a one-shot ping and holds no state', async () => { + await manager.test(cfg()) + expect(pgDriver.testConnection).toHaveBeenCalledTimes(1) + expect(manager.getConnection('c1')).toBeUndefined() + expect(manager.getStatus('c1').status).toBe('idle') + }) + + describe('introspection', () => { + it('throws db_not_connected when no live connection is held', async () => { + await expect(manager.introspectSchemas('missing')).rejects.toThrow('db_not_connected') + await expect(manager.introspectTables('missing', 'public')).rejects.toThrow( + 'db_not_connected' + ) + await expect( + manager.introspectColumns('missing', { schema: 'public', table: 't' }) + ).rejects.toThrow('db_not_connected') + }) + + it('delegates to the connection engine driver with caps', async () => { + await manager.connect(cfg()) + await manager.introspectSchemas('c1') + await manager.introspectTables('c1', 'public') + await manager.introspectColumns('c1', { schema: 'public', table: 'users' }) + expect(pgDriver.introspectSchemas).toHaveBeenCalledWith( + expect.objectContaining({ id: 'c1' }), + expect.any(Number) + ) + expect(pgDriver.introspectTables).toHaveBeenCalledWith( + expect.anything(), + 'public', + expect.any(Number) + ) + expect(pgDriver.introspectColumns).toHaveBeenCalledWith(expect.anything(), { + schema: 'public', + table: 'users' + }) + }) + }) + + describe('query + cancel', () => { + it('runs a query via the driver and returns its result', async () => { + await manager.connect(cfg()) + const result = await manager.query('c1', 'SELECT 1', { + rowLimit: 100, + timeoutMs: 1000, + allowWrite: false + }) + expect(pgDriver.query).toHaveBeenCalledTimes(1) + expect(result.truncated).toBe(false) + }) + + it('cancels the in-flight query with the captured backend handle', async () => { + await manager.connect(cfg()) + let release: (() => void) | undefined + pgDriver.query.mockImplementationOnce( + (conn: LiveConnection, _sql: string, _opts: unknown, onStart: (h: QueryHandle) => void) => { + onStart({ connectionId: conn.id, backendPid: 456 }) + return new Promise((resolve) => { + release = () => + resolve({ columns: [], rows: [], rowCount: 0, truncated: false, durationMs: 1 }) + }) + } + ) + const running = manager.query('c1', 'SELECT 1', { + rowLimit: 100, + timeoutMs: 1000, + allowWrite: false + }) + await manager.cancelQuery('c1') + expect(pgDriver.cancel).toHaveBeenCalledWith( + expect.objectContaining({ id: 'c1' }), + { connectionId: 'c1', backendPid: 456 } + ) + release?.() + await running + }) + + it('cancelQuery is a no-op when nothing is running', async () => { + await manager.connect(cfg()) + await manager.cancelQuery('c1') + expect(pgDriver.cancel).not.toHaveBeenCalled() + }) + + it('rejects multi-statement input on a read-only connection (H1)', async () => { + await manager.connect(cfg()) + await expect( + manager.query('c1', 'SET TRANSACTION READ WRITE; DELETE FROM t', { + rowLimit: 100, + timeoutMs: 1000, + allowWrite: false + }) + ).rejects.toThrow('db_read_only_multi_statement') + expect(pgDriver.query).not.toHaveBeenCalled() + }) + + it('allows multi-statement on a writable connection', async () => { + await manager.connect(cfg()) + await manager.query('c1', 'INSERT INTO t VALUES (1); INSERT INTO t VALUES (2)', { + rowLimit: 100, + timeoutMs: 1000, + allowWrite: true + }) + expect(pgDriver.query).toHaveBeenCalledTimes(1) + }) + + it('clears the in-flight handle when a connection is dropped mid-query', async () => { + await manager.connect(cfg()) + let release: (() => void) | undefined + pgDriver.query.mockImplementationOnce( + (conn: LiveConnection, _sql: string, _opts: unknown, onStart: (h: QueryHandle) => void) => { + onStart({ connectionId: conn.id, backendPid: 9 }) + return new Promise((resolve) => { + release = () => + resolve({ columns: [], rows: [], rowCount: 0, truncated: false, durationMs: 1 }) + }) + } + ) + const opts = { rowLimit: 100, timeoutMs: 1000, allowWrite: false } + const running = manager.query('c1', 'SELECT 1', opts) + // Connection drops while the query is still in-flight. + state.capturedOnError?.(Object.assign(new Error('reset'), { code: 'ECONNRESET' })) + expect(manager.getStatus('c1').status).toBe('lost') + // Reconnect and run again: the stale in-flight handle must not block it + // with db_query_in_progress or misdirect a later cancel. + await manager.connect(cfg()) + await expect(manager.query('c1', 'SELECT 2', opts)).resolves.toBeDefined() + release?.() + await running + }) + + it('rejects a second concurrent query on the same connection (L3)', async () => { + await manager.connect(cfg()) + let release: (() => void) | undefined + pgDriver.query.mockImplementationOnce( + (conn: LiveConnection, _sql: string, _opts: unknown, onStart: (h: QueryHandle) => void) => { + onStart({ connectionId: conn.id, backendPid: 1 }) + return new Promise((resolve) => { + release = () => + resolve({ columns: [], rows: [], rowCount: 0, truncated: false, durationMs: 1 }) + }) + } + ) + const opts = { rowLimit: 100, timeoutMs: 1000, allowWrite: false } + const first = manager.query('c1', 'SELECT 1', opts) + await expect(manager.query('c1', 'SELECT 2', opts)).rejects.toThrow('db_query_in_progress') + release?.() + await first + // Once the first completes the connection is free again. + await manager.query('c1', 'SELECT 3', opts) + }) + }) + + describe('execute + executeBatch', () => { + const stmt = { sql: 'SELECT * FROM t LIMIT 100 OFFSET 0', params: [] } + const opts = { rowLimit: 1000, timeoutMs: 1000, allowWrite: true } + + it('execute delegates a parameterized statement to the engine driver', async () => { + await manager.connect(cfg()) + await manager.execute('c1', stmt, { ...opts, allowWrite: false }) + expect(pgDriver.execute).toHaveBeenCalledWith( + expect.objectContaining({ id: 'c1' }), + stmt, + expect.objectContaining({ allowWrite: false }), + expect.any(Function) + ) + }) + + it('execute throws db_not_connected when no live connection is held', async () => { + await expect(manager.execute('missing', stmt, opts)).rejects.toThrow('db_not_connected') + }) + + it('executeBatch rejects on a read-only connection before touching the driver', async () => { + await manager.connect(cfg()) + await expect( + manager.executeBatch('c1', [stmt], { ...opts, allowWrite: false }) + ).rejects.toThrow('db_read_only_write_blocked') + expect(pgDriver.executeBatch).not.toHaveBeenCalled() + }) + + it('executeBatch applies statements and returns a row count per statement', async () => { + await manager.connect(cfg()) + const counts = await manager.executeBatch('c1', [stmt, stmt], opts) + expect(pgDriver.executeBatch).toHaveBeenCalledTimes(1) + expect(counts).toEqual([1, 1]) + }) + + it('execute shares the one-op-per-connection guard', async () => { + await manager.connect(cfg()) + let release: (() => void) | undefined + pgDriver.execute.mockImplementationOnce( + (conn: LiveConnection, _s: unknown, _o: unknown, onStart: (h: QueryHandle) => void) => { + onStart({ connectionId: conn.id, backendPid: 5 }) + return new Promise((resolve) => { + release = () => + resolve({ columns: [], rows: [], rowCount: 0, truncated: false, durationMs: 1 }) + }) + } + ) + const first = manager.execute('c1', stmt, opts) + await expect(manager.execute('c1', stmt, opts)).rejects.toThrow('db_query_in_progress') + release?.() + await first + }) + }) +}) diff --git a/src/main/database/db-connection-manager.ts b/src/main/database/db-connection-manager.ts new file mode 100644 index 00000000000..dd8e730dc54 --- /dev/null +++ b/src/main/database/db-connection-manager.ts @@ -0,0 +1,250 @@ +// Owns live database connections with full lifecycle resilience: a connect-race +// guard, a driver 'error' listener that degrades a dropped connection to `lost` +// (never crashing the main process), and quit-time disposal. Mirrors +// ssh-connection-manager's pool/guard shape; drivers lazy-import pg/mysql2. + +import type { + DbColumn, + DbConnectionRuntimeState, + DbEngine, + DbSchemaTree, + DbStatement, + DbTableList, + DbTableRef, + QueryHandle, + QueryOptions, + QueryResult +} from '../../shared/database-types' +import { + DB_MAX_SCHEMAS, + DB_MAX_TABLES_PER_SCHEMA, + normalizeDbError, + type DbDriver, + type LiveConnection, + type ResolvedDbConfig +} from './db-driver' +import { isMultiStatement } from '../../shared/sql-statement-classifier' +import { postgresDriver } from './postgres-driver' +import { mysqlDriver } from './mysql-driver' + +function getDriver(engine: DbEngine): DbDriver { + return engine === 'postgres' ? postgresDriver : mysqlDriver +} + +type StatusListener = (state: DbConnectionRuntimeState) => void + +export class DbConnectionManager { + private connections = new Map() + private statuses = new Map() + // Backend PID of the query currently running on each connection, so a cancel + // request can target the right server-side query. + private inFlight = new Map() + // Why: two concurrent connect(sameId) calls would both create a pool and orphan + // the first — this guard rejects the second while one is in progress (SSH F12). + private connectingTargets = new Set() + private statusListener: StatusListener = () => {} + + setStatusListener(listener: StatusListener): void { + this.statusListener = listener + } + + private setStatus( + id: string, + status: DbConnectionRuntimeState['status'], + error?: DbConnectionRuntimeState['error'] + ): void { + const state: DbConnectionRuntimeState = error ? { id, status, error } : { id, status } + this.statuses.set(id, state) + this.statusListener(state) + } + + getStatus(id: string): DbConnectionRuntimeState { + return this.statuses.get(id) ?? { id, status: 'idle' } + } + + getAllStatuses(): DbConnectionRuntimeState[] { + return Array.from(this.statuses.values()) + } + + getConnection(id: string): LiveConnection | undefined { + return this.connections.get(id) + } + + // One-shot ping; holds no state and never touches the live-connection map. + async test(cfg: ResolvedDbConfig): Promise { + await getDriver(cfg.engine).testConnection(cfg) + } + + async connect(cfg: ResolvedDbConfig): Promise { + if (this.connections.has(cfg.id)) { + return this.getStatus(cfg.id) + } + if (this.connectingTargets.has(cfg.id)) { + throw new Error('db_connect_in_progress') + } + this.connectingTargets.add(cfg.id) + this.setStatus(cfg.id, 'connecting') + // Bind the driver's 'error' listener to this specific connection: a late + // error from a pool that was disconnected/reconnected must not tear down the + // live connection that replaced it under the same id. + let liveConn: LiveConnection | null = null + try { + const conn = await getDriver(cfg.engine).connect(cfg, (err) => { + if (liveConn && this.connections.get(cfg.id) === liveConn) { + this.handleConnectionError(cfg.id, err) + } + }) + liveConn = conn + this.connections.set(cfg.id, conn) + this.setStatus(cfg.id, 'connected') + return this.getStatus(cfg.id) + } catch (err) { + this.setStatus(cfg.id, 'error', normalizeDbError(err)) + throw err + } finally { + this.connectingTargets.delete(cfg.id) + } + } + + // Red-team F4: driver 'error' after a live connection drops. Mark lost, drop + // it from the map, and close it in the background — never re-throw (that would + // crash the process and kill every PTY/SSH/terminal). + private handleConnectionError(id: string, err: unknown): void { + const conn = this.connections.get(id) + this.connections.delete(id) + // Drop any in-flight handle: a reconnect under this id must not inherit the + // dead query's backend PID (cancel would target the wrong backend) or trip + // the concurrency guard against a query that can no longer settle. + this.inFlight.delete(id) + this.setStatus(id, 'lost', normalizeDbError(err)) + if (conn) { + void getDriver(conn.engine) + .close(conn) + .catch(() => {}) + } + } + + // Why: introspection needs a held connection. Throw a fixed code (mapped to a + // safe message at the IPC boundary) rather than dereferencing undefined. + private requireLive(id: string): LiveConnection { + const conn = this.connections.get(id) + if (!conn) { + throw new Error('db_not_connected') + } + return conn + } + + // async so the requireLive guard surfaces as a rejection, not a sync throw. + async introspectSchemas(id: string): Promise { + const conn = this.requireLive(id) + return getDriver(conn.engine).introspectSchemas(conn, DB_MAX_SCHEMAS) + } + + async introspectTables(id: string, schema: string): Promise { + const conn = this.requireLive(id) + return getDriver(conn.engine).introspectTables(conn, schema, DB_MAX_TABLES_PER_SCHEMA) + } + + async introspectColumns(id: string, ref: DbTableRef): Promise { + const conn = this.requireLive(id) + return getDriver(conn.engine).introspectColumns(conn, ref) + } + + async query(id: string, sql: string, opts: QueryOptions): Promise { + const conn = this.requireLive(id) + // L3: one query per connection at a time — a second concurrent query would + // overwrite the in-flight handle and make cancel target the wrong query. + if (this.inFlight.has(id)) { + throw new Error('db_query_in_progress') + } + // H1 (red-team): a read-only connection must reject multi-statement input. + // A Postgres simple query runs every statement, so a multi-statement string + // could flip `SET TRANSACTION READ WRITE` before any query and defeat the + // read-only transaction. The DB read-only txn covers single-statement writes + // and writing CTEs; this closes the multi-statement gap. + if (!opts.allowWrite && isMultiStatement(sql)) { + throw new Error('db_read_only_multi_statement') + } + // Reserve synchronously so the concurrency guard holds before the driver's + // async backend-PID capture; onStart replaces this with the real handle. + this.inFlight.set(id, { connectionId: id, backendPid: null }) + try { + return await getDriver(conn.engine).query(conn, sql, opts, (handle) => { + this.inFlight.set(id, handle) + }) + } finally { + this.inFlight.delete(id) + } + } + + // Parameterized single statement (Data-tab select/count or wrapped free-form + // re-query). Shares the one-query-per-connection guard so cancel targets the + // right backend and a concurrent op can't overwrite the in-flight handle. + async execute(id: string, statement: DbStatement, opts: QueryOptions): Promise { + const conn = this.requireLive(id) + if (this.inFlight.has(id)) { + throw new Error('db_query_in_progress') + } + this.inFlight.set(id, { connectionId: id, backendPid: null }) + try { + return await getDriver(conn.engine).execute(conn, statement, opts, (handle) => { + this.inFlight.set(id, handle) + }) + } finally { + this.inFlight.delete(id) + } + } + + // Atomic staged-edit apply. Writes only: a read-only connection is rejected + // here (defense in depth — the UI already disables editing on read-only). + async executeBatch(id: string, statements: DbStatement[], opts: QueryOptions): Promise { + const conn = this.requireLive(id) + if (!opts.allowWrite) { + throw new Error('db_read_only_write_blocked') + } + if (this.inFlight.has(id)) { + throw new Error('db_query_in_progress') + } + this.inFlight.set(id, { connectionId: id, backendPid: null }) + try { + return await getDriver(conn.engine).executeBatch(conn, statements, opts, (handle) => { + this.inFlight.set(id, handle) + }) + } finally { + this.inFlight.delete(id) + } + } + + // Cancel the in-flight query on a connection via the driver's side connection. + // No-op if nothing is running or the backend PID was never captured. + async cancelQuery(id: string): Promise { + const handle = this.inFlight.get(id) + const conn = this.connections.get(id) + if (!handle || !conn) { + return + } + await getDriver(conn.engine).cancel(conn, handle) + } + + async disconnect(id: string): Promise { + const conn = this.connections.get(id) + this.connections.delete(id) + this.inFlight.delete(id) + this.setStatus(id, 'idle') + if (conn) { + await getDriver(conn.engine).close(conn) + } + } + + // Quit-time disposal (red-team F12): SSH's own manager is never disposed on + // quit, so this must be wired explicitly into index.ts will-quit. + async disconnectAll(): Promise { + const live = Array.from(this.connections.values()) + this.connections.clear() + this.inFlight.clear() + await Promise.allSettled(live.map((conn) => getDriver(conn.engine).close(conn))) + } +} + +// Single main-process instance shared by the IPC handlers and quit hook. +export const dbConnectionManager = new DbConnectionManager() diff --git a/src/main/database/db-credential-store.test.ts b/src/main/database/db-credential-store.test.ts new file mode 100644 index 00000000000..79de8165ad6 --- /dev/null +++ b/src/main/database/db-credential-store.test.ts @@ -0,0 +1,311 @@ +import { beforeEach, afterEach, describe, it, expect, vi } from 'vitest' +import { + getDbEncryptionStatus, + encryptDbSecret, + decryptDbSecret, + ensureDbSecretAtRest, + isDbSecretAtRest +} from './db-credential-store' + +// Mock electron safeStorage and control platform +const { isEncryptionAvailableMock, encryptStringMock, decryptStringMock, getSelectedStorageBackendMock } = + vi.hoisted(() => ({ + isEncryptionAvailableMock: vi.fn(() => true), + encryptStringMock: vi.fn((plaintext: string) => + Buffer.from(`mock-encrypted:${plaintext}`, 'utf-8') + ), + decryptStringMock: vi.fn((ciphertext: Buffer) => { + const decoded = ciphertext.toString('utf-8') + if (!decoded.startsWith('mock-encrypted:')) { + throw new Error('decryption_failed') + } + return decoded.slice('mock-encrypted:'.length) + }), + getSelectedStorageBackendMock: vi.fn(() => 'gnome_libsecret') + })) + +vi.mock('electron', () => ({ + safeStorage: { + isEncryptionAvailable: isEncryptionAvailableMock, + encryptString: encryptStringMock, + decryptString: decryptStringMock, + getSelectedStorageBackend: getSelectedStorageBackendMock + } +})) + +describe('db-credential-store', () => { + beforeEach(() => { + vi.clearAllMocks() + isEncryptionAvailableMock.mockReturnValue(true) + encryptStringMock.mockImplementation((plaintext: string) => + Buffer.from(`mock-encrypted:${plaintext}`, 'utf-8') + ) + decryptStringMock.mockImplementation((ciphertext: Buffer) => { + const decoded = ciphertext.toString('utf-8') + if (!decoded.startsWith('mock-encrypted:')) { + throw new Error('decryption_failed') + } + return decoded.slice('mock-encrypted:'.length) + }) + getSelectedStorageBackendMock.mockReturnValue('gnome_libsecret') + }) + + afterEach(() => { + const originalPlatform = process.platform + Object.defineProperty(process, 'platform', { + configurable: true, + value: originalPlatform + }) + }) + + describe('getDbEncryptionStatus', () => { + it('returns strong backend for gnome_libsecret on linux', () => { + Object.defineProperty(process, 'platform', { + configurable: true, + value: 'linux' + }) + getSelectedStorageBackendMock.mockReturnValue('gnome_libsecret') + + const status = getDbEncryptionStatus() + + expect(status.backend).toBe('gnome_libsecret') + expect(status.isStrong).toBe(true) + }) + + it('returns strong backend for kwallet on linux', () => { + Object.defineProperty(process, 'platform', { + configurable: true, + value: 'linux' + }) + getSelectedStorageBackendMock.mockReturnValue('kwallet5') + + const status = getDbEncryptionStatus() + + expect(status.backend).toBe('kwallet5') + expect(status.isStrong).toBe(true) + }) + + it('returns strong backend for keychain on darwin', () => { + Object.defineProperty(process, 'platform', { + configurable: true, + value: 'darwin' + }) + + const status = getDbEncryptionStatus() + + expect(status.backend).toBe('keychain') + expect(status.isStrong).toBe(true) + }) + + it('returns strong backend for dpapi on win32', () => { + Object.defineProperty(process, 'platform', { + configurable: true, + value: 'win32' + }) + + const status = getDbEncryptionStatus() + + expect(status.backend).toBe('dpapi') + expect(status.isStrong).toBe(true) + }) + + it('returns weak backend for basic_text', () => { + Object.defineProperty(process, 'platform', { + configurable: true, + value: 'linux' + }) + getSelectedStorageBackendMock.mockReturnValue('basic_text') + + const status = getDbEncryptionStatus() + + expect(status.backend).toBe('basic_text') + expect(status.isStrong).toBe(false) + }) + + it('returns unavailable backend when encryption is not available', () => { + isEncryptionAvailableMock.mockReturnValue(false) + + const status = getDbEncryptionStatus() + + expect(status.backend).toBe('unavailable') + expect(status.isStrong).toBe(false) + }) + }) + + describe('encryptDbSecret', () => { + it('returns empty string unchanged', () => { + const result = encryptDbSecret('') + + expect(result).toBe('') + }) + + it('encrypts plaintext with ENC prefix when encryption is available', () => { + const result = encryptDbSecret('mysecret') + + expect(result).toMatch(/^db\.safeStorage\.v1:/) + expect(result).not.toContain('mysecret') + expect(encryptStringMock).toHaveBeenCalledWith('mysecret') + }) + + it('returns RAW-prefixed plaintext when encryption is unavailable', () => { + isEncryptionAvailableMock.mockReturnValue(false) + + const result = encryptDbSecret('mysecret') + + expect(result).toBe('db.plaintext.v1:mysecret') + }) + + it('throws when encryption fails on strong backend', () => { + encryptStringMock.mockImplementation(() => { + throw new Error('encryption_failed') + }) + + expect(() => encryptDbSecret('mysecret')).toThrow('db_secret_encrypt_failed') + }) + }) + + describe('decryptDbSecret', () => { + it('returns empty string unchanged', () => { + const result = decryptDbSecret('') + + expect(result).toBe('') + }) + + it('decrypts ENC-prefixed ciphertext', () => { + const encrypted = encryptDbSecret('mysecret') + + const result = decryptDbSecret(encrypted) + + expect(result).toBe('mysecret') + }) + + it('decrypts RAW-prefixed plaintext', () => { + const result = decryptDbSecret('db.plaintext.v1:mysecret') + + expect(result).toBe('mysecret') + }) + + it('throws on untagged value (fail-closed)', () => { + expect(() => decryptDbSecret('untagged-secret')).toThrow('db_secret_unknown_format') + }) + + it('throws when safeStorage.decryptString fails (fail-closed)', () => { + decryptStringMock.mockImplementation(() => { + throw new Error('decryption_failed') + }) + const encrypted = encryptDbSecret('mysecret') + + expect(() => decryptDbSecret(encrypted)).toThrow() + }) + + it('throws on corrupt base64 in ENC-prefixed value', () => { + decryptStringMock.mockImplementation(() => { + throw new Error('decryption_failed') + }) + + expect(() => decryptDbSecret('db.safeStorage.v1:not-valid-base64!!!')).toThrow() + }) + }) + + describe('isDbSecretAtRest', () => { + it('returns false for undefined', () => { + expect(isDbSecretAtRest(undefined)).toBe(false) + }) + + it('returns false for empty string', () => { + expect(isDbSecretAtRest('')).toBe(false) + }) + + it('returns true for ENC-prefixed value', () => { + expect(isDbSecretAtRest('db.safeStorage.v1:xyz')).toBe(true) + }) + + it('returns true for RAW-prefixed value', () => { + expect(isDbSecretAtRest('db.plaintext.v1:xyz')).toBe(true) + }) + + it('returns false for untagged plaintext', () => { + expect(isDbSecretAtRest('plaintext-secret')).toBe(false) + }) + }) + + describe('ensureDbSecretAtRest', () => { + it('returns undefined unchanged', () => { + expect(ensureDbSecretAtRest(undefined)).toBeUndefined() + }) + + it('returns empty string unchanged', () => { + expect(ensureDbSecretAtRest('')).toBe('') + }) + + it('passes through already-tagged ENC value (idempotent)', () => { + const tagged = 'db.safeStorage.v1:base64-ciphertext' + + const result = ensureDbSecretAtRest(tagged) + + expect(result).toBe(tagged) + expect(encryptStringMock).not.toHaveBeenCalled() + }) + + it('passes through already-tagged RAW value (idempotent)', () => { + const tagged = 'db.plaintext.v1:mysecret' + + const result = ensureDbSecretAtRest(tagged) + + expect(result).toBe(tagged) + expect(encryptStringMock).not.toHaveBeenCalled() + }) + + it('encrypts untagged plaintext', () => { + const result = ensureDbSecretAtRest('mysecret') + + expect(result).toMatch(/^db\.safeStorage\.v1:/) + expect(result).not.toContain('mysecret') + expect(encryptStringMock).toHaveBeenCalledWith('mysecret') + }) + + it('encrypts untagged plaintext to RAW when encryption unavailable', () => { + isEncryptionAvailableMock.mockReturnValue(false) + + const result = ensureDbSecretAtRest('mysecret') + + expect(result).toBe('db.plaintext.v1:mysecret') + }) + }) + + describe('round-trip encryption/decryption', () => { + it('encrypts and decrypts back to original with strong backend', () => { + const plaintext = 'mypassword123' + + const encrypted = encryptDbSecret(plaintext) + const decrypted = decryptDbSecret(encrypted) + + expect(decrypted).toBe(plaintext) + }) + + it('encrypts and decrypts back to original with weak backend', () => { + isEncryptionAvailableMock.mockReturnValue(false) + const plaintext = 'mypassword123' + + const encrypted = encryptDbSecret(plaintext) + const decrypted = decryptDbSecret(encrypted) + + expect(decrypted).toBe(plaintext) + }) + + it('ensures idempotency: double-ensuring returns same value', () => { + const plaintext = 'mypassword123' + + const encrypted1 = ensureDbSecretAtRest(plaintext) + const encrypted2 = ensureDbSecretAtRest(encrypted1) + + expect(encrypted1).toBe(encrypted2) + + const decrypted1 = decryptDbSecret(encrypted1 ?? '') + const decrypted2 = decryptDbSecret(encrypted2 ?? '') + + expect(decrypted1).toBe(plaintext) + expect(decrypted2).toBe(plaintext) + }) + }) +}) diff --git a/src/main/database/db-credential-store.ts b/src/main/database/db-credential-store.ts new file mode 100644 index 00000000000..ddf7f0c4b4a --- /dev/null +++ b/src/main/database/db-credential-store.ts @@ -0,0 +1,95 @@ +import { safeStorage } from 'electron' +import type { DbEncryptionStatus } from '../../shared/database-types' + +// Tagged at-rest representation so decrypt is unambiguous and fail-closed: +// ENC_PREFIX → real safeStorage ciphertext (base64) +// RAW_PREFIX → warn-and-store plaintext (no OS crypto backend available) +const ENC_PREFIX = 'db.safeStorage.v1:' +const RAW_PREFIX = 'db.plaintext.v1:' + +// Why: a connection's password is recoverable from disk unless the OS exposes a +// real keystore. basic_text uses a hardcoded key (Electron still reports +// isEncryptionAvailable() === true for it), so it is NOT strong. macOS Keychain +// and Windows DPAPI are strong whenever encryption is available; +// getSelectedStorageBackend() is Linux-only. +const KNOWN_STRONG_BACKENDS = new Set([ + 'gnome_libsecret', + 'kwallet', + 'kwallet5', + 'kwallet6', + 'keychain', + 'dpapi' +]) + +function selectedBackend(): string { + if (!safeStorage.isEncryptionAvailable()) { + return 'unavailable' + } + if (process.platform === 'darwin') { + return 'keychain' + } + if (process.platform === 'win32') { + return 'dpapi' + } + try { + return safeStorage.getSelectedStorageBackend() + } catch { + return 'unknown' + } +} + +export function getDbEncryptionStatus(): DbEncryptionStatus { + const backend = selectedBackend() + return { backend, isStrong: KNOWN_STRONG_BACKENDS.has(backend) } +} + +export function isDbSecretAtRest(value: string | undefined): boolean { + return !!value && (value.startsWith(ENC_PREFIX) || value.startsWith(RAW_PREFIX)) +} + +// Encrypt a password into its tagged at-rest form. Uses the OS keystore whenever +// available (strong OR basic_text — both decrypt via safeStorage); only falls +// back to tagged plaintext when no backend exists at all (warn-and-store). A +// strong backend that throws mid-encrypt FAILS CLOSED rather than silently +// downgrading to a recoverable secret. +export function encryptDbSecret(plaintext: string): string { + if (!plaintext) { + return plaintext + } + if (safeStorage.isEncryptionAvailable()) { + try { + return ENC_PREFIX + safeStorage.encryptString(plaintext).toString('base64') + } catch { + throw new Error('db_secret_encrypt_failed') + } + } + return RAW_PREFIX + plaintext +} + +// Idempotent: encrypt only if the value is not already in tagged at-rest form. +// Used by the persistence write paths to guarantee no plaintext password is ever +// written to disk, even if one leaked into in-memory state. +export function ensureDbSecretAtRest(value: string | undefined): string | undefined { + if (!value || isDbSecretAtRest(value)) { + return value + } + return encryptDbSecret(value) +} + +// Strict, fail-closed decrypt. Unlike the cookie-grade decrypt (which returns the +// ciphertext on failure), a corrupt or keychain-changed value THROWS so callers +// surface "password could not be decrypted on this machine" instead of handing a +// bogus credential to a driver. +export function decryptDbSecret(stored: string): string { + if (!stored) { + return stored + } + if (stored.startsWith(RAW_PREFIX)) { + return stored.slice(RAW_PREFIX.length) + } + if (stored.startsWith(ENC_PREFIX)) { + return safeStorage.decryptString(Buffer.from(stored.slice(ENC_PREFIX.length), 'base64')) + } + // This feature never wrote an untagged secret; treat anything else as corrupt. + throw new Error('db_secret_unknown_format') +} diff --git a/src/main/database/db-driver.test.ts b/src/main/database/db-driver.test.ts new file mode 100644 index 00000000000..578f88f0eb1 --- /dev/null +++ b/src/main/database/db-driver.test.ts @@ -0,0 +1,167 @@ +import { afterEach, describe, expect, it, vi } from 'vitest' +import type { DbConnection } from '../../shared/database-types' +import { + applyCap, + DB_CONNECT_TIMEOUT_MS, + DbTimeoutError, + isLocalHost, + normalizeDbError, + raceWithTimeout, + resolveDbConfig, + resolveSslMode +} from './db-driver' + +function makeConnection(overrides: Partial = {}): DbConnection { + return { + id: 'c1', + name: 'db', + engine: 'postgres', + host: 'db.example.com', + port: 5432, + database: 'app', + user: 'admin', + readOnly: false, + createdAt: 0, + updatedAt: 0, + ...overrides + } +} + +describe('isLocalHost', () => { + it.each(['localhost', '127.0.0.1', '::1', '[::1]', '0.0.0.0', 'LOCALHOST'])( + 'treats %s as local', + (host) => { + expect(isLocalHost(host)).toBe(true) + } + ) + + it.each(['db.example.com', '10.0.0.5', 'postgres.internal'])('treats %s as remote', (host) => { + expect(isLocalHost(host)).toBe(false) + }) +}) + +describe('resolveSslMode (smart-by-host)', () => { + it('defaults localhost to disable when ssl is unset', () => { + expect(resolveSslMode(undefined, 'localhost')).toBe('disable') + }) + + it('defaults a remote host to verify-full when ssl is unset', () => { + expect(resolveSslMode(undefined, 'db.example.com')).toBe('verify-full') + }) + + it('honors an explicit mode over the host heuristic', () => { + expect(resolveSslMode('insecure-no-verify', 'db.example.com')).toBe('insecure-no-verify') + expect(resolveSslMode('verify-full', 'localhost')).toBe('verify-full') + }) +}) + +describe('resolveDbConfig', () => { + it('collapses smart-by-host SSL and carries the decrypted password at point-of-use', () => { + const cfg = resolveDbConfig(makeConnection(), 's3cr3t') + expect(cfg.ssl).toBe('verify-full') + expect(cfg.password).toBe('s3cr3t') + expect(cfg.readOnly).toBe(false) + }) + + it('defaults a remote host with no explicit ssl to verify-full', () => { + expect(resolveDbConfig(makeConnection({ host: 'remote.db' }), undefined).ssl).toBe( + 'verify-full' + ) + }) +}) + +describe('normalizeDbError (no credential leak)', () => { + it('maps known driver codes to safe messages', () => { + expect(normalizeDbError({ code: '28P01' }).code).toBe('auth_failed') + expect(normalizeDbError({ code: 'ER_ACCESS_DENIED_ERROR' }).code).toBe('auth_failed') + expect(normalizeDbError({ code: '3D000' }).code).toBe('database_not_found') + expect(normalizeDbError({ code: 'ECONNREFUSED' }).code).toBe('connection_refused') + expect(normalizeDbError({ code: 'ENOTFOUND' }).code).toBe('host_unreachable') + expect(normalizeDbError({ code: 'ETIMEDOUT' }).code).toBe('timeout') + expect(normalizeDbError({ code: 'DEPTH_ZERO_SELF_SIGNED_CERT' }).code).toBe('tls_error') + }) + + it('maps a DbTimeoutError to the timeout code', () => { + expect(normalizeDbError(new DbTimeoutError()).code).toBe('timeout') + }) + + it('maps fail-closed credential-store errors to decrypt_failed', () => { + expect(normalizeDbError(new Error('db_secret_unknown_format')).code).toBe('decrypt_failed') + }) + + it('maps internal guard errors to their safe codes', () => { + expect(normalizeDbError(new Error('db_read_only_multi_statement')).code).toBe('read_only_blocked') + expect(normalizeDbError(new Error('db_query_in_progress')).code).toBe('busy') + expect(normalizeDbError(new Error('db_not_connected')).code).toBe('not_connected') + }) + + it('falls back to unknown for unrecognized errors', () => { + expect(normalizeDbError({ code: 'WAT' }).code).toBe('unknown') + expect(normalizeDbError('a string').code).toBe('unknown') + }) + + it('never forwards the raw message (DSN/password) in the safe payload', () => { + const raw = Object.assign( + new Error('password authentication failed for "admin" postgres://admin:s3cr3t@db:5432/app'), + { code: '28P01' } + ) + const safe = normalizeDbError(raw) + const serialized = JSON.stringify(safe) + expect(serialized).not.toContain('s3cr3t') + expect(serialized).not.toContain('postgres://') + expect(serialized).not.toContain('admin') + }) +}) + +describe('applyCap (introspection overflow)', () => { + it('marks truncated and slices when rows exceed the cap', () => { + // Level queried with cap+1, so 3 rows for a cap of 2 means overflow. + expect(applyCap([1, 2, 3], 2)).toEqual({ kept: [1, 2], truncated: true }) + }) + + it('keeps all rows and is not truncated at or below the cap', () => { + expect(applyCap([1, 2], 2)).toEqual({ kept: [1, 2], truncated: false }) + expect(applyCap([1], 2)).toEqual({ kept: [1], truncated: false }) + }) +}) + +describe('raceWithTimeout', () => { + afterEach(() => { + vi.useRealTimers() + }) + + it('resolves with the wrapped value when it settles first', async () => { + await expect(raceWithTimeout(Promise.resolve('ok'), DB_CONNECT_TIMEOUT_MS)).resolves.toBe('ok') + }) + + it('propagates a wrapped rejection', async () => { + await expect( + raceWithTimeout(Promise.reject(new Error('boom')), DB_CONNECT_TIMEOUT_MS) + ).rejects.toThrow('boom') + }) + + it('rejects with DbTimeoutError and runs onTimeout when the deadline passes', async () => { + vi.useFakeTimers() + const onTimeout = vi.fn() + // A promise that never settles — only the timeout can win. + const pending = raceWithTimeout(new Promise(() => {}), 1000, onTimeout) + const assertion = expect(pending).rejects.toBeInstanceOf(DbTimeoutError) + await vi.advanceTimersByTimeAsync(1000) + await assertion + expect(onTimeout).toHaveBeenCalledTimes(1) + }) + + it('still rejects deterministically when onTimeout throws', async () => { + vi.useFakeTimers() + const onTimeout = vi.fn(() => { + throw new Error('cleanup failed') + }) + // A throwing cleanup callback must not escape the timer or leave the caller + // hanging — the timeout rejection must still surface. + const pending = raceWithTimeout(new Promise(() => {}), 1000, onTimeout) + const assertion = expect(pending).rejects.toBeInstanceOf(DbTimeoutError) + await vi.advanceTimersByTimeAsync(1000) + await assertion + expect(onTimeout).toHaveBeenCalledTimes(1) + }) +}) diff --git a/src/main/database/db-driver.ts b/src/main/database/db-driver.ts new file mode 100644 index 00000000000..29a18b08cd8 --- /dev/null +++ b/src/main/database/db-driver.ts @@ -0,0 +1,281 @@ +// Driver abstraction for the in-app database client. Postgres/MySQL drivers +// implement this interface with lazy (`await import`) module loading so pg/mysql2 +// stay out of the main-process startup require-graph until the first connect. + +import type { + DbColumn, + DbConnection, + DbEngine, + DbSafeError, + DbSchemaTree, + DbSslMode, + DbStatement, + DbTableList, + DbTableRef, + QueryHandle, + QueryOptions, + QueryResult +} from '../../shared/database-types' + +// Mirrors SSH's readyTimeout (ssh-connection-utils CONNECT_TIMEOUT_MS): a dead +// host must never hang the IPC forever. pg's default is wait-forever. +export const DB_CONNECT_TIMEOUT_MS = 30_000 + +// Introspection caps (red-team F9): a level with more objects than its cap is +// returned truncated rather than buffering the whole catalog into main + IPC. +export const DB_MAX_SCHEMAS = 500 +export const DB_MAX_TABLES_PER_SCHEMA = 2_000 + +// Query caps (red-team F9): a server-side cursor fetches at most DB_MAX_ROWS + 1 +// so a huge result never buffers whole; the statement timeout bounds runtime. +export const DB_MAX_ROWS = 1_000 +export const DB_STATEMENT_TIMEOUT_MS = 30_000 + +// SSL after smart-by-host resolution — always a concrete mode, never unset. +export type ResolvedSslMode = 'disable' | 'verify-full' | 'insecure-no-verify' + +// Ready-to-dial config: password decrypted at point-of-use (never persisted in +// this shape) and SSL collapsed from smart-by-host to a concrete mode. +export type ResolvedDbConfig = { + id: string + engine: DbEngine + host: string + port: number + database: string + user: string + password?: string + ssl: ResolvedSslMode + readOnly: boolean +} + +// Opaque live handle held by the manager. `raw` is the driver-native pool; the +// driver attaches an 'error' listener (forwarding to the manager) before this is +// returned, so a dropped connection degrades to `lost` instead of crashing. +// `config` is retained so cancel can open a short-lived side connection to issue +// pg_cancel_backend / KILL QUERY without competing for a pooled slot. +export type LiveConnection = { + id: string + engine: DbEngine + raw: unknown + config: ResolvedDbConfig +} + +export type DbDriver = { + // Bounded by the connect timeout; throws on failure. Holds no state. + testConnection(cfg: ResolvedDbConfig): Promise + // Pool-backed; attaches the mandatory 'error' listener before returning. + connect(cfg: ResolvedDbConfig, onError: (err: unknown) => void): Promise + // Lazy, capped introspection. Each runs on a pooled connection so it never + // contends with a running query (red-team F11). Query/cancel come in Phase 5. + introspectSchemas(conn: LiveConnection, maxSchemas: number): Promise + introspectTables( + conn: LiveConnection, + schema: string, + maxTables: number + ): Promise + introspectColumns(conn: LiveConnection, ref: DbTableRef): Promise + // Runs SQL on a dedicated pooled connection: read-only DB transaction when + // !allowWrite, statement timeout, cursor bounded to rowLimit+1. `onStart` + // reports the backend PID so the query can be cancelled while running. + query( + conn: LiveConnection, + sql: string, + opts: QueryOptions, + onStart: (handle: QueryHandle) => void + ): Promise + // Runs a single parameterized statement (Data-tab select/count or a wrapped + // free-form re-query). Read-only DB transaction when !allowWrite; statement + // timeout; rows defensively capped to rowLimit. onStart reports the backend PID. + execute( + conn: LiveConnection, + statement: DbStatement, + opts: QueryOptions, + onStart: (handle: QueryHandle) => void + ): Promise + // Applies staged writes atomically in one transaction (BEGIN … COMMIT; ROLLBACK + // + DbBatchError on any failure). Requires opts.allowWrite. Returns the affected + // row count per statement, positional to `statements`. + executeBatch( + conn: LiveConnection, + statements: DbStatement[], + opts: QueryOptions, + onStart: (handle: QueryHandle) => void + ): Promise + // Server-side cancel via a short-lived side connection (pg_cancel_backend / + // KILL QUERY). No-op if the backend PID was never captured. + cancel(conn: LiveConnection, handle: QueryHandle): Promise + // Releases the pool. + close(conn: LiveConnection): Promise +} + +// Thrown by executeBatch when a statement in the transaction fails, carrying the +// 0-based index so the UI can point at the offending staged change. The original +// driver error is attached as `cause` for redaction at the IPC boundary. +export class DbBatchError extends Error { + constructor( + readonly failedIndex: number, + readonly cause: unknown + ) { + super('db_batch_failed') + this.name = 'DbBatchError' + } +} + +// Shared cap logic: query with `cap + 1`, then a level is truncated when the +// server returned more than the cap allowed. Returns the kept slice + the flag. +export function applyCap(rows: T[], cap: number): { kept: T[]; truncated: boolean } { + if (rows.length > cap) { + return { kept: rows.slice(0, cap), truncated: true } + } + return { kept: rows, truncated: false } +} + +// Why: a localhost server usually has no cert to verify, while any remote host +// should verify by default. An explicit user-chosen mode always wins. +export function isLocalHost(host: string): boolean { + const h = host.trim().toLowerCase().replace(/^\[|\]$/g, '') + return h === 'localhost' || h === '127.0.0.1' || h === '::1' || h === '0.0.0.0' +} + +export function resolveSslMode(ssl: DbSslMode | undefined, host: string): ResolvedSslMode { + if (ssl) { + return ssl + } + return isLocalHost(host) ? 'disable' : 'verify-full' +} + +// Collapse a persisted connection + its decrypted password into a dial-ready +// config. Smart-by-host SSL is resolved here so drivers see a concrete mode. +export function resolveDbConfig( + connection: DbConnection, + decryptedPassword: string | undefined +): ResolvedDbConfig { + return { + id: connection.id, + engine: connection.engine, + host: connection.host, + port: connection.port, + database: connection.database, + user: connection.user, + password: decryptedPassword, + ssl: resolveSslMode(connection.ssl, connection.host), + readOnly: connection.readOnly + } +} + +export class DbTimeoutError extends Error { + constructor() { + super('db_connect_timeout') + this.name = 'DbTimeoutError' + } +} + +// Belt-and-suspenders around the driver-native connect timeout: guarantees the +// IPC settles even if a driver ignores its own timeout option. `onTimeout` lets +// the caller tear down the half-open socket the losing promise still owns. +export function raceWithTimeout( + promise: Promise, + ms: number, + onTimeout?: () => void +): Promise { + return new Promise((resolve, reject) => { + const timer = setTimeout(() => { + // A throwing cleanup callback must not escape the timer (uncaught) and + // leave the caller hanging — always reject deterministically on timeout. + try { + onTimeout?.() + } catch { + // Best-effort teardown; the operation still fails closed below. + } + reject(new DbTimeoutError()) + }, ms) + promise.then( + (value) => { + clearTimeout(timer) + resolve(value) + }, + (err) => { + clearTimeout(timer) + reject(err) + } + ) + }) +} + +// ── Error normalization (red-team F6) ──────────────────────────────── +// +// Raw driver errors embed the DSN, host, user, and sometimes the password; +// there is no IPC redactor. Map known driver error codes to a fixed, credential- +// free { code, safeMessage } and fall back to a generic message otherwise — the +// raw message is NEVER forwarded. + +const SAFE_MESSAGES: Record = { + auth_failed: 'Authentication failed. Check the user name and password.', + database_not_found: 'The specified database does not exist.', + host_unreachable: 'Could not resolve or reach the database host.', + connection_refused: 'Connection refused by the host.', + timeout: 'Connection timed out.', + tls_error: 'TLS/SSL negotiation failed.', + decrypt_failed: 'The stored password could not be decrypted on this machine.', + not_connected: 'The connection is not open.', + read_only_blocked: 'Read-only connection: run one statement at a time.', + read_only_write: 'This connection is read-only; changes cannot be saved.', + busy: 'A query is already running on this connection.', + unknown: 'Could not connect to the database.' +} + +// Internal error strings raised by our own code (credential store, manager); +// surface them as clear codes rather than a bare "unknown". +const MESSAGE_CODE_MAP: Record = { + db_secret_unknown_format: 'decrypt_failed', + db_secret_encrypt_failed: 'decrypt_failed', + db_not_connected: 'not_connected', + db_read_only_multi_statement: 'read_only_blocked', + db_read_only_write_blocked: 'read_only_write', + db_query_in_progress: 'busy' +} + +// Driver error code → safe code. Covers pg (SQLSTATE + libpq errno) and mysql2 +// (ER_*/PROTOCOL_*) plus shared Node socket/TLS errnos. +const DRIVER_CODE_MAP: Record = { + // Postgres SQLSTATE + '28P01': 'auth_failed', + '28000': 'auth_failed', + '3D000': 'database_not_found', + // MySQL ER_* + ER_ACCESS_DENIED_ERROR: 'auth_failed', + ER_DBACCESS_DENIED_ERROR: 'auth_failed', + ER_BAD_DB_ERROR: 'database_not_found', + PROTOCOL_SEQUENCE_TIMEOUT: 'timeout', + PROTOCOL_CONNECTION_LOST: 'connection_refused', + HANDSHAKE_SSL_ERROR: 'tls_error', + // Shared Node socket/DNS/TLS errnos + ECONNREFUSED: 'connection_refused', + ECONNRESET: 'connection_refused', + ENOTFOUND: 'host_unreachable', + EAI_AGAIN: 'host_unreachable', + EHOSTUNREACH: 'host_unreachable', + ENETUNREACH: 'host_unreachable', + ETIMEDOUT: 'timeout', + DEPTH_ZERO_SELF_SIGNED_CERT: 'tls_error', + SELF_SIGNED_CERT_IN_CHAIN: 'tls_error', + UNABLE_TO_VERIFY_LEAF_SIGNATURE: 'tls_error', + ERR_TLS_CERT_ALTNAME_INVALID: 'tls_error', + CERT_HAS_EXPIRED: 'tls_error' +} + +export function normalizeDbError(err: unknown): DbSafeError { + if (err instanceof DbTimeoutError) { + return { code: 'timeout', safeMessage: SAFE_MESSAGES.timeout } + } + const rawCode = + typeof err === 'object' && err !== null && 'code' in err + ? String((err as { code: unknown }).code) + : undefined + const rawMessage = err instanceof Error ? err.message : undefined + const safeCode = + (rawCode && DRIVER_CODE_MAP[rawCode]) || + (rawMessage && MESSAGE_CODE_MAP[rawMessage]) || + 'unknown' + return { code: safeCode, safeMessage: SAFE_MESSAGES[safeCode] } +} diff --git a/src/main/database/mysql-driver.test.ts b/src/main/database/mysql-driver.test.ts new file mode 100644 index 00000000000..f3a7caf0cb2 --- /dev/null +++ b/src/main/database/mysql-driver.test.ts @@ -0,0 +1,167 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest' +import type { ResolvedDbConfig } from './db-driver' + +type ConnListener = (conn: unknown) => void + +// Fake mysql2/promise pool. The core (callback) pool is exposed via `.pool`, which +// is where the driver attaches the 'connection' listener that wires per-connection +// 'error' forwarding (red-team F4). +class FakePool { + static instances: FakePool[] = [] + static nextGetConnection: (() => Promise) | null = null + // Rows returned by pool-level introspection queries (mysql2 returns [rows]). + static queryRows: unknown[] = [] + config: Record + connectionListener: ConnListener | null = null + ended = false + pool = { + on: (event: string, cb: ConnListener): void => { + if (event === 'connection') { + this.connectionListener = cb + } + } + } + getConnectionImpl: () => Promise + + constructor(config: Record) { + this.config = config + this.getConnectionImpl = + FakePool.nextGetConnection ?? + (() => Promise.resolve({ query: vi.fn().mockResolvedValue([[]]), release: vi.fn() })) + FakePool.nextGetConnection = null + FakePool.instances.push(this) + } + + getConnection(): Promise { + return this.getConnectionImpl() + } + + query(): Promise<[unknown[], unknown[]]> { + return Promise.resolve([FakePool.queryRows, []]) + } + + async end(): Promise { + this.ended = true + } +} + +vi.mock('mysql2/promise', () => ({ + default: { createPool: (config: Record) => new FakePool(config) } +})) + +import { buildMysqlPoolConfig, buildMysqlSsl, mysqlDriver } from './mysql-driver' + +function cfg(overrides: Partial = {}): ResolvedDbConfig { + return { + id: 'c1', + engine: 'mysql', + host: 'db.example.com', + port: 3306, + database: 'app', + user: 'admin', + password: 'pw', + ssl: 'verify-full', + readOnly: false, + ...overrides + } +} + +describe('buildMysqlSsl', () => { + it('omits TLS for ssl=disable', () => { + expect(buildMysqlSsl('disable')).toBeUndefined() + }) + + it('verifies certs and hostname for ssl=verify-full', () => { + expect(buildMysqlSsl('verify-full')).toEqual({ rejectUnauthorized: true, verifyIdentity: true }) + }) + + it('does not verify certs or hostname for ssl=insecure-no-verify', () => { + expect(buildMysqlSsl('insecure-no-verify')).toEqual({ + rejectUnauthorized: false, + verifyIdentity: false + }) + }) +}) + +describe('buildMysqlPoolConfig', () => { + it('disables the LOCAL INFILE and multi-statement vectors', () => { + const config = buildMysqlPoolConfig(cfg()) + expect('infileStreamFactory' in config).toBe(true) + expect(config.infileStreamFactory).toBeUndefined() + expect(config.multipleStatements).toBe(false) + }) + + it('sets a bounded connect timeout and a small pool', () => { + const config = buildMysqlPoolConfig(cfg()) + expect(config.connectTimeout).toBeGreaterThan(0) + expect(config.connectionLimit).toBe(2) + expect(config.ssl).toEqual({ rejectUnauthorized: true, verifyIdentity: true }) + }) +}) + +describe('mysqlDriver', () => { + beforeEach(() => { + FakePool.instances = [] + FakePool.nextGetConnection = null + FakePool.queryRows = [] + }) + + it('connect wires per-connection error forwarding and returns a live connection', async () => { + const onError = vi.fn() + const conn = await mysqlDriver.connect(cfg(), onError) + const pool = FakePool.instances[0] + expect(pool.connectionListener).toBeTypeOf('function') + expect(conn.engine).toBe('mysql') + + // A pooled connection's 'error' must reach the manager, not crash the process. + const listeners: Record void> = {} + pool.connectionListener?.({ + on: (event: string, cb: (err: unknown) => void) => { + listeners[event] = cb + } + }) + listeners.error?.(new Error('socket dropped')) + expect(onError).toHaveBeenCalledTimes(1) + }) + + it('connect ends the pool and rethrows when validation fails', async () => { + FakePool.nextGetConnection = () => + Promise.resolve({ + query: vi.fn().mockRejectedValue( + Object.assign(new Error('denied'), { code: 'ER_ACCESS_DENIED_ERROR' }) + ), + release: vi.fn() + }) + await expect(mysqlDriver.connect(cfg(), vi.fn())).rejects.toThrow() + expect(FakePool.instances[0].ended).toBe(true) + }) + + it('testConnection ends the pool even on success', async () => { + await mysqlDriver.testConnection(cfg()) + expect(FakePool.instances[0].ended).toBe(true) + }) + + it('introspectTables maps kinds and flags truncation past the cap', async () => { + const conn = await mysqlDriver.connect(cfg(), vi.fn()) + FakePool.queryRows = [ + { name: 'a', type: 'BASE TABLE' }, + { name: 'b', type: 'VIEW' }, + { name: 'c', type: 'BASE TABLE' } + ] + const list = await mysqlDriver.introspectTables(conn, 'app', 2) + expect(list.tables).toEqual([ + { name: 'a', kind: 'table' }, + { name: 'b', kind: 'view' } + ]) + expect(list.truncated).toBe(true) + }) + + it('introspectColumns maps column_key=PRI to primary key', async () => { + const conn = await mysqlDriver.connect(cfg(), vi.fn()) + FakePool.queryRows = [ + { name: 'id', data_type: 'int', is_nullable: 'NO', column_key: 'PRI' } + ] + const columns = await mysqlDriver.introspectColumns(conn, { schema: 'app', table: 'users' }) + expect(columns).toEqual([{ name: 'id', dataType: 'int', nullable: false, isPrimaryKey: true }]) + }) +}) diff --git a/src/main/database/mysql-driver.ts b/src/main/database/mysql-driver.ts new file mode 100644 index 00000000000..89429aafe6a --- /dev/null +++ b/src/main/database/mysql-driver.ts @@ -0,0 +1,231 @@ +// MySQL driver: lazy-imports `mysql2/promise` on first use, dials over a small +// pool with a bounded connect timeout, disables the client-side LOCAL INFILE and +// multi-statement vectors, and forwards pooled-connection 'error' events so a +// dropped connection degrades to `lost` instead of crashing the main process. + +import { + applyCap, + DB_CONNECT_TIMEOUT_MS, + DB_STATEMENT_TIMEOUT_MS, + raceWithTimeout, + type DbDriver, + type LiveConnection, + type ResolvedDbConfig, + type ResolvedSslMode +} from './db-driver' +import { + mapColumnRows, + mapSchemaRows, + mapTableRows, + MYSQL_COLUMNS_SQL, + MYSQL_SCHEMAS_SQL, + MYSQL_TABLES_SQL +} from './mysql-introspection-queries' +import { cancelMysqlQuery, runMysqlBatch, runMysqlExecute, runMysqlQuery } from './mysql-query' +import type { + DbColumn, + DbSchemaTree, + DbStatement, + DbTableList, + DbTableRef, + QueryHandle, + QueryOptions, + QueryResult +} from '../../shared/database-types' +import type { Connection, ConnectionOptions, Pool, PoolOptions, RowDataPacket } from 'mysql2/promise' + +// 1 query + 1 introspection connection (red-team F11); small to bound backends. +const POOL_MAX = 2 +const POOL_IDLE_TIMEOUT_MS = 30_000 + +// disable → no TLS. verify-full → verify cert chain. insecure-no-verify → TLS +// without verification (explicit opt-in only). +export function buildMysqlSsl(ssl: ResolvedSslMode): PoolOptions['ssl'] { + if (ssl === 'disable') { + return undefined + } + // mysql2 only runs the TLS hostname check when `verifyIdentity` is set; with + // `rejectUnauthorized` alone it validates the chain but NOT the hostname, so a + // CA-signed cert issued for any other host would pass. verify-full must verify + // both (matches pg, which sets `servername` so Node checks the hostname). + return { + rejectUnauthorized: ssl === 'verify-full', + verifyIdentity: ssl === 'verify-full' + } +} + +// Connection-level config shared by the pool and the short-lived cancel session. +export function buildMysqlClientConfig(cfg: ResolvedDbConfig): ConnectionOptions { + return { + host: cfg.host, + port: cfg.port, + database: cfg.database, + user: cfg.user, + password: cfg.password, + ssl: buildMysqlSsl(cfg.ssl), + connectTimeout: DB_CONNECT_TIMEOUT_MS, + // Red-team F7: mysql2 keeps LOCAL INFILE (client file-read) disabled unless a + // stream factory is supplied — pin it undefined so the vector stays closed. + infileStreamFactory: undefined, + // Red-team F3/F7: single-statement only — the write guard and cancel logic + // both assume one statement per query; multi-statement would bypass them. + multipleStatements: false, + // BIGINT/DECIMAL exceed JS number precision (> 2^53); return them as strings + // so large values round-trip exactly — otherwise a PK-keyed edit could bind a + // rounded value and update the wrong row (or none). pg already returns bigint + // as a string. + supportBigNumbers: true, + bigNumberStrings: true + } +} + +export function buildMysqlPoolConfig(cfg: ResolvedDbConfig): PoolOptions { + return { + ...buildMysqlClientConfig(cfg), + connectionLimit: POOL_MAX, + idleTimeout: POOL_IDLE_TIMEOUT_MS + } +} + +// mysql2/promise is CommonJS; interop may hand back the module under `.default`. +// The core (callback) pool is reachable via `.pool` for the 'connection' event. +type MysqlPool = Pool & { pool?: { on(event: 'connection', cb: (c: unknown) => void): void } } +type MysqlModule = { + createPool: (config: PoolOptions) => MysqlPool + createConnection: (config: ConnectionOptions) => Promise +} +async function loadMysql(): Promise { + const mod = (await import('mysql2/promise')) as unknown as MysqlModule & { + default?: MysqlModule + } + return mod.default ?? mod +} + +async function validatePool(pool: Pool): Promise { + const conn = await raceWithTimeout(pool.getConnection(), DB_CONNECT_TIMEOUT_MS) + try { + // Bound the ping too: connectTimeout only covers the TCP dial, so a server + // that accepts the socket but never answers would otherwise hang connect(). + await raceWithTimeout(conn.query('SELECT 1'), DB_CONNECT_TIMEOUT_MS) + } finally { + conn.release() + } +} + +// Attach an 'error' listener to every pooled connection. mysql2 pools remove a +// broken connection internally, but the connection is still an EventEmitter — an +// 'error' with no listener re-throws and crashes the process (red-team F4). +function wireConnectionErrors(pool: MysqlPool, onError: (err: unknown) => void): void { + pool.pool?.on('connection', (conn) => { + const emitter = conn as { on?: (event: 'error', cb: (err: unknown) => void) => void } + emitter.on?.('error', (err) => onError(err)) + }) +} + +export const mysqlDriver: DbDriver = { + async testConnection(cfg: ResolvedDbConfig): Promise { + const mysql = await loadMysql() + const pool = mysql.createPool(buildMysqlPoolConfig(cfg)) + wireConnectionErrors(pool, () => {}) + try { + await validatePool(pool) + } finally { + await pool.end().catch(() => {}) + } + }, + + async connect( + cfg: ResolvedDbConfig, + onError: (err: unknown) => void + ): Promise { + const mysql = await loadMysql() + const pool = mysql.createPool(buildMysqlPoolConfig(cfg)) + wireConnectionErrors(pool, onError) + try { + await validatePool(pool) + } catch (err) { + await pool.end().catch(() => {}) + throw err + } + return { id: cfg.id, engine: 'mysql', raw: pool, config: cfg } + }, + + async introspectSchemas(conn: LiveConnection, maxSchemas: number): Promise { + // Query cap+1 so an overflow beyond the cap is detectable (red-team F9). + const [rows] = await (conn.raw as Pool).query({ + sql: MYSQL_SCHEMAS_SQL, + values: [maxSchemas + 1], + // Red-team M2: bound introspection runtime (client-side timeout). + timeout: DB_STATEMENT_TIMEOUT_MS + }) + const { kept, truncated } = applyCap(mapSchemaRows(rows as { name: string }[]), maxSchemas) + return { schemas: kept, truncated } + }, + + async introspectTables( + conn: LiveConnection, + schema: string, + maxTables: number + ): Promise { + const [rows] = await (conn.raw as Pool).query({ + sql: MYSQL_TABLES_SQL, + values: [schema, maxTables + 1], + timeout: DB_STATEMENT_TIMEOUT_MS + }) + const { kept, truncated } = applyCap( + mapTableRows(rows as { name: string; type: string }[]), + maxTables + ) + return { tables: kept, truncated } + }, + + async introspectColumns(conn: LiveConnection, ref: DbTableRef): Promise { + const [rows] = await (conn.raw as Pool).query({ + sql: MYSQL_COLUMNS_SQL, + values: [ref.schema, ref.table], + timeout: DB_STATEMENT_TIMEOUT_MS + }) + return mapColumnRows( + rows as { name: string; data_type: string; is_nullable: string; column_key: string }[] + ) + }, + + query( + conn: LiveConnection, + sql: string, + opts: QueryOptions, + onStart: (handle: QueryHandle) => void + ): Promise { + return runMysqlQuery(conn.raw as Pool, conn.id, sql, opts, onStart) + }, + + execute( + conn: LiveConnection, + statement: DbStatement, + opts: QueryOptions, + onStart: (handle: QueryHandle) => void + ): Promise { + return runMysqlExecute(conn.raw as Pool, conn.id, statement, opts, onStart) + }, + + executeBatch( + conn: LiveConnection, + statements: DbStatement[], + opts: QueryOptions, + onStart: (handle: QueryHandle) => void + ): Promise { + return runMysqlBatch(conn.raw as Pool, conn.id, statements, opts, onStart) + }, + + async cancel(conn: LiveConnection, handle: QueryHandle): Promise { + if (handle.backendPid == null) { + return + } + const mysql = await loadMysql() + await cancelMysqlQuery(await mysql.createConnection(buildMysqlClientConfig(conn.config)), handle.backendPid) + }, + + async close(conn: LiveConnection): Promise { + await (conn.raw as Pool).end() + } +} diff --git a/src/main/database/mysql-introspection-queries.test.ts b/src/main/database/mysql-introspection-queries.test.ts new file mode 100644 index 00000000000..61ec113359b --- /dev/null +++ b/src/main/database/mysql-introspection-queries.test.ts @@ -0,0 +1,71 @@ +import { describe, expect, it } from 'vitest' +import { + mapColumnRows, + mapSchemaRows, + mapTableRows, + MYSQL_COLUMNS_SQL, + MYSQL_SCHEMAS_SQL, + MYSQL_TABLES_SQL +} from './mysql-introspection-queries' + +describe('mysql introspection SQL', () => { + it('schemas query excludes system databases and caps via ?', () => { + expect(MYSQL_SCHEMAS_SQL).toContain('information_schema.schemata') + expect(MYSQL_SCHEMAS_SQL).toContain( + "NOT IN ('mysql', 'information_schema', 'performance_schema', 'sys')" + ) + expect(MYSQL_SCHEMAS_SQL).toContain('LIMIT ?') + }) + + it('tables query filters by schema and caps via ?', () => { + expect(MYSQL_TABLES_SQL).toContain('information_schema.tables') + expect(MYSQL_TABLES_SQL).toContain('table_schema = ?') + expect(MYSQL_TABLES_SQL).toContain('LIMIT ?') + }) + + it('columns query selects column_key and filters by schema+table', () => { + expect(MYSQL_COLUMNS_SQL).toContain('information_schema.columns') + expect(MYSQL_COLUMNS_SQL).toContain('column_key AS column_key') + expect(MYSQL_COLUMNS_SQL).toContain('table_schema = ? AND table_name = ?') + expect(MYSQL_COLUMNS_SQL).toContain('ORDER BY ordinal_position') + }) + + it('uses no window functions (MySQL 5.7 compatible)', () => { + for (const sql of [MYSQL_SCHEMAS_SQL, MYSQL_TABLES_SQL, MYSQL_COLUMNS_SQL]) { + expect(sql.toUpperCase()).not.toContain('ROW_NUMBER') + expect(sql.toUpperCase()).not.toContain('OVER (') + } + }) +}) + +describe('mysql row mappers', () => { + it('maps schema rows to names', () => { + expect(mapSchemaRows([{ name: 'app' }, { name: 'reporting' }])).toEqual(['app', 'reporting']) + }) + + it('treats BASE TABLE as table and any *VIEW as view', () => { + expect( + mapTableRows([ + { name: 't', type: 'BASE TABLE' }, + { name: 'v', type: 'VIEW' }, + { name: 'sv', type: 'SYSTEM VIEW' } + ]) + ).toEqual([ + { name: 't', kind: 'table' }, + { name: 'v', kind: 'view' }, + { name: 'sv', kind: 'view' } + ]) + }) + + it('maps column_key=PRI to primary key and is_nullable to nullable', () => { + expect( + mapColumnRows([ + { name: 'id', data_type: 'int', is_nullable: 'NO', column_key: 'PRI' }, + { name: 'email', data_type: 'varchar', is_nullable: 'YES', column_key: 'UNI' } + ]) + ).toEqual([ + { name: 'id', dataType: 'int', nullable: false, isPrimaryKey: true }, + { name: 'email', dataType: 'varchar', nullable: true, isPrimaryKey: false } + ]) + }) +}) diff --git a/src/main/database/mysql-introspection-queries.ts b/src/main/database/mysql-introspection-queries.ts new file mode 100644 index 00000000000..ba86dd0527e --- /dev/null +++ b/src/main/database/mysql-introspection-queries.ts @@ -0,0 +1,59 @@ +// MySQL introspection SQL + row mappers. Named/aliased columns keep the result +// keys stable (mysql2 returns keys as written in the SELECT) so the mappers read +// deterministic snake_case fields. Uses `?` placeholders and a plain LIMIT so it +// works on MySQL 5.7+ (no window functions). + +import type { DbColumn, DbTable } from '../../shared/database-types' + +// MySQL databases are the browsable schema level; exclude the system catalogs. +// ? = cap. +export const MYSQL_SCHEMAS_SQL = `SELECT schema_name AS name +FROM information_schema.schemata +WHERE schema_name NOT IN ('mysql', 'information_schema', 'performance_schema', 'sys') +ORDER BY schema_name +LIMIT ?` + +// Tables + views in one schema (database). ? = schema, ? = cap. +export const MYSQL_TABLES_SQL = `SELECT table_name AS name, table_type AS type +FROM information_schema.tables +WHERE table_schema = ? +ORDER BY table_name +LIMIT ?` + +// Columns of one table; column_key = 'PRI' marks a primary-key member. +// ? = schema, ? = table. +export const MYSQL_COLUMNS_SQL = `SELECT column_name AS name, + data_type AS data_type, + is_nullable AS is_nullable, + column_key AS column_key +FROM information_schema.columns +WHERE table_schema = ? AND table_name = ? +ORDER BY ordinal_position` + +export function mapSchemaRows(rows: { name: string }[]): string[] { + return rows.map((row) => row.name) +} + +export function mapTableRows(rows: { name: string; type: string }[]): DbTable[] { + // information_schema.table_type is 'BASE TABLE' | 'VIEW' | 'SYSTEM VIEW'. + return rows.map((row) => ({ + name: row.name, + kind: row.type.includes('VIEW') ? 'view' : 'table' + })) +} + +export function mapColumnRows( + rows: { + name: string + data_type: string + is_nullable: string + column_key: string + }[] +): DbColumn[] { + return rows.map((row) => ({ + name: row.name, + dataType: row.data_type, + nullable: row.is_nullable === 'YES', + isPrimaryKey: row.column_key === 'PRI' + })) +} diff --git a/src/main/database/mysql-query.test.ts b/src/main/database/mysql-query.test.ts new file mode 100644 index 00000000000..5208a1ff32a --- /dev/null +++ b/src/main/database/mysql-query.test.ts @@ -0,0 +1,332 @@ +import { Readable } from 'node:stream' +import { describe, expect, it, vi } from 'vitest' +import type { QueryOptions } from '../../shared/database-types' +import { DbBatchError } from './db-driver' +import { cancelMysqlQuery, runMysqlBatch, runMysqlExecute, runMysqlQuery } from './mysql-query' + +// Fake core query object: fires 'fields' when the stream starts, then streams +// each row as a chunk (object mode). +function makeCoreQuery(rows: unknown[][], fields: { name: string }[]) { + const listeners: Record void> = {} + return { + on(event: string, cb: (arg: unknown) => void) { + listeners[event] = cb + return this + }, + stream() { + listeners.fields?.(fields as unknown) + return Readable.from(rows) + } + } +} + +function makeConnection( + rows: unknown[][], + fields: { name: string }[], + options: { commitRejects?: boolean } = {} +) { + const calls: string[] = [] + const core = { query: vi.fn(() => makeCoreQuery(rows, fields)) } + const conn = { + query: vi.fn((arg: string | { sql: string }): Promise<[unknown, unknown]> => { + if (typeof arg === 'string') { + calls.push(arg) + if (arg === 'COMMIT' && options.commitRejects) { + return Promise.reject(new Error('commands out of sync')) + } + if (arg.includes('CONNECTION_ID')) { + return Promise.resolve([[{ id: 55 }], []]) + } + return Promise.resolve([[], []]) + } + calls.push(arg.sql) + return Promise.resolve([{ affectedRows: 3 }, []]) + }), + connection: core, + release: vi.fn(), + destroy: vi.fn() + } + return { calls, core, conn } +} + +function makePool(conn: unknown) { + return { getConnection: vi.fn(async () => conn) } as never +} + +function opts(overrides: Partial = {}): QueryOptions { + return { rowLimit: 2, timeoutMs: 30_000, allowWrite: false, ...overrides } +} + +describe('runMysqlQuery', () => { + it('streams a read in a read-only transaction and bounds the row count', async () => { + const { calls, core, conn } = makeConnection([[1], [2], [3]], [{ name: 'n' }]) + const onStart = vi.fn() + const result = await runMysqlQuery(makePool(conn), 'c1', 'SELECT * FROM t', opts(), onStart) + + expect(calls).toContain('START TRANSACTION READ ONLY') + expect(calls).toContain('SET SESSION max_execution_time = 30000') + // H2: COMMIT must NOT run on a poisoned (early-torn-down) stream connection. + expect(calls).not.toContain('COMMIT') + // The user's SQL streams verbatim — no appended LIMIT. + expect(core.query).toHaveBeenCalledWith({ sql: 'SELECT * FROM t', rowsAsArray: true }) + expect(result.rows).toEqual([[1], [2]]) + expect(result.truncated).toBe(true) + expect(onStart).toHaveBeenCalledWith({ connectionId: 'c1', backendPid: 55 }) + // A truncated stream poisons the connection → dropped, not returned to the pool. + expect(conn.destroy).toHaveBeenCalledTimes(1) + expect(conn.release).not.toHaveBeenCalled() + }) + + it('does not lose truncated rows even if COMMIT would fail (H2)', async () => { + // If COMMIT were issued on the poisoned connection it would reject; the fix + // skips it, so the valid truncated rows still come back. + const { conn } = makeConnection([[1], [2], [3]], [{ name: 'n' }], { commitRejects: true }) + const result = await runMysqlQuery(makePool(conn), 'c1', 'SELECT * FROM t', opts(), vi.fn()) + expect(result.rows).toEqual([[1], [2]]) + expect(result.truncated).toBe(true) + expect(conn.destroy).toHaveBeenCalledTimes(1) + }) + + it('returns all rows and releases the connection when under the cap', async () => { + const { conn } = makeConnection([[1], [2]], [{ name: 'n' }]) + const result = await runMysqlQuery(makePool(conn), 'c1', 'SELECT * FROM t', opts({ rowLimit: 5 }), vi.fn()) + expect(result.truncated).toBe(false) + expect(result.rows).toEqual([[1], [2]]) + expect(conn.release).toHaveBeenCalledTimes(1) + expect(conn.destroy).not.toHaveBeenCalled() + }) + + it('runs a write directly under a plain transaction (no read-only, no stream)', async () => { + const { calls, core, conn } = makeConnection([], []) + const result = await runMysqlQuery( + makePool(conn), + 'c1', + 'INSERT INTO t VALUES (1)', + opts({ allowWrite: true }), + vi.fn() + ) + expect(calls).toContain('START TRANSACTION') + expect(calls).not.toContain('START TRANSACTION READ ONLY') + expect(core.query).not.toHaveBeenCalled() + expect(result.rowCount).toBe(3) + }) +}) + +// Fake connection for the parameterized execute/batch paths: SELECT → rows, +// other DML → an affectedRows header. Records SQL + bind values per call. +function makeParamConnection(selectRows: unknown[][], fields: { name: string }[]) { + const calls: { sql: string; values?: unknown[] }[] = [] + // Core (callback) connection used by the bounded streaming read path. + const core = { + query: vi.fn((arg: { sql: string; values?: unknown[]; rowsAsArray: boolean }) => { + calls.push({ sql: arg.sql, values: arg.values }) + return makeCoreQuery(selectRows, fields) + }) + } + const conn = { + query: vi.fn((arg: string | { sql: string; values?: unknown[] }): Promise<[unknown, unknown]> => { + if (typeof arg === 'string') { + calls.push({ sql: arg }) + if (arg.includes('CONNECTION_ID')) { + return Promise.resolve([[{ id: 55 }], []]) + } + return Promise.resolve([[], []]) + } + // Non-cursorable statements (writes/DDL) run directly and return a header. + calls.push({ sql: arg.sql, values: arg.values }) + return Promise.resolve([{ affectedRows: 2 }, []]) + }), + connection: core, + release: vi.fn(), + destroy: vi.fn() + } + return { calls, conn } +} + +describe('runMysqlExecute', () => { + it('runs a read in a read-only transaction, threads params, and commits', async () => { + const { calls, conn } = makeParamConnection([[1], [2]], [{ name: 'n' }]) + const result = await runMysqlExecute( + makePool(conn), + 'c1', + { sql: 'SELECT * FROM t WHERE a = ? LIMIT 100 OFFSET 0', params: [9] }, + opts({ rowLimit: 1000 }), + vi.fn() + ) + const sqls = calls.map((c) => c.sql) + expect(sqls).toContain('START TRANSACTION READ ONLY') + expect(sqls).toContain('COMMIT') + const select = calls.find((c) => c.sql.startsWith('SELECT * FROM t')) + expect(select?.values).toEqual([9]) + expect(result.rows).toEqual([[1], [2]]) + expect(result.columns).toEqual([{ name: 'n' }]) + expect(conn.release).toHaveBeenCalledTimes(1) + }) + + it('uses a plain (writable) transaction when allowWrite', async () => { + const { calls, conn } = makeParamConnection([], []) + await runMysqlExecute( + makePool(conn), + 'c1', + { sql: 'UPDATE t SET a = ? WHERE id = ?', params: [1, 2] }, + opts({ allowWrite: true }), + vi.fn() + ) + const sqls = calls.map((c) => c.sql) + expect(sqls).toContain('START TRANSACTION') + expect(sqls).not.toContain('START TRANSACTION READ ONLY') + expect(sqls).toContain('COMMIT') + }) + + it('bounds a streamed read to rowLimit and drops the poisoned connection', async () => { + const { conn } = makeParamConnection([[1], [2], [3]], [{ name: 'n' }]) + const result = await runMysqlExecute( + makePool(conn), + 'c1', + { sql: 'SELECT * FROM t', params: [] }, + opts({ rowLimit: 2 }), + vi.fn() + ) + expect(result.rows).toEqual([[1], [2]]) + expect(result.truncated).toBe(true) + expect(conn.destroy).toHaveBeenCalledTimes(1) + }) + + it('destroys the connection when ROLLBACK fails after a write error', async () => { + const conn = { + query: vi.fn((arg: string | { sql: string; values?: unknown[] }): Promise<[unknown, unknown]> => { + const sql = typeof arg === 'string' ? arg : arg.sql + if (sql.includes('CONNECTION_ID')) { + return Promise.resolve([[{ id: 55 }], []]) + } + if (sql === 'ROLLBACK') { + return Promise.reject(new Error('rollback boom')) + } + // The write statement (object arg) itself fails. + if (typeof arg !== 'string') { + return Promise.reject(new Error('write boom')) + } + return Promise.resolve([[], []]) + }), + release: vi.fn(), + destroy: vi.fn() + } + await expect( + runMysqlExecute( + makePool(conn), + 'c1', + { sql: 'UPDATE t SET a = ?', params: [1] }, + opts({ allowWrite: true }), + vi.fn() + ) + ).rejects.toThrow() + expect(conn.destroy).toHaveBeenCalledTimes(1) + expect(conn.release).not.toHaveBeenCalled() + }) +}) + +describe('runMysqlBatch', () => { + it('applies every statement in one transaction and returns affectedRows per statement', async () => { + const { calls, conn } = makeParamConnection([], []) + const counts = await runMysqlBatch( + makePool(conn), + 'c1', + [ + { sql: 'UPDATE t SET a = ? WHERE id = ?', params: [1, 10] }, + { sql: 'DELETE FROM t WHERE id = ?', params: [11] } + ], + opts({ allowWrite: true }), + vi.fn() + ) + const sqls = calls.map((c) => c.sql) + expect(sqls).toContain('START TRANSACTION') + expect(sqls).toContain('COMMIT') + expect(sqls).not.toContain('ROLLBACK') + expect(counts).toEqual([2, 2]) + }) + + it('uses a read-only transaction when allowWrite is false', async () => { + const { calls, conn } = makeParamConnection([], []) + await runMysqlBatch( + makePool(conn), + 'c1', + [{ sql: 'UPDATE t SET a = ? WHERE id = ?', params: [1, 2] }], + opts({ allowWrite: false }), + vi.fn() + ) + expect(calls.map((c) => c.sql)).toContain('START TRANSACTION READ ONLY') + }) + + it('rolls back and throws DbBatchError with the failing index', async () => { + const calls: string[] = [] + const conn = { + query: vi.fn((arg: string | { sql: string; values?: unknown[] }): Promise<[unknown, unknown]> => { + const sql = typeof arg === 'string' ? arg : arg.sql + calls.push(sql) + if (sql.includes('CONNECTION_ID')) { + return Promise.resolve([[{ id: 55 }], []]) + } + if (sql.startsWith('DELETE')) { + return Promise.reject(Object.assign(new Error('fk'), { code: 'ER_ROW_IS_REFERENCED' })) + } + return Promise.resolve([{ affectedRows: 1 }, []]) + }), + release: vi.fn(), + destroy: vi.fn() + } + const err = await runMysqlBatch( + makePool(conn), + 'c1', + [ + { sql: 'UPDATE t SET a = ? WHERE id = ?', params: [1, 10] }, + { sql: 'DELETE FROM t WHERE id = ?', params: [11] } + ], + opts({ allowWrite: true }), + vi.fn() + ).catch((e) => e) + expect(err).toBeInstanceOf(DbBatchError) + expect((err as DbBatchError).failedIndex).toBe(1) + expect(calls).toContain('ROLLBACK') + expect(calls).not.toContain('COMMIT') + expect(conn.release).toHaveBeenCalledTimes(1) + }) + + it('destroys the connection when COMMIT fails', async () => { + const conn = { + query: vi.fn((arg: string | { sql: string; values?: unknown[] }): Promise<[unknown, unknown]> => { + const sql = typeof arg === 'string' ? arg : arg.sql + if (sql.includes('CONNECTION_ID')) { + return Promise.resolve([[{ id: 55 }], []]) + } + // COMMIT and its recovery ROLLBACK both fail → unknown txn state → drop it. + if (sql === 'COMMIT' || sql === 'ROLLBACK') { + return Promise.reject(new Error('boom')) + } + return Promise.resolve([{ affectedRows: 1 }, []]) + }), + release: vi.fn(), + destroy: vi.fn() + } + await expect( + runMysqlBatch( + makePool(conn), + 'c1', + [{ sql: 'UPDATE t SET a = ? WHERE id = ?', params: [1, 10] }], + opts({ allowWrite: true }), + vi.fn() + ) + ).rejects.toThrow() + expect(conn.destroy).toHaveBeenCalledTimes(1) + expect(conn.release).not.toHaveBeenCalled() + }) +}) + +describe('cancelMysqlQuery', () => { + it('issues KILL QUERY with the captured thread id and closes the session', async () => { + const query = vi.fn().mockResolvedValue([[], []]) + const end = vi.fn().mockResolvedValue(undefined) + await cancelMysqlQuery({ query, end } as never, 77) + expect(query).toHaveBeenCalledWith('KILL QUERY 77') + expect(end).toHaveBeenCalledTimes(1) + }) +}) diff --git a/src/main/database/mysql-query.ts b/src/main/database/mysql-query.ts new file mode 100644 index 00000000000..8e849cccd39 --- /dev/null +++ b/src/main/database/mysql-query.ts @@ -0,0 +1,276 @@ +// MySQL query execution + cancellation, split out of the driver so the streaming +// / transaction / cancel logic stays focused and testable. +// +// Read-only enforcement (red-team F3): reads run inside `START TRANSACTION READ +// ONLY`, so the database rejects writes — not a keyword check. Multi-statement is +// already off (multipleStatements:false). +// Result bounding (red-team F9): SELECTs stream row-by-row and stop after +// rowLimit+1, so a huge result never buffers whole; the user's SQL is never +// rewritten (no appended LIMIT). +// Cancellation (red-team F10): the connection id is captured at start and killed +// from a separate short-lived connection via KILL QUERY. + +import type { Readable } from 'node:stream' +import { isCursorableRead } from '../../shared/sql-statement-classifier' +import { DbBatchError } from './db-driver' +import type { + DbStatement, + QueryHandle, + QueryOptions, + QueryResult +} from '../../shared/database-types' +import type { Connection, Pool, PoolConnection } from 'mysql2/promise' + +type BoundedRows = { columns: { name: string }[]; rows: unknown[][]; rowCount: number; truncated: boolean } + +// Minimal shape of the core (callback) connection reachable via `.connection`, +// used only for row-by-row streaming (the promise API buffers whole results). +type StreamingCore = { + query(opts: { sql: string; values?: unknown[]; rowsAsArray: boolean }): { + on(event: 'fields', cb: (fields: { name: string }[]) => void): unknown + on(event: 'error', cb: (err: unknown) => void): unknown + stream(): Readable + } +} + +function normalizeFields(fields: { name: string }[] | undefined): { name: string }[] { + return (fields ?? []).map((f) => ({ name: f.name })) +} + +// Stream a SELECT, keeping at most rowLimit rows; the (rowLimit+1)-th row only +// flips `truncated` and triggers an early stream teardown. +function streamBounded( + core: StreamingCore, + sql: string, + rowLimit: number, + onEarlyStop: () => void, + values?: unknown[] +): Promise { + return new Promise((resolve, reject) => { + let columns: { name: string }[] = [] + const rows: unknown[][] = [] + let settled = false + const settle = (fn: (value: never) => void, value: unknown): void => { + if (!settled) { + settled = true + ;(fn as (v: unknown) => void)(value) + } + } + const q = core.query(values ? { sql, values, rowsAsArray: true } : { sql, rowsAsArray: true }) + q.on('fields', (fields) => { + columns = normalizeFields(fields) + }) + q.on('error', (err) => settle(reject, err)) + const stream = q.stream() + stream.on('data', (row: unknown[]) => { + if (rows.length >= rowLimit) { + onEarlyStop() + stream.destroy() + settle(resolve, { columns, rows, rowCount: rows.length, truncated: true }) + return + } + rows.push(row) + }) + stream.on('error', (err) => settle(reject, err)) + stream.on('end', () => settle(resolve, { columns, rows, rowCount: rows.length, truncated: false })) + }) +} + +// Non-cursorable statements (writes/DDL) run directly; they return a header, not +// rows. rowsAsArray keeps any returned rows positional. A client-side `timeout` +// bounds writes (red-team L1) — `max_execution_time` only covers SELECT. +async function runDirect( + conn: PoolConnection, + sql: string, + timeoutMs: number, + values?: unknown[] +): Promise { + const [result, fields] = await conn.query( + values + ? { sql, values, rowsAsArray: true, timeout: timeoutMs } + : { sql, rowsAsArray: true, timeout: timeoutMs } + ) + if (Array.isArray(result)) { + const rows = result as unknown[][] + return { columns: normalizeFields(fields as { name: string }[]), rows, rowCount: rows.length, truncated: false } + } + const affected = (result as { affectedRows?: number }).affectedRows ?? 0 + return { columns: [], rows: [], rowCount: affected, truncated: false } +} + +export async function runMysqlQuery( + pool: Pool, + connectionId: string, + sql: string, + opts: QueryOptions, + onStart: (handle: QueryHandle) => void +): Promise { + const conn = await pool.getConnection() + let poisoned = false + try { + const [pidRows] = await conn.query('SELECT CONNECTION_ID() AS id') + const backendPid = Number((pidRows as { id?: number }[])[0]?.id) || null + onStart({ connectionId, backendPid }) + + // max_execution_time bounds SELECTs (ms); the read-only transaction is the + // write boundary when !allowWrite. + await conn.query(`SET SESSION max_execution_time = ${Math.trunc(opts.timeoutMs)}`) + await conn.query(opts.allowWrite ? 'START TRANSACTION' : 'START TRANSACTION READ ONLY') + const startedAt = Date.now() + try { + const bounded = isCursorableRead(sql) + ? await streamBounded( + (conn as unknown as { connection: StreamingCore }).connection, + sql, + opts.rowLimit, + () => { + poisoned = true + } + ) + : await runDirect(conn, sql, opts.timeoutMs) + // H2 (red-team): a stream torn down early leaves undrained packets on the + // socket — COMMIT here would either desync (losing the valid truncated + // rows behind a spurious error) or drain the full result (defeating the + // bound). Skip it: the read-only transaction is abandoned when `finally` + // destroys the poisoned connection. + if (!poisoned) { + await conn.query('COMMIT') + } + return { ...bounded, durationMs: Date.now() - startedAt } + } catch (err) { + // A failed ROLLBACK leaves the transaction state unknown — poison the + // connection so `finally` drops it instead of returning it to the pool. + await conn.query('ROLLBACK').catch(() => { + poisoned = true + }) + throw err + } + } finally { + // A stream torn down mid-result leaves unread packets on the socket; drop the + // connection instead of returning a poisoned one to the pool. + if (poisoned) { + conn.destroy() + } else { + conn.release() + } + } +} + +// Runs one parameterized statement (Data-tab select/count/mutation or wrapped +// free-form re-query) inside a transaction — read-only when !allowWrite. The +// statement's own LIMIT bounds the page; rowLimit is a defensive server-side cap. +export async function runMysqlExecute( + pool: Pool, + connectionId: string, + statement: DbStatement, + opts: QueryOptions, + onStart: (handle: QueryHandle) => void +): Promise { + const conn = await pool.getConnection() + let poisoned = false + try { + const [pidRows] = await conn.query('SELECT CONNECTION_ID() AS id') + onStart({ connectionId, backendPid: Number((pidRows as { id?: number }[])[0]?.id) || null }) + await conn.query(`SET SESSION max_execution_time = ${Math.trunc(opts.timeoutMs)}`) + await conn.query(opts.allowWrite ? 'START TRANSACTION' : 'START TRANSACTION READ ONLY') + const startedAt = Date.now() + try { + // Row-producing statements stream bounded to rowLimit+1 so a statement that + // ever ships without its own LIMIT can't materialize the whole result set; + // writes/DDL run directly (they return a header, not rows). + const bounded = isCursorableRead(statement.sql) + ? await streamBounded( + (conn as unknown as { connection: StreamingCore }).connection, + statement.sql, + opts.rowLimit, + () => { + poisoned = true + }, + statement.params + ) + : await runDirect(conn, statement.sql, opts.timeoutMs, statement.params) + // H2: skip COMMIT on a stream torn down early (undrained packets); the txn + // is abandoned when `finally` destroys the poisoned connection. + if (!poisoned) { + await conn.query('COMMIT') + } + return { ...bounded, durationMs: Date.now() - startedAt } + } catch (err) { + // A failed ROLLBACK leaves the transaction state unknown — poison the + // connection so `finally` drops it instead of returning it to the pool. + await conn.query('ROLLBACK').catch(() => { + poisoned = true + }) + throw err + } + } finally { + if (poisoned) { + conn.destroy() + } else { + conn.release() + } + } +} + +// Applies staged writes atomically: START TRANSACTION, run each statement in +// order, COMMIT. Any failure rolls back and throws DbBatchError(failedIndex). +export async function runMysqlBatch( + pool: Pool, + connectionId: string, + statements: DbStatement[], + opts: QueryOptions, + onStart: (handle: QueryHandle) => void +): Promise { + const conn = await pool.getConnection() + let poisoned = false + try { + const [pidRows] = await conn.query('SELECT CONNECTION_ID() AS id') + onStart({ connectionId, backendPid: Number((pidRows as { id?: number }[])[0]?.id) || null }) + await conn.query(`SET SESSION max_execution_time = ${Math.trunc(opts.timeoutMs)}`) + // Defense in depth: keep a read-only connection DB-enforced on this write path + // too (the manager already rejects !allowWrite before reaching here). + await conn.query(opts.allowWrite ? 'START TRANSACTION' : 'START TRANSACTION READ ONLY') + const rowCounts: number[] = [] + for (let i = 0; i < statements.length; i++) { + try { + const [result] = await conn.query({ + sql: statements[i].sql, + values: statements[i].params, + timeout: opts.timeoutMs + }) + rowCounts.push((result as { affectedRows?: number }).affectedRows ?? 0) + } catch (err) { + await conn.query('ROLLBACK').catch(() => { + poisoned = true + }) + throw new DbBatchError(i, err) + } + } + // A failing COMMIT leaves the transaction in an unknown state — roll back and + // poison the connection on failure rather than returning it to the pool. + try { + await conn.query('COMMIT') + } catch (err) { + await conn.query('ROLLBACK').catch(() => { + poisoned = true + }) + throw err + } + return rowCounts + } finally { + if (poisoned) { + conn.destroy() + } else { + conn.release() + } + } +} + +export async function cancelMysqlQuery(conn: Connection, threadId: number): Promise { + try { + // KILL QUERY takes no placeholder; threadId is a captured integer. + await conn.query(`KILL QUERY ${Math.trunc(threadId)}`) + } finally { + await conn.end().catch(() => {}) + } +} diff --git a/src/main/database/postgres-driver.test.ts b/src/main/database/postgres-driver.test.ts new file mode 100644 index 00000000000..5a282220ac7 --- /dev/null +++ b/src/main/database/postgres-driver.test.ts @@ -0,0 +1,141 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest' +import type { ResolvedDbConfig } from './db-driver' + +// Fake pg.Pool capturing config + listeners so we can assert the driver wires +// the mandatory 'error' listener and validates before returning. +const okConnect = (): Promise => + Promise.resolve({ query: vi.fn().mockResolvedValue({ rows: [] }), release: vi.fn() }) + +class FakePool { + static instances: FakePool[] = [] + // Constructor-time override so a test can arm failing validation before the + // driver's internal `await import('pg')` resolves and creates the pool. + static nextConnect: (() => Promise) | null = null + // Rows returned by pool-level introspection queries (pg returns { rows }). + static queryRows: unknown[] = [] + config: Record + listeners: Record void> = {} + ended = false + connectImpl: () => Promise + + constructor(config: Record) { + this.config = config + this.connectImpl = FakePool.nextConnect ?? okConnect + FakePool.nextConnect = null + FakePool.instances.push(this) + } + + on(event: string, cb: (arg: unknown) => void): void { + this.listeners[event] = cb + } + + connect(): Promise { + return this.connectImpl() + } + + query(): Promise<{ rows: unknown[] }> { + return Promise.resolve({ rows: FakePool.queryRows }) + } + + async end(): Promise { + this.ended = true + } +} + +vi.mock('pg', () => ({ default: { Pool: FakePool } })) + +import { buildPgPoolConfig, buildPgSsl, postgresDriver } from './postgres-driver' + +function cfg(overrides: Partial = {}): ResolvedDbConfig { + return { + id: 'c1', + engine: 'postgres', + host: 'db.example.com', + port: 5432, + database: 'app', + user: 'admin', + password: 'pw', + ssl: 'verify-full', + readOnly: false, + ...overrides + } +} + +describe('buildPgSsl', () => { + it('disables TLS for ssl=disable', () => { + expect(buildPgSsl('disable')).toBe(false) + }) + + it('verifies certs for ssl=verify-full', () => { + expect(buildPgSsl('verify-full')).toEqual({ rejectUnauthorized: true }) + }) + + it('does not verify certs for ssl=insecure-no-verify', () => { + expect(buildPgSsl('insecure-no-verify')).toEqual({ rejectUnauthorized: false }) + }) +}) + +describe('buildPgPoolConfig', () => { + it('sets a bounded connect timeout and a small pool', () => { + const pool = buildPgPoolConfig(cfg()) + expect(pool.connectionTimeoutMillis).toBeGreaterThan(0) + expect(pool.max).toBe(2) + expect(pool.ssl).toEqual({ rejectUnauthorized: true }) + }) +}) + +describe('postgresDriver', () => { + beforeEach(() => { + FakePool.instances = [] + FakePool.nextConnect = null + FakePool.queryRows = [] + }) + + it('connect attaches an error listener, validates, and returns a live connection', async () => { + const onError = vi.fn() + const conn = await postgresDriver.connect(cfg(), onError) + const pool = FakePool.instances[0] + expect(pool.listeners.error).toBeTypeOf('function') + expect(conn.engine).toBe('postgres') + expect(conn.raw).toBe(pool) + + // The wired listener forwards a dropped-connection error to the manager. + pool.listeners.error(new Error('idle client error')) + expect(onError).toHaveBeenCalledTimes(1) + }) + + it('connect ends the pool and rethrows when validation fails', async () => { + // Arm the next pool to fail its validation query before connect creates it. + FakePool.nextConnect = () => + Promise.resolve({ + query: vi.fn().mockRejectedValue(Object.assign(new Error('nope'), { code: '28P01' })), + release: vi.fn() + }) + await expect(postgresDriver.connect(cfg(), vi.fn())).rejects.toThrow() + expect(FakePool.instances[0].ended).toBe(true) + }) + + it('testConnection ends the pool even on success (holds no state)', async () => { + await postgresDriver.testConnection(cfg()) + expect(FakePool.instances[0].ended).toBe(true) + }) + + it('introspectSchemas maps rows and flags truncation past the cap', async () => { + const conn = await postgresDriver.connect(cfg(), vi.fn()) + FakePool.queryRows = [{ name: 'public' }, { name: 'app' }, { name: 'extra' }] + const tree = await postgresDriver.introspectSchemas(conn, 2) + expect(tree.schemas).toEqual(['public', 'app']) + expect(tree.truncated).toBe(true) + }) + + it('introspectColumns maps nullability and primary-key membership', async () => { + const conn = await postgresDriver.connect(cfg(), vi.fn()) + FakePool.queryRows = [ + { name: 'id', data_type: 'integer', is_nullable: 'NO', is_primary_key: true } + ] + const columns = await postgresDriver.introspectColumns(conn, { schema: 'public', table: 't' }) + expect(columns).toEqual([ + { name: 'id', dataType: 'integer', nullable: false, isPrimaryKey: true } + ]) + }) +}) diff --git a/src/main/database/postgres-driver.ts b/src/main/database/postgres-driver.ts new file mode 100644 index 00000000000..507283985ce --- /dev/null +++ b/src/main/database/postgres-driver.ts @@ -0,0 +1,195 @@ +// Postgres driver: lazy-imports `pg` on first use (kept out of the startup +// require-graph), dials over a small pool with a bounded connect timeout, and +// attaches the mandatory pool 'error' listener so a dropped idle client degrades +// to `lost` instead of crashing the main process. + +import { + applyCap, + DB_CONNECT_TIMEOUT_MS, + DB_STATEMENT_TIMEOUT_MS, + raceWithTimeout, + type DbDriver, + type LiveConnection, + type ResolvedDbConfig, + type ResolvedSslMode +} from './db-driver' +import { + mapColumnRows, + mapSchemaRows, + mapTableRows, + PG_COLUMNS_SQL, + PG_SCHEMAS_SQL, + PG_TABLES_SQL +} from './postgres-introspection-queries' +import { + cancelPostgresBackend, + runPostgresBatch, + runPostgresExecute, + runPostgresQuery +} from './postgres-query' +import type { + DbColumn, + DbSchemaTree, + DbStatement, + DbTableList, + DbTableRef, + QueryHandle, + QueryOptions, + QueryResult +} from '../../shared/database-types' +import type { Client, ClientConfig, Pool, PoolConfig } from 'pg' + +// Why: 1 query connection + 1 introspection connection so introspect (P4) and +// query (P5) don't serialize on a single socket (red-team F11). Kept small to +// bound the server-side backend count. +const POOL_MAX = 2 +const POOL_IDLE_TIMEOUT_MS = 30_000 + +// disable → no TLS. verify-full → TLS with cert + hostname verification. +// insecure-no-verify → TLS without verification (explicit opt-in only). +export function buildPgSsl(ssl: ResolvedSslMode): PoolConfig['ssl'] { + if (ssl === 'disable') { + return false + } + return { rejectUnauthorized: ssl === 'verify-full' } +} + +// Connection-level config, shared by the pool and the short-lived cancel client. +export function buildPgClientConfig(cfg: ResolvedDbConfig): ClientConfig { + return { + host: cfg.host, + port: cfg.port, + database: cfg.database, + user: cfg.user, + password: cfg.password, + ssl: buildPgSsl(cfg.ssl), + connectionTimeoutMillis: DB_CONNECT_TIMEOUT_MS, + // Red-team M2: bound every statement on the connection (introspection + + // validate), not just connect. The query path re-SETs this per transaction. + statement_timeout: DB_STATEMENT_TIMEOUT_MS + } +} + +export function buildPgPoolConfig(cfg: ResolvedDbConfig): PoolConfig { + return { + ...buildPgClientConfig(cfg), + max: POOL_MAX, + idleTimeoutMillis: POOL_IDLE_TIMEOUT_MS + } +} + +// pg is CommonJS; interop may hand back the module under `.default`. +type PgModule = { + Pool: new (config?: PoolConfig) => Pool + Client: new (config?: ClientConfig) => Client +} +async function loadPg(): Promise { + const mod = (await import('pg')) as unknown as PgModule & { default?: PgModule } + return mod.default ?? mod +} + +async function validatePool(pool: Pool): Promise { + const client = await raceWithTimeout(pool.connect(), DB_CONNECT_TIMEOUT_MS) + try { + await client.query('SELECT 1') + } finally { + client.release() + } +} + +export const postgresDriver: DbDriver = { + async testConnection(cfg: ResolvedDbConfig): Promise { + const pg = await loadPg() + const pool = new pg.Pool(buildPgPoolConfig(cfg)) + // Why: even a throwaway test pool is an EventEmitter — an 'error' with no + // listener re-throws and crashes the process. + pool.on('error', () => {}) + try { + await validatePool(pool) + } finally { + await pool.end().catch(() => {}) + } + }, + + async connect( + cfg: ResolvedDbConfig, + onError: (err: unknown) => void + ): Promise { + const pg = await loadPg() + const pool = new pg.Pool(buildPgPoolConfig(cfg)) + // Red-team F4 (Critical): a pooled idle client whose socket drops emits + // 'error' on the pool; unhandled, it crashes every PTY/SSH/terminal. + pool.on('error', (err) => onError(err)) + try { + await validatePool(pool) + } catch (err) { + await pool.end().catch(() => {}) + throw err + } + return { id: cfg.id, engine: 'postgres', raw: pool, config: cfg } + }, + + async introspectSchemas(conn: LiveConnection, maxSchemas: number): Promise { + // Query cap+1 so an overflow beyond the cap is detectable (red-team F9). + const result = await (conn.raw as Pool).query(PG_SCHEMAS_SQL, [maxSchemas + 1]) + const { kept, truncated } = applyCap(mapSchemaRows(result.rows), maxSchemas) + return { schemas: kept, truncated } + }, + + async introspectTables( + conn: LiveConnection, + schema: string, + maxTables: number + ): Promise { + const result = await (conn.raw as Pool).query(PG_TABLES_SQL, [schema, maxTables + 1]) + const { kept, truncated } = applyCap(mapTableRows(result.rows), maxTables) + return { tables: kept, truncated } + }, + + async introspectColumns(conn: LiveConnection, ref: DbTableRef): Promise { + const result = await (conn.raw as Pool).query(PG_COLUMNS_SQL, [ref.schema, ref.table]) + return mapColumnRows(result.rows) + }, + + query( + conn: LiveConnection, + sql: string, + opts: QueryOptions, + onStart: (handle: QueryHandle) => void + ): Promise { + return runPostgresQuery(conn.raw as Pool, conn.id, sql, opts, onStart) + }, + + execute( + conn: LiveConnection, + statement: DbStatement, + opts: QueryOptions, + onStart: (handle: QueryHandle) => void + ): Promise { + return runPostgresExecute(conn.raw as Pool, conn.id, statement, opts, onStart) + }, + + executeBatch( + conn: LiveConnection, + statements: DbStatement[], + opts: QueryOptions, + onStart: (handle: QueryHandle) => void + ): Promise { + return runPostgresBatch(conn.raw as Pool, conn.id, statements, opts, onStart) + }, + + async cancel(conn: LiveConnection, handle: QueryHandle): Promise { + if (handle.backendPid == null) { + return + } + const pg = await loadPg() + await cancelPostgresBackend( + new pg.Client(buildPgClientConfig(conn.config)), + handle.backendPid + ) + }, + + async close(conn: LiveConnection): Promise { + await (conn.raw as Pool).end() + } +} diff --git a/src/main/database/postgres-introspection-queries.test.ts b/src/main/database/postgres-introspection-queries.test.ts new file mode 100644 index 00000000000..b591068f24f --- /dev/null +++ b/src/main/database/postgres-introspection-queries.test.ts @@ -0,0 +1,61 @@ +import { describe, expect, it } from 'vitest' +import { + mapColumnRows, + mapSchemaRows, + mapTableRows, + PG_COLUMNS_SQL, + PG_SCHEMAS_SQL, + PG_TABLES_SQL +} from './postgres-introspection-queries' + +describe('postgres introspection SQL', () => { + it('schemas query excludes system namespaces and caps via $1', () => { + expect(PG_SCHEMAS_SQL).toContain('information_schema.schemata') + expect(PG_SCHEMAS_SQL).toContain("NOT IN ('pg_catalog', 'information_schema')") + expect(PG_SCHEMAS_SQL).toContain("NOT LIKE 'pg_%'") + expect(PG_SCHEMAS_SQL).toContain('LIMIT $1') + }) + + it('tables query filters by schema ($1) and caps via $2', () => { + expect(PG_TABLES_SQL).toContain('information_schema.tables') + expect(PG_TABLES_SQL).toContain('table_schema = $1') + expect(PG_TABLES_SQL).toContain('LIMIT $2') + }) + + it('columns query joins primary-key membership and filters by schema+table', () => { + expect(PG_COLUMNS_SQL).toContain('information_schema.columns') + expect(PG_COLUMNS_SQL).toContain("constraint_type = 'PRIMARY KEY'") + expect(PG_COLUMNS_SQL).toContain('c.table_schema = $1 AND c.table_name = $2') + expect(PG_COLUMNS_SQL).toContain('ORDER BY c.ordinal_position') + }) +}) + +describe('postgres row mappers', () => { + it('maps schema rows to names', () => { + expect(mapSchemaRows([{ name: 'public' }, { name: 'app' }])).toEqual(['public', 'app']) + }) + + it('maps table_type to table/view kind', () => { + expect( + mapTableRows([ + { name: 't', type: 'BASE TABLE' }, + { name: 'v', type: 'VIEW' } + ]) + ).toEqual([ + { name: 't', kind: 'table' }, + { name: 'v', kind: 'view' } + ]) + }) + + it('maps column rows including nullability and primary key', () => { + expect( + mapColumnRows([ + { name: 'id', data_type: 'integer', is_nullable: 'NO', is_primary_key: true }, + { name: 'note', data_type: 'text', is_nullable: 'YES', is_primary_key: false } + ]) + ).toEqual([ + { name: 'id', dataType: 'integer', nullable: false, isPrimaryKey: true }, + { name: 'note', dataType: 'text', nullable: true, isPrimaryKey: false } + ]) + }) +}) diff --git a/src/main/database/postgres-introspection-queries.ts b/src/main/database/postgres-introspection-queries.ts new file mode 100644 index 00000000000..f56f97fe467 --- /dev/null +++ b/src/main/database/postgres-introspection-queries.ts @@ -0,0 +1,68 @@ +// Postgres introspection SQL + row mappers. Kept pure and separate so the SQL is +// unit-testable and the driver just supplies params (schema name, cap+1) and maps +// rows. Aliases are snake_case because Postgres folds unquoted identifiers to +// lowercase — the mappers translate to the camelCase shared types. + +import type { DbColumn, DbTable } from '../../shared/database-types' + +// User schemas in the connected database, excluding system namespaces. $1 = cap. +export const PG_SCHEMAS_SQL = `SELECT schema_name AS name +FROM information_schema.schemata +WHERE schema_name NOT IN ('pg_catalog', 'information_schema') + AND schema_name NOT LIKE 'pg_%' +ORDER BY schema_name +LIMIT $1` + +// Tables + views in one schema. $1 = schema, $2 = cap. +export const PG_TABLES_SQL = `SELECT table_name AS name, table_type AS type +FROM information_schema.tables +WHERE table_schema = $1 +ORDER BY table_name +LIMIT $2` + +// Columns of one table with primary-key membership. $1 = schema, $2 = table. +export const PG_COLUMNS_SQL = `SELECT c.column_name AS name, + c.data_type AS data_type, + c.is_nullable AS is_nullable, + (pk.column_name IS NOT NULL) AS is_primary_key +FROM information_schema.columns c +LEFT JOIN ( + 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 + AND tc.table_name = kcu.table_name + WHERE tc.constraint_type = 'PRIMARY KEY' + AND tc.table_schema = $1 + AND tc.table_name = $2 +) pk ON pk.column_name = c.column_name +WHERE c.table_schema = $1 AND c.table_name = $2 +ORDER BY c.ordinal_position` + +export function mapSchemaRows(rows: { name: string }[]): string[] { + return rows.map((row) => row.name) +} + +export function mapTableRows(rows: { name: string; type: string }[]): DbTable[] { + return rows.map((row) => ({ + name: row.name, + kind: row.type === 'VIEW' ? 'view' : 'table' + })) +} + +export function mapColumnRows( + rows: { + name: string + data_type: string + is_nullable: string + is_primary_key: boolean + }[] +): DbColumn[] { + return rows.map((row) => ({ + name: row.name, + dataType: row.data_type, + nullable: row.is_nullable === 'YES', + isPrimaryKey: row.is_primary_key === true + })) +} diff --git a/src/main/database/postgres-query.test.ts b/src/main/database/postgres-query.test.ts new file mode 100644 index 00000000000..f67babe8e50 --- /dev/null +++ b/src/main/database/postgres-query.test.ts @@ -0,0 +1,337 @@ +import { describe, expect, it, vi } from 'vitest' +import type { QueryOptions } from '../../shared/database-types' +import { DbBatchError } from './db-driver' +import { + cancelPostgresBackend, + runPostgresBatch, + runPostgresExecute, + runPostgresQuery +} from './postgres-query' + +type QueryArg = string | { text: string; rowMode?: string } +function textOf(arg: QueryArg): string { + return typeof arg === 'string' ? arg : arg.text +} + +// Fake pooled client recording every query. Returns rows for pid + FETCH + direct. +function makeClient(fetchRows: unknown[][], fields: { name: string }[]) { + const calls: string[] = [] + const release = vi.fn() + const query = vi.fn((arg: QueryArg) => { + const text = textOf(arg) + calls.push(text) + if (text.includes('pg_backend_pid')) { + return Promise.resolve({ rows: [{ pid: 42 }] }) + } + if (text.startsWith('FETCH FORWARD')) { + return Promise.resolve({ rows: fetchRows, fields }) + } + if (typeof arg === 'object') { + // Direct (non-cursor) execution path. + return Promise.resolve({ rows: fetchRows, fields, rowCount: fetchRows.length }) + } + return Promise.resolve({ rows: [], rowCount: 0 }) + }) + return { calls, client: { query, release } } +} + +function makePool(client: { query: unknown; release: unknown }) { + return { connect: vi.fn(async () => client) } as never +} + +function opts(overrides: Partial = {}): QueryOptions { + return { rowLimit: 2, timeoutMs: 30_000, allowWrite: false, ...overrides } +} + +describe('runPostgresQuery', () => { + it('runs a read in a read-only transaction and bounds via a cursor', async () => { + const { calls, client } = makeClient([[1], [2], [3]], [{ name: 'n' }]) + const onStart = vi.fn() + const result = await runPostgresQuery(makePool(client), 'c1', 'SELECT * FROM t', opts(), onStart) + + expect(calls).toContain('BEGIN') + expect(calls).toContain('SET TRANSACTION READ ONLY') + expect(calls).toContain('COMMIT') + // Cursor fetches rowLimit+1 to detect overflow; only rowLimit rows returned. + expect(calls).toContain('FETCH FORWARD 3 FROM orca_query_cursor') + expect(result.rows).toEqual([[1], [2]]) + expect(result.truncated).toBe(true) + expect(result.rowCount).toBe(2) + expect(onStart).toHaveBeenCalledWith({ connectionId: 'c1', backendPid: 42 }) + expect(client.release).toHaveBeenCalledTimes(1) + }) + + it('embeds the user SQL verbatim in DECLARE — never appends LIMIT', async () => { + const { calls, client } = makeClient([[1]], [{ name: 'n' }]) + await runPostgresQuery(makePool(client), 'c1', 'SELECT * FROM t', opts(), vi.fn()) + const declare = calls.find((c) => c.startsWith('DECLARE')) + expect(declare).toBe('DECLARE orca_query_cursor NO SCROLL CURSOR FOR SELECT * FROM t') + expect(calls.some((c) => /\bLIMIT\b/i.test(c))).toBe(false) + }) + + it('runs a writable non-cursorable statement in autocommit (no transaction)', async () => { + const { calls, client } = makeClient([], []) + await runPostgresQuery( + makePool(client), + 'c1', + 'INSERT INTO t VALUES (1)', + opts({ allowWrite: true }), + vi.fn() + ) + expect(calls).not.toContain('SET TRANSACTION READ ONLY') + // A write is non-cursorable → runs directly, no DECLARE. + expect(calls.some((c) => c.startsWith('DECLARE'))).toBe(false) + // No explicit transaction wraps it, so a transaction-block-unsafe command + // (VACUUM, CREATE DATABASE, REINDEX CONCURRENTLY) isn't rejected. + expect(calls).not.toContain('BEGIN') + expect(calls).not.toContain('COMMIT') + }) + + it('does not open a transaction for VACUUM on a writable connection', async () => { + const { calls, client } = makeClient([], []) + await runPostgresQuery( + makePool(client), + 'c1', + 'VACUUM ANALYZE t', + opts({ allowWrite: true }), + vi.fn() + ) + expect(calls).not.toContain('BEGIN') + expect(calls).not.toContain('COMMIT') + }) + + it('still wraps a writable cursorable read in a transaction for the cursor', async () => { + const { calls, client } = makeClient([[1]], [{ name: 'n' }]) + await runPostgresQuery( + makePool(client), + 'c1', + 'SELECT * FROM t', + opts({ allowWrite: true }), + vi.fn() + ) + expect(calls).toContain('BEGIN') + expect(calls).toContain('COMMIT') + // Writable, so no read-only downgrade — but the cursor still bounds the read. + expect(calls).not.toContain('SET TRANSACTION READ ONLY') + expect(calls.some((c) => c.startsWith('DECLARE'))).toBe(true) + }) + + it('rolls back and rethrows when the query fails', async () => { + const calls: string[] = [] + const client = { + release: vi.fn(), + query: vi.fn((arg: QueryArg) => { + const text = textOf(arg) + calls.push(text) + if (text.includes('pg_backend_pid')) { + return Promise.resolve({ rows: [{ pid: 7 }] }) + } + if (text.startsWith('DECLARE')) { + return Promise.reject(Object.assign(new Error('boom'), { code: '42P01' })) + } + return Promise.resolve({ rows: [], rowCount: 0 }) + }) + } + await expect( + runPostgresQuery(makePool(client), 'c1', 'SELECT * FROM missing', opts(), vi.fn()) + ).rejects.toThrow() + expect(calls).toContain('ROLLBACK') + expect(client.release).toHaveBeenCalledTimes(1) + }) + + it('evicts the client from the pool when ROLLBACK also fails', async () => { + const client = { + release: vi.fn(), + query: vi.fn((arg: QueryArg) => { + const text = textOf(arg) + if (text.includes('pg_backend_pid')) { + return Promise.resolve({ rows: [{ pid: 7 }] }) + } + // Fail the query and its ROLLBACK: the client may still hold an open + // transaction and must be evicted, not returned to the pool. + if (text.startsWith('DECLARE') || text === 'ROLLBACK') { + return Promise.reject(new Error('boom')) + } + return Promise.resolve({ rows: [], rowCount: 0 }) + }) + } + await expect( + runPostgresQuery(makePool(client), 'c1', 'SELECT * FROM missing', opts(), vi.fn()) + ).rejects.toThrow() + expect(client.release).toHaveBeenCalledWith(true) + }) +}) + +// Fake client for the parameterized execute/batch paths: records SQL + values. +function makeParamClient(rows: unknown[][], fields: { name: string }[]) { + const calls: { text: string; values?: unknown[] }[] = [] + const release = vi.fn() + const query = vi.fn((arg: QueryArg) => { + const text = textOf(arg) + const values = typeof arg === 'object' ? (arg as { values?: unknown[] }).values : undefined + calls.push({ text, values }) + if (text.includes('pg_backend_pid')) { + return Promise.resolve({ rows: [{ pid: 7 }] }) + } + if (typeof arg === 'object' && 'rowMode' in arg) { + return Promise.resolve({ rows, fields, rowCount: rows.length }) + } + return Promise.resolve({ rows: [], rowCount: 1 }) + }) + return { calls, client: { query, release } } +} + +describe('runPostgresExecute', () => { + it('runs a read in a read-only transaction and threads bind params', async () => { + const { calls, client } = makeParamClient([[1], [2]], [{ name: 'n' }]) + const result = await runPostgresExecute( + makePool(client), + 'c1', + { sql: 'SELECT * FROM t WHERE a = $1 LIMIT 100 OFFSET 0', params: [5] }, + opts({ rowLimit: 1000 }), + vi.fn() + ) + const texts = calls.map((c) => c.text) + expect(texts).toContain('BEGIN') + expect(texts).toContain('SET TRANSACTION READ ONLY') + expect(texts).toContain('COMMIT') + // A read is bounded through a cursor; the bind params attach to the DECLARE. + const declare = calls.find((c) => c.text.startsWith('DECLARE')) + expect(declare?.text).toContain('SELECT * FROM t WHERE a = $1') + expect(declare?.values).toEqual([5]) + expect(texts.some((t) => t.startsWith('FETCH FORWARD'))).toBe(true) + expect(result.rows).toEqual([[1], [2]]) + expect(result.columns).toEqual([{ name: 'n' }]) + }) + + it('omits READ ONLY on a writable connection', async () => { + const { calls, client } = makeParamClient([], []) + await runPostgresExecute( + makePool(client), + 'c1', + { sql: 'UPDATE t SET a = $1 WHERE id = $2', params: [1, 2] }, + opts({ allowWrite: true }), + vi.fn() + ) + expect(calls.map((c) => c.text)).not.toContain('SET TRANSACTION READ ONLY') + expect(calls.map((c) => c.text)).toContain('COMMIT') + }) + + it('bounds a read through the cursor even if the statement lacks its own LIMIT', async () => { + const { client } = makeParamClient([[1], [2], [3]], [{ name: 'n' }]) + const result = await runPostgresExecute( + makePool(client), + 'c1', + { sql: 'SELECT * FROM t', params: [] }, + opts({ rowLimit: 2 }), + vi.fn() + ) + expect(result.rows).toEqual([[1], [2]]) + expect(result.truncated).toBe(true) + }) +}) + +describe('runPostgresBatch', () => { + it('applies every statement in one transaction and returns per-statement counts', async () => { + const { calls, client } = makeParamClient([], []) + const counts = await runPostgresBatch( + makePool(client), + 'c1', + [ + { sql: 'UPDATE t SET a = $1 WHERE id = $2', params: [1, 10] }, + { sql: 'DELETE FROM t WHERE id = $1', params: [11] } + ], + opts({ allowWrite: true }), + vi.fn() + ) + const texts = calls.map((c) => c.text) + expect(texts).toContain('BEGIN') + expect(texts).toContain('COMMIT') + expect(texts).not.toContain('ROLLBACK') + expect(counts).toEqual([1, 1]) + }) + + it('downgrades the batch transaction to READ ONLY when allowWrite is false', async () => { + const { calls, client } = makeParamClient([], []) + await runPostgresBatch( + makePool(client), + 'c1', + [{ sql: 'UPDATE t SET a = $1 WHERE id = $2', params: [1, 2] }], + opts({ allowWrite: false }), + vi.fn() + ) + expect(calls.map((c) => c.text)).toContain('SET TRANSACTION READ ONLY') + }) + + it('rolls back the whole batch and throws DbBatchError with the failing index', async () => { + const calls: string[] = [] + const client = { + release: vi.fn(), + query: vi.fn((arg: QueryArg) => { + const text = textOf(arg) + calls.push(text) + if (text.includes('pg_backend_pid')) { + return Promise.resolve({ rows: [{ pid: 7 }] }) + } + if (text.startsWith('DELETE')) { + return Promise.reject(Object.assign(new Error('fk'), { code: '23503' })) + } + return Promise.resolve({ rows: [], rowCount: 1 }) + }) + } + const err = await runPostgresBatch( + makePool(client), + 'c1', + [ + { sql: 'UPDATE t SET a = $1 WHERE id = $2', params: [1, 10] }, + { sql: 'DELETE FROM t WHERE id = $1', params: [11] } + ], + opts({ allowWrite: true }), + vi.fn() + ).catch((e) => e) + expect(err).toBeInstanceOf(DbBatchError) + expect((err as DbBatchError).failedIndex).toBe(1) + expect(calls).toContain('ROLLBACK') + expect(calls).not.toContain('COMMIT') + expect(client.release).toHaveBeenCalledTimes(1) + }) + + it('evicts the client when COMMIT fails', async () => { + const client = { + release: vi.fn(), + query: vi.fn((arg: QueryArg) => { + const text = textOf(arg) + if (text.includes('pg_backend_pid')) { + return Promise.resolve({ rows: [{ pid: 7 }] }) + } + // COMMIT and its recovery ROLLBACK both fail → unknown txn state → evict. + if (text === 'COMMIT' || text === 'ROLLBACK') { + return Promise.reject(new Error('boom')) + } + return Promise.resolve({ rows: [], rowCount: 1 }) + }) + } + await expect( + runPostgresBatch( + makePool(client), + 'c1', + [{ sql: 'UPDATE t SET a = 1', params: [] }], + opts({ allowWrite: true }), + vi.fn() + ) + ).rejects.toThrow() + expect(client.release).toHaveBeenCalledWith(true) + }) +}) + +describe('cancelPostgresBackend', () => { + it('issues pg_cancel_backend on a short-lived connection', async () => { + const query = vi.fn().mockResolvedValue({ rows: [] }) + const end = vi.fn().mockResolvedValue(undefined) + const connect = vi.fn().mockResolvedValue(undefined) + await cancelPostgresBackend({ connect, query, end } as never, 99) + expect(connect).toHaveBeenCalledTimes(1) + expect(query).toHaveBeenCalledWith('SELECT pg_cancel_backend($1)', [99]) + expect(end).toHaveBeenCalledTimes(1) + }) +}) diff --git a/src/main/database/postgres-query.ts b/src/main/database/postgres-query.ts new file mode 100644 index 00000000000..8537a200a60 --- /dev/null +++ b/src/main/database/postgres-query.ts @@ -0,0 +1,219 @@ +// Postgres query execution + cancellation, kept out of the driver file so the +// transaction / cursor / cancel logic stays focused and testable. +// +// Read-only enforcement (red-team F3): reads run inside `BEGIN` + +// `SET TRANSACTION READ ONLY`, so the database — not a keyword check — rejects +// single-statement writes and writing CTEs. Multi-statement input is rejected +// up-front for read-only connections by the manager (a simple query runs every +// statement, so it could otherwise flip the txn back to read-write first). +// A writable, non-cursorable statement runs in autocommit (no explicit BEGIN) +// so transaction-block-unsafe commands (VACUUM, CREATE DATABASE, REINDEX +// CONCURRENTLY) aren't rejected with "cannot run inside a transaction block". +// Result bounding (red-team F9): a SELECT runs through a server-side cursor, +// fetching only rowLimit+1 rows; the user's SQL is never rewritten (no appended +// LIMIT), so trailing semicolons / existing LIMIT / multi-statement stay intact. +// Cancellation (red-team F10): the backend PID is captured at start and cancelled +// from a separate short-lived connection. + +import { isCursorableRead } from '../../shared/sql-statement-classifier' +import { DbBatchError } from './db-driver' +import type { + DbStatement, + QueryHandle, + QueryOptions, + QueryResult +} from '../../shared/database-types' +import type { Client, Pool, PoolClient } from 'pg' + +const CURSOR_NAME = 'orca_query_cursor' + +type BoundedRows = { columns: { name: string }[]; rows: unknown[][]; rowCount: number; truncated: boolean } + +async function fetchBounded( + client: PoolClient, + sql: string, + rowLimit: number, + params?: unknown[] +): Promise { + if (isCursorableRead(sql)) { + // DECLARE the query as a cursor, then FETCH one more than the cap to detect + // overflow. Any bind params attach to the cursor's query via the extended + // protocol; the SQL text is embedded verbatim (the cursor query itself can't + // be otherwise parameterized) — safe: it's the user's own SQL on their own DB. + await client.query({ text: `DECLARE ${CURSOR_NAME} NO SCROLL CURSOR FOR ${sql}`, values: params }) + const fetched = await client.query({ + text: `FETCH FORWARD ${rowLimit + 1} FROM ${CURSOR_NAME}`, + rowMode: 'array' + }) + await client.query(`CLOSE ${CURSOR_NAME}`) + const truncated = fetched.rows.length > rowLimit + const rows = truncated ? fetched.rows.slice(0, rowLimit) : fetched.rows + return { + columns: (fetched.fields ?? []).map((f) => ({ name: f.name })), + rows, + rowCount: rows.length, + truncated + } + } + // Writes / non-cursorable statements: run directly (they return few/no rows). + const result = await client.query({ text: sql, values: params, rowMode: 'array' }) + const rows = (result.rows ?? []) as unknown[][] + return { + columns: (result.fields ?? []).map((f) => ({ name: f.name })), + rows, + rowCount: result.rowCount ?? rows.length, + truncated: false + } +} + +export async function runPostgresQuery( + pool: Pool, + connectionId: string, + sql: string, + opts: QueryOptions, + onStart: (handle: QueryHandle) => void +): Promise { + const client = await pool.connect() + // Evict the client from the pool if ROLLBACK fails (it may still hold an open + // transaction) rather than returning a poisoned connection. + let evict = false + try { + const pidResult = await client.query('SELECT pg_backend_pid() AS pid') + const backendPid = Number(pidResult.rows[0]?.pid) || null + onStart({ connectionId, backendPid }) + + await client.query(`SET statement_timeout = ${Math.trunc(opts.timeoutMs)}`) + + // A cursor must live inside a transaction, and a read-only connection must + // run inside a read-only transaction (the write guard). A writable, + // non-cursorable statement gets no explicit transaction so it may be a + // transaction-block-unsafe command. + if (opts.allowWrite && !isCursorableRead(sql)) { + const startedAt = Date.now() + const bounded = await fetchBounded(client, sql, opts.rowLimit) + return { ...bounded, durationMs: Date.now() - startedAt } + } + + await client.query('BEGIN') + if (!opts.allowWrite) { + await client.query('SET TRANSACTION READ ONLY') + } + const startedAt = Date.now() + try { + const bounded = await fetchBounded(client, sql, opts.rowLimit) + await client.query('COMMIT') + return { ...bounded, durationMs: Date.now() - startedAt } + } catch (err) { + await client.query('ROLLBACK').catch(() => { + evict = true + }) + throw err + } + } finally { + client.release(evict) + } +} + +// Runs one parameterized statement (Data-tab select/count/mutation or wrapped +// free-form re-query). Always inside a transaction — read-only when !allowWrite — +// so a value-bound read on a read-only connection is DB-enforced. The statement's +// own LIMIT bounds the page; rowLimit is a defensive server-side cap. +export async function runPostgresExecute( + pool: Pool, + connectionId: string, + statement: DbStatement, + opts: QueryOptions, + onStart: (handle: QueryHandle) => void +): Promise { + const client = await pool.connect() + // Evict the client from the pool if ROLLBACK fails (it may still hold an open + // transaction) rather than returning a poisoned connection. + let evict = false + try { + const pidResult = await client.query('SELECT pg_backend_pid() AS pid') + onStart({ connectionId, backendPid: Number(pidResult.rows[0]?.pid) || null }) + await client.query(`SET statement_timeout = ${Math.trunc(opts.timeoutMs)}`) + await client.query('BEGIN') + if (!opts.allowWrite) { + await client.query('SET TRANSACTION READ ONLY') + } + const startedAt = Date.now() + try { + // Bound the fetch (cursor for row-producing reads) so a statement that ever + // ships without its own LIMIT still can't materialize the whole result set. + const bounded = await fetchBounded(client, statement.sql, opts.rowLimit, statement.params) + await client.query('COMMIT') + return { ...bounded, durationMs: Date.now() - startedAt } + } catch (err) { + await client.query('ROLLBACK').catch(() => { + evict = true + }) + throw err + } + } finally { + client.release(evict) + } +} + +// Applies staged writes atomically: BEGIN, run each statement in order, COMMIT. +// Any failure rolls back the whole batch and throws DbBatchError(failedIndex). +export async function runPostgresBatch( + pool: Pool, + connectionId: string, + statements: DbStatement[], + opts: QueryOptions, + onStart: (handle: QueryHandle) => void +): Promise { + const client = await pool.connect() + // Evict the client from the pool if ROLLBACK fails (it may still hold an open + // transaction) rather than returning a poisoned connection. + let evict = false + try { + const pidResult = await client.query('SELECT pg_backend_pid() AS pid') + onStart({ connectionId, backendPid: Number(pidResult.rows[0]?.pid) || null }) + await client.query(`SET statement_timeout = ${Math.trunc(opts.timeoutMs)}`) + await client.query('BEGIN') + // Defense in depth: a read-only connection stays DB-enforced even if this + // write path is somehow reached (the manager already rejects !allowWrite). + if (!opts.allowWrite) { + await client.query('SET TRANSACTION READ ONLY') + } + const rowCounts: number[] = [] + for (let i = 0; i < statements.length; i++) { + try { + const result = await client.query({ + text: statements[i].sql, + values: statements[i].params + }) + rowCounts.push(result.rowCount ?? 0) + } catch (err) { + await client.query('ROLLBACK').catch(() => { + evict = true + }) + throw new DbBatchError(i, err) + } + } + // A failing COMMIT leaves the transaction in an unknown state — roll back and + // evict on failure rather than returning the client to the pool. + try { + await client.query('COMMIT') + } catch (err) { + await client.query('ROLLBACK').catch(() => { + evict = true + }) + throw err + } + return rowCounts + } finally { + client.release(evict) + } +} + +export async function cancelPostgresBackend(client: Client, backendPid: number): Promise { + await client.connect() + try { + await client.query('SELECT pg_cancel_backend($1)', [backendPid]) + } finally { + await client.end().catch(() => {}) + } +} diff --git a/src/main/index.ts b/src/main/index.ts index 8bc44bdc36d..275268e1ce9 100644 --- a/src/main/index.ts +++ b/src/main/index.ts @@ -23,6 +23,7 @@ import { killAllPty } from './ipc/pty' import { initDaemonPtyProvider, disconnectDaemon, shutdownDaemon } from './daemon/daemon-init' import { closeAllWatchers } from './ipc/filesystem-watcher' import { disposeWorktreeBaseDirectoryWatchers } from './ipc/worktree-base-directory-watcher' +import { dbConnectionManager } from './database/db-connection-manager' import { registerCoreHandlers } from './ipc/register-core-handlers' import { initObservability, shutdownObservability } from './observability' import { startSpan } from './observability/tracer' @@ -2045,6 +2046,11 @@ app.on('will-quit', (e) => { const emulatorShutdown = runtime?.getEmulatorBridge()?.destroyAllSessions() ?? Promise.resolve() serveSimStateWatcher.stop() killAllPty() + // Why (red-team F12): SSH's manager is never disposed on quit; the DB manager + // must be wired explicitly so held pools/sockets are torn down on exit. Awaited + // in the teardown chain below so pools close before the process exits + // (disconnectAll never rejects and is idempotent on the guarded second pass). + const dbShutdown = dbConnectionManager.disconnectAll() const watcherShutdown = shutdownWatchersOnce() store?.flush() @@ -2094,7 +2100,7 @@ app.on('will-quit', (e) => { // Why: normal quits preserve the detached daemon for warm reattach, but a // dev parent dying means the temp/dev profile has no owner left to reattach. const daemonTeardown = isDevParentShutdownRequested() ? shutdownDaemon() : disconnectDaemon() - Promise.allSettled([daemonTeardown, rpcStopAndClear, watcherShutdown, emulatorShutdown]) + Promise.allSettled([daemonTeardown, rpcStopAndClear, watcherShutdown, emulatorShutdown, dbShutdown]) .then(() => shutdownTelemetry()) .then(() => shutdownObservability()) .catch(() => { diff --git a/src/main/ipc/database.test.ts b/src/main/ipc/database.test.ts new file mode 100644 index 00000000000..4a4af1cce6b --- /dev/null +++ b/src/main/ipc/database.test.ts @@ -0,0 +1,900 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest' +import type { Store } from '../persistence' +import type { + DbConnection, + DbConnectionInput, + DbConnectionRuntimeState, + DbConnectionSummary, + DbConnectionUpdate +} from '../../shared/database-types' + +const { handleMock, removeHandlerMock } = vi.hoisted(() => ({ + handleMock: vi.fn(), + removeHandlerMock: vi.fn() +})) + +vi.mock('electron', () => ({ + ipcMain: { + handle: handleMock, + removeHandler: removeHandlerMock + }, + BrowserWindow: { + getAllWindows: () => [] + } +})) + +const { getDbEncryptionStatusMock, decryptDbSecretMock } = vi.hoisted(() => ({ + getDbEncryptionStatusMock: vi.fn(() => ({ + backend: 'mock-backend', + isStrong: true + })), + // Strip the at-rest tag so resolved configs carry the plaintext secret. + decryptDbSecretMock: vi.fn((stored: string) => + stored.replace(/^db\.(safeStorage|plaintext)\.v1:/, '') + ) +})) + +vi.mock('../database/db-credential-store', () => ({ + getDbEncryptionStatus: getDbEncryptionStatusMock, + decryptDbSecret: decryptDbSecretMock +})) + +// Manager double: the real manager would load pg/mysql2 and dial a live server. +const { managerMock } = vi.hoisted(() => ({ + managerMock: { + setStatusListener: vi.fn(), + test: vi.fn(async () => {}), + connect: vi.fn(async (cfg: { id: string }) => ({ id: cfg.id, status: 'connected' as const })), + disconnect: vi.fn(async () => {}), + getAllStatuses: vi.fn((): DbConnectionRuntimeState[] => []), + introspectSchemas: vi.fn(async () => ({ schemas: ['public'], truncated: false })), + introspectTables: vi.fn(async () => ({ tables: [{ name: 'users', kind: 'table' }], truncated: false })), + introspectColumns: vi.fn(async () => [ + { name: 'id', dataType: 'int', nullable: false, isPrimaryKey: true } + ]), + query: vi.fn(async () => ({ + columns: [{ name: 'n' }], + rows: [[1]], + rowCount: 1, + truncated: false, + durationMs: 3 + })), + cancelQuery: vi.fn(async () => {}), + execute: vi.fn(async () => ({ + columns: [{ name: 'n' }], + rows: [[1]], + rowCount: 1, + truncated: false, + durationMs: 2 + })), + executeBatch: vi.fn(async () => [1]) + } +})) + +vi.mock('../database/db-connection-manager', () => ({ dbConnectionManager: managerMock })) + +const { isTrustedUIRendererMock } = vi.hoisted(() => ({ + isTrustedUIRendererMock: vi.fn( + (sender: Record) => sender.isTrusted === true + ) +})) + +vi.mock('./ui', () => ({ + isTrustedUIRenderer: isTrustedUIRendererMock +})) + +import { registerDatabaseHandlers } from './database' +import { DbBatchError } from '../database/db-driver' + +function makeStore(overrides: Partial = {}): Store { + return { + getDbConnections: vi.fn(() => []), + getDbConnection: vi.fn(), + addDbConnection: vi.fn(), + updateDbConnection: vi.fn(), + removeDbConnection: vi.fn(), + ...overrides + } as unknown as Store +} + +function makeEvent(senderOverrides: Record = {}) { + return { + sender: { + id: 1, + ...senderOverrides + } + } +} + +function makeDbConnection(overrides: Partial = {}): DbConnection { + return { + id: 'conn-1', + name: 'test-db', + engine: 'postgres' as const, + host: 'localhost', + port: 5432, + database: 'testdb', + user: 'testuser', + password: 'db.safeStorage.v1:encrypted-secret', + ssl: 'verify-full' as const, + readOnly: false, + createdAt: Date.now(), + updatedAt: Date.now(), + ...overrides + } +} + +describe('registerDatabaseHandlers', () => { + beforeEach(() => { + vi.clearAllMocks() + handleMock.mockReset() + removeHandlerMock.mockReset() + }) + + describe('initialization', () => { + it('registers all database IPC channels', () => { + const store = makeStore() + + registerDatabaseHandlers(store) + + const expected = [ + 'database:list', + 'database:add', + 'database:update', + 'database:remove', + 'database:encryptionStatus', + 'database:test', + 'database:connect', + 'database:disconnect', + 'database:statuses', + 'database:introspect', + 'database:introspectSchemaTables', + 'database:introspectTableColumns', + 'database:query', + 'database:cancelQuery', + 'database:execute', + 'database:executeBatch' + ] + for (const channel of expected) { + expect(removeHandlerMock).toHaveBeenCalledWith(channel) + } + expect(handleMock).toHaveBeenCalledTimes(expected.length) + expect(handleMock.mock.calls.map(([channel]) => channel)).toEqual(expected) + }) + + it('can be called multiple times without error', () => { + const store = makeStore() + + expect(() => { + registerDatabaseHandlers(store) + registerDatabaseHandlers(store) + }).not.toThrow() + }) + }) + + describe('database:list', () => { + it('returns empty list when no connections', () => { + const store = makeStore({ + getDbConnections: vi.fn(() => []) + }) + registerDatabaseHandlers(store) + + const handler = handleMock.mock.calls.find(([ch]) => ch === 'database:list')?.[1] + const event = makeEvent({ isTrusted: true }) + + const result = handler?.(event) + + expect(result).toEqual([]) + }) + + it('returns connections as summaries without password field', () => { + const conn = makeDbConnection({ + id: 'id-1', + name: 'db1', + password: 'db.safeStorage.v1:secret' + }) + const store = makeStore({ + getDbConnections: vi.fn(() => [conn]) + }) + registerDatabaseHandlers(store) + + const handler = handleMock.mock.calls.find(([ch]) => ch === 'database:list')?.[1] + const event = makeEvent({ isTrusted: true }) + + const result = handler?.(event) as DbConnectionSummary[] + + expect(result).toHaveLength(1) + expect(result[0].id).toBe('id-1') + expect(result[0].name).toBe('db1') + expect(result[0].hasPassword).toBe(true) + expect('password' in result[0]).toBe(false) + }) + + it('sets hasPassword=false when connection has no password', () => { + const conn = makeDbConnection({ password: undefined }) + const store = makeStore({ + getDbConnections: vi.fn(() => [conn]) + }) + registerDatabaseHandlers(store) + + const handler = handleMock.mock.calls.find(([ch]) => ch === 'database:list')?.[1] + const event = makeEvent({ isTrusted: true }) + + const result = handler?.(event) as DbConnectionSummary[] + + expect(result[0].hasPassword).toBe(false) + }) + + it('rejects untrusted sender', () => { + const store = makeStore() + registerDatabaseHandlers(store) + + const handler = handleMock.mock.calls.find(([ch]) => ch === 'database:list')?.[1] + const event = makeEvent({ isTrusted: false }) + + expect(() => handler?.(event)).toThrow('untrusted_sender') + }) + }) + + describe('database:add', () => { + it('sanitizes input and adds connection', () => { + const addedConn = makeDbConnection({ id: 'new-id' }) + const store = makeStore({ + addDbConnection: vi.fn(() => addedConn) + }) + registerDatabaseHandlers(store) + + const handler = handleMock.mock.calls.find(([ch]) => ch === 'database:add')?.[1] + const event = makeEvent({ isTrusted: true }) + const input: DbConnectionInput = { + name: 'mydb', + engine: 'postgres', + host: 'db.example.com', + port: 5432, + database: 'mydb', + user: 'user', + password: 'secret' + } + + const result = handler?.(event, { input }) as DbConnectionSummary + + expect(store.addDbConnection).toHaveBeenCalled() + expect(result.id).toBe('new-id') + expect('password' in result).toBe(false) + expect(result.hasPassword).toBe(true) + }) + + it('coerces port to integer', () => { + const addedConn = makeDbConnection() + const store = makeStore({ + addDbConnection: vi.fn(() => addedConn) + }) + registerDatabaseHandlers(store) + + const handler = handleMock.mock.calls.find(([ch]) => ch === 'database:add')?.[1] + const event = makeEvent({ isTrusted: true }) + const input = { + name: 'db', + engine: 'postgres' as const, + host: 'localhost', + port: '5432' as unknown as number, + database: 'db', + user: 'user' + } + + handler?.(event, { input }) + + const callArg = vi.mocked(store.addDbConnection).mock.calls[0][0] + expect(callArg.port).toBe(5432) + expect(typeof callArg.port).toBe('number') + }) + + it('rejects invalid port', () => { + const store = makeStore() + registerDatabaseHandlers(store) + + const handler = handleMock.mock.calls.find(([ch]) => ch === 'database:add')?.[1] + const event = makeEvent({ isTrusted: true }) + const input = { + name: 'db', + engine: 'postgres' as const, + host: 'localhost', + port: 99999, + database: 'db', + user: 'user' + } + + expect(() => handler?.(event, { input })).toThrow('invalid_port') + }) + + it('rejects invalid engine', () => { + const store = makeStore() + registerDatabaseHandlers(store) + + const handler = handleMock.mock.calls.find(([ch]) => ch === 'database:add')?.[1] + const event = makeEvent({ isTrusted: true }) + const input = { + name: 'db', + engine: 'sqlite' as unknown as DbConnectionInput['engine'], + host: 'localhost', + port: 5432, + database: 'db', + user: 'user' + } + + expect(() => handler?.(event, { input })).toThrow('invalid_engine') + }) + + it('rejects untrusted sender', () => { + const store = makeStore() + registerDatabaseHandlers(store) + + const handler = handleMock.mock.calls.find(([ch]) => ch === 'database:add')?.[1] + const event = makeEvent({ isTrusted: false }) + const input = makeDbConnection() + + expect(() => handler?.(event, { input })).toThrow('untrusted_sender') + expect(store.addDbConnection).not.toHaveBeenCalled() + }) + }) + + describe('database:update', () => { + it('returns password-stripped summary on success', () => { + const updated = makeDbConnection({ id: 'id-1', name: 'updated' }) + const store = makeStore({ + updateDbConnection: vi.fn(() => updated) + }) + registerDatabaseHandlers(store) + + const handler = handleMock.mock.calls.find(([ch]) => ch === 'database:update')?.[1] + const event = makeEvent({ isTrusted: true }) + const updates: DbConnectionUpdate = { name: 'updated' } + + const result = handler?.(event, { id: 'id-1', updates }) as DbConnectionSummary + + expect(result.name).toBe('updated') + expect('password' in result).toBe(false) + expect(result.hasPassword).toBe(true) + }) + + it('returns null if connection not found', () => { + const store = makeStore({ + updateDbConnection: vi.fn(() => null) + }) + registerDatabaseHandlers(store) + + const handler = handleMock.mock.calls.find(([ch]) => ch === 'database:update')?.[1] + const event = makeEvent({ isTrusted: true }) + + const result = handler?.(event, { id: 'nonexistent', updates: {} }) + + expect(result).toBeNull() + }) + + it('sanitizes and coerces update fields before persisting', () => { + const updated = makeDbConnection({ id: 'id-1' }) + const store = makeStore({ updateDbConnection: vi.fn(() => updated) }) + registerDatabaseHandlers(store) + + const handler = handleMock.mock.calls.find(([ch]) => ch === 'database:update')?.[1] + const event = makeEvent({ isTrusted: true }) + handler?.(event, { + id: 'id-1', + updates: { name: ' spaced ', port: '6543' as unknown as number } + }) + + const arg = vi.mocked(store.updateDbConnection).mock.calls[0][1] + expect(arg.name).toBe('spaced') + expect(arg.port).toBe(6543) + expect(typeof arg.port).toBe('number') + // Absent fields must not be injected into the persisted record. + expect('engine' in arg).toBe(false) + expect('host' in arg).toBe(false) + }) + + it('rejects an invalid engine in an update', () => { + const store = makeStore() + registerDatabaseHandlers(store) + + const handler = handleMock.mock.calls.find(([ch]) => ch === 'database:update')?.[1] + const event = makeEvent({ isTrusted: true }) + + expect(() => + handler?.(event, { + id: 'id-1', + updates: { engine: 'sqlite' as unknown as DbConnectionInput['engine'] } + }) + ).toThrow('invalid_engine') + expect(store.updateDbConnection).not.toHaveBeenCalled() + }) + + it('rejects an invalid port in an update', () => { + const store = makeStore() + registerDatabaseHandlers(store) + + const handler = handleMock.mock.calls.find(([ch]) => ch === 'database:update')?.[1] + const event = makeEvent({ isTrusted: true }) + + expect(() => handler?.(event, { id: 'id-1', updates: { port: 0 } })).toThrow('invalid_port') + expect(store.updateDbConnection).not.toHaveBeenCalled() + }) + + it('rejects untrusted sender', () => { + const store = makeStore() + registerDatabaseHandlers(store) + + const handler = handleMock.mock.calls.find(([ch]) => ch === 'database:update')?.[1] + const event = makeEvent({ isTrusted: false }) + + expect(() => handler?.(event, { id: 'id-1', updates: {} })).toThrow('untrusted_sender') + expect(store.updateDbConnection).not.toHaveBeenCalled() + }) + }) + + describe('database:remove', () => { + it('calls store.removeDbConnection and returns void', () => { + const store = makeStore({ + removeDbConnection: vi.fn() + }) + registerDatabaseHandlers(store) + + const handler = handleMock.mock.calls.find(([ch]) => ch === 'database:remove')?.[1] + const event = makeEvent({ isTrusted: true }) + + const result = handler?.(event, { id: 'id-1' }) + + expect(store.removeDbConnection).toHaveBeenCalledWith('id-1') + expect(result).toBeUndefined() + }) + + it('rejects untrusted sender', () => { + const store = makeStore() + registerDatabaseHandlers(store) + + const handler = handleMock.mock.calls.find(([ch]) => ch === 'database:remove')?.[1] + const event = makeEvent({ isTrusted: false }) + + expect(() => handler?.(event, { id: 'id-1' })).toThrow('untrusted_sender') + expect(store.removeDbConnection).not.toHaveBeenCalled() + }) + }) + + describe('database:encryptionStatus', () => { + it('returns encryption status from credential store', () => { + const store = makeStore() + getDbEncryptionStatusMock.mockReturnValue({ + backend: 'dpapi', + isStrong: true + }) + registerDatabaseHandlers(store) + + const handler = handleMock.mock.calls.find(([ch]) => ch === 'database:encryptionStatus')?.[1] + const event = makeEvent({ isTrusted: true }) + + const result = handler?.(event) + + expect(result).toEqual({ backend: 'dpapi', isStrong: true }) + }) + + it('rejects untrusted sender', () => { + const store = makeStore() + registerDatabaseHandlers(store) + + const handler = handleMock.mock.calls.find(([ch]) => ch === 'database:encryptionStatus')?.[1] + const event = makeEvent({ isTrusted: false }) + + expect(() => handler?.(event)).toThrow('untrusted_sender') + }) + }) + + describe('untrusted sender isolation', () => { + it('prevents untrusted sender from mutating store', () => { + const store = makeStore({ + addDbConnection: vi.fn(), + updateDbConnection: vi.fn(), + removeDbConnection: vi.fn() + }) + registerDatabaseHandlers(store) + + const addHandler = handleMock.mock.calls.find(([ch]) => ch === 'database:add')?.[1] + const updateHandler = handleMock.mock.calls.find(([ch]) => ch === 'database:update')?.[1] + const removeHandler = handleMock.mock.calls.find(([ch]) => ch === 'database:remove')?.[1] + const event = makeEvent({ isTrusted: false }) + + expect(() => addHandler?.(event, { input: {} })).toThrow('untrusted_sender') + expect(() => updateHandler?.(event, { id: 'x', updates: {} })).toThrow('untrusted_sender') + expect(() => removeHandler?.(event, { id: 'x' })).toThrow('untrusted_sender') + + expect(store.addDbConnection).not.toHaveBeenCalled() + expect(store.updateDbConnection).not.toHaveBeenCalled() + expect(store.removeDbConnection).not.toHaveBeenCalled() + }) + + it('allows trusted sender to mutate store', () => { + const addedConn = makeDbConnection() + const store = makeStore({ + addDbConnection: vi.fn(() => addedConn), + removeDbConnection: vi.fn() + }) + registerDatabaseHandlers(store) + + const addHandler = handleMock.mock.calls.find(([ch]) => ch === 'database:add')?.[1] + const removeHandler = handleMock.mock.calls.find(([ch]) => ch === 'database:remove')?.[1] + const event = makeEvent({ isTrusted: true }) + const input = makeDbConnection() + + addHandler?.(event, { input }) + removeHandler?.(event, { id: 'id-1' }) + + expect(store.addDbConnection).toHaveBeenCalled() + expect(store.removeDbConnection).toHaveBeenCalledWith('id-1') + }) + }) + + describe('password stripping', () => { + it('strips password from all summary returns', () => { + const connWithPassword = makeDbConnection({ + password: 'db.safeStorage.v1:secret' + }) + const connWithoutPassword = makeDbConnection({ + password: undefined + }) + + const store = makeStore({ + getDbConnections: vi.fn(() => [connWithPassword, connWithoutPassword]), + addDbConnection: vi.fn(() => connWithPassword), + updateDbConnection: vi.fn(() => connWithPassword) + }) + registerDatabaseHandlers(store) + + const listHandler = handleMock.mock.calls.find(([ch]) => ch === 'database:list')?.[1] + const addHandler = handleMock.mock.calls.find(([ch]) => ch === 'database:add')?.[1] + const updateHandler = handleMock.mock.calls.find(([ch]) => ch === 'database:update')?.[1] + const event = makeEvent({ isTrusted: true }) + const validInput: DbConnectionInput = { + name: 'db', + engine: 'postgres', + host: 'localhost', + port: 5432, + database: 'db', + user: 'user' + } + + const listResult = listHandler?.(event) as DbConnectionSummary[] + const addResult = addHandler?.(event, { input: validInput }) as DbConnectionSummary + const updateResult = updateHandler?.(event, { id: 'x', updates: {} }) as DbConnectionSummary + + expect(listResult.every((r) => !('password' in r))).toBe(true) + expect(!('password' in addResult)).toBe(true) + expect(!('password' in updateResult)).toBe(true) + + expect(listResult[0].hasPassword).toBe(true) + expect(listResult[1].hasPassword).toBe(false) + expect(addResult.hasPassword).toBe(true) + expect(updateResult.hasPassword).toBe(true) + }) + }) + + describe('lifecycle handlers (test/connect/disconnect/statuses)', () => { + const trusted = makeEvent({ isTrusted: true }) + + function getHandler(channel: string) { + return handleMock.mock.calls.find(([ch]) => ch === channel)?.[1] + } + + const formInput: DbConnectionInput = { + name: 'db', + engine: 'postgres', + host: 'db.example.com', + port: 5432, + database: 'app', + user: 'admin', + password: 'typed-pw' + } + + it('database:test returns ok and resolves smart-by-host SSL for the config', async () => { + registerDatabaseHandlers(makeStore()) + const result = await getHandler('database:test')?.(trusted, { input: formInput }) + expect(result).toEqual({ ok: true }) + expect(managerMock.test).toHaveBeenCalledWith( + expect.objectContaining({ id: 'db-test', ssl: 'verify-full', password: 'typed-pw' }) + ) + }) + + it('database:test returns a redacted error and never leaks the DSN/password', async () => { + registerDatabaseHandlers(makeStore()) + managerMock.test.mockRejectedValueOnce( + Object.assign(new Error('auth failed postgres://admin:s3cr3t@db:5432/app'), { + code: '28P01' + }) + ) + const result = await getHandler('database:test')?.(trusted, { input: formInput }) + expect(result.ok).toBe(false) + expect(result.error.code).toBe('auth_failed') + const serialized = JSON.stringify(result) + expect(serialized).not.toContain('s3cr3t') + expect(serialized).not.toContain('postgres://') + }) + + it('database:connect decrypts the stored secret and returns connected state', async () => { + const store = makeStore({ + getDbConnection: vi.fn(() => + makeDbConnection({ id: 'conn-1', password: 'db.safeStorage.v1:stored-pw' }) + ) + }) + registerDatabaseHandlers(store) + const result = await getHandler('database:connect')?.(trusted, { id: 'conn-1' }) + expect(result).toEqual({ id: 'conn-1', status: 'connected' }) + expect(decryptDbSecretMock).toHaveBeenCalledWith('db.safeStorage.v1:stored-pw') + expect(managerMock.connect).toHaveBeenCalledWith( + expect.objectContaining({ id: 'conn-1', password: 'stored-pw' }) + ) + }) + + it('database:connect returns a redacted error state when the dial fails', async () => { + const store = makeStore({ + getDbConnection: vi.fn(() => makeDbConnection({ id: 'conn-1' })) + }) + registerDatabaseHandlers(store) + managerMock.connect.mockRejectedValueOnce( + Object.assign(new Error('nope postgres://admin:s3cr3t@db'), { code: 'ECONNREFUSED' }) + ) + const result = await getHandler('database:connect')?.(trusted, { id: 'conn-1' }) + expect(result.status).toBe('error') + expect(result.error.code).toBe('connection_refused') + expect(JSON.stringify(result)).not.toContain('s3cr3t') + }) + + it('database:disconnect delegates to the manager', async () => { + registerDatabaseHandlers(makeStore()) + await getHandler('database:disconnect')?.(trusted, { id: 'conn-1' }) + expect(managerMock.disconnect).toHaveBeenCalledWith('conn-1') + }) + + it('database:statuses returns the manager snapshot', async () => { + managerMock.getAllStatuses.mockReturnValueOnce([{ id: 'conn-1', status: 'connected' }]) + registerDatabaseHandlers(makeStore()) + const result = await getHandler('database:statuses')?.(trusted) + expect(result).toEqual([{ id: 'conn-1', status: 'connected' }]) + }) + + it('rejects untrusted senders on every lifecycle channel', async () => { + registerDatabaseHandlers(makeStore()) + const untrusted = makeEvent({ isTrusted: false }) + await expect(getHandler('database:test')?.(untrusted, { input: formInput })).rejects.toThrow( + 'untrusted_sender' + ) + await expect(getHandler('database:connect')?.(untrusted, { id: 'x' })).rejects.toThrow( + 'untrusted_sender' + ) + await expect(getHandler('database:disconnect')?.(untrusted, { id: 'x' })).rejects.toThrow( + 'untrusted_sender' + ) + expect(() => getHandler('database:statuses')?.(untrusted)).toThrow('untrusted_sender') + expect(managerMock.test).not.toHaveBeenCalled() + expect(managerMock.connect).not.toHaveBeenCalled() + expect(managerMock.disconnect).not.toHaveBeenCalled() + }) + }) + + describe('introspection handlers', () => { + const trusted = makeEvent({ isTrusted: true }) + function getHandler(channel: string) { + return handleMock.mock.calls.find(([ch]) => ch === channel)?.[1] + } + + it('database:introspect returns the schema tree', async () => { + registerDatabaseHandlers(makeStore()) + const result = await getHandler('database:introspect')?.(trusted, { id: 'c1' }) + expect(result).toEqual({ ok: true, tree: { schemas: ['public'], truncated: false } }) + expect(managerMock.introspectSchemas).toHaveBeenCalledWith('c1') + }) + + it('database:introspectSchemaTables returns the table list', async () => { + registerDatabaseHandlers(makeStore()) + const result = await getHandler('database:introspectSchemaTables')?.(trusted, { + id: 'c1', + schema: 'public' + }) + expect(result.ok).toBe(true) + expect(managerMock.introspectTables).toHaveBeenCalledWith('c1', 'public') + }) + + it('database:introspectTableColumns returns columns', async () => { + registerDatabaseHandlers(makeStore()) + const result = await getHandler('database:introspectTableColumns')?.(trusted, { + id: 'c1', + ref: { schema: 'public', table: 'users' } + }) + expect(result.ok).toBe(true) + expect(managerMock.introspectColumns).toHaveBeenCalledWith('c1', { + schema: 'public', + table: 'users' + }) + }) + + it('returns a redacted error (never a raw throw) when introspection fails', async () => { + registerDatabaseHandlers(makeStore()) + managerMock.introspectSchemas.mockRejectedValueOnce( + Object.assign(new Error('boom postgres://admin:s3cr3t@db'), { code: 'ECONNRESET' }) + ) + const result = await getHandler('database:introspect')?.(trusted, { id: 'c1' }) + expect(result.ok).toBe(false) + expect(result.error.code).toBe('connection_refused') + expect(JSON.stringify(result)).not.toContain('s3cr3t') + }) + + it('rejects untrusted senders on introspection channels', async () => { + registerDatabaseHandlers(makeStore()) + const untrusted = makeEvent({ isTrusted: false }) + await expect(getHandler('database:introspect')?.(untrusted, { id: 'c1' })).rejects.toThrow( + 'untrusted_sender' + ) + expect(managerMock.introspectSchemas).not.toHaveBeenCalled() + }) + }) + + describe('query handlers', () => { + const trusted = makeEvent({ isTrusted: true }) + function getHandler(channel: string) { + return handleMock.mock.calls.find(([ch]) => ch === channel)?.[1] + } + + it('derives allowWrite=false from a read-only stored connection', async () => { + const store = makeStore({ + getDbConnection: vi.fn(() => makeDbConnection({ id: 'c1', readOnly: true })) + }) + registerDatabaseHandlers(store) + await getHandler('database:query')?.(trusted, { id: 'c1', sql: 'SELECT 1' }) + expect(managerMock.query).toHaveBeenCalledWith( + 'c1', + 'SELECT 1', + expect.objectContaining({ allowWrite: false }) + ) + }) + + it('derives allowWrite=true from a writable stored connection', async () => { + const store = makeStore({ + getDbConnection: vi.fn(() => makeDbConnection({ id: 'c1', readOnly: false })) + }) + registerDatabaseHandlers(store) + await getHandler('database:query')?.(trusted, { id: 'c1', sql: 'SELECT 1' }) + expect(managerMock.query).toHaveBeenCalledWith( + 'c1', + 'SELECT 1', + expect.objectContaining({ allowWrite: true }) + ) + }) + + it('ignores any allowWrite the renderer tries to send (server uses readOnly)', async () => { + const store = makeStore({ + getDbConnection: vi.fn(() => makeDbConnection({ id: 'c1', readOnly: true })) + }) + registerDatabaseHandlers(store) + // Renderer attempts to force writes; handler must not honor it. + await getHandler('database:query')?.(trusted, { + id: 'c1', + sql: 'DROP TABLE x', + allowWrite: true + }) + expect(managerMock.query).toHaveBeenCalledWith( + 'c1', + 'DROP TABLE x', + expect.objectContaining({ allowWrite: false }) + ) + }) + + it('treats a missing connection as read-only', async () => { + const store = makeStore({ getDbConnection: vi.fn(() => undefined) }) + registerDatabaseHandlers(store) + await getHandler('database:query')?.(trusted, { id: 'missing', sql: 'SELECT 1' }) + expect(managerMock.query).toHaveBeenCalledWith( + 'missing', + 'SELECT 1', + expect.objectContaining({ allowWrite: false }) + ) + }) + + it('returns a redacted error when the query fails', async () => { + const store = makeStore({ + getDbConnection: vi.fn(() => makeDbConnection({ id: 'c1' })) + }) + registerDatabaseHandlers(store) + managerMock.query.mockRejectedValueOnce( + Object.assign(new Error('boom postgres://admin:s3cr3t@db'), { code: 'ECONNRESET' }) + ) + const result = await getHandler('database:query')?.(trusted, { id: 'c1', sql: 'SELECT 1' }) + expect(result.ok).toBe(false) + expect(JSON.stringify(result)).not.toContain('s3cr3t') + }) + + it('cancelQuery delegates to the manager and rejects untrusted senders', async () => { + registerDatabaseHandlers(makeStore()) + await getHandler('database:cancelQuery')?.(trusted, { id: 'c1' }) + expect(managerMock.cancelQuery).toHaveBeenCalledWith('c1') + + const untrusted = makeEvent({ isTrusted: false }) + await expect(getHandler('database:query')?.(untrusted, { id: 'c1', sql: 'x' })).rejects.toThrow( + 'untrusted_sender' + ) + }) + }) + + describe('execute + executeBatch handlers', () => { + const trusted = makeEvent({ isTrusted: true }) + const statement = { sql: 'SELECT * FROM t LIMIT 100 OFFSET 0', params: [] } + function getHandler(channel: string) { + return handleMock.mock.calls.find(([ch]) => ch === channel)?.[1] + } + + it('execute derives allowWrite from readOnly and forwards the statement', async () => { + const store = makeStore({ + getDbConnection: vi.fn(() => makeDbConnection({ id: 'c1', readOnly: true })) + }) + registerDatabaseHandlers(store) + const result = await getHandler('database:execute')?.(trusted, { id: 'c1', statement }) + expect(result.ok).toBe(true) + expect(managerMock.execute).toHaveBeenCalledWith( + 'c1', + statement, + expect.objectContaining({ allowWrite: false }) + ) + }) + + it('execute returns a redacted error and never leaks the DSN/password', async () => { + const store = makeStore({ getDbConnection: vi.fn(() => makeDbConnection({ id: 'c1' })) }) + registerDatabaseHandlers(store) + managerMock.execute.mockRejectedValueOnce( + Object.assign(new Error('boom postgres://admin:s3cr3t@db'), { code: 'ECONNRESET' }) + ) + const result = await getHandler('database:execute')?.(trusted, { id: 'c1', statement }) + expect(result.ok).toBe(false) + expect(JSON.stringify(result)).not.toContain('s3cr3t') + }) + + it('executeBatch derives allowWrite=true on a writable connection and returns rowCounts', async () => { + const store = makeStore({ + getDbConnection: vi.fn(() => makeDbConnection({ id: 'c1', readOnly: false })) + }) + registerDatabaseHandlers(store) + const result = await getHandler('database:executeBatch')?.(trusted, { + id: 'c1', + statements: [statement] + }) + expect(result).toEqual({ ok: true, rowCounts: [1] }) + expect(managerMock.executeBatch).toHaveBeenCalledWith( + 'c1', + [statement], + expect.objectContaining({ allowWrite: true }) + ) + }) + + it('executeBatch surfaces failedIndex and a redacted error on a DbBatchError', async () => { + const store = makeStore({ getDbConnection: vi.fn(() => makeDbConnection({ id: 'c1' })) }) + registerDatabaseHandlers(store) + managerMock.executeBatch.mockRejectedValueOnce( + new DbBatchError(1, Object.assign(new Error('fk postgres://admin:s3cr3t@db'), { code: '23503' })) + ) + const result = await getHandler('database:executeBatch')?.(trusted, { + id: 'c1', + statements: [statement, statement] + }) + expect(result.ok).toBe(false) + expect(result.failedIndex).toBe(1) + expect(JSON.stringify(result)).not.toContain('s3cr3t') + }) + + it('rejects untrusted senders on both channels', async () => { + registerDatabaseHandlers(makeStore()) + const untrusted = makeEvent({ isTrusted: false }) + await expect( + getHandler('database:execute')?.(untrusted, { id: 'c1', statement }) + ).rejects.toThrow('untrusted_sender') + await expect( + getHandler('database:executeBatch')?.(untrusted, { id: 'c1', statements: [statement] }) + ).rejects.toThrow('untrusted_sender') + expect(managerMock.execute).not.toHaveBeenCalled() + expect(managerMock.executeBatch).not.toHaveBeenCalled() + }) + }) +}) diff --git a/src/main/ipc/database.ts b/src/main/ipc/database.ts new file mode 100644 index 00000000000..7501609bd80 --- /dev/null +++ b/src/main/ipc/database.ts @@ -0,0 +1,369 @@ +import { BrowserWindow, ipcMain, type WebContents } from 'electron' +import type { Store } from '../persistence' +import { decryptDbSecret, getDbEncryptionStatus } from '../database/db-credential-store' +import { dbConnectionManager } from '../database/db-connection-manager' +import { + DB_MAX_ROWS, + DB_STATEMENT_TIMEOUT_MS, + DbBatchError, + normalizeDbError, + resolveDbConfig, + type ResolvedDbConfig +} from '../database/db-driver' +import { isTrustedUIRenderer } from './ui' +import type { + DbBatchResult, + DbColumnListResult, + DbConnection, + DbConnectionInput, + DbConnectionRuntimeState, + DbConnectionSummary, + DbConnectionUpdate, + DbEncryptionStatus, + DbEngine, + DbExecuteResult, + DbIntrospectResult, + DbQueryResult, + DbStatement, + DbTableListResult, + DbTableRef, + DbTestResult +} from '../../shared/database-types' + +// Phase 2 surface: connection CRUD + encryption posture. Phase 3 adds the live +// lifecycle (test/connect/disconnect/statuses); introspect/query come later. +const DATABASE_IPC_CHANNELS = [ + 'database:list', + 'database:add', + 'database:update', + 'database:remove', + 'database:encryptionStatus', + 'database:test', + 'database:connect', + 'database:disconnect', + 'database:statuses', + 'database:introspect', + 'database:introspectSchemaTables', + 'database:introspectTableColumns', + 'database:query', + 'database:cancelQuery', + 'database:execute', + 'database:executeBatch' +] as const + +const VALID_ENGINES = new Set(['postgres', 'mysql']) + +// Why: never hand the stored secret back to the renderer — it only needs to know +// whether a password exists. +function toSummary(connection: DbConnection): DbConnectionSummary { + const { password: _password, ...rest } = connection + return { ...rest, hasPassword: !!connection.password } +} + +// Why: the sender is already gated to the trusted UI renderer, but coerce the +// core fields so a malformed payload can't poison persisted state. +function sanitizeInput(input: DbConnectionInput): DbConnectionInput { + if (!VALID_ENGINES.has(input.engine)) { + throw new Error('invalid_engine') + } + const port = Number(input.port) + if (!Number.isInteger(port) || port <= 0 || port > 65535) { + throw new Error('invalid_port') + } + return { + name: String(input.name ?? '').trim(), + engine: input.engine, + host: String(input.host ?? '').trim(), + port, + database: String(input.database ?? '').trim(), + user: String(input.user ?? '').trim(), + password: input.password ? String(input.password) : undefined, + ssl: input.ssl, + readOnly: input.readOnly ?? false, + sshTunnel: input.sshTunnel + } +} + +// Why: like sanitizeInput, but for a partial update — only coerce/validate the +// fields actually present so a malformed update can't poison the stored record +// (database:update forwards renderer fields straight into persistence otherwise). +function sanitizeUpdate(updates: DbConnectionUpdate): DbConnectionUpdate { + const sanitized: DbConnectionUpdate = {} + if (updates.name !== undefined) sanitized.name = String(updates.name).trim() + if (updates.host !== undefined) sanitized.host = String(updates.host).trim() + if (updates.database !== undefined) sanitized.database = String(updates.database).trim() + if (updates.user !== undefined) sanitized.user = String(updates.user).trim() + if (updates.password !== undefined) { + sanitized.password = updates.password ? String(updates.password) : undefined + } + if (updates.engine !== undefined) { + if (!VALID_ENGINES.has(updates.engine)) { + throw new Error('invalid_engine') + } + sanitized.engine = updates.engine + } + if (updates.port !== undefined) { + const port = Number(updates.port) + if (!Number.isInteger(port) || port <= 0 || port > 65535) { + throw new Error('invalid_port') + } + sanitized.port = port + } + if (updates.readOnly !== undefined) sanitized.readOnly = updates.readOnly === true + if (updates.ssl !== undefined) sanitized.ssl = updates.ssl + if (updates.sshTunnel !== undefined) sanitized.sshTunnel = updates.sshTunnel + return sanitized +} + +// Why: decrypt is strict/fail-closed — a keychain-changed secret throws here +// rather than handing a bogus credential to a driver. Resolves the saved record +// into a dial-ready config (password decrypted at point-of-use; SSL smart-by-host). +function resolveSavedConfig(store: Store, id: string): ResolvedDbConfig { + const connection = store.getDbConnection(id) + if (!connection) { + throw new Error('db_connection_not_found') + } + const password = connection.password ? decryptDbSecret(connection.password) : undefined + return resolveDbConfig(connection, password) +} + +// Build a config for a one-shot Test from the form: a typed password wins; an +// empty field on an existing connection falls back to its stored secret. +function resolveTestConfig( + store: Store, + input: DbConnectionInput, + id: string | undefined +): ResolvedDbConfig { + const existing = id ? store.getDbConnection(id) : undefined + let password = input.password + if (!password && existing?.password) { + password = decryptDbSecret(existing.password) + } + return resolveDbConfig( + { + id: id ?? 'db-test', + name: input.name, + engine: input.engine, + host: input.host, + port: input.port, + database: input.database, + user: input.user, + ssl: input.ssl, + readOnly: input.readOnly ?? false, + createdAt: 0, + updatedAt: 0 + }, + password + ) +} + +export function registerDatabaseHandlers(store: Store): void { + // Why: on macOS re-activation this can be called again; ipcMain.handle throws + // on a duplicate channel, so clear any prior handlers first (mirrors SSH). + for (const channel of DATABASE_IPC_CHANNELS) { + ipcMain.removeHandler(channel) + } + + // Why (mirrors ui.ts onUIChanged): the DB view is desktop-main-window only, so + // broadcast every runtime status change to all live windows — the renderer + // re-hydrates from these instead of polling. + dbConnectionManager.setStatusListener((state) => { + for (const window of BrowserWindow.getAllWindows()) { + if (!window.isDestroyed()) { + window.webContents.send('database:status-changed', state) + } + } + }) + + // Why (red-team F15): the renderer is sandboxed but the channel is reachable; + // reject any sender that is not the trusted main-window UI renderer. + const requireTrusted = (sender: WebContents): void => { + if (!isTrustedUIRenderer(sender)) { + throw new Error('untrusted_sender') + } + } + + ipcMain.handle('database:list', (event): DbConnectionSummary[] => { + requireTrusted(event.sender) + return store.getDbConnections().map(toSummary) + }) + + ipcMain.handle( + 'database:add', + (event, args: { input: DbConnectionInput }): DbConnectionSummary => { + requireTrusted(event.sender) + return toSummary(store.addDbConnection(sanitizeInput(args.input))) + } + ) + + ipcMain.handle( + 'database:update', + (event, args: { id: string; updates: DbConnectionUpdate }): DbConnectionSummary | null => { + requireTrusted(event.sender) + const updated = store.updateDbConnection(args.id, sanitizeUpdate(args.updates)) + return updated ? toSummary(updated) : null + } + ) + + ipcMain.handle('database:remove', (event, args: { id: string }): void => { + requireTrusted(event.sender) + store.removeDbConnection(args.id) + }) + + ipcMain.handle('database:encryptionStatus', (event): DbEncryptionStatus => { + requireTrusted(event.sender) + return getDbEncryptionStatus() + }) + + // Why: Test returns its result (never throws a raw rejection) so a failure + // carries only the redacted { code, safeMessage } — the raw driver error + // embeds the DSN/password. + ipcMain.handle( + 'database:test', + async (event, args: { input: DbConnectionInput; id?: string }): Promise => { + requireTrusted(event.sender) + try { + await dbConnectionManager.test(resolveTestConfig(store, args.input, args.id)) + return { ok: true } + } catch (err) { + return { ok: false, error: normalizeDbError(err) } + } + } + ) + + ipcMain.handle( + 'database:connect', + async (event, args: { id: string }): Promise => { + requireTrusted(event.sender) + try { + return await dbConnectionManager.connect(resolveSavedConfig(store, args.id)) + } catch (err) { + // Why: the manager already pushed an `error` status via broadcast; return + // the redacted state so the caller never sees the raw rejection. + return { id: args.id, status: 'error', error: normalizeDbError(err) } + } + } + ) + + ipcMain.handle('database:disconnect', async (event, args: { id: string }): Promise => { + requireTrusted(event.sender) + await dbConnectionManager.disconnect(args.id) + }) + + ipcMain.handle('database:statuses', (event): DbConnectionRuntimeState[] => { + requireTrusted(event.sender) + return dbConnectionManager.getAllStatuses() + }) + + // Why: introspection results are returned (never thrown) so a bad catalog read + // surfaces a redacted error inline instead of crashing the schema tree. + ipcMain.handle( + 'database:introspect', + async (event, args: { id: string }): Promise => { + requireTrusted(event.sender) + try { + return { ok: true, tree: await dbConnectionManager.introspectSchemas(args.id) } + } catch (err) { + return { ok: false, error: normalizeDbError(err) } + } + } + ) + + ipcMain.handle( + 'database:introspectSchemaTables', + async (event, args: { id: string; schema: string }): Promise => { + requireTrusted(event.sender) + try { + return { ok: true, list: await dbConnectionManager.introspectTables(args.id, args.schema) } + } catch (err) { + return { ok: false, error: normalizeDbError(err) } + } + } + ) + + ipcMain.handle( + 'database:introspectTableColumns', + async (event, args: { id: string; ref: DbTableRef }): Promise => { + requireTrusted(event.sender) + try { + return { ok: true, columns: await dbConnectionManager.introspectColumns(args.id, args.ref) } + } catch (err) { + return { ok: false, error: normalizeDbError(err) } + } + } + ) + + ipcMain.handle( + 'database:query', + async (event, args: { id: string; sql: string }): Promise => { + requireTrusted(event.sender) + try { + // Why (red-team F3): allowWrite is derived from the stored connection's + // readOnly server-side — never trusted from the renderer. A missing + // connection is treated as read-only (safe default). + const connection = store.getDbConnection(args.id) + const allowWrite = connection ? !connection.readOnly : false + const result = await dbConnectionManager.query(args.id, args.sql, { + rowLimit: DB_MAX_ROWS, + timeoutMs: DB_STATEMENT_TIMEOUT_MS, + allowWrite + }) + return { ok: true, result } + } catch (err) { + return { ok: false, error: normalizeDbError(err) } + } + } + ) + + ipcMain.handle('database:cancelQuery', async (event, args: { id: string }): Promise => { + requireTrusted(event.sender) + await dbConnectionManager.cancelQuery(args.id) + }) + + // Why: a single parameterized statement (Data-tab select/count or a wrapped + // free-form re-query). Like database:query, allowWrite is derived from the + // stored readOnly server-side; a missing connection defaults to read-only. + ipcMain.handle( + 'database:execute', + async (event, args: { id: string; statement: DbStatement }): Promise => { + requireTrusted(event.sender) + try { + const connection = store.getDbConnection(args.id) + const allowWrite = connection ? !connection.readOnly : false + const result = await dbConnectionManager.execute(args.id, args.statement, { + rowLimit: DB_MAX_ROWS, + timeoutMs: DB_STATEMENT_TIMEOUT_MS, + allowWrite + }) + return { ok: true, result } + } catch (err) { + return { ok: false, error: normalizeDbError(err) } + } + } + ) + + // Why: staged edits applied atomically. On a per-statement failure the redacted + // error carries the 0-based failedIndex so the grid can flag the offending change. + ipcMain.handle( + 'database:executeBatch', + async (event, args: { id: string; statements: DbStatement[] }): Promise => { + requireTrusted(event.sender) + try { + const connection = store.getDbConnection(args.id) + const allowWrite = connection ? !connection.readOnly : false + const rowCounts = await dbConnectionManager.executeBatch(args.id, args.statements, { + rowLimit: DB_MAX_ROWS, + timeoutMs: DB_STATEMENT_TIMEOUT_MS, + allowWrite + }) + return { ok: true, rowCounts } + } catch (err) { + // DbBatchError wraps the offending statement's raw driver error; redact + // that (never the batch wrapper) and surface which change failed. + const failedIndex = err instanceof DbBatchError ? err.failedIndex : -1 + const raw = err instanceof DbBatchError ? err.cause : err + return { ok: false, error: normalizeDbError(raw), failedIndex } + } + } + ) +} diff --git a/src/main/ipc/register-core-handlers.test.ts b/src/main/ipc/register-core-handlers.test.ts index f1fea8f7765..70e01af1e02 100644 --- a/src/main/ipc/register-core-handlers.test.ts +++ b/src/main/ipc/register-core-handlers.test.ts @@ -51,6 +51,7 @@ const { registerOnboardingHandlersMock, registerSpeechHandlersMock, registerSkillsHandlersMock, + registerDatabaseHandlersMock, registerWorkspaceSpaceHandlersMock, registerWorkspacePortHandlersMock, registerLocalhostWorktreeLabelHandlersMock, @@ -106,6 +107,7 @@ const { registerOnboardingHandlersMock: vi.fn(), registerSpeechHandlersMock: vi.fn(), registerSkillsHandlersMock: vi.fn(), + registerDatabaseHandlersMock: vi.fn(), registerWorkspaceSpaceHandlersMock: vi.fn(), registerWorkspacePortHandlersMock: vi.fn(), registerLocalhostWorktreeLabelHandlersMock: vi.fn(), @@ -186,6 +188,10 @@ vi.mock('./skills', () => ({ registerSkillsHandlers: registerSkillsHandlersMock })) +vi.mock('./database', () => ({ + registerDatabaseHandlers: registerDatabaseHandlersMock +})) + vi.mock('./workspace-space', () => ({ registerWorkspaceSpaceHandlers: registerWorkspaceSpaceHandlersMock })) diff --git a/src/main/ipc/register-core-handlers.ts b/src/main/ipc/register-core-handlers.ts index b54a54493ec..20478ed01dc 100644 --- a/src/main/ipc/register-core-handlers.ts +++ b/src/main/ipc/register-core-handlers.ts @@ -36,6 +36,7 @@ import { registerSessionHandlers } from './session' import { registerSettingsHandlers } from './settings' import { registerDiagnosticsHandlers } from './diagnostics' import { registerSkillsHandlers } from './skills' +import { registerDatabaseHandlers } from './database' import { registerWorkspaceSpaceHandlers } from './workspace-space' import { registerWorkspacePortHandlers } from './workspace-ports' import { registerLocalhostWorktreeLabelHandlers } from './localhost-worktree-labels' @@ -142,6 +143,7 @@ export function registerCoreHandlers( registerComputerUsePermissionHandlers() registerSettingsHandlers(store, agentAwakeService) registerSkillsHandlers(store) + registerDatabaseHandlers(store) if (automations) { registerAutomationHandlers(store, automations) } diff --git a/src/main/ipc/ui.ts b/src/main/ipc/ui.ts index efe6e300f85..c8161894ee0 100644 --- a/src/main/ipc/ui.ts +++ b/src/main/ipc/ui.ts @@ -58,7 +58,7 @@ export function registerUIHandlers(store: Store): void { }) } -function isTrustedUIRenderer(sender: WebContents): boolean { +export function isTrustedUIRenderer(sender: WebContents): boolean { if (sender.isDestroyed() || sender.getType() !== 'window') { return false } diff --git a/src/main/persistence-db-connections.test.ts b/src/main/persistence-db-connections.test.ts new file mode 100644 index 00000000000..82282c617ed --- /dev/null +++ b/src/main/persistence-db-connections.test.ts @@ -0,0 +1,527 @@ +import { describe, it, expect, vi, beforeEach, afterEach } from 'vitest' +import { + writeFileSync, + readFileSync, + mkdirSync, + rmSync +} from 'fs' +import { join } from 'path' +import { tmpdir } from 'os' +import type { DbConnection, DbConnectionInput } from '../shared/database-types' + +// Shared mutable state so the electron mock can reference a per-test directory +const testState = { dir: '' } + +const { isEncryptionAvailableMock, encryptStringMock, decryptStringMock } = vi.hoisted(() => ({ + isEncryptionAvailableMock: vi.fn(() => true), + encryptStringMock: vi.fn((plaintext: string) => + Buffer.from(`mock-encrypted:${plaintext}`, 'utf-8') + ), + decryptStringMock: vi.fn((ciphertext: Buffer) => { + const decoded = ciphertext.toString('utf-8') + if (!decoded.startsWith('mock-encrypted:')) { + throw new Error('invalid ciphertext') + } + return decoded.slice('mock-encrypted:'.length) + }) +})) + +vi.mock('electron', () => ({ + app: { + getPath: () => testState.dir + }, + safeStorage: { + isEncryptionAvailable: isEncryptionAvailableMock, + encryptString: encryptStringMock, + decryptString: decryptStringMock + } +})) + +vi.mock('./ssh/ssh-config-parser', () => ({ + loadUserSshConfig: vi.fn(), + sshConfigHostsToTargets: vi.fn() +})) + +vi.mock('./git/repo', () => ({ + getGitUsername: vi.fn().mockReturnValue('testuser') +})) + +vi.mock('./telemetry/client', () => ({ + track: vi.fn() +})) + +vi.mock('./telemetry/cohort-classifier', () => ({ + getCohortAtEmit: vi.fn() +})) + +// Reset modules and dynamically import Store so the data-file path picks up testState.dir +async function createStore() { + vi.resetModules() + const { Store, initDataPath } = await import('./persistence') + initDataPath() + return new Store() +} + +function dataFile(): string { + return join(testState.dir, 'orca-data.json') +} + +function writeDataFile(data: unknown): void { + mkdirSync(testState.dir, { recursive: true }) + writeFileSync(dataFile(), JSON.stringify(data, null, 2), 'utf-8') +} + +function readDataFile(): unknown { + return JSON.parse(readFileSync(dataFile(), 'utf-8')) +} + +const makeDbConnectionInput = (overrides: Partial = {}): DbConnectionInput => ({ + name: 'test-db', + engine: 'postgres', + host: 'localhost', + port: 5432, + database: 'testdb', + user: 'testuser', + password: 'testpass', + ...overrides +}) + +describe('Store DB Connections CRUD', () => { + beforeEach(() => { + testState.dir = join(tmpdir(), `orca-test-${Date.now()}-${Math.random().toString(36).slice(2)}`) + vi.clearAllMocks() + isEncryptionAvailableMock.mockReturnValue(true) + encryptStringMock.mockImplementation((plaintext: string) => + Buffer.from(`mock-encrypted:${plaintext}`, 'utf-8') + ) + decryptStringMock.mockImplementation((ciphertext: Buffer) => { + const decoded = ciphertext.toString('utf-8') + if (!decoded.startsWith('mock-encrypted:')) { + throw new Error('invalid ciphertext') + } + return decoded.slice('mock-encrypted:'.length) + }) + }) + + afterEach(() => { + if (testState.dir) { + rmSync(testState.dir, { recursive: true, force: true }) + } + }) + + describe('addDbConnection', () => { + it('assigns a uuid id + createdAt/updatedAt timestamps', async () => { + const store = await createStore() + const input = makeDbConnectionInput() + + const connection = store.addDbConnection(input) + + expect(connection.id).toBeDefined() + expect(connection.id).toMatch( + /^[0-9a-f]{8}-[0-9a-f]{4}-4[0-9a-f]{3}-[89ab][0-9a-f]{3}-[0-9a-f]{12}$/i + ) + expect(connection.createdAt).toBeGreaterThan(0) + expect(connection.updatedAt).toEqual(connection.createdAt) + }) + + it('defaults readOnly to false', async () => { + const store = await createStore() + const input = makeDbConnectionInput({ readOnly: undefined }) + + const connection = store.addDbConnection(input) + + expect(connection.readOnly).toBe(false) + }) + + it('stores password in tagged ENC form (not plaintext)', async () => { + const store = await createStore() + const input = makeDbConnectionInput({ password: 'mysecret' }) + + const connection = store.addDbConnection(input) + + expect(connection.password).toMatch(/^db\.safeStorage\.v1:/) + expect(connection.password).not.toContain('mysecret') + }) + + it('stores no password field if password is omitted', async () => { + const store = await createStore() + const input = makeDbConnectionInput({ password: undefined }) + + const connection = store.addDbConnection(input) + + expect(connection.password).toBeUndefined() + }) + + it('returns connection in getDbConnections', async () => { + const store = await createStore() + const input = makeDbConnectionInput() + + const added = store.addDbConnection(input) + const list = store.getDbConnections() + + expect(list).toHaveLength(1) + expect(list[0].id).toBe(added.id) + expect(list[0].name).toBe(input.name) + }) + + it('triggers a save on addDbConnection', async () => { + const store = await createStore() + const scheduleSaveSpy = vi.spyOn(store as unknown as { scheduleSave: () => void }, 'scheduleSave') + const input = makeDbConnectionInput({ password: 'mysecret' }) + + store.addDbConnection(input) + + expect(scheduleSaveSpy).toHaveBeenCalled() + }) + + it('falls back to RAW prefix when encryption unavailable', async () => { + isEncryptionAvailableMock.mockReturnValue(false) + const store = await createStore() + const input = makeDbConnectionInput({ password: 'mysecret' }) + + const connection = store.addDbConnection(input) + + expect(connection.password).toBe('db.plaintext.v1:mysecret') + }) + }) + + describe('updateDbConnection', () => { + it('omitting password leaves stored secret unchanged', async () => { + const store = await createStore() + const input = makeDbConnectionInput({ password: 'original' }) + const added = store.addDbConnection(input) + const originalPassword = added.password + + const updated = store.updateDbConnection(added.id, { name: 'renamed' }) + + expect(updated?.password).toBe(originalPassword) + }) + + it('passing new password replaces it in tagged form', async () => { + const store = await createStore() + const input = makeDbConnectionInput({ password: 'original' }) + const added = store.addDbConnection(input) + + const updated = store.updateDbConnection(added.id, { password: 'newsecret' }) + + expect(updated?.password).toMatch(/^db\.safeStorage\.v1:/) + expect(updated?.password).not.toContain('newsecret') + // Check that it was indeed stored + expect(updated?.password).toBeDefined() + }) + + it('passing empty password string keeps existing secret', async () => { + const store = await createStore() + const input = makeDbConnectionInput({ password: 'original' }) + const added = store.addDbConnection(input) + const originalPassword = added.password + + const updated = store.updateDbConnection(added.id, { password: '' }) + + expect(updated?.password).toBe(originalPassword) + }) + + it('updates non-password fields', async () => { + const store = await createStore() + const input = makeDbConnectionInput() + const added = store.addDbConnection(input) + + const updated = store.updateDbConnection(added.id, { + name: 'newname', + host: 'newhost', + port: 3306 + }) + + expect(updated?.name).toBe('newname') + expect(updated?.host).toBe('newhost') + expect(updated?.port).toBe(3306) + }) + + it('advances updatedAt timestamp', async () => { + const store = await createStore() + const input = makeDbConnectionInput() + const added = store.addDbConnection(input) + const originalUpdatedAt = added.updatedAt + + // Small delay to ensure time difference + await new Promise((r) => setTimeout(r, 10)) + + const updated = store.updateDbConnection(added.id, { name: 'newname' }) + + expect(updated?.updatedAt).toBeGreaterThan(originalUpdatedAt) + }) + + it('returns null if connection not found', async () => { + const store = await createStore() + + const result = store.updateDbConnection('nonexistent-id', { name: 'test' }) + + expect(result).toBeNull() + }) + }) + + describe('removeDbConnection', () => { + it('removes connection from store', async () => { + const store = await createStore() + const input1 = makeDbConnectionInput({ name: 'db1' }) + const input2 = makeDbConnectionInput({ name: 'db2' }) + const added1 = store.addDbConnection(input1) + const added2 = store.addDbConnection(input2) + + store.removeDbConnection(added1.id) + + const list = store.getDbConnections() + expect(list).toHaveLength(1) + expect(list[0].id).toBe(added2.id) + }) + + it('triggers save when removing connection', async () => { + const store = await createStore() + const input = makeDbConnectionInput() + const added = store.addDbConnection(input) + const scheduleSaveSpy = vi.spyOn(store as unknown as { scheduleSave: () => void }, 'scheduleSave') + + store.removeDbConnection(added.id) + + expect(scheduleSaveSpy).toHaveBeenCalled() + }) + + it('silently ignores removal of nonexistent id', async () => { + const store = await createStore() + const input = makeDbConnectionInput() + store.addDbConnection(input) + + expect(() => store.removeDbConnection('nonexistent-id')).not.toThrow() + + const list = store.getDbConnections() + expect(list).toHaveLength(1) + }) + }) + + describe('normalizeDbConnection on load', () => { + it('loads missing readOnly as false', async () => { + const now = Date.now() + const persistedData = { + dbConnections: [ + { + id: 'test-1', + name: 'test', + engine: 'postgres' as const, + host: 'localhost', + port: 5432, + database: 'db', + user: 'user', + createdAt: now, + updatedAt: now + // readOnly intentionally omitted + } as unknown as DbConnection + ] + } + writeDataFile(persistedData) + + const store = await createStore() + const list = store.getDbConnections() + + expect(list[0].readOnly).toBe(false) + }) + + it('loads invalid ssl as undefined', async () => { + const now = Date.now() + const persistedData = { + dbConnections: [ + { + id: 'test-1', + name: 'test', + engine: 'postgres' as const, + host: 'localhost', + port: 5432, + database: 'db', + user: 'user', + ssl: 'invalid-mode', + createdAt: now, + updatedAt: now, + readOnly: false + } + ] + } + writeDataFile(persistedData) + + const store = await createStore() + const list = store.getDbConnections() + + expect(list[0].ssl).toBeUndefined() + }) + + it('loads valid ssl value unchanged', async () => { + const now = Date.now() + const persistedData = { + dbConnections: [ + { + id: 'test-1', + name: 'test', + engine: 'postgres' as const, + host: 'localhost', + port: 5432, + database: 'db', + user: 'user', + ssl: 'verify-full' as const, + createdAt: now, + updatedAt: now, + readOnly: false + } + ] + } + writeDataFile(persistedData) + + const store = await createStore() + const list = store.getDbConnections() + + expect(list[0].ssl).toBe('verify-full') + }) + + it('preserves password field through round-trip', async () => { + const now = Date.now() + const persistedData = { + dbConnections: [ + { + id: 'test-1', + name: 'test', + engine: 'postgres' as const, + host: 'localhost', + port: 5432, + database: 'db', + user: 'user', + password: 'db.safeStorage.v1:mock-encrypted:mysecret', + createdAt: now, + updatedAt: now, + readOnly: false + } + ] + } + writeDataFile(persistedData) + + const store = await createStore() + const list = store.getDbConnections() + + expect(list[0].password).toBe('db.safeStorage.v1:mock-encrypted:mysecret') + }) + }) + + describe('getDbConnection by id', () => { + it('returns connection by id', async () => { + const store = await createStore() + const input = makeDbConnectionInput() + const added = store.addDbConnection(input) + + const retrieved = store.getDbConnection(added.id) + + expect(retrieved).toBeDefined() + expect(retrieved?.id).toBe(added.id) + expect(retrieved?.name).toBe(input.name) + }) + + it('returns undefined for nonexistent id', async () => { + const store = await createStore() + + const retrieved = store.getDbConnection('nonexistent-id') + + expect(retrieved).toBeUndefined() + }) + }) + + describe('data file persistence', () => { + it('maintains connections in in-memory state', async () => { + const store = await createStore() + store.addDbConnection(makeDbConnectionInput({ password: 'mysecret' })) + + const list = store.getDbConnections() + expect(list).toHaveLength(1) + expect(list[0].password).toMatch(/^db\.safeStorage\.v1:/) + }) + + it('maintains connections without password in state', async () => { + const store = await createStore() + store.addDbConnection(makeDbConnectionInput({ password: undefined })) + + const list = store.getDbConnections() + expect(list).toHaveLength(1) + expect(list[0].password).toBeUndefined() + }) + }) + + describe('encryption on disk vs in memory', () => { + it('keeps tagged password in memory after adding', async () => { + const store = await createStore() + const inMemory = store.addDbConnection(makeDbConnectionInput({ password: 'mysecret' })) + + // In memory, password is tagged (encrypted or plaintext prefix) — never raw. + expect(inMemory.password).toMatch(/^db\.(safeStorage\.v1:|plaintext\.v1:)/) + }) + + it('plaintext passwords are encrypted before storage (in-memory return)', async () => { + const store = await createStore() + const plaintext = 'mysecret' + const stored = store.addDbConnection(makeDbConnectionInput({ password: plaintext })) + + expect(stored.password).not.toBe(plaintext) + expect(stored.password).toMatch(/^db\./) + }) + + // Why: the real security guarantee is what reaches orca-data.json. BOTH the + // sync (flush/shutdown) and async (debounced) write paths must persist the + // tagged ciphertext, never the plaintext (red-team F2 / both-write-paths). + it('sync write path (flush) persists tagged ciphertext, not plaintext', async () => { + const store = await createStore() + const plaintext = 'on-disk-secret-sync' + store.addDbConnection(makeDbConnectionInput({ password: plaintext })) + + store.flush() + + const persisted = readDataFile() as { dbConnections?: { password?: string }[] } + expect(persisted.dbConnections).toHaveLength(1) + expect(persisted.dbConnections?.[0]?.password).toMatch(/^db\.safeStorage\.v1:/) + expect(persisted.dbConnections?.[0]?.password).not.toContain(plaintext) + }) + + it('async write path (debounced) persists tagged ciphertext, not plaintext', async () => { + const store = await createStore() + const plaintext = 'on-disk-secret-async' + store.addDbConnection(makeDbConnectionInput({ password: plaintext })) + + // Poll until the debounced write actually lands: a fixed sleep can race a + // slow CI timer (the 300ms debounce may fire later than any wall-clock + // guess, and waitForPendingWrite is a no-op until the timer sets it). + let persisted: { dbConnections?: { password?: string }[] } = {} + for (let i = 0; i < 200; i++) { + await store.waitForPendingWrite() + // The file may not exist until the first debounced write lands. + try { + persisted = readDataFile() as { dbConnections?: { password?: string }[] } + } catch { + persisted = {} + } + if (persisted.dbConnections?.[0]?.password) { + break + } + await new Promise((resolve) => setTimeout(resolve, 10)) + } + + expect(persisted.dbConnections?.[0]?.password).toMatch(/^db\.safeStorage\.v1:/) + expect(persisted.dbConnections?.[0]?.password).not.toContain(plaintext) + }) + + it('warn-and-store (no OS backend) persists RAW-tagged value on disk', async () => { + isEncryptionAvailableMock.mockReturnValue(false) + const store = await createStore() + const plaintext = 'weak-backend-secret' + store.addDbConnection(makeDbConnectionInput({ password: plaintext })) + + store.flush() + + const persisted = readDataFile() as { dbConnections?: { password?: string }[] } + // No keystore → warn-and-store: value is RAW-tagged (the form banner warned + // the user it is recoverable at rest); it is NOT silently dropped. + expect(persisted.dbConnections?.[0]?.password).toBe(`db.plaintext.v1:${plaintext}`) + }) + }) +}) diff --git a/src/main/persistence.ts b/src/main/persistence.ts index 5d9ef610f48..b7f5e4216d0 100644 --- a/src/main/persistence.ts +++ b/src/main/persistence.ts @@ -80,6 +80,15 @@ import { import type { MigrationUnsupportedPtyEntry } from '../shared/agent-status-types' import { MOBILE_PAIRING_USERDATA_FILES } from './runtime/mobile-pairing-files' import { hardenExistingSecureFile } from '../shared/secure-file' +import { + encryptDbSecret, + ensureDbSecretAtRest +} from './database/db-credential-store' +import type { + DbConnection, + DbConnectionInput, + DbConnectionUpdate +} from '../shared/database-types' import { LEGACY_DEFAULT_SSH_RELAY_GRACE_PERIOD_SECONDS, type SshRemotePtyLease, @@ -1019,6 +1028,40 @@ type LegacySshTarget = SshTarget & { // Why: old persisted targets predate configHost. Default to label-based lookup // so imported SSH aliases keep resolving through ssh -G after upgrade. +const DB_SSL_MODES = new Set(['disable', 'verify-full', 'insecure-no-verify']) + +// Why: older/partial records must load without crashing. readOnly defaults to +// FALSE (writable — validated decision; the Phase 5 confirm dialog is the write +// safety net). An invalid/absent ssl is dropped to undefined, which means +// smart-by-host (localhost → disable, remote → verify-full) at connect time. +function normalizeDbConnection(c: DbConnection): DbConnection { + return { + ...c, + readOnly: c.readOnly ?? false, + ssl: c.ssl && DB_SSL_MODES.has(c.ssl) ? c.ssl : undefined + } +} + +// Why: guarantee no plaintext password reaches disk. Normal flow already holds +// tagged at-rest values (encrypt-on-mutation), so this is a pass-through; the +// try/catch fails CLOSED — a strong-backend encrypt hiccup drops the secret from +// disk (re-enter next session) rather than persisting it recoverable. +function encryptDbConnectionForDisk(c: DbConnection): DbConnection { + if (!c.password) { + return c + } + try { + return { ...c, password: ensureDbSecretAtRest(c.password) } + } catch (err) { + console.error( + '[persistence] DB secret encryption failed; omitting password on disk for', + c.id, + err + ) + return { ...c, password: undefined } + } +} + function normalizeSshTarget(t: SshTarget): SshTarget { const target = { ...(t as LegacySshTarget) } const legacySyncEnabled = target.remoteWorkspaceSyncEnabled @@ -3200,6 +3243,10 @@ export class Store { defaults.workspaceSession ), sshTargets: (parsed.sshTargets ?? []).map(normalizeSshTarget), + // Why: passwords stay in tagged at-rest form in memory and are + // decrypted strictly at point-of-use (connect), so load never touches + // the keystore and can't crash on a keychain reset. + dbConnections: (parsed.dbConnections ?? []).map(normalizeDbConnection), sshRemotePtyLeases: (parsed.sshRemotePtyLeases ?? []) .map(normalizeSshRemotePtyLease) .filter((lease): lease is SshRemotePtyLease => lease !== null), @@ -3403,6 +3450,16 @@ export class Store { } } + // Why: a warn-and-store DB password (weak/absent OS crypto backend) lands as + // recoverable text in orca-data.json, which the plain write path leaves + // world-readable. Restrict the published file's ACL to the current user. + // hardenExistingSecureFile caches the applied path, so repeated saves are cheap. + private maybeHardenDataFileForDbSecrets(dataFile: string): void { + if (this.state.dbConnections.some((c) => c.password)) { + hardenExistingSecureFile(dataFile) + } + } + private computeStateHash(): string { return createHash('sha1').update(JSON.stringify(this.state)).digest('hex') } @@ -3423,7 +3480,8 @@ export class Store { ui: { ...this.state.ui, browserKagiSessionLink: encryptOptionalSecret(this.state.ui.browserKagiSessionLink) - } + }, + dbConnections: this.state.dbConnections.map(encryptDbConnectionForDisk) } return JSON.stringify(stateToSave, null, 2) } @@ -3470,6 +3528,9 @@ export class Store { await rm(tmpFile).catch(() => {}) } } + if (renamed) { + this.maybeHardenDataFileForDbSecrets(dataFile) + } // Why (issue #1158): rotate only after the atomic rename succeeded; then // re-check the generation so a concurrent flush owns any backup rotation. if (this.writeGeneration !== gen) { @@ -3520,6 +3581,9 @@ export class Store { } } } + if (renamed) { + this.maybeHardenDataFileForDbSecrets(dataFile) + } const now = Date.now() if (this.shouldRotateBackups(now, dataFile)) { this.rotateBackupsSync(dataFile) @@ -5741,6 +5805,61 @@ export class Store { this.scheduleSave() } + // ── Database Connections ─────────────────────────────────────────── + // Passwords are kept in tagged at-rest form (encrypt-on-mutation) so they are + // never plaintext in memory; decrypt strictly at point-of-use (connect). + + getDbConnections(): DbConnection[] { + return (this.state.dbConnections ?? []).map(normalizeDbConnection) + } + + getDbConnection(id: string): DbConnection | undefined { + const connection = this.state.dbConnections?.find((c) => c.id === id) + return connection ? normalizeDbConnection(connection) : undefined + } + + addDbConnection(input: DbConnectionInput): DbConnection { + const now = Date.now() + const connection = normalizeDbConnection({ + ...input, + id: randomUUID(), + readOnly: input.readOnly ?? false, + password: input.password ? encryptDbSecret(input.password) : undefined, + createdAt: now, + updatedAt: now + }) + this.state.dbConnections ??= [] + this.state.dbConnections.push(connection) + this.scheduleSave() + return connection + } + + updateDbConnection(id: string, updates: DbConnectionUpdate): DbConnection | null { + const connection = this.state.dbConnections?.find((c) => c.id === id) + if (!connection) { + return null + } + const { password, ...rest } = updates + Object.assign(connection, rest) + // Why: an omitted password leaves the stored secret unchanged; a non-empty + // password replaces it (re-encrypted to tagged at-rest form). + if (password) { + connection.password = encryptDbSecret(password) + } + connection.updatedAt = Date.now() + Object.assign(connection, normalizeDbConnection(connection)) + this.scheduleSave() + return normalizeDbConnection(connection) + } + + removeDbConnection(id: string): void { + if (!this.state.dbConnections) { + return + } + this.state.dbConnections = this.state.dbConnections.filter((c) => c.id !== id) + this.scheduleSave() + } + // ── SSH Remote PTY Leases ────────────────────────────────────────── getSshRemotePtyLeases(targetId?: string): SshRemotePtyLease[] { diff --git a/src/preload/api-types.ts b/src/preload/api-types.ts index 3f31cc560e2..16ea8d38097 100644 --- a/src/preload/api-types.ts +++ b/src/preload/api-types.ts @@ -351,6 +351,22 @@ import type { PortForwardEntry, EnrichedDetectedPort } from '../shared/ssh-types' +import type { + DbBatchResult, + DbColumnListResult, + DbConnectionInput, + DbConnectionRuntimeState, + DbConnectionSummary, + DbConnectionUpdate, + DbEncryptionStatus, + DbExecuteResult, + DbIntrospectResult, + DbQueryResult, + DbStatement, + DbTableListResult, + DbTableRef, + DbTestResult +} from '../shared/database-types' import type { CodexUsageBreakdownKind, CodexUsageBreakdownRow, @@ -2026,6 +2042,28 @@ export type PreloadApi = { skills: { discover: (target?: SkillDiscoveryTarget) => Promise } + database: { + list: () => Promise + add: (args: { input: DbConnectionInput }) => Promise + update: (args: { + id: string + updates: DbConnectionUpdate + }) => Promise + remove: (args: { id: string }) => Promise + encryptionStatus: () => Promise + test: (args: { input: DbConnectionInput; id?: string }) => Promise + connect: (args: { id: string }) => Promise + disconnect: (args: { id: string }) => Promise + statuses: () => Promise + onStatusChanged: (callback: (state: DbConnectionRuntimeState) => void) => () => void + introspect: (args: { id: string }) => Promise + introspectSchemaTables: (args: { id: string; schema: string }) => Promise + introspectTableColumns: (args: { id: string; ref: DbTableRef }) => Promise + query: (args: { id: string; sql: string }) => Promise + cancelQuery: (args: { id: string }) => Promise + execute: (args: { id: string; statement: DbStatement }) => Promise + executeBatch: (args: { id: string; statements: DbStatement[] }) => Promise + } pet: { import: () => Promise importPetBundle: () => Promise diff --git a/src/preload/index.ts b/src/preload/index.ts index 9a162aa5473..b440c231caf 100644 --- a/src/preload/index.ts +++ b/src/preload/index.ts @@ -123,6 +123,22 @@ import type { PortForwardEntry, EnrichedDetectedPort } from '../shared/ssh-types' +import type { + DbBatchResult, + DbColumnListResult, + DbConnectionInput, + DbConnectionRuntimeState, + DbConnectionSummary, + DbConnectionUpdate, + DbEncryptionStatus, + DbExecuteResult, + DbIntrospectResult, + DbQueryResult, + DbStatement, + DbTableListResult, + DbTableRef, + DbTestResult +} from '../shared/database-types' import type { AgentStatusIpcPayload, MigrationUnsupportedPtyEntry @@ -2004,6 +2020,50 @@ const api = { ipcRenderer.invoke('skills:discover', target) }, + database: { + list: (): Promise => ipcRenderer.invoke('database:list'), + add: (args: { input: DbConnectionInput }): Promise => + ipcRenderer.invoke('database:add', args), + update: (args: { + id: string + updates: DbConnectionUpdate + }): Promise => ipcRenderer.invoke('database:update', args), + remove: (args: { id: string }): Promise => ipcRenderer.invoke('database:remove', args), + encryptionStatus: (): Promise => + ipcRenderer.invoke('database:encryptionStatus'), + test: (args: { input: DbConnectionInput; id?: string }): Promise => + ipcRenderer.invoke('database:test', args), + connect: (args: { id: string }): Promise => + ipcRenderer.invoke('database:connect', args), + disconnect: (args: { id: string }): Promise => + ipcRenderer.invoke('database:disconnect', args), + statuses: (): Promise => ipcRenderer.invoke('database:statuses'), + onStatusChanged: ( + callback: (state: DbConnectionRuntimeState) => void + ): (() => void) => { + const listener = ( + _event: Electron.IpcRendererEvent, + state: DbConnectionRuntimeState + ): void => callback(state) + ipcRenderer.on('database:status-changed', listener) + return () => ipcRenderer.removeListener('database:status-changed', listener) + }, + introspect: (args: { id: string }): Promise => + ipcRenderer.invoke('database:introspect', args), + introspectSchemaTables: (args: { id: string; schema: string }): Promise => + ipcRenderer.invoke('database:introspectSchemaTables', args), + introspectTableColumns: (args: { id: string; ref: DbTableRef }): Promise => + ipcRenderer.invoke('database:introspectTableColumns', args), + query: (args: { id: string; sql: string }): Promise => + ipcRenderer.invoke('database:query', args), + cancelQuery: (args: { id: string }): Promise => + ipcRenderer.invoke('database:cancelQuery', args), + execute: (args: { id: string; statement: DbStatement }): Promise => + ipcRenderer.invoke('database:execute', args), + executeBatch: (args: { id: string; statements: DbStatement[] }): Promise => + ipcRenderer.invoke('database:executeBatch', args) + }, + pet: { import: (): Promise => ipcRenderer.invoke('pet:import'), importPetBundle: (): Promise => ipcRenderer.invoke('pet:importPetBundle'), diff --git a/src/renderer/src/App.tsx b/src/renderer/src/App.tsx index 697d0dd45f5..2ee0fff805e 100644 --- a/src/renderer/src/App.tsx +++ b/src/renderer/src/App.tsx @@ -284,6 +284,7 @@ const Settings = lazy(() => import('./components/settings/Settings')) const SkillsPage = lazy(() => import('./components/skills/SkillsPage')) const WorkspaceSpacePage = lazy(() => import('./components/workspace-space/WorkspaceSpacePage')) const MobilePage = lazy(() => import('./components/mobile/MobilePage')) +const DatabasePage = lazy(() => import('./components/database/DatabasePage')) const QuickOpen = lazy(() => import('./components/QuickOpen')) const WorktreeJumpPalette = lazy(() => import('./components/WorktreeJumpPalette')) const WorkspaceCleanupDialog = lazy( @@ -1413,7 +1414,8 @@ function App(): React.JSX.Element { activeView !== 'settings' && activeView !== 'activity' && activeView !== 'space' && - activeView !== 'skills' + activeView !== 'skills' && + activeView !== 'database' // Why: Tasks/Landing keep the full titlebar only when the sidebar is // collapsed; with it open, mirror workspace view so titlebar-left sits flush // above nav. Creation layout suppresses the full-width titlebar. @@ -2303,6 +2305,7 @@ function App(): React.JSX.Element { {activeView === 'activity' ? : null} {activeView === 'space' ? : null} {activeView === 'mobile' ? : null} + {activeView === 'database' ? : null} {activeView === 'terminal' && creationLayoutActive && activePendingCreationId ? ( diff --git a/src/renderer/src/components/database/ConnectionForm.tsx b/src/renderer/src/components/database/ConnectionForm.tsx new file mode 100644 index 00000000000..a5f48e4c8bc --- /dev/null +++ b/src/renderer/src/components/database/ConnectionForm.tsx @@ -0,0 +1,417 @@ +import React, { useEffect, useState } from 'react' +import { toast } from 'sonner' +import { Button } from '@/components/ui/button' +import { Checkbox } from '@/components/ui/checkbox' +import { + Dialog, + DialogContent, + DialogDescription, + DialogFooter, + DialogHeader, + DialogTitle +} from '@/components/ui/dialog' +import { Input } from '@/components/ui/input' +import { Label } from '@/components/ui/label' +import { + Select, + SelectContent, + SelectItem, + SelectTrigger, + SelectValue +} from '@/components/ui/select' +import { useAppStore } from '@/store' +import { useMountedRef } from '@/hooks/useMountedRef' +import { translate } from '@/i18n/i18n' +import type { + DbConnectionInput, + DbConnectionSummary, + DbConnectionUpdate, + DbEngine, + DbSslMode +} from '../../../../shared/database-types' +import { DB_DEFAULT_PORT } from '../../../../shared/database-types' +import { ConsentCheckbox, EncryptionWarningBanner } from './connection-encryption-warning' +import { buildInitialState, type SslFieldValue } from './connection-form-defaults' + +type ConnectionFormProps = { + open: boolean + onOpenChange: (open: boolean) => void + connection?: DbConnectionSummary +} + +export function ConnectionForm({ + open, + onOpenChange, + connection +}: ConnectionFormProps): React.JSX.Element { + const addDbConnection = useAppStore((s) => s.addDbConnection) + const updateDbConnection = useAppStore((s) => s.updateDbConnection) + const testDbConnection = useAppStore((s) => s.testDbConnection) + const dbEncryptionStatus = useAppStore((s) => s.dbEncryptionStatus) + const mountedRef = useMountedRef() + + const [name, setName] = useState('') + const [engine, setEngine] = useState('postgres') + const [host, setHost] = useState('') + const [port, setPort] = useState(DB_DEFAULT_PORT.postgres.toString()) + const [database, setDatabase] = useState('') + const [user, setUser] = useState('') + const [password, setPassword] = useState('') + const [ssl, setSsl] = useState('auto') + const [readOnly, setReadOnly] = useState(false) + const [saving, setSaving] = useState(false) + const [testing, setTesting] = useState(false) + const [consentChecked, setConsentChecked] = useState(false) + + // Reset all fields whenever the dialog opens or the target connection changes. + useEffect(() => { + if (!open) { return } + const initial = buildInitialState(connection) + setName(initial.name) + setEngine(initial.engine) + setHost(initial.host) + setPort(initial.port) + setDatabase(initial.database) + setUser(initial.user) + setPassword('') + setSsl(initial.ssl) + setReadOnly(initial.readOnly) + setConsentChecked(false) + setSaving(false) + setTesting(false) + }, [open, connection]) + + // Build the current form fields into a create payload — shared by Test and Save. + function currentInput(): DbConnectionInput { + const sslValue: DbSslMode | undefined = ssl === 'auto' ? undefined : ssl + return { + name: name.trim(), + engine, + host: host.trim(), + port: parsedPort, + database: database.trim(), + user: user.trim(), + ssl: sslValue, + readOnly, + ...(password.trim() ? { password: password.trim() } : {}) + } + } + + async function handleTest(): Promise { + if (!isValid || testing) { return } + setTesting(true) + try { + // An empty password field on an existing connection falls back to the + // stored secret (resolved in the main process), so pass the id along. + const result = await testDbConnection(currentInput(), connection?.id) + if (!mountedRef.current) { return } + if (result.ok) { + toast.success( + translate('auto.components.database.ConnectionForm.testSuccess', 'Connection succeeded') + ) + } else { + toast.error(result.error.safeMessage) + } + } catch { + if (mountedRef.current) { + toast.error( + translate('auto.components.database.ConnectionForm.testError', 'Test failed') + ) + } + } finally { + if (mountedRef.current) { setTesting(false) } + } + } + + function handleEngineChange(value: string): void { + const next = value as DbEngine + const prevDefault = DB_DEFAULT_PORT[engine].toString() + // Auto-advance port only when the user hasn't customised it away from the previous default. + if (port === prevDefault) { + setPort(DB_DEFAULT_PORT[next].toString()) + } + setEngine(next) + } + + const parsedPort = parseInt(port, 10) + const isValidPort = !Number.isNaN(parsedPort) && parsedPort >= 1 && parsedPort <= 65535 + const isValid = + name.trim().length > 0 && + host.trim().length > 0 && + database.trim().length > 0 && + user.trim().length > 0 && + isValidPort + + const isWeakEncryption = dbEncryptionStatus !== null && !dbEncryptionStatus.isStrong + + // "entering/keeping a password": create → any password text; edit → text entered OR existing secret kept. + const hasPasswordIntent = + password.trim().length > 0 || (connection !== undefined && (connection.hasPassword || false)) + + const needsConsent = isWeakEncryption && hasPasswordIntent + // While encryption status is still loading (null) and a password will be stored, + // block Save so the consent gate can't be bypassed during the load window. + const encryptionStatusPending = dbEncryptionStatus === null && hasPasswordIntent + const saveEnabled = isValid && !saving && !encryptionStatusPending && (!needsConsent || consentChecked) + + const isEditMode = connection !== undefined + const dialogTitle = isEditMode + ? translate('auto.components.database.ConnectionForm.titleEdit', 'Edit connection') + : translate('auto.components.database.ConnectionForm.titleCreate', 'Add connection') + const dialogDescription = isEditMode + ? translate( + 'auto.components.database.ConnectionForm.descriptionEdit', + 'Update the database connection settings.' + ) + : translate( + 'auto.components.database.ConnectionForm.descriptionCreate', + 'Configure a new database connection.' + ) + + async function handleSubmit(event: React.FormEvent): Promise { + event.preventDefault() + if (!saveEnabled) { return } + setSaving(true) + try { + // An omitted password leaves the stored secret unchanged (store contract). + await (isEditMode + ? updateDbConnection(connection.id, currentInput() as DbConnectionUpdate) + : addDbConnection(currentInput())) + + if (mountedRef.current) { + onOpenChange(false) + } + } catch (error) { + console.error('Failed to save connection:', error) + if (mountedRef.current) { + toast.error( + isEditMode + ? translate( + 'auto.components.database.ConnectionForm.errorUpdate', + 'Failed to update connection' + ) + : translate( + 'auto.components.database.ConnectionForm.errorCreate', + 'Failed to add connection' + ) + ) + setSaving(false) + } + } + } + + return ( + + + + {dialogTitle} + {dialogDescription} + + + {/* Warn-and-store banner: visible whenever the OS lacks a strong secret store. */} + {isWeakEncryption ? : null} + +
void handleSubmit(e)} className="flex flex-col gap-4"> +
+ + setName(e.target.value)} + placeholder={translate( + 'auto.components.database.ConnectionForm.namePlaceholder', + 'My database' + )} + autoComplete="off" + /> +
+ +
+
+ + +
+
+ + +
+
+ +
+
+ + setHost(e.target.value)} + placeholder={translate( + 'auto.components.database.ConnectionForm.hostPlaceholder', + 'localhost' + )} + autoComplete="off" + /> +
+
+ + setPort(e.target.value)} + min={1} + max={65535} + autoComplete="off" + /> +
+
+ +
+ + setDatabase(e.target.value)} + autoComplete="off" + /> +
+ +
+ + setUser(e.target.value)} + autoComplete="off" + /> +
+ + {/* Password — write-only; in edit mode never pre-fills the stored secret. */} +
+ + setPassword(e.target.value)} + placeholder={ + isEditMode && connection.hasPassword + ? translate( + 'auto.components.database.ConnectionForm.passwordPlaceholderEdit', + '•••••• (unchanged — type to replace)' + ) + : undefined + } + autoComplete="new-password" + /> +
+ +
+ setReadOnly(checked === true)} + /> + +
+ + {/* Consent required when weak encryption backend AND a password will be stored. */} + {needsConsent ? ( + + ) : null} + + + +
+ + +
+
+ +
+
+ ) +} diff --git a/src/renderer/src/components/database/ConnectionList.tsx b/src/renderer/src/components/database/ConnectionList.tsx new file mode 100644 index 00000000000..44cb2556cfc --- /dev/null +++ b/src/renderer/src/components/database/ConnectionList.tsx @@ -0,0 +1,202 @@ +import React, { useState } from 'react' +import { Database, Plus } from 'lucide-react' +import { toast } from 'sonner' +import { Button } from '@/components/ui/button' +import { + Dialog, + DialogContent, + DialogDescription, + DialogFooter, + DialogHeader, + DialogTitle +} from '@/components/ui/dialog' +import { useAppStore } from '@/store' +import { translate } from '@/i18n/i18n' +import type { DbConnectionSummary } from '../../../../shared/database-types' +import { ConnectionForm } from './ConnectionForm' +import { ConnectionRow } from './connection-row' + +function EmptyConnectionState({ onAddClick }: { onAddClick: () => void }): React.JSX.Element { + return ( +
+
+ +
+

+ {translate( + 'auto.components.database.ConnectionList.emptyTitle', + 'No connections yet' + )} +

+

+ {translate( + 'auto.components.database.ConnectionList.emptyBody', + 'Add a Postgres or MySQL connection to get started.' + )} +

+
+ +
+
+ ) +} + +export function ConnectionList(): React.JSX.Element { + const dbConnections = useAppStore((s) => s.dbConnections) + const removeDbConnection = useAppStore((s) => s.removeDbConnection) + const dbStatuses = useAppStore((s) => s.dbStatuses) + const connectDbConnection = useAppStore((s) => s.connectDbConnection) + const disconnectDbConnection = useAppStore((s) => s.disconnectDbConnection) + const activeDbConnectionId = useAppStore((s) => s.activeDbConnectionId) + const setActiveDbConnection = useAppStore((s) => s.setActiveDbConnection) + + const [formOpen, setFormOpen] = useState(false) + const [editingConnection, setEditingConnection] = useState( + undefined + ) + const [deletingId, setDeletingId] = useState(null) + + const deletingConnection = dbConnections.find((c) => c.id === deletingId) ?? null + + async function handleConnect(id: string): Promise { + const result = await connectDbConnection(id) + if (result.status === 'error' || result.status === 'lost') { + toast.error( + result.error?.safeMessage ?? + translate( + 'auto.components.database.ConnectionList.errorConnect', + 'Failed to connect' + ) + ) + } + } + + async function handleDisconnect(id: string): Promise { + try { + await disconnectDbConnection(id) + } catch { + toast.error( + translate( + 'auto.components.database.ConnectionList.errorDisconnect', + 'Failed to disconnect' + ) + ) + } + } + + function openCreate(): void { + setEditingConnection(undefined) + setFormOpen(true) + } + + function openEdit(connection: DbConnectionSummary): void { + setEditingConnection(connection) + setFormOpen(true) + } + + async function handleDelete(): Promise { + if (!deletingId) { return } + const id = deletingId + setDeletingId(null) + try { + await removeDbConnection(id) + } catch (error) { + console.error('Failed to remove connection:', error) + toast.error( + translate( + 'auto.components.database.ConnectionList.errorDelete', + 'Failed to remove connection' + ) + ) + } + } + + return ( + <> + {dbConnections.length === 0 ? ( + + ) : ( +
+
+ +
+
+
+ {dbConnections.map((connection) => ( + void handleConnect(id)} + onDisconnect={(id) => void handleDisconnect(id)} + onEdit={openEdit} + onDelete={setDeletingId} + /> + ))} +
+
+
+ )} + + { + if (!open) { setDeletingId(null) } + }} + > + + + + {translate( + 'auto.components.database.ConnectionList.confirmDeleteTitle', + 'Delete connection' + )} + + {deletingConnection ? ( + + {translate( + 'auto.components.database.ConnectionList.confirmDeleteBody', + 'Permanently remove "{{name}}"? This cannot be undone.', + { name: deletingConnection.name } + )} + + ) : null} + + + + + + + + + + + ) +} diff --git a/src/renderer/src/components/database/DataGridColumnHeader.tsx b/src/renderer/src/components/database/DataGridColumnHeader.tsx new file mode 100644 index 00000000000..2ed82b734b7 --- /dev/null +++ b/src/renderer/src/components/database/DataGridColumnHeader.tsx @@ -0,0 +1,62 @@ +import React from 'react' +import { ArrowDown, ArrowUp, ChevronsUpDown, KeyRound } from 'lucide-react' +import type { DbColumnFilter, DbSortDirection } from '../../../../shared/database-types' +import { DataGridColumnFilter } from './data-grid-column-filter' + +// One sortable/filterable header cell, shared by the table Data grid and (Phase 4) +// the free-form results grid. Sort is a tri-state toggle; the funnel opens the +// per-column filter popover. Both are opt-out via `sortable`/`filterable`. +export function DataGridColumnHeader({ + name, + dataType, + isPrimaryKey, + sortDirection, + onSort, + filter, + onFilter, + sortable = true, + filterable = true +}: { + name: string + dataType?: string + isPrimaryKey?: boolean + sortDirection: DbSortDirection | null + onSort?: () => void + filter?: DbColumnFilter + onFilter?: (filter: DbColumnFilter | null) => void + sortable?: boolean + filterable?: boolean +}): React.JSX.Element { + const SortIcon = + sortDirection === 'asc' ? ArrowUp : sortDirection === 'desc' ? ArrowDown : ChevronsUpDown + return ( +
+ + {filterable && onFilter ? ( + + ) : null} +
+ ) +} diff --git a/src/renderer/src/components/database/DatabasePage.tsx b/src/renderer/src/components/database/DatabasePage.tsx new file mode 100644 index 00000000000..7bf870a051e --- /dev/null +++ b/src/renderer/src/components/database/DatabasePage.tsx @@ -0,0 +1,110 @@ +import { useEffect } from 'react' +import { ArrowLeft, Database } from 'lucide-react' +import { Badge } from '@/components/ui/badge' +import { Button } from '@/components/ui/button' +import { useAppStore } from '@/store' +import { translate } from '@/i18n/i18n' +import { ConnectionList } from './ConnectionList' +import { SchemaTree } from './SchemaTree' +import { DatabaseWorkspace } from './DatabaseWorkspace' + +// Phase 2: connection list + form mount into the body slot below the header. +export default function DatabasePage(): React.JSX.Element { + const closeDatabasePage = useAppStore((s) => s.closeDatabasePage) + const loadDbConnections = useAppStore((s) => s.loadDbConnections) + const subscribeDbStatusChanges = useAppStore((s) => s.subscribeDbStatusChanges) + const activeDbConnectionId = useAppStore((s) => s.activeDbConnectionId) + + useEffect(() => { + void loadDbConnections() + }, [loadDbConnections]) + + // Why: live status (connected/lost/error) is pushed from the main process; keep + // the subscription for the page's lifetime so a dropped connection updates the UI. + useEffect(() => subscribeDbStatusChanges(), [subscribeDbStatusChanges]) + + useEffect(() => { + const hasVisibleOverlay = (): boolean => + Array.from( + document.querySelectorAll('[role="dialog"], [role="listbox"], [role="menu"]') + ).some((element) => element instanceof HTMLElement) + + const handleKeyDown = (event: KeyboardEvent): void => { + if (event.key !== 'Escape') { + return + } + // Why: menus and dialogs own Escape before page-level navigation. + if (hasVisibleOverlay()) { + return + } + const target = event.target as HTMLElement | null + if ( + target?.matches('input, textarea, select, [contenteditable="true"], [contenteditable=""]') + ) { + return + } + event.preventDefault() + closeDatabasePage() + } + + // Why: capture keeps page-level back navigation reliable when no overlay is active. + window.addEventListener('keydown', handleKeyDown, { capture: true }) + return () => window.removeEventListener('keydown', handleKeyDown, { capture: true }) + }, [closeDatabasePage]) + + return ( +
+
+ +
+ +
+
+

+ {translate('auto.components.database.DatabasePage.title', 'Database')} +

+ + {translate('auto.components.database.DatabasePage.beta', 'Beta')} + +
+

+ {translate( + 'auto.components.database.DatabasePage.subtitle', + 'Connect to Postgres and MySQL servers' + )} +

+
+
+
+ +
+
+ +
+ {/* Schema browser + query workspace for the active (connected) connection. */} + {activeDbConnectionId ? ( + <> +
+ +
+
+ +
+ + ) : null} +
+
+ ) +} diff --git a/src/renderer/src/components/database/DatabaseWorkspace.tsx b/src/renderer/src/components/database/DatabaseWorkspace.tsx new file mode 100644 index 00000000000..4e5f714c26c --- /dev/null +++ b/src/renderer/src/components/database/DatabaseWorkspace.tsx @@ -0,0 +1,160 @@ +import React, { useState } from 'react' +import { Table2, TerminalSquare, X } from 'lucide-react' +import { Button } from '@/components/ui/button' +import { + Dialog, + DialogContent, + DialogDescription, + DialogFooter, + DialogHeader, + DialogTitle +} from '@/components/ui/dialog' +import { useAppStore } from '@/store' +import { translate } from '@/i18n/i18n' +import { DB_QUERY_TAB_ID, defaultWorkspaceTabs } from '@/store/database-workspace-tabs' +import { isBufferDirty } from './table-data-edit-buffer' +import { QueryWorkspace } from './QueryWorkspace' +import { TableDataView } from './TableDataView' + +// Tabbed workspace for one connection: a permanent free-form "Query" tab plus a +// "Data" tab per table opened from the schema tree. Only the active tab's content +// mounts — per-tab state lives in the store, so switching never loses work. +export function DatabaseWorkspace({ connectionId }: { connectionId: string }): React.JSX.Element { + const ws = useAppStore((s) => s.dbWorkspaceTabs[connectionId]) + const setActiveDbTab = useAppStore((s) => s.setActiveDbTab) + const closeDbTab = useAppStore((s) => s.closeDbTab) + const dbTableData = useAppStore((s) => s.dbTableData) + + const [closingTabId, setClosingTabId] = useState(null) + + const tabs = ws?.tabs ?? defaultWorkspaceTabs().tabs + const activeTabId = ws?.activeTabId ?? DB_QUERY_TAB_ID + + // Check for unsaved edits before closing a Data tab and prompt for confirmation. + const handleCloseTab = (tabId: string): void => { + const edit = dbTableData[connectionId]?.[tabId]?.edit + if (edit && isBufferDirty(edit)) { + setClosingTabId(tabId) + } else { + closeDbTab(connectionId, tabId) + } + } + + return ( + <> +
+
+ {tabs.map((tab) => { + const active = tab.tabId === activeTabId + const isQuery = tab.kind === 'query' + return ( +
setActiveDbTab(connectionId, tab.tabId)} + onKeyDown={(event) => { + if (event.key === 'Enter' || event.key === ' ') { + event.preventDefault() + setActiveDbTab(connectionId, tab.tabId) + } + }} + className={`group flex h-8 shrink-0 cursor-pointer items-center gap-1.5 border-b-2 px-2 text-xs outline-none ${ + active + ? 'border-primary text-foreground' + : 'border-transparent text-muted-foreground hover:text-foreground' + }`} + > + {isQuery ? ( + + ) : ( + + )} + + {tab.kind === 'query' + ? translate('auto.components.database.DatabaseWorkspace.queryTab', 'Query') + : tab.table} + + {tab.kind === 'table-data' ? ( + + ) : null} +
+ ) + })} +
+ +
+ {activeTabId === DB_QUERY_TAB_ID ? ( + + ) : ( + + )} +
+
+ + {/* Confirm discard when a Data tab with unsaved edits is closed. */} + { + if (!open) setClosingTabId(null) + }} + > + + + + {translate( + 'auto.components.database.DatabaseWorkspace.discardTitle', + 'Discard unsaved changes?' + )} + + + {translate( + 'auto.components.database.DatabaseWorkspace.discardBody', + 'Closing this tab will discard unsaved edits. This cannot be undone.' + )} + + + + + + + + + + ) +} diff --git a/src/renderer/src/components/database/QueryEditor.tsx b/src/renderer/src/components/database/QueryEditor.tsx new file mode 100644 index 00000000000..737e8a63c26 --- /dev/null +++ b/src/renderer/src/components/database/QueryEditor.tsx @@ -0,0 +1,55 @@ +import React, { useCallback, useEffect, useRef } from 'react' +import Editor, { type OnMount } from '@monaco-editor/react' +import '@/lib/monaco-setup' +import { ensureSqlLanguageRegistered } from '@/lib/monaco-sql-language' +import { useAppStore } from '@/store' + +type QueryEditorProps = { + value: string + onChange: (value: string) => void + onRunShortcut: () => void +} + +// Monaco SQL editor. Cmd/Ctrl+Enter runs — KeyMod.CtrlCmd is platform-aware +// (Cmd on macOS, Ctrl on Linux/Windows), satisfying the cross-platform rule. +export function QueryEditor({ value, onChange, onRunShortcut }: QueryEditorProps): React.JSX.Element { + const settings = useAppStore((s) => s.settings) + const isDark = + settings?.theme === 'dark' || + (settings?.theme === 'system' && window.matchMedia('(prefers-color-scheme: dark)').matches) + + // Keep the latest run handler without rebinding the Monaco command. + const runRef = useRef(onRunShortcut) + runRef.current = onRunShortcut + + useEffect(() => { + void ensureSqlLanguageRegistered() + }, []) + + const handleMount: OnMount = useCallback((editor, monaco) => { + void ensureSqlLanguageRegistered() + editor.addCommand(monaco.KeyMod.CtrlCmd | monaco.KeyCode.Enter, () => runRef.current()) + }, []) + + return ( + onChange(next ?? '')} + onMount={handleMount} + options={{ + minimap: { enabled: false }, + scrollBeyondLastLine: false, + fontSize: 13, + fontFamily: settings?.terminalFontFamily || 'monospace', + lineNumbers: 'on', + automaticLayout: true, + tabSize: 2, + wordWrap: 'on', + padding: { top: 8 } + }} + /> + ) +} diff --git a/src/renderer/src/components/database/QueryWorkspace.tsx b/src/renderer/src/components/database/QueryWorkspace.tsx new file mode 100644 index 00000000000..0719b92d123 --- /dev/null +++ b/src/renderer/src/components/database/QueryWorkspace.tsx @@ -0,0 +1,152 @@ +import React, { useCallback, useState } from 'react' +import { Loader2, Play, Square } from 'lucide-react' +import { Badge } from '@/components/ui/badge' +import { Button } from '@/components/ui/button' +import { + Dialog, + DialogContent, + DialogDescription, + DialogFooter, + DialogHeader, + DialogTitle +} from '@/components/ui/dialog' +import { useAppStore } from '@/store' +import { translate } from '@/i18n/i18n' +import { needsWriteConfirm } from '../../../../shared/sql-statement-classifier' +import { setColumnFilter } from './data-grid-filters' +import { QueryEditor } from './QueryEditor' +import { ResultsGrid } from './ResultsGrid' + +export function QueryWorkspace({ connectionId }: { connectionId: string }): React.JSX.Element { + const text = useAppStore((s) => s.dbQueryText[connectionId] ?? '') + const queryState = useAppStore((s) => s.dbQueryState[connectionId]) + const readOnly = useAppStore( + (s) => s.dbConnections.find((c) => c.id === connectionId)?.readOnly ?? false + ) + const setDbQueryText = useAppStore((s) => s.setDbQueryText) + const runDbQuery = useAppStore((s) => s.runDbQuery) + const cancelDbQuery = useAppStore((s) => s.cancelDbQuery) + const setDbQuerySort = useAppStore((s) => s.setDbQuerySort) + const setDbQueryFilters = useAppStore((s) => s.setDbQueryFilters) + const setDbQueryPage = useAppStore((s) => s.setDbQueryPage) + + const running = queryState?.running ?? false + const [pendingSql, setPendingSql] = useState(null) + + const execute = useCallback( + (sql: string) => { + void runDbQuery(connectionId, sql) + }, + [connectionId, runDbQuery] + ) + + const handleRun = useCallback(() => { + const sql = text.trim() + if (!sql || running) { + return + } + // Writable connections have no DB backstop, so a destructive/ambiguous + // statement must be confirmed first (read-only connections are DB-enforced). + if (!readOnly && needsWriteConfirm(sql)) { + setPendingSql(sql) + return + } + execute(sql) + }, [text, running, readOnly, execute]) + + return ( +
+
+ {running ? ( + + ) : ( + + )} + {running ? : null} + {readOnly ? ( + + {translate('auto.components.database.QueryWorkspace.readOnly', 'Read-only')} + + ) : null} + + {translate('auto.components.database.QueryWorkspace.runHint', '{{mod}} + Enter to run', { + // Show the platform's real run-shortcut modifier (⌘ on Mac, Ctrl elsewhere). + mod: navigator.userAgent.includes('Mac') ? '⌘' : 'Ctrl' + })} + +
+ +
+ setDbQueryText(connectionId, next)} + onRunShortcut={handleRun} + /> +
+ + setDbQuerySort(connectionId, ordinal), + onFilter: (column, filter) => + setDbQueryFilters( + connectionId, + setColumnFilter(queryState.refine?.filters ?? [], column, filter) + ), + onPage: (delta) => setDbQueryPage(connectionId, delta) + } + : undefined + } + /> + + !open && setPendingSql(null)}> + + + + {translate('auto.components.database.QueryWorkspace.confirmTitle', 'Run this statement?')} + + + {translate( + 'auto.components.database.QueryWorkspace.confirmBody', + 'This looks like it may modify data or schema. Run it against a writable connection?' + )} + + + + + + + + +
+ ) +} diff --git a/src/renderer/src/components/database/ResultsGrid.tsx b/src/renderer/src/components/database/ResultsGrid.tsx new file mode 100644 index 00000000000..d4ad7e7efda --- /dev/null +++ b/src/renderer/src/components/database/ResultsGrid.tsx @@ -0,0 +1,211 @@ +import React, { useRef } from 'react' +import { useVirtualizer } from '@tanstack/react-virtual' +import { ChevronLeft, ChevronRight, Loader2 } from 'lucide-react' +import { Button } from '@/components/ui/button' +import { translate } from '@/i18n/i18n' +import type { DbColumnFilter, DbSafeError, QueryResult } from '../../../../shared/database-types' +import type { DbQueryRefine } from '@/store/slices/database' +import { formatCell } from './data-grid-cell-format' +import { DataGridColumnHeader } from './DataGridColumnHeader' +import { filterFor } from './data-grid-filters' +import { ordinalSortDirectionFor } from './data-grid-sort-state' + +const ROW_HEIGHT = 24 +const OVERSCAN = 16 +const COL_MIN_PX = 120 +const COL_MAX_PX = 320 + +// Handlers that turn the free-form results grid into a server-side sort/filter +// surface (wrapping the last read). Omitted → plain read-only headers. +export type ResultsGridRefine = { + refine: DbQueryRefine + onSort: (ordinal: number) => void + onFilter: (column: string, filter: DbColumnFilter | null) => void + onPage: (delta: number) => void +} + +export function ResultsGrid({ + result, + error, + running, + refine +}: { + result?: QueryResult + error?: DbSafeError + running: boolean + refine?: ResultsGridRefine +}): React.JSX.Element { + const scrollRef = useRef(null) + const rows = result?.rows ?? [] + const virtualizer = useVirtualizer({ + count: rows.length, + getScrollElement: () => scrollRef.current, + estimateSize: () => ROW_HEIGHT, + overscan: OVERSCAN, + getItemKey: (index) => index + }) + + if (running) { + return + } + if (error) { + return ( +
+

{error.safeMessage}

+
+ ) + } + if (!result) { + return ( + + ) + } + + const gridTemplate = + result.columns.length > 0 + ? result.columns.map(() => `minmax(${COL_MIN_PX}px, ${COL_MAX_PX}px)`).join(' ') + : '1fr' + // Filtering a wrapped subquery is by name → ambiguous for duplicate names; only + // uniquely-named columns get a filter control. + const nameCounts = new Map() + for (const col of result.columns) { + nameCounts.set(col.name, (nameCounts.get(col.name) ?? 0) + 1) + } + + return ( +
+
+
+
+ {result.columns.map((col, i) => + refine ? ( + refine.onSort(i + 1)} + filter={filterFor(refine.refine.filters, col.name)} + onFilter={(filter) => refine.onFilter(col.name, filter)} + filterable={(nameCounts.get(col.name) ?? 0) === 1} + /> + ) : ( +
+ {col.name} +
+ ) + )} +
+ {rows.length === 0 ? ( +
+ {translate('auto.components.database.ResultsGrid.noRows', 'No rows returned')} +
+ ) : ( +
+ {virtualizer.getVirtualItems().map((item) => { + const row = rows[item.index] + return ( +
+ {result.columns.map((_col, ci) => { + const { text, isNull } = formatCell(row?.[ci]) + return ( +
+ {text} +
+ ) + })} +
+ ) + })} +
+ )} +
+
+
+ + {translate('auto.components.database.ResultsGrid.rowCount', '{{count}} rows', { + count: result.rowCount + })} + + + {translate('auto.components.database.ResultsGrid.durationMs', '{{ms}} ms', { + ms: result.durationMs + })} + + {refine?.refine.engaged ? ( +
+ + + {translate('auto.components.database.ResultsGrid.pageRange', '{{from}}–{{to}}', { + from: result.rowCount === 0 ? 0 : refine.refine.offset + 1, + to: refine.refine.offset + result.rowCount + })} + + +
+ ) : result.truncated ? ( + + {translate( + 'auto.components.database.ResultsGrid.truncated', + 'Truncated — showing the first rows' + )} + + ) : null} +
+
+ ) +} + +function GridPlaceholder({ text, spinning }: { text: string; spinning?: boolean }): React.JSX.Element { + return ( +
+ {spinning ? : null} +

{text}

+
+ ) +} diff --git a/src/renderer/src/components/database/SchemaTree.tsx b/src/renderer/src/components/database/SchemaTree.tsx new file mode 100644 index 00000000000..128f6725c22 --- /dev/null +++ b/src/renderer/src/components/database/SchemaTree.tsx @@ -0,0 +1,458 @@ +import React, { useCallback, useEffect, useMemo, useRef, useState } from 'react' +import { useVirtualizer } from '@tanstack/react-virtual' +import { + ChevronRight, + Columns3, + Database, + Eye, + KeyRound, + Loader2, + RefreshCw, + Table2 +} from 'lucide-react' +import { Button } from '@/components/ui/button' +import { useAppStore } from '@/store' +import { dbColumnKey } from '@/store/slices/database' +import { translate } from '@/i18n/i18n' +import { buildSchemaRows, type NodeLoadState, type SchemaTreeRow } from './schema-tree-rows' + +const ROW_HEIGHT = 26 +const OVERSCAN = 12 +const INDENT_PX = 14 +const BASE_PAD_PX = 8 + +type RootState = 'idle' | 'loading' | 'error' + +// Toggle a value in a Set immutably (new Set so React sees the change). +function toggleSet(set: ReadonlySet, value: string): Set { + const next = new Set(set) + if (!next.delete(value)) { + next.add(value) + } + return next +} + +export function SchemaTree({ connectionId }: { connectionId: string }): React.JSX.Element { + const cache = useAppStore((s) => s.dbSchemaCache[connectionId]) + const status = useAppStore((s) => s.dbStatuses[connectionId]?.status ?? 'idle') + const loadDbSchemas = useAppStore((s) => s.loadDbSchemas) + const loadDbSchemaTables = useAppStore((s) => s.loadDbSchemaTables) + const loadDbTableColumns = useAppStore((s) => s.loadDbTableColumns) + const openDbTableTab = useAppStore((s) => s.openDbTableTab) + + const [expandedSchemas, setExpandedSchemas] = useState>(new Set()) + const [expandedTables, setExpandedTables] = useState>(new Set()) + const [nodeState, setNodeState] = useState>({}) + const [rootState, setRootState] = useState('idle') + const [selectedIndex, setSelectedIndex] = useState(0) + + const scrollRef = useRef(null) + + const runRootLoad = useCallback(() => { + setRootState('loading') + void loadDbSchemas(connectionId).then((result) => { + setRootState(result.ok ? 'idle' : 'error') + }) + }, [connectionId, loadDbSchemas]) + + // Load the schema list once the connection is live and nothing is cached yet. + useEffect(() => { + if (status === 'connected' && !cache && rootState === 'idle') { + runRootLoad() + } + }, [status, cache, rootState, runRootLoad]) + + const refresh = useCallback(() => { + setExpandedSchemas(new Set()) + setExpandedTables(new Set()) + setNodeState({}) + runRootLoad() + }, [runRootLoad]) + + // Expand a schema (lazy-loading its tables) or collapse it. + const toggleSchema = useCallback( + (schema: string) => { + // Derive willExpand before updating state: the early-return guard needs the + // current snapshot. Reading expandedSchemas after setExpandedSchemas would + // see the stale closure value, not the post-update state. + const willExpand = !expandedSchemas.has(schema) + setExpandedSchemas((prev) => toggleSet(prev, schema)) + // Skip fetch when collapsing, already cached, or an in-flight fetch is + // pending — nodeState 'loading' guard prevents a double-click double-fetch. + if (!willExpand || cache?.tables[schema] || nodeState[schema] === 'loading') { + return + } + setNodeState((prev) => ({ ...prev, [schema]: 'loading' })) + void loadDbSchemaTables(connectionId, schema).then((result) => { + setNodeState((prev) => { + const next = { ...prev } + if (result.ok) { + delete next[schema] + } else { + next[schema] = 'error' + } + return next + }) + }) + }, + [cache, connectionId, expandedSchemas, loadDbSchemaTables, nodeState] + ) + + const toggleTable = useCallback( + (schema: string, table: string) => { + const key = dbColumnKey(schema, table) + setExpandedTables((prev) => toggleSet(prev, key)) + if (expandedTables.has(key) || cache?.columns[key]) { + return + } + setNodeState((prev) => ({ ...prev, [key]: 'loading' })) + void loadDbTableColumns(connectionId, schema, table).then((result) => { + setNodeState((prev) => { + const next = { ...prev } + if (result.ok) { + delete next[key] + } else { + next[key] = 'error' + } + return next + }) + }) + }, + [cache, connectionId, expandedTables, loadDbTableColumns] + ) + + const rows: SchemaTreeRow[] = useMemo( + () => (cache ? buildSchemaRows(cache, expandedSchemas, expandedTables, nodeState) : []), + [cache, expandedSchemas, expandedTables, nodeState] + ) + + const virtualizer = useVirtualizer({ + count: rows.length, + getScrollElement: () => scrollRef.current, + estimateSize: () => ROW_HEIGHT, + overscan: OVERSCAN, + getItemKey: (index) => rows[index]?.key ?? index + }) + + // Clamp selectedIndex to the visible row count when rows shrink after a collapse + // so the cursor never points past the end of the list. + useEffect(() => { + if (rows.length > 0 && selectedIndex >= rows.length) { + setSelectedIndex(rows.length - 1) + } + }, [rows.length, selectedIndex]) + + // Scroll the virtualizer when selection changes — kept separate from the + // setSelectedIndex updater because React state updaters must be pure (no side + // effects). Calling scrollToIndex inside an updater would run during reconcile. + useEffect(() => { + if (rows.length > 0) { + virtualizer.scrollToIndex(selectedIndex, { align: 'auto' }) + } + }, [selectedIndex, rows.length, virtualizer]) + + // Expand/collapse a schema or table (tables/columns lazy-load on first open). + // Bound to the chevron and the Arrow keys so it stays separate from a row's + // primary action. + const toggleExpand = useCallback( + (row: SchemaTreeRow) => { + if (row.type === 'schema') { + toggleSchema(row.schema) + } else if (row.type === 'table') { + toggleTable(row.schema, row.table.name) + } + }, + [toggleSchema, toggleTable] + ) + + // A row's primary action (click / Enter): schemas expand; a table/view opens + // (or focuses) its Data tab in the workspace. + const activateRow = useCallback( + (row: SchemaTreeRow) => { + if (row.type === 'schema') { + toggleSchema(row.schema) + } else if (row.type === 'table') { + openDbTableTab(connectionId, row.schema, row.table.name) + } + }, + [connectionId, openDbTableTab, toggleSchema] + ) + + const moveSelection = useCallback( + (delta: number) => { + setSelectedIndex((prev) => Math.max(0, Math.min(rows.length - 1, prev + delta))) + }, + [rows.length] + ) + + const handleKeyDown = useCallback( + (event: React.KeyboardEvent) => { + const row = rows[selectedIndex] + switch (event.key) { + case 'ArrowDown': + event.preventDefault() + moveSelection(1) + break + case 'ArrowUp': + event.preventDefault() + moveSelection(-1) + break + case 'ArrowRight': + if (row && (row.type === 'schema' || row.type === 'table') && !row.expanded) { + event.preventDefault() + toggleExpand(row) + } + break + case 'ArrowLeft': + if (row) { + event.preventDefault() + if ((row.type === 'schema' || row.type === 'table') && row.expanded) { + // Collapse the expanded node. + toggleExpand(row) + } else if (row.depth > 0) { + // Navigate to the nearest ancestor row (first row above with a + // lower depth value — e.g. table → schema, column → table). + let parentIdx = selectedIndex - 1 + while (parentIdx >= 0 && rows[parentIdx].depth >= row.depth) { + parentIdx-- + } + if (parentIdx >= 0) { + setSelectedIndex(parentIdx) + } + } + } + break + case 'Enter': + if (row) { + event.preventDefault() + activateRow(row) + } + break + default: + break + } + }, + [rows, selectedIndex, moveSelection, activateRow, toggleExpand] + ) + + if (status !== 'connected') { + return + } + if (rootState === 'loading' && !cache) { + return + } + if (rootState === 'error' && !cache) { + return ( + + ) + } + + return ( +
+
+ + {translate('auto.components.database.SchemaTree.title', 'Schema')} + + +
+ {rows.length === 0 ? ( + + ) : ( +
+
+ {virtualizer.getVirtualItems().map((virtualItem) => { + const row = rows[virtualItem.index] + return ( +
+ setSelectedIndex(virtualItem.index)} + onActivate={() => activateRow(row)} + onToggleExpand={() => toggleExpand(row)} + /> +
+ ) + })} +
+
+ )} +
+ ) +} + +function disconnectedMessage(status: string): string { + return status === 'lost' + ? translate('auto.components.database.SchemaTree.lost', 'Connection lost — reconnect to browse') + : translate('auto.components.database.SchemaTree.notConnected', 'Connect to browse the schema') +} + +function TreePlaceholder({ + text, + spinning, + onRetry +}: { + text: string + spinning?: boolean + onRetry?: () => void +}): React.JSX.Element { + return ( +
+ {spinning ? : null} +

{text}

+ {onRetry ? ( + + ) : null} +
+ ) +} + +function SchemaRowView({ + row, + selected, + onSelect, + onActivate, + onToggleExpand +}: { + row: SchemaTreeRow + selected: boolean + onSelect: () => void + onActivate: () => void + onToggleExpand: () => void +}): React.JSX.Element { + const padLeft = BASE_PAD_PX + row.depth * INDENT_PX + const selectedClass = selected ? 'bg-accent' : 'hover:bg-accent/50' + return ( +
+ +
+ ) +} + +function RowContent({ + row, + onToggleExpand +}: { + row: SchemaTreeRow + onToggleExpand: () => void +}): React.JSX.Element { + if (row.type === 'schema') { + return ( + <> + + + {row.schema} + + ) + } + if (row.type === 'table') { + const Icon = row.table.kind === 'view' ? Eye : Table2 + return ( + <> + + + {row.table.name} + + ) + } + if (row.type === 'column') { + return ( + <> + + + {row.column.name} + {row.column.isPrimaryKey ? ( +