diff --git a/Cargo.lock b/Cargo.lock index d7d2501..af5c4a0 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -349,6 +349,19 @@ dependencies = [ "syn 3.0.3", ] +[[package]] +name = "asynchronous-codec" +version = "0.6.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4057f2c32adbb2fc158e22fb38433c8e9bbf76b75a4732c7c0cbaf695fb65568" +dependencies = [ + "bytes", + "futures-sink", + "futures-util", + "memchr", + "pin-project-lite", +] + [[package]] name = "atk" version = "0.18.2" @@ -384,6 +397,29 @@ version = "1.5.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53" +[[package]] +name = "aws-lc-rs" +version = "1.17.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "00bdb5da18dac48ca2cc7cd4a98e533e8635a58e2361d13a1a4ee3888e0d72f1" +dependencies = [ + "aws-lc-sys", + "zeroize", +] + +[[package]] +name = "aws-lc-sys" +version = "0.43.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "43103168cc76fe62678a375e722fc9cb3a0146159ac5828bc4f0dfd755c2224c" +dependencies = [ + "cc", + "cmake", + "dunce", + "fs_extra", + "pkg-config", +] + [[package]] name = "axum" version = "0.8.9" @@ -468,7 +504,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "144e573728da132683b9488acd528274c790e07fc06ff81ee29f9d8f8b1041e0" dependencies = [ "blowfish 0.10.0", - "pbkdf2", + "pbkdf2 0.13.0", "sha2 0.11.0", ] @@ -757,6 +793,8 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5add81bb678e6cb321aff7fa0dc7689ad82b112dbc032cea19f91d6b8e3582b9" dependencies = [ "find-msvc-tools", + "jobserver", + "libc", "shlex", ] @@ -883,10 +921,12 @@ dependencies = [ "futures-util", "hex", "mysql_async", + "oracle-rs", "prost", "quick-xml 0.37.5", "russh", "rustix", + "rustls 0.23.42", "serde", "serde_json", "sha1 0.10.7", @@ -894,7 +934,10 @@ dependencies = [ "sqlparser", "tempfile", "thiserror 2.0.19", + "tiberius", "tokio", + "tokio-postgres", + "tokio-postgres-rustls", "tokio-util", "tracing", "url", @@ -1138,6 +1181,15 @@ version = "1.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c8d4a3bb8b1e0c1050499d1815f5ab16d04f0959b233085fb31653fbfc9d98f9" +[[package]] +name = "cmake" +version = "0.1.58" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c0f78a02292a74a88ac736019ab962ece0bc380e3f977bf72e376c5d78ff0678" +dependencies = [ + "cc", +] + [[package]] name = "cmov" version = "0.5.4" @@ -1169,6 +1221,18 @@ dependencies = [ "crossbeam-utils", ] +[[package]] +name = "connection-string" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "510ca239cf13b7f8d16a2b48f263de7b4f8c566f0af58d901031473c76afb1e3" + +[[package]] +name = "const-oid" +version = "0.9.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c2459377285ad874054d797f3ccebf984978aa39129f6eafde5cdc8315b612f8" + [[package]] name = "const-oid" version = "0.10.2" @@ -1560,17 +1624,41 @@ dependencies = [ "syn 2.0.119", ] +[[package]] +name = "der" +version = "0.7.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e7c1832837b905bbfb5101e07cc24c8deddf52f93225eee6ead5f4d63d53ddcb" +dependencies = [ + "const-oid 0.9.6", + "der_derive", + "flagset", + "pem-rfc7468 0.7.0", + "zeroize", +] + [[package]] name = "der" version = "0.8.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a69dedd701da44b0536442edf09c81a64b0ab97a7a4a5e3d1971f00027cbc63d" dependencies = [ - "const-oid", - "pem-rfc7468", + "const-oid 0.10.2", + "pem-rfc7468 1.0.0", "zeroize", ] +[[package]] +name = "der_derive" +version = "0.7.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8034092389675178f570469e6c3b0465d3d30b4505c294a6550db47f3c17ad18" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "deranged" version = "0.5.8" @@ -1652,7 +1740,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f1dd6dbb5841937940781866fa1281a1ff7bd3bf827091440879f9994983d5c2" dependencies = [ "block-buffer 0.12.1", - "const-oid", + "const-oid 0.10.2", "crypto-common 0.2.2", "ctutils", ] @@ -1820,12 +1908,12 @@ version = "0.17.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c0681a4fc24c767085329728d8dfba959af91228aa4610cca4f8ce317ba46ae0" dependencies = [ - "der", + "der 0.8.1", "digest 0.11.3", "elliptic-curve", "rfc6979", "signature", - "spki", + "spki 0.8.0", "zeroize", ] @@ -1835,7 +1923,7 @@ version = "3.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "29fcf32e6c73d1079f83ab4d782de2d81620346a5f38c6237a86a22f8368980a" dependencies = [ - "pkcs8", + "pkcs8 0.11.0", "signature", ] @@ -1875,8 +1963,8 @@ dependencies = [ "group", "hkdf 0.13.0", "hybrid-array", - "pem-rfc7468", - "pkcs8", + "pem-rfc7468 1.0.0", + "pkcs8 0.11.0", "rand_core 0.10.1", "sec1", "subtle", @@ -2010,6 +2098,12 @@ dependencies = [ "pin-project-lite", ] +[[package]] +name = "fallible-iterator" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4443176a9f2c162692bd3d352d745ef9413eec5782a80d8fd6f8a1ac692a07f7" + [[package]] name = "fallible-iterator" version = "0.3.0" @@ -2075,6 +2169,12 @@ version = "0.5.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1d674e81391d1e1ab681a28d99df07927c6d4aa5b027d7da16ba32d1d21ecd99" +[[package]] +name = "flagset" +version = "0.4.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b7ac824320a75a52197e8f2d787f6a38b6718bb6897a35142d749af3c0e8f4fe" + [[package]] name = "flate2" version = "1.1.9" @@ -2161,6 +2261,12 @@ dependencies = [ "winapi", ] +[[package]] +name = "fs_extra" +version = "1.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c" + [[package]] name = "futf" version = "0.1.5" @@ -2703,6 +2809,17 @@ dependencies = [ "digest 0.11.3", ] +[[package]] +name = "hostname" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "617aaa3557aef3810a6369d0a99fac8a080891b68bd9f9812a1eeda0c0730cbd" +dependencies = [ + "cfg-if", + "libc", + "windows-link 0.2.1", +] + [[package]] name = "html5ever" version = "0.29.1" @@ -2818,11 +2935,11 @@ dependencies = [ "http", "hyper", "hyper-util", - "rustls", + "rustls 0.23.42", "tokio", - "tokio-rustls", + "tokio-rustls 0.26.4", "tower-service", - "webpki-roots", + "webpki-roots 1.0.9", ] [[package]] @@ -3190,6 +3307,16 @@ dependencies = [ "syn 2.0.119", ] +[[package]] +name = "jobserver" +version = "0.1.35" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1c00acbd29eabad4a2392fa0e921c874934dbbf4194312ad20f04a0ed67a3cb3" +dependencies = [ + "getrandom 0.4.3", + "libc", +] + [[package]] name = "js-sys" version = "0.3.103" @@ -3476,6 +3603,26 @@ version = "0.3.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4facc753ae494aeb6e3c22f839b158aebd4f9270f55cd3c79906c45476c47ab4" +[[package]] +name = "md-5" +version = "0.10.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d89e7ee0cfbedfc4da3340218492196241d89eefb6dab27de5df917a6d2e78cf" +dependencies = [ + "cfg-if", + "digest 0.10.7", +] + +[[package]] +name = "md-5" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "69b6441f590336821bb897fb28fc622898ccceb1d6cea3fde5ea86b090c4de98" +dependencies = [ + "cfg-if", + "digest 0.11.3", +] + [[package]] name = "md5" version = "0.8.1" @@ -3549,7 +3696,7 @@ dependencies = [ "hybrid-array", "kem", "module-lattice", - "pkcs8", + "pkcs8 0.11.0", "rand_core 0.10.1", "sha3 0.11.0", ] @@ -3645,16 +3792,16 @@ dependencies = [ "mysql_common", "percent-encoding", "rand 0.10.2", - "rustls", + "rustls 0.23.42", "serde", "socket2", "thiserror 2.0.19", "tokio", - "tokio-rustls", + "tokio-rustls 0.26.4", "tokio-util", "twox-hash", "url", - "webpki-roots", + "webpki-roots 1.0.9", ] [[package]] @@ -4054,6 +4201,15 @@ dependencies = [ "objc2-core-foundation", ] +[[package]] +name = "objc2-system-configuration" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7216bd11cbda54ccabcab84d523dc93b858ec75ecfb3a7d89513fa22464da396" +dependencies = [ + "objc2-core-foundation", +] + [[package]] name = "objc2-ui-kit" version = "0.3.2" @@ -4120,12 +4276,57 @@ dependencies = [ "libc", ] +[[package]] +name = "openssl-probe" +version = "0.1.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d05e27ee213611ffe7d6348b942e8f942b37114c00cc03cec254295a4a17852e" + +[[package]] +name = "openssl-probe" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7c87def4c32ab89d880effc9e097653c8da5d6ef28e6b539d313baaacfbafcbe" + [[package]] name = "option-ext" version = "0.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "04744f49eae99ab78e0d5c0b603ab218f515ea8cfe5a456d7629ad883a3b6e7d" +[[package]] +name = "oracle-rs" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5c9ca364b441f92b717658c62e85207158b30285deb01811ccd923017f61612e" +dependencies = [ + "aes 0.8.4", + "async-trait", + "bytes", + "cbc 0.1.2", + "chrono", + "hex", + "hmac 0.12.1", + "hostname", + "indexmap 2.14.0", + "md-5 0.10.6", + "pbkdf2 0.12.2", + "pkcs8 0.10.2", + "rand 0.8.7", + "rustls 0.23.42", + "rustls-pemfile 2.2.0", + "rustls-pki-types", + "serde", + "serde_json", + "sha1 0.10.7", + "sha2 0.10.9", + "thiserror 1.0.69", + "tokio", + "tokio-rustls 0.26.4", + "tracing", + "webpki-roots 0.26.11", +] + [[package]] name = "ordered-stream" version = "0.2.0" @@ -4272,6 +4473,16 @@ version = "0.2.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2ee67f1008b1ba2321834326597b8e186293b049a023cdef258527550b9935b4" +[[package]] +name = "pbkdf2" +version = "0.12.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8ed6a7761f76e3b9f92dfb0a60a6a6477c61024b775147ff0973a02653abaf2" +dependencies = [ + "digest 0.10.7", + "hmac 0.12.1", +] + [[package]] name = "pbkdf2" version = "0.13.0" @@ -4282,6 +4493,15 @@ dependencies = [ "hmac 0.13.0", ] +[[package]] +name = "pem-rfc7468" +version = "0.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "88b39c9bfcfc231068454382784bb460aae594343fb030d46e9f50a645418412" +dependencies = [ + "base64ct", +] + [[package]] name = "pem-rfc7468" version = "1.0.0" @@ -4514,8 +4734,23 @@ version = "0.8.0-rc.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "986d2e952779af96ea048f160fd9194e1751b4faea78bcf3ceb456efe008088e" dependencies = [ - "der", - "spki", + "der 0.8.1", + "spki 0.8.0", +] + +[[package]] +name = "pkcs5" +version = "0.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e847e2c91a18bfa887dd028ec33f2fe6f25db77db3619024764914affe8b69a6" +dependencies = [ + "aes 0.8.4", + "cbc 0.1.2", + "der 0.7.10", + "pbkdf2 0.12.2", + "scrypt 0.11.0", + "sha2 0.10.9", + "spki 0.7.3", ] [[package]] @@ -4527,12 +4762,24 @@ dependencies = [ "aes 0.9.2", "aes-gcm 0.11.0", "cbc 0.2.1", - "der", - "pbkdf2", + "der 0.8.1", + "pbkdf2 0.13.0", "rand_core 0.10.1", - "scrypt", + "scrypt 0.12.0", "sha2 0.11.0", - "spki", + "spki 0.8.0", +] + +[[package]] +name = "pkcs8" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f950b2377845cebe5cf8b5165cb3cc1a5e0fa5cfa3e1f7f55707d8fd82e0a7b7" +dependencies = [ + "der 0.7.10", + "pkcs5 0.7.1", + "rand_core 0.6.4", + "spki 0.7.3", ] [[package]] @@ -4541,10 +4788,10 @@ version = "0.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "451913da69c775a56034ea8d9003d27ee8948e12443eae7c038ba100a4f21cb7" dependencies = [ - "der", - "pkcs5", + "der 0.8.1", + "pkcs5 0.8.1", "rand_core 0.10.1", - "spki", + "spki 0.8.0", ] [[package]] @@ -4627,6 +4874,39 @@ dependencies = [ "universal-hash 0.6.1", ] +[[package]] +name = "postgres-protocol" +version = "0.6.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08808e3c483c46e999108051c78334f473d5adb59d78bb80a1268c7e6aa6c514" +dependencies = [ + "base64 0.22.1", + "byteorder", + "bytes", + "fallible-iterator 0.2.0", + "hmac 0.13.0", + "md-5 0.11.0", + "memchr", + "rand 0.10.2", + "sha2 0.11.0", + "stringprep", +] + +[[package]] +name = "postgres-types" +version = "0.2.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "851ca9db4932932d69f3ea811b1abe63087a0f740a47692619dd40d4899b68be" +dependencies = [ + "bytes", + "chrono", + "fallible-iterator 0.2.0", + "postgres-protocol", + "serde_core", + "serde_json", + "uuid", +] + [[package]] name = "potential_utf" version = "0.1.5" @@ -4657,6 +4937,12 @@ version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "925383efa346730478fb4838dbe9137d2a47675ad789c546d150a6e1dd4ab31c" +[[package]] +name = "pretty-hex" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c6fa0831dd7cc608c38a5e323422a0077678fa5744aa2be4ad91c4ece8eec8d5" + [[package]] name = "prettyplease" version = "0.2.37" @@ -4928,7 +5214,7 @@ dependencies = [ "quinn-proto", "quinn-udp", "rustc-hash", - "rustls", + "rustls 0.23.42", "socket2", "thiserror 2.0.19", "tokio", @@ -4949,7 +5235,7 @@ dependencies = [ "rand_pcg 0.10.2", "ring", "rustc-hash", - "rustls", + "rustls 0.23.42", "rustls-pki-types", "slab", "thiserror 2.0.19", @@ -5245,14 +5531,14 @@ dependencies = [ "percent-encoding", "pin-project-lite", "quinn", - "rustls", + "rustls 0.23.42", "rustls-pki-types", "serde", "serde_json", "serde_urlencoded", "sync_wrapper", "tokio", - "tokio-rustls", + "tokio-rustls 0.26.4", "tokio-util", "tower", "tower-http 0.6.11", @@ -5262,7 +5548,7 @@ dependencies = [ "wasm-bindgen-futures", "wasm-streams", "web-sys", - "webpki-roots", + "webpki-roots 1.0.9", ] [[package]] @@ -5356,16 +5642,16 @@ version = "0.10.0-rc.18" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "30b2aa4ba0d89f73d1e332df05be0eeab8840351c36ca5654341dfdb57bb3caf" dependencies = [ - "const-oid", + "const-oid 0.10.2", "crypto-bigint", "crypto-primes", "digest 0.11.3", "pkcs1", - "pkcs8", + "pkcs8 0.11.0", "rand_core 0.10.1", "sha2 0.11.0", "signature", - "spki", + "spki 0.8.0", "zeroize", ] @@ -5376,7 +5662,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "165ca6e57b20e1351573e3729b958bc62f0e48025386970b6e4d29e7a7e71f3f" dependencies = [ "bitflags 2.13.1", - "fallible-iterator", + "fallible-iterator 0.3.0", "fallible-streaming-iterator", "hashlink", "libsqlite3-sys", @@ -5401,7 +5687,7 @@ dependencies = [ "curve25519-dalek", "data-encoding", "delegate", - "der", + "der 0.8.1", "digest 0.11.3", "ecdsa", "ed25519-dalek", @@ -5425,10 +5711,10 @@ dependencies = [ "p384", "p521", "pageant", - "pbkdf2", + "pbkdf2 0.13.0", "pkcs1", - "pkcs5", - "pkcs8", + "pkcs5 0.8.1", + "pkcs8 0.11.0", "polyval 0.7.3", "rand 0.10.2", "rand_core 0.10.1", @@ -5436,14 +5722,14 @@ dependencies = [ "rsa", "russh-cryptovec", "russh-util", - "salsa20", - "scrypt", + "salsa20 0.11.0", + "scrypt 0.12.0", "sec1", "sha1 0.11.0", "sha2 0.11.0", "sha3 0.12.0", "signature", - "spki", + "spki 0.8.0", "ssh-encoding", "ssh-key", "subtle", @@ -5516,20 +5802,76 @@ dependencies = [ "rustix", ] +[[package]] +name = "rustls" +version = "0.21.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f56a14d1f48b391359b22f731fd4bd7e43c97f3c50eee276f3aa09c94784d3e" +dependencies = [ + "log", + "ring", + "rustls-webpki 0.101.7", + "sct", +] + [[package]] name = "rustls" version = "0.23.42" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3c54fcab019b409d04215d3a17cb438fd7fbf192ee61461f20f4fe18704bc138" dependencies = [ + "aws-lc-rs", + "log", "once_cell", "ring", "rustls-pki-types", - "rustls-webpki", + "rustls-webpki 0.103.13", "subtle", "zeroize", ] +[[package]] +name = "rustls-native-certs" +version = "0.6.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a9aace74cb666635c918e9c12bc0d348266037aa8eb599b5cba565709a8dff00" +dependencies = [ + "openssl-probe 0.1.6", + "rustls-pemfile 1.0.4", + "schannel", + "security-framework 2.11.1", +] + +[[package]] +name = "rustls-native-certs" +version = "0.8.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dab5152771c58876a2146916e53e35057e1a4dfa2b9df0f0305b07f611fdea4d" +dependencies = [ + "openssl-probe 0.2.1", + "rustls-pki-types", + "schannel", + "security-framework 3.7.0", +] + +[[package]] +name = "rustls-pemfile" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1c74cae0a4cf6ccbbf5f359f08efdf8ee7e1dc532573bf0db71968cb56b1448c" +dependencies = [ + "base64 0.21.7", +] + +[[package]] +name = "rustls-pemfile" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dce314e5fee3f39953d46bb63bb8a46d40c2f8fb7cc5a3b6cab2bde9721d6e50" +dependencies = [ + "rustls-pki-types", +] + [[package]] name = "rustls-pki-types" version = "1.15.1" @@ -5540,12 +5882,23 @@ dependencies = [ "zeroize", ] +[[package]] +name = "rustls-webpki" +version = "0.101.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b6275d1ee7a1cd780b64aca7726599a1dbc893b1e64144529e55c3c2f745765" +dependencies = [ + "ring", + "untrusted", +] + [[package]] name = "rustls-webpki" version = "0.103.13" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "61c429a8649f110dddef65e2a5ad240f747e85f7758a6bccc7e5777bd33f756e" dependencies = [ + "aws-lc-rs", "ring", "rustls-pki-types", "untrusted", @@ -5563,6 +5916,15 @@ version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" +[[package]] +name = "salsa20" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "97a22f5af31f73a954c10289c93e8a50cc23d971e80ee446f1f6f7137a088213" +dependencies = [ + "cipher 0.4.4", +] + [[package]] name = "salsa20" version = "0.11.0" @@ -5588,6 +5950,15 @@ version = "0.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ece8e78b2f38ec51c51f5d475df0a7187ba5111b2a28bdc761ee05b075d40a71" +[[package]] +name = "schannel" +version = "0.1.29" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91c1b7e4904c873ef0710c1f407dde2e6287de2bebc1bbbf7d430bb7cbffd939" +dependencies = [ + "windows-sys 0.61.2", +] + [[package]] name = "schemars" version = "0.8.22" @@ -5665,6 +6036,17 @@ version = "1.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" +[[package]] +name = "scrypt" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0516a385866c09368f0b5bcd1caff3366aace790fcd46e2bb032697bb172fd1f" +dependencies = [ + "pbkdf2 0.12.2", + "salsa20 0.10.2", + "sha2 0.10.9", +] + [[package]] name = "scrypt" version = "0.12.0" @@ -5672,11 +6054,21 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d87af57419b594aa23fa95f09f0e06d80d84ba01c26148c43844cad6ff4485f0" dependencies = [ "cfg-if", - "pbkdf2", - "salsa20", + "pbkdf2 0.13.0", + "salsa20 0.11.0", "sha2 0.11.0", ] +[[package]] +name = "sct" +version = "0.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "da046153aa2352493d6cb7da4b6e5c0c057d8a1d0a9aa8560baffdd945acd414" +dependencies = [ + "ring", + "untrusted", +] + [[package]] name = "sec1" version = "0.8.1" @@ -5685,7 +6077,7 @@ checksum = "d56d437c2f19203ce5f7122e507831de96f3d2d4d3be5af44a0b0a09d8a80e4d" dependencies = [ "base16ct", "ctutils", - "der", + "der 0.8.1", "hybrid-array", "subtle", "zeroize", @@ -6188,6 +6580,16 @@ version = "0.9.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3763264f6b73151db08c50ff20d7d8a0b8796e021cdea7ceedad07b80155fa0e" +[[package]] +name = "spki" +version = "0.7.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d91ed6c858b01f942cd56b37a94b3e0a1798290327d1236e4d9cf4eaca44d29d" +dependencies = [ + "base64ct", + "der 0.7.10", +] + [[package]] name = "spki" version = "0.8.0" @@ -6195,7 +6597,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1d9efca8738c78ee9484207732f728b1ef517bbb1833d6fc0879ca898a522f6f" dependencies = [ "base64ct", - "der", + "der 0.8.1", ] [[package]] @@ -6243,7 +6645,7 @@ dependencies = [ "crypto-bigint", "ctutils", "digest 0.11.3", - "pem-rfc7468", + "pem-rfc7468 1.0.0", "zeroize", ] @@ -6347,6 +6749,17 @@ dependencies = [ "quote", ] +[[package]] +name = "stringprep" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7b4df3d392d81bd458a8a621b8bffbd2302a12ffe288a9d931670948749463b1" +dependencies = [ + "unicode-bidi", + "unicode-normalization", + "unicode-properties", +] + [[package]] name = "strsim" version = "0.11.1" @@ -6873,6 +7286,34 @@ dependencies = [ "cfg-if", ] +[[package]] +name = "tiberius" +version = "0.12.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a1446cb4198848d1562301a3340424b4f425ef79f35ef9ee034769a9dd92c10d" +dependencies = [ + "async-trait", + "asynchronous-codec", + "byteorder", + "bytes", + "chrono", + "connection-string", + "encoding_rs", + "enumflags2", + "futures-util", + "num-traits", + "once_cell", + "pin-project-lite", + "pretty-hex", + "rustls-native-certs 0.6.3", + "rustls-pemfile 1.0.4", + "thiserror 1.0.69", + "tokio-rustls 0.24.1", + "tokio-util", + "tracing", + "uuid", +] + [[package]] name = "time" version = "0.3.54" @@ -6928,6 +7369,27 @@ version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" +[[package]] +name = "tls_codec" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0de2e01245e2bb89d6f05801c564fa27624dbd7b1846859876c7dad82e90bf6b" +dependencies = [ + "tls_codec_derive", + "zeroize", +] + +[[package]] +name = "tls_codec_derive" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2d2e76690929402faae40aebdda620a2c0e25dd6d3b9afe48867dfd95991f4bd" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "tokio" version = "1.53.1" @@ -6956,13 +7418,64 @@ dependencies = [ "syn 2.0.119", ] +[[package]] +name = "tokio-postgres" +version = "0.7.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a528f7d280f6d5b9cd149635c8705b0dd049754bc67d81d31fa25169a93809d3" +dependencies = [ + "async-trait", + "byteorder", + "bytes", + "fallible-iterator 0.2.0", + "futures-channel", + "futures-util", + "log", + "parking_lot", + "percent-encoding", + "phf 0.13.1", + "pin-project-lite", + "postgres-protocol", + "postgres-types", + "rand 0.10.2", + "socket2", + "tokio", + "tokio-util", + "whoami", +] + +[[package]] +name = "tokio-postgres-rustls" +version = "0.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4c2ad44aa0ae96db89c4742212ed41645b2f597311ff6e1945542a4d9fadc2fb" +dependencies = [ + "rustls 0.23.42", + "rustls-native-certs 0.8.4", + "sha2 0.11.0", + "tokio", + "tokio-postgres", + "tokio-rustls 0.26.4", + "x509-cert", +] + +[[package]] +name = "tokio-rustls" +version = "0.24.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c28327cf380ac148141087fbfb9de9d7bd4e84ab5d2c28fbc911d753de8a7081" +dependencies = [ + "rustls 0.21.12", + "tokio", +] + [[package]] name = "tokio-rustls" version = "0.26.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1729aa945f29d91ba541258c8df89027d5792d85a8841fb65e8bf0f4ede4ef61" dependencies = [ - "rustls", + "rustls 0.23.42", "tokio", ] @@ -6985,6 +7498,7 @@ checksum = "494815d09bf52b5548659851081238f0ca39ff638363907596da739561c62c52" dependencies = [ "bytes", "futures-core", + "futures-io", "futures-sink", "futures-util", "libc", @@ -7346,12 +7860,33 @@ version = "2.9.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "dbc4bc3a9f746d862c45cb89d705aa10f187bb96c76001afab07a0d35ce60142" +[[package]] +name = "unicode-bidi" +version = "0.3.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5c1cb5db39152898a79168971543b1cb5020dff7fe43c8dc468b0885f5e29df5" + [[package]] name = "unicode-ident" version = "1.0.24" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" +[[package]] +name = "unicode-normalization" +version = "0.1.25" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5fd4f6878c9cb28d874b009da9e8d183b5abc80117c40bbd187a1fde336be6e8" +dependencies = [ + "tinyvec", +] + +[[package]] +name = "unicode-properties" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7df058c713841ad818f1dc5d3fd88063241cc61f49f5fbea4b951e8cf5a8d71d" + [[package]] name = "unicode-segmentation" version = "1.13.3" @@ -7557,6 +8092,15 @@ version = "0.11.1+wasi-snapshot-preview1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b" +[[package]] +name = "wasi" +version = "0.14.7+wasi-0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "883478de20367e224c0090af9cf5f9fa85bed63a95c1abf3afc5c083ebc06e8c" +dependencies = [ + "wasip2", +] + [[package]] name = "wasip2" version = "1.0.4+wasi-0.2.12" @@ -7566,6 +8110,15 @@ dependencies = [ "wit-bindgen", ] +[[package]] +name = "wasite" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "66fe902b4a6b8028a753d5424909b764ccf79b7a209eac9bf97e59cda9f71a42" +dependencies = [ + "wasi 0.14.7+wasi-0.2.4", +] + [[package]] name = "wasm-bindgen" version = "0.2.126" @@ -7770,6 +8323,15 @@ dependencies = [ "system-deps", ] +[[package]] +name = "webpki-roots" +version = "0.26.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "521bc38abb08001b01866da9f51eb7c5d647a19260e00054a8c7fd5f9e57f7a9" +dependencies = [ + "webpki-roots 1.0.9", +] + [[package]] name = "webpki-roots" version = "1.0.9" @@ -7815,6 +8377,19 @@ dependencies = [ "windows-core 0.61.2", ] +[[package]] +name = "whoami" +version = "2.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "998767ef88740d1f5b0682a9c53c24431453923962269c2db68ee43788c5a40d" +dependencies = [ + "libc", + "libredox", + "objc2-system-configuration", + "wasite", + "web-sys", +] + [[package]] name = "winapi" version = "0.3.9" @@ -8463,6 +9038,18 @@ dependencies = [ "windows-version", ] +[[package]] +name = "x509-cert" +version = "0.2.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1301e935010a701ae5f8655edc0ad17c44bad3ac5ce8c39185f75453b720ae94" +dependencies = [ + "const-oid 0.9.6", + "der 0.7.10", + "spki 0.7.3", + "tls_codec", +] + [[package]] name = "xdg-home" version = "1.3.0" diff --git a/Cargo.toml b/Cargo.toml index 04cdc2c..bdd564b 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -43,6 +43,7 @@ http-body-util = "0.1" hex = "0.4.3" keyring = { version = "3.6.3", default-features = false } mysql_async = { version = "=0.37.0", default-features = false, features = ["default-rustls-ring"] } +oracle-rs = "=0.1.7" prost = "0.14" quick-xml = "0.37.5" rand = "0.9.2" @@ -50,6 +51,7 @@ reqwest = { version = "0.12.24", default-features = false, features = ["json", " rmcp = { version = "=2.2.0", default-features = false, features = ["server", "macros", "transport-io", "elicitation", "schemars"] } rusqlite = { version = "0.37", features = ["bundled"] } russh = { version = "0.62.5", default-features = false, features = ["ring", "rsa"] } +rustls = { version = "0.23", default-features = false, features = ["ring", "std", "tls12"] } rustix = { version = "1", features = ["fs", "process"] } serde = { version = "1", features = ["derive"] } serde_json = "1" @@ -64,7 +66,10 @@ tauri-plugin-opener = "=2.2.7" thiserror = "2" tokio = { version = "1", features = ["fs", "io-std", "io-util", "macros", "net", "process", "rt-multi-thread", "signal", "sync", "time"] } tokio-stream = { version = "0.1", features = ["sync"] } -tokio-util = { version = "0.7.16", features = ["io", "rt"] } +tokio-postgres = { version = "=0.7.18", features = ["with-chrono-0_4", "with-serde_json-1", "with-uuid-1"] } +tokio-postgres-rustls = { version = "=0.14.0", features = ["native-certs", "ring"] } +tokio-util = { version = "0.7.16", features = ["compat", "io", "rt"] } +tiberius = { version = "=0.12.3", default-features = false, features = ["chrono", "rustls", "tds73"] } tower = { version = "0.5", features = ["util"] } tower-http = { version = "0.7", features = ["fs", "trace"] } tracing = "0.1" diff --git a/README.md b/README.md index 128d04a..e8677bd 100644 --- a/README.md +++ b/README.md @@ -2,11 +2,12 @@ Source-available implementation of the Chat2DB Community hybrid runtime. -Chat2DB Rust owns the product runtime in Rust, uses a native Rust path for the -current MySQL browser and Console data plane, and retains the broader public Chat2DB -Community database compatibility layer behind a supervised Java process. The -repository is under active development and is not yet a stable end-user -release. +Chat2DB Rust owns the product runtime in Rust, uses native Rust paths for the +MySQL, PostgreSQL, SQL Server, and Oracle browser and Console data planes, and +uses a Rust-owned DM adapter over the generic JDBC bridge. The broader public +Chat2DB Community database compatibility layer remains behind a supervised Java +process. The repository is under active development and is not yet a stable +end-user release. ## Clone @@ -119,7 +120,19 @@ surface reached by the pinned Community frontend: - complete native MySQL datasource lifecycle, SSH tunneling, portability, metadata, editable DML and DDL, views and routines, transfer tasks, account administration, schema diff, pins, ER layout, workspace persistence, and - SQLite-backed Dashboard/Chart CRUD with native read-only chart refresh; and + SQLite-backed Dashboard/Chart CRUD with native read-only chart refresh; +- native PostgreSQL through `tokio-postgres 0.7.18`, with connection and SSH, + retained queries and typed parameters, cancellation and limits, Console, + relational and programmability metadata, table DDL and preview, ER metadata, + and native schema/namespace/DML builders; +- native SQL Server through `tiberius 0.12.3`, with the same relational + workbench slice, TDS-aware result handling, direct-batch Console semantics, + conservative write cancellation, and native schema/namespace/DML builders; +- native Oracle through the pure-Rust `oracle-rs 0.1.7` protocol client, with + connection and SSH, retained queries, Console, metadata, DDL, preview, ER, + and native schema/namespace/DML builders without OCI, ODPI-C, JDBC, or Java; + unsupported lossy Oracle result types fail closed instead of fabricating + values; and - the pinned Community AI workspace routes plus confirmed Agent, CLI, and MCP writes, with explicit approval, read-only enforcement, single-statement validation, and conservative unknown-outcome handling. @@ -161,16 +174,33 @@ output, `CHART` operation history, fixture cleanup, and Java dormant. The complete `rtk make verify` gate then passed with the Dashboard/Chart increment included. +On 2026-08-07 the added native PostgreSQL, SQL Server, and Oracle paths passed +real product verticals against PostgreSQL 17, Azure SQL Edge's SQL Server TDS +endpoint, and Oracle Free 23 respectively. The verticals cover connection, +metadata, retained query, Console, preview, DDL/dialect behavior, read-only +enforcement, cancellation, bounded values, cleanup, and continuous proof that +Java remained dormant. Core passed `357` all-target tests with `9` ignored plus +strict Clippy under both the default toolchain and the minimum Rust `1.88.0`. +For Oracle, `oracle-rs 0.1.7` has a fixed 100-row initial prefetch whose legacy +decoder cannot safely expose `BINARY_FLOAT`, `BINARY_DOUBLE`, `ROWID`, or +`UROWID`; those result columns return `oracle_result_type_not_supported` when +the driver exposes their type, and the project does not fork or vendor the +upstream crate. + Stage 6 and the Stage 7A-7M foundations are complete. Web and desktop own the product runtime and publish its owner-only local endpoint; CLI and MCP attach to that host and never contact Java directly. The pinned Community frontend's complete MySQL workbench surface is mapped through the shared Axum/Tauri legacy dispatcher. Native MySQL connections, metadata, Console, mutations, transfer, class generation, accounts, schema diff, chart refresh, and workspace operations -remain in Rust and do not acquire a Java lease. Dashboard and chart documents -remain in SQLite. Community parser, formatter, completion, SQL-builder, and -exact plugin compatibility operations remain Java-backed and start the -supervised process only on demand. +remain in Rust and do not acquire a Java lease. PostgreSQL, SQL Server, and +Oracle connection, relational workbench, and dialect-builder operations also +remain in Rust when their explicit native driver ids are persisted. Existing +managed JDBC datasources continue through Java; registering a native driver +never silently changes their execution engine. Dashboard and chart documents +remain in SQLite. Community parser, formatter, completion, and exact plugin +compatibility operations remain Java-backed and start the supervised process +only on demand. The Console compatibility path uses SQLite migrations 3 and 4 for saved Consoles and durable execution history. Historical `/api/operation/saved/*`, diff --git a/apps/chat2db-desktop/src/lib.rs b/apps/chat2db-desktop/src/lib.rs index 5fbf309..b04eb14 100644 --- a/apps/chat2db-desktop/src/lib.rs +++ b/apps/chat2db-desktop/src/lib.rs @@ -3154,7 +3154,7 @@ mod tests { } #[tokio::test] - async fn table_preview_command_maps_unavailable_engine_errors() { + async fn table_preview_command_maps_unavailable_storage_errors() { let error = start_community_table_preview_for( &Application::new(), StartCommunityTablePreviewRequest { @@ -3167,9 +3167,9 @@ mod tests { }, ) .await - .expect_err("table preview without an engine must fail"); + .expect_err("table preview without storage must fail"); - assert_eq!(error.code, "database_engine_unavailable"); + assert_eq!(error.code, "storage_unavailable"); } #[tokio::test] diff --git a/apps/chat2db-web/src/lib.rs b/apps/chat2db-web/src/lib.rs index 7bbf5b7..ab4f73e 100644 --- a/apps/chat2db-web/src/lib.rs +++ b/apps/chat2db-web/src/lib.rs @@ -238,11 +238,21 @@ mod tests { assert_eq!(response.status(), StatusCode::OK); let inventory: JdbcDriverList = response_json(response).await; - assert_eq!(inventory.items.len(), 1); - let mysql = &inventory.items[0]; - assert_eq!(mysql.driver_id, "mysql"); - assert_eq!(mysql.driver_class, "rust:mysql_async"); - assert_eq!(mysql.artifact_count, 0); + assert_eq!(inventory.items.len(), 4); + for (driver_id, driver_class) in [ + ("mysql", "rust:mysql_async"), + ("postgresql", "rust:tokio-postgres"), + ("sqlserver", "rust:tiberius"), + ("oracle", "rust:oracle-rs"), + ] { + let driver = inventory + .items + .iter() + .find(|driver| driver.driver_id == driver_id) + .expect("every built-in native connection driver must be advertised"); + assert_eq!(driver.driver_class, driver_class); + assert_eq!(driver.artifact_count, 0); + } } #[tokio::test] @@ -959,7 +969,7 @@ mod tests { "rowLimit": 200 }), ), - "database_engine_unavailable", + "storage_unavailable", ), ( json_request( diff --git a/crates/chat2db-core/Cargo.toml b/crates/chat2db-core/Cargo.toml index f2afeb9..f84dba1 100644 --- a/crates/chat2db-core/Cargo.toml +++ b/crates/chat2db-core/Cargo.toml @@ -22,11 +22,14 @@ chat2db-storage = { path = "../chat2db-storage" } chrono.workspace = true csv.workspace = true directories.workspace = true +futures-util.workspace = true hex.workspace = true mysql_async.workspace = true +oracle-rs.workspace = true prost.workspace = true quick-xml.workspace = true russh.workspace = true +rustls.workspace = true serde.workspace = true serde_json.workspace = true sha1.workspace = true @@ -36,16 +39,16 @@ rustix.workspace = true tempfile = "3" thiserror.workspace = true tokio.workspace = true +tokio-postgres.workspace = true +tokio-postgres-rustls.workspace = true tokio-util.workspace = true +tiberius.workspace = true tracing.workspace = true url.workspace = true uuid.workspace = true xls.workspace = true zip.workspace = true -[dev-dependencies] -futures-util.workspace = true - [[test]] name = "java_h2_product" path = "tests/java_h2_product.rs" diff --git a/crates/chat2db-core/src/community.rs b/crates/chat2db-core/src/community.rs index ceeb0c2..db1a431 100644 --- a/crates/chat2db-core/src/community.rs +++ b/crates/chat2db-core/src/community.rs @@ -1,4 +1,4 @@ -use std::future::Future; +use std::{future::Future, sync::Arc}; use base64::{Engine as _, engine::general_purpose::STANDARD}; use chat2db_contract::{ @@ -99,10 +99,48 @@ use crate::{ AppError, Application, datasource_session::{SessionReadOnly, open_datasource_session, resolve_datasource_connection}, engine_manager::EngineLease, - native_driver::native_capability_not_supported, + native_driver::{NativeDriver, native_capability_not_supported}, + storage_call, }; impl Application { + async fn native_driver_for_community_datasource( + &self, + datasource_id: &str, + database_type: &str, + ) -> Result>, AppError> { + let storage = self.require_storage()?; + let datasource_id = datasource_id.to_owned(); + let datasource = storage_call(move || storage.get_datasource(&datasource_id)) + .await? + .ok_or_else(|| { + AppError::not_found("datasource_not_found", "The datasource does not exist") + })?; + let Some(driver) = self.native_driver_for_datasource_driver_id(&datasource.driver_id) + else { + return Ok(None); + }; + let database_type_matches = driver + .descriptor() + .database_types + .iter() + .any(|candidate| candidate.eq_ignore_ascii_case(database_type.trim())); + if database_type_matches { + return Ok(Some(driver)); + } + if driver + .descriptor() + .id + .eq_ignore_ascii_case(datasource.driver_id.trim()) + { + return Err(AppError::invalid( + "datasource_database_type_mismatch", + "The datasource native Rust driver does not support the requested database type", + )); + } + Ok(None) + } + /// Lists the plugins discovered from the fixed Community classpath. /// /// # Errors @@ -128,7 +166,10 @@ impl Application { &self, request: ListCommunitySchemasRequest, ) -> Result { - if let Some(driver) = self.native_driver_for_database_type(&request.database_type) { + if let Some(driver) = self + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? + { driver.metadata().ok_or_else(|| { native_capability_not_supported(&request.database_type, "schema metadata") })?; @@ -173,7 +214,10 @@ impl Application { &self, request: ListCommunityDatabasesRequest, ) -> Result { - if let Some(driver) = self.native_driver_for_database_type(&request.database_type) { + if let Some(driver) = self + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? + { driver.metadata().ok_or_else(|| { native_capability_not_supported(&request.database_type, "database metadata") })?; @@ -217,7 +261,10 @@ impl Application { &self, request: ListCommunityTablesRequest, ) -> Result { - if let Some(driver) = self.native_driver_for_database_type(&request.database_type) { + if let Some(driver) = self + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? + { driver.metadata().ok_or_else(|| { native_capability_not_supported(&request.database_type, "table metadata") })?; @@ -271,7 +318,10 @@ impl Application { &self, request: ListCommunityColumnsRequest, ) -> Result { - if let Some(driver) = self.native_driver_for_database_type(&request.database_type) { + if let Some(driver) = self + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? + { driver.metadata().ok_or_else(|| { native_capability_not_supported(&request.database_type, "column metadata") })?; @@ -377,7 +427,10 @@ impl Application { &self, request: ListCommunityIndexesRequest, ) -> Result { - if let Some(driver) = self.native_driver_for_database_type(&request.database_type) { + if let Some(driver) = self + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? + { let metadata = driver.metadata().ok_or_else(|| { native_capability_not_supported(&request.database_type, "metadata") })?; @@ -430,7 +483,10 @@ impl Application { &self, request: ListCommunityViewsRequest, ) -> Result { - if let Some(driver) = self.native_driver_for_database_type(&request.database_type) { + if let Some(driver) = self + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? + { let metadata = driver.metadata().ok_or_else(|| { native_capability_not_supported(&request.database_type, "metadata") })?; @@ -484,7 +540,10 @@ impl Application { request: ListCommunityViewsRequest, ) -> Result { let view_name = request.view_name_pattern.clone(); - if let Some(driver) = self.native_driver_for_database_type(&request.database_type) { + if let Some(driver) = self + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? + { let metadata = driver.metadata().ok_or_else(|| { native_capability_not_supported(&request.database_type, "metadata") })?; @@ -515,7 +574,10 @@ impl Application { &self, request: ListCommunityTableKeysRequest, ) -> Result { - if let Some(driver) = self.native_driver_for_database_type(&request.database_type) { + if let Some(driver) = self + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? + { let metadata = driver.metadata().ok_or_else(|| { native_capability_not_supported(&request.database_type, "metadata") })?; @@ -568,7 +630,10 @@ impl Application { &self, request: ListCommunityTableKeysRequest, ) -> Result { - if let Some(driver) = self.native_driver_for_database_type(&request.database_type) { + if let Some(driver) = self + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? + { let metadata = driver.metadata().ok_or_else(|| { native_capability_not_supported(&request.database_type, "metadata") })?; @@ -621,7 +686,10 @@ impl Application { &self, request: ListCommunityTableKeysRequest, ) -> Result { - if let Some(driver) = self.native_driver_for_database_type(&request.database_type) { + if let Some(driver) = self + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? + { let metadata = driver.metadata().ok_or_else(|| { native_capability_not_supported(&request.database_type, "metadata") })?; @@ -674,7 +742,10 @@ impl Application { &self, request: ListCommunityFunctionsRequest, ) -> Result { - if let Some(driver) = self.native_driver_for_database_type(&request.database_type) { + if let Some(driver) = self + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? + { let metadata = driver.metadata().ok_or_else(|| { native_capability_not_supported(&request.database_type, "metadata") })?; @@ -719,7 +790,10 @@ impl Application { &self, request: GetCommunityFunctionRequest, ) -> Result { - if let Some(driver) = self.native_driver_for_database_type(&request.database_type) { + if let Some(driver) = self + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? + { let metadata = driver.metadata().ok_or_else(|| { native_capability_not_supported(&request.database_type, "metadata") })?; @@ -770,7 +844,10 @@ impl Application { &self, request: GetCommunityFunctionRequest, ) -> Result { - if let Some(driver) = self.native_driver_for_database_type(&request.database_type) { + if let Some(driver) = self + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? + { let metadata = driver.metadata().ok_or_else(|| { native_capability_not_supported(&request.database_type, "metadata") })?; @@ -826,7 +903,10 @@ impl Application { &self, request: ListCommunityProceduresRequest, ) -> Result { - if let Some(driver) = self.native_driver_for_database_type(&request.database_type) { + if let Some(driver) = self + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? + { let metadata = driver.metadata().ok_or_else(|| { native_capability_not_supported(&request.database_type, "metadata") })?; @@ -871,7 +951,10 @@ impl Application { &self, request: GetCommunityProcedureRequest, ) -> Result { - if let Some(driver) = self.native_driver_for_database_type(&request.database_type) { + if let Some(driver) = self + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? + { let metadata = driver.metadata().ok_or_else(|| { native_capability_not_supported(&request.database_type, "metadata") })?; @@ -922,7 +1005,10 @@ impl Application { &self, request: GetCommunityProcedureRequest, ) -> Result { - if let Some(driver) = self.native_driver_for_database_type(&request.database_type) { + if let Some(driver) = self + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? + { let metadata = driver.metadata().ok_or_else(|| { native_capability_not_supported(&request.database_type, "metadata") })?; @@ -978,8 +1064,16 @@ impl Application { &self, request: PreviewCommunityRoutineInvocationRequest, ) -> Result { + self.native_driver_for_database_type(&request.database_type) + .ok_or_else(|| { + AppError::invalid( + "invalid_community_routine_invocation_request", + "routine invocation preview requires a native Rust driver", + ) + })?; let driver = self - .native_driver_for_database_type(&request.database_type) + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? .ok_or_else(|| { AppError::invalid( "invalid_community_routine_invocation_request", @@ -1041,7 +1135,8 @@ impl Application { request: CommunityRoutineMigrationRequest, ) -> Result { let driver = self - .native_driver_for_database_type(&request.database_type) + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? .ok_or_else(|| { AppError::invalid( "invalid_community_routine_migration_request", @@ -1070,7 +1165,10 @@ impl Application { &self, request: ListCommunityTriggersRequest, ) -> Result { - if let Some(driver) = self.native_driver_for_database_type(&request.database_type) { + if let Some(driver) = self + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? + { let metadata = driver.metadata().ok_or_else(|| { native_capability_not_supported(&request.database_type, "metadata") })?; @@ -1115,7 +1213,10 @@ impl Application { &self, request: GetCommunityTriggerRequest, ) -> Result { - if let Some(driver) = self.native_driver_for_database_type(&request.database_type) { + if let Some(driver) = self + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? + { let metadata = driver.metadata().ok_or_else(|| { native_capability_not_supported(&request.database_type, "metadata") })?; @@ -1258,7 +1359,8 @@ impl Application { )); } if self - .native_driver_for_database_type(&request.database_type) + .native_driver_for_community_datasource(&request.datasource_id, &request.database_type) + .await? .is_some() { let database_type = request.database_type.clone(); @@ -2397,9 +2499,10 @@ mod tests { CommunityPrimaryKey, CommunityProcedure, CommunityProcedureParameter, CommunitySchema, CommunitySqlAnalysis, CommunitySqlDiagnostic, CommunitySqlValidation, CommunityTable, CommunityTableColumn, CommunityTableIndex, CommunityTableIndexColumn, CommunityTrigger, - DatasourceConnection, ListCommunityColumnsRequest, ListCommunityDatabasesRequest, - ListCommunityIndexesRequest, ListCommunitySchemasRequest, ListCommunityTableKeysRequest, - ListCommunityTablesRequest, ListCommunityViewsRequest, + CreateDatasourceRequest, DatasourceConnection, ListCommunityColumnsRequest, + ListCommunityDatabasesRequest, ListCommunityIndexesRequest, ListCommunitySchemasRequest, + ListCommunityTableKeysRequest, ListCommunityTablesRequest, ListCommunityViewsRequest, + StartCommunityTablePreviewRequest, }; use chat2db_java_bridge::{ CommunityDatabase as BridgeCommunityDatabase, @@ -3457,44 +3560,149 @@ mod tests { ); let storage = Storage::open(directory.path(), vault).expect("test storage must open"); let application = Application::with_storage(storage); - let database_engine_error = application + let database_error = application .list_community_databases(database_request()) .await - .expect_err("database metadata requires the engine"); - assert_eq!(database_engine_error.kind(), AppErrorKind::Unavailable); - assert_eq!( - database_engine_error.api_error().code, - "database_engine_unavailable" - ); - let table_engine_error = application + .expect_err("database metadata requires an existing datasource"); + assert_eq!(database_error.kind(), AppErrorKind::NotFound); + assert_eq!(database_error.api_error().code, "datasource_not_found"); + let table_error = application .list_community_tables(table_request()) .await - .expect_err("table metadata requires the engine"); - assert_eq!(table_engine_error.kind(), AppErrorKind::Unavailable); - assert_eq!( - table_engine_error.api_error().code, - "database_engine_unavailable" - ); - let column_engine_error = application + .expect_err("table metadata requires an existing datasource"); + assert_eq!(table_error.kind(), AppErrorKind::NotFound); + assert_eq!(table_error.api_error().code, "datasource_not_found"); + let column_error = application .list_community_columns(column_request()) .await - .expect_err("column metadata requires the engine"); - assert_eq!(column_engine_error.kind(), AppErrorKind::Unavailable); + .expect_err("column metadata requires an existing datasource"); + assert_eq!(column_error.kind(), AppErrorKind::NotFound); + assert_eq!(column_error.api_error().code, "datasource_not_found"); + let index_error = application + .list_community_indexes(index_request()) + .await + .expect_err("index metadata requires an existing datasource"); + assert_eq!(index_error.kind(), AppErrorKind::NotFound); + assert_eq!(index_error.api_error().code, "datasource_not_found"); + } + + #[tokio::test] + async fn explicit_native_datasource_rejects_mismatched_community_database_type() { + let directory = TempDir::new().expect("temporary data directory must open"); + let vault = Arc::new( + EncryptedFileVault::new(directory.path(), [0x5d; 32]).expect("test vault must open"), + ); + let storage = Storage::open(directory.path(), vault).expect("test storage must open"); + let application = Application::with_storage(storage); + let datasource = application + .create_datasource(CreateDatasourceRequest { + name: "Native PostgreSQL".to_owned(), + driver_id: "postgresql".to_owned(), + connection: Some(DatasourceConnection { + jdbc_url: "jdbc:postgresql://127.0.0.1:1/app".to_owned(), + properties: Vec::new(), + read_only: false, + ssh: None, + }), + }) + .await + .expect("native datasource must persist"); + + let metadata_error = application + .list_community_databases(ListCommunityDatabasesRequest { + datasource_id: datasource.id.clone(), + database_type: "ORACLE".to_owned(), + }) + .await + .expect_err("mismatched metadata type must fail closed"); + assert_eq!(metadata_error.kind(), AppErrorKind::InvalidRequest); assert_eq!( - column_engine_error.api_error().code, - "database_engine_unavailable" + metadata_error.api_error().code, + "datasource_database_type_mismatch" ); - let index_engine_error = application - .list_community_indexes(index_request()) + + let preview_error = application + .start_community_table_preview(StartCommunityTablePreviewRequest { + datasource_id: datasource.id, + database_type: "ORACLE".to_owned(), + database_name: "app".to_owned(), + schema_name: "APP".to_owned(), + table_name: "items".to_owned(), + row_limit: Some(10), + }) .await - .expect_err("index metadata requires the engine"); - assert_eq!(index_engine_error.kind(), AppErrorKind::Unavailable); + .expect_err("mismatched preview type must fail closed"); + assert_eq!(preview_error.kind(), AppErrorKind::InvalidRequest); assert_eq!( - index_engine_error.api_error().code, - "database_engine_unavailable" + preview_error.api_error().code, + "datasource_database_type_mismatch" ); } + #[tokio::test] + async fn managed_relational_datasources_keep_java_metadata_and_preview_routes() { + let directory = TempDir::new().expect("temporary data directory must open"); + let vault = Arc::new( + EncryptedFileVault::new(directory.path(), [0x6d; 32]).expect("test vault must open"), + ); + let storage = Storage::open(directory.path(), vault).expect("test storage must open"); + let application = Application::with_storage(storage); + + for (driver_class, database_type, jdbc_url) in [ + ( + "org.postgresql.Driver", + "POSTGRESQL", + "jdbc:postgresql://127.0.0.1:1/app", + ), + ( + "com.microsoft.sqlserver.jdbc.SQLServerDriver", + "SQLSERVER", + "jdbc:sqlserver://127.0.0.1:1;databaseName=app;encrypt=false", + ), + ( + "oracle.jdbc.OracleDriver", + "ORACLE", + "jdbc:oracle:thin:@127.0.0.1:1/FREEPDB1", + ), + ] { + let datasource = application + .create_datasource(CreateDatasourceRequest { + name: format!("Managed {database_type}"), + driver_id: driver_class.to_owned(), + connection: Some(DatasourceConnection { + jdbc_url: jdbc_url.to_owned(), + properties: Vec::new(), + read_only: false, + ssh: None, + }), + }) + .await + .expect("legacy managed datasource must persist"); + + let metadata_error = application + .list_community_databases(ListCommunityDatabasesRequest { + datasource_id: datasource.id.clone(), + database_type: database_type.to_owned(), + }) + .await + .expect_err("managed metadata must route to the unavailable Java engine"); + assert_unavailable(&metadata_error, "database_engine_unavailable"); + + let preview_error = application + .start_community_table_preview(StartCommunityTablePreviewRequest { + datasource_id: datasource.id, + database_type: database_type.to_owned(), + database_name: "app".to_owned(), + schema_name: "public".to_owned(), + table_name: "items".to_owned(), + row_limit: Some(10), + }) + .await + .expect_err("managed preview must route to the unavailable Java engine"); + assert_unavailable(&preview_error, "database_engine_unavailable"); + } + } + #[tokio::test] async fn community_relation_services_report_unconfigured_dependencies_safely() { let application = Application::new(); @@ -3529,26 +3737,31 @@ mod tests { let view_error = application .list_community_views(view_request()) .await - .expect_err("view metadata requires the engine"); - assert_unavailable(&view_error, "database_engine_unavailable"); + .expect_err("view metadata requires an existing datasource"); + assert_not_found(&view_error, "datasource_not_found"); for error in [ application .list_community_imported_keys(key_request()) .await - .expect_err("imported-key metadata requires the engine"), + .expect_err("imported-key metadata requires an existing datasource"), application .list_community_exported_keys(key_request()) .await - .expect_err("exported-key metadata requires the engine"), + .expect_err("exported-key metadata requires an existing datasource"), application .list_community_primary_keys(key_request()) .await - .expect_err("primary-key metadata requires the engine"), + .expect_err("primary-key metadata requires an existing datasource"), ] { - assert_unavailable(&error, "database_engine_unavailable"); + assert_not_found(&error, "datasource_not_found"); } } + fn assert_not_found(error: &AppError, code: &str) { + assert_eq!(error.kind(), AppErrorKind::NotFound); + assert_eq!(error.api_error().code, code); + } + fn assert_unavailable(error: &AppError, code: &str) { assert_eq!(error.kind(), AppErrorKind::Unavailable); assert_eq!(error.api_error().code, code); diff --git a/crates/chat2db-core/src/datasource_compatibility.rs b/crates/chat2db-core/src/datasource_compatibility.rs index 8698696..0b0f741 100644 --- a/crates/chat2db-core/src/datasource_compatibility.rs +++ b/crates/chat2db-core/src/datasource_compatibility.rs @@ -287,6 +287,9 @@ impl Application { pub(crate) fn jdbc_driver_from_descriptor(descriptor: &NativeDriverDescriptor) -> JdbcDriver { let display_name = match descriptor.id.to_ascii_lowercase().as_str() { "mysql" => "MySQL".to_owned(), + "oracle" => "Oracle".to_owned(), + "postgresql" => "PostgreSQL".to_owned(), + "sqlserver" => "SQL Server".to_owned(), _ => descriptor .database_types .first() @@ -311,17 +314,22 @@ pub(crate) fn native_driver_for_datasource_driver_id( datasource_driver_id: &str, managed_drivers: &[JdbcDriver], ) -> Option> { - if let Some(driver) = registry.driver_for_datasource_driver_id(datasource_driver_id) { - return Some(driver); - } - let managed_descriptor = managed_drivers.iter().find(|driver| { driver .driver_id .eq_ignore_ascii_case(datasource_driver_id.trim()) - })?; + }); + let Some(managed_descriptor) = managed_descriptor else { + let driver = registry.driver_for_datasource_driver_id(datasource_driver_id)?; + return (driver + .descriptor() + .id + .eq_ignore_ascii_case(datasource_driver_id.trim()) + || driver.can_replace_managed_jdbc_datasource()) + .then_some(driver); + }; let mut matches = registry - .descriptors() + .managed_jdbc_replacement_descriptors() .filter(|descriptor| jdbc_driver_matches_descriptor(managed_descriptor, descriptor)); let driver_id = matches.next()?.id; if matches.next().is_some() { @@ -834,23 +842,129 @@ mod tests { ); } + #[test] + fn native_descriptor_uses_product_display_names_for_new_relational_drivers() { + let descriptor = |id, implementation, database_type| NativeDriverDescriptor { + id, + implementation, + database_types: database_type, + compatibility_aliases: &[], + }; + + for (descriptor, expected_name) in [ + ( + descriptor("postgresql", "tokio-postgres", &["POSTGRESQL"]), + "PostgreSQL (native Rust)", + ), + ( + descriptor("sqlserver", "tiberius", &["SQLSERVER"]), + "SQL Server (native Rust)", + ), + ( + descriptor("oracle", "oracle-rs", &["ORACLE"]), + "Oracle (native Rust)", + ), + ] { + assert_eq!(jdbc_driver_from_descriptor(&descriptor).name, expected_name); + } + } + #[test] fn managed_jdbc_driver_id_resolves_through_the_datasource_compatibility_boundary() { let registry = NativeDriverRegistry::built_in(); - let managed_drivers = vec![JdbcDriver { - pack_id: "mysql-connector-j".to_owned(), - name: "MySQL JDBC".to_owned(), - version: "9".to_owned(), - driver_id: "managed-mysql".to_owned(), - driver_class: "com.mysql.cj.jdbc.Driver".to_owned(), - artifact_count: 1, - artifact_bytes: "1".to_owned(), - }]; + let managed_drivers = vec![ + JdbcDriver { + pack_id: "mysql-connector-j".to_owned(), + name: "MySQL JDBC".to_owned(), + version: "9".to_owned(), + driver_id: "managed-mysql".to_owned(), + driver_class: "com.mysql.cj.jdbc.Driver".to_owned(), + artifact_count: 1, + artifact_bytes: "1".to_owned(), + }, + JdbcDriver { + pack_id: "dm-jdbc".to_owned(), + name: "DM JDBC".to_owned(), + version: "8".to_owned(), + driver_id: "managed-dm".to_owned(), + driver_class: "dm.jdbc.driver.DmDriver".to_owned(), + artifact_count: 1, + artifact_bytes: "1".to_owned(), + }, + ]; let driver = native_driver_for_datasource_driver_id(®istry, "managed-mysql", &managed_drivers) .expect("managed MySQL descriptor resolves to the native implementation"); assert_eq!(driver.descriptor().id, "mysql"); + let driver = + native_driver_for_datasource_driver_id(®istry, "managed-dm", &managed_drivers) + .expect("managed DM descriptor resolves to the opted-in native implementation"); + assert_eq!(driver.descriptor().id, "dm"); + } + + #[test] + fn new_native_drivers_do_not_implicitly_replace_managed_jdbc_datasources() { + let registry = NativeDriverRegistry::built_in(); + let managed_drivers = [ + ("managed-postgresql", "PostgreSQL", "org.postgresql.Driver"), + ( + "managed-sqlserver", + "SQL Server", + "com.microsoft.sqlserver.jdbc.SQLServerDriver", + ), + ("managed-oracle", "Oracle", "oracle.jdbc.OracleDriver"), + ] + .into_iter() + .map(|(driver_id, name, driver_class)| JdbcDriver { + pack_id: format!("{driver_id}-pack"), + name: name.to_owned(), + version: "1".to_owned(), + driver_id: driver_id.to_owned(), + driver_class: driver_class.to_owned(), + artifact_count: 1, + artifact_bytes: "1".to_owned(), + }) + .collect::>(); + + for driver_id in ["managed-postgresql", "managed-sqlserver", "managed-oracle"] { + assert!( + native_driver_for_datasource_driver_id(®istry, driver_id, &managed_drivers) + .is_none(), + "persisted {driver_id} JDBC datasource must keep its managed execution engine" + ); + } + } + + #[test] + fn jdbc_class_aliases_cannot_bypass_replacement_policy_without_managed_inventory() { + let registry = NativeDriverRegistry::built_in(); + for jdbc_class in [ + "org.postgresql.Driver", + "com.microsoft.sqlserver.jdbc.SQLServerDriver", + "oracle.jdbc.OracleDriver", + ] { + assert!( + native_driver_for_datasource_driver_id(®istry, jdbc_class, &[]).is_none(), + "{jdbc_class} must not silently become a native datasource" + ); + } + + for native_driver_id in ["postgresql", "sqlserver", "oracle"] { + assert!( + native_driver_for_datasource_driver_id(®istry, native_driver_id, &[]).is_some(), + "the explicit native driver id {native_driver_id} must remain routable" + ); + } + assert!( + native_driver_for_datasource_driver_id(®istry, "com.mysql", &[],).is_some(), + "MySQL explicitly opts into managed JDBC replacement" + ); + assert!( + native_driver_for_datasource_driver_id(®istry, "dm.jdbc.driver.DmDriver", &[],) + .is_some(), + "DM explicitly opts into managed JDBC replacement" + ); } #[test] diff --git a/crates/chat2db-core/src/lib.rs b/crates/chat2db-core/src/lib.rs index cb0e52f..dbc53c7 100644 --- a/crates/chat2db-core/src/lib.rs +++ b/crates/chat2db-core/src/lib.rs @@ -23,7 +23,10 @@ mod native_dm; mod native_driver; mod native_driver_types; mod native_mysql; +mod native_oracle; +mod native_postgres; mod native_schema_diff_types; +mod native_sqlserver; mod operation; mod query; mod ssh; diff --git a/crates/chat2db-core/src/native_driver.rs b/crates/chat2db-core/src/native_driver.rs index 1265dd9..c5164fe 100644 --- a/crates/chat2db-core/src/native_driver.rs +++ b/crates/chat2db-core/src/native_driver.rs @@ -31,7 +31,10 @@ use crate::{ TablePreviewAccepted, TablePreviewRequest, TriggerList, TriggerMetadata, ViewList, }, native_mysql, + native_oracle::OracleNativeDriver, + native_postgres::PostgresNativeDriver, native_schema_diff_types::{SchemaDiffRequest, SchemaDiffSql}, + native_sqlserver::SqlServerNativeDriver, operation::CancellationRequest, query::{ DatabaseWriteError, NativeConsoleRequest, NativeConsoleResult, PreparedQuery, @@ -348,6 +351,14 @@ pub(crate) trait NativeSchemaDiffDriver: Send + Sync { pub(crate) trait NativeDriver: Send + Sync { fn descriptor(&self) -> &'static NativeDriverDescriptor; + /// Whether this driver may transparently replace a persisted managed JDBC datasource. + /// + /// Native drivers default to an explicit opt-in migration boundary so adding a Rust driver + /// cannot silently change the execution engine of existing JDBC datasources. + fn can_replace_managed_jdbc_datasource(&self) -> bool { + false + } + fn connection(&self) -> Option<&dyn NativeConnectionDriver> { None } @@ -393,8 +404,14 @@ pub(crate) struct NativeDriverRegistry { impl NativeDriverRegistry { pub(crate) fn built_in() -> Self { - Self::try_new(vec![Arc::new(MysqlNativeDriver), Arc::new(DmNativeDriver)]) - .expect("built-in native drivers must have unique identities") + Self::try_new(vec![ + Arc::new(MysqlNativeDriver), + Arc::new(DmNativeDriver), + Arc::new(PostgresNativeDriver), + Arc::new(SqlServerNativeDriver), + Arc::new(OracleNativeDriver), + ]) + .expect("built-in native drivers must have unique identities") } pub(crate) fn try_new(drivers: Vec>) -> Result { @@ -457,6 +474,7 @@ impl NativeDriverRegistry { }) } + #[cfg(test)] pub(crate) fn descriptors(&self) -> impl Iterator + '_ { self.drivers.iter().map(|driver| driver.descriptor()) } @@ -470,6 +488,15 @@ impl NativeDriverRegistry { .map(|driver| driver.descriptor()) } + pub(crate) fn managed_jdbc_replacement_descriptors( + &self, + ) -> impl Iterator + '_ { + self.drivers + .iter() + .filter(|driver| driver.can_replace_managed_jdbc_datasource()) + .map(|driver| driver.descriptor()) + } + /// Resolves a persisted datasource driver ID to its native implementation. pub(crate) fn driver_for_datasource_driver_id( &self, @@ -629,6 +656,10 @@ impl NativeDriver for MysqlNativeDriver { &MYSQL_DRIVER_DESCRIPTOR } + fn can_replace_managed_jdbc_datasource(&self) -> bool { + true + } + fn connection(&self) -> Option<&dyn NativeConnectionDriver> { Some(self) } @@ -671,6 +702,10 @@ impl NativeDriver for DmNativeDriver { &native_dm::DM_DRIVER_DESCRIPTOR } + fn can_replace_managed_jdbc_datasource(&self) -> bool { + true + } + fn metadata(&self) -> Option<&dyn NativeMetadataDriver> { Some(self) } @@ -1412,6 +1447,9 @@ mod tests { include_str!("native_driver_types.rs"), ), ("native_mysql.rs", include_str!("native_mysql.rs")), + ("native_oracle.rs", include_str!("native_oracle.rs")), + ("native_postgres.rs", include_str!("native_postgres.rs")), + ("native_sqlserver.rs", include_str!("native_sqlserver.rs")), ( "native_administration_types.rs", include_str!("native_administration_types.rs"), @@ -1544,6 +1582,22 @@ mod tests { ); } + #[test] + fn built_in_registry_exposes_every_owned_database_driver() { + let registry = NativeDriverRegistry::built_in(); + let ids = registry + .descriptors() + .map(|descriptor| descriptor.id) + .collect::>(); + + assert_eq!(ids, ["mysql", "dm", "postgresql", "sqlserver", "oracle"]); + let replacement_ids = registry + .managed_jdbc_replacement_descriptors() + .map(|descriptor| descriptor.id) + .collect::>(); + assert_eq!(replacement_ids, ["mysql", "dm"]); + } + #[test] fn application_dispatches_a_postgres_capability_through_the_registry() { let registry = NativeDriverRegistry::try_new(vec![Arc::new(FakePostgresDriver)]) diff --git a/crates/chat2db-core/src/native_oracle.rs b/crates/chat2db-core/src/native_oracle.rs new file mode 100644 index 0000000..db6cb9e --- /dev/null +++ b/crates/chat2db-core/src/native_oracle.rs @@ -0,0 +1,4806 @@ +use std::{ + collections::BTreeMap, + fmt::Write as _, + future::Future, + mem::size_of, + time::{Duration, Instant}, +}; + +use async_trait::async_trait; +use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64_STANDARD}; +use chat2db_contract::{ + ApiError, ColumnNullability, DatasourceConnection, JdbcValue, JdbcValueType, QueryLimits, + ResultColumn, ResultMetadata, ResultRow, StartQueryRequest, +}; +use chat2db_engine_protocol::wire; +use chat2db_storage::Storage; +use chrono::{DateTime, Datelike, FixedOffset, NaiveDate, NaiveDateTime, NaiveTime, Timelike}; +use oracle_rs::{ + Config, Connection, Error as OracleError, LobData, LobValue, OracleType, + QueryResult as OracleQueryResult, Row as OracleRow, Value as OracleValue, + config::ServiceMethod, + statement::ColumnInfo, + types::{OracleDate, OracleNumber, OracleTimestamp}, +}; +use prost::Message; +use sqlparser::{ast::Statement, dialect::OracleDialect, parser::Parser}; +use tokio::sync::watch; +use tokio_util::sync::CancellationToken; +use url::Url; + +use crate::{ + AppError, AppErrorKind, Application, + datasource_session::{ResolvedDatasourceConnection, resolve_datasource_connection}, + native_driver::{ + NativeConnectionDriver, NativeDialectDriver, NativeDriver, NativeMetadataDriver, + NativeQueryDriver, NativeTableDriver, + }, + native_driver_types::{ + BuiltSql, ColumnList, ColumnMetadata, CreateSchemaSqlRequest, DatabaseDefinition, + DatabaseList, DatabaseMetadata, DmlAssignment, DmlColumn, DmlRow, DmlSqlRequest, + DmlStatement, DmlTarget, DmlTemporalKind, DmlValue, EntityRelationColumn, + EntityRelationForeignKey, EntityRelationTable, ForeignKeyList, ForeignKeyMetadata, + FunctionList, FunctionMetadata, FunctionParameterList, FunctionParameterMetadata, + IndexColumnMetadata, IndexList, IndexMetadata, ListColumnsRequest, ListDatabasesRequest, + ListIndexesRequest, ListRoutinesRequest, ListSchemasRequest, ListTableKeysRequest, + ListTablesRequest, ListTriggersRequest, ListViewsRequest, MetadataObjectRef, MetadataScope, + NamespaceSqlOperation, NamespaceSqlRequest, NativeDriverDescriptor, PrimaryKeyList, + PrimaryKeyMetadata, ProcedureList, ProcedureMetadata, ProcedureParameterList, + ProcedureParameterMetadata, SchemaList, SchemaMetadata, TableList, TableMetadata, + TablePreviewAccepted, TablePreviewRequest, TableRef, TriggerList, TriggerMetadata, + ViewList, + }, + operation::CancellationRequest, + query::{ + DatabaseValue, DatabaseWriteError, NativeConsoleRequest, NativeConsoleResult, + PreparedQuery, QueryExecutionOptions, QueryParameter, QueryTaskError, RetainedWriter, + }, + ssh::{SshTunnel, SshTunnelIdentity}, +}; + +const ORACLE_DATABASE_TYPE: &str = "ORACLE"; +const CONNECT_TIMEOUT: Duration = Duration::from_secs(15); +const OPERATION_TIMEOUT: Duration = Duration::from_secs(30); +const CLOSE_TIMEOUT: Duration = Duration::from_secs(5); +const DEFAULT_PORT: u16 = 1_521; +const DEFAULT_BATCH_ROWS: u32 = 256; +const DEFAULT_BATCH_BYTES: u32 = 256 * 1024; +const DEFAULT_RESULT_BYTES: u64 = wire::JdbcResultByteLimit::DefaultResultBytes as u64; +const MAX_RESULT_BYTES: u64 = wire::JdbcResultByteLimit::MaxResultBytes as u64; +const MAX_BATCH_ROWS: u32 = wire::JdbcProtocolLimit::MaxBatchRows as u32; +const MAX_BATCH_BYTES: u32 = wire::JdbcProtocolLimit::MaxBatchBytes as u32; +const MAX_COLUMNS: usize = wire::JdbcProtocolLimit::MaxColumns as usize; +const MAX_PARAMETERS: usize = wire::JdbcProtocolLimit::MaxParameters as usize; +const MAX_SQL_BYTES: usize = wire::JdbcProtocolLimit::MaxSqlBytes as usize; +const MAX_SCALAR_BYTES: usize = wire::JdbcProtocolLimit::MaxScalarBytes as usize; +const MAX_IDENTIFIER_BYTES: usize = 128; +const MAX_METADATA_ROWS: usize = 100_000; +const MAX_METADATA_RESULT_BYTES: u64 = 16 * 1024 * 1024; +const MAX_CONSOLE_STATEMENTS: usize = 1_000; +const MAX_CONSOLE_PAGE_SIZE: u32 = 10_000; +const MAX_CONSOLE_RESULT_BYTES: u64 = DEFAULT_RESULT_BYTES; +const MAX_CONSOLE_ROWS: u64 = 1_000_000; +const MAX_TABLE_PREVIEW_ROWS: u32 = 1_000; +const FETCH_ROWS: u32 = 1; + +const ORACLE_SYSTEM_SCHEMAS: &[&str] = &[ + "ANONYMOUS", + "AUDSYS", + "CTXSYS", + "DBSNMP", + "DIP", + "DVF", + "DVSYS", + "GGSYS", + "GSMADMIN_INTERNAL", + "LBACSYS", + "MDSYS", + "OJVMSYS", + "OLAPSYS", + "ORDDATA", + "ORDSYS", + "OUTLN", + "SYS", + "SYSBACKUP", + "SYSDG", + "SYSKM", + "SYSTEM", + "WMSYS", + "XDB", +]; + +pub(crate) struct OracleNativeDriver; + +pub(crate) const ORACLE_DRIVER_DESCRIPTOR: NativeDriverDescriptor = NativeDriverDescriptor { + id: "oracle", + implementation: "oracle-rs", + database_types: &["ORACLE"], + compatibility_aliases: &[ + "oracle", + "oracle-rs", + "oracle.jdbc.OracleDriver", + "oracle.jdbc.driver.OracleDriver", + ], +}; + +impl NativeDriver for OracleNativeDriver { + fn descriptor(&self) -> &'static NativeDriverDescriptor { + &ORACLE_DRIVER_DESCRIPTOR + } + + fn connection(&self) -> Option<&dyn NativeConnectionDriver> { + Some(self) + } + + fn query(&self) -> Option<&dyn NativeQueryDriver> { + Some(self) + } + + fn metadata(&self) -> Option<&dyn NativeMetadataDriver> { + Some(self) + } + + fn tables(&self) -> Option<&dyn NativeTableDriver> { + Some(self) + } + + fn dialect(&self) -> Option<&dyn NativeDialectDriver> { + Some(self) + } +} + +impl NativeDialectDriver for OracleNativeDriver { + fn build_create_schema(&self, request: CreateSchemaSqlRequest) -> Result { + build_oracle_create_schema(request) + } + + fn build_namespace_sql(&self, request: NamespaceSqlRequest) -> Result { + build_oracle_namespace_sql(request) + } + + fn build_dml(&self, request: DmlSqlRequest) -> Result { + build_oracle_dml(request) + } +} + +#[async_trait] +impl NativeConnectionDriver for OracleNativeDriver { + async fn test_connection(&self, connection: &DatasourceConnection) -> Result<(), AppError> { + test_connection(connection).await.map(|_| ()) + } + + async fn test_connection_with_local_port( + &self, + connection: &DatasourceConnection, + ) -> Result, AppError> { + test_connection(connection).await + } +} + +#[async_trait] +impl NativeQueryDriver for OracleNativeDriver { + fn is_read_candidate(&self, sql: &str) -> Result { + is_read_candidate(sql) + } + + fn validate_query(&self, query: &PreparedQuery) -> Result<(), AppError> { + validate_query(query) + } + + async fn execute_query_task( + &self, + application: &Application, + operation_id: &str, + cancellation: watch::Receiver, + query: PreparedQuery, + storage: Storage, + resolved: ResolvedDatasourceConnection, + ) -> Result { + execute_query_task( + application, + operation_id, + cancellation, + query, + storage, + resolved, + ) + .await + } + + async fn execute_update( + &self, + resolved: ResolvedDatasourceConnection, + sql: String, + cancellation: CancellationToken, + ) -> Result { + execute_update(resolved, sql, cancellation).await + } + + async fn execute_console( + &self, + application: &Application, + request: NativeConsoleRequest, + cancellation: watch::Receiver, + force_read_only: bool, + ) -> Result, AppError> { + execute_console(application, request, cancellation, force_read_only).await + } +} + +#[async_trait] +impl NativeMetadataDriver for OracleNativeDriver { + async fn list_schemas( + &self, + application: &Application, + request: ListSchemasRequest, + ) -> Result { + list_schemas(application, request).await + } + + async fn list_databases( + &self, + application: &Application, + request: ListDatabasesRequest, + ) -> Result { + list_databases(application, request).await + } + + async fn list_tables( + &self, + application: &Application, + request: ListTablesRequest, + ) -> Result { + list_tables(application, request).await + } + + async fn list_columns( + &self, + application: &Application, + request: ListColumnsRequest, + ) -> Result { + list_columns(application, request).await + } + + async fn list_indexes( + &self, + application: &Application, + request: ListIndexesRequest, + ) -> Result { + list_indexes(application, request).await + } + + async fn list_views( + &self, + application: &Application, + request: ListViewsRequest, + ) -> Result { + list_views(application, request).await + } + + async fn get_view( + &self, + application: &Application, + request: MetadataObjectRef, + ) -> Result { + get_view(application, request).await + } + + async fn list_imported_keys( + &self, + application: &Application, + request: ListTableKeysRequest, + ) -> Result { + list_imported_keys(application, request).await + } + + async fn list_exported_keys( + &self, + application: &Application, + request: ListTableKeysRequest, + ) -> Result { + list_exported_keys(application, request).await + } + + async fn list_primary_keys( + &self, + application: &Application, + request: ListTableKeysRequest, + ) -> Result { + list_primary_keys(application, request).await + } + + async fn list_functions( + &self, + application: &Application, + request: ListRoutinesRequest, + ) -> Result { + list_functions(application, request).await + } + + async fn get_function( + &self, + application: &Application, + request: MetadataObjectRef, + ) -> Result { + get_function(application, request).await + } + + async fn list_function_parameters( + &self, + application: &Application, + request: MetadataObjectRef, + ) -> Result { + list_function_parameters(application, request).await + } + + async fn list_procedures( + &self, + application: &Application, + request: ListRoutinesRequest, + ) -> Result { + list_procedures(application, request).await + } + + async fn get_procedure( + &self, + application: &Application, + request: MetadataObjectRef, + ) -> Result { + get_procedure(application, request).await + } + + async fn list_procedure_parameters( + &self, + application: &Application, + request: MetadataObjectRef, + ) -> Result { + list_procedure_parameters(application, request).await + } + + async fn list_triggers( + &self, + application: &Application, + request: ListTriggersRequest, + ) -> Result { + list_triggers(application, request).await + } + + async fn get_trigger( + &self, + application: &Application, + request: MetadataObjectRef, + ) -> Result { + get_trigger(application, request).await + } +} + +#[async_trait] +impl NativeTableDriver for OracleNativeDriver { + async fn load_er_tables( + &self, + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + ) -> Result, AppError> { + load_er_tables(application, datasource_id, database_name, schema_name).await + } + + async fn validate_column_reorder( + &self, + _application: &Application, + _datasource_id: &str, + _database_name: &str, + _table_name: &str, + _column_names: &[String], + ) -> Result<(), AppError> { + Err(capability_not_supported("physical column reordering")) + } + + async fn table_ddl( + &self, + application: &Application, + datasource_id: &str, + _database_name: &str, + schema_name: &str, + table_name: &str, + ) -> Result { + object_ddl(application, datasource_id, schema_name, table_name, "TABLE").await + } + + async fn start_table_preview( + &self, + application: &Application, + request: TablePreviewRequest, + row_limit: u32, + ) -> Result { + start_table_preview(application, request, row_limit).await + } +} + +struct PreparedOracleConnection { + config: Config, + tunnel: Option, +} + +struct ManagedOracleConnection { + connection: Connection, + tunnel: Option, + local_port: Option, +} + +impl ManagedOracleConnection { + async fn close(self) -> Result<(), AppError> { + let close_result = tokio::time::timeout(CLOSE_TIMEOUT, self.connection.close()).await; + let connection_result = match close_result { + Ok(Ok(())) => Ok(()), + Ok(Err(error)) => Err(oracle_connection_error(&error)), + Err(_) => Err(AppError::unavailable( + "oracle_connection_close_timeout", + "The Oracle connection could not be closed in time", + )), + }; + let tunnel_result = match self.tunnel { + Some(tunnel) => match tokio::time::timeout(CLOSE_TIMEOUT, tunnel.close()).await { + Ok(result) => result, + Err(_) => Err(AppError::unavailable( + "oracle_ssh_tunnel_close_timeout", + "The Oracle SSH tunnel could not be closed in time", + )), + }, + None => Ok(()), + }; + connection_result.and(tunnel_result) + } + + async fn abandon(self) { + self.connection.mark_closed(); + drop(self.connection); + if let Some(tunnel) = self.tunnel + && let Err(error) = tunnel.close().await + { + tracing::warn!(error = %error, "Oracle SSH tunnel cleanup failed after connection abandonment"); + } + } +} + +async fn test_connection(connection: &DatasourceConnection) -> Result, AppError> { + let managed = open_connection(connection, SshTunnelIdentity::Ephemeral).await?; + let local_port = managed.local_port; + let ping = tokio::time::timeout(OPERATION_TIMEOUT, managed.connection.ping()).await; + match ping { + Ok(Ok(())) => { + managed.close().await?; + Ok(local_port) + } + Ok(Err(error)) => { + managed.abandon().await; + Err(oracle_connection_error(&error)) + } + Err(_) => { + managed.abandon().await; + Err(oracle_operation_timeout("connection test")) + } + } +} + +async fn open_resolved_connection( + resolved: &ResolvedDatasourceConnection, +) -> Result { + open_connection( + &resolved.connection, + SshTunnelIdentity::Datasource { + datasource_id: &resolved.datasource_id, + revision: resolved.datasource_revision, + }, + ) + .await +} + +async fn open_connection( + connection: &DatasourceConnection, + identity: SshTunnelIdentity<'_>, +) -> Result { + let prepared = prepare_connection(connection, identity).await?; + let local_port = prepared.tunnel.as_ref().map(SshTunnel::local_port); + let connect = Connection::connect_with_config(prepared.config); + match tokio::time::timeout(CONNECT_TIMEOUT, connect).await { + Ok(Ok(connection)) => Ok(ManagedOracleConnection { + connection, + tunnel: prepared.tunnel, + local_port, + }), + Ok(Err(error)) => { + if let Some(tunnel) = prepared.tunnel + && let Err(close_error) = tunnel.close().await + { + tracing::warn!(error = %close_error, "Oracle SSH tunnel cleanup failed after connection rejection"); + } + Err(oracle_connection_error(&error)) + } + Err(_) => { + if let Some(tunnel) = prepared.tunnel + && let Err(error) = tunnel.close().await + { + tracing::warn!(error = %error, "Oracle SSH tunnel cleanup failed after connection timeout"); + } + Err(AppError::unavailable( + "oracle_connection_timeout", + "The Oracle server did not accept the connection in time", + )) + } + } +} + +async fn prepare_connection( + connection: &DatasourceConnection, + identity: SshTunnelIdentity<'_>, +) -> Result { + let mut config = connection_config(connection)?; + let Some(ssh) = connection.ssh.as_ref() else { + return Ok(PreparedOracleConnection { + config, + tunnel: None, + }); + }; + let tunnel = SshTunnel::open(identity, ssh, config.host.clone(), config.port).await?; + apply_ssh_forward(&mut config, tunnel.local_port()); + Ok(PreparedOracleConnection { + config, + tunnel: Some(tunnel), + }) +} + +fn connection_config(connection: &DatasourceConnection) -> Result { + let (mut config, url_username, url_password, tls) = parse_oracle_url(&connection.jdbc_url)?; + let username = connection_property(connection, &["user", "username"])? + .or(url_username) + .ok_or_else(|| invalid_connection_property("username"))?; + let password = connection_property(connection, &["password"])? + .or(url_password) + .ok_or_else(|| invalid_connection_property("password"))?; + validate_credential(&username, "username")?; + validate_credential(&password, "password")?; + config.set_username(username); + config.set_password(password); + if tls { + ensure_rustls_crypto_provider()?; + config = config + .with_tls() + .map_err(|error| oracle_connection_error(&error))?; + } + Ok(config) +} + +fn ensure_rustls_crypto_provider() -> Result<(), AppError> { + if rustls::crypto::CryptoProvider::get_default().is_some() { + return Ok(()); + } + let _ = rustls::crypto::ring::default_provider().install_default(); + if rustls::crypto::CryptoProvider::get_default().is_some() { + Ok(()) + } else { + Err(AppError::internal()) + } +} + +fn apply_ssh_forward(config: &mut Config, local_port: u16) { + let original_host = config.host.clone(); + if let Some(tls) = config.tls_config.as_mut() { + tls.server_name.get_or_insert(original_host); + } + "127.0.0.1".clone_into(&mut config.host); + config.port = local_port; +} + +fn parse_oracle_url( + jdbc_url: &str, +) -> Result<(Config, Option, Option, bool), AppError> { + let value = jdbc_url.trim(); + if let Some(rest) = value.strip_prefix("jdbc:oracle:thin:@") { + return parse_jdbc_oracle_target(rest); + } + if let Some(rest) = value.strip_prefix("jdbc:oracle:@") { + return parse_jdbc_oracle_target(rest); + } + if value + .get(..9) + .is_some_and(|prefix| prefix.eq_ignore_ascii_case("oracle://")) + { + return parse_oracle_url_target(value); + } + Err(invalid_connection_url()) +} + +fn parse_jdbc_oracle_target( + target: &str, +) -> Result<(Config, Option, Option, bool), AppError> { + let target = target.trim(); + if target.is_empty() || target.starts_with('(') { + return Err(invalid_connection_url()); + } + let (target, query) = target.split_once('?').map_or((target, ""), |parts| parts); + let target = target.trim_start_matches('/'); + if target.is_empty() || (!target.contains('/') && target.matches(':').count() != 2) { + return Err(invalid_connection_url()); + } + let config = target + .parse::() + .map_err(|_| invalid_connection_url())?; + validate_config_target(&config)?; + Ok((config, None, None, tls_requested(query)?)) +} + +fn parse_oracle_url_target( + value: &str, +) -> Result<(Config, Option, Option, bool), AppError> { + let url = Url::parse(value).map_err(|_| invalid_connection_url())?; + if url.fragment().is_some() { + return Err(invalid_connection_url()); + } + let host = url.host_str().ok_or_else(invalid_connection_url)?; + if host.contains(':') { + return Err(AppError::invalid( + "invalid_oracle_connection", + "IPv6 Oracle connection targets are not supported by the selected protocol driver", + )); + } + let service = url.path().trim_matches('/'); + if service.is_empty() || service.contains('/') { + return Err(invalid_connection_url()); + } + validate_connection_component(host, "host", 255)?; + validate_connection_component(service, "serviceName", MAX_IDENTIFIER_BYTES)?; + let port = url.port().unwrap_or(DEFAULT_PORT); + let mut sid = None; + let mut tls = false; + for (key, value) in url.query_pairs() { + if key.eq_ignore_ascii_case("sid") { + if sid.replace(value.into_owned()).is_some() { + return Err(invalid_connection_property("sid")); + } + } else if key.eq_ignore_ascii_case("ssl") || key.eq_ignore_ascii_case("tcps") { + tls = parse_bool(&value).ok_or_else(|| invalid_connection_property(&key))?; + } else { + return Err(invalid_connection_property(&key)); + } + } + let username = (!url.username().is_empty()).then(|| url.username().to_owned()); + let password = url.password().map(ToOwned::to_owned); + let config = if let Some(sid) = sid { + validate_connection_component(&sid, "sid", MAX_IDENTIFIER_BYTES)?; + Config::with_sid(host, port, sid, "", "") + } else { + Config::new(host, port, service, "", "") + }; + Ok((config, username, password, tls)) +} + +fn tls_requested(query: &str) -> Result { + let mut tls = false; + for (key, value) in url::form_urlencoded::parse(query.as_bytes()) { + if key.eq_ignore_ascii_case("ssl") || key.eq_ignore_ascii_case("tcps") { + tls = parse_bool(&value).ok_or_else(|| invalid_connection_property(&key))?; + } else if !key.is_empty() { + return Err(invalid_connection_property(&key)); + } + } + Ok(tls) +} + +fn parse_bool(value: &str) -> Option { + if value.eq_ignore_ascii_case("true") || value == "1" { + Some(true) + } else if value.eq_ignore_ascii_case("false") || value == "0" { + Some(false) + } else { + None + } +} + +fn validate_config_target(config: &Config) -> Result<(), AppError> { + validate_connection_component(&config.host, "host", 255)?; + match &config.service { + ServiceMethod::ServiceName(service) => { + validate_connection_component(service, "serviceName", MAX_IDENTIFIER_BYTES) + } + ServiceMethod::Sid(sid) => validate_connection_component(sid, "sid", MAX_IDENTIFIER_BYTES), + } +} + +fn validate_connection_component( + value: &str, + field: &str, + max_bytes: usize, +) -> Result<(), AppError> { + if value.trim().is_empty() + || value.len() > max_bytes + || value.contains('\0') + || value.chars().any(char::is_control) + { + return Err(invalid_connection_property(field)); + } + Ok(()) +} + +fn validate_credential(value: &str, field: &str) -> Result<(), AppError> { + if value.is_empty() || value.len() > 64 * 1024 || value.contains('\0') { + return Err(invalid_connection_property(field)); + } + Ok(()) +} + +fn connection_property( + connection: &DatasourceConnection, + keys: &[&str], +) -> Result, AppError> { + let mut result = None; + for property in &connection.properties { + if keys + .iter() + .any(|key| property.key.eq_ignore_ascii_case(key)) + && result.replace(property.value.clone()).is_some() + { + return Err(invalid_connection_property(keys[0])); + } + } + Ok(result) +} + +fn invalid_connection_url() -> AppError { + AppError::invalid( + "invalid_oracle_connection", + "A valid jdbc:oracle:thin:@host:port/service or oracle://host:port/service URL is required", + ) +} + +fn invalid_connection_property(property: &str) -> AppError { + AppError::invalid( + "invalid_oracle_connection", + format!("The Oracle connection property {property} is invalid"), + ) +} + +fn oracle_connection_error(error: &OracleError) -> AppError { + match error { + OracleError::InvalidConnectionString(_) => invalid_connection_url(), + OracleError::InvalidCredentials + | OracleError::AuthenticationFailed(_) + | OracleError::InvalidServiceName { .. } + | OracleError::InvalidSid { .. } + | OracleError::ConnectionRefused { .. } + | OracleError::OracleError { .. } + | OracleError::ServerError { .. } => AppError::new( + AppErrorKind::InvalidRequest, + ApiError::new("oracle_connection_rejected", error.to_string()), + ), + OracleError::ProtocolVersionNotSupported(_, _) + | OracleError::FeatureNotSupported(_) + | OracleError::NativeNetworkEncryptionRequired => AppError::new( + AppErrorKind::InvalidRequest, + ApiError::new( + "oracle_protocol_not_supported", + "The native Oracle driver supports Oracle Database 12.1 or newer and does not support this server configuration", + ), + ), + _ => AppError::unavailable( + "oracle_connection_failed", + "The Oracle server could not be reached or the protocol session ended unexpectedly", + ), + } +} + +fn oracle_query_error(error: &OracleError) -> AppError { + match error { + OracleError::OracleError { .. } + | OracleError::ServerError { .. } + | OracleError::SqlError(_) + | OracleError::NoDataFound => AppError::new( + AppErrorKind::InvalidRequest, + ApiError::new("oracle_query_failed", error.to_string()), + ), + OracleError::FeatureNotSupported(_) => AppError::new( + AppErrorKind::InvalidRequest, + ApiError::new("oracle_query_not_supported", error.to_string()), + ), + _ => oracle_connection_error(error), + } +} + +fn oracle_operation_timeout(operation: &str) -> AppError { + AppError::unavailable( + "oracle_operation_timeout", + format!( + "The Oracle {operation} did not complete in time; its dedicated connection was discarded" + ), + ) +} + +fn capability_not_supported(capability: &str) -> AppError { + AppError::invalid( + "native_driver_capability_not_supported", + format!("The Oracle driver does not implement {capability}"), + ) +} + +async fn resolve_native_connection( + application: &Application, + datasource_id: &str, +) -> Result { + let storage = application.require_storage()?; + let resolved = resolve_datasource_connection(&storage, datasource_id).await?; + if application + .native_driver_for_datasource_driver_id(&resolved.driver_id) + .is_none_or(|driver| driver.descriptor().id != ORACLE_DRIVER_DESCRIPTOR.id) + { + return Err(AppError::invalid( + "oracle_driver_mismatch", + "The datasource is not configured with the native Oracle driver", + )); + } + Ok(resolved) +} + +pub(crate) fn is_read_candidate(sql: &str) -> Result { + Ok(matches!( + first_sql_keyword(sql)?.as_deref(), + Some("SELECT" | "WITH") + )) +} + +pub(crate) fn validate_query(query: &PreparedQuery) -> Result<(), AppError> { + validate_read_sql(&query.sql)?; + let _ = oracle_query_parameters(&query.parameters)?; + validate_query_options(query.options) +} + +fn validate_read_sql(sql: &str) -> Result<(), AppError> { + validate_sql_text(sql)?; + if split_oracle_script(sql)?.len() != 1 { + return Err(AppError::invalid( + "oracle_native_query_unsupported", + "Native Oracle read execution accepts exactly one statement", + )); + } + if !is_read_candidate(sql)? { + return Err(AppError::invalid( + "oracle_native_query_unsupported", + "Native Oracle read execution accepts SELECT or WITH queries", + )); + } + let words = oracle_sql_words(sql)?; + if words + .windows(2) + .any(|window| matches!(window, [first, second] if first == "FOR" && second == "UPDATE")) + { + return Err(AppError::invalid( + "oracle_native_query_unsupported", + "Native Oracle read execution does not accept SELECT FOR UPDATE", + )); + } + if let Ok(statements) = Parser::parse_sql(&OracleDialect {}, sql) + && !matches!(statements.as_slice(), [Statement::Query(_)]) + { + return Err(AppError::invalid( + "oracle_native_query_unsupported", + "Native Oracle read execution accepts one query statement", + )); + } + Ok(()) +} + +fn validate_sql_text(sql: &str) -> Result<(), AppError> { + if sql.trim().is_empty() || sql.len() > MAX_SQL_BYTES || sql.contains('\0') { + return Err(AppError::invalid( + "invalid_query_request", + format!("Oracle SQL must be non-empty and at most {MAX_SQL_BYTES} UTF-8 bytes"), + )); + } + Ok(()) +} + +fn validate_query_options(options: QueryExecutionOptions) -> Result<(), AppError> { + if options.target_batch_rows > MAX_BATCH_ROWS { + return Err(AppError::invalid( + "invalid_query_limits", + format!("batchRows must be at most {MAX_BATCH_ROWS}"), + )); + } + if options.target_batch_bytes != 0 + && !(1_024..=MAX_BATCH_BYTES).contains(&options.target_batch_bytes) + { + return Err(AppError::invalid( + "invalid_query_limits", + format!("batchBytes must be zero or between 1024 and {MAX_BATCH_BYTES}"), + )); + } + if options.max_result_bytes > MAX_RESULT_BYTES { + return Err(AppError::invalid( + "invalid_query_limits", + format!("maxResultBytes must be at most {MAX_RESULT_BYTES}"), + )); + } + Ok(()) +} + +fn oracle_query_parameters(parameters: &[QueryParameter]) -> Result, AppError> { + if parameters.len() > MAX_PARAMETERS { + return Err(AppError::invalid( + "invalid_query_parameter_count", + format!("Oracle queries accept at most {MAX_PARAMETERS} parameters"), + )); + } + let mut ordered = parameters.iter().collect::>(); + ordered.sort_unstable_by_key(|parameter| parameter.position); + ordered + .into_iter() + .enumerate() + .map(|(index, parameter)| { + let expected = u32::try_from(index + 1).map_err(|_| AppError::internal())?; + if parameter.position != expected { + return Err(AppError::invalid( + "invalid_query_parameter", + "Oracle parameter positions must be unique and contiguous from 1", + )); + } + oracle_query_value(¶meter.value) + }) + .collect() +} + +fn oracle_query_value(value: &DatabaseValue) -> Result { + match value { + DatabaseValue::Null => Ok(OracleValue::Null), + DatabaseValue::Boolean(value) => Ok(OracleValue::Boolean(*value)), + DatabaseValue::SignedInteger(value) => Ok(OracleValue::Integer(*value)), + DatabaseValue::UnsignedInteger(value) => i64::try_from(*value) + .map(OracleValue::Integer) + .or_else(|_| oracle_decimal_value(&value.to_string())), + DatabaseValue::Float32(value) => Ok(OracleValue::Float(f64::from(*value))), + DatabaseValue::Float64(value) => Ok(OracleValue::Float(*value)), + DatabaseValue::Decimal(value) => oracle_decimal_value(value), + DatabaseValue::Text(value) => oracle_string_value(value, "text"), + DatabaseValue::Binary(value) => { + validate_scalar_bytes(value.len(), "binary")?; + Ok(OracleValue::Bytes(value.clone())) + } + DatabaseValue::Date(value) => oracle_date_value(value), + DatabaseValue::Time(value) => oracle_string_value(value, "time"), + DatabaseValue::Timestamp(value) => oracle_timestamp_value(value), + DatabaseValue::TimestampWithTimeZone(value) => oracle_timestamp_tz_value(value), + DatabaseValue::Json(value) => { + validate_scalar_bytes(value.len(), "JSON")?; + serde_json::from_str(value) + .map(OracleValue::Json) + .map_err(|_| invalid_query_parameter("JSON")) + } + DatabaseValue::Uuid(value) => { + validate_scalar_bytes(value.len(), "UUID")?; + uuid::Uuid::parse_str(value).map_err(|_| invalid_query_parameter("UUID"))?; + Ok(OracleValue::String(value.clone())) + } + } +} + +fn oracle_decimal_value(value: &str) -> Result { + validate_scalar_bytes(value.len(), "decimal")?; + oracle_rs::types::encode_oracle_number(value) + .map_err(|_| invalid_query_parameter("decimal"))?; + Ok(OracleValue::Number(OracleNumber::new(value))) +} + +fn oracle_string_value(value: &str, label: &str) -> Result { + validate_scalar_bytes(value.len(), label)?; + Ok(OracleValue::String(value.to_owned())) +} + +fn oracle_date_value(value: &str) -> Result { + validate_scalar_bytes(value.len(), "date")?; + let date = NaiveDate::parse_from_str(value, "%Y-%m-%d") + .map_err(|_| invalid_query_parameter("date"))?; + Ok(OracleValue::Date(OracleDate::date( + date.year(), + u8::try_from(date.month()).map_err(|_| AppError::internal())?, + u8::try_from(date.day()).map_err(|_| AppError::internal())?, + ))) +} + +fn oracle_timestamp_value(value: &str) -> Result { + validate_scalar_bytes(value.len(), "timestamp")?; + let value = ["%Y-%m-%dT%H:%M:%S%.f", "%Y-%m-%d %H:%M:%S%.f"] + .into_iter() + .find_map(|format| NaiveDateTime::parse_from_str(value, format).ok()) + .ok_or_else(|| invalid_query_parameter("timestamp"))?; + Ok(OracleValue::Timestamp(OracleTimestamp::new( + value.year(), + u8::try_from(value.month()).map_err(|_| AppError::internal())?, + u8::try_from(value.day()).map_err(|_| AppError::internal())?, + u8::try_from(value.hour()).map_err(|_| AppError::internal())?, + u8::try_from(value.minute()).map_err(|_| AppError::internal())?, + u8::try_from(value.second()).map_err(|_| AppError::internal())?, + value.nanosecond() / 1_000, + ))) +} + +fn oracle_timestamp_tz_value(value: &str) -> Result { + validate_scalar_bytes(value.len(), "timestamp with time zone")?; + let value = DateTime::parse_from_rfc3339(value) + .map_err(|_| invalid_query_parameter("timestamp with time zone"))?; + let offset_seconds = value.offset().local_minus_utc(); + let offset_hours = offset_seconds / 3_600; + let offset_minutes = (offset_seconds % 3_600) / 60; + Ok(OracleValue::Timestamp(OracleTimestamp::with_timezone( + value.year(), + u8::try_from(value.month()).map_err(|_| AppError::internal())?, + u8::try_from(value.day()).map_err(|_| AppError::internal())?, + u8::try_from(value.hour()).map_err(|_| AppError::internal())?, + u8::try_from(value.minute()).map_err(|_| AppError::internal())?, + u8::try_from(value.second()).map_err(|_| AppError::internal())?, + value.nanosecond() / 1_000, + i8::try_from(offset_hours) + .map_err(|_| invalid_query_parameter("timestamp with time zone"))?, + i8::try_from(offset_minutes) + .map_err(|_| invalid_query_parameter("timestamp with time zone"))?, + ))) +} + +fn validate_scalar_bytes(size: usize, label: &str) -> Result<(), AppError> { + if size > MAX_SCALAR_BYTES { + return Err(AppError::invalid( + "invalid_query_parameter", + format!("The Oracle {label} parameter exceeds {MAX_SCALAR_BYTES} bytes"), + )); + } + Ok(()) +} + +fn invalid_query_parameter(label: &str) -> AppError { + AppError::invalid( + "invalid_query_parameter", + format!("The Oracle {label} parameter is invalid"), + ) +} + +enum AwaitOutcome { + Completed(T), + Cancelled(Option), +} + +async fn await_with_cancellation( + cancellation: &mut watch::Receiver, + future: F, +) -> AwaitOutcome +where + F: Future, +{ + tokio::pin!(future); + let mut cancellation_open = true; + loop { + tokio::select! { + biased; + changed = cancellation.changed(), if cancellation_open => { + if changed.is_err() { + cancellation_open = false; + continue; + } + let request = { cancellation.borrow().clone() }; + if let CancellationRequest::Requested { reason } = request { + return AwaitOutcome::Cancelled(reason); + } + } + output = &mut future => return AwaitOutcome::Completed(output), + } + } +} + +async fn start_read_only_transaction( + managed: &ManagedOracleConnection, + cancellation: &mut watch::Receiver, +) -> AwaitOutcome> { + // Autonomous transactions are independent of this transaction. Preventing side effects from + // autonomous routines requires a database account whose grants do not permit those writes. + match await_with_cancellation( + cancellation, + tokio::time::timeout( + OPERATION_TIMEOUT, + managed.connection.execute("SET TRANSACTION READ ONLY", &[]), + ), + ) + .await + { + AwaitOutcome::Completed(Ok(Ok(_))) => AwaitOutcome::Completed(Ok(())), + AwaitOutcome::Completed(Ok(Err(error))) => { + AwaitOutcome::Completed(Err(oracle_query_error(&error))) + } + AwaitOutcome::Completed(Err(_)) => { + AwaitOutcome::Completed(Err(oracle_operation_timeout("read-only transaction setup"))) + } + AwaitOutcome::Cancelled(reason) => AwaitOutcome::Cancelled(reason), + } +} + +#[allow(clippy::too_many_lines)] +pub(crate) async fn execute_query_task( + application: &Application, + operation_id: &str, + mut cancellation: watch::Receiver, + query: PreparedQuery, + storage: Storage, + resolved: ResolvedDatasourceConnection, +) -> Result { + if let CancellationRequest::Requested { reason } = cancellation.borrow().clone() { + return Err(QueryTaskError::Cancelled(reason)); + } + validate_query(&query)?; + let parameters = oracle_query_parameters(&query.parameters)?; + let open = open_resolved_connection(&resolved); + let managed = match await_with_cancellation(&mut cancellation, open).await { + AwaitOutcome::Completed(result) => result?, + AwaitOutcome::Cancelled(reason) => return Err(QueryTaskError::Cancelled(reason)), + }; + match start_read_only_transaction(&managed, &mut cancellation).await { + AwaitOutcome::Completed(Ok(())) => {} + AwaitOutcome::Completed(Err(error)) => { + managed.abandon().await; + return Err(error.into()); + } + AwaitOutcome::Cancelled(reason) => { + managed.abandon().await; + return Err(QueryTaskError::Cancelled(reason)); + } + } + let mut page = match await_with_cancellation( + &mut cancellation, + managed.connection.query(&query.sql, ¶meters), + ) + .await + { + AwaitOutcome::Completed(Ok(result)) => result, + AwaitOutcome::Completed(Err(error)) => { + managed.abandon().await; + return Err(oracle_query_error(&error).into()); + } + AwaitOutcome::Cancelled(reason) => { + managed.abandon().await; + return Err(QueryTaskError::Cancelled(reason)); + } + }; + let columns = page.columns.clone(); + if let Err(error) = validate_result_columns(&columns) { + managed.abandon().await; + return Err(error.into()); + } + if columns.len() > MAX_COLUMNS { + managed.abandon().await; + return Err(resource_error( + "oracle_result_too_wide", + format!("Oracle returned more than {MAX_COLUMNS} columns"), + ) + .into()); + } + let schema = wire::QueryStarted { + columns: columns + .iter() + .enumerate() + .map(|(index, column)| oracle_column(index, column)) + .collect::>()?, + }; + let mut writer = match RetainedWriter::begin(storage, schema, query.retention).await { + Ok(writer) => writer, + Err(error) => { + managed.abandon().await; + return Err(error.into()); + } + }; + if let Err(error) = application.inner.operations.started(operation_id).await { + abort_writer(&mut writer).await; + managed.abandon().await; + return Err(error.into()); + } + + let max_rows = query.options.max_rows; + let max_result_bytes = if query.options.max_result_bytes == 0 { + DEFAULT_RESULT_BYTES + } else { + query.options.max_result_bytes + }; + let batch_rows = if query.options.target_batch_rows == 0 { + DEFAULT_BATCH_ROWS + } else { + query.options.target_batch_rows + }; + let batch_bytes = if query.options.target_batch_bytes == 0 { + DEFAULT_BATCH_BYTES + } else { + query.options.target_batch_bytes + }; + let cursor_id = page.cursor_id; + let mut pending_rows = Vec::new(); + let mut pending_bytes = 0_u64; + let mut row_count = 0_u64; + let mut result_bytes = 0_u64; + let mut truncated_by_max_rows = false; + let mut truncated_by_max_result_bytes = false; + let mut discard_connection = false; + + 'pages: loop { + let has_more = page.has_more_rows; + for row in std::mem::take(&mut page.rows) { + if max_rows != 0 && row_count >= max_rows { + truncated_by_max_rows = true; + discard_connection = true; + break 'pages; + } + let converted = match await_with_cancellation( + &mut cancellation, + oracle_row(&managed.connection, row, &columns), + ) + .await + { + AwaitOutcome::Completed(Ok(row)) => row, + AwaitOutcome::Completed(Err(error)) => { + abort_writer(&mut writer).await; + managed.abandon().await; + return Err(error.into()); + } + AwaitOutcome::Cancelled(reason) => { + abort_writer(&mut writer).await; + managed.abandon().await; + return Err(QueryTaskError::Cancelled(reason)); + } + }; + let row_bytes = u64::try_from(converted.encoded_len()) + .map_err(|_| QueryTaskError::Failed(AppError::internal()))?; + if result_bytes.saturating_add(row_bytes) > max_result_bytes { + truncated_by_max_result_bytes = true; + discard_connection = true; + break 'pages; + } + let entry_bytes = row_batch_entry_bytes(&converted)?; + let candidate_bytes = pending_bytes + .saturating_add(if pending_rows.is_empty() { + row_batch_prefix_bytes(row_count) + } else { + 0 + }) + .saturating_add(entry_bytes); + if !pending_rows.is_empty() + && (pending_rows.len() + >= usize::try_from(batch_rows) + .map_err(|_| QueryTaskError::Failed(AppError::internal()))? + || candidate_bytes > u64::from(batch_bytes)) + { + flush_rows( + application, + operation_id, + &mut writer, + &mut pending_rows, + row_count, + ) + .await?; + pending_bytes = 0; + } + if pending_rows.is_empty() { + pending_bytes = row_batch_prefix_bytes(row_count); + } + pending_rows.push(converted); + pending_bytes = pending_bytes.saturating_add(entry_bytes); + row_count = row_count + .checked_add(1) + .ok_or_else(|| QueryTaskError::Failed(AppError::internal()))?; + result_bytes = result_bytes + .checked_add(row_bytes) + .ok_or_else(|| QueryTaskError::Failed(AppError::internal()))?; + } + if !has_more { + break; + } + page = match await_with_cancellation( + &mut cancellation, + managed + .connection + .fetch_more(cursor_id, &columns, FETCH_ROWS), + ) + .await + { + AwaitOutcome::Completed(Ok(result)) => result, + AwaitOutcome::Completed(Err(error)) => { + abort_writer(&mut writer).await; + managed.abandon().await; + return Err(oracle_query_error(&error).into()); + } + AwaitOutcome::Cancelled(reason) => { + abort_writer(&mut writer).await; + managed.abandon().await; + return Err(QueryTaskError::Cancelled(reason)); + } + }; + } + + if let Err(error) = flush_rows( + application, + operation_id, + &mut writer, + &mut pending_rows, + row_count, + ) + .await + { + abort_writer(&mut writer).await; + managed.abandon().await; + return Err(error); + } + let metadata = match writer + .finish(wire::QueryCompleted { + row_count, + truncated_by_max_rows, + truncated_by_max_result_bytes, + }) + .await + { + Ok(metadata) => metadata, + Err(error) => { + abort_writer(&mut writer).await; + managed.abandon().await; + return Err(error.into()); + } + }; + if discard_connection { + managed.abandon().await; + } else { + finish_read_only_connection(managed).await; + } + Ok(metadata) +} + +fn row_batch_prefix_bytes(start_row_offset: u64) -> u64 { + if start_row_offset == 0 { + 0 + } else { + 1_u64.saturating_add( + u64::try_from(prost::encoding::encoded_len_varint(start_row_offset)) + .unwrap_or(u64::MAX), + ) + } +} + +fn row_batch_entry_bytes(row: &wire::JdbcRow) -> Result { + let row_bytes = row.encoded_len(); + let length_bytes = prost::encoding::length_delimiter_len(row_bytes); + u64::try_from( + 1_usize + .saturating_add(length_bytes) + .saturating_add(row_bytes), + ) + .map_err(|_| QueryTaskError::Failed(AppError::internal())) +} + +async fn flush_rows( + application: &Application, + operation_id: &str, + writer: &mut RetainedWriter, + rows: &mut Vec, + row_count: u64, +) -> Result<(), QueryTaskError> { + if rows.is_empty() { + return Ok(()); + } + let row_len = u64::try_from(rows.len()).map_err(|_| AppError::internal())?; + let start_row_offset = row_count + .checked_sub(row_len) + .ok_or_else(AppError::internal)?; + let batch = wire::RowBatch { + start_row_offset, + rows: std::mem::take(rows), + }; + if batch.encoded_len() > usize::try_from(MAX_BATCH_BYTES).unwrap_or(usize::MAX) { + return Err(resource_error( + "oracle_result_batch_too_large", + "One Oracle result row exceeds the retained-result batch limit", + ) + .into()); + } + let byte_count = writer.append(batch).await?; + application + .inner + .operations + .progress(operation_id, row_count, byte_count) + .await?; + Ok(()) +} + +async fn abort_writer(writer: &mut RetainedWriter) { + if let Err(error) = writer.abort().await { + tracing::warn!(error = %error, "Oracle retained-result cleanup failed"); + } +} + +async fn finish_read_only_connection(managed: ManagedOracleConnection) { + match tokio::time::timeout(OPERATION_TIMEOUT, managed.connection.rollback()).await { + Ok(Ok(())) => {} + Ok(Err(error)) => { + tracing::warn!(error = %error, "Oracle read-only transaction rollback failed"); + managed.abandon().await; + return; + } + Err(_) => { + tracing::warn!("Oracle read-only transaction rollback timed out"); + managed.abandon().await; + return; + } + } + if let Err(error) = managed.close().await { + tracing::warn!(error = %error, "Oracle connection cleanup failed"); + } +} + +fn resource_error(code: impl Into, message: impl Into) -> AppError { + AppError::new( + AppErrorKind::ResourceExhausted, + ApiError::new(code, message), + ) +} + +fn oracle_column(index: usize, column: &ColumnInfo) -> Result { + let ordinal = u32::try_from(index) + .ok() + .and_then(|index| index.checked_add(1)) + .ok_or_else(AppError::internal)?; + let value_type = oracle_value_type(column); + let precision = u32::try_from(column.precision) + .ok() + .filter(|value| *value > 0); + let scale = (column.scale != 0).then(|| i32::from(column.scale)); + let display_size = (column.data_size > 0).then_some(column.data_size); + Ok(wire::JdbcColumn { + ordinal, + name: column.name.clone(), + label: column.name.clone(), + jdbc_type: oracle_jdbc_type(column.oracle_type), + jdbc_type_name: oracle_type_name(column.oracle_type).to_owned(), + value_type: value_type as i32, + nullability: if column.nullable { + wire::ColumnNullability::Nullable as i32 + } else { + wire::ColumnNullability::NoNulls as i32 + }, + precision, + scale, + display_size, + signed: oracle_numeric_type(column.oracle_type).then_some(true), + catalog_name: None, + schema_name: column.type_schema.clone(), + table_name: None, + }) +} + +fn validate_result_columns(columns: &[ColumnInfo]) -> Result<(), AppError> { + if let Some(column) = columns + .iter() + .find(|column| !oracle_result_type_supported(column.oracle_type)) + { + return Err(result_type_not_supported(column.oracle_type)); + } + Ok(()) +} + +const fn oracle_result_type_supported(oracle_type: OracleType) -> bool { + !matches!( + oracle_type, + OracleType::BinaryFloat + | OracleType::BinaryDouble + | OracleType::Rowid + | OracleType::Urowid + | OracleType::Bfile + | OracleType::Cursor + | OracleType::Object + | OracleType::Vector + | OracleType::IntervalYm + | OracleType::IntervalDs + ) +} + +fn oracle_value_type(column: &ColumnInfo) -> wire::JdbcValueType { + match column.oracle_type { + OracleType::Number | OracleType::BinaryInteger + if column.scale == 0 && (1..=18).contains(&column.precision) => + { + wire::JdbcValueType::SignedInteger + } + OracleType::Number | OracleType::BinaryInteger => wire::JdbcValueType::Decimal, + OracleType::BinaryFloat => wire::JdbcValueType::Float32, + OracleType::BinaryDouble => wire::JdbcValueType::Float64, + OracleType::Raw | OracleType::LongRaw | OracleType::Blob => wire::JdbcValueType::Binary, + OracleType::Date | OracleType::Timestamp | OracleType::TimestampLtz => { + wire::JdbcValueType::Timestamp + } + OracleType::TimestampTz => wire::JdbcValueType::TimestampWithTimeZone, + OracleType::Json => wire::JdbcValueType::Json, + OracleType::Boolean => wire::JdbcValueType::Boolean, + OracleType::Varchar + | OracleType::Char + | OracleType::Long + | OracleType::Clob + | OracleType::Rowid + | OracleType::Urowid => wire::JdbcValueType::Text, + OracleType::Bfile + | OracleType::Cursor + | OracleType::Object + | OracleType::Vector + | OracleType::IntervalYm + | OracleType::IntervalDs => wire::JdbcValueType::Opaque, + } +} + +const fn oracle_numeric_type(oracle_type: OracleType) -> bool { + matches!( + oracle_type, + OracleType::Number + | OracleType::BinaryInteger + | OracleType::BinaryFloat + | OracleType::BinaryDouble + ) +} + +const fn oracle_jdbc_type(oracle_type: OracleType) -> i32 { + match oracle_type { + OracleType::Varchar => 12, + OracleType::Number => 2, + OracleType::BinaryInteger => 4, + OracleType::Long => -1, + OracleType::Rowid | OracleType::Urowid => -8, + OracleType::Date | OracleType::Timestamp | OracleType::TimestampLtz => 93, + OracleType::Raw => -3, + OracleType::LongRaw => -4, + OracleType::Char => 1, + OracleType::BinaryFloat => 6, + OracleType::BinaryDouble => 8, + OracleType::Cursor => -10, + OracleType::Object => 2_002, + OracleType::Clob => 2_005, + OracleType::Blob => 2_004, + OracleType::Bfile => -13, + OracleType::Json | OracleType::Vector | OracleType::IntervalYm | OracleType::IntervalDs => { + 1_111 + } + OracleType::TimestampTz => 2_014, + OracleType::Boolean => 16, + } +} + +const fn oracle_type_name(oracle_type: OracleType) -> &'static str { + match oracle_type { + OracleType::Varchar => "VARCHAR2", + OracleType::Number => "NUMBER", + OracleType::BinaryInteger => "BINARY_INTEGER", + OracleType::Long => "LONG", + OracleType::Rowid => "ROWID", + OracleType::Date => "DATE", + OracleType::Raw => "RAW", + OracleType::LongRaw => "LONG RAW", + OracleType::Char => "CHAR", + OracleType::BinaryFloat => "BINARY_FLOAT", + OracleType::BinaryDouble => "BINARY_DOUBLE", + OracleType::Cursor => "REF CURSOR", + OracleType::Object => "OBJECT", + OracleType::Clob => "CLOB", + OracleType::Blob => "BLOB", + OracleType::Bfile => "BFILE", + OracleType::Json => "JSON", + OracleType::Vector => "VECTOR", + OracleType::Timestamp => "TIMESTAMP", + OracleType::TimestampTz => "TIMESTAMP WITH TIME ZONE", + OracleType::IntervalYm => "INTERVAL YEAR TO MONTH", + OracleType::IntervalDs => "INTERVAL DAY TO SECOND", + OracleType::Urowid => "UROWID", + OracleType::TimestampLtz => "TIMESTAMP WITH LOCAL TIME ZONE", + OracleType::Boolean => "BOOLEAN", + } +} + +async fn oracle_row( + connection: &Connection, + row: OracleRow, + columns: &[ColumnInfo], +) -> Result { + if row.len() != columns.len() { + return Err(AppError::internal()); + } + let mut values = Vec::with_capacity(columns.len()); + for (value, column) in row.into_values().into_iter().zip(columns) { + values.push(oracle_wire_value(connection, value, column).await?); + } + Ok(wire::JdbcRow { values }) +} + +async fn oracle_wire_value( + connection: &Connection, + value: OracleValue, + column: &ColumnInfo, +) -> Result { + use wire::jdbc_value::Value as WireValue; + if matches!(value, OracleValue::Null | OracleValue::Lob(LobValue::Null)) { + return Ok(wire_value(WireValue::NullValue(wire::JdbcNull {}))); + } + let value_type = oracle_value_type(column); + let value = match value_type { + wire::JdbcValueType::Boolean => { + WireValue::BooleanValue(oracle_bool(&value).ok_or_else(result_decode_error)?) + } + wire::JdbcValueType::SignedInteger => { + WireValue::SignedIntegerValue(oracle_i64(&value).ok_or_else(result_decode_error)?) + } + wire::JdbcValueType::UnsignedInteger => WireValue::UnsignedIntegerValue( + u64::try_from(oracle_i64(&value).ok_or_else(result_decode_error)?) + .map_err(|_| result_decode_error())?, + ), + wire::JdbcValueType::Float32 => { + WireValue::Float32Value(oracle_f32(&value).ok_or_else(result_decode_error)?) + } + wire::JdbcValueType::Float64 => { + WireValue::Float64Value(oracle_f64(&value).ok_or_else(result_decode_error)?) + } + wire::JdbcValueType::Decimal => WireValue::DecimalValue(oracle_decimal_text(&value)?), + wire::JdbcValueType::Text => { + WireValue::TextValue(oracle_text(connection, value, column.oracle_type).await?) + } + wire::JdbcValueType::Binary => { + WireValue::BinaryValue(oracle_binary(connection, value).await?) + } + wire::JdbcValueType::Date => { + WireValue::DateValue(oracle_text(connection, value, column.oracle_type).await?) + } + wire::JdbcValueType::Time => { + WireValue::TimeValue(oracle_text(connection, value, column.oracle_type).await?) + } + wire::JdbcValueType::Timestamp => { + WireValue::TimestampValue(oracle_temporal_text(&value, false)?) + } + wire::JdbcValueType::TimestampWithTimeZone => { + WireValue::TimestampWithTimeZoneValue(oracle_temporal_text(&value, true)?) + } + wire::JdbcValueType::Json => { + let text = match value { + OracleValue::Json(value) => value.to_string(), + other => oracle_text(connection, other, column.oracle_type).await?, + }; + validate_result_scalar(text.len())?; + WireValue::JsonValue(text) + } + wire::JdbcValueType::Uuid => { + WireValue::UuidValue(oracle_text(connection, value, column.oracle_type).await?) + } + wire::JdbcValueType::Opaque | wire::JdbcValueType::Unspecified => { + return Err(result_type_not_supported(column.oracle_type)); + } + }; + Ok(wire_value(value)) +} + +fn wire_value(value: wire::jdbc_value::Value) -> wire::JdbcValue { + wire::JdbcValue { value: Some(value) } +} + +fn oracle_bool(value: &OracleValue) -> Option { + match value { + OracleValue::Boolean(value) => Some(*value), + OracleValue::Integer(value) => Some(*value != 0), + OracleValue::Number(value) => value.to_i64().ok().map(|value| value != 0), + OracleValue::String(value) if value == "1" || value.eq_ignore_ascii_case("true") => { + Some(true) + } + OracleValue::String(value) if value == "0" || value.eq_ignore_ascii_case("false") => { + Some(false) + } + OracleValue::String(value) if value.as_bytes() == [1, 1] => Some(true), + OracleValue::String(value) if value.as_bytes() == [1, 0] => Some(false), + _ => None, + } +} + +fn oracle_i64(value: &OracleValue) -> Option { + match value { + OracleValue::Integer(value) => Some(*value), + OracleValue::Float(value) => value.to_string().parse().ok(), + OracleValue::Number(value) => value.to_i64().ok(), + OracleValue::String(value) => value.parse().ok(), + _ => None, + } +} + +fn oracle_f32(value: &OracleValue) -> Option { + match value { + OracleValue::Float(value) => value.to_string().parse().ok(), + OracleValue::Integer(value) => value.to_string().parse().ok(), + OracleValue::Number(value) => value.as_str().parse().ok(), + OracleValue::String(value) => value.parse().ok(), + _ => None, + } +} + +fn oracle_f64(value: &OracleValue) -> Option { + match value { + OracleValue::Float(value) => Some(*value), + OracleValue::Integer(value) => value.to_string().parse().ok(), + OracleValue::Number(value) => value.to_f64().ok(), + OracleValue::String(value) => value.parse().ok(), + _ => None, + } +} + +fn oracle_decimal_text(value: &OracleValue) -> Result { + let value = match value { + OracleValue::Number(value) => value.as_str().to_owned(), + OracleValue::Integer(value) => value.to_string(), + OracleValue::Float(value) => value.to_string(), + OracleValue::String(value) => value.clone(), + _ => return Err(result_decode_error()), + }; + validate_result_scalar(value.len())?; + Ok(value) +} + +async fn oracle_text( + connection: &Connection, + value: OracleValue, + oracle_type: OracleType, +) -> Result { + let value = match value { + OracleValue::String(value) => value, + OracleValue::Bytes(value) => String::from_utf8(value).map_err(|_| result_decode_error())?, + OracleValue::Integer(value) => value.to_string(), + OracleValue::Float(value) => value.to_string(), + OracleValue::Number(value) => value.as_str().to_owned(), + OracleValue::Date(value) => format_oracle_date(value), + OracleValue::Timestamp(value) if value.has_timezone() => format_oracle_timestamp_tz(value)?, + OracleValue::Timestamp(value) => format_oracle_timestamp(value), + OracleValue::RowId(value) => format!("{value}"), + OracleValue::Boolean(value) => value.to_string(), + OracleValue::Json(value) => value.to_string(), + OracleValue::Lob(value) => oracle_lob_text(connection, value, oracle_type).await?, + OracleValue::Null => return Err(result_decode_error()), + _ => return Err(result_type_not_supported(oracle_type)), + }; + validate_result_scalar(value.len())?; + Ok(value) +} + +async fn oracle_binary(connection: &Connection, value: OracleValue) -> Result, AppError> { + let value = match value { + OracleValue::Bytes(value) => value, + OracleValue::String(value) => value.into_bytes(), + OracleValue::Lob(LobValue::Inline(value)) => { + validate_result_scalar(value.len())?; + value.to_vec() + } + OracleValue::Lob(LobValue::Empty) => Vec::new(), + OracleValue::Lob(LobValue::Locator(locator)) => { + if locator.size() > u64::try_from(MAX_SCALAR_BYTES).unwrap_or(u64::MAX) { + return Err(result_scalar_too_large()); + } + match connection + .read_lob(&locator) + .await + .map_err(|error| oracle_query_error(&error))? + { + LobData::Bytes(value) => { + validate_result_scalar(value.len())?; + value.to_vec() + } + LobData::String(value) => { + validate_result_scalar(value.len())?; + value.into_bytes() + } + } + } + OracleValue::Lob(LobValue::Null) | OracleValue::Null => { + return Err(result_decode_error()); + } + _ => return Err(result_decode_error()), + }; + validate_result_scalar(value.len())?; + Ok(value) +} + +async fn oracle_lob_text( + connection: &Connection, + value: LobValue, + oracle_type: OracleType, +) -> Result { + match value { + LobValue::Inline(value) => { + validate_result_scalar(value.len())?; + String::from_utf8(value.to_vec()).map_err(|_| result_decode_error()) + } + LobValue::Empty => Ok(String::new()), + LobValue::Null => Err(result_decode_error()), + LobValue::Locator(locator) => { + if locator.size() > u64::try_from(MAX_SCALAR_BYTES).unwrap_or(u64::MAX) { + return Err(result_scalar_too_large()); + } + match connection + .read_lob(&locator) + .await + .map_err(|error| oracle_query_error(&error))? + { + LobData::String(value) => { + validate_result_scalar(value.len())?; + Ok(value) + } + LobData::Bytes(value) if matches!(oracle_type, OracleType::Clob) => { + validate_result_scalar(value.len())?; + String::from_utf8(value.to_vec()).map_err(|_| result_decode_error()) + } + LobData::Bytes(value) => { + let encoded_len = value.len().saturating_mul(4).div_ceil(3); + validate_result_scalar(encoded_len)?; + Ok(BASE64_STANDARD.encode(value)) + } + } + } + } +} + +fn oracle_temporal_text(value: &OracleValue, require_timezone: bool) -> Result { + let value = match value { + OracleValue::Date(value) if !require_timezone => format_oracle_date(*value), + OracleValue::Timestamp(value) if require_timezone => format_oracle_timestamp_tz(*value)?, + OracleValue::Timestamp(value) => format_oracle_timestamp(*value), + OracleValue::String(value) => value.clone(), + _ => return Err(result_decode_error()), + }; + validate_result_scalar(value.len())?; + Ok(value) +} + +fn format_oracle_date(value: OracleDate) -> String { + format!( + "{:04}-{:02}-{:02}T{:02}:{:02}:{:02}", + value.year, value.month, value.day, value.hour, value.minute, value.second + ) +} + +fn format_oracle_timestamp(value: OracleTimestamp) -> String { + let mut result = format!( + "{:04}-{:02}-{:02}T{:02}:{:02}:{:02}", + value.year, value.month, value.day, value.hour, value.minute, value.second + ); + if value.microsecond != 0 { + let _ = write!(&mut result, ".{:06}", value.microsecond); + } + result +} + +fn format_oracle_timestamp_tz(value: OracleTimestamp) -> Result { + let offset_minutes = i32::from(value.tz_minute_offset); + let offset_hours = i32::from(value.tz_hour_offset); + if offset_minutes.unsigned_abs() > 59 + || (offset_hours > 0 && offset_minutes < 0) + || (offset_hours < 0 && offset_minutes > 0) + { + return Err(result_decode_error()); + } + let offset_seconds = offset_hours + .checked_mul(3_600) + .and_then(|seconds| { + offset_minutes + .checked_mul(60) + .and_then(|minutes| seconds.checked_add(minutes)) + }) + .filter(|seconds| (-12 * 3_600..=14 * 3_600).contains(seconds)) + .ok_or_else(result_decode_error)?; + let offset = FixedOffset::east_opt(offset_seconds).ok_or_else(result_decode_error)?; + let utc = NaiveDate::from_ymd_opt(value.year, u32::from(value.month), u32::from(value.day)) + .and_then(|date| { + date.and_hms_micro_opt( + u32::from(value.hour), + u32::from(value.minute), + u32::from(value.second), + value.microsecond, + ) + }) + .ok_or_else(result_decode_error)?; + // Oracle's TSTZ wire value carries UTC fields plus the original numeric offset. + let local = utc + .checked_add_signed(chrono::TimeDelta::seconds(i64::from(offset_seconds))) + .ok_or_else(result_decode_error)?; + Ok(format!( + "{}{}", + local.format("%Y-%m-%dT%H:%M:%S%.f"), + offset + )) +} + +fn oracle_display(value: &OracleValue) -> String { + let display = value.to_string(); + if display.len() <= MAX_SCALAR_BYTES { + display + } else { + "[value exceeds display limit]".to_owned() + } +} + +fn validate_result_scalar(size: usize) -> Result<(), AppError> { + if size > MAX_SCALAR_BYTES { + Err(result_scalar_too_large()) + } else { + Ok(()) + } +} + +fn result_scalar_too_large() -> AppError { + resource_error( + "oracle_scalar_too_large", + format!("An Oracle value exceeds {MAX_SCALAR_BYTES} bytes"), + ) +} + +fn result_decode_error() -> AppError { + AppError::unavailable( + "oracle_result_decode_failed", + "An Oracle result value could not be decoded safely", + ) +} + +fn result_type_not_supported(oracle_type: OracleType) -> AppError { + AppError::invalid( + "oracle_result_type_not_supported", + format!( + "The native Oracle driver does not support result type {}", + oracle_type_name(oracle_type) + ), + ) +} + +fn console_column(index: usize, column: &ColumnInfo) -> Result { + let column = oracle_column(index, column)?; + let value_type = match wire::JdbcValueType::try_from(column.value_type) { + Ok(wire::JdbcValueType::Boolean) => JdbcValueType::Boolean, + Ok(wire::JdbcValueType::SignedInteger) => JdbcValueType::SignedInteger, + Ok(wire::JdbcValueType::UnsignedInteger) => JdbcValueType::UnsignedInteger, + Ok(wire::JdbcValueType::Float32) => JdbcValueType::Float32, + Ok(wire::JdbcValueType::Float64) => JdbcValueType::Float64, + Ok(wire::JdbcValueType::Decimal) => JdbcValueType::Decimal, + Ok(wire::JdbcValueType::Text) => JdbcValueType::Text, + Ok(wire::JdbcValueType::Binary) => JdbcValueType::Binary, + Ok(wire::JdbcValueType::Date) => JdbcValueType::Date, + Ok(wire::JdbcValueType::Time) => JdbcValueType::Time, + Ok(wire::JdbcValueType::Timestamp) => JdbcValueType::Timestamp, + Ok(wire::JdbcValueType::TimestampWithTimeZone) => JdbcValueType::TimestampWithTimeZone, + Ok(wire::JdbcValueType::Json) => JdbcValueType::Json, + Ok(wire::JdbcValueType::Uuid) => JdbcValueType::Uuid, + Ok(wire::JdbcValueType::Opaque) => JdbcValueType::Opaque, + Ok(wire::JdbcValueType::Unspecified) | Err(_) => return Err(AppError::internal()), + }; + let nullability = match wire::ColumnNullability::try_from(column.nullability) { + Ok(wire::ColumnNullability::Unknown) => ColumnNullability::Unknown, + Ok(wire::ColumnNullability::NoNulls) => ColumnNullability::NoNulls, + Ok(wire::ColumnNullability::Nullable) => ColumnNullability::Nullable, + Err(_) => return Err(AppError::internal()), + }; + Ok(ResultColumn { + ordinal: column.ordinal, + label: column.label, + name: column.name, + jdbc_type: column.jdbc_type, + jdbc_type_name: column.jdbc_type_name, + value_type, + nullability, + precision: column.precision, + scale: column.scale, + display_size: column.display_size, + signed: column.signed, + catalog_name: column.catalog_name, + schema_name: column.schema_name, + table_name: column.table_name, + }) +} + +async fn console_row( + connection: &Connection, + row: OracleRow, + columns: &[ColumnInfo], +) -> Result { + let row = oracle_row(connection, row, columns).await?; + Ok(ResultRow { + values: row + .values + .into_iter() + .map(contract_value) + .collect::>()?, + }) +} + +fn contract_value(value: wire::JdbcValue) -> Result { + use wire::jdbc_value::Value; + match value.value.ok_or_else(AppError::internal)? { + Value::NullValue(_) => Ok(JdbcValue::Null), + Value::BooleanValue(value) => Ok(JdbcValue::Boolean { value }), + Value::SignedIntegerValue(value) => Ok(JdbcValue::SignedInteger { + value: value.to_string(), + }), + Value::UnsignedIntegerValue(value) => Ok(JdbcValue::UnsignedInteger { + value: value.to_string(), + }), + Value::Float32Value(value) => Ok(JdbcValue::Float32 { + value: value.to_string(), + }), + Value::Float64Value(value) => Ok(JdbcValue::Float64 { + value: value.to_string(), + }), + Value::DecimalValue(value) => Ok(JdbcValue::Decimal { value }), + Value::TextValue(value) => Ok(JdbcValue::Text { value }), + Value::BinaryValue(value) => Ok(JdbcValue::Binary { + value: BASE64_STANDARD.encode(value), + }), + Value::DateValue(value) => Ok(JdbcValue::Date { value }), + Value::TimeValue(value) => Ok(JdbcValue::Time { value }), + Value::TimestampValue(value) => Ok(JdbcValue::Timestamp { value }), + Value::TimestampWithTimeZoneValue(value) => Ok(JdbcValue::TimestampWithTimeZone { value }), + Value::JsonValue(value) => Ok(JdbcValue::Json { value }), + Value::UuidValue(value) => Ok(JdbcValue::Uuid { value }), + Value::OpaqueValue(value) => Ok(JdbcValue::Opaque { + type_name: value.type_name, + display_value: value.display_value, + }), + } +} + +fn console_row_retained_bytes(row: &ResultRow) -> u64 { + let mut bytes = u64::try_from(size_of::()).unwrap_or(u64::MAX); + bytes = bytes.saturating_add( + u64::try_from(row.values.capacity()) + .unwrap_or(u64::MAX) + .saturating_mul(u64::try_from(size_of::()).unwrap_or(u64::MAX)), + ); + for value in &row.values { + let value_bytes = match value { + JdbcValue::Null | JdbcValue::Boolean { .. } => 0, + JdbcValue::SignedInteger { value } + | JdbcValue::UnsignedInteger { value } + | JdbcValue::Float32 { value } + | JdbcValue::Float64 { value } + | JdbcValue::Decimal { value } + | JdbcValue::Text { value } + | JdbcValue::Binary { value } + | JdbcValue::Date { value } + | JdbcValue::Time { value } + | JdbcValue::Timestamp { value } + | JdbcValue::TimestampWithTimeZone { value } + | JdbcValue::Json { value } + | JdbcValue::Uuid { value } => value.capacity(), + JdbcValue::Opaque { + type_name, + display_value, + } => type_name + .capacity() + .saturating_add(display_value.capacity()), + }; + bytes = bytes.saturating_add(u64::try_from(value_bytes).unwrap_or(u64::MAX)); + } + bytes +} + +fn reserve_console_result_bytes(total: &mut u64, row: &ResultRow) -> Result<(), AppError> { + let next = total.saturating_add(console_row_retained_bytes(row)); + if next > MAX_CONSOLE_RESULT_BYTES { + return Err(resource_error( + "oracle_console_result_too_large", + format!( + "Oracle Console results are limited to {MAX_CONSOLE_RESULT_BYTES} retained bytes" + ), + )); + } + *total = next; + Ok(()) +} + +pub(crate) async fn execute_update( + resolved: ResolvedDatasourceConnection, + sql: String, + cancellation: CancellationToken, +) -> Result { + if cancellation.is_cancelled() { + return Err(DatabaseWriteError::not_started(oracle_write_cancelled( + false, + ))); + } + let sql = validate_single_write_sql(&sql).map_err(DatabaseWriteError::not_started)?; + if resolved.connection.read_only { + return Err(DatabaseWriteError::not_started(AppError::new( + AppErrorKind::Conflict, + ApiError::new( + "datasource_read_only", + "The datasource connection is configured as read-only", + ), + ))); + } + let open = open_resolved_connection(&resolved); + tokio::pin!(open); + let managed = tokio::select! { + biased; + () = cancellation.cancelled() => { + return Err(DatabaseWriteError::not_started(oracle_write_cancelled(false))); + } + result = &mut open => result.map_err(DatabaseWriteError::not_started)?, + }; + if cancellation.is_cancelled() { + managed.abandon().await; + return Err(DatabaseWriteError::not_started(oracle_write_cancelled( + false, + ))); + } + + let Some(result) = await_with_token(&cancellation, managed.connection.execute(&sql, &[])).await + else { + managed.abandon().await; + return Err(DatabaseWriteError::unknown(oracle_write_cancelled(true))); + }; + let result = match result { + Ok(result) => result, + Err(error) => { + managed.abandon().await; + return Err(DatabaseWriteError::unknown(AppError::unavailable( + "database_write_outcome_unknown", + format!( + "Oracle reported an error after write dispatch; partial effects cannot be excluded, so do not retry blindly: {}", + safe_oracle_error(&error) + ), + ))); + } + }; + let affected_rows = result.rows_affected; + let Some(commit_result) = await_with_token(&cancellation, managed.connection.commit()).await + else { + managed.abandon().await; + return Err(DatabaseWriteError::unknown(oracle_write_cancelled(true))); + }; + if let Err(error) = commit_result { + managed.abandon().await; + return Err(DatabaseWriteError::unknown(AppError::unavailable( + "database_write_outcome_unknown", + format!( + "The Oracle commit outcome is unknown; do not retry blindly: {}", + safe_oracle_error(&error) + ), + ))); + } + if let Err(error) = managed.close().await { + tracing::warn!(error = %error, "Oracle write connection cleanup failed after commit"); + } + Ok(affected_rows) +} + +async fn await_with_token(cancellation: &CancellationToken, future: F) -> Option +where + F: Future, +{ + tokio::pin!(future); + tokio::select! { + biased; + () = cancellation.cancelled() => None, + output = &mut future => Some(output), + } +} + +fn validate_single_write_sql(sql: &str) -> Result { + validate_sql_text(sql)?; + let mut statements = split_oracle_script(sql)?; + if statements.len() != 1 { + return Err(AppError::invalid( + "invalid_database_write", + "Exactly one Oracle write statement is required", + )); + } + let statement = statements.pop().expect("length checked above"); + let first = first_sql_keyword(&statement)?; + if !matches!( + first.as_deref(), + Some( + "INSERT" + | "UPDATE" + | "DELETE" + | "MERGE" + | "CREATE" + | "ALTER" + | "DROP" + | "TRUNCATE" + | "RENAME" + | "GRANT" + | "REVOKE" + | "ANALYZE" + | "CALL" + | "BEGIN" + | "DECLARE" + ) + ) { + return Err(AppError::invalid( + "database_write_statement_required", + "The confirmed Oracle write surface accepts one DML, DDL, grant, call, or PL/SQL statement", + )); + } + Ok(statement) +} + +fn oracle_write_cancelled(dispatched: bool) -> AppError { + if dispatched { + AppError::unavailable( + "database_write_outcome_unknown", + "The Oracle write was interrupted after dispatch; do not retry it blindly", + ) + } else { + AppError::new( + AppErrorKind::Conflict, + ApiError::new( + "database_write_cancelled", + "The Oracle write was cancelled before dispatch", + ), + ) + } +} + +fn safe_oracle_error(error: &OracleError) -> String { + match error { + OracleError::OracleError { .. } + | OracleError::ServerError { .. } + | OracleError::SqlError(_) => error.to_string(), + _ => "the protocol session ended unexpectedly".to_owned(), + } +} + +#[allow(clippy::too_many_lines)] +pub(crate) async fn execute_console( + application: &Application, + request: NativeConsoleRequest, + mut cancellation: watch::Receiver, + force_read_only: bool, +) -> Result, AppError> { + let (statements, page_offset, page_end) = prepare_console_statements(&request)?; + if let CancellationRequest::Requested { reason } = cancellation.borrow().clone() { + return Err(oracle_console_cancelled(reason)); + } + if force_read_only { + validate_read_only_statements(&statements, "agent read query")?; + } + let resolved = resolve_native_connection(application, &request.datasource_id).await?; + if resolved.connection.read_only && !force_read_only { + validate_read_only_statements(&statements, "read-only datasource")?; + } + let effective_read_only = force_read_only || resolved.connection.read_only; + let managed = + match await_with_cancellation(&mut cancellation, open_resolved_connection(&resolved)).await + { + AwaitOutcome::Completed(result) => result?, + AwaitOutcome::Cancelled(reason) => return Err(oracle_console_cancelled(reason)), + }; + if effective_read_only { + match start_read_only_transaction(&managed, &mut cancellation).await { + AwaitOutcome::Completed(Ok(())) => {} + AwaitOutcome::Completed(Err(error)) => { + managed.abandon().await; + return Err(error); + } + AwaitOutcome::Cancelled(reason) => { + managed.abandon().await; + return Err(oracle_console_cancelled(reason)); + } + } + } + + let mut results = Vec::new(); + let mut retained_result_bytes = 0_u64; + for (index, sql) in statements.into_iter().enumerate() { + let statement_read_only = validate_read_sql(&sql).is_ok(); + let statement_sequence = u32::try_from(index) + .ok() + .and_then(|index| index.checked_add(1)) + .ok_or_else(AppError::internal)?; + let started = Instant::now(); + let execution = + await_with_cancellation(&mut cancellation, managed.connection.execute(&sql, &[])).await; + let mut query_result = match execution { + AwaitOutcome::Completed(Ok(result)) => result, + AwaitOutcome::Completed(Err(error)) if oracle_database_error(&error) => { + let error = oracle_query_error(&error); + results.push(console_failure_result( + statement_sequence, + sql, + &error, + elapsed_millis(started), + )); + if !request.error_continue { + break; + } + continue; + } + AwaitOutcome::Completed(Err(error)) => { + managed.abandon().await; + return Err(oracle_console_connection_error(statement_read_only, &error)); + } + AwaitOutcome::Cancelled(reason) => { + managed.abandon().await; + return Err(oracle_console_interrupted(statement_read_only, reason)); + } + }; + let columns = query_result.columns.clone(); + if let Err(error) = validate_result_columns(&columns) { + managed.abandon().await; + return Err(oracle_console_post_dispatch_error( + statement_read_only, + error, + )); + } + if columns.len() > MAX_COLUMNS { + managed.abandon().await; + return Err(resource_error( + "oracle_result_too_wide", + format!("Oracle returned more than {MAX_COLUMNS} columns"), + )); + } + let tabular = !columns.is_empty(); + if tabular { + let retain = request.result_set_id.is_none_or(|selected| selected == 1); + let converted_columns = if retain { + columns + .iter() + .enumerate() + .map(|(index, column)| console_column(index, column)) + .collect::, _>>()? + } else { + Vec::new() + }; + let cursor_id = query_result.cursor_id; + let mut row_count = 0_u64; + let mut rows = Vec::new(); + 'pages: loop { + let has_more = query_result.has_more_rows; + for row in std::mem::take(&mut query_result.rows) { + if row_count >= MAX_CONSOLE_ROWS { + managed.abandon().await; + return Err(resource_error( + "oracle_console_row_limit_exceeded", + format!("Oracle Console results cannot exceed {MAX_CONSOLE_ROWS} rows"), + )); + } + if retain && (page_offset..page_end).contains(&row_count) { + let row = match await_with_cancellation( + &mut cancellation, + console_row(&managed.connection, row, &columns), + ) + .await + { + AwaitOutcome::Completed(Ok(row)) => row, + AwaitOutcome::Completed(Err(error)) => { + managed.abandon().await; + return Err(oracle_console_post_dispatch_error( + statement_read_only, + error, + )); + } + AwaitOutcome::Cancelled(reason) => { + managed.abandon().await; + return Err(oracle_console_interrupted( + statement_read_only, + reason, + )); + } + }; + reserve_console_result_bytes(&mut retained_result_bytes, &row)?; + rows.push(row); + } + row_count = row_count.checked_add(1).ok_or_else(AppError::internal)?; + } + if !has_more { + break 'pages; + } + query_result = match await_with_cancellation( + &mut cancellation, + managed + .connection + .fetch_more(cursor_id, &columns, FETCH_ROWS), + ) + .await + { + AwaitOutcome::Completed(Ok(result)) => result, + AwaitOutcome::Completed(Err(error)) => { + managed.abandon().await; + return Err(oracle_console_connection_error(statement_read_only, &error)); + } + AwaitOutcome::Cancelled(reason) => { + managed.abandon().await; + return Err(oracle_console_interrupted(statement_read_only, reason)); + } + }; + } + if retain { + results.push(NativeConsoleResult { + statement_sequence, + result_set_id: Some(1), + sql, + success: true, + message: "Statement executed successfully".to_owned(), + update_count: 0, + columns: converted_columns, + rows, + row_count, + has_more: row_count > page_end, + duration_ms: elapsed_millis(started), + error: None, + }); + } + } else if request.result_set_id.is_none() { + let update_count = query_result.rows_affected; + if !effective_read_only { + match await_with_cancellation( + &mut cancellation, + tokio::time::timeout(OPERATION_TIMEOUT, managed.connection.commit()), + ) + .await + { + AwaitOutcome::Completed(Ok(Ok(()))) => {} + AwaitOutcome::Completed(Ok(Err(error))) => { + managed.abandon().await; + return Err(oracle_console_write_outcome_unknown(Some(&error))); + } + AwaitOutcome::Completed(Err(_)) | AwaitOutcome::Cancelled(_) => { + managed.abandon().await; + return Err(oracle_console_write_outcome_unknown(None)); + } + } + } + results.push(NativeConsoleResult { + statement_sequence, + result_set_id: None, + sql, + success: true, + message: "Statement executed successfully".to_owned(), + update_count, + columns: Vec::new(), + rows: Vec::new(), + row_count: 0, + has_more: false, + duration_ms: elapsed_millis(started), + error: None, + }); + } + } + + if effective_read_only { + finish_read_only_connection(managed).await; + } else if let Err(error) = managed.close().await { + tracing::warn!(error = %error, "Oracle Console connection cleanup failed"); + } + Ok(results) +} + +fn prepare_console_statements( + request: &NativeConsoleRequest, +) -> Result<(Vec, u64, u64), AppError> { + if request.page_no == 0 { + return Err(AppError::invalid( + "invalid_oracle_console_request", + "pageNo must be greater than zero", + )); + } + let page_size = if request.page_size_all { + MAX_CONSOLE_PAGE_SIZE + } else { + request.page_size + }; + if page_size == 0 || page_size > MAX_CONSOLE_PAGE_SIZE { + return Err(AppError::invalid( + "invalid_oracle_console_request", + format!("pageSize must be between 1 and {MAX_CONSOLE_PAGE_SIZE}"), + )); + } + validate_sql_text(&request.sql)?; + let mut statements = if request.single || looks_like_plsql(&request.sql) { + vec![normalize_preserved_statement(&request.sql)?] + } else { + split_oracle_script(&request.sql)? + }; + if statements.is_empty() || statements.len() > MAX_CONSOLE_STATEMENTS { + return Err(AppError::invalid( + "invalid_oracle_console_request", + format!("Oracle Console accepts between 1 and {MAX_CONSOLE_STATEMENTS} statements"), + )); + } + if request.explain { + for statement in &mut statements { + if !is_read_candidate(statement)? { + return Err(AppError::invalid( + "invalid_oracle_console_request", + "Oracle EXPLAIN accepts query statements only", + )); + } + *statement = format!("EXPLAIN PLAN FOR {statement}"); + } + } + let page_offset = u64::from(request.page_no - 1) + .checked_mul(u64::from(page_size)) + .ok_or_else(|| { + AppError::invalid( + "invalid_oracle_console_request", + "The requested result page is too large", + ) + })?; + let page_end = page_offset + .checked_add(u64::from(page_size)) + .ok_or_else(|| { + AppError::invalid( + "invalid_oracle_console_request", + "The requested result page is too large", + ) + })?; + Ok((statements, page_offset, page_end)) +} + +fn validate_read_only_statements(statements: &[String], source: &str) -> Result<(), AppError> { + for statement in statements { + validate_read_sql(statement).map_err(|_| { + AppError::new( + AppErrorKind::Conflict, + ApiError::new( + "datasource_read_only", + format!("The Oracle {source} accepts read-only SELECT statements"), + ), + ) + })?; + } + Ok(()) +} + +fn oracle_database_error(error: &OracleError) -> bool { + matches!( + error, + OracleError::OracleError { .. } + | OracleError::ServerError { .. } + | OracleError::SqlError(_) + | OracleError::NoDataFound + ) +} + +fn console_failure_result( + statement_sequence: u32, + sql: String, + error: &AppError, + duration_ms: u64, +) -> NativeConsoleResult { + NativeConsoleResult { + statement_sequence, + result_set_id: None, + sql, + success: false, + message: error.api_error().message, + update_count: 0, + columns: Vec::new(), + rows: Vec::new(), + row_count: 0, + has_more: false, + duration_ms, + error: Some(error.api_error()), + } +} + +fn oracle_console_cancelled(reason: Option) -> AppError { + AppError::new( + AppErrorKind::Conflict, + ApiError::new( + "oracle_console_cancelled", + reason.unwrap_or_else(|| "The Oracle Console execution was cancelled".to_owned()), + ), + ) +} + +fn oracle_console_write_outcome_unknown(error: Option<&OracleError>) -> AppError { + let message = error.map_or_else( + || "The Oracle write was interrupted after dispatch; do not retry it blindly".to_owned(), + |error| { + format!( + "The Oracle write outcome is unknown after dispatch; do not retry it blindly: {}", + safe_oracle_error(error) + ) + }, + ); + AppError::unavailable("database_write_outcome_unknown", message) +} + +fn oracle_console_interrupted(statement_read_only: bool, reason: Option) -> AppError { + if statement_read_only { + oracle_console_cancelled(reason) + } else { + oracle_console_write_outcome_unknown(None) + } +} + +fn oracle_console_connection_error(statement_read_only: bool, error: &OracleError) -> AppError { + if statement_read_only { + oracle_query_error(error) + } else { + oracle_console_write_outcome_unknown(Some(error)) + } +} + +fn oracle_console_post_dispatch_error(statement_read_only: bool, error: AppError) -> AppError { + if statement_read_only { + error + } else { + oracle_console_write_outcome_unknown(None) + } +} + +fn elapsed_millis(started: Instant) -> u64 { + u64::try_from(started.elapsed().as_millis()).unwrap_or(u64::MAX) +} + +fn normalize_preserved_statement(sql: &str) -> Result { + let mut sql = sql.trim().to_owned(); + if sql.lines().last().is_some_and(|line| line.trim() == "/") { + let index = sql.rfind('/').ok_or_else(AppError::internal)?; + sql.truncate(index); + sql = sql.trim_end().to_owned(); + } + if sql.is_empty() { + return Err(AppError::invalid( + "invalid_oracle_console_request", + "sql must contain at least one Oracle statement", + )); + } + Ok(sql) +} + +fn looks_like_plsql(sql: &str) -> bool { + let words = oracle_sql_words(sql).unwrap_or_default(); + matches!(words.first().map(String::as_str), Some("BEGIN" | "DECLARE")) + || (words.first().is_some_and(|word| word == "CREATE") + && words.iter().take(8).any(|word| { + matches!( + word.as_str(), + "FUNCTION" | "PROCEDURE" | "TRIGGER" | "PACKAGE" | "TYPE" + ) + })) +} + +fn first_sql_keyword(sql: &str) -> Result, AppError> { + Ok(oracle_sql_words(sql)?.into_iter().next()) +} + +fn oracle_sql_words(sql: &str) -> Result, AppError> { + let bytes = sql.as_bytes(); + let mut words = Vec::new(); + let mut index = 0; + while index < bytes.len() { + match bytes[index] { + b'\'' => skip_quoted(bytes, &mut index, b'\'', b'\'')?, + b'"' => skip_quoted(bytes, &mut index, b'"', b'"')?, + b'-' if bytes.get(index + 1) == Some(&b'-') => skip_line_comment(bytes, &mut index), + b'/' if bytes.get(index + 1) == Some(&b'*') => { + skip_block_comment(bytes, &mut index)?; + } + b'q' | b'Q' if bytes.get(index + 1) == Some(&b'\'') => { + skip_oracle_q_quote(bytes, &mut index)?; + } + byte if byte.is_ascii_alphabetic() || byte == b'_' => { + let start = index; + index += 1; + while bytes.get(index).is_some_and(|byte| { + byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'$' | b'#') + }) { + index += 1; + } + words.push(sql[start..index].to_ascii_uppercase()); + } + _ => index += 1, + } + } + Ok(words) +} + +fn split_oracle_script(sql: &str) -> Result, AppError> { + if looks_like_plsql_without_split(sql) { + return Ok(vec![normalize_preserved_statement(sql)?]); + } + let bytes = sql.as_bytes(); + let mut statements = Vec::new(); + let mut start = 0; + let mut index = 0; + while index < bytes.len() { + match bytes[index] { + b'\'' => skip_quoted(bytes, &mut index, b'\'', b'\'')?, + b'"' => skip_quoted(bytes, &mut index, b'"', b'"')?, + b'-' if bytes.get(index + 1) == Some(&b'-') => skip_line_comment(bytes, &mut index), + b'/' if bytes.get(index + 1) == Some(&b'*') => { + skip_block_comment(bytes, &mut index)?; + } + b'q' | b'Q' if bytes.get(index + 1) == Some(&b'\'') => { + skip_oracle_q_quote(bytes, &mut index)?; + } + b';' => { + let statement = sql[start..index].trim(); + if !statement.is_empty() { + statements.push(statement.to_owned()); + } + index += 1; + start = index; + } + _ => index += 1, + } + } + let statement = sql[start..].trim(); + if !statement.is_empty() { + statements.push(statement.to_owned()); + } + Ok(statements) +} + +fn looks_like_plsql_without_split(sql: &str) -> bool { + let prefix = sql + .trim_start_matches(|character: char| character.is_whitespace()) + .get(..256) + .unwrap_or(sql) + .to_ascii_uppercase(); + prefix.starts_with("BEGIN") + || prefix.starts_with("DECLARE") + || (prefix.starts_with("CREATE") + && [" FUNCTION", " PROCEDURE", " TRIGGER", " PACKAGE", " TYPE"] + .iter() + .any(|keyword| prefix.contains(keyword))) +} + +fn skip_quoted( + bytes: &[u8], + index: &mut usize, + delimiter: u8, + escaped_delimiter: u8, +) -> Result<(), AppError> { + *index += 1; + while *index < bytes.len() { + if bytes[*index] == delimiter { + if bytes.get(*index + 1) == Some(&escaped_delimiter) { + *index += 2; + continue; + } + *index += 1; + return Ok(()); + } + *index += 1; + } + Err(AppError::invalid( + "invalid_query_request", + "Oracle SQL contains an unterminated quoted value", + )) +} + +fn skip_oracle_q_quote(bytes: &[u8], index: &mut usize) -> Result<(), AppError> { + let opener = *bytes.get(*index + 2).ok_or_else(|| { + AppError::invalid( + "invalid_query_request", + "Oracle SQL contains an invalid alternative quote", + ) + })?; + let closer = match opener { + b'[' => b']', + b'{' => b'}', + b'(' => b')', + b'<' => b'>', + other => other, + }; + *index += 3; + while *index + 1 < bytes.len() { + if bytes[*index] == closer && bytes[*index + 1] == b'\'' { + *index += 2; + return Ok(()); + } + *index += 1; + } + Err(AppError::invalid( + "invalid_query_request", + "Oracle SQL contains an unterminated alternative quote", + )) +} + +fn skip_line_comment(bytes: &[u8], index: &mut usize) { + *index += 2; + while *index < bytes.len() && !matches!(bytes[*index], b'\r' | b'\n') { + *index += 1; + } +} + +fn skip_block_comment(bytes: &[u8], index: &mut usize) -> Result<(), AppError> { + *index += 2; + while *index + 1 < bytes.len() { + if bytes[*index] == b'*' && bytes[*index + 1] == b'/' { + *index += 2; + return Ok(()); + } + *index += 1; + } + Err(AppError::invalid( + "invalid_query_request", + "Oracle SQL contains an unterminated block comment", + )) +} + +async fn metadata_query( + application: &Application, + datasource_id: &str, + sql: &str, + parameters: Vec, +) -> Result { + let resolved = resolve_native_connection(application, datasource_id).await?; + let managed = open_resolved_connection(&resolved).await?; + let result = tokio::time::timeout( + OPERATION_TIMEOUT, + query_all(&managed.connection, sql, ¶meters, MAX_METADATA_ROWS), + ) + .await; + match result { + Ok(Ok(result)) => { + if let Err(error) = managed.close().await { + tracing::warn!(error = %error, "Oracle metadata connection cleanup failed"); + } + Ok(result) + } + Ok(Err(error)) => { + managed.abandon().await; + Err(error) + } + Err(_) => { + managed.abandon().await; + Err(oracle_operation_timeout("metadata query")) + } + } +} + +async fn query_all( + connection: &Connection, + sql: &str, + parameters: &[OracleValue], + max_rows: usize, +) -> Result { + let mut page = connection + .query(sql, parameters) + .await + .map_err(|error| oracle_query_error(&error))?; + let columns = page.columns.clone(); + validate_result_columns(&columns)?; + if columns.len() > MAX_COLUMNS { + return Err(resource_error( + "oracle_result_too_wide", + format!("Oracle returned more than {MAX_COLUMNS} columns"), + )); + } + let cursor_id = page.cursor_id; + let rows_affected = page.rows_affected; + let mut rows = Vec::new(); + let mut result_bytes = 0_u64; + loop { + if rows.len().saturating_add(page.rows.len()) > max_rows { + return Err(resource_error( + "oracle_metadata_too_large", + format!("Oracle metadata exceeded the {max_rows} row safety limit"), + )); + } + for row in &page.rows { + result_bytes = result_bytes + .checked_add(oracle_metadata_row_size(row)?) + .ok_or_else(|| { + resource_error( + "oracle_metadata_too_large", + "Oracle metadata exceeded its byte safety limit", + ) + })?; + if result_bytes > MAX_METADATA_RESULT_BYTES { + return Err(resource_error( + "oracle_metadata_too_large", + format!( + "Oracle metadata exceeded the {MAX_METADATA_RESULT_BYTES} byte safety limit" + ), + )); + } + } + rows.append(&mut page.rows); + if !page.has_more_rows { + break; + } + page = connection + .fetch_more(cursor_id, &columns, FETCH_ROWS) + .await + .map_err(|error| oracle_query_error(&error))?; + } + Ok(OracleQueryResult { + columns, + rows, + rows_affected, + has_more_rows: false, + cursor_id, + }) +} + +fn oracle_metadata_row_size(row: &OracleRow) -> Result { + let mut size = 0_u64; + for index in 0..row.len() { + let value = row.get(index).ok_or_else(result_decode_error)?; + size = size + .checked_add(oracle_metadata_value_size(value)?) + .and_then(|size| size.checked_add(8)) + .ok_or_else(|| { + resource_error( + "oracle_metadata_too_large", + "Oracle metadata exceeded its byte safety limit", + ) + })?; + } + Ok(size) +} + +fn oracle_metadata_value_size(value: &OracleValue) -> Result { + let size = match value { + OracleValue::Null | OracleValue::Lob(LobValue::Null | LobValue::Empty) => 0, + OracleValue::String(value) => { + u64::try_from(value.len()).map_err(|_| AppError::internal())? + } + OracleValue::Bytes(value) => { + u64::try_from(value.len()).map_err(|_| AppError::internal())? + } + OracleValue::Integer(_) | OracleValue::Float(_) => 8, + OracleValue::Number(value) => { + u64::try_from(value.as_str().len()).map_err(|_| AppError::internal())? + } + OracleValue::Date(_) => 7, + OracleValue::Timestamp(_) => 13, + OracleValue::RowId(value) => { + u64::try_from(format!("{value}").len()).map_err(|_| AppError::internal())? + } + OracleValue::Boolean(_) => 1, + OracleValue::Lob(LobValue::Inline(value)) => { + u64::try_from(value.len()).map_err(|_| AppError::internal())? + } + OracleValue::Lob(LobValue::Locator(locator)) => locator.size(), + OracleValue::Json(value) => { + u64::try_from(value.to_string().len()).map_err(|_| AppError::internal())? + } + OracleValue::Vector(_) | OracleValue::Cursor(_) | OracleValue::Collection(_) => { + return Err(AppError::invalid( + "oracle_metadata_type_not_supported", + "Oracle metadata returned a value type that the native driver does not support", + )); + } + }; + Ok(size) +} + +fn validate_metadata_identifier(value: &str, field: &str) -> Result<(), AppError> { + if value.trim().is_empty() + || value.len() > MAX_IDENTIFIER_BYTES + || value.contains('\0') + || value.chars().any(char::is_control) + { + return Err(AppError::invalid( + "invalid_oracle_metadata_request", + format!("{field} is invalid"), + )); + } + Ok(()) +} + +fn validate_name_pattern(value: &str) -> Result<(), AppError> { + if value.len() > MAX_IDENTIFIER_BYTES * 4 + || value.contains('\0') + || value.chars().any(char::is_control) + { + return Err(AppError::invalid( + "invalid_oracle_metadata_request", + "namePattern is invalid", + )); + } + Ok(()) +} + +fn required_text(row: &OracleRow, index: usize) -> Result { + optional_text(row, index)?.ok_or_else(result_decode_error) +} + +fn optional_text(row: &OracleRow, index: usize) -> Result, AppError> { + let Some(value) = row.get(index) else { + return Err(result_decode_error()); + }; + let value = match value { + OracleValue::Null | OracleValue::Lob(LobValue::Null) => return Ok(None), + OracleValue::String(value) => value.clone(), + OracleValue::Bytes(value) => { + String::from_utf8(value.clone()).map_err(|_| result_decode_error())? + } + OracleValue::Integer(value) => value.to_string(), + OracleValue::Float(value) => value.to_string(), + OracleValue::Number(value) => value.as_str().to_owned(), + OracleValue::Date(value) => format_oracle_date(*value), + OracleValue::Timestamp(value) if value.has_timezone() => { + format_oracle_timestamp_tz(*value)? + } + OracleValue::Timestamp(value) => format_oracle_timestamp(*value), + OracleValue::RowId(value) => format!("{value}"), + OracleValue::Boolean(value) => value.to_string(), + OracleValue::Json(value) => value.to_string(), + OracleValue::Lob(LobValue::Inline(value)) => { + String::from_utf8(value.to_vec()).map_err(|_| result_decode_error())? + } + OracleValue::Lob(LobValue::Empty) => String::new(), + OracleValue::Lob(LobValue::Locator(_)) => return Err(result_decode_error()), + other => oracle_display(other), + }; + validate_result_scalar(value.len())?; + Ok(Some(value)) +} + +fn optional_i32(row: &OracleRow, index: usize) -> Result, AppError> { + optional_text(row, index)? + .map(|value| value.parse::().map_err(|_| result_decode_error())) + .transpose() +} + +fn optional_bool(row: &OracleRow, index: usize) -> Result, AppError> { + optional_text(row, index)? + .map(|value| match value.to_ascii_uppercase().as_str() { + "Y" | "YES" | "TRUE" | "1" => Ok(true), + "N" | "NO" | "FALSE" | "0" => Ok(false), + _ => Err(result_decode_error()), + }) + .transpose() +} + +pub(crate) async fn list_databases( + application: &Application, + request: ListDatabasesRequest, +) -> Result { + let result = metadata_query( + application, + &request.datasource_id, + "SELECT SYS_CONTEXT('USERENV', 'DB_NAME'), SYS_CONTEXT('USERENV', 'CURRENT_USER') FROM DUAL", + Vec::new(), + ) + .await?; + let row = result.rows.first().ok_or_else(result_decode_error)?; + Ok(DatabaseList { + items: vec![DatabaseMetadata { + name: required_text(row, 0)?, + owner: required_text(row, 1)?, + ..DatabaseMetadata::default() + }], + }) +} + +pub(crate) async fn list_schemas( + application: &Application, + request: ListSchemasRequest, +) -> Result { + let result = metadata_query( + application, + &request.datasource_id, + "SELECT USERNAME FROM ALL_USERS ORDER BY USERNAME", + Vec::new(), + ) + .await?; + let items = result + .rows + .iter() + .map(|row| { + let name = required_text(row, 0)?; + Ok(SchemaMetadata { + database_name: request.database_name.clone(), + owner: name.clone(), + system: oracle_system_schema(&name), + name, + ..SchemaMetadata::default() + }) + }) + .collect::>()?; + Ok(SchemaList { items }) +} + +fn oracle_system_schema(name: &str) -> bool { + ORACLE_SYSTEM_SCHEMAS + .iter() + .any(|candidate| candidate.eq_ignore_ascii_case(name)) +} + +pub(crate) async fn list_tables( + application: &Application, + request: ListTablesRequest, +) -> Result { + validate_metadata_identifier(&request.scope.schema_name, "schemaName")?; + validate_name_pattern(&request.name_pattern)?; + let (sql, parameters) = if request.name_pattern.is_empty() { + ( + "SELECT T.OWNER, T.TABLE_NAME, C.COMMENTS, T.TABLESPACE_NAME, T.NUM_ROWS, T.BLOCKS \ + FROM ALL_TABLES T LEFT JOIN ALL_TAB_COMMENTS C \ + ON C.OWNER = T.OWNER AND C.TABLE_NAME = T.TABLE_NAME AND C.TABLE_TYPE = 'TABLE' \ + WHERE T.OWNER = :1 ORDER BY T.TABLE_NAME", + vec![OracleValue::String(request.scope.schema_name.clone())], + ) + } else { + ( + "SELECT T.OWNER, T.TABLE_NAME, C.COMMENTS, T.TABLESPACE_NAME, T.NUM_ROWS, T.BLOCKS \ + FROM ALL_TABLES T LEFT JOIN ALL_TAB_COMMENTS C \ + ON C.OWNER = T.OWNER AND C.TABLE_NAME = T.TABLE_NAME AND C.TABLE_TYPE = 'TABLE' \ + WHERE T.OWNER = :1 AND T.TABLE_NAME LIKE :2 ORDER BY T.TABLE_NAME", + vec![ + OracleValue::String(request.scope.schema_name.clone()), + OracleValue::String(request.name_pattern.clone()), + ], + ) + }; + let result = metadata_query(application, &request.scope.datasource_id, sql, parameters).await?; + let items = result + .rows + .iter() + .map(|row| table_metadata(&request.scope.database_name, row, "TABLE")) + .collect::>()?; + Ok(TableList { items }) +} + +fn table_metadata( + database_name: &str, + row: &OracleRow, + table_type: &str, +) -> Result { + Ok(TableMetadata { + database_name: database_name.to_owned(), + schema_name: required_text(row, 0)?, + name: required_text(row, 1)?, + table_type: table_type.to_owned(), + comment: optional_text(row, 2)?.unwrap_or_default(), + database_type: ORACLE_DATABASE_TYPE.to_owned(), + tablespace: optional_text(row, 3)?.unwrap_or_default(), + rows: optional_text(row, 4)?, + data_length: optional_text(row, 5)?, + ..TableMetadata::default() + }) +} + +pub(crate) async fn list_columns( + application: &Application, + request: ListColumnsRequest, +) -> Result { + let scope = &request.table.scope; + validate_metadata_identifier(&scope.schema_name, "schemaName")?; + validate_metadata_identifier(&request.table.table_name, "tableName")?; + let result = metadata_query( + application, + &scope.datasource_id, + "SELECT C.COLUMN_NAME, C.DATA_TYPE, C.DATA_DEFAULT_VC, CC.COMMENTS, C.NULLABLE, \ + C.COLUMN_ID, C.DATA_LENGTH, C.DATA_PRECISION, C.DATA_SCALE, C.CHAR_LENGTH, \ + C.CHAR_USED, PK.CONSTRAINT_NAME, PK.POSITION, C.IDENTITY_COLUMN, C.VIRTUAL_COLUMN \ + FROM ALL_TAB_COLS C \ + LEFT JOIN ALL_COL_COMMENTS CC ON CC.OWNER = C.OWNER \ + AND CC.TABLE_NAME = C.TABLE_NAME AND CC.COLUMN_NAME = C.COLUMN_NAME \ + LEFT JOIN (SELECT AC.OWNER, ACC.TABLE_NAME, ACC.COLUMN_NAME, AC.CONSTRAINT_NAME, ACC.POSITION \ + FROM ALL_CONSTRAINTS AC JOIN ALL_CONS_COLUMNS ACC \ + ON ACC.OWNER = AC.OWNER AND ACC.CONSTRAINT_NAME = AC.CONSTRAINT_NAME \ + WHERE AC.CONSTRAINT_TYPE = 'P') PK \ + ON PK.OWNER = C.OWNER AND PK.TABLE_NAME = C.TABLE_NAME AND PK.COLUMN_NAME = C.COLUMN_NAME \ + WHERE C.OWNER = :1 AND C.TABLE_NAME = :2 AND C.HIDDEN_COLUMN = 'NO' \ + ORDER BY C.COLUMN_ID", + vec![ + OracleValue::String(scope.schema_name.clone()), + OracleValue::String(request.table.table_name.clone()), + ], + ) + .await?; + let items = result + .rows + .iter() + .map(|row| column_metadata(&request, row)) + .collect::>()?; + Ok(ColumnList { items }) +} + +fn column_metadata( + request: &ListColumnsRequest, + row: &OracleRow, +) -> Result { + let column_type = required_text(row, 1)?; + let primary_key_name = optional_text(row, 11)?.unwrap_or_default(); + let char_used = optional_text(row, 10)?.unwrap_or_default(); + Ok(ColumnMetadata { + database_name: request.table.scope.database_name.clone(), + schema_name: request.table.scope.schema_name.clone(), + table_name: request.table.table_name.clone(), + name: required_text(row, 0)?, + data_type: Some(oracle_metadata_jdbc_type(&column_type)), + column_type, + default_value: optional_text(row, 2)?, + auto_increment: optional_bool(row, 13)?, + comment: optional_text(row, 3)?.unwrap_or_default(), + primary_key: Some(!primary_key_name.is_empty()), + primary_key_name, + primary_key_order: optional_i32(row, 12)?.unwrap_or_default(), + column_size: optional_i32(row, 7)?.or(optional_i32(row, 9)?), + buffer_length: optional_i32(row, 6)?, + decimal_digits: optional_i32(row, 8)?, + char_octet_length: optional_i32(row, 6)?, + ordinal_position: optional_i32(row, 5)?, + nullable: optional_text(row, 4)?.map(|value| i32::from(value.eq_ignore_ascii_case("Y"))), + generated_column: optional_bool(row, 14)?, + unit: match char_used.as_str() { + "C" => "CHAR".to_owned(), + "B" => "BYTE".to_owned(), + _ => String::new(), + }, + ..ColumnMetadata::default() + }) +} + +fn oracle_metadata_jdbc_type(data_type: &str) -> i32 { + match data_type.to_ascii_uppercase().as_str() { + "VARCHAR2" | "VARCHAR" => 12, + "NVARCHAR2" => -9, + "CHAR" => 1, + "NCHAR" => -15, + "NUMBER" | "DECIMAL" | "NUMERIC" => 2, + "FLOAT" | "BINARY_FLOAT" => 6, + "BINARY_DOUBLE" => 8, + "DATE" | "TIMESTAMP" | "TIMESTAMP WITH LOCAL TIME ZONE" => 93, + "TIMESTAMP WITH TIME ZONE" => 2_014, + "RAW" => -3, + "LONG RAW" => -4, + "LONG" => -1, + "CLOB" | "NCLOB" => 2_005, + "BLOB" => 2_004, + "BFILE" => -13, + "ROWID" | "UROWID" => -8, + "BOOLEAN" => 16, + _ => 1_111, + } +} + +pub(crate) async fn list_views( + application: &Application, + request: ListViewsRequest, +) -> Result { + validate_metadata_identifier(&request.scope.schema_name, "schemaName")?; + validate_name_pattern(&request.name_pattern)?; + let (sql, parameters) = if request.name_pattern.is_empty() { + ( + "SELECT V.OWNER, V.VIEW_NAME, C.COMMENTS, CAST(NULL AS VARCHAR2(128)), \ + CAST(NULL AS NUMBER), CAST(NULL AS NUMBER) \ + FROM ALL_VIEWS V LEFT JOIN ALL_TAB_COMMENTS C \ + ON C.OWNER = V.OWNER AND C.TABLE_NAME = V.VIEW_NAME AND C.TABLE_TYPE = 'VIEW' \ + WHERE V.OWNER = :1 ORDER BY V.VIEW_NAME", + vec![OracleValue::String(request.scope.schema_name.clone())], + ) + } else { + ( + "SELECT V.OWNER, V.VIEW_NAME, C.COMMENTS, CAST(NULL AS VARCHAR2(128)), \ + CAST(NULL AS NUMBER), CAST(NULL AS NUMBER) \ + FROM ALL_VIEWS V LEFT JOIN ALL_TAB_COMMENTS C \ + ON C.OWNER = V.OWNER AND C.TABLE_NAME = V.VIEW_NAME AND C.TABLE_TYPE = 'VIEW' \ + WHERE V.OWNER = :1 AND V.VIEW_NAME LIKE :2 ORDER BY V.VIEW_NAME", + vec![ + OracleValue::String(request.scope.schema_name.clone()), + OracleValue::String(request.name_pattern.clone()), + ], + ) + }; + let result = metadata_query(application, &request.scope.datasource_id, sql, parameters).await?; + let items = result + .rows + .iter() + .map(|row| table_metadata(&request.scope.database_name, row, "VIEW")) + .collect::>()?; + Ok(ViewList { items }) +} + +pub(crate) async fn get_view( + application: &Application, + request: MetadataObjectRef, +) -> Result { + validate_metadata_identifier(&request.scope.schema_name, "schemaName")?; + validate_metadata_identifier(&request.object_name, "viewName")?; + let result = metadata_query( + application, + &request.scope.datasource_id, + "SELECT V.OWNER, V.VIEW_NAME, C.COMMENTS, CAST(NULL AS VARCHAR2(128)), \ + CAST(NULL AS NUMBER), CAST(NULL AS NUMBER) \ + FROM ALL_VIEWS V LEFT JOIN ALL_TAB_COMMENTS C \ + ON C.OWNER = V.OWNER AND C.TABLE_NAME = V.VIEW_NAME AND C.TABLE_TYPE = 'VIEW' \ + WHERE V.OWNER = :1 AND V.VIEW_NAME = :2", + vec![ + OracleValue::String(request.scope.schema_name.clone()), + OracleValue::String(request.object_name.clone()), + ], + ) + .await?; + let mut metadata = table_metadata( + &request.scope.database_name, + result + .rows + .first() + .ok_or_else(|| metadata_not_found("view", &request))?, + "VIEW", + )?; + metadata.ddl = object_ddl( + application, + &request.scope.datasource_id, + &request.scope.schema_name, + &request.object_name, + "VIEW", + ) + .await?; + Ok(metadata) +} + +pub(crate) async fn list_indexes( + application: &Application, + request: ListIndexesRequest, +) -> Result { + validate_metadata_identifier(&request.table.scope.schema_name, "schemaName")?; + validate_metadata_identifier(&request.table.table_name, "tableName")?; + let result = metadata_query( + application, + &request.table.scope.datasource_id, + "SELECT I.OWNER, I.TABLE_NAME, I.INDEX_NAME, I.INDEX_TYPE, I.UNIQUENESS, \ + I.TABLESPACE_NAME, I.STATUS, C.COLUMN_POSITION, C.COLUMN_NAME, C.DESCEND \ + FROM ALL_INDEXES I JOIN ALL_IND_COLUMNS C \ + ON C.INDEX_OWNER = I.OWNER AND C.INDEX_NAME = I.INDEX_NAME \ + WHERE I.TABLE_OWNER = :1 AND I.TABLE_NAME = :2 \ + ORDER BY I.INDEX_NAME, C.COLUMN_POSITION", + vec![ + OracleValue::String(request.table.scope.schema_name.clone()), + OracleValue::String(request.table.table_name.clone()), + ], + ) + .await?; + let mut indexes = BTreeMap::::new(); + for row in &result.rows { + let name = required_text(row, 2)?; + let unique = required_text(row, 4)?.eq_ignore_ascii_case("UNIQUE"); + let column = IndexColumnMetadata { + database_name: request.table.scope.database_name.clone(), + schema_name: request.table.scope.schema_name.clone(), + table_name: request.table.table_name.clone(), + index_name: name.clone(), + column_name: optional_text(row, 8)?.unwrap_or_default(), + ordinal_position: optional_i32(row, 7)?, + non_unique: Some(!unique), + index_qualifier: required_text(row, 0)?, + sort_order: optional_text(row, 9)?.unwrap_or_default(), + ..IndexColumnMetadata::default() + }; + indexes + .entry(name.clone()) + .or_insert_with(|| IndexMetadata { + database_name: request.table.scope.database_name.clone(), + schema_name: request.table.scope.schema_name.clone(), + table_name: request.table.table_name.clone(), + name, + index_type: if unique { "Unique" } else { "Normal" }.to_owned(), + unique: Some(unique), + method: optional_text(row, 3).ok().flatten().unwrap_or_default(), + comment: optional_text(row, 6).ok().flatten().unwrap_or_default(), + ..IndexMetadata::default() + }) + .columns + .push(column); + } + Ok(IndexList { + items: indexes.into_values().collect(), + }) +} + +async fn object_ddl( + application: &Application, + datasource_id: &str, + schema_name: &str, + object_name: &str, + object_type: &str, +) -> Result { + validate_metadata_identifier(schema_name, "schemaName")?; + validate_metadata_identifier(object_name, "objectName")?; + let resolved = resolve_native_connection(application, datasource_id).await?; + let managed = open_resolved_connection(&resolved).await?; + let ddl = tokio::time::timeout(OPERATION_TIMEOUT, async { + let result = query_all( + &managed.connection, + "SELECT DBMS_METADATA.GET_DDL(:1, :2, :3) FROM DUAL", + &[ + OracleValue::String(object_type.to_owned()), + OracleValue::String(object_name.to_owned()), + OracleValue::String(schema_name.to_owned()), + ], + 2, + ) + .await?; + let row = result.rows.first().ok_or_else(|| { + AppError::not_found("oracle_metadata_not_found", "Oracle object does not exist") + })?; + let value = row.get(0).cloned().ok_or_else(result_decode_error)?; + let oracle_type = result + .columns + .first() + .map_or(OracleType::Clob, |column| column.oracle_type); + oracle_text(&managed.connection, value, oracle_type).await + }) + .await; + match ddl { + Ok(Ok(ddl)) => { + if let Err(error) = managed.close().await { + tracing::warn!(error = %error, "Oracle DDL connection cleanup failed"); + } + Ok(ddl) + } + Ok(Err(error)) => { + managed.abandon().await; + Err(error) + } + Err(_) => { + managed.abandon().await; + Err(oracle_operation_timeout("DDL lookup and LOB read")) + } + } +} + +fn metadata_not_found(kind: &str, request: &MetadataObjectRef) -> AppError { + AppError::not_found( + "oracle_metadata_not_found", + format!( + "Oracle {kind} {}.{} does not exist", + request.scope.schema_name, request.object_name + ), + ) +} + +pub(crate) async fn list_imported_keys( + application: &Application, + request: ListTableKeysRequest, +) -> Result { + list_foreign_keys(application, &request, false).await +} + +pub(crate) async fn list_exported_keys( + application: &Application, + request: ListTableKeysRequest, +) -> Result { + list_foreign_keys(application, &request, true).await +} + +async fn list_foreign_keys( + application: &Application, + request: &ListTableKeysRequest, + exported: bool, +) -> Result { + let scope = &request.table.scope; + validate_metadata_identifier(&scope.schema_name, "schemaName")?; + validate_metadata_identifier(&request.table.table_name, "tableName")?; + let filter = if exported { + "PC.OWNER = :1 AND PC.TABLE_NAME = :2" + } else { + "FC.OWNER = :1 AND FC.TABLE_NAME = :2" + }; + let sql = format!( + "SELECT PC.OWNER, PC.TABLE_NAME, PCC.COLUMN_NAME, \ + FC.OWNER, FC.TABLE_NAME, FCC.COLUMN_NAME, FCC.POSITION, \ + FC.DELETE_RULE, FC.CONSTRAINT_NAME, PC.CONSTRAINT_NAME, \ + FC.DEFERRABLE, FC.DEFERRED \ + FROM ALL_CONSTRAINTS FC \ + JOIN ALL_CONS_COLUMNS FCC ON FCC.OWNER = FC.OWNER \ + AND FCC.CONSTRAINT_NAME = FC.CONSTRAINT_NAME \ + JOIN ALL_CONSTRAINTS PC ON PC.OWNER = FC.R_OWNER \ + AND PC.CONSTRAINT_NAME = FC.R_CONSTRAINT_NAME \ + JOIN ALL_CONS_COLUMNS PCC ON PCC.OWNER = PC.OWNER \ + AND PCC.CONSTRAINT_NAME = PC.CONSTRAINT_NAME \ + AND PCC.POSITION = FCC.POSITION \ + WHERE FC.CONSTRAINT_TYPE = 'R' AND {filter} \ + ORDER BY FC.CONSTRAINT_NAME, FCC.POSITION" + ); + let result = metadata_query( + application, + &scope.datasource_id, + &sql, + vec![ + OracleValue::String(scope.schema_name.clone()), + OracleValue::String(request.table.table_name.clone()), + ], + ) + .await?; + let items = result + .rows + .iter() + .map(|row| foreign_key_metadata(&scope.database_name, row)) + .collect::>()?; + Ok(ForeignKeyList { items }) +} + +fn foreign_key_metadata( + database_name: &str, + row: &OracleRow, +) -> Result { + let delete_rule = match required_text(row, 7)?.to_ascii_uppercase().as_str() { + "CASCADE" => 0, + "RESTRICT" => 1, + "SET NULL" => 2, + "SET DEFAULT" => 4, + _ => 3, + }; + let deferrability = if required_text(row, 10)?.eq_ignore_ascii_case("NOT DEFERRABLE") { + 7 + } else if required_text(row, 11)?.eq_ignore_ascii_case("DEFERRED") { + 5 + } else { + 6 + }; + Ok(ForeignKeyMetadata { + primary_table_database: database_name.to_owned(), + primary_table_schema: required_text(row, 0)?, + primary_table_name: required_text(row, 1)?, + primary_column_name: required_text(row, 2)?, + foreign_table_database: database_name.to_owned(), + foreign_table_schema: required_text(row, 3)?, + foreign_table_name: required_text(row, 4)?, + foreign_column_name: required_text(row, 5)?, + key_sequence: optional_i32(row, 6)?.unwrap_or_default(), + update_rule: 3, + delete_rule, + foreign_key_name: required_text(row, 8)?, + primary_key_name: required_text(row, 9)?, + deferrability, + }) +} + +pub(crate) async fn list_primary_keys( + application: &Application, + request: ListTableKeysRequest, +) -> Result { + let scope = &request.table.scope; + validate_metadata_identifier(&scope.schema_name, "schemaName")?; + validate_metadata_identifier(&request.table.table_name, "tableName")?; + let result = metadata_query( + application, + &scope.datasource_id, + "SELECT AC.OWNER, ACC.TABLE_NAME, ACC.COLUMN_NAME, AC.CONSTRAINT_NAME \ + FROM ALL_CONSTRAINTS AC JOIN ALL_CONS_COLUMNS ACC \ + ON ACC.OWNER = AC.OWNER AND ACC.CONSTRAINT_NAME = AC.CONSTRAINT_NAME \ + WHERE AC.CONSTRAINT_TYPE = 'P' AND AC.OWNER = :1 AND ACC.TABLE_NAME = :2 \ + ORDER BY ACC.POSITION", + vec![ + OracleValue::String(scope.schema_name.clone()), + OracleValue::String(request.table.table_name.clone()), + ], + ) + .await?; + let items = result + .rows + .iter() + .map(|row| { + Ok(PrimaryKeyMetadata { + database_name: scope.database_name.clone(), + schema_name: required_text(row, 0)?, + table_name: required_text(row, 1)?, + column_name: required_text(row, 2)?, + name: required_text(row, 3)?, + }) + }) + .collect::>()?; + Ok(PrimaryKeyList { items }) +} + +pub(crate) async fn list_functions( + application: &Application, + request: ListRoutinesRequest, +) -> Result { + validate_metadata_identifier(&request.scope.schema_name, "schemaName")?; + let result = metadata_query( + application, + &request.scope.datasource_id, + "SELECT OWNER, OBJECT_NAME FROM ALL_OBJECTS \ + WHERE OWNER = :1 AND OBJECT_TYPE = 'FUNCTION' ORDER BY OBJECT_NAME", + vec![OracleValue::String(request.scope.schema_name.clone())], + ) + .await?; + let items = result + .rows + .iter() + .map(|row| { + let name = required_text(row, 1)?; + Ok(FunctionMetadata { + database_name: request.scope.database_name.clone(), + schema_name: required_text(row, 0)?, + name: name.clone(), + function_type: Some(1), + specific_name: name, + ..FunctionMetadata::default() + }) + }) + .collect::>()?; + Ok(FunctionList { items }) +} + +pub(crate) async fn get_function( + application: &Application, + request: MetadataObjectRef, +) -> Result { + let body = routine_source(application, &request, "FUNCTION").await?; + Ok(FunctionMetadata { + database_name: request.scope.database_name, + schema_name: request.scope.schema_name, + name: request.object_name.clone(), + function_type: Some(1), + specific_name: request.object_name, + body, + ..FunctionMetadata::default() + }) +} + +pub(crate) async fn list_function_parameters( + application: &Application, + request: MetadataObjectRef, +) -> Result { + let rows = routine_parameters(application, &request).await?; + let items = rows + .iter() + .map(|row| function_parameter_metadata(&request, row)) + .collect::>()?; + Ok(FunctionParameterList { items }) +} + +fn function_parameter_metadata( + request: &MetadataObjectRef, + row: &OracleRow, +) -> Result { + let position = optional_i32(row, 1)?.unwrap_or_default(); + let mode = optional_text(row, 2)?.unwrap_or_default(); + Ok(FunctionParameterMetadata { + function_database: request.scope.database_name.clone(), + function_schema: request.scope.schema_name.clone(), + function_name: request.object_name.clone(), + column_name: optional_text(row, 0)?.unwrap_or_default(), + column_type: Some(if position == 0 { + 4 + } else { + routine_parameter_mode(&mode, false) + }), + data_type: Some(oracle_metadata_jdbc_type(&required_text(row, 3)?)), + type_name: required_text(row, 3)?, + length: optional_i32(row, 4)?, + precision: optional_i32(row, 5)?, + scale: optional_i32(row, 6)?, + radix: optional_i32(row, 7)?, + nullable: Some(2), + ordinal_position: Some(position), + is_nullable: String::new(), + specific_name: request.object_name.clone(), + ..FunctionParameterMetadata::default() + }) +} + +pub(crate) async fn list_procedures( + application: &Application, + request: ListRoutinesRequest, +) -> Result { + validate_metadata_identifier(&request.scope.schema_name, "schemaName")?; + let result = metadata_query( + application, + &request.scope.datasource_id, + "SELECT OWNER, OBJECT_NAME FROM ALL_OBJECTS \ + WHERE OWNER = :1 AND OBJECT_TYPE = 'PROCEDURE' ORDER BY OBJECT_NAME", + vec![OracleValue::String(request.scope.schema_name.clone())], + ) + .await?; + let items = result + .rows + .iter() + .map(|row| { + let name = required_text(row, 1)?; + Ok(ProcedureMetadata { + database_name: request.scope.database_name.clone(), + schema_name: required_text(row, 0)?, + name: name.clone(), + procedure_type: Some(1), + specific_name: name, + ..ProcedureMetadata::default() + }) + }) + .collect::>()?; + Ok(ProcedureList { items }) +} + +pub(crate) async fn get_procedure( + application: &Application, + request: MetadataObjectRef, +) -> Result { + let body = routine_source(application, &request, "PROCEDURE").await?; + Ok(ProcedureMetadata { + database_name: request.scope.database_name, + schema_name: request.scope.schema_name, + name: request.object_name.clone(), + procedure_type: Some(1), + specific_name: request.object_name, + body, + ..ProcedureMetadata::default() + }) +} + +pub(crate) async fn list_procedure_parameters( + application: &Application, + request: MetadataObjectRef, +) -> Result { + let rows = routine_parameters(application, &request).await?; + let items = rows + .iter() + .filter(|row| optional_i32(row, 1).ok().flatten().unwrap_or_default() > 0) + .map(|row| procedure_parameter_metadata(&request, row)) + .collect::>()?; + Ok(ProcedureParameterList { items }) +} + +fn procedure_parameter_metadata( + request: &MetadataObjectRef, + row: &OracleRow, +) -> Result { + let mode = optional_text(row, 2)?.unwrap_or_default(); + Ok(ProcedureParameterMetadata { + procedure_database: request.scope.database_name.clone(), + procedure_schema: request.scope.schema_name.clone(), + procedure_name: request.object_name.clone(), + column_name: optional_text(row, 0)?.unwrap_or_default(), + column_type: Some(routine_parameter_mode(&mode, true)), + data_type: Some(oracle_metadata_jdbc_type(&required_text(row, 3)?)), + type_name: required_text(row, 3)?, + length: optional_i32(row, 4)?, + precision: optional_i32(row, 5)?, + scale: optional_i32(row, 6)?, + radix: optional_i32(row, 7)?, + nullable: Some(2), + ordinal_position: optional_i32(row, 1)?, + is_nullable: String::new(), + specific_name: request.object_name.clone(), + ..ProcedureParameterMetadata::default() + }) +} + +const fn routine_parameter_mode(mode: &str, procedure: bool) -> i32 { + match (mode.as_bytes(), procedure) { + (b"IN", _) => 1, + (b"IN/OUT" | b"IN OUT", _) => 2, + (b"OUT", true) => 4, + (b"OUT", false) => 3, + _ => 0, + } +} + +async fn routine_parameters( + application: &Application, + request: &MetadataObjectRef, +) -> Result, AppError> { + validate_metadata_identifier(&request.scope.schema_name, "schemaName")?; + validate_metadata_identifier(&request.object_name, "routineName")?; + let result = metadata_query( + application, + &request.scope.datasource_id, + "SELECT ARGUMENT_NAME, POSITION, IN_OUT, DATA_TYPE, DATA_LENGTH, \ + DATA_PRECISION, DATA_SCALE, RADIX, DEFAULTED \ + FROM ALL_ARGUMENTS WHERE OWNER = :1 AND OBJECT_NAME = :2 \ + AND PACKAGE_NAME IS NULL ORDER BY SEQUENCE", + vec![ + OracleValue::String(request.scope.schema_name.clone()), + OracleValue::String(request.object_name.clone()), + ], + ) + .await?; + Ok(result.rows) +} + +async fn routine_source( + application: &Application, + request: &MetadataObjectRef, + routine_type: &str, +) -> Result { + validate_metadata_identifier(&request.scope.schema_name, "schemaName")?; + validate_metadata_identifier(&request.object_name, "routineName")?; + let result = metadata_query( + application, + &request.scope.datasource_id, + "SELECT TEXT FROM ALL_SOURCE WHERE OWNER = :1 AND NAME = :2 AND TYPE = :3 ORDER BY LINE", + vec![ + OracleValue::String(request.scope.schema_name.clone()), + OracleValue::String(request.object_name.clone()), + OracleValue::String(routine_type.to_owned()), + ], + ) + .await?; + if result.rows.is_empty() { + return Err(metadata_not_found( + &routine_type.to_ascii_lowercase(), + request, + )); + } + let mut body = String::new(); + for row in &result.rows { + body.push_str(&required_text(row, 0)?); + if body.len() > MAX_SQL_BYTES { + return Err(resource_error( + "oracle_routine_source_too_large", + format!("Oracle routine source exceeds {MAX_SQL_BYTES} bytes"), + )); + } + } + Ok(body) +} + +pub(crate) async fn list_triggers( + application: &Application, + request: ListTriggersRequest, +) -> Result { + validate_metadata_identifier(&request.scope.schema_name, "schemaName")?; + let result = metadata_query( + application, + &request.scope.datasource_id, + "SELECT OWNER, TRIGGER_NAME, TRIGGERING_EVENT FROM ALL_TRIGGERS \ + WHERE OWNER = :1 ORDER BY TRIGGER_NAME", + vec![OracleValue::String(request.scope.schema_name.clone())], + ) + .await?; + let items = result + .rows + .iter() + .map(|row| { + Ok(TriggerMetadata { + database_name: request.scope.database_name.clone(), + schema_name: required_text(row, 0)?, + name: required_text(row, 1)?, + event_manipulation: required_text(row, 2)?, + body: String::new(), + }) + }) + .collect::>()?; + Ok(TriggerList { items }) +} + +pub(crate) async fn get_trigger( + application: &Application, + request: MetadataObjectRef, +) -> Result { + validate_metadata_identifier(&request.scope.schema_name, "schemaName")?; + validate_metadata_identifier(&request.object_name, "triggerName")?; + let result = metadata_query( + application, + &request.scope.datasource_id, + "SELECT OWNER, TRIGGER_NAME, TRIGGERING_EVENT FROM ALL_TRIGGERS \ + WHERE OWNER = :1 AND TRIGGER_NAME = :2", + vec![ + OracleValue::String(request.scope.schema_name.clone()), + OracleValue::String(request.object_name.clone()), + ], + ) + .await?; + let row = result + .rows + .first() + .ok_or_else(|| metadata_not_found("trigger", &request))?; + Ok(TriggerMetadata { + database_name: request.scope.database_name.clone(), + schema_name: required_text(row, 0)?, + name: required_text(row, 1)?, + event_manipulation: required_text(row, 2)?, + body: object_ddl( + application, + &request.scope.datasource_id, + &request.scope.schema_name, + &request.object_name, + "TRIGGER", + ) + .await?, + }) +} + +pub(crate) async fn load_er_tables( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, +) -> Result, AppError> { + let tables = list_tables( + application, + ListTablesRequest { + scope: MetadataScope { + datasource_id: datasource_id.to_owned(), + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + }, + name_pattern: String::new(), + }, + ) + .await?; + let mut result = Vec::with_capacity(tables.items.len()); + for table in tables.items { + let table_ref = TableRef { + scope: MetadataScope { + datasource_id: datasource_id.to_owned(), + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + }, + table_name: table.name.clone(), + }; + let columns = list_columns( + application, + ListColumnsRequest { + table: table_ref.clone(), + }, + ) + .await?; + let foreign_keys = + list_imported_keys(application, ListTableKeysRequest { table: table_ref }).await?; + result.push(EntityRelationTable { + name: table.name, + comment: table.comment, + columns: columns + .items + .into_iter() + .map(|column| EntityRelationColumn { + name: column.name, + column_type: column.column_type, + primary_key: column.primary_key.unwrap_or(false), + comment: column.comment, + }) + .collect(), + foreign_keys: foreign_keys + .items + .into_iter() + .map(|key| EntityRelationForeignKey { + primary_table: key.primary_table_name, + primary_column: key.primary_column_name, + foreign_table: key.foreign_table_name, + foreign_column: key.foreign_column_name, + }) + .collect(), + }); + } + Ok(result) +} + +pub(crate) async fn start_table_preview( + application: &Application, + request: TablePreviewRequest, + row_limit: u32, +) -> Result { + if row_limit == 0 || row_limit > MAX_TABLE_PREVIEW_ROWS { + return Err(AppError::invalid( + "invalid_table_preview_request", + format!("rowLimit must be between 1 and {MAX_TABLE_PREVIEW_ROWS}"), + )); + } + let schema_name = quote_identifier(&request.table.scope.schema_name, "schemaName")?; + let table_name = quote_identifier(&request.table.table_name, "tableName")?; + let sql = format!("SELECT * FROM {schema_name}.{table_name} FETCH FIRST {row_limit} ROWS ONLY"); + let accepted = application + .start_read_query(StartQueryRequest { + datasource_id: request.table.scope.datasource_id, + sql: sql.clone(), + parameters: Vec::new(), + limits: QueryLimits { + max_rows: row_limit.to_string(), + max_result_bytes: (8 * 1024 * 1024_u64).to_string(), + batch_rows: row_limit.min(200), + batch_bytes: 1024 * 1024, + result_ttl_seconds: 60 * 60, + }, + }) + .await?; + Ok(TablePreviewAccepted { + operation_id: accepted.operation_id, + sql, + row_limit, + }) +} + +fn build_oracle_create_schema(request: CreateSchemaSqlRequest) -> Result { + let schema = request.schema; + if !schema.owner.trim().is_empty() && !schema.owner.eq_ignore_ascii_case(&schema.name) { + return Err(AppError::invalid( + "oracle_schema_owner_mismatch", + "An Oracle schema name must match its authorization user", + )); + } + let name = quote_identifier(&schema.name, "schemaName")?; + Ok(BuiltSql { + sql: format!("CREATE SCHEMA AUTHORIZATION {name};"), + }) +} + +fn build_oracle_namespace_sql(request: NamespaceSqlRequest) -> Result { + let sql = match request.operation { + NamespaceSqlOperation::CreateDatabase { database } => { + build_oracle_create_database(&database)? + } + NamespaceSqlOperation::AlterDatabase { .. } => { + return Err(oracle_namespace_unsupported( + "oracle_database_alter_unsupported", + "Oracle does not support renaming a database through a connected SQL session", + )); + } + NamespaceSqlOperation::DropDatabase { database_name } => format!( + "DROP DATABASE {};", + quote_identifier(&database_name, "databaseName")? + ), + NamespaceSqlOperation::UseDatabase { .. } => { + return Err(oracle_namespace_unsupported( + "oracle_database_switch_unsupported", + "Oracle selects a database service when opening the connection and cannot switch it with SQL", + )); + } + NamespaceSqlOperation::CreateSchema { schema } => { + return build_oracle_create_schema(CreateSchemaSqlRequest { schema }); + } + NamespaceSqlOperation::AlterSchema { .. } => { + return Err(oracle_namespace_unsupported( + "oracle_schema_rename_unsupported", + "Oracle schemas are database users and cannot be renamed", + )); + } + NamespaceSqlOperation::DropSchema { schema_name } => format!( + "DROP USER {} CASCADE;", + quote_identifier(&schema_name, "schemaName")? + ), + }; + Ok(BuiltSql { sql }) +} + +fn build_oracle_create_database(database: &DatabaseDefinition) -> Result { + let mut sql = format!( + "CREATE DATABASE {}", + quote_identifier(&database.name, "databaseName")? + ); + if !database.charset.trim().is_empty() { + validate_oracle_keyword(&database.charset, "charset")?; + write!(&mut sql, " CHARACTER SET {}", database.charset) + .map_err(|_| AppError::internal())?; + } + sql.push(';'); + Ok(sql) +} + +fn validate_oracle_keyword(value: &str, field: &str) -> Result<(), AppError> { + if value.is_empty() + || value.len() > MAX_IDENTIFIER_BYTES + || !value + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'$' | b'#')) + { + return Err(AppError::invalid( + "invalid_oracle_namespace_request", + format!("{field} is invalid"), + )); + } + Ok(()) +} + +fn oracle_namespace_unsupported(code: &'static str, message: &'static str) -> AppError { + AppError::invalid(code, message) +} + +fn build_oracle_dml(request: DmlSqlRequest) -> Result { + let target = oracle_dml_target(&request.target)?; + let sql = match request.statement { + DmlStatement::SingleInsert { columns, row } => { + oracle_insert_sql(&target, &columns, std::slice::from_ref(&row))? + } + DmlStatement::MultiInsert { columns, rows } => oracle_insert_sql(&target, &columns, &rows)?, + DmlStatement::Update { + assignments, + predicates, + } => oracle_update_sql(&target, &assignments, &predicates)?, + }; + if sql.len() > MAX_SQL_BYTES { + return Err(invalid_oracle_dml( + "The generated Oracle DML exceeds the SQL byte limit", + )); + } + Ok(BuiltSql { sql }) +} + +fn oracle_dml_target(target: &DmlTarget) -> Result { + let table = quote_identifier(&target.table_name, "tableName")?; + let qualifier = target + .schema_name + .as_deref() + .filter(|value| !value.is_empty()) + .or_else(|| { + target + .database_name + .as_deref() + .filter(|value| !value.is_empty()) + }); + qualifier.map_or(Ok(table.clone()), |schema| { + Ok(format!( + "{}.{}", + quote_identifier(schema, "schemaName")?, + table + )) + }) +} + +fn oracle_insert_sql( + target: &str, + columns: &[DmlColumn], + rows: &[DmlRow], +) -> Result { + if columns.is_empty() || rows.is_empty() { + return Err(invalid_oracle_dml( + "Oracle INSERT requires at least one column and row", + )); + } + let column_sql = columns + .iter() + .map(|column| quote_identifier(&column.name, "columnName")) + .collect::, _>>()? + .join(", "); + let values = rows + .iter() + .map(|row| oracle_dml_row(row, columns)) + .collect::, _>>()?; + if let [values] = values.as_slice() { + return Ok(format!( + "INSERT INTO {target} ({column_sql}) VALUES ({values});" + )); + } + let into_clauses = values + .iter() + .map(|values| format!(" INTO {target} ({column_sql}) VALUES ({values})")) + .collect::>() + .join("\n"); + Ok(format!("INSERT ALL\n{into_clauses}\nSELECT 1 FROM DUAL;")) +} + +fn oracle_dml_row(row: &DmlRow, columns: &[DmlColumn]) -> Result { + if row.values.len() != columns.len() { + return Err(invalid_oracle_dml( + "Each Oracle INSERT row must match the selected column count", + )); + } + row.values + .iter() + .zip(columns) + .map(|(value, column)| oracle_dml_value(value, column)) + .collect::, _>>() + .map(|values| values.join(", ")) +} + +fn oracle_update_sql( + target: &str, + assignments: &[DmlAssignment], + predicates: &[DmlAssignment], +) -> Result { + if assignments.is_empty() || predicates.is_empty() { + return Err(invalid_oracle_dml( + "Oracle UPDATE requires assignments and key predicates", + )); + } + let assignments = assignments + .iter() + .map(|assignment| oracle_dml_assignment(assignment, false)) + .collect::, _>>()? + .join(", "); + let predicates = predicates + .iter() + .map(|predicate| oracle_dml_assignment(predicate, true)) + .collect::, _>>()? + .join(" AND "); + Ok(format!( + "UPDATE {target} SET {assignments} WHERE {predicates};" + )) +} + +fn oracle_dml_assignment(assignment: &DmlAssignment, predicate: bool) -> Result { + let column = quote_identifier(&assignment.column.name, "columnName")?; + if predicate && matches!(assignment.value, DmlValue::Null) { + return Ok(format!("{column} IS NULL")); + } + Ok(format!( + "{column} = {}", + oracle_dml_value(&assignment.value, &assignment.column)? + )) +} + +fn oracle_dml_value(value: &DmlValue, column: &DmlColumn) -> Result { + match value { + DmlValue::Null => Ok("NULL".to_owned()), + DmlValue::String(value) => quote_oracle_literal(value), + DmlValue::Decimal(value) => { + validate_oracle_decimal(value)?; + Ok(value.clone()) + } + DmlValue::Boolean(value) => { + if column.data_type_name.eq_ignore_ascii_case("BOOLEAN") { + Ok(if *value { "TRUE" } else { "FALSE" }.to_owned()) + } else { + Ok(if *value { "1" } else { "0" }.to_owned()) + } + } + DmlValue::Temporal { kind, iso8601 } => oracle_dml_temporal(*kind, iso8601), + DmlValue::Binary(value) => oracle_dml_binary(value, column), + } +} + +fn validate_oracle_decimal(value: &str) -> Result<(), AppError> { + if value.len() > MAX_SCALAR_BYTES || oracle_rs::types::encode_oracle_number(value).is_err() { + return Err(invalid_oracle_dml("The Oracle decimal value is invalid")); + } + Ok(()) +} + +fn quote_oracle_literal(value: &str) -> Result { + if value.len() > MAX_SCALAR_BYTES || value.contains('\0') { + return Err(invalid_oracle_dml("The Oracle string value is invalid")); + } + let escaped_length = value + .len() + .checked_add(value.bytes().filter(|byte| *byte == b'\'').count()) + .and_then(|length| length.checked_add(2)) + .filter(|length| *length <= MAX_SCALAR_BYTES) + .ok_or_else(|| invalid_oracle_dml("The escaped Oracle string value is too large"))?; + let mut escaped = String::with_capacity(escaped_length); + for character in value.chars() { + if character == '\'' { + escaped.push('\''); + } + escaped.push(character); + } + Ok(format!("'{escaped}'")) +} + +fn oracle_dml_temporal(kind: DmlTemporalKind, value: &str) -> Result { + match kind { + DmlTemporalKind::Date => NaiveDate::parse_from_str(value, "%Y-%m-%d") + .map(|value| format!("DATE '{}'", value.format("%Y-%m-%d"))) + .map_err(|_| invalid_oracle_dml("The Oracle date value is invalid")), + DmlTemporalKind::Time => parse_oracle_dml_time(value).map(|value| { + let format = if value.nanosecond() == 0 { + "YYYY-MM-DD HH24:MI:SS" + } else { + "YYYY-MM-DD HH24:MI:SS.FF" + }; + format!( + "TO_TIMESTAMP('1970-01-01 {}', '{format}')", + value.format("%H:%M:%S%.f") + ) + }), + DmlTemporalKind::LocalDatetime => parse_oracle_dml_timestamp(value) + .map(|value| format!("TIMESTAMP '{}'", value.format("%Y-%m-%d %H:%M:%S%.f"))), + DmlTemporalKind::OffsetDatetime => DateTime::parse_from_rfc3339(value) + .map(|value| { + format!( + "TIMESTAMP '{} {}'", + value.format("%Y-%m-%d %H:%M:%S%.f"), + value.format("%:z") + ) + }) + .map_err(|_| { + invalid_oracle_dml("The Oracle timestamp with time zone value is invalid") + }), + } +} + +fn parse_oracle_dml_time(value: &str) -> Result { + ["%H:%M:%S%.f", "%H:%M:%S"] + .into_iter() + .find_map(|format| NaiveTime::parse_from_str(value, format).ok()) + .ok_or_else(|| invalid_oracle_dml("The Oracle time value is invalid")) +} + +fn parse_oracle_dml_timestamp(value: &str) -> Result { + ["%Y-%m-%dT%H:%M:%S%.f", "%Y-%m-%d %H:%M:%S%.f"] + .into_iter() + .find_map(|format| NaiveDateTime::parse_from_str(value, format).ok()) + .ok_or_else(|| invalid_oracle_dml("The Oracle timestamp value is invalid")) +} + +fn oracle_dml_binary(value: &[u8], column: &DmlColumn) -> Result { + const MAX_RAW_LITERAL_BYTES: usize = 2_000; + if value.len() > MAX_RAW_LITERAL_BYTES { + return Err(invalid_oracle_dml( + "Oracle inline binary DML values cannot exceed 2000 bytes", + )); + } + let raw = format!("HEXTORAW('{}')", hex::encode_upper(value)); + if column.data_type_name.eq_ignore_ascii_case("BLOB") { + Ok(format!("TO_BLOB({raw})")) + } else { + Ok(raw) + } +} + +fn invalid_oracle_dml(message: impl Into) -> AppError { + AppError::invalid("invalid_oracle_dml", message) +} + +fn quote_identifier(value: &str, field: &str) -> Result { + validate_metadata_identifier(value, field)?; + Ok(format!("\"{}\"", value.replace('"', "\"\""))) +} + +#[cfg(test)] +mod tests { + use chat2db_contract::{DatasourceConnection, DatasourceConnectionProperty}; + use oracle_rs::{ + Error as OracleError, OracleType, Value as OracleValue, config::ServiceMethod, + types::OracleTimestamp, + }; + + use super::{ + OracleNativeDriver, apply_ssh_forward, build_oracle_create_schema, build_oracle_dml, + build_oracle_namespace_sql, connection_config, format_oracle_timestamp_tz, oracle_bool, + oracle_console_connection_error, oracle_console_interrupted, oracle_f32, oracle_f64, + oracle_i64, oracle_query_parameters, oracle_result_type_supported, + parse_jdbc_oracle_target, parse_oracle_url_target, quote_identifier, split_oracle_script, + validate_metadata_identifier, validate_read_sql, + }; + use crate::native_driver::NativeDriver as _; + use crate::native_driver_types::{ + CreateSchemaSqlRequest, DatabaseDefinition, DmlAssignment, DmlColumn, DmlRow, + DmlSqlRequest, DmlStatement, DmlTarget, DmlTemporalKind, DmlValue, NamespaceSqlOperation, + NamespaceSqlRequest, SchemaDefinition, + }; + use crate::query::{DatabaseValue, QueryParameter}; + + #[test] + fn oracle_urls_support_service_name_sid_credentials_and_tls() { + let (service, _, _, tls) = parse_jdbc_oracle_target("db.example:1522/FREEPDB1?ssl=true") + .expect("JDBC service-name URL must parse"); + assert_eq!(service.host, "db.example"); + assert_eq!(service.port, 1_522); + assert_eq!( + service.service, + ServiceMethod::ServiceName("FREEPDB1".to_owned()) + ); + assert!(tls); + + let (sid, _, _, tls) = + parse_jdbc_oracle_target("db.example:1523:ORCL").expect("JDBC SID URL must parse"); + assert_eq!(sid.host, "db.example"); + assert_eq!(sid.port, 1_523); + assert_eq!(sid.service, ServiceMethod::Sid("ORCL".to_owned())); + assert!(!tls); + + let (native, username, password, tls) = + parse_oracle_url_target("oracle://scott:tiger@db.example:2484/ignored?sid=ORCL&tcps=1") + .expect("native SID URL must parse"); + assert_eq!(native.service, ServiceMethod::Sid("ORCL".to_owned())); + assert_eq!(username.as_deref(), Some("scott")); + assert_eq!(password.as_deref(), Some("tiger")); + assert!(tls); + } + + #[test] + fn oracle_connection_properties_reject_duplicates_and_unknown_url_options() { + let duplicate_user = DatasourceConnection { + jdbc_url: "jdbc:oracle:thin:@localhost:1521/FREEPDB1".to_owned(), + properties: vec![ + property("user", "app", false), + property("username", "duplicate", false), + property("password", "secret", true), + ], + read_only: false, + ssh: None, + }; + let error = connection_config(&duplicate_user) + .expect_err("duplicate username aliases must be rejected"); + assert_eq!(error.api_error().code, "invalid_oracle_connection"); + + let error = + parse_oracle_url_target("oracle://localhost/FREEPDB1?wallet=/tmp/not-supported") + .expect_err("unknown native URL properties must be rejected"); + assert_eq!(error.api_error().code, "invalid_oracle_connection"); + + let error = parse_oracle_url_target("oracle://localhost/FREEPDB1?sid=A&sid=B") + .expect_err("duplicate SID properties must be rejected"); + assert_eq!(error.api_error().code, "invalid_oracle_connection"); + } + + #[test] + fn oracle_tcps_configuration_never_panics_and_keeps_tls_identity_through_ssh() { + let connection = DatasourceConnection { + jdbc_url: "jdbc:oracle:thin:@db.example:2484/FREEPDB1?tcps=true".to_owned(), + properties: vec![ + property("user", "app", false), + property("password", "secret", true), + ], + read_only: false, + ssh: None, + }; + let configured = std::panic::catch_unwind(|| connection_config(&connection)); + assert!(configured.is_ok(), "TCPS configuration must not panic"); + let mut config = configured + .expect("panic checked above") + .expect("TCPS configuration must build"); + assert!(config.is_tls_enabled()); + + apply_ssh_forward(&mut config, 31_521); + assert_eq!(config.host, "127.0.0.1"); + assert_eq!(config.port, 31_521); + assert_eq!( + config + .tls_config + .as_ref() + .and_then(|tls| tls.server_name.as_deref()), + Some("db.example") + ); + } + + #[test] + fn oracle_console_only_reports_plain_cancellation_for_read_only_statements() { + assert_eq!( + oracle_console_interrupted(true, Some("cancelled".to_owned())) + .api_error() + .code, + "oracle_console_cancelled" + ); + assert_eq!( + oracle_console_interrupted(false, Some("cancelled".to_owned())) + .api_error() + .code, + "database_write_outcome_unknown" + ); + let connection_error = OracleError::ConnectionClosed; + assert_eq!( + oracle_console_connection_error(true, &connection_error) + .api_error() + .code, + "oracle_connection_failed" + ); + assert_eq!( + oracle_console_connection_error(false, &connection_error) + .api_error() + .code, + "database_write_outcome_unknown" + ); + } + + #[test] + fn oracle_unsupported_result_types_fail_closed() { + for oracle_type in [ + OracleType::BinaryFloat, + OracleType::BinaryDouble, + OracleType::Rowid, + OracleType::Urowid, + OracleType::Bfile, + OracleType::Cursor, + OracleType::Object, + OracleType::Vector, + OracleType::IntervalYm, + OracleType::IntervalDs, + ] { + assert!(!oracle_result_type_supported(oracle_type)); + } + for oracle_type in [ + OracleType::Number, + OracleType::Long, + OracleType::LongRaw, + OracleType::Blob, + OracleType::Clob, + OracleType::Json, + OracleType::Boolean, + ] { + assert!(oracle_result_type_supported(oracle_type)); + } + } + + #[test] + fn oracle_console_splitter_respects_strings_comments_q_quotes_and_plsql() { + let statements = split_oracle_script( + "SELECT ';' FROM DUAL; -- keep ; in comment\n\ + SELECT q'[a;b]' FROM DUAL; SELECT \"odd;name\" FROM DUAL", + ) + .expect("valid Oracle script must split"); + assert_eq!(statements.len(), 3); + assert_eq!(statements[0], "SELECT ';' FROM DUAL"); + assert!(statements[1].ends_with("SELECT q'[a;b]' FROM DUAL")); + assert_eq!(statements[2], "SELECT \"odd;name\" FROM DUAL"); + + let block = "BEGIN\n DBMS_OUTPUT.PUT_LINE(q'[a;b]');\nEND;\n/"; + assert_eq!( + split_oracle_script(block).expect("PL/SQL block must stay intact"), + ["BEGIN\n DBMS_OUTPUT.PUT_LINE(q'[a;b]');\nEND;"] + ); + + for invalid in [ + "SELECT 'unterminated", + "SELECT q'[unterminated", + "SELECT /* open", + ] { + assert!( + split_oracle_script(invalid).is_err(), + "{invalid} must be rejected" + ); + } + } + + #[test] + fn oracle_native_reads_reject_writes_locking_and_multiple_statements() { + for sql in [ + "SELECT 1 FROM DUAL", + "/* leading */ SELECT 'FOR UPDATE' FROM DUAL", + "WITH item AS (SELECT 1 value FROM DUAL) SELECT value FROM item", + ] { + validate_read_sql(sql) + .unwrap_or_else(|error| panic!("{sql} should be accepted: {error}")); + } + + for sql in [ + "UPDATE items SET value = 1", + "SELECT * FROM items FOR UPDATE", + "SELECT * FROM items FOR/**/UPDATE", + "SELECT 1 FROM DUAL; SELECT 2 FROM DUAL", + ] { + let error = validate_read_sql(sql).expect_err("unsafe read SQL must be rejected"); + assert_eq!(error.api_error().code, "oracle_native_query_unsupported"); + } + } + + #[test] + fn oracle_parameters_are_sorted_and_must_be_contiguous() { + let parameters = vec![ + QueryParameter { + position: 2, + value: DatabaseValue::Text("second".to_owned()), + }, + QueryParameter { + position: 1, + value: DatabaseValue::SignedInteger(1), + }, + ]; + let converted = oracle_query_parameters(¶meters) + .expect("out-of-order contiguous parameters must be sorted"); + assert!(matches!(&converted[0], OracleValue::Integer(1))); + assert!(matches!(&converted[1], OracleValue::String(value) if value == "second")); + + for positions in [[1, 1], [1, 3], [0, 1]] { + let parameters = positions.map(|position| QueryParameter { + position, + value: DatabaseValue::Null, + }); + let error = oracle_query_parameters(¶meters) + .expect_err("duplicate, missing, and zero positions must be rejected"); + assert_eq!(error.api_error().code, "invalid_query_parameter"); + } + } + + #[test] + fn oracle_number_strings_convert_to_expected_numeric_wire_values() { + let one = OracleValue::String("1".to_owned()); + assert_eq!(oracle_i64(&one), Some(1)); + assert_eq!(oracle_f32(&one), Some(1.0)); + assert_eq!(oracle_f64(&one), Some(1.0)); + + let fractional = OracleValue::Float(1.5); + assert_eq!(oracle_i64(&fractional), None); + } + + #[test] + fn oracle_boolean_decoding_accepts_only_known_wire_shapes() { + assert_eq!(oracle_bool(&OracleValue::Boolean(true)), Some(true)); + assert_eq!( + oracle_bool(&OracleValue::String("true".to_owned())), + Some(true) + ); + assert_eq!( + oracle_bool(&OracleValue::String("false".to_owned())), + Some(false) + ); + assert_eq!( + oracle_bool(&OracleValue::String("\u{1}\u{1}".to_owned())), + Some(true) + ); + assert_eq!( + oracle_bool(&OracleValue::String("\u{1}\0".to_owned())), + Some(false) + ); + assert_eq!( + oracle_bool(&OracleValue::String("\u{1}\u{2}".to_owned())), + None + ); + } + + #[test] + fn oracle_timestamp_timezone_conversion_is_checked_and_crosses_days() { + let positive = OracleTimestamp::with_timezone(2026, 12, 31, 20, 30, 0, 123_456, 8, 0); + assert_eq!( + format_oracle_timestamp_tz(positive).expect("positive Oracle offset"), + "2027-01-01T04:30:00.123456+08:00" + ); + let negative = OracleTimestamp::with_timezone(2026, 1, 1, 1, 15, 0, 0, -5, -30); + assert_eq!( + format_oracle_timestamp_tz(negative).expect("negative Oracle offset"), + "2025-12-31T19:45:00-05:30" + ); + let invalid = OracleTimestamp::with_timezone(2026, 1, 1, 1, 15, 0, 0, 1, -30); + assert!(format_oracle_timestamp_tz(invalid).is_err()); + } + + #[test] + fn oracle_dialect_builds_schema_and_supported_namespace_sql() { + assert!(OracleNativeDriver.dialect().is_some()); + let schema = SchemaDefinition { + database_name: "FREEPDB1".to_owned(), + name: "APP".to_owned(), + comment: String::new(), + owner: "app".to_owned(), + system: false, + }; + assert_eq!( + build_oracle_create_schema(CreateSchemaSqlRequest { + schema: schema.clone(), + }) + .expect("Oracle schema SQL") + .sql, + "CREATE SCHEMA AUTHORIZATION \"APP\";" + ); + assert_eq!( + build_oracle_namespace_sql(NamespaceSqlRequest { + operation: NamespaceSqlOperation::CreateSchema { schema }, + }) + .expect("Oracle namespace schema SQL") + .sql, + "CREATE SCHEMA AUTHORIZATION \"APP\";" + ); + assert_eq!( + build_oracle_namespace_sql(NamespaceSqlRequest { + operation: NamespaceSqlOperation::DropSchema { + schema_name: "Odd\"User".to_owned(), + }, + }) + .expect("Oracle drop schema SQL") + .sql, + "DROP USER \"Odd\"\"User\" CASCADE;" + ); + assert_eq!( + build_oracle_namespace_sql(NamespaceSqlRequest { + operation: NamespaceSqlOperation::CreateDatabase { + database: DatabaseDefinition { + name: "APPDB".to_owned(), + comment: String::new(), + charset: "AL32UTF8".to_owned(), + collation: String::new(), + owner: String::new(), + system: false, + }, + }, + }) + .expect("Oracle create database SQL") + .sql, + "CREATE DATABASE \"APPDB\" CHARACTER SET AL32UTF8;" + ); + let error = build_oracle_namespace_sql(NamespaceSqlRequest { + operation: NamespaceSqlOperation::UseDatabase { + database_name: "APPDB".to_owned(), + }, + }) + .expect_err("Oracle database switching must fail explicitly"); + assert_eq!(error.api_error().code, "oracle_database_switch_unsupported"); + } + + #[test] + fn oracle_dialect_builds_typed_single_and_multi_insert_sql() { + let target = DmlTarget { + database_name: Some("FREEPDB1".to_owned()), + schema_name: Some("APP".to_owned()), + table_name: "ITEMS".to_owned(), + }; + let columns = vec![ + dml_column("ID", "NUMBER"), + dml_column("LABEL", "VARCHAR2"), + dml_column("ACTIVE", "BOOLEAN"), + dml_column("PAYLOAD", "BLOB"), + dml_column("CREATED_AT", "TIMESTAMP"), + ]; + let row = DmlRow { + values: vec![ + DmlValue::Decimal("1".to_owned()), + DmlValue::String("owner's".to_owned()), + DmlValue::Boolean(true), + DmlValue::Binary(vec![0, 255]), + DmlValue::Temporal { + kind: DmlTemporalKind::LocalDatetime, + iso8601: "2026-08-07T12:34:56.123456".to_owned(), + }, + ], + }; + let sql = build_oracle_dml(DmlSqlRequest { + target: target.clone(), + statement: DmlStatement::SingleInsert { + columns: columns.clone(), + row, + }, + }) + .expect("Oracle single INSERT SQL") + .sql; + assert_eq!( + sql, + "INSERT INTO \"APP\".\"ITEMS\" (\"ID\", \"LABEL\", \"ACTIVE\", \"PAYLOAD\", \"CREATED_AT\") VALUES (1, 'owner''s', TRUE, TO_BLOB(HEXTORAW('00FF')), TIMESTAMP '2026-08-07 12:34:56.123456');" + ); + + let sql = build_oracle_dml(DmlSqlRequest { + target, + statement: DmlStatement::MultiInsert { + columns: vec![dml_column("ID", "NUMBER")], + rows: vec![ + DmlRow { + values: vec![DmlValue::Decimal("1".to_owned())], + }, + DmlRow { + values: vec![DmlValue::Decimal("2".to_owned())], + }, + ], + }, + }) + .expect("Oracle multi INSERT SQL") + .sql; + assert_eq!( + sql, + "INSERT ALL\n INTO \"APP\".\"ITEMS\" (\"ID\") VALUES (1)\n INTO \"APP\".\"ITEMS\" (\"ID\") VALUES (2)\nSELECT 1 FROM DUAL;" + ); + } + + #[test] + fn oracle_dialect_builds_bounded_update_sql() { + let sql = build_oracle_dml(DmlSqlRequest { + target: DmlTarget { + database_name: None, + schema_name: Some("APP".to_owned()), + table_name: "ITEMS".to_owned(), + }, + statement: DmlStatement::Update { + assignments: vec![DmlAssignment { + column: dml_column("LABEL", "VARCHAR2"), + value: DmlValue::String("next".to_owned()), + }], + predicates: vec![ + DmlAssignment { + column: dml_column("ID", "NUMBER"), + value: DmlValue::Decimal("7".to_owned()), + }, + DmlAssignment { + column: dml_column("DELETED_AT", "TIMESTAMP"), + value: DmlValue::Null, + }, + ], + }, + }) + .expect("Oracle UPDATE SQL") + .sql; + assert_eq!( + sql, + "UPDATE \"APP\".\"ITEMS\" SET \"LABEL\" = 'next' WHERE \"ID\" = 7 AND \"DELETED_AT\" IS NULL;" + ); + } + + #[test] + fn oracle_identifiers_are_quoted_and_bounded() { + assert_eq!( + quote_identifier("Odd\"Name", "tableName").expect("identifier must quote"), + "\"Odd\"\"Name\"" + ); + for invalid in [String::new(), "bad\nname".to_owned(), "x".repeat(129)] { + assert!( + validate_metadata_identifier(&invalid, "tableName").is_err(), + "invalid identifier must be rejected" + ); + } + } + + fn property(key: &str, value: &str, sensitive: bool) -> DatasourceConnectionProperty { + DatasourceConnectionProperty { + key: key.to_owned(), + value: value.to_owned(), + sensitive, + } + } + + fn dml_column(name: &str, data_type_name: &str) -> DmlColumn { + DmlColumn { + name: name.to_owned(), + data_type_name: data_type_name.to_owned(), + precision: None, + scale: None, + } + } +} diff --git a/crates/chat2db-core/src/native_postgres.rs b/crates/chat2db-core/src/native_postgres.rs new file mode 100644 index 0000000..e2e8349 --- /dev/null +++ b/crates/chat2db-core/src/native_postgres.rs @@ -0,0 +1,5768 @@ +use std::{ + collections::{BTreeMap, HashMap}, + error::Error, + fmt::Write as _, + future::Future, + mem::size_of, + net::{IpAddr, Ipv4Addr, Ipv6Addr}, + str::FromStr, + time::{Duration, Instant}, +}; + +use async_trait::async_trait; +use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64_STANDARD}; +use chat2db_contract::{ + ApiError, ColumnNullability, DatasourceConnection, JdbcValue, JdbcValueType, QueryLimits, + ResultColumn, ResultMetadata, ResultRow, StartQueryRequest, +}; +use chat2db_engine_protocol::wire; +use chat2db_storage::Storage; +use chrono::{DateTime, NaiveDate, NaiveDateTime, NaiveTime, TimeDelta}; +use futures_util::StreamExt; +use prost::Message; +use sqlparser::{dialect::PostgreSqlDialect, parser::Parser}; +use tokio::{sync::watch, task::JoinHandle}; +use tokio_postgres::{ + Client, Config, Error as PostgresError, NoTls, Row, + config::SslMode, + types::{Format, FromSql, IsNull, Kind, ToSql, Type, private::BytesMut}, +}; +use tokio_postgres_rustls::MakeRustlsConnect; +use tokio_util::sync::CancellationToken; +use url::Url; + +use crate::{ + AppError, AppErrorKind, Application, + datasource_session::{ResolvedDatasourceConnection, resolve_datasource_connection}, + native_driver::{ + NativeConnectionDriver, NativeDialectDriver, NativeDriver, NativeMetadataDriver, + NativeQueryDriver, NativeTableDriver, + }, + native_driver_types::{ + BuiltSql, ColumnList, ColumnMetadata, CreateSchemaSqlRequest, DatabaseList, + DatabaseMetadata, DmlAssignment, DmlColumn, DmlRow, DmlSqlRequest, DmlStatement, DmlTarget, + DmlTemporalKind, DmlValue, EntityRelationColumn, EntityRelationForeignKey, + EntityRelationTable, ForeignKeyList, ForeignKeyMetadata, FunctionList, FunctionMetadata, + FunctionParameterList, FunctionParameterMetadata, IndexColumnMetadata, IndexList, + IndexMetadata, ListColumnsRequest, ListDatabasesRequest, ListIndexesRequest, + ListRoutinesRequest, ListSchemasRequest, ListTableKeysRequest, ListTablesRequest, + ListTriggersRequest, ListViewsRequest, MetadataObjectRef, NamespaceSqlOperation, + NamespaceSqlRequest, NativeDriverDescriptor, PrimaryKeyList, PrimaryKeyMetadata, + ProcedureList, ProcedureMetadata, ProcedureParameterList, ProcedureParameterMetadata, + SchemaList, SchemaMetadata, TableList, TableMetadata, TablePreviewAccepted, + TablePreviewRequest, TriggerList, TriggerMetadata, ViewList, + }, + operation::CancellationRequest, + query::{ + DatabaseValue, DatabaseWriteError, NativeConsoleRequest, NativeConsoleResult, + PreparedQuery, QueryExecutionOptions, QueryParameter, QueryTaskError, RetainedWriter, + }, + ssh::{SshTunnel, SshTunnelIdentity}, +}; + +const POSTGRES_SCHEME: &str = "postgresql://"; +const JDBC_POSTGRES_SCHEME: &str = "jdbc:postgresql://"; +const POSTGRES_DEFAULT_PORT: u16 = 5_432; +const CONNECT_TIMEOUT: Duration = Duration::from_secs(15); +const DISCONNECT_TIMEOUT: Duration = Duration::from_secs(5); +const METADATA_TIMEOUT: Duration = Duration::from_secs(30); +const CONSOLE_STATEMENT_TIMEOUT: Duration = Duration::from_secs(5 * 60); +const DEFAULT_BATCH_ROWS: u32 = 256; +const DEFAULT_BATCH_BYTES: u32 = 256 * 1024; +const DEFAULT_RESULT_BYTES: u64 = wire::JdbcResultByteLimit::DefaultResultBytes as u64; +const MAX_RESULT_BYTES: u64 = wire::JdbcResultByteLimit::MaxResultBytes as u64; +const MAX_BATCH_ROWS: u32 = wire::JdbcProtocolLimit::MaxBatchRows as u32; +const MAX_BATCH_BYTES: u32 = wire::JdbcProtocolLimit::MaxBatchBytes as u32; +const MAX_COLUMNS: usize = wire::JdbcProtocolLimit::MaxColumns as usize; +const MAX_PARAMETERS: usize = wire::JdbcProtocolLimit::MaxParameters as usize; +const MAX_SQL_BYTES: usize = wire::JdbcProtocolLimit::MaxSqlBytes as usize; +const MAX_SCALAR_BYTES: usize = wire::JdbcProtocolLimit::MaxScalarBytes as usize; +const MAX_IDENTIFIER_BYTES: usize = 63; +const MAX_CONSOLE_RESULT_BYTES: u64 = DEFAULT_RESULT_BYTES; +const MAX_CONSOLE_PAGE_SIZE: u32 = 10_000; +const MAX_CONSOLE_STATEMENTS: usize = 1_000; +const MAX_CONSOLE_SCANNED_ROWS: u64 = 1_000_000; +const MAX_CONSOLE_SCANNED_BYTES: u64 = MAX_RESULT_BYTES; +const POSTGRES_ARRAY_MAX_DIMENSIONS: usize = 6; + +pub(crate) const POSTGRES_DRIVER_DESCRIPTOR: NativeDriverDescriptor = NativeDriverDescriptor { + id: "postgresql", + implementation: "tokio-postgres", + database_types: &["POSTGRESQL", "POSTGRES"], + compatibility_aliases: &["postgresql", "postgres", "tokio-postgres"], +}; + +pub(crate) struct PostgresNativeDriver; + +impl NativeDriver for PostgresNativeDriver { + fn descriptor(&self) -> &'static NativeDriverDescriptor { + &POSTGRES_DRIVER_DESCRIPTOR + } + + fn connection(&self) -> Option<&dyn NativeConnectionDriver> { + Some(self) + } + + fn query(&self) -> Option<&dyn NativeQueryDriver> { + Some(self) + } + + fn metadata(&self) -> Option<&dyn NativeMetadataDriver> { + Some(self) + } + + fn tables(&self) -> Option<&dyn NativeTableDriver> { + Some(self) + } + + fn dialect(&self) -> Option<&dyn NativeDialectDriver> { + Some(self) + } +} + +#[async_trait] +impl NativeConnectionDriver for PostgresNativeDriver { + async fn test_connection(&self, connection: &DatasourceConnection) -> Result<(), AppError> { + self.test_connection_with_local_port(connection) + .await + .map(|_| ()) + } + + async fn test_connection_with_local_port( + &self, + connection: &DatasourceConnection, + ) -> Result, AppError> { + let connection = open_connection(connection).await?; + let local_port = connection.local_tunnel_port(); + let result = postgres_timeout( + CONNECT_TIMEOUT, + "postgres_connection_timeout", + "The PostgreSQL connection test timed out", + connection.client().simple_query("SELECT 1"), + ) + .await + .map(|_| ()); + finish_connection(connection, result).await?; + Ok(local_port) + } +} + +#[async_trait] +impl NativeQueryDriver for PostgresNativeDriver { + fn is_read_candidate(&self, sql: &str) -> Result { + is_native_read_candidate(sql) + } + + fn validate_query(&self, query: &PreparedQuery) -> Result<(), AppError> { + validate_query(query) + } + + async fn execute_query_task( + &self, + application: &Application, + operation_id: &str, + cancellation: watch::Receiver, + query: PreparedQuery, + storage: Storage, + resolved: ResolvedDatasourceConnection, + ) -> Result { + execute_query_task( + application, + operation_id, + cancellation, + query, + storage, + resolved, + ) + .await + } + + async fn execute_update( + &self, + resolved: ResolvedDatasourceConnection, + sql: String, + cancellation: CancellationToken, + ) -> Result { + execute_update(resolved, sql, cancellation).await + } + + async fn execute_console( + &self, + application: &Application, + request: NativeConsoleRequest, + cancellation: watch::Receiver, + force_read_only: bool, + ) -> Result, AppError> { + execute_console(application, request, cancellation, force_read_only).await + } +} + +#[async_trait] +impl NativeMetadataDriver for PostgresNativeDriver { + async fn list_schemas( + &self, + application: &Application, + request: ListSchemasRequest, + ) -> Result { + list_schemas(application, &request.datasource_id, &request.database_name).await + } + + async fn list_databases( + &self, + application: &Application, + request: ListDatabasesRequest, + ) -> Result { + list_databases(application, &request.datasource_id).await + } + + async fn list_tables( + &self, + application: &Application, + request: ListTablesRequest, + ) -> Result { + list_tables( + application, + &request.scope.datasource_id, + &request.scope.database_name, + &request.scope.schema_name, + &request.name_pattern, + ) + .await + } + + async fn list_columns( + &self, + application: &Application, + request: ListColumnsRequest, + ) -> Result { + list_columns( + application, + &request.table.scope.datasource_id, + &request.table.scope.database_name, + &request.table.scope.schema_name, + &request.table.table_name, + ) + .await + } + + async fn list_indexes( + &self, + application: &Application, + request: ListIndexesRequest, + ) -> Result { + list_indexes( + application, + &request.table.scope.datasource_id, + &request.table.scope.database_name, + &request.table.scope.schema_name, + &request.table.table_name, + ) + .await + } + + async fn list_views( + &self, + application: &Application, + request: ListViewsRequest, + ) -> Result { + list_views( + application, + &request.scope.datasource_id, + &request.scope.database_name, + &request.scope.schema_name, + &request.name_pattern, + ) + .await + } + + async fn get_view( + &self, + application: &Application, + request: MetadataObjectRef, + ) -> Result { + get_view( + application, + &request.scope.datasource_id, + &request.scope.database_name, + &request.scope.schema_name, + &request.object_name, + ) + .await + } + + async fn list_imported_keys( + &self, + application: &Application, + request: ListTableKeysRequest, + ) -> Result { + list_foreign_keys(application, &request, ForeignKeyDirection::Imported).await + } + + async fn list_exported_keys( + &self, + application: &Application, + request: ListTableKeysRequest, + ) -> Result { + list_foreign_keys(application, &request, ForeignKeyDirection::Exported).await + } + + async fn list_primary_keys( + &self, + application: &Application, + request: ListTableKeysRequest, + ) -> Result { + list_primary_keys(application, &request).await + } + + async fn list_functions( + &self, + application: &Application, + request: ListRoutinesRequest, + ) -> Result { + list_functions(application, &request).await + } + + async fn get_function( + &self, + application: &Application, + request: MetadataObjectRef, + ) -> Result { + get_function(application, &request).await + } + + async fn list_function_parameters( + &self, + application: &Application, + request: MetadataObjectRef, + ) -> Result { + list_function_parameters(application, &request).await + } + + async fn list_procedures( + &self, + application: &Application, + request: ListRoutinesRequest, + ) -> Result { + list_procedures(application, &request).await + } + + async fn get_procedure( + &self, + application: &Application, + request: MetadataObjectRef, + ) -> Result { + get_procedure(application, &request).await + } + + async fn list_procedure_parameters( + &self, + application: &Application, + request: MetadataObjectRef, + ) -> Result { + list_procedure_parameters(application, &request).await + } + + async fn list_triggers( + &self, + application: &Application, + request: ListTriggersRequest, + ) -> Result { + list_triggers(application, &request).await + } + + async fn get_trigger( + &self, + application: &Application, + request: MetadataObjectRef, + ) -> Result { + get_trigger(application, &request).await + } +} + +#[async_trait] +impl NativeTableDriver for PostgresNativeDriver { + async fn load_er_tables( + &self, + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + ) -> Result, AppError> { + load_er_tables(application, datasource_id, database_name, schema_name).await + } + + async fn validate_column_reorder( + &self, + _application: &Application, + _datasource_id: &str, + _database_name: &str, + _table_name: &str, + _column_names: &[String], + ) -> Result<(), AppError> { + Err(postgres_capability_not_supported( + "physical column reordering", + )) + } + + async fn table_ddl( + &self, + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + table_name: &str, + ) -> Result { + table_ddl( + application, + datasource_id, + database_name, + schema_name, + table_name, + ) + .await + } + + async fn start_table_preview( + &self, + application: &Application, + request: TablePreviewRequest, + row_limit: u32, + ) -> Result { + start_table_preview(application, request, row_limit).await + } +} + +impl NativeDialectDriver for PostgresNativeDriver { + fn build_create_schema(&self, request: CreateSchemaSqlRequest) -> Result { + build_create_schema(request) + } + + fn build_namespace_sql(&self, request: NamespaceSqlRequest) -> Result { + build_namespace_sql(request) + } + + fn build_dml(&self, request: DmlSqlRequest) -> Result { + build_dml(request) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum PostgresTlsMode { + Disable, + Prefer, + Require, +} + +struct PreparedPostgresConnection { + config: Config, + tls_mode: PostgresTlsMode, + tunnel: Option, +} + +struct ManagedPostgresConnection { + client: Option, + task: Option>>, + tunnel: Option, +} + +impl ManagedPostgresConnection { + fn new( + client: Client, + task: JoinHandle>, + tunnel: Option, + ) -> Self { + Self { + client: Some(client), + task: Some(task), + tunnel, + } + } + + fn client(&self) -> &Client { + self.client + .as_ref() + .expect("managed PostgreSQL client exists until cleanup") + } + + fn local_tunnel_port(&self) -> Option { + self.tunnel.as_ref().map(SshTunnel::local_port) + } + + async fn abort(mut self) { + self.client.take(); + if let Some(task) = self.task.take() { + task.abort(); + } + if let Some(tunnel) = self.tunnel.take() + && let Err(error) = tunnel.close().await + { + tracing::warn!(error = %error, "SSH tunnel cleanup failed after PostgreSQL abort"); + } + } +} + +impl Drop for ManagedPostgresConnection { + fn drop(&mut self) { + self.client.take(); + if let Some(task) = self.task.take() { + task.abort(); + } + } +} + +async fn open_connection( + connection: &DatasourceConnection, +) -> Result { + open_prepared_connection( + prepare_connection_with_identity(connection, SshTunnelIdentity::Ephemeral, None).await?, + ) + .await +} + +async fn open_resolved_connection( + resolved: &ResolvedDatasourceConnection, + database_name: Option<&str>, +) -> Result { + open_prepared_connection( + prepare_connection_with_identity( + &resolved.connection, + SshTunnelIdentity::Datasource { + datasource_id: &resolved.datasource_id, + revision: resolved.datasource_revision, + }, + database_name, + ) + .await?, + ) + .await +} + +async fn prepare_connection_with_identity( + connection: &DatasourceConnection, + identity: SshTunnelIdentity<'_>, + database_name: Option<&str>, +) -> Result { + let (config, tls_mode, target_host, target_port) = + connection_config(connection, database_name, None)?; + let Some(ssh) = connection.ssh.as_ref() else { + return Ok(PreparedPostgresConnection { + config, + tls_mode, + tunnel: None, + }); + }; + + let tunnel = SshTunnel::open(identity, ssh, target_host, target_port).await?; + let (mut config, tls_mode, _, _) = + connection_config(connection, database_name, Some(tunnel.local_port()))?; + config.hostaddr(IpAddr::V4(Ipv4Addr::LOCALHOST)); + Ok(PreparedPostgresConnection { + config, + tls_mode, + tunnel: Some(tunnel), + }) +} + +async fn open_prepared_connection( + mut prepared: PreparedPostgresConnection, +) -> Result { + let connect = async { + match prepared.tls_mode { + PostgresTlsMode::Disable => { + let (client, connection) = prepared + .config + .connect(NoTls) + .await + .map_err(postgres_connection_error)?; + Ok(ManagedPostgresConnection::new( + client, + tokio::spawn(connection), + prepared.tunnel.take(), + )) + } + PostgresTlsMode::Prefer | PostgresTlsMode::Require => { + ensure_postgres_rustls_provider()?; + let (connector, errors) = MakeRustlsConnect::with_native_certs().map_err(|_| { + AppError::unavailable( + "postgres_tls_roots_unavailable", + "No trusted system certificate roots are available for PostgreSQL TLS", + ) + })?; + if !errors.is_empty() { + tracing::warn!( + error_count = errors.len(), + "some native PostgreSQL TLS certificates could not be loaded" + ); + } + let (client, connection) = prepared + .config + .connect(connector) + .await + .map_err(postgres_connection_error)?; + Ok(ManagedPostgresConnection::new( + client, + tokio::spawn(connection), + prepared.tunnel.take(), + )) + } + } + }; + + match tokio::time::timeout(CONNECT_TIMEOUT, connect).await { + Ok(Ok(connection)) => Ok(connection), + Ok(Err(error)) => { + close_tunnel_quietly(prepared.tunnel.take()).await; + Err(error) + } + Err(_) => { + close_tunnel_quietly(prepared.tunnel.take()).await; + Err(AppError::unavailable( + "postgres_connection_timeout", + "The PostgreSQL connection attempt timed out", + )) + } + } +} + +fn ensure_postgres_rustls_provider() -> Result<(), AppError> { + if rustls::crypto::CryptoProvider::get_default().is_some() { + return Ok(()); + } + let _ = rustls::crypto::ring::default_provider().install_default(); + if rustls::crypto::CryptoProvider::get_default().is_some() { + Ok(()) + } else { + Err(AppError::unavailable( + "postgres_tls_provider_unavailable", + "A cryptographic provider could not be initialized for PostgreSQL TLS", + )) + } +} + +async fn finish_connection( + mut connection: ManagedPostgresConnection, + result: Result, +) -> Result { + connection.client.take(); + let close_result = match connection.task.take() { + Some(mut task) => match tokio::time::timeout(DISCONNECT_TIMEOUT, &mut task).await { + Ok(Ok(Ok(()))) => Ok(()), + Ok(Ok(Err(error))) => Err(postgres_connection_error(error)), + Ok(Err(_)) => Err(AppError::internal()), + Err(_) => { + task.abort(); + Err(AppError::unavailable( + "postgres_disconnect_timeout", + "The PostgreSQL connection did not close in time", + )) + } + }, + None => Ok(()), + }; + let tunnel_result = match connection.tunnel.take() { + Some(tunnel) => tunnel.close().await, + None => Ok(()), + }; + let cleanup_result = close_result.and(tunnel_result); + match result { + Ok(value) => cleanup_result.map(|()| value), + Err(primary) => { + if let Err(cleanup_error) = cleanup_result { + tracing::warn!(error = %cleanup_error, "PostgreSQL cleanup also failed"); + } + Err(primary) + } + } +} + +async fn close_tunnel_quietly(tunnel: Option) { + if let Some(tunnel) = tunnel + && let Err(error) = tunnel.close().await + { + tracing::warn!(error = %error, "SSH tunnel cleanup failed"); + } +} + +fn connection_config( + connection: &DatasourceConnection, + database_name: Option<&str>, + connect_port: Option, +) -> Result<(Config, PostgresTlsMode, String, u16), AppError> { + let normalized = normalize_postgres_url(&connection.jdbc_url)?; + let parsed = Url::parse(&normalized).map_err(|_| invalid_connection_url())?; + if parsed.scheme() != "postgresql" || parsed.host_str().is_none() || parsed.fragment().is_some() + { + return Err(invalid_connection_url()); + } + let target_host = parsed + .host_str() + .ok_or_else(invalid_connection_url)? + .to_owned(); + let target_port = parsed.port().unwrap_or(POSTGRES_DEFAULT_PORT); + let query_properties = parsed + .query_pairs() + .map(|(key, value)| (key.into_owned(), value.into_owned())) + .collect::>(); + let mut base = parsed; + base.set_query(None); + if let Some(connect_port) = connect_port { + base.set_port(Some(connect_port)) + .map_err(|()| invalid_connection_url())?; + } + let mut config = Config::from_str(base.as_str()).map_err(|_| invalid_connection_url())?; + config.connect_timeout(CONNECT_TIMEOUT); + config.application_name("Chat2DB-Rust"); + + let mut tls_mode = PostgresTlsMode::Prefer; + for (key, value) in query_properties + .iter() + .map(|(key, value)| (key.as_str(), value.as_str())) + .chain( + connection + .properties + .iter() + .map(|property| (property.key.as_str(), property.value.as_str())), + ) + { + apply_connection_property(&mut config, &mut tls_mode, key, value)?; + } + if let Some(database_name) = database_name { + validate_identifier(database_name, "databaseName")?; + config.dbname(database_name); + } + config.ssl_mode(match tls_mode { + PostgresTlsMode::Disable => SslMode::Disable, + PostgresTlsMode::Prefer => SslMode::Prefer, + PostgresTlsMode::Require => SslMode::Require, + }); + Ok((config, tls_mode, target_host, target_port)) +} + +fn apply_connection_property( + config: &mut Config, + tls_mode: &mut PostgresTlsMode, + key: &str, + value: &str, +) -> Result<(), AppError> { + match key.trim().to_ascii_lowercase().as_str() { + "user" | "username" => { + config.user(value); + } + "password" => { + config.password(value.as_bytes()); + } + "database" | "databasename" => { + config.dbname(value); + } + "applicationname" | "application_name" => { + config.application_name(value); + } + "options" => { + config.options(value); + } + "connecttimeout" | "connect_timeout" => { + let seconds = value + .trim() + .parse::() + .ok() + .filter(|seconds| *seconds > 0) + .ok_or_else(|| invalid_connection_property("connectTimeout"))?; + config.connect_timeout(Duration::from_secs(seconds.min(300))); + } + "ssl" => { + *tls_mode = if parse_bool(value) { + PostgresTlsMode::Require + } else { + PostgresTlsMode::Disable + }; + } + "sslmode" => { + *tls_mode = match value.trim().to_ascii_lowercase().as_str() { + "disable" | "disabled" | "false" => PostgresTlsMode::Disable, + "allow" | "prefer" => PostgresTlsMode::Prefer, + "require" | "verify-ca" | "verify-full" | "true" => PostgresTlsMode::Require, + _ => return Err(invalid_connection_property("sslMode")), + }; + } + "sslrootcert" | "sslcert" | "sslkey" if !value.trim().is_empty() => { + return Err(AppError::invalid( + "postgres_tls_property_not_supported", + "Custom PostgreSQL TLS certificate files are not supported; use the system trust store", + )); + } + "currentschema" | "current_schema" => { + validate_identifier(value, "currentSchema")?; + config.options(format!("-c search_path={}", quote_config_value(value)?)); + } + _ => { + return Err(AppError::invalid( + "postgres_connection_property_not_supported", + format!("The PostgreSQL connection property {key} is not supported"), + )); + } + } + Ok(()) +} + +fn quote_config_value(value: &str) -> Result { + if value + .bytes() + .any(|byte| byte.is_ascii_whitespace() || byte == b'\\') + { + return Err(invalid_connection_property("currentSchema")); + } + Ok(value.to_owned()) +} + +fn normalize_postgres_url(value: &str) -> Result { + let value = value.trim(); + if value + .get(..JDBC_POSTGRES_SCHEME.len()) + .is_some_and(|prefix| prefix.eq_ignore_ascii_case(JDBC_POSTGRES_SCHEME)) + { + return Ok(format!( + "{POSTGRES_SCHEME}{}", + &value[JDBC_POSTGRES_SCHEME.len()..] + )); + } + for scheme in [POSTGRES_SCHEME, "postgres://"] { + if value + .get(..scheme.len()) + .is_some_and(|prefix| prefix.eq_ignore_ascii_case(scheme)) + { + return Ok(format!("{POSTGRES_SCHEME}{}", &value[scheme.len()..])); + } + } + Err(invalid_connection_url()) +} + +fn parse_bool(value: &str) -> bool { + matches!( + value.trim().to_ascii_lowercase().as_str(), + "1" | "true" | "yes" | "on" | "required" + ) +} + +fn invalid_connection_url() -> AppError { + AppError::invalid( + "invalid_postgres_connection", + "A valid jdbc:postgresql://, postgresql://, or postgres:// connection URL is required", + ) +} + +fn invalid_connection_property(property: &str) -> AppError { + AppError::invalid( + "invalid_postgres_connection", + format!("The PostgreSQL connection property {property} is invalid"), + ) +} + +#[allow( + clippy::needless_pass_by_value, + reason = "the owned signature lets this mapper be passed directly to map_err" +)] +fn postgres_connection_error(error: PostgresError) -> AppError { + if let Some(database) = error.as_db_error() { + return AppError::new( + AppErrorKind::InvalidRequest, + ApiError::new("postgres_connection_rejected", database.message()), + ); + } + AppError::unavailable( + "postgres_connection_failed", + "The PostgreSQL server could not be reached", + ) +} + +#[allow( + clippy::needless_pass_by_value, + reason = "the owned signature lets this mapper be passed directly to map_err" +)] +fn postgres_query_error(error: PostgresError) -> AppError { + if let Some(database) = error.as_db_error() { + return AppError::new( + AppErrorKind::InvalidRequest, + ApiError::new("postgres_query_rejected", database.message()), + ); + } + AppError::unavailable( + "postgres_query_failed", + "The PostgreSQL query could not be completed", + ) +} + +async fn postgres_timeout( + duration: Duration, + code: &'static str, + message: &'static str, + future: F, +) -> Result +where + F: Future>, +{ + tokio::time::timeout(duration, future) + .await + .map_err(|_| AppError::unavailable(code, message))? + .map_err(postgres_query_error) +} + +async fn resolve_native_connection( + application: &Application, + datasource_id: &str, +) -> Result { + let storage = application.require_storage()?; + let resolved = resolve_datasource_connection(&storage, datasource_id).await?; + if application + .native_driver_for_datasource_driver_id(&resolved.driver_id) + .is_none_or(|driver| driver.descriptor().id != POSTGRES_DRIVER_DESCRIPTOR.id) + { + return Err(AppError::invalid( + "postgres_driver_mismatch", + "The datasource is not configured with a PostgreSQL driver", + )); + } + Ok(resolved) +} + +fn postgres_capability_not_supported(capability: &'static str) -> AppError { + AppError::invalid( + "native_driver_capability_not_supported", + format!("The PostgreSQL driver does not implement {capability}"), + ) +} + +async fn metadata_rows( + application: &Application, + datasource_id: &str, + database_name: Option<&str>, + sql: &str, + parameters: &[&(dyn ToSql + Sync)], +) -> Result, AppError> { + let resolved = resolve_native_connection(application, datasource_id).await?; + let connection = open_resolved_connection(&resolved, database_name).await?; + let result = postgres_timeout( + METADATA_TIMEOUT, + "postgres_metadata_timeout", + "The PostgreSQL metadata query did not finish in time", + connection.client().query(sql, parameters), + ) + .await; + finish_connection(connection, result).await +} + +async fn list_databases( + application: &Application, + datasource_id: &str, +) -> Result { + let rows = metadata_rows( + application, + datasource_id, + None, + "SELECT d.datname, pg_encoding_to_char(d.encoding), d.datcollate, \ + pg_get_userbyid(d.datdba), COALESCE(obj_description(d.oid, 'pg_database'), ''), \ + d.datistemplate OR NOT d.datallowconn \ + FROM pg_database d ORDER BY d.datname", + &[], + ) + .await?; + Ok(DatabaseList { + items: rows + .into_iter() + .map(|row| { + Ok(DatabaseMetadata { + name: row.try_get(0).map_err(postgres_query_error)?, + charset: row.try_get(1).map_err(postgres_query_error)?, + collation: row.try_get(2).map_err(postgres_query_error)?, + owner: row.try_get(3).map_err(postgres_query_error)?, + comment: row.try_get(4).map_err(postgres_query_error)?, + system: row.try_get(5).map_err(postgres_query_error)?, + }) + }) + .collect::>()?, + }) +} + +async fn list_schemas( + application: &Application, + datasource_id: &str, + database_name: &str, +) -> Result { + validate_identifier(database_name, "databaseName")?; + let rows = metadata_rows( + application, + datasource_id, + Some(database_name), + "SELECT current_database(), n.nspname, \ + COALESCE(obj_description(n.oid, 'pg_namespace'), ''), \ + pg_get_userbyid(n.nspowner), \ + n.nspname LIKE 'pg\\_%' ESCAPE '\\' OR n.nspname = 'information_schema' \ + FROM pg_namespace n ORDER BY n.nspname", + &[], + ) + .await?; + Ok(SchemaList { + items: rows + .into_iter() + .map(|row| { + Ok(SchemaMetadata { + database_name: row.try_get(0).map_err(postgres_query_error)?, + name: row.try_get(1).map_err(postgres_query_error)?, + comment: row.try_get(2).map_err(postgres_query_error)?, + owner: row.try_get(3).map_err(postgres_query_error)?, + system: row.try_get(4).map_err(postgres_query_error)?, + }) + }) + .collect::>()?, + }) +} + +async fn list_tables( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + name_pattern: &str, +) -> Result { + validate_metadata_scope(database_name, schema_name)?; + let pattern = name_pattern.trim().to_owned(); + let rows = metadata_rows( + application, + datasource_id, + Some(database_name), + "SELECT current_database(), n.nspname, c.relname, c.relkind::text, \ + COALESCE(obj_description(c.oid, 'pg_class'), ''), \ + COALESCE(ts.spcname, ''), c.reltuples::bigint::text, \ + pg_total_relation_size(c.oid)::bigint::text, c.relpersistence::text \ + FROM pg_class c \ + JOIN pg_namespace n ON n.oid = c.relnamespace \ + LEFT JOIN pg_tablespace ts ON ts.oid = c.reltablespace \ + WHERE n.nspname = $1 AND c.relkind IN ('r', 'p', 'f') \ + AND ($2 = '' OR c.relname LIKE $2 ESCAPE '\\') \ + ORDER BY c.relname", + &[&schema_name, &pattern], + ) + .await?; + Ok(TableList { + items: rows + .iter() + .map(postgres_table_metadata) + .collect::>()?, + }) +} + +fn postgres_table_metadata(row: &Row) -> Result { + let relation_kind: String = row.try_get(3).map_err(postgres_query_error)?; + let persistence: String = row.try_get(8).map_err(postgres_query_error)?; + let engine = match (relation_kind.as_str(), persistence.as_str()) { + ("p", _) => "PARTITIONED", + ("f", _) => "FOREIGN", + (_, "u") => "UNLOGGED", + (_, "t") => "TEMPORARY", + _ => "HEAP", + }; + Ok(TableMetadata { + database_name: row.try_get(0).map_err(postgres_query_error)?, + schema_name: row.try_get(1).map_err(postgres_query_error)?, + name: row.try_get(2).map_err(postgres_query_error)?, + table_type: "TABLE".to_owned(), + comment: row.try_get(4).map_err(postgres_query_error)?, + database_type: "POSTGRESQL".to_owned(), + engine: engine.to_owned(), + tablespace: row.try_get(5).map_err(postgres_query_error)?, + rows: row.try_get(6).map_err(postgres_query_error)?, + data_length: row.try_get(7).map_err(postgres_query_error)?, + ..TableMetadata::default() + }) +} + +async fn list_columns( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + table_name: &str, +) -> Result { + validate_metadata_table(database_name, schema_name, table_name)?; + let rows = metadata_rows( + application, + datasource_id, + Some(database_name), + "SELECT current_database(), n.nspname, c.relname, a.attname, \ + format_type(a.atttypid, a.atttypmod), t.typname, \ + pg_get_expr(ad.adbin, ad.adrelid), a.attnotnull, \ + a.attidentity::text, a.attgenerated::text, \ + COALESCE(col_description(c.oid, a.attnum), ''), a.attnum::int4, \ + information_schema._pg_numeric_precision(a.atttypid, a.atttypmod)::int4, \ + information_schema._pg_numeric_scale(a.atttypid, a.atttypmod)::int4, \ + information_schema._pg_char_max_length(a.atttypid, a.atttypmod)::int4, \ + COALESCE(coll.collname, ''), COALESCE(pk.conname, ''), \ + COALESCE(array_position(pk.conkey, a.attnum), 0)::int4 \ + FROM pg_class c \ + JOIN pg_namespace n ON n.oid = c.relnamespace \ + JOIN pg_attribute a ON a.attrelid = c.oid AND a.attnum > 0 AND NOT a.attisdropped \ + JOIN pg_type t ON t.oid = a.atttypid \ + LEFT JOIN pg_attrdef ad ON ad.adrelid = c.oid AND ad.adnum = a.attnum \ + LEFT JOIN pg_collation coll ON coll.oid = a.attcollation AND a.attcollation <> 0 \ + LEFT JOIN pg_constraint pk ON pk.conrelid = c.oid AND pk.contype = 'p' \ + AND a.attnum = ANY(pk.conkey) \ + WHERE n.nspname = $1 AND c.relname = $2 \ + ORDER BY a.attnum", + &[&schema_name, &table_name], + ) + .await?; + Ok(ColumnList { + items: rows + .iter() + .map(postgres_column_metadata) + .collect::>()?, + }) +} + +fn postgres_column_metadata(row: &Row) -> Result { + let column_type: String = row.try_get(4).map_err(postgres_query_error)?; + let type_name: String = row.try_get(5).map_err(postgres_query_error)?; + let default_value: Option = row.try_get(6).map_err(postgres_query_error)?; + let not_null: bool = row.try_get(7).map_err(postgres_query_error)?; + let identity: String = row.try_get(8).map_err(postgres_query_error)?; + let generated: String = row.try_get(9).map_err(postgres_query_error)?; + let primary_key_name: String = row.try_get(16).map_err(postgres_query_error)?; + let primary_key_order: i32 = row.try_get(17).map_err(postgres_query_error)?; + Ok(ColumnMetadata { + database_name: row.try_get(0).map_err(postgres_query_error)?, + schema_name: row.try_get(1).map_err(postgres_query_error)?, + table_name: row.try_get(2).map_err(postgres_query_error)?, + name: row.try_get(3).map_err(postgres_query_error)?, + column_type, + data_type: Some(postgres_jdbc_type_name(&type_name)), + default_value, + auto_increment: Some(!identity.is_empty()), + comment: row.try_get(10).map_err(postgres_query_error)?, + primary_key: Some(primary_key_order > 0), + primary_key_name, + primary_key_order, + column_size: row.try_get(14).map_err(postgres_query_error)?, + decimal_digits: row.try_get(13).map_err(postgres_query_error)?, + num_prec_radix: row + .try_get::<_, Option>(12) + .map_err(postgres_query_error)? + .map(|_| 10), + ordinal_position: Some(row.try_get(11).map_err(postgres_query_error)?), + nullable: Some(i32::from(!not_null)), + generated_column: Some(!generated.is_empty()), + extent: if !identity.is_empty() { + format!( + "GENERATED {} AS IDENTITY", + if identity == "a" { + "ALWAYS" + } else { + "BY DEFAULT" + } + ) + } else if !generated.is_empty() { + "GENERATED ALWAYS".to_owned() + } else { + String::new() + }, + collation: row.try_get(15).map_err(postgres_query_error)?, + ..ColumnMetadata::default() + }) +} + +async fn list_indexes( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + table_name: &str, +) -> Result { + validate_metadata_table(database_name, schema_name, table_name)?; + let rows = metadata_rows( + application, + datasource_id, + Some(database_name), + "SELECT idx.relname, i.indisunique, am.amname, \ + COALESCE(obj_description(idx.oid, 'pg_class'), ''), \ + key.ordinality::int4, \ + COALESCE(att.attname, pg_get_indexdef(i.indexrelid, key.ordinality, true)), \ + COALESCE(pg_get_expr(i.indpred, i.indrelid), ''), \ + CASE WHEN ((i.indoption::smallint[])[key.ordinality - 1] & 1) = 1 \ + THEN 'D' ELSE 'A' END \ + FROM pg_index i \ + JOIN pg_class tbl ON tbl.oid = i.indrelid \ + JOIN pg_namespace n ON n.oid = tbl.relnamespace \ + JOIN pg_class idx ON idx.oid = i.indexrelid \ + JOIN pg_am am ON am.oid = idx.relam \ + CROSS JOIN LATERAL generate_series(1, i.indnatts) key(ordinality) \ + LEFT JOIN pg_attribute att ON att.attrelid = tbl.oid \ + AND att.attnum = (i.indkey::smallint[])[key.ordinality - 1] \ + WHERE n.nspname = $1 AND tbl.relname = $2 \ + ORDER BY idx.relname, key.ordinality", + &[&schema_name, &table_name], + ) + .await?; + + let mut indexes = BTreeMap::::new(); + for row in rows { + let name: String = row.try_get(0).map_err(postgres_query_error)?; + let unique: bool = row.try_get(1).map_err(postgres_query_error)?; + let method: String = row.try_get(2).map_err(postgres_query_error)?; + let comment: String = row.try_get(3).map_err(postgres_query_error)?; + let ordinal_position: i32 = row.try_get(4).map_err(postgres_query_error)?; + let column_name: String = row.try_get(5).map_err(postgres_query_error)?; + let filter_condition: String = row.try_get(6).map_err(postgres_query_error)?; + let sort_order: String = row.try_get(7).map_err(postgres_query_error)?; + let index = indexes + .entry(name.clone()) + .or_insert_with(|| IndexMetadata { + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + table_name: table_name.to_owned(), + name: name.clone(), + index_type: if name.ends_with("_pkey") { + "PRIMARY".to_owned() + } else if unique { + "UNIQUE".to_owned() + } else { + "INDEX".to_owned() + }, + unique: Some(unique), + comment, + method: method.clone(), + ..IndexMetadata::default() + }); + index.columns.push(IndexColumnMetadata { + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + table_name: table_name.to_owned(), + index_name: name, + column_name, + ordinal_position: Some(ordinal_position), + non_unique: Some(!unique), + sort_order, + filter_condition, + ..IndexColumnMetadata::default() + }); + } + Ok(IndexList { + items: indexes.into_values().collect(), + }) +} + +async fn list_views( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + name_pattern: &str, +) -> Result { + validate_metadata_scope(database_name, schema_name)?; + let pattern = name_pattern.trim().to_owned(); + let rows = metadata_rows( + application, + datasource_id, + Some(database_name), + "SELECT current_database(), n.nspname, c.relname, c.relkind::text, \ + COALESCE(obj_description(c.oid, 'pg_class'), ''), \ + pg_get_viewdef(c.oid, true) \ + FROM pg_class c JOIN pg_namespace n ON n.oid = c.relnamespace \ + WHERE n.nspname = $1 AND c.relkind IN ('v', 'm') \ + AND ($2 = '' OR c.relname LIKE $2 ESCAPE '\\') \ + ORDER BY c.relname", + &[&schema_name, &pattern], + ) + .await?; + Ok(ViewList { + items: rows + .iter() + .map(postgres_view_metadata) + .collect::>()?, + }) +} + +async fn get_view( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + view_name: &str, +) -> Result { + validate_metadata_table(database_name, schema_name, view_name)?; + let mut rows = metadata_rows( + application, + datasource_id, + Some(database_name), + "SELECT current_database(), n.nspname, c.relname, c.relkind::text, \ + COALESCE(obj_description(c.oid, 'pg_class'), ''), \ + pg_get_viewdef(c.oid, true) \ + FROM pg_class c JOIN pg_namespace n ON n.oid = c.relnamespace \ + WHERE n.nspname = $1 AND c.relname = $2 AND c.relkind IN ('v', 'm')", + &[&schema_name, &view_name], + ) + .await?; + let row = rows + .pop() + .ok_or_else(|| metadata_not_found("view", database_name, schema_name, view_name))?; + postgres_view_metadata(&row) +} + +fn postgres_view_metadata(row: &Row) -> Result { + let database_name: String = row.try_get(0).map_err(postgres_query_error)?; + let schema_name: String = row.try_get(1).map_err(postgres_query_error)?; + let name: String = row.try_get(2).map_err(postgres_query_error)?; + let kind: String = row.try_get(3).map_err(postgres_query_error)?; + let definition: String = row.try_get(5).map_err(postgres_query_error)?; + let materialized = kind == "m"; + let ddl = format!( + "CREATE {}VIEW {}.{} AS\n{};", + if materialized { "MATERIALIZED " } else { "" }, + quote_identifier(&schema_name, "schemaName")?, + quote_identifier(&name, "viewName")?, + definition.trim_end_matches(';') + ); + Ok(TableMetadata { + database_name, + schema_name, + name, + table_type: if materialized { + "MATERIALIZED VIEW".to_owned() + } else { + "VIEW".to_owned() + }, + comment: row.try_get(4).map_err(postgres_query_error)?, + database_type: "POSTGRESQL".to_owned(), + ddl, + ..TableMetadata::default() + }) +} + +#[derive(Clone, Copy)] +enum ForeignKeyDirection { + Imported, + Exported, +} + +async fn list_foreign_keys( + application: &Application, + request: &ListTableKeysRequest, + direction: ForeignKeyDirection, +) -> Result { + let scope = &request.table.scope; + validate_metadata_table( + &scope.database_name, + &scope.schema_name, + &request.table.table_name, + )?; + let filter = match direction { + ForeignKeyDirection::Imported => "fn.nspname = $1 AND ft.relname = $2", + ForeignKeyDirection::Exported => "pn.nspname = $1 AND pt.relname = $2", + }; + let sql = format!( + "SELECT current_database(), pn.nspname, pt.relname, pa.attname, \ + fn.nspname, ft.relname, fa.attname, pos.n::int4, \ + con.confupdtype::text, con.confdeltype::text, con.conname, \ + COALESCE(pkidx.relname, ''), con.condeferrable, con.condeferred \ + FROM pg_constraint con \ + JOIN pg_class ft ON ft.oid = con.conrelid \ + JOIN pg_namespace fn ON fn.oid = ft.relnamespace \ + JOIN pg_class pt ON pt.oid = con.confrelid \ + JOIN pg_namespace pn ON pn.oid = pt.relnamespace \ + JOIN LATERAL generate_subscripts(con.conkey, 1) pos(n) ON true \ + JOIN pg_attribute fa ON fa.attrelid = ft.oid AND fa.attnum = con.conkey[pos.n] \ + JOIN pg_attribute pa ON pa.attrelid = pt.oid AND pa.attnum = con.confkey[pos.n] \ + LEFT JOIN pg_class pkidx ON pkidx.oid = con.conindid \ + WHERE con.contype = 'f' AND {filter} \ + ORDER BY con.conname, pos.n" + ); + let rows = metadata_rows( + application, + &scope.datasource_id, + Some(&scope.database_name), + &sql, + &[&scope.schema_name, &request.table.table_name], + ) + .await?; + Ok(ForeignKeyList { + items: rows + .into_iter() + .map(|row| { + let deferrable: bool = row.try_get(12).map_err(postgres_query_error)?; + let deferred: bool = row.try_get(13).map_err(postgres_query_error)?; + Ok(ForeignKeyMetadata { + primary_table_database: row.try_get(0).map_err(postgres_query_error)?, + primary_table_schema: row.try_get(1).map_err(postgres_query_error)?, + primary_table_name: row.try_get(2).map_err(postgres_query_error)?, + primary_column_name: row.try_get(3).map_err(postgres_query_error)?, + foreign_table_database: row.try_get(0).map_err(postgres_query_error)?, + foreign_table_schema: row.try_get(4).map_err(postgres_query_error)?, + foreign_table_name: row.try_get(5).map_err(postgres_query_error)?, + foreign_column_name: row.try_get(6).map_err(postgres_query_error)?, + key_sequence: row.try_get(7).map_err(postgres_query_error)?, + update_rule: postgres_referential_rule( + &row.try_get::<_, String>(8).map_err(postgres_query_error)?, + ), + delete_rule: postgres_referential_rule( + &row.try_get::<_, String>(9).map_err(postgres_query_error)?, + ), + foreign_key_name: row.try_get(10).map_err(postgres_query_error)?, + primary_key_name: row.try_get(11).map_err(postgres_query_error)?, + deferrability: if !deferrable { + 7 + } else if deferred { + 5 + } else { + 6 + }, + }) + }) + .collect::>()?, + }) +} + +async fn list_primary_keys( + application: &Application, + request: &ListTableKeysRequest, +) -> Result { + let scope = &request.table.scope; + validate_metadata_table( + &scope.database_name, + &scope.schema_name, + &request.table.table_name, + )?; + let rows = metadata_rows( + application, + &scope.datasource_id, + Some(&scope.database_name), + "SELECT current_database(), n.nspname, c.relname, a.attname, con.conname \ + FROM pg_constraint con \ + JOIN pg_class c ON c.oid = con.conrelid \ + JOIN pg_namespace n ON n.oid = c.relnamespace \ + JOIN LATERAL unnest(con.conkey) WITH ORDINALITY key(attnum, ordinality) ON true \ + JOIN pg_attribute a ON a.attrelid = c.oid AND a.attnum = key.attnum \ + WHERE con.contype = 'p' AND n.nspname = $1 AND c.relname = $2 \ + ORDER BY key.ordinality", + &[&scope.schema_name, &request.table.table_name], + ) + .await?; + Ok(PrimaryKeyList { + items: rows + .into_iter() + .map(|row| { + Ok(PrimaryKeyMetadata { + database_name: row.try_get(0).map_err(postgres_query_error)?, + schema_name: row.try_get(1).map_err(postgres_query_error)?, + table_name: row.try_get(2).map_err(postgres_query_error)?, + column_name: row.try_get(3).map_err(postgres_query_error)?, + name: row.try_get(4).map_err(postgres_query_error)?, + }) + }) + .collect::>()?, + }) +} + +fn postgres_referential_rule(value: &str) -> i32 { + match value { + "c" => 0, + "r" => 1, + "n" => 2, + "d" => 4, + _ => 3, + } +} + +fn validate_metadata_scope(database_name: &str, schema_name: &str) -> Result<(), AppError> { + validate_identifier(database_name, "databaseName")?; + validate_identifier(schema_name, "schemaName") +} + +fn validate_metadata_table( + database_name: &str, + schema_name: &str, + table_name: &str, +) -> Result<(), AppError> { + validate_metadata_scope(database_name, schema_name)?; + validate_identifier(table_name, "tableName") +} + +fn validate_identifier(value: &str, field: &str) -> Result<(), AppError> { + if value.trim().is_empty() || value.len() > MAX_IDENTIFIER_BYTES || value.contains('\0') { + return Err(AppError::invalid( + "invalid_postgres_metadata_request", + format!("{field} is invalid"), + )); + } + Ok(()) +} + +fn quote_identifier(value: &str, field: &str) -> Result { + validate_identifier(value, field)?; + Ok(format!("\"{}\"", value.replace('"', "\"\""))) +} + +fn metadata_not_found( + kind: &str, + database_name: &str, + schema_name: &str, + object_name: &str, +) -> AppError { + AppError::not_found( + "postgres_metadata_not_found", + format!("PostgreSQL {kind} {database_name}.{schema_name}.{object_name} was not found"), + ) +} + +async fn list_functions( + application: &Application, + request: &ListRoutinesRequest, +) -> Result { + validate_metadata_scope(&request.scope.database_name, &request.scope.schema_name)?; + let rows = metadata_rows( + application, + &request.scope.datasource_id, + Some(&request.scope.database_name), + "SELECT current_database(), n.nspname, p.proname, \ + COALESCE(obj_description(p.oid, 'pg_proc'), ''), p.oid::text, \ + pg_get_function_identity_arguments(p.oid) \ + FROM pg_proc p JOIN pg_namespace n ON n.oid = p.pronamespace \ + WHERE n.nspname = $1 AND p.prokind = 'f' \ + ORDER BY p.proname, pg_get_function_identity_arguments(p.oid)", + &[&request.scope.schema_name], + ) + .await?; + Ok(FunctionList { + items: rows + .into_iter() + .map(|row| { + let name: String = row.try_get(2).map_err(postgres_query_error)?; + let oid: String = row.try_get(4).map_err(postgres_query_error)?; + let arguments: String = row.try_get(5).map_err(postgres_query_error)?; + Ok(FunctionMetadata { + database_name: row.try_get(0).map_err(postgres_query_error)?, + schema_name: row.try_get(1).map_err(postgres_query_error)?, + name: name.clone(), + remarks: row.try_get(3).map_err(postgres_query_error)?, + function_type: Some(1), + specific_name: format!("{name}_{oid}"), + template: format!("{name}({arguments})"), + ..FunctionMetadata::default() + }) + }) + .collect::>()?, + }) +} + +async fn get_function( + application: &Application, + request: &MetadataObjectRef, +) -> Result { + let routine = resolve_routine(application, request, "f", "function").await?; + Ok(FunctionMetadata { + database_name: request.scope.database_name.clone(), + schema_name: request.scope.schema_name.clone(), + name: routine.name.clone(), + remarks: routine.remarks, + function_type: Some(1), + specific_name: format!("{}_{}", routine.name, routine.oid), + body: routine.definition, + template: format!("{}({})", routine.name, routine.identity_arguments), + }) +} + +async fn list_function_parameters( + application: &Application, + request: &MetadataObjectRef, +) -> Result { + let routine = resolve_routine(application, request, "f", "function").await?; + let rows = routine_parameter_rows(application, request, routine.oid, true).await?; + Ok(FunctionParameterList { + items: rows + .into_iter() + .map(|row| { + let ordinal: i32 = row.try_get(0).map_err(postgres_query_error)?; + let mode: String = row.try_get(2).map_err(postgres_query_error)?; + let type_name: String = row.try_get(3).map_err(postgres_query_error)?; + Ok(FunctionParameterMetadata { + function_database: request.scope.database_name.clone(), + function_schema: request.scope.schema_name.clone(), + function_name: routine.name.clone(), + column_name: row.try_get(1).map_err(postgres_query_error)?, + column_type: Some(postgres_function_column_type(&mode, ordinal)), + data_type: Some(postgres_jdbc_type_name(&type_name)), + type_name, + ordinal_position: Some(ordinal), + nullable: Some(2), + is_nullable: String::new(), + specific_name: format!("{}_{}", routine.name, routine.oid), + ..FunctionParameterMetadata::default() + }) + }) + .collect::>()?, + }) +} + +async fn list_procedures( + application: &Application, + request: &ListRoutinesRequest, +) -> Result { + validate_metadata_scope(&request.scope.database_name, &request.scope.schema_name)?; + let rows = metadata_rows( + application, + &request.scope.datasource_id, + Some(&request.scope.database_name), + "SELECT current_database(), n.nspname, p.proname, \ + COALESCE(obj_description(p.oid, 'pg_proc'), ''), p.oid::text \ + FROM pg_proc p JOIN pg_namespace n ON n.oid = p.pronamespace \ + WHERE n.nspname = $1 AND p.prokind = 'p' \ + ORDER BY p.proname, pg_get_function_identity_arguments(p.oid)", + &[&request.scope.schema_name], + ) + .await?; + Ok(ProcedureList { + items: rows + .into_iter() + .map(|row| { + let name: String = row.try_get(2).map_err(postgres_query_error)?; + let oid: String = row.try_get(4).map_err(postgres_query_error)?; + Ok(ProcedureMetadata { + database_name: row.try_get(0).map_err(postgres_query_error)?, + schema_name: row.try_get(1).map_err(postgres_query_error)?, + name: name.clone(), + remarks: row.try_get(3).map_err(postgres_query_error)?, + procedure_type: Some(2), + specific_name: format!("{name}_{oid}"), + body: String::new(), + }) + }) + .collect::>()?, + }) +} + +async fn get_procedure( + application: &Application, + request: &MetadataObjectRef, +) -> Result { + let routine = resolve_routine(application, request, "p", "procedure").await?; + Ok(ProcedureMetadata { + database_name: request.scope.database_name.clone(), + schema_name: request.scope.schema_name.clone(), + name: routine.name.clone(), + remarks: routine.remarks, + procedure_type: Some(2), + specific_name: format!("{}_{}", routine.name, routine.oid), + body: routine.definition, + }) +} + +async fn list_procedure_parameters( + application: &Application, + request: &MetadataObjectRef, +) -> Result { + let routine = resolve_routine(application, request, "p", "procedure").await?; + let rows = routine_parameter_rows(application, request, routine.oid, false).await?; + Ok(ProcedureParameterList { + items: rows + .into_iter() + .map(|row| { + let ordinal: i32 = row.try_get(0).map_err(postgres_query_error)?; + let mode: String = row.try_get(2).map_err(postgres_query_error)?; + let type_name: String = row.try_get(3).map_err(postgres_query_error)?; + Ok(ProcedureParameterMetadata { + procedure_database: request.scope.database_name.clone(), + procedure_schema: request.scope.schema_name.clone(), + procedure_name: routine.name.clone(), + column_name: row.try_get(1).map_err(postgres_query_error)?, + column_type: Some(postgres_procedure_column_type(&mode)), + data_type: Some(postgres_jdbc_type_name(&type_name)), + type_name, + ordinal_position: Some(ordinal), + nullable: Some(2), + specific_name: format!("{}_{}", routine.name, routine.oid), + ..ProcedureParameterMetadata::default() + }) + }) + .collect::>()?, + }) +} + +struct ResolvedRoutine { + oid: u32, + name: String, + remarks: String, + definition: String, + identity_arguments: String, +} + +async fn resolve_routine( + application: &Application, + request: &MetadataObjectRef, + kind: &str, + label: &str, +) -> Result { + validate_metadata_scope(&request.scope.database_name, &request.scope.schema_name)?; + validate_identifier(&request.object_name, "routineName")?; + let rows = metadata_rows( + application, + &request.scope.datasource_id, + Some(&request.scope.database_name), + "SELECT p.oid, p.proname, COALESCE(obj_description(p.oid, 'pg_proc'), ''), \ + pg_get_functiondef(p.oid), pg_get_function_identity_arguments(p.oid) \ + FROM pg_proc p JOIN pg_namespace n ON n.oid = p.pronamespace \ + WHERE n.nspname = $1 AND p.prokind::text = $2 \ + AND (p.proname = $3 OR p.proname || '_' || p.oid::text = $3) \ + ORDER BY p.oid LIMIT 2", + &[&request.scope.schema_name, &kind, &request.object_name], + ) + .await?; + if rows.is_empty() { + return Err(metadata_not_found( + label, + &request.scope.database_name, + &request.scope.schema_name, + &request.object_name, + )); + } + if rows.len() > 1 { + return Err(AppError::invalid( + "postgres_routine_ambiguous", + format!( + "PostgreSQL {label} {} is overloaded; use its specificName", + request.object_name + ), + )); + } + let row = &rows[0]; + Ok(ResolvedRoutine { + oid: row.try_get(0).map_err(postgres_query_error)?, + name: row.try_get(1).map_err(postgres_query_error)?, + remarks: row.try_get(2).map_err(postgres_query_error)?, + definition: row.try_get(3).map_err(postgres_query_error)?, + identity_arguments: row.try_get(4).map_err(postgres_query_error)?, + }) +} + +async fn routine_parameter_rows( + application: &Application, + request: &MetadataObjectRef, + oid: u32, + include_return: bool, +) -> Result, AppError> { + let return_filter = if include_return { "" } else { "WHERE false" }; + let sql = format!( + "WITH target AS (SELECT * FROM pg_proc WHERE oid = $1), \ + arguments AS ( \ + SELECT ordinality::int4 AS ordinal, \ + COALESCE(p.proargnames[ordinality], '') AS name, \ + COALESCE(p.proargmodes[ordinality], 'i')::text AS mode, \ + format_type(CASE WHEN p.proallargtypes IS NULL \ + THEN p.proargtypes[ordinality - 1] \ + ELSE p.proallargtypes[ordinality] END, NULL) AS type_name \ + FROM target p, LATERAL generate_series( \ + 1, COALESCE(array_length(p.proallargtypes, 1), p.pronargs) \ + ) AS ordinality \ + ) \ + SELECT 0::int4, ''::text, 'r'::text, format_type(prorettype, NULL) FROM target {return_filter} \ + UNION ALL SELECT ordinal, name, mode, type_name FROM arguments \ + ORDER BY 1" + ); + metadata_rows( + application, + &request.scope.datasource_id, + Some(&request.scope.database_name), + &sql, + &[&oid], + ) + .await +} + +fn postgres_function_column_type(mode: &str, ordinal: i32) -> i32 { + if ordinal == 0 || mode == "r" { + return 4; + } + match mode { + "i" | "v" => 1, + "b" => 2, + "o" | "t" => 3, + _ => 0, + } +} + +fn postgres_procedure_column_type(mode: &str) -> i32 { + match mode { + "i" | "v" => 1, + "b" => 2, + "o" | "t" => 4, + _ => 0, + } +} + +async fn list_triggers( + application: &Application, + request: &ListTriggersRequest, +) -> Result { + validate_metadata_scope(&request.scope.database_name, &request.scope.schema_name)?; + let rows = metadata_rows( + application, + &request.scope.datasource_id, + Some(&request.scope.database_name), + "SELECT current_database(), n.nspname, t.tgname, \ + concat_ws(',', \ + CASE WHEN (t.tgtype & 4) <> 0 THEN 'INSERT' END, \ + CASE WHEN (t.tgtype & 8) <> 0 THEN 'DELETE' END, \ + CASE WHEN (t.tgtype & 16) <> 0 THEN 'UPDATE' END, \ + CASE WHEN (t.tgtype & 32) <> 0 THEN 'TRUNCATE' END), \ + pg_get_triggerdef(t.oid, true) \ + FROM pg_trigger t \ + JOIN pg_class c ON c.oid = t.tgrelid \ + JOIN pg_namespace n ON n.oid = c.relnamespace \ + WHERE NOT t.tgisinternal AND n.nspname = $1 \ + ORDER BY t.tgname", + &[&request.scope.schema_name], + ) + .await?; + Ok(TriggerList { + items: rows + .iter() + .map(postgres_trigger_metadata) + .collect::>()?, + }) +} + +async fn get_trigger( + application: &Application, + request: &MetadataObjectRef, +) -> Result { + validate_metadata_scope(&request.scope.database_name, &request.scope.schema_name)?; + validate_identifier(&request.object_name, "triggerName")?; + let mut rows = metadata_rows( + application, + &request.scope.datasource_id, + Some(&request.scope.database_name), + "SELECT current_database(), n.nspname, t.tgname, \ + concat_ws(',', \ + CASE WHEN (t.tgtype & 4) <> 0 THEN 'INSERT' END, \ + CASE WHEN (t.tgtype & 8) <> 0 THEN 'DELETE' END, \ + CASE WHEN (t.tgtype & 16) <> 0 THEN 'UPDATE' END, \ + CASE WHEN (t.tgtype & 32) <> 0 THEN 'TRUNCATE' END), \ + pg_get_triggerdef(t.oid, true) \ + FROM pg_trigger t \ + JOIN pg_class c ON c.oid = t.tgrelid \ + JOIN pg_namespace n ON n.oid = c.relnamespace \ + WHERE NOT t.tgisinternal AND n.nspname = $1 AND t.tgname = $2", + &[&request.scope.schema_name, &request.object_name], + ) + .await?; + let row = rows.pop().ok_or_else(|| { + metadata_not_found( + "trigger", + &request.scope.database_name, + &request.scope.schema_name, + &request.object_name, + ) + })?; + postgres_trigger_metadata(&row) +} + +fn postgres_trigger_metadata(row: &Row) -> Result { + Ok(TriggerMetadata { + database_name: row.try_get(0).map_err(postgres_query_error)?, + schema_name: row.try_get(1).map_err(postgres_query_error)?, + name: row.try_get(2).map_err(postgres_query_error)?, + event_manipulation: row.try_get(3).map_err(postgres_query_error)?, + body: row.try_get(4).map_err(postgres_query_error)?, + }) +} + +async fn load_er_tables( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, +) -> Result, AppError> { + validate_metadata_scope(database_name, schema_name)?; + let tables = list_tables(application, datasource_id, database_name, schema_name, "").await?; + let column_rows = metadata_rows( + application, + datasource_id, + Some(database_name), + "SELECT c.relname, a.attname, format_type(a.atttypid, a.atttypmod), \ + EXISTS (SELECT 1 FROM pg_constraint pk \ + WHERE pk.conrelid = c.oid AND pk.contype = 'p' \ + AND a.attnum = ANY(pk.conkey)), \ + COALESCE(col_description(c.oid, a.attnum), '') \ + FROM pg_class c \ + JOIN pg_namespace n ON n.oid = c.relnamespace \ + JOIN pg_attribute a ON a.attrelid = c.oid AND a.attnum > 0 AND NOT a.attisdropped \ + WHERE n.nspname = $1 AND c.relkind IN ('r', 'p', 'f') \ + ORDER BY c.relname, a.attnum", + &[&schema_name], + ) + .await?; + let foreign_rows = metadata_rows( + application, + datasource_id, + Some(database_name), + "SELECT pt.relname, pa.attname, ft.relname, fa.attname \ + FROM pg_constraint con \ + JOIN pg_class ft ON ft.oid = con.conrelid \ + JOIN pg_namespace fn ON fn.oid = ft.relnamespace \ + JOIN pg_class pt ON pt.oid = con.confrelid \ + JOIN pg_namespace pn ON pn.oid = pt.relnamespace \ + JOIN LATERAL generate_subscripts(con.conkey, 1) pos(n) ON true \ + JOIN pg_attribute fa ON fa.attrelid = ft.oid AND fa.attnum = con.conkey[pos.n] \ + JOIN pg_attribute pa ON pa.attrelid = pt.oid AND pa.attnum = con.confkey[pos.n] \ + WHERE con.contype = 'f' AND fn.nspname = $1 AND pn.nspname = $1 \ + ORDER BY ft.relname, con.conname, pos.n", + &[&schema_name], + ) + .await?; + + let mut result = tables + .items + .into_iter() + .map(|table| EntityRelationTable { + name: table.name, + comment: table.comment, + columns: Vec::new(), + foreign_keys: Vec::new(), + }) + .collect::>(); + let indexes = result + .iter() + .enumerate() + .map(|(index, table)| (table.name.clone(), index)) + .collect::>(); + for row in column_rows { + let table_name: String = row.try_get(0).map_err(postgres_query_error)?; + if let Some(index) = indexes.get(&table_name) { + result[*index].columns.push(EntityRelationColumn { + name: row.try_get(1).map_err(postgres_query_error)?, + column_type: row.try_get(2).map_err(postgres_query_error)?, + primary_key: row.try_get(3).map_err(postgres_query_error)?, + comment: row.try_get(4).map_err(postgres_query_error)?, + }); + } + } + for row in foreign_rows { + let foreign_table: String = row.try_get(2).map_err(postgres_query_error)?; + if let Some(index) = indexes.get(&foreign_table) { + result[*index].foreign_keys.push(EntityRelationForeignKey { + primary_table: row.try_get(0).map_err(postgres_query_error)?, + primary_column: row.try_get(1).map_err(postgres_query_error)?, + foreign_table, + foreign_column: row.try_get(3).map_err(postgres_query_error)?, + }); + } + } + Ok(result) +} + +async fn table_ddl( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + table_name: &str, +) -> Result { + validate_metadata_table(database_name, schema_name, table_name)?; + let resolved = resolve_native_connection(application, datasource_id).await?; + let connection = open_resolved_connection(&resolved, Some(database_name)).await?; + let result = build_table_ddl(&connection, database_name, schema_name, table_name).await; + finish_connection(connection, result).await +} + +#[allow( + clippy::too_many_lines, + reason = "table DDL reconstruction intentionally assembles all PostgreSQL table clauses in catalog order" +)] +async fn build_table_ddl( + connection: &ManagedPostgresConnection, + database_name: &str, + schema_name: &str, + table_name: &str, +) -> Result { + let header_rows = metadata_query_on( + connection, + "SELECT c.relkind::text, c.relpersistence::text, \ + COALESCE(obj_description(c.oid, 'pg_class'), ''), \ + COALESCE(ts.spcname, ''), COALESCE(pg_get_partkeydef(c.oid), ''), \ + COALESCE(fs.srvname, ''), \ + COALESCE(array_to_string(ft.ftoptions, E'\\n'), '') \ + FROM pg_class c \ + JOIN pg_namespace n ON n.oid = c.relnamespace \ + LEFT JOIN pg_tablespace ts ON ts.oid = c.reltablespace \ + LEFT JOIN pg_foreign_table ft ON ft.ftrelid = c.oid \ + LEFT JOIN pg_foreign_server fs ON fs.oid = ft.ftserver \ + WHERE n.nspname = $1 AND c.relname = $2 AND c.relkind IN ('r', 'p', 'f')", + &[&schema_name, &table_name], + ) + .await?; + let header = header_rows + .first() + .ok_or_else(|| metadata_not_found("table", database_name, schema_name, table_name))?; + let kind: String = header.try_get(0).map_err(postgres_query_error)?; + let persistence: String = header.try_get(1).map_err(postgres_query_error)?; + let table_comment: String = header.try_get(2).map_err(postgres_query_error)?; + let tablespace: String = header.try_get(3).map_err(postgres_query_error)?; + let partition_key: String = header.try_get(4).map_err(postgres_query_error)?; + let foreign_server: String = header.try_get(5).map_err(postgres_query_error)?; + let foreign_options: String = header.try_get(6).map_err(postgres_query_error)?; + let columns = metadata_query_on( + connection, + "SELECT a.attname, format_type(a.atttypid, a.atttypmod), a.attnotnull, \ + pg_get_expr(ad.adbin, ad.adrelid), a.attidentity::text, \ + a.attgenerated::text, COALESCE(coll_ns.nspname, ''), \ + COALESCE(coll.collname, ''), COALESCE(col_description(c.oid, a.attnum), '') \ + FROM pg_class c \ + JOIN pg_namespace n ON n.oid = c.relnamespace \ + JOIN pg_attribute a ON a.attrelid = c.oid AND a.attnum > 0 AND NOT a.attisdropped \ + LEFT JOIN pg_attrdef ad ON ad.adrelid = c.oid AND ad.adnum = a.attnum \ + LEFT JOIN pg_collation coll ON coll.oid = a.attcollation AND a.attcollation <> 0 \ + LEFT JOIN pg_namespace coll_ns ON coll_ns.oid = coll.collnamespace \ + WHERE n.nspname = $1 AND c.relname = $2 ORDER BY a.attnum", + &[&schema_name, &table_name], + ) + .await?; + let constraints = metadata_query_on( + connection, + "SELECT con.conname, pg_get_constraintdef(con.oid, true) \ + FROM pg_constraint con \ + JOIN pg_class c ON c.oid = con.conrelid \ + JOIN pg_namespace n ON n.oid = c.relnamespace \ + WHERE n.nspname = $1 AND c.relname = $2 \ + ORDER BY CASE con.contype WHEN 'p' THEN 0 WHEN 'u' THEN 1 \ + WHEN 'f' THEN 2 WHEN 'c' THEN 3 ELSE 4 END, con.conname", + &[&schema_name, &table_name], + ) + .await?; + let indexes = metadata_query_on( + connection, + "SELECT pg_get_indexdef(i.indexrelid) \ + FROM pg_index i \ + JOIN pg_class c ON c.oid = i.indrelid \ + JOIN pg_namespace n ON n.oid = c.relnamespace \ + LEFT JOIN pg_constraint con ON con.conindid = i.indexrelid \ + WHERE n.nspname = $1 AND c.relname = $2 AND con.oid IS NULL \ + ORDER BY i.indexrelid::regclass::text", + &[&schema_name, &table_name], + ) + .await?; + + let qualified = format!( + "{}.{}", + quote_identifier(schema_name, "schemaName")?, + quote_identifier(table_name, "tableName")? + ); + let mut ddl = String::new(); + let prefix = match kind.as_str() { + "f" => "CREATE FOREIGN TABLE", + _ if persistence == "u" => "CREATE UNLOGGED TABLE", + _ => "CREATE TABLE", + }; + writeln!(&mut ddl, "{prefix} {qualified} (").map_err(|_| AppError::internal())?; + let mut definitions = Vec::new(); + let mut column_comments = Vec::new(); + for row in columns { + let name: String = row.try_get(0).map_err(postgres_query_error)?; + let data_type: String = row.try_get(1).map_err(postgres_query_error)?; + let not_null: bool = row.try_get(2).map_err(postgres_query_error)?; + let default_value: Option = row.try_get(3).map_err(postgres_query_error)?; + let identity: String = row.try_get(4).map_err(postgres_query_error)?; + let generated: String = row.try_get(5).map_err(postgres_query_error)?; + let collation_schema: String = row.try_get(6).map_err(postgres_query_error)?; + let collation_name: String = row.try_get(7).map_err(postgres_query_error)?; + let comment: String = row.try_get(8).map_err(postgres_query_error)?; + let mut definition = format!(" {} {data_type}", quote_identifier(&name, "columnName")?); + if !collation_name.is_empty() { + write!( + &mut definition, + " COLLATE {}.{}", + quote_identifier(&collation_schema, "collationSchema")?, + quote_identifier(&collation_name, "collationName")? + ) + .map_err(|_| AppError::internal())?; + } + if generated == "s" { + if let Some(expression) = default_value.as_deref() { + write!( + &mut definition, + " GENERATED ALWAYS AS ({expression}) STORED" + ) + .map_err(|_| AppError::internal())?; + } + } else if !identity.is_empty() { + write!( + &mut definition, + " GENERATED {} AS IDENTITY", + if identity == "a" { + "ALWAYS" + } else { + "BY DEFAULT" + } + ) + .map_err(|_| AppError::internal())?; + } else if let Some(default_value) = default_value { + write!(&mut definition, " DEFAULT {default_value}") + .map_err(|_| AppError::internal())?; + } + if not_null { + definition.push_str(" NOT NULL"); + } + definitions.push(definition); + if !comment.is_empty() { + column_comments.push((name, comment)); + } + } + for row in constraints { + let name: String = row.try_get(0).map_err(postgres_query_error)?; + let definition: String = row.try_get(1).map_err(postgres_query_error)?; + definitions.push(format!( + " CONSTRAINT {} {definition}", + quote_identifier(&name, "constraintName")? + )); + } + ddl.push_str(&definitions.join(",\n")); + ddl.push_str("\n)"); + if kind == "p" && !partition_key.is_empty() { + write!(&mut ddl, " PARTITION BY {partition_key}").map_err(|_| AppError::internal())?; + } + if kind == "f" { + write!( + &mut ddl, + " SERVER {}", + quote_identifier(&foreign_server, "foreignServer")? + ) + .map_err(|_| AppError::internal())?; + let options = render_foreign_options(&foreign_options)?; + if !options.is_empty() { + write!(&mut ddl, " OPTIONS ({options})").map_err(|_| AppError::internal())?; + } + } + if !tablespace.is_empty() && kind != "f" { + write!( + &mut ddl, + " TABLESPACE {}", + quote_identifier(&tablespace, "tablespace")? + ) + .map_err(|_| AppError::internal())?; + } + ddl.push_str(";\n"); + for row in indexes { + let definition: String = row.try_get(0).map_err(postgres_query_error)?; + writeln!(&mut ddl, "{};", definition.trim_end_matches(';')) + .map_err(|_| AppError::internal())?; + } + if !table_comment.is_empty() { + writeln!( + &mut ddl, + "COMMENT ON TABLE {qualified} IS {};", + quote_literal(&table_comment)? + ) + .map_err(|_| AppError::internal())?; + } + for (name, comment) in column_comments { + writeln!( + &mut ddl, + "COMMENT ON COLUMN {qualified}.{} IS {};", + quote_identifier(&name, "columnName")?, + quote_literal(&comment)? + ) + .map_err(|_| AppError::internal())?; + } + Ok(ddl.trim_end().to_owned()) +} + +async fn metadata_query_on( + connection: &ManagedPostgresConnection, + sql: &str, + parameters: &[&(dyn ToSql + Sync)], +) -> Result, AppError> { + postgres_timeout( + METADATA_TIMEOUT, + "postgres_metadata_timeout", + "The PostgreSQL metadata query did not finish in time", + connection.client().query(sql, parameters), + ) + .await +} + +fn render_foreign_options(options: &str) -> Result { + options + .lines() + .filter(|line| !line.is_empty()) + .map(|option| { + let (name, value) = option.split_once('=').ok_or_else(AppError::internal)?; + Ok(format!( + "{} {}", + quote_identifier(name, "foreignOption")?, + quote_literal(value)? + )) + }) + .collect::, AppError>>() + .map(|values| values.join(", ")) +} + +fn quote_literal(value: &str) -> Result { + if value.len() > MAX_SCALAR_BYTES || value.contains('\0') { + return Err(AppError::invalid( + "invalid_postgres_literal", + "The PostgreSQL literal is invalid", + )); + } + let escaped_length = value.chars().try_fold(0_usize, |length, character| { + length.checked_add(match character { + '\\' | '\'' => 2, + _ => character.len_utf8(), + }) + }); + let Some(escaped_length) = escaped_length + .and_then(|length| length.checked_add(3)) + .filter(|length| *length <= MAX_SCALAR_BYTES) + else { + return Err(AppError::invalid( + "invalid_postgres_literal", + "The escaped PostgreSQL literal exceeds the scalar byte limit", + )); + }; + let mut escaped = String::with_capacity(escaped_length); + for character in value.chars() { + match character { + '\\' => escaped.push_str("\\\\"), + '\'' => escaped.push_str("''"), + _ => escaped.push(character), + } + } + Ok(format!("E'{escaped}'")) +} + +async fn start_table_preview( + application: &Application, + request: TablePreviewRequest, + row_limit: u32, +) -> Result { + if row_limit == 0 || row_limit > MAX_CONSOLE_PAGE_SIZE { + return Err(AppError::invalid( + "invalid_table_preview_request", + format!("rowLimit must be between 1 and {MAX_CONSOLE_PAGE_SIZE}"), + )); + } + validate_metadata_table( + &request.table.scope.database_name, + &request.table.scope.schema_name, + &request.table.table_name, + )?; + let sql = format!( + "SELECT * FROM {}.{} LIMIT {row_limit}", + quote_identifier(&request.table.scope.schema_name, "schemaName")?, + quote_identifier(&request.table.table_name, "tableName")? + ); + let accepted = application + .start_read_query(StartQueryRequest { + datasource_id: request.table.scope.datasource_id, + sql: sql.clone(), + parameters: Vec::new(), + limits: QueryLimits { + max_rows: row_limit.to_string(), + max_result_bytes: (8 * 1024 * 1024_u64).to_string(), + batch_rows: row_limit.min(200), + batch_bytes: 1024 * 1024, + result_ttl_seconds: 60 * 60, + }, + }) + .await?; + Ok(TablePreviewAccepted { + operation_id: accepted.operation_id, + sql, + row_limit, + }) +} + +fn build_create_schema(request: CreateSchemaSqlRequest) -> Result { + let schema = request.schema; + let name = quote_identifier(&schema.name, "schemaName")?; + let mut sql = format!("CREATE SCHEMA {name}"); + if !schema.owner.trim().is_empty() { + write!( + &mut sql, + " AUTHORIZATION {}", + quote_identifier(&schema.owner, "owner")? + ) + .map_err(|_| AppError::internal())?; + } + sql.push(';'); + if !schema.comment.is_empty() { + write!( + &mut sql, + "\nCOMMENT ON SCHEMA {name} IS {};", + quote_literal(&schema.comment)? + ) + .map_err(|_| AppError::internal())?; + } + Ok(BuiltSql { sql }) +} + +fn build_namespace_sql(request: NamespaceSqlRequest) -> Result { + let sql = match request.operation { + NamespaceSqlOperation::CreateDatabase { database } => build_create_database(&database)?, + NamespaceSqlOperation::AlterDatabase { + old_database, + new_database, + } => build_alter_database(&old_database, &new_database)?, + NamespaceSqlOperation::DropDatabase { database_name } => format!( + "DROP DATABASE {};", + quote_identifier(&database_name, "databaseName")? + ), + NamespaceSqlOperation::UseDatabase { .. } => { + return Err(AppError::invalid( + "postgres_database_switch_unsupported", + "PostgreSQL selects a database when opening the connection and cannot switch it with SQL", + )); + } + NamespaceSqlOperation::CreateSchema { schema } => { + return build_create_schema(CreateSchemaSqlRequest { schema }); + } + NamespaceSqlOperation::AlterSchema { + old_schema_name, + new_schema_name, + } => format!( + "ALTER SCHEMA {} RENAME TO {};", + quote_identifier(&old_schema_name, "schemaName")?, + quote_identifier(&new_schema_name, "schemaName")? + ), + NamespaceSqlOperation::DropSchema { schema_name } => format!( + "DROP SCHEMA {};", + quote_identifier(&schema_name, "schemaName")? + ), + }; + Ok(BuiltSql { sql }) +} + +fn build_create_database( + database: &crate::native_driver_types::DatabaseDefinition, +) -> Result { + let name = quote_identifier(&database.name, "databaseName")?; + let mut sql = format!("CREATE DATABASE {name}"); + if !database.owner.trim().is_empty() { + write!( + &mut sql, + " OWNER {}", + quote_identifier(&database.owner, "owner")? + ) + .map_err(|_| AppError::internal())?; + } + if !database.charset.trim().is_empty() { + write!(&mut sql, " ENCODING {}", quote_literal(&database.charset)?) + .map_err(|_| AppError::internal())?; + } + if !database.collation.trim().is_empty() { + write!( + &mut sql, + " LC_COLLATE {} LC_CTYPE {}", + quote_literal(&database.collation)?, + quote_literal(&database.collation)? + ) + .map_err(|_| AppError::internal())?; + } + sql.push(';'); + if !database.comment.is_empty() { + write!( + &mut sql, + "\nCOMMENT ON DATABASE {name} IS {};", + quote_literal(&database.comment)? + ) + .map_err(|_| AppError::internal())?; + } + Ok(sql) +} + +fn build_alter_database( + old_database: &crate::native_driver_types::DatabaseDefinition, + new_database: &crate::native_driver_types::DatabaseDefinition, +) -> Result { + if old_database.charset != new_database.charset + || old_database.collation != new_database.collation + { + return Err(AppError::invalid( + "postgres_database_alter_unsupported", + "PostgreSQL cannot alter a database encoding or collation in place", + )); + } + let old_name = quote_identifier(&old_database.name, "databaseName")?; + let new_name = quote_identifier(&new_database.name, "databaseName")?; + let mut statements = Vec::new(); + if old_database.name != new_database.name { + statements.push(format!("ALTER DATABASE {old_name} RENAME TO {new_name};")); + } + let active_name = if old_database.name == new_database.name { + old_name + } else { + new_name + }; + if old_database.owner != new_database.owner && !new_database.owner.trim().is_empty() { + statements.push(format!( + "ALTER DATABASE {active_name} OWNER TO {};", + quote_identifier(&new_database.owner, "owner")? + )); + } + if old_database.comment != new_database.comment { + statements.push(format!( + "COMMENT ON DATABASE {active_name} IS {};", + quote_literal(&new_database.comment)? + )); + } + if statements.is_empty() { + return Err(AppError::invalid( + "postgres_database_alter_empty", + "The PostgreSQL database definition has no supported changes", + )); + } + Ok(statements.join("\n")) +} + +fn build_dml(request: DmlSqlRequest) -> Result { + let target = postgres_dml_target(&request.target)?; + let sql = match request.statement { + DmlStatement::SingleInsert { columns, row } => { + postgres_insert_sql(&target, &columns, std::slice::from_ref(&row))? + } + DmlStatement::MultiInsert { columns, rows } => { + postgres_insert_sql(&target, &columns, &rows)? + } + DmlStatement::Update { + assignments, + predicates, + } => postgres_update_sql(&target, &assignments, &predicates)?, + }; + Ok(BuiltSql { sql }) +} + +fn postgres_dml_target(target: &DmlTarget) -> Result { + let table = quote_identifier(&target.table_name, "tableName")?; + match target + .schema_name + .as_deref() + .filter(|value| !value.is_empty()) + { + Some(schema) => Ok(format!( + "{}.{}", + quote_identifier(schema, "schemaName")?, + table + )), + None => Ok(table), + } +} + +fn postgres_insert_sql( + target: &str, + columns: &[DmlColumn], + rows: &[DmlRow], +) -> Result { + if columns.is_empty() || rows.is_empty() { + return Err(AppError::invalid( + "invalid_postgres_dml", + "PostgreSQL INSERT requires at least one column and row", + )); + } + let column_sql = columns + .iter() + .map(|column| quote_identifier(&column.name, "columnName")) + .collect::, _>>()? + .join(", "); + let value_rows = rows + .iter() + .map(|row| { + if row.values.len() != columns.len() { + return Err(AppError::invalid( + "invalid_postgres_dml", + "Each PostgreSQL INSERT row must match the selected column count", + )); + } + row.values + .iter() + .zip(columns) + .map(|(value, column)| postgres_dml_value(value, column)) + .collect::, _>>() + .map(|values| format!("({})", values.join(", "))) + }) + .collect::, _>>()? + .join(",\n"); + Ok(format!( + "INSERT INTO {target} ({column_sql}) VALUES\n{value_rows};" + )) +} + +fn postgres_update_sql( + target: &str, + assignments: &[DmlAssignment], + predicates: &[DmlAssignment], +) -> Result { + if assignments.is_empty() || predicates.is_empty() { + return Err(AppError::invalid( + "invalid_postgres_dml", + "PostgreSQL UPDATE requires assignments and key predicates", + )); + } + let assignments = assignments + .iter() + .map(|assignment| { + Ok(format!( + "{} = {}", + quote_identifier(&assignment.column.name, "columnName")?, + postgres_dml_value(&assignment.value, &assignment.column)? + )) + }) + .collect::, AppError>>()? + .join(", "); + let predicates = predicates + .iter() + .map(|predicate| { + let column = quote_identifier(&predicate.column.name, "columnName")?; + match predicate.value { + DmlValue::Null => Ok(format!("{column} IS NULL")), + _ => Ok(format!( + "{column} = {}", + postgres_dml_value(&predicate.value, &predicate.column)? + )), + } + }) + .collect::, AppError>>()? + .join(" AND "); + Ok(format!( + "UPDATE {target} SET {assignments} WHERE {predicates};" + )) +} + +fn postgres_dml_value(value: &DmlValue, column: &DmlColumn) -> Result { + match value { + DmlValue::Null => Ok("NULL".to_owned()), + DmlValue::String(value) => quote_literal(value), + DmlValue::Decimal(value) => { + validate_decimal(value)?; + Ok(value.clone()) + } + DmlValue::Boolean(value) => Ok(if *value { "TRUE" } else { "FALSE" }.to_owned()), + DmlValue::Temporal { kind, iso8601 } => { + validate_temporal(*kind, iso8601)?; + let prefix = match kind { + DmlTemporalKind::Date => "DATE", + DmlTemporalKind::Time => "TIME", + DmlTemporalKind::LocalDatetime => "TIMESTAMP", + DmlTemporalKind::OffsetDatetime => "TIMESTAMPTZ", + }; + Ok(format!("{prefix} {}", quote_literal(iso8601)?)) + } + DmlValue::Binary(value) => { + if value.len() > MAX_SCALAR_BYTES { + return Err(AppError::invalid( + "invalid_postgres_dml", + "The PostgreSQL binary value exceeds the scalar limit", + )); + } + Ok(format!("decode('{}', 'hex')", hex::encode(value))) + } + } + .map(|sql| { + if matches!(value, DmlValue::String(_)) + && column.data_type_name.eq_ignore_ascii_case("uuid") + { + format!("{sql}::uuid") + } else { + sql + } + }) +} + +fn is_native_read_candidate(sql: &str) -> Result { + let words = postgres_words(sql)?; + Ok(matches!( + words.first().map(String::as_str), + Some("SELECT" | "WITH" | "VALUES" | "TABLE" | "EXPLAIN") + )) +} + +fn validate_query(query: &PreparedQuery) -> Result<(), AppError> { + if query.sql.len() > MAX_SQL_BYTES { + return Err(AppError::invalid( + "invalid_query_request", + format!("SQL cannot exceed {MAX_SQL_BYTES} UTF-8 bytes"), + )); + } + validate_read_sql(&query.sql)?; + let _ = postgres_query_parameters(&query.parameters)?; + validate_query_options(query.options) +} + +fn validate_read_sql(sql: &str) -> Result<(), AppError> { + let statements = split_postgres_script(sql)?; + if statements.len() != 1 { + return Err(AppError::invalid( + "postgres_native_query_unsupported", + "Native PostgreSQL accepts exactly one read statement", + )); + } + let words = postgres_words(&statements[0])?; + if !matches!( + words.first().map(String::as_str), + Some("SELECT" | "WITH" | "VALUES" | "TABLE" | "EXPLAIN") + ) { + return Err(AppError::invalid( + "postgres_native_query_unsupported", + "Native PostgreSQL supports SELECT, WITH, VALUES, TABLE, and EXPLAIN read statements", + )); + } + let forbidden = words.iter().any(|word| { + matches!( + word.as_str(), + "INSERT" + | "UPDATE" + | "DELETE" + | "MERGE" + | "CREATE" + | "ALTER" + | "DROP" + | "TRUNCATE" + | "COPY" + | "CALL" + | "DO" + | "GRANT" + | "REVOKE" + | "LOCK" + ) + }) || words.windows(2).any(|words| { + matches!( + words, + [first, second] + if (first == "FOR" && matches!(second.as_str(), "UPDATE" | "SHARE")) + || (first == "SELECT" && second == "INTO") + ) + }) || words.windows(3).any(|words| { + matches!( + words, + [first, second, third] + if first == "FOR" + && ((second == "KEY" && third == "SHARE") + || (second == "NO" && third == "KEY")) + ) + }); + if forbidden { + return Err(AppError::invalid( + "postgres_native_query_unsupported", + "Native PostgreSQL read queries must not write data, create objects, or lock rows", + )); + } + Parser::parse_sql(&PostgreSqlDialect {}, &statements[0]).map_err(|_| { + AppError::invalid( + "postgres_native_query_unsupported", + "Native PostgreSQL requires one valid read statement", + ) + })?; + Ok(()) +} + +fn validate_query_options(options: QueryExecutionOptions) -> Result<(), AppError> { + if options.target_batch_rows > MAX_BATCH_ROWS { + return Err(AppError::invalid( + "invalid_query_limits", + format!("batchRows cannot exceed {MAX_BATCH_ROWS}"), + )); + } + if options.target_batch_bytes != 0 + && !(1024..=MAX_BATCH_BYTES).contains(&options.target_batch_bytes) + { + return Err(AppError::invalid( + "invalid_query_limits", + format!("batchBytes must be zero or between 1024 and {MAX_BATCH_BYTES}"), + )); + } + if options.max_result_bytes > MAX_RESULT_BYTES { + return Err(AppError::invalid( + "invalid_query_limits", + format!("maxResultBytes cannot exceed {MAX_RESULT_BYTES}"), + )); + } + Ok(()) +} + +#[derive(Debug)] +struct PostgresParameter(Option); + +impl ToSql for PostgresParameter { + fn to_sql( + &self, + _ty: &Type, + out: &mut BytesMut, + ) -> Result> { + match &self.0 { + Some(value) => { + out.extend_from_slice(value.as_bytes()); + Ok(IsNull::No) + } + None => Ok(IsNull::Yes), + } + } + + fn accepts(_ty: &Type) -> bool { + true + } + + fn encode_format(&self, _ty: &Type) -> Format { + Format::Text + } + + tokio_postgres::types::to_sql_checked!(); +} + +fn postgres_query_parameters( + parameters: &[QueryParameter], +) -> Result, AppError> { + if parameters.len() > MAX_PARAMETERS { + return Err(AppError::invalid( + "invalid_query_parameter_count", + format!("PostgreSQL queries accept at most {MAX_PARAMETERS} parameters"), + )); + } + let mut ordered = parameters.iter().collect::>(); + ordered.sort_unstable_by_key(|parameter| parameter.position); + ordered + .into_iter() + .enumerate() + .map(|(index, parameter)| { + let expected = u32::try_from(index + 1).map_err(|_| AppError::internal())?; + if parameter.position != expected { + return Err(AppError::invalid( + "invalid_query_parameter", + "PostgreSQL parameter positions must be unique and contiguous from 1", + )); + } + postgres_query_parameter(¶meter.value) + }) + .collect() +} + +fn postgres_query_parameter(value: &DatabaseValue) -> Result { + let text = match value { + DatabaseValue::Null => None, + DatabaseValue::Boolean(value) => Some(if *value { "true" } else { "false" }.to_owned()), + DatabaseValue::SignedInteger(value) => Some(value.to_string()), + DatabaseValue::UnsignedInteger(value) => Some(value.to_string()), + DatabaseValue::Float32(value) => Some(value.to_string()), + DatabaseValue::Float64(value) => Some(value.to_string()), + DatabaseValue::Decimal(value) => { + validate_decimal(value)?; + Some(value.clone()) + } + DatabaseValue::Text(value) => Some(validate_parameter_text(value, "text")?), + DatabaseValue::Binary(value) => { + if value.len() > MAX_SCALAR_BYTES { + return Err(invalid_parameter("binary")); + } + Some(format!("\\x{}", hex::encode(value))) + } + DatabaseValue::Date(value) => { + NaiveDate::parse_from_str(value, "%Y-%m-%d").map_err(|_| invalid_parameter("date"))?; + Some(validate_parameter_text(value, "date")?) + } + DatabaseValue::Time(value) => { + parse_postgres_time(value)?; + Some(validate_parameter_text(value, "time")?) + } + DatabaseValue::Timestamp(value) => { + parse_postgres_timestamp(value)?; + Some(validate_parameter_text(value, "timestamp")?) + } + DatabaseValue::TimestampWithTimeZone(value) => { + DateTime::parse_from_rfc3339(value) + .map_err(|_| invalid_parameter("timestamp with time zone"))?; + Some(validate_parameter_text(value, "timestamp with time zone")?) + } + DatabaseValue::Json(value) => { + serde_json::from_str::(value) + .map_err(|_| invalid_parameter("JSON"))?; + Some(validate_parameter_text(value, "JSON")?) + } + DatabaseValue::Uuid(value) => { + uuid::Uuid::parse_str(value).map_err(|_| invalid_parameter("UUID"))?; + Some(validate_parameter_text(value, "UUID")?) + } + }; + Ok(PostgresParameter(text)) +} + +fn validate_parameter_text(value: &str, label: &str) -> Result { + if value.len() > MAX_SCALAR_BYTES || value.contains('\0') { + return Err(invalid_parameter(label)); + } + Ok(value.to_owned()) +} + +fn validate_decimal(value: &str) -> Result<(), AppError> { + let unsigned = value.strip_prefix(['+', '-']).unwrap_or(value); + let (mantissa, exponent) = unsigned + .split_once(['e', 'E']) + .map_or((unsigned, None), |(mantissa, exponent)| { + (mantissa, Some(exponent)) + }); + let mut digits = 0_usize; + let mut points = 0_u8; + for byte in mantissa.bytes() { + if byte.is_ascii_digit() { + digits += 1; + } else if byte == b'.' { + points += 1; + } else { + return Err(invalid_parameter("decimal")); + } + } + if digits == 0 || points > 1 { + return Err(invalid_parameter("decimal")); + } + if let Some(exponent) = exponent { + let exponent = exponent.strip_prefix(['+', '-']).unwrap_or(exponent); + if exponent.is_empty() || !exponent.bytes().all(|byte| byte.is_ascii_digit()) { + return Err(invalid_parameter("decimal")); + } + } + Ok(()) +} + +fn validate_temporal(kind: DmlTemporalKind, value: &str) -> Result<(), AppError> { + match kind { + DmlTemporalKind::Date => { + NaiveDate::parse_from_str(value, "%Y-%m-%d").map_err(|_| invalid_parameter("date"))?; + } + DmlTemporalKind::Time => { + parse_postgres_time(value)?; + } + DmlTemporalKind::LocalDatetime => { + parse_postgres_timestamp(value)?; + } + DmlTemporalKind::OffsetDatetime => { + DateTime::parse_from_rfc3339(value) + .map_err(|_| invalid_parameter("timestamp with time zone"))?; + } + } + Ok(()) +} + +fn parse_postgres_time(value: &str) -> Result { + ["%H:%M:%S%.f", "%H:%M:%S"] + .into_iter() + .find_map(|format| NaiveTime::parse_from_str(value, format).ok()) + .ok_or_else(|| invalid_parameter("time")) +} + +fn parse_postgres_timestamp(value: &str) -> Result { + ["%Y-%m-%dT%H:%M:%S%.f", "%Y-%m-%d %H:%M:%S%.f"] + .into_iter() + .find_map(|format| NaiveDateTime::parse_from_str(value, format).ok()) + .ok_or_else(|| invalid_parameter("timestamp")) +} + +fn invalid_parameter(label: &str) -> AppError { + AppError::invalid( + "invalid_query_parameter", + format!("The PostgreSQL {label} parameter is invalid"), + ) +} + +#[derive(Debug)] +enum RawPostgresValue { + Bytes(Vec), + TooLarge { byte_count: usize }, +} + +impl<'a> FromSql<'a> for RawPostgresValue { + fn from_sql(_ty: &Type, raw: &'a [u8]) -> Result> { + if raw.len() > MAX_SCALAR_BYTES { + Ok(Self::TooLarge { + byte_count: raw.len(), + }) + } else { + Ok(Self::Bytes(raw.to_vec())) + } + } + + fn accepts(_ty: &Type) -> bool { + true + } +} + +fn postgres_column( + index: usize, + column: &tokio_postgres::Column, +) -> Result { + let ordinal = u32::try_from(index) + .ok() + .and_then(|index| index.checked_add(1)) + .ok_or_else(AppError::internal)?; + let data_type = postgres_base_type(column.type_()); + let value_type = postgres_value_type(data_type); + Ok(wire::JdbcColumn { + ordinal, + label: column.name().to_owned(), + name: column.name().to_owned(), + jdbc_type: postgres_jdbc_type(data_type), + jdbc_type_name: column.type_().name().to_ascii_uppercase(), + value_type: value_type as i32, + nullability: wire::ColumnNullability::Unknown as i32, + precision: postgres_type_precision(data_type), + scale: postgres_type_scale(data_type), + display_size: postgres_type_display_size(data_type), + signed: postgres_type_signed(data_type), + catalog_name: None, + schema_name: None, + table_name: None, + }) +} + +fn postgres_row(row: &Row, columns: &[tokio_postgres::Column]) -> Result { + if row.len() != columns.len() { + return Err(AppError::internal()); + } + Ok(wire::JdbcRow { + values: columns + .iter() + .enumerate() + .map(|(index, column)| postgres_wire_value(row, index, column.type_())) + .collect::>()?, + }) +} + +fn postgres_wire_value( + row: &Row, + index: usize, + data_type: &Type, +) -> Result { + use wire::jdbc_value::Value; + let raw_value = row + .try_get::<_, Option>(index) + .map_err(postgres_query_error)?; + let value = match raw_value { + None => Value::NullValue(wire::JdbcNull {}), + Some(RawPostgresValue::Bytes(bytes)) => decode_postgres_value(data_type, &bytes)?, + Some(RawPostgresValue::TooLarge { byte_count }) => { + return Err(postgres_scalar_too_large(byte_count)); + } + }; + Ok(wire::JdbcValue { value: Some(value) }) +} + +fn decode_postgres_value( + data_type: &Type, + raw: &[u8], +) -> Result { + use wire::jdbc_value::Value; + ensure_postgres_scalar_size(raw.len())?; + if let Kind::Domain(inner) = data_type.kind() { + return decode_postgres_value(inner, raw); + } + if matches!(data_type.kind(), Kind::Array(_)) { + return Ok(Value::OpaqueValue(wire::OpaqueValue { + type_name: data_type.name().to_owned(), + display_value: decode_postgres_array(data_type, raw)?, + })); + } + if matches!(data_type.kind(), Kind::Enum(_)) { + return Ok(Value::TextValue(postgres_utf8(raw)?)); + } + let value = if *data_type == Type::BOOL { + Value::BooleanValue(raw.first().copied() == Some(1)) + } else if *data_type == Type::INT2 { + Value::SignedIntegerValue(i64::from(read_i16(raw)?)) + } else if *data_type == Type::INT4 { + Value::SignedIntegerValue(i64::from(read_i32(raw)?)) + } else if *data_type == Type::INT8 { + Value::SignedIntegerValue(read_i64(raw)?) + } else if postgres_unsigned_type(data_type) { + Value::UnsignedIntegerValue(u64::from(read_u32(raw)?)) + } else if *data_type == Type::FLOAT4 { + Value::Float32Value(f32::from_bits(read_u32(raw)?)) + } else if *data_type == Type::FLOAT8 { + Value::Float64Value(f64::from_bits(read_u64(raw)?)) + } else if *data_type == Type::NUMERIC { + Value::DecimalValue(decode_postgres_numeric(raw)?) + } else if *data_type == Type::MONEY { + Value::OpaqueValue(wire::OpaqueValue { + type_name: data_type.name().to_owned(), + display_value: format!("raw_units={}", read_i64(raw)?), + }) + } else if *data_type == Type::BYTEA { + Value::BinaryValue(raw.to_vec()) + } else if *data_type == Type::DATE { + Value::DateValue(decode_postgres_date(raw)?) + } else if *data_type == Type::TIME { + Value::TimeValue(decode_postgres_time(raw)?) + } else if *data_type == Type::TIMESTAMP { + Value::TimestampValue(decode_postgres_timestamp(raw)?) + } else if *data_type == Type::TIMESTAMPTZ { + Value::TimestampWithTimeZoneValue(format!("{}Z", decode_postgres_timestamp(raw)?)) + } else if *data_type == Type::TIMETZ { + Value::OpaqueValue(wire::OpaqueValue { + type_name: data_type.name().to_owned(), + display_value: decode_postgres_timetz(raw)?, + }) + } else if *data_type == Type::JSON { + Value::JsonValue(postgres_utf8(raw)?) + } else if *data_type == Type::JSONB { + let body = raw.strip_prefix(&[1]).ok_or_else(result_decode_error)?; + Value::JsonValue(postgres_utf8(body)?) + } else if *data_type == Type::UUID { + Value::UuidValue( + uuid::Uuid::from_slice(raw) + .map_err(|_| result_decode_error())? + .to_string(), + ) + } else if *data_type == Type::BIT || *data_type == Type::VARBIT { + Value::OpaqueValue(wire::OpaqueValue { + type_name: data_type.name().to_owned(), + display_value: decode_postgres_bits(raw)?, + }) + } else if *data_type == Type::INET || *data_type == Type::CIDR { + Value::TextValue(decode_postgres_network(data_type, raw)?) + } else if postgres_text_type(data_type) { + Value::TextValue(postgres_utf8(raw)?) + } else { + postgres_opaque_value(data_type.name(), raw)? + }; + Ok(value) +} + +fn postgres_base_type(data_type: &Type) -> &Type { + match data_type.kind() { + Kind::Domain(inner) => postgres_base_type(inner), + _ => data_type, + } +} + +fn postgres_value_type(data_type: &Type) -> wire::JdbcValueType { + if matches!( + data_type.kind(), + Kind::Array(_) | Kind::Range(_) | Kind::Multirange(_) | Kind::Composite(_) + ) { + return wire::JdbcValueType::Opaque; + } + if matches!(data_type.kind(), Kind::Enum(_)) { + return wire::JdbcValueType::Text; + } + if *data_type == Type::BOOL { + wire::JdbcValueType::Boolean + } else if matches!(*data_type, Type::INT2 | Type::INT4 | Type::INT8) { + wire::JdbcValueType::SignedInteger + } else if postgres_unsigned_type(data_type) { + wire::JdbcValueType::UnsignedInteger + } else if *data_type == Type::FLOAT4 { + wire::JdbcValueType::Float32 + } else if *data_type == Type::FLOAT8 { + wire::JdbcValueType::Float64 + } else if *data_type == Type::NUMERIC { + wire::JdbcValueType::Decimal + } else if *data_type == Type::MONEY { + wire::JdbcValueType::Opaque + } else if *data_type == Type::BYTEA { + wire::JdbcValueType::Binary + } else if *data_type == Type::DATE { + wire::JdbcValueType::Date + } else if *data_type == Type::TIME { + wire::JdbcValueType::Time + } else if *data_type == Type::TIMESTAMP { + wire::JdbcValueType::Timestamp + } else if *data_type == Type::TIMESTAMPTZ { + wire::JdbcValueType::TimestampWithTimeZone + } else if *data_type == Type::JSON || *data_type == Type::JSONB { + wire::JdbcValueType::Json + } else if *data_type == Type::UUID { + wire::JdbcValueType::Uuid + } else if postgres_text_type(data_type) || *data_type == Type::INET || *data_type == Type::CIDR + { + wire::JdbcValueType::Text + } else { + wire::JdbcValueType::Opaque + } +} + +fn postgres_jdbc_type(data_type: &Type) -> i32 { + if matches!(data_type.kind(), Kind::Array(_)) { + return 2_003; + } + if *data_type == Type::BOOL { + 16 + } else if *data_type == Type::INT2 { + 5 + } else if *data_type == Type::INT4 { + 4 + } else if *data_type == Type::INT8 { + -5 + } else if *data_type == Type::FLOAT4 { + 7 + } else if *data_type == Type::FLOAT8 { + 8 + } else if *data_type == Type::NUMERIC { + 2 + } else if *data_type == Type::MONEY { + 1_111 + } else if *data_type == Type::BYTEA { + -2 + } else if *data_type == Type::DATE { + 91 + } else if *data_type == Type::TIME { + 92 + } else if *data_type == Type::TIMETZ { + 2_013 + } else if *data_type == Type::TIMESTAMP { + 93 + } else if *data_type == Type::TIMESTAMPTZ { + 2_014 + } else if *data_type == Type::CHAR || *data_type == Type::BPCHAR { + 1 + } else if *data_type == Type::VARCHAR { + 12 + } else if *data_type == Type::TEXT || *data_type == Type::JSON || *data_type == Type::JSONB { + -1 + } else if *data_type == Type::BIT { + -7 + } else if *data_type == Type::VARBIT { + -3 + } else { + 1_111 + } +} + +fn postgres_jdbc_type_name(type_name: &str) -> i32 { + let normalized = type_name.trim().to_ascii_lowercase(); + if normalized.starts_with('_') + || normalized.ends_with("[]") + || normalized.eq_ignore_ascii_case("array") + { + return 2_003; + } + let normalized = normalized.trim_start_matches('_').trim_end_matches("[]"); + let mut without_modifiers = String::with_capacity(normalized.len()); + let mut modifier_depth = 0_u32; + for character in normalized.chars() { + match character { + '(' => modifier_depth = modifier_depth.saturating_add(1), + ')' if modifier_depth > 0 => modifier_depth -= 1, + _ if modifier_depth == 0 => without_modifiers.push(character), + _ => {} + } + } + let canonical = without_modifiers + .split_whitespace() + .collect::>() + .join(" "); + match canonical.as_str() { + "bool" | "boolean" => 16, + "int2" | "smallint" | "smallserial" => 5, + "int4" | "integer" | "serial" => 4, + "int8" | "bigint" | "bigserial" => -5, + "float4" | "real" => 7, + "float8" | "double" | "double precision" => 8, + "numeric" | "decimal" => 2, + "bytea" => -2, + "date" => 91, + "time" | "time without time zone" => 92, + "timetz" | "time with time zone" => 2_013, + "timestamp" | "timestamp without time zone" => 93, + "timestamptz" | "timestamp with time zone" => 2_014, + "char" | "bpchar" | "character" => 1, + "varchar" | "name" | "character varying" => 12, + "text" | "json" | "jsonb" | "xml" => -1, + "bit" => -7, + "varbit" | "bit varying" => -3, + _ => 1_111, + } +} + +fn postgres_unsigned_type(data_type: &Type) -> bool { + matches!( + *data_type, + Type::OID + | Type::REGPROC + | Type::REGPROCEDURE + | Type::REGOPER + | Type::REGOPERATOR + | Type::REGCLASS + | Type::REGTYPE + | Type::REGROLE + | Type::REGNAMESPACE + ) +} + +fn postgres_text_type(data_type: &Type) -> bool { + matches!( + *data_type, + Type::CHAR + | Type::NAME + | Type::TEXT + | Type::BPCHAR + | Type::VARCHAR + | Type::UNKNOWN + | Type::XML + ) +} + +fn postgres_type_precision(data_type: &Type) -> Option { + match *data_type { + Type::INT2 => Some(5), + Type::INT4 => Some(10), + Type::INT8 => Some(19), + Type::FLOAT4 => Some(8), + Type::FLOAT8 => Some(17), + _ => None, + } +} + +fn postgres_type_scale(data_type: &Type) -> Option { + matches!(*data_type, Type::INT2 | Type::INT4 | Type::INT8).then_some(0) +} + +fn postgres_type_display_size(data_type: &Type) -> Option { + match *data_type { + Type::BOOL => Some(5), + Type::INT2 => Some(6), + Type::INT4 => Some(11), + Type::INT8 => Some(20), + Type::DATE => Some(10), + Type::TIME => Some(15), + Type::TIMESTAMP => Some(29), + Type::TIMESTAMPTZ => Some(35), + Type::UUID => Some(36), + _ => None, + } +} + +fn postgres_type_signed(data_type: &Type) -> Option { + if matches!( + *data_type, + Type::INT2 | Type::INT4 | Type::INT8 | Type::FLOAT4 | Type::FLOAT8 | Type::NUMERIC + ) { + Some(true) + } else if postgres_unsigned_type(data_type) { + Some(false) + } else { + None + } +} + +fn read_i16(raw: &[u8]) -> Result { + raw.try_into() + .map(i16::from_be_bytes) + .map_err(|_| result_decode_error()) +} + +fn read_u16(raw: &[u8]) -> Result { + raw.try_into() + .map(u16::from_be_bytes) + .map_err(|_| result_decode_error()) +} + +fn read_i32(raw: &[u8]) -> Result { + raw.try_into() + .map(i32::from_be_bytes) + .map_err(|_| result_decode_error()) +} + +fn read_u32(raw: &[u8]) -> Result { + raw.try_into() + .map(u32::from_be_bytes) + .map_err(|_| result_decode_error()) +} + +fn read_i64(raw: &[u8]) -> Result { + raw.try_into() + .map(i64::from_be_bytes) + .map_err(|_| result_decode_error()) +} + +fn read_u64(raw: &[u8]) -> Result { + raw.try_into() + .map(u64::from_be_bytes) + .map_err(|_| result_decode_error()) +} + +fn take_bytes<'a>(raw: &mut &'a [u8], length: usize) -> Result<&'a [u8], AppError> { + let (head, tail) = raw + .split_at_checked(length) + .ok_or_else(result_decode_error)?; + *raw = tail; + Ok(head) +} + +fn take_i32(raw: &mut &[u8]) -> Result { + read_i32(take_bytes(raw, 4)?) +} + +fn take_u32(raw: &mut &[u8]) -> Result { + read_u32(take_bytes(raw, 4)?) +} + +fn ensure_postgres_scalar_size(byte_count: usize) -> Result<(), AppError> { + if byte_count > MAX_SCALAR_BYTES { + return Err(postgres_scalar_too_large(byte_count)); + } + Ok(()) +} + +fn postgres_scalar_too_large(byte_count: usize) -> AppError { + resource_error( + "postgres_scalar_too_large", + format!( + "A PostgreSQL scalar contains {byte_count} bytes; the limit is {MAX_SCALAR_BYTES} bytes" + ), + ) +} + +fn postgres_opaque_value(type_name: &str, raw: &[u8]) -> Result { + use wire::jdbc_value::Value; + ensure_postgres_scalar_size(raw.len())?; + let display_bytes = raw + .len() + .checked_mul(2) + .and_then(|length| length.checked_add(2)) + .ok_or_else(|| postgres_scalar_too_large(raw.len()))?; + ensure_postgres_scalar_size(display_bytes)?; + Ok(Value::OpaqueValue(wire::OpaqueValue { + type_name: type_name.to_owned(), + display_value: format!("\\x{}", hex::encode(raw)), + })) +} + +fn postgres_utf8(raw: &[u8]) -> Result { + ensure_postgres_scalar_size(raw.len())?; + std::str::from_utf8(raw) + .map(str::to_owned) + .map_err(|_| result_decode_error()) +} + +fn decode_postgres_network(data_type: &Type, raw: &[u8]) -> Result { + let [family, prefix_bits, is_cidr, address_length, address @ ..] = raw else { + return Err(result_decode_error()); + }; + let (address, maximum_bits) = match (*family, usize::from(*address_length), address) { + (2, 4, address) => { + let octets: [u8; 4] = address.try_into().map_err(|_| result_decode_error())?; + (Ipv4Addr::from(octets).to_string(), 32_u8) + } + (3, 16, address) => { + let octets: [u8; 16] = address.try_into().map_err(|_| result_decode_error())?; + (Ipv6Addr::from(octets).to_string(), 128_u8) + } + _ => return Err(result_decode_error()), + }; + if *prefix_bits > maximum_bits || *is_cidr > 1 { + return Err(result_decode_error()); + } + if (*data_type == Type::CIDR && *is_cidr != 1) || (*data_type == Type::INET && *is_cidr != 0) { + return Err(result_decode_error()); + } + if *data_type == Type::INET && *prefix_bits == maximum_bits { + Ok(address) + } else { + Ok(format!("{address}/{prefix_bits}")) + } +} + +fn decode_postgres_date(raw: &[u8]) -> Result { + let days = read_i32(raw)?; + match days { + i32::MAX => Ok("infinity".to_owned()), + i32::MIN => Ok("-infinity".to_owned()), + days => postgres_epoch() + .checked_add_signed(TimeDelta::days(i64::from(days))) + .map(|date| date.format("%Y-%m-%d").to_string()) + .ok_or_else(result_decode_error), + } +} + +fn decode_postgres_time(raw: &[u8]) -> Result { + format_time_micros(read_i64(raw)?) +} + +fn decode_postgres_timestamp(raw: &[u8]) -> Result { + let micros = read_i64(raw)?; + match micros { + i64::MAX => Ok("infinity".to_owned()), + i64::MIN => Ok("-infinity".to_owned()), + micros => postgres_epoch() + .and_hms_opt(0, 0, 0) + .and_then(|epoch| epoch.checked_add_signed(TimeDelta::microseconds(micros))) + .map(|value| value.format("%Y-%m-%dT%H:%M:%S%.6f").to_string()) + .ok_or_else(result_decode_error), + } +} + +fn decode_postgres_timetz(raw: &[u8]) -> Result { + let (time, zone) = raw.split_at_checked(8).ok_or_else(result_decode_error)?; + let time = format_time_micros(read_i64(time)?)?; + let seconds_west = read_i32(zone)?; + let seconds_east = seconds_west.checked_neg().ok_or_else(result_decode_error)?; + let sign = if seconds_east < 0 { '-' } else { '+' }; + let absolute = seconds_east.unsigned_abs(); + Ok(format!( + "{time}{sign}{:02}:{:02}", + absolute / 3_600, + (absolute % 3_600) / 60 + )) +} + +fn format_time_micros(micros: i64) -> Result { + if !(0..=86_400_000_000).contains(µs) { + return Err(result_decode_error()); + } + if micros == 86_400_000_000 { + return Ok("24:00:00".to_owned()); + } + let hours = micros / 3_600_000_000; + let minutes = (micros / 60_000_000) % 60; + let seconds = (micros / 1_000_000) % 60; + let fraction = micros % 1_000_000; + if fraction == 0 { + Ok(format!("{hours:02}:{minutes:02}:{seconds:02}")) + } else { + Ok(format!( + "{hours:02}:{minutes:02}:{seconds:02}.{fraction:06}" + )) + } +} + +fn postgres_epoch() -> NaiveDate { + NaiveDate::from_ymd_opt(2000, 1, 1).expect("PostgreSQL epoch is a valid date") +} + +fn decode_postgres_bits(raw: &[u8]) -> Result { + let (length, bytes) = raw.split_at_checked(4).ok_or_else(result_decode_error)?; + let length = usize::try_from(read_i32(length)?).map_err(|_| result_decode_error())?; + if bytes.len() != length.div_ceil(8) || length > MAX_SCALAR_BYTES { + return Err(result_decode_error()); + } + let mut result = String::with_capacity(length); + for index in 0..length { + let byte = bytes[index / 8]; + let bit = (byte >> (7 - index % 8)) & 1; + result.push(if bit == 0 { '0' } else { '1' }); + } + Ok(result) +} + +fn decode_postgres_numeric(raw: &[u8]) -> Result { + if raw.len() < 8 || !raw.len().is_multiple_of(2) { + return Err(result_decode_error()); + } + let digit_count = usize::try_from(read_i16(&raw[0..2])?).map_err(|_| result_decode_error())?; + let weight = i32::from(read_i16(&raw[2..4])?); + let sign = read_u16(&raw[4..6])?; + let scale = usize::from(read_u16(&raw[6..8])?); + if raw.len() != 8 + digit_count.saturating_mul(2) { + return Err(result_decode_error()); + } + match sign { + 0xC000 => return Ok("NaN".to_owned()), + 0xD000 => return Ok("Infinity".to_owned()), + 0xF000 => return Ok("-Infinity".to_owned()), + 0x0000 | 0x4000 => {} + _ => return Err(result_decode_error()), + } + let integer_groups = + usize::try_from(weight.saturating_add(1).max(0)).map_err(|_| result_decode_error())?; + let maximum_display_bytes = integer_groups + .max(1) + .checked_mul(4) + .and_then(|length| length.checked_add(scale)) + .and_then(|length| length.checked_add(2)) + .ok_or_else(|| postgres_scalar_too_large(raw.len()))?; + ensure_postgres_scalar_size(maximum_display_bytes)?; + let digits = raw[8..] + .chunks_exact(2) + .map(read_u16) + .collect::, _>>()?; + if digits.iter().any(|digit| *digit > 9_999) { + return Err(result_decode_error()); + } + let mut integer = String::new(); + if integer_groups == 0 { + integer.push('0'); + } else { + for group in 0..integer_groups { + let digit = digits.get(group).copied().unwrap_or(0); + if group == 0 { + write!(&mut integer, "{digit}").map_err(|_| AppError::internal())?; + } else { + write!(&mut integer, "{digit:04}").map_err(|_| AppError::internal())?; + } + } + } + let mut fraction = String::new(); + if scale > 0 { + let leading_groups = + usize::try_from((-weight - 1).max(0)).map_err(|_| result_decode_error())?; + for _ in 0..leading_groups { + fraction.push_str("0000"); + } + let start = if weight >= 0 { + usize::try_from(weight + 1).map_err(|_| result_decode_error())? + } else { + 0 + }; + for digit in digits.iter().skip(start) { + write!(&mut fraction, "{digit:04}").map_err(|_| AppError::internal())?; + } + while fraction.len() < scale { + fraction.push('0'); + } + fraction.truncate(scale); + } + let negative = sign == 0x4000 && (integer != "0" || fraction.bytes().any(|byte| byte != b'0')); + let normalized_integer = integer.trim_start_matches('0'); + let normalized_integer = if normalized_integer.is_empty() { + "0" + } else { + normalized_integer + }; + Ok(format!( + "{}{}{}", + if negative { "-" } else { "" }, + normalized_integer, + if scale == 0 { + String::new() + } else { + format!(".{fraction}") + } + )) +} + +fn decode_postgres_array(data_type: &Type, raw: &[u8]) -> Result { + let Kind::Array(element_type) = data_type.kind() else { + return Err(result_decode_error()); + }; + let mut cursor = raw; + let dimensions = usize::try_from(take_i32(&mut cursor)?).map_err(|_| result_decode_error())?; + let _has_null = take_i32(&mut cursor)?; + let element_oid = take_u32(&mut cursor)?; + if element_oid != element_type.oid() || dimensions > POSTGRES_ARRAY_MAX_DIMENSIONS { + return Err(result_decode_error()); + } + let mut value_count = 1_usize; + let mut shape = Vec::with_capacity(dimensions); + for _ in 0..dimensions { + let length = usize::try_from(take_i32(&mut cursor)?).map_err(|_| result_decode_error())?; + let lower_bound = take_i32(&mut cursor)?; + value_count = value_count + .checked_mul(length) + .ok_or_else(result_decode_error)?; + shape.push(PostgresArrayDimension { + length, + lower_bound, + }); + } + if dimensions == 0 { + value_count = 0; + } + if value_count > MAX_CONSOLE_PAGE_SIZE as usize * MAX_COLUMNS { + return Err(resource_error( + "postgres_array_too_large", + "The PostgreSQL array exceeds the display element limit", + )); + } + if value_count > cursor.len() / size_of::() { + return Err(result_decode_error()); + } + let mut values = Vec::with_capacity(value_count); + for _ in 0..value_count { + let length = take_i32(&mut cursor)?; + if length == -1 { + values.push(None); + continue; + } + if length < -1 { + return Err(result_decode_error()); + } + let value = take_bytes( + &mut cursor, + usize::try_from(length).map_err(|_| result_decode_error())?, + )?; + values.push(Some(postgres_display_value(decode_postgres_value( + element_type, + value, + )?)?)); + } + if !cursor.is_empty() { + return Err(result_decode_error()); + } + render_postgres_array(&shape, &values) +} + +#[derive(Debug, Clone, Copy)] +struct PostgresArrayDimension { + length: usize, + lower_bound: i32, +} + +fn render_postgres_array( + shape: &[PostgresArrayDimension], + values: &[Option], +) -> Result { + if shape.is_empty() { + if !values.is_empty() { + return Err(result_decode_error()); + } + return Ok("{}".to_owned()); + } + let mut output = String::new(); + if shape.iter().any(|dimension| dimension.lower_bound != 1) { + for dimension in shape { + let upper_bound = if dimension.length == 0 { + dimension.lower_bound.saturating_sub(1) + } else { + dimension + .lower_bound + .checked_add( + i32::try_from(dimension.length - 1).map_err(|_| result_decode_error())?, + ) + .ok_or_else(result_decode_error)? + }; + push_postgres_array_text( + &mut output, + &format!("[{}:{upper_bound}]", dimension.lower_bound), + )?; + } + push_postgres_array_text(&mut output, "=")?; + } + let mut value_index = 0_usize; + render_postgres_array_dimension(&mut output, shape, 0, values, &mut value_index)?; + if value_index != values.len() { + return Err(result_decode_error()); + } + Ok(output) +} + +fn render_postgres_array_dimension( + output: &mut String, + shape: &[PostgresArrayDimension], + depth: usize, + values: &[Option], + value_index: &mut usize, +) -> Result<(), AppError> { + push_postgres_array_text(output, "{")?; + for index in 0..shape[depth].length { + if index > 0 { + push_postgres_array_text(output, ",")?; + } + if depth + 1 == shape.len() { + let value = values.get(*value_index).ok_or_else(result_decode_error)?; + *value_index = value_index.checked_add(1).ok_or_else(AppError::internal)?; + match value { + None => push_postgres_array_text(output, "NULL")?, + Some(value) => push_postgres_array_element(output, value)?, + } + } else { + render_postgres_array_dimension(output, shape, depth + 1, values, value_index)?; + } + } + push_postgres_array_text(output, "}") +} + +fn push_postgres_array_element(output: &mut String, value: &str) -> Result<(), AppError> { + let quote = value.is_empty() + || value.eq_ignore_ascii_case("NULL") + || value.bytes().any(|byte| { + matches!(byte, b',' | b'{' | b'}' | b'"' | b'\\') || byte.is_ascii_whitespace() + }); + if !quote { + return push_postgres_array_text(output, value); + } + push_postgres_array_text(output, "\"")?; + for character in value.chars() { + if matches!(character, '"' | '\\') { + push_postgres_array_text(output, "\\")?; + } + let mut encoded = [0_u8; 4]; + push_postgres_array_text(output, character.encode_utf8(&mut encoded))?; + } + push_postgres_array_text(output, "\"") +} + +fn push_postgres_array_text(output: &mut String, value: &str) -> Result<(), AppError> { + if output.len().saturating_add(value.len()) > MAX_SCALAR_BYTES { + return Err(postgres_scalar_too_large( + output.len().saturating_add(value.len()), + )); + } + output.push_str(value); + Ok(()) +} + +fn postgres_display_value(value: wire::jdbc_value::Value) -> Result { + use wire::jdbc_value::Value; + let display = match value { + Value::NullValue(_) => "NULL".to_owned(), + Value::BooleanValue(value) => value.to_string(), + Value::SignedIntegerValue(value) => value.to_string(), + Value::UnsignedIntegerValue(value) => value.to_string(), + Value::Float32Value(value) => value.to_string(), + Value::Float64Value(value) => value.to_string(), + Value::DecimalValue(value) + | Value::TextValue(value) + | Value::DateValue(value) + | Value::TimeValue(value) + | Value::TimestampValue(value) + | Value::TimestampWithTimeZoneValue(value) + | Value::JsonValue(value) + | Value::UuidValue(value) => value, + Value::BinaryValue(value) => { + let display_bytes = value + .len() + .checked_mul(2) + .and_then(|length| length.checked_add(2)) + .ok_or_else(|| postgres_scalar_too_large(value.len()))?; + ensure_postgres_scalar_size(display_bytes)?; + format!("\\x{}", hex::encode(value)) + } + Value::OpaqueValue(value) => value.display_value, + }; + ensure_postgres_scalar_size(display.len())?; + Ok(display) +} + +fn result_decode_error() -> AppError { + AppError::unavailable( + "postgres_result_decode_failed", + "A PostgreSQL result value could not be decoded safely", + ) +} + +fn resource_error(code: impl Into, message: impl Into) -> AppError { + AppError::new( + AppErrorKind::ResourceExhausted, + ApiError::new(code, message), + ) +} + +#[allow( + clippy::too_many_lines, + reason = "the retained-query lifecycle keeps dispatch, streaming limits, persistence, and cleanup visible in one transaction" +)] +async fn execute_query_task( + application: &Application, + operation_id: &str, + mut cancellation: watch::Receiver, + query: PreparedQuery, + storage: Storage, + resolved: ResolvedDatasourceConnection, +) -> Result { + if let CancellationRequest::Requested { reason } = cancellation.borrow().clone() { + return Err(QueryTaskError::Cancelled(reason)); + } + validate_query(&query)?; + let parameters = postgres_query_parameters(&query.parameters)?; + let connection = open_query_connection(&resolved, None, &mut cancellation).await?; + if let Err(error) = connection + .client() + .batch_execute("BEGIN TRANSACTION READ ONLY") + .await + .map_err(postgres_query_error) + { + connection.abort().await; + return Err(error.into()); + } + + let statement = match cancellable_postgres( + connection.client().prepare(&query.sql), + &mut cancellation, + ) + .await + { + Ok(statement) => statement, + Err(PostgresCancellableError::Cancelled(reason)) => { + connection.abort().await; + return Err(QueryTaskError::Cancelled(reason)); + } + Err(PostgresCancellableError::Failed(error)) => { + connection.abort().await; + return Err(error.into()); + } + }; + if statement.params().len() != parameters.len() { + connection.abort().await; + return Err(AppError::invalid( + "invalid_query_parameter_count", + format!( + "The PostgreSQL statement expects {} parameters but {} were supplied", + statement.params().len(), + parameters.len() + ), + ) + .into()); + } + let columns = statement.columns(); + if columns.len() > MAX_COLUMNS { + connection.abort().await; + return Err(resource_error( + "postgres_result_too_wide", + format!("PostgreSQL returned more than {MAX_COLUMNS} columns"), + ) + .into()); + } + let schema = wire::QueryStarted { + columns: columns + .iter() + .enumerate() + .map(|(index, column)| postgres_column(index, column)) + .collect::>()?, + }; + let mut writer = RetainedWriter::begin(storage, schema, query.retention).await?; + if let Err(error) = application.inner.operations.started(operation_id).await { + abort_writer(&mut writer).await; + connection.abort().await; + return Err(error.into()); + } + let parameter_refs = parameters + .iter() + .map(|parameter| parameter as &(dyn ToSql + Sync)) + .collect::>(); + let max_rows = query.options.max_rows; + let max_result_bytes = if query.options.max_result_bytes == 0 { + DEFAULT_RESULT_BYTES + } else { + query.options.max_result_bytes + }; + let batch_rows = if query.options.target_batch_rows == 0 { + DEFAULT_BATCH_ROWS + } else { + query.options.target_batch_rows + }; + let batch_bytes = if query.options.target_batch_bytes == 0 { + DEFAULT_BATCH_BYTES + } else { + query.options.target_batch_bytes + }; + let consumption = async { + let stream = match cancellable_postgres( + connection.client().query_raw(&statement, parameter_refs), + &mut cancellation, + ) + .await + { + Ok(stream) => stream, + Err(PostgresCancellableError::Cancelled(reason)) => { + return Err(QueryTaskError::Cancelled(reason)); + } + Err(PostgresCancellableError::Failed(error)) => return Err(error.into()), + }; + tokio::pin!(stream); + let mut pending_rows = Vec::new(); + let mut pending_bytes = 0_u64; + let mut row_count = 0_u64; + let mut result_bytes = 0_u64; + let mut truncated_rows = false; + let mut truncated_bytes = false; + let mut cancellation_open = true; + loop { + let next = tokio::select! { + biased; + changed = cancellation.changed(), if cancellation_open => { + if let Ok(()) = changed { + if let CancellationRequest::Requested { reason } = cancellation.borrow().clone() { + return Err(QueryTaskError::Cancelled(reason)); + } + } else { + cancellation_open = false; + } + continue; + } + row = stream.next() => row, + }; + let Some(row) = next else { + break; + }; + let row = row.map_err(postgres_query_error)?; + if max_rows != 0 && row_count >= max_rows { + truncated_rows = true; + break; + } + let row = postgres_row(&row, columns)?; + let row_bytes = u64::try_from(row.encoded_len()).map_err(|_| AppError::internal())?; + if result_bytes.saturating_add(row_bytes) > max_result_bytes { + truncated_bytes = true; + break; + } + let entry_bytes = row_batch_entry_bytes(&row)?; + let candidate_bytes = pending_bytes + .saturating_add(if pending_rows.is_empty() { + row_batch_prefix_bytes(row_count) + } else { + 0 + }) + .saturating_add(entry_bytes); + if !pending_rows.is_empty() + && (pending_rows.len() >= usize::try_from(batch_rows).unwrap_or(usize::MAX) + || candidate_bytes > u64::from(batch_bytes)) + { + flush_rows( + application, + operation_id, + &mut writer, + &mut pending_rows, + row_count, + ) + .await?; + pending_bytes = 0; + } + if pending_rows.is_empty() { + pending_bytes = row_batch_prefix_bytes(row_count); + } + pending_rows.push(row); + pending_bytes = pending_bytes.saturating_add(entry_bytes); + row_count = row_count.checked_add(1).ok_or_else(AppError::internal)?; + result_bytes = result_bytes + .checked_add(row_bytes) + .ok_or_else(AppError::internal)?; + } + flush_rows( + application, + operation_id, + &mut writer, + &mut pending_rows, + row_count, + ) + .await?; + Ok::<_, QueryTaskError>((row_count, truncated_rows, truncated_bytes)) + } + .await; + let (row_count, truncated_rows, truncated_bytes) = match consumption { + Ok(outcome) => outcome, + Err(error) => { + abort_writer(&mut writer).await; + connection.abort().await; + return Err(error); + } + }; + let metadata = match writer + .finish(wire::QueryCompleted { + row_count, + truncated_by_max_rows: truncated_rows, + truncated_by_max_result_bytes: truncated_bytes, + }) + .await + { + Ok(metadata) => metadata, + Err(error) => { + abort_writer(&mut writer).await; + connection.abort().await; + return Err(error.into()); + } + }; + if truncated_rows || truncated_bytes { + connection.abort().await; + } else { + let rollback = connection.client().batch_execute("ROLLBACK").await; + let result = rollback.map_err(postgres_query_error).map(|()| metadata); + return finish_connection(connection, result) + .await + .map_err(QueryTaskError::from); + } + Ok(metadata) +} + +async fn open_query_connection( + resolved: &ResolvedDatasourceConnection, + database_name: Option<&str>, + cancellation: &mut watch::Receiver, +) -> Result { + let open = open_resolved_connection(resolved, database_name); + tokio::pin!(open); + let mut cancellation_open = true; + loop { + tokio::select! { + biased; + changed = cancellation.changed(), if cancellation_open => { + if let Ok(()) = changed { + if let CancellationRequest::Requested { reason } = cancellation.borrow().clone() { + return Err(QueryTaskError::Cancelled(reason)); + } + } else { + cancellation_open = false; + } + } + result = &mut open => return result.map_err(QueryTaskError::from), + } + } +} + +enum PostgresCancellableError { + Cancelled(Option), + Failed(AppError), +} + +async fn cancellable_postgres( + future: F, + cancellation: &mut watch::Receiver, +) -> Result +where + F: Future>, +{ + tokio::pin!(future); + let mut cancellation_open = true; + loop { + tokio::select! { + biased; + changed = cancellation.changed(), if cancellation_open => { + if let Ok(()) = changed { + if let CancellationRequest::Requested { reason } = cancellation.borrow().clone() { + return Err(PostgresCancellableError::Cancelled(reason)); + } + } else { + cancellation_open = false; + } + } + result = &mut future => { + return result + .map_err(postgres_query_error) + .map_err(PostgresCancellableError::Failed); + } + } + } +} + +fn row_batch_prefix_bytes(start_row_offset: u64) -> u64 { + if start_row_offset == 0 { + 0 + } else { + 1_u64.saturating_add( + u64::try_from(prost::encoding::encoded_len_varint(start_row_offset)) + .unwrap_or(u64::MAX), + ) + } +} + +fn row_batch_entry_bytes(row: &wire::JdbcRow) -> Result { + let row_bytes = row.encoded_len(); + let length_bytes = prost::encoding::length_delimiter_len(row_bytes); + u64::try_from( + 1_usize + .saturating_add(length_bytes) + .saturating_add(row_bytes), + ) + .map_err(|_| QueryTaskError::Failed(AppError::internal())) +} + +async fn flush_rows( + application: &Application, + operation_id: &str, + writer: &mut RetainedWriter, + rows: &mut Vec, + row_count: u64, +) -> Result<(), QueryTaskError> { + if rows.is_empty() { + return Ok(()); + } + let row_len = u64::try_from(rows.len()).map_err(|_| AppError::internal())?; + let start_row_offset = row_count + .checked_sub(row_len) + .ok_or_else(AppError::internal)?; + let batch = wire::RowBatch { + start_row_offset, + rows: std::mem::take(rows), + }; + if batch.encoded_len() > usize::try_from(MAX_BATCH_BYTES).unwrap_or(usize::MAX) { + return Err(resource_error( + "postgres_result_batch_too_large", + "One PostgreSQL result row exceeds the retained-result batch limit", + ) + .into()); + } + let byte_count = writer.append(batch).await?; + application + .inner + .operations + .progress(operation_id, row_count, byte_count) + .await?; + Ok(()) +} + +async fn abort_writer(writer: &mut RetainedWriter) { + if let Err(error) = writer.abort().await { + tracing::warn!(error = %error, "native PostgreSQL retained-result cleanup failed"); + } +} + +async fn execute_update( + resolved: ResolvedDatasourceConnection, + sql: String, + cancellation: CancellationToken, +) -> Result { + if cancellation.is_cancelled() { + return Err(DatabaseWriteError::not_started(postgres_write_cancelled())); + } + let sql = validate_single_write_sql(&sql).map_err(DatabaseWriteError::not_started)?; + if resolved.connection.read_only { + return Err(DatabaseWriteError::not_started(AppError::new( + AppErrorKind::Conflict, + ApiError::new( + "datasource_read_only", + "The datasource connection is configured as read-only", + ), + ))); + } + let connection = tokio::select! { + biased; + () = cancellation.cancelled() => { + return Err(DatabaseWriteError::not_started(postgres_write_cancelled())); + } + result = open_resolved_connection(&resolved, None) => { + result.map_err(DatabaseWriteError::not_started)? + } + }; + let statement = tokio::select! { + biased; + () = cancellation.cancelled() => { + connection.abort().await; + return Err(DatabaseWriteError::not_started(postgres_write_cancelled())); + } + result = connection.client().prepare(&sql) => { + result.map_err(postgres_query_error).map_err(DatabaseWriteError::not_started)? + } + }; + if !statement.params().is_empty() { + connection.abort().await; + return Err(DatabaseWriteError::not_started(AppError::invalid( + "invalid_database_write", + "The confirmed PostgreSQL write does not accept unbound parameters", + ))); + } + let result = tokio::select! { + biased; + () = cancellation.cancelled() => None, + result = connection.client().execute(&statement, &[]) => Some(result), + }; + let Some(result) = result else { + connection.abort().await; + return Err(DatabaseWriteError::unknown(AppError::unavailable( + "database_write_outcome_unknown", + "The PostgreSQL write was interrupted after dispatch; do not retry it blindly", + ))); + }; + match result { + Ok(affected_rows) => finish_connection(connection, Ok(affected_rows)) + .await + .map_err(DatabaseWriteError::unknown), + Err(error) => { + connection.abort().await; + tracing::warn!(error = %error, "PostgreSQL rejected a dispatched write"); + Err(DatabaseWriteError::unknown(AppError::unavailable( + "database_write_outcome_unknown", + "PostgreSQL reported an error after write dispatch; do not retry it blindly", + ))) + } + } +} + +fn validate_single_write_sql(sql: &str) -> Result { + if sql.len() > MAX_SQL_BYTES { + return Err(AppError::invalid( + "invalid_database_write", + format!("SQL cannot exceed {MAX_SQL_BYTES} UTF-8 bytes"), + )); + } + let mut statements = split_postgres_script(sql)?; + if statements.len() != 1 { + return Err(AppError::invalid( + "invalid_database_write", + "Exactly one PostgreSQL write statement is required", + )); + } + let statement = statements.pop().expect("statement length checked"); + let words = postgres_words(&statement)?; + if !matches!( + words.first().map(String::as_str), + Some( + "INSERT" + | "UPDATE" + | "DELETE" + | "MERGE" + | "CREATE" + | "ALTER" + | "DROP" + | "TRUNCATE" + | "GRANT" + | "REVOKE" + | "COMMENT" + | "ANALYZE" + | "VACUUM" + | "CALL" + | "DO" + | "REFRESH" + | "REINDEX" + | "CLUSTER" + ) + ) { + return Err(AppError::invalid( + "database_write_statement_required", + "The confirmed PostgreSQL write surface accepts one DML, DDL, grant, maintenance, or routine statement", + )); + } + Ok(statement) +} + +fn postgres_write_cancelled() -> AppError { + AppError::new( + AppErrorKind::Conflict, + ApiError::new( + "database_write_cancelled", + "The PostgreSQL write was cancelled before dispatch", + ), + ) +} + +struct ConsoleStatementExecution { + result: Option, + failure: Option, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum ConsoleStatementKind { + ReadOnly, + Write, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum ConsoleFailure { + Cancelled, + TimedOut, + ResultProcessing, + Driver { server_rejected: bool }, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +struct ConsoleDispatchState { + statement_kind: ConsoleStatementKind, + dispatched: bool, +} + +impl ConsoleDispatchState { + fn classify(sql: &str) -> Self { + Self { + statement_kind: if validate_read_sql(sql).is_ok() { + ConsoleStatementKind::ReadOnly + } else { + ConsoleStatementKind::Write + }, + dispatched: false, + } + } + + fn mark_dispatched(&mut self) { + self.dispatched = true; + } + + fn requires_unknown_outcome(self, failure: ConsoleFailure) -> bool { + self.statement_kind == ConsoleStatementKind::Write + && self.dispatched + && !matches!( + failure, + ConsoleFailure::Driver { + server_rejected: true + } + ) + } +} + +enum ConsoleExecutionError { + Cancelled(Option), + Fatal(AppError), +} + +enum ConsolePostgresError { + Cancelled(Option), + Failed(PostgresError), +} + +async fn execute_console( + application: &Application, + request: NativeConsoleRequest, + mut cancellation: watch::Receiver, + force_read_only: bool, +) -> Result, AppError> { + let (statements, page_offset, page_end) = prepare_console_statements(&request)?; + if let CancellationRequest::Requested { reason } = cancellation.borrow().clone() { + return Err(postgres_console_cancelled(reason)); + } + let resolved = resolve_native_connection(application, &request.datasource_id).await?; + let read_only = force_read_only || resolved.connection.read_only; + if read_only { + validate_read_only_console(&statements, force_read_only)?; + } + let target_database = + (!request.database_name.trim().is_empty()).then_some(request.database_name.as_str()); + let connection = match open_query_connection(&resolved, target_database, &mut cancellation) + .await + { + Ok(connection) => connection, + Err(QueryTaskError::Cancelled(reason)) => return Err(postgres_console_cancelled(reason)), + Err(QueryTaskError::Failed(error)) => return Err(error), + }; + if read_only && let Err(error) = begin_read_only_console_transaction(&connection).await { + connection.abort().await; + return Err(error); + } + + let mut results = Vec::new(); + let mut retained_result_bytes = 0_u64; + let mut dispatched_write = false; + for (index, statement) in statements.into_iter().enumerate() { + let mut dispatch = ConsoleDispatchState::classify(&statement); + let statement_sequence = u32::try_from(index) + .ok() + .and_then(|index| index.checked_add(1)) + .ok_or_else(AppError::internal)?; + let started = Instant::now(); + let Ok(execution) = tokio::time::timeout( + CONSOLE_STATEMENT_TIMEOUT, + execute_console_statement( + &connection, + &statement, + statement_sequence, + page_offset, + page_end, + request.result_set_id, + &mut retained_result_bytes, + &mut cancellation, + &mut dispatch, + ), + ) + .await + else { + connection.abort().await; + return Err(postgres_console_timeout(dispatch)); + }; + dispatched_write |= dispatch.requires_unknown_outcome(ConsoleFailure::Driver { + server_rejected: false, + }); + match execution { + Ok(execution) => { + if let Some(result) = execution.result { + results.push(result); + } + if let Some(error) = execution.failure { + results.push(console_failure_result( + statement_sequence, + statement, + &error, + elapsed_millis(started), + )); + if !request.error_continue { + break; + } + } + } + Err(ConsoleExecutionError::Cancelled(reason)) => { + connection.abort().await; + return Err(postgres_console_interrupted(dispatch, reason)); + } + Err(ConsoleExecutionError::Fatal(error)) => { + connection.abort().await; + return Err(postgres_console_fatal(dispatch, error)); + } + } + } + if read_only && let Err(error) = rollback_read_only_console_transaction(&connection).await { + connection.abort().await; + return Err(error); + } + finish_console_connection(connection, results, dispatched_write).await +} + +fn postgres_console_timeout(dispatch: ConsoleDispatchState) -> AppError { + if dispatch.requires_unknown_outcome(ConsoleFailure::TimedOut) { + postgres_console_write_outcome_unknown(ConsoleFailure::TimedOut) + } else { + postgres_console_statement_timeout() + } +} + +fn postgres_console_interrupted( + dispatch: ConsoleDispatchState, + reason: Option, +) -> AppError { + if dispatch.requires_unknown_outcome(ConsoleFailure::Cancelled) { + postgres_console_write_outcome_unknown(ConsoleFailure::Cancelled) + } else { + postgres_console_cancelled(reason) + } +} + +fn postgres_console_fatal(dispatch: ConsoleDispatchState, error: AppError) -> AppError { + if dispatch.requires_unknown_outcome(ConsoleFailure::ResultProcessing) { + postgres_console_write_outcome_unknown(ConsoleFailure::ResultProcessing) + } else { + error + } +} + +async fn finish_console_connection( + connection: ManagedPostgresConnection, + results: Vec, + dispatched_write: bool, +) -> Result, AppError> { + match finish_connection(connection, Ok(results)).await { + Err(error) if dispatched_write => { + tracing::warn!(error = %error, "PostgreSQL connection failed after a Console write dispatch"); + Err(postgres_console_write_outcome_unknown( + ConsoleFailure::Driver { + server_rejected: false, + }, + )) + } + result => result, + } +} + +async fn begin_read_only_console_transaction( + connection: &ManagedPostgresConnection, +) -> Result<(), AppError> { + match tokio::time::timeout( + CONSOLE_STATEMENT_TIMEOUT, + connection + .client() + .batch_execute("BEGIN TRANSACTION READ ONLY"), + ) + .await + { + Ok(result) => result.map_err(postgres_query_error), + Err(_) => Err(postgres_console_statement_timeout()), + } +} + +async fn rollback_read_only_console_transaction( + connection: &ManagedPostgresConnection, +) -> Result<(), AppError> { + match tokio::time::timeout( + DISCONNECT_TIMEOUT, + connection.client().batch_execute("ROLLBACK"), + ) + .await + { + Ok(result) => result.map_err(postgres_query_error), + Err(_) => Err(AppError::unavailable( + "postgres_console_rollback_timeout", + "The PostgreSQL read-only Console transaction did not roll back in time", + )), + } +} + +fn postgres_console_statement_timeout() -> AppError { + AppError::unavailable( + "postgres_console_statement_timeout", + format!( + "A PostgreSQL Console statement exceeded the {} second execution limit", + CONSOLE_STATEMENT_TIMEOUT.as_secs() + ), + ) +} + +fn postgres_console_write_outcome_unknown(failure: ConsoleFailure) -> AppError { + let message = match failure { + ConsoleFailure::Cancelled => { + "The PostgreSQL Console write was cancelled after dispatch; its outcome is unknown, so do not retry it blindly" + } + ConsoleFailure::TimedOut => { + "The PostgreSQL Console write timed out after dispatch; its outcome is unknown, so do not retry it blindly" + } + ConsoleFailure::ResultProcessing => { + "The PostgreSQL Console could not finish processing a write result after dispatch; its outcome is unknown, so do not retry it blindly" + } + ConsoleFailure::Driver { .. } => { + "The PostgreSQL Console connection failed after write dispatch; its outcome is unknown, so do not retry it blindly" + } + }; + AppError::new( + AppErrorKind::Unavailable, + ApiError::new("database_write_outcome_unknown", message), + ) +} + +fn prepare_console_statements( + request: &NativeConsoleRequest, +) -> Result<(Vec, u64, u64), AppError> { + let (page_offset, page_end) = validate_console_request(request)?; + let mut statements = if request.single { + vec![request.sql.trim().to_owned()] + } else { + split_postgres_script(&request.sql)? + }; + if statements.is_empty() { + return Err(AppError::invalid( + "invalid_postgres_console_request", + "sql must contain at least one PostgreSQL statement", + )); + } + if request.explain { + for statement in &mut statements { + *statement = format!("EXPLAIN {statement}"); + } + } + Ok((statements, page_offset, page_end)) +} + +fn validate_console_request(request: &NativeConsoleRequest) -> Result<(u64, u64), AppError> { + if request.datasource_id.trim().is_empty() || request.sql.trim().is_empty() { + return Err(AppError::invalid( + "invalid_postgres_console_request", + "dataSourceId and sql cannot be empty", + )); + } + if request.sql.len() > MAX_SQL_BYTES { + return Err(resource_error( + "postgres_console_script_too_large", + format!("PostgreSQL Console scripts are limited to {MAX_SQL_BYTES} bytes"), + )); + } + if !request.database_name.trim().is_empty() { + validate_identifier(&request.database_name, "databaseName")?; + } + if request.page_no == 0 || request.page_size == 0 || request.page_size > MAX_CONSOLE_PAGE_SIZE { + return Err(AppError::invalid( + "invalid_postgres_console_request", + format!( + "pageNo must be positive and pageSize must be between 1 and {MAX_CONSOLE_PAGE_SIZE}" + ), + )); + } + if request.result_set_id == Some(0) { + return Err(AppError::invalid( + "invalid_postgres_console_request", + "resultSetId must be a positive one-based integer", + )); + } + if request.page_size_all { + Ok((0, u64::from(MAX_CONSOLE_PAGE_SIZE))) + } else { + let offset = u64::from(request.page_no - 1) * u64::from(request.page_size); + let end = offset + .checked_add(u64::from(request.page_size)) + .ok_or_else(AppError::internal)?; + Ok((offset, end)) + } +} + +fn validate_forced_read_console(statements: &[String]) -> Result<(), AppError> { + for statement in statements { + validate_read_sql(statement).map_err(|_| { + AppError::invalid( + "chart_query_must_be_read_only", + "Chart refresh accepts only PostgreSQL read statements without writes or row locks", + ) + })?; + } + Ok(()) +} + +fn validate_read_only_console( + statements: &[String], + force_read_only: bool, +) -> Result<(), AppError> { + if force_read_only { + return validate_forced_read_console(statements); + } + for statement in statements { + validate_read_sql(statement).map_err(|_| { + AppError::invalid( + "postgres_console_must_be_read_only", + "This PostgreSQL datasource accepts only read statements without writes or row locks", + ) + })?; + } + Ok(()) +} + +#[allow(clippy::too_many_arguments, clippy::too_many_lines)] +async fn execute_console_statement( + connection: &ManagedPostgresConnection, + sql: &str, + statement_sequence: u32, + page_offset: u64, + page_end: u64, + selected_result_set_id: Option, + retained_result_bytes: &mut u64, + cancellation: &mut watch::Receiver, + dispatch: &mut ConsoleDispatchState, +) -> Result { + let started = Instant::now(); + let statement = + match cancellable_console_postgres(connection.client().prepare(sql), cancellation).await { + Ok(statement) => statement, + Err(ConsolePostgresError::Cancelled(reason)) => { + return Err(ConsoleExecutionError::Cancelled(reason)); + } + Err(ConsolePostgresError::Failed(error)) => { + return console_postgres_failure(*dispatch, error); + } + }; + if !statement.params().is_empty() { + return Ok(ConsoleStatementExecution { + result: None, + failure: Some(AppError::invalid( + "invalid_postgres_console_request", + "PostgreSQL Console statements cannot contain unbound parameters", + )), + }); + } + if statement.columns().is_empty() { + dispatch.mark_dispatched(); + let update_count = match cancellable_console_postgres( + connection.client().execute(&statement, &[]), + cancellation, + ) + .await + { + Ok(count) => count, + Err(ConsolePostgresError::Cancelled(reason)) => { + return Err(ConsoleExecutionError::Cancelled(reason)); + } + Err(ConsolePostgresError::Failed(error)) => { + return console_postgres_failure(*dispatch, error); + } + }; + return Ok(ConsoleStatementExecution { + result: Some(NativeConsoleResult { + statement_sequence, + result_set_id: None, + sql: sql.to_owned(), + success: true, + message: "Statement executed successfully".to_owned(), + update_count, + columns: Vec::new(), + rows: Vec::new(), + row_count: 0, + has_more: false, + duration_ms: elapsed_millis(started), + error: None, + }), + failure: None, + }); + } + let retain = selected_result_set_id.is_none_or(|selected| selected == 1); + let columns = statement.columns(); + if columns.len() > MAX_COLUMNS { + return Err(ConsoleExecutionError::Fatal(resource_error( + "postgres_result_too_wide", + format!("PostgreSQL returned more than {MAX_COLUMNS} columns"), + ))); + } + let converted_columns = if retain { + columns + .iter() + .enumerate() + .map(|(index, column)| console_column(index, column)) + .collect::, _>>() + .map_err(ConsoleExecutionError::Fatal)? + } else { + Vec::new() + }; + dispatch.mark_dispatched(); + let stream = match cancellable_console_postgres( + connection + .client() + .query_raw(&statement, std::iter::empty::<&(dyn ToSql + Sync)>()), + cancellation, + ) + .await + { + Ok(stream) => stream, + Err(ConsolePostgresError::Cancelled(reason)) => { + return Err(ConsoleExecutionError::Cancelled(reason)); + } + Err(ConsolePostgresError::Failed(error)) => { + return console_postgres_failure(*dispatch, error); + } + }; + tokio::pin!(stream); + let mut rows = Vec::new(); + let mut row_count = 0_u64; + let mut scanned_bytes = 0_u64; + let mut cancellation_open = true; + loop { + let next = tokio::select! { + biased; + changed = cancellation.changed(), if cancellation_open => { + if let Ok(()) = changed { + if let CancellationRequest::Requested { reason } = cancellation.borrow().clone() { + return Err(ConsoleExecutionError::Cancelled(reason)); + } + } else { + cancellation_open = false; + } + continue; + } + row = stream.next() => row, + }; + let Some(row) = next else { + break; + }; + let row = match row { + Ok(row) => row, + Err(error) => { + return console_postgres_failure(*dispatch, error); + } + }; + let wire_row = postgres_row(&row, columns).map_err(ConsoleExecutionError::Fatal)?; + let next_row_count = row_count + .checked_add(1) + .ok_or_else(|| ConsoleExecutionError::Fatal(AppError::internal()))?; + if next_row_count > MAX_CONSOLE_SCANNED_ROWS { + return Err(ConsoleExecutionError::Fatal(resource_error( + "postgres_console_scan_row_limit_exceeded", + format!( + "A PostgreSQL Console result exceeds the {MAX_CONSOLE_SCANNED_ROWS} row scan limit" + ), + ))); + } + let row_bytes = u64::try_from(wire_row.encoded_len()) + .map_err(|_| ConsoleExecutionError::Fatal(AppError::internal()))?; + let next_scanned_bytes = scanned_bytes + .checked_add(row_bytes) + .ok_or_else(|| ConsoleExecutionError::Fatal(AppError::internal()))?; + if next_scanned_bytes > MAX_CONSOLE_SCANNED_BYTES { + return Err(ConsoleExecutionError::Fatal(resource_error( + "postgres_console_scan_byte_limit_exceeded", + format!( + "A PostgreSQL Console result exceeds the {MAX_CONSOLE_SCANNED_BYTES} byte scan limit" + ), + ))); + } + if retain && (page_offset..page_end).contains(&row_count) { + let retained_row = + console_row_from_wire(wire_row).map_err(ConsoleExecutionError::Fatal)?; + reserve_console_result_bytes(retained_result_bytes, &retained_row) + .map_err(ConsoleExecutionError::Fatal)?; + rows.push(retained_row); + } + row_count = next_row_count; + scanned_bytes = next_scanned_bytes; + } + Ok(ConsoleStatementExecution { + result: retain.then_some(NativeConsoleResult { + statement_sequence, + result_set_id: Some(1), + sql: sql.to_owned(), + success: true, + message: "Statement executed successfully".to_owned(), + update_count: 0, + columns: converted_columns, + rows, + row_count, + has_more: row_count > page_end, + duration_ms: elapsed_millis(started), + error: None, + }), + failure: None, + }) +} + +async fn cancellable_console_postgres( + future: F, + cancellation: &mut watch::Receiver, +) -> Result +where + F: Future>, +{ + tokio::pin!(future); + let mut cancellation_open = true; + loop { + tokio::select! { + biased; + changed = cancellation.changed(), if cancellation_open => { + if let Ok(()) = changed { + if let CancellationRequest::Requested { reason } = cancellation.borrow().clone() { + return Err(ConsolePostgresError::Cancelled(reason)); + } + } else { + cancellation_open = false; + } + } + result = &mut future => return result.map_err(ConsolePostgresError::Failed), + } + } +} + +fn console_postgres_failure( + dispatch: ConsoleDispatchState, + error: PostgresError, +) -> Result { + let failure = ConsoleFailure::Driver { + server_rejected: error.as_db_error().is_some(), + }; + if dispatch.requires_unknown_outcome(failure) { + return Err(ConsoleExecutionError::Fatal( + postgres_console_write_outcome_unknown(failure), + )); + } + Ok(ConsoleStatementExecution { + result: None, + failure: Some(postgres_query_error(error)), + }) +} + +fn console_column(index: usize, column: &tokio_postgres::Column) -> Result { + let column = postgres_column(index, column)?; + let value_type = match wire::JdbcValueType::try_from(column.value_type) { + Ok(wire::JdbcValueType::Boolean) => JdbcValueType::Boolean, + Ok(wire::JdbcValueType::SignedInteger) => JdbcValueType::SignedInteger, + Ok(wire::JdbcValueType::UnsignedInteger) => JdbcValueType::UnsignedInteger, + Ok(wire::JdbcValueType::Float32) => JdbcValueType::Float32, + Ok(wire::JdbcValueType::Float64) => JdbcValueType::Float64, + Ok(wire::JdbcValueType::Decimal) => JdbcValueType::Decimal, + Ok(wire::JdbcValueType::Text) => JdbcValueType::Text, + Ok(wire::JdbcValueType::Binary) => JdbcValueType::Binary, + Ok(wire::JdbcValueType::Date) => JdbcValueType::Date, + Ok(wire::JdbcValueType::Time) => JdbcValueType::Time, + Ok(wire::JdbcValueType::Timestamp) => JdbcValueType::Timestamp, + Ok(wire::JdbcValueType::TimestampWithTimeZone) => JdbcValueType::TimestampWithTimeZone, + Ok(wire::JdbcValueType::Json) => JdbcValueType::Json, + Ok(wire::JdbcValueType::Uuid) => JdbcValueType::Uuid, + Ok(wire::JdbcValueType::Opaque) => JdbcValueType::Opaque, + Ok(wire::JdbcValueType::Unspecified) | Err(_) => return Err(AppError::internal()), + }; + Ok(ResultColumn { + ordinal: column.ordinal, + label: column.label, + name: column.name, + jdbc_type: column.jdbc_type, + jdbc_type_name: column.jdbc_type_name, + value_type, + nullability: ColumnNullability::Unknown, + precision: column.precision, + scale: column.scale, + display_size: column.display_size, + signed: column.signed, + catalog_name: column.catalog_name, + schema_name: column.schema_name, + table_name: column.table_name, + }) +} + +fn console_row_from_wire(wire: wire::JdbcRow) -> Result { + Ok(ResultRow { + values: wire + .values + .into_iter() + .map(console_value) + .collect::>()?, + }) +} + +fn console_value(value: wire::JdbcValue) -> Result { + use wire::jdbc_value::Value; + Ok(match value.value.ok_or_else(AppError::internal)? { + Value::NullValue(_) => JdbcValue::Null, + Value::BooleanValue(value) => JdbcValue::Boolean { value }, + Value::SignedIntegerValue(value) => JdbcValue::SignedInteger { + value: value.to_string(), + }, + Value::UnsignedIntegerValue(value) => JdbcValue::UnsignedInteger { + value: value.to_string(), + }, + Value::Float32Value(value) => JdbcValue::Float32 { + value: display_float32(value), + }, + Value::Float64Value(value) => JdbcValue::Float64 { + value: display_float64(value), + }, + Value::DecimalValue(value) => JdbcValue::Decimal { value }, + Value::TextValue(value) => JdbcValue::Text { value }, + Value::BinaryValue(value) => JdbcValue::Binary { + value: BASE64_STANDARD.encode(value), + }, + Value::DateValue(value) => JdbcValue::Date { value }, + Value::TimeValue(value) => JdbcValue::Time { value }, + Value::TimestampValue(value) => JdbcValue::Timestamp { value }, + Value::TimestampWithTimeZoneValue(value) => JdbcValue::TimestampWithTimeZone { value }, + Value::JsonValue(value) => JdbcValue::Json { value }, + Value::UuidValue(value) => JdbcValue::Uuid { value }, + Value::OpaqueValue(value) => JdbcValue::Opaque { + type_name: value.type_name, + display_value: value.display_value, + }, + }) +} + +fn reserve_console_result_bytes(total: &mut u64, row: &ResultRow) -> Result<(), AppError> { + let next = total.saturating_add(console_row_retained_bytes(row)); + if next > MAX_CONSOLE_RESULT_BYTES { + return Err(resource_error( + "postgres_console_result_too_large", + format!( + "PostgreSQL Console results are limited to {MAX_CONSOLE_RESULT_BYTES} retained bytes" + ), + )); + } + *total = next; + Ok(()) +} + +fn console_row_retained_bytes(row: &ResultRow) -> u64 { + let mut bytes = u64::try_from(size_of::()).unwrap_or(u64::MAX); + bytes = bytes.saturating_add( + u64::try_from(row.values.capacity()) + .unwrap_or(u64::MAX) + .saturating_mul(u64::try_from(size_of::()).unwrap_or(u64::MAX)), + ); + for value in &row.values { + let value_bytes = match value { + JdbcValue::Null | JdbcValue::Boolean { .. } => 0, + JdbcValue::SignedInteger { value } + | JdbcValue::UnsignedInteger { value } + | JdbcValue::Float32 { value } + | JdbcValue::Float64 { value } + | JdbcValue::Decimal { value } + | JdbcValue::Text { value } + | JdbcValue::Binary { value } + | JdbcValue::Date { value } + | JdbcValue::Time { value } + | JdbcValue::Timestamp { value } + | JdbcValue::TimestampWithTimeZone { value } + | JdbcValue::Json { value } + | JdbcValue::Uuid { value } => value.capacity(), + JdbcValue::Opaque { + type_name, + display_value, + } => type_name + .capacity() + .saturating_add(display_value.capacity()), + }; + bytes = bytes.saturating_add(u64::try_from(value_bytes).unwrap_or(u64::MAX)); + } + bytes +} + +fn console_failure_result( + statement_sequence: u32, + sql: String, + error: &AppError, + duration_ms: u64, +) -> NativeConsoleResult { + let api_error = error.api_error(); + NativeConsoleResult { + statement_sequence, + result_set_id: None, + sql, + success: false, + message: api_error.message.clone(), + update_count: 0, + columns: Vec::new(), + rows: Vec::new(), + row_count: 0, + has_more: false, + duration_ms, + error: Some(api_error), + } +} + +fn postgres_console_cancelled(reason: Option) -> AppError { + AppError::new( + AppErrorKind::Conflict, + ApiError::new( + "postgres_console_cancelled", + reason.unwrap_or_else(|| "The PostgreSQL Console execution was cancelled".to_owned()), + ), + ) +} + +fn elapsed_millis(started: Instant) -> u64 { + u64::try_from(started.elapsed().as_millis()).unwrap_or(u64::MAX) +} + +fn display_float32(value: f32) -> String { + if value.is_nan() { + "NaN".to_owned() + } else if value == f32::INFINITY { + "Infinity".to_owned() + } else if value == f32::NEG_INFINITY { + "-Infinity".to_owned() + } else { + value.to_string() + } +} + +fn display_float64(value: f64) -> String { + if value.is_nan() { + "NaN".to_owned() + } else if value == f64::INFINITY { + "Infinity".to_owned() + } else if value == f64::NEG_INFINITY { + "-Infinity".to_owned() + } else { + value.to_string() + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +enum PostgresScriptState { + Normal, + SingleQuote, + DoubleQuote, + DollarQuote(String), + LineComment, + BlockComment(usize), +} + +fn split_postgres_script(script: &str) -> Result, AppError> { + let bytes = script.as_bytes(); + let mut statements = Vec::new(); + let mut state = PostgresScriptState::Normal; + let mut statement_start = 0_usize; + let mut index = 0_usize; + while index < bytes.len() { + match &mut state { + PostgresScriptState::Normal => match bytes[index] { + b'\'' => state = PostgresScriptState::SingleQuote, + b'"' => state = PostgresScriptState::DoubleQuote, + b'$' => { + if let Some(tag) = postgres_dollar_tag(script, index) { + index = index.saturating_add(tag.len().saturating_sub(1)); + state = PostgresScriptState::DollarQuote(tag.to_owned()); + } + } + b'-' if bytes.get(index + 1) == Some(&b'-') => { + index += 1; + state = PostgresScriptState::LineComment; + } + b'/' if bytes.get(index + 1) == Some(&b'*') => { + index += 1; + state = PostgresScriptState::BlockComment(1); + } + b';' => { + push_postgres_statement(&mut statements, &script[statement_start..index])?; + statement_start = index + 1; + } + _ => {} + }, + PostgresScriptState::SingleQuote => { + if bytes[index] == b'\\' { + index = index.saturating_add(1); + } else if bytes[index] == b'\'' { + if bytes.get(index + 1) == Some(&b'\'') { + index += 1; + } else { + state = PostgresScriptState::Normal; + } + } + } + PostgresScriptState::DoubleQuote => { + if bytes[index] == b'"' { + if bytes.get(index + 1) == Some(&b'"') { + index += 1; + } else { + state = PostgresScriptState::Normal; + } + } + } + PostgresScriptState::DollarQuote(tag) => { + if script[index..].starts_with(tag.as_str()) { + index = index.saturating_add(tag.len().saturating_sub(1)); + state = PostgresScriptState::Normal; + } + } + PostgresScriptState::LineComment => { + if bytes[index] == b'\n' { + state = PostgresScriptState::Normal; + } + } + PostgresScriptState::BlockComment(depth) => { + if bytes[index] == b'/' && bytes.get(index + 1) == Some(&b'*') { + *depth = depth.saturating_add(1); + index += 1; + } else if bytes[index] == b'*' && bytes.get(index + 1) == Some(&b'/') { + *depth -= 1; + index += 1; + if *depth == 0 { + state = PostgresScriptState::Normal; + } + } + } + } + index += 1; + } + match state { + PostgresScriptState::Normal | PostgresScriptState::LineComment => {} + PostgresScriptState::SingleQuote => { + return Err(invalid_console_script("unterminated string literal")); + } + PostgresScriptState::DoubleQuote => { + return Err(invalid_console_script("unterminated quoted identifier")); + } + PostgresScriptState::DollarQuote(_) => { + return Err(invalid_console_script("unterminated dollar-quoted body")); + } + PostgresScriptState::BlockComment(_) => { + return Err(invalid_console_script("unterminated block comment")); + } + } + push_postgres_statement(&mut statements, &script[statement_start..])?; + Ok(statements) +} + +fn postgres_words(sql: &str) -> Result, AppError> { + let _ = split_postgres_script(sql)?; + let bytes = sql.as_bytes(); + let mut words = Vec::new(); + let mut state = PostgresScriptState::Normal; + let mut index = 0_usize; + while index < bytes.len() { + match &mut state { + PostgresScriptState::Normal => match bytes[index] { + b'\'' => state = PostgresScriptState::SingleQuote, + b'"' => state = PostgresScriptState::DoubleQuote, + b'$' => { + if let Some(tag) = postgres_dollar_tag(sql, index) { + index = index.saturating_add(tag.len().saturating_sub(1)); + state = PostgresScriptState::DollarQuote(tag.to_owned()); + } else { + index += 1; + while bytes.get(index).is_some_and(u8::is_ascii_digit) { + index += 1; + } + index = index.saturating_sub(1); + } + } + b'-' if bytes.get(index + 1) == Some(&b'-') => { + index += 1; + state = PostgresScriptState::LineComment; + } + b'/' if bytes.get(index + 1) == Some(&b'*') => { + index += 1; + state = PostgresScriptState::BlockComment(1); + } + byte if byte.is_ascii_alphabetic() || byte == b'_' => { + let start = index; + index += 1; + while bytes.get(index).is_some_and(|byte| { + byte.is_ascii_alphanumeric() || *byte == b'_' || *byte == b'$' + }) { + index += 1; + } + words.push(sql[start..index].to_ascii_uppercase()); + index -= 1; + } + _ => {} + }, + PostgresScriptState::SingleQuote => { + if bytes[index] == b'\\' { + index = index.saturating_add(1); + } else if bytes[index] == b'\'' { + if bytes.get(index + 1) == Some(&b'\'') { + index += 1; + } else { + state = PostgresScriptState::Normal; + } + } + } + PostgresScriptState::DoubleQuote => { + if bytes[index] == b'"' { + if bytes.get(index + 1) == Some(&b'"') { + index += 1; + } else { + state = PostgresScriptState::Normal; + } + } + } + PostgresScriptState::DollarQuote(tag) => { + if sql[index..].starts_with(tag.as_str()) { + index = index.saturating_add(tag.len().saturating_sub(1)); + state = PostgresScriptState::Normal; + } + } + PostgresScriptState::LineComment => { + if bytes[index] == b'\n' { + state = PostgresScriptState::Normal; + } + } + PostgresScriptState::BlockComment(depth) => { + if bytes[index] == b'/' && bytes.get(index + 1) == Some(&b'*') { + *depth = depth.saturating_add(1); + index += 1; + } else if bytes[index] == b'*' && bytes.get(index + 1) == Some(&b'/') { + *depth -= 1; + index += 1; + if *depth == 0 { + state = PostgresScriptState::Normal; + } + } + } + } + index += 1; + } + Ok(words) +} + +fn postgres_dollar_tag(sql: &str, start: usize) -> Option<&str> { + let bytes = sql.as_bytes(); + if bytes.get(start) != Some(&b'$') { + return None; + } + let mut end = start + 1; + while let Some(byte) = bytes.get(end) { + if *byte == b'$' { + return Some(&sql[start..=end]); + } + if end - start > 64 + || !(*byte == b'_' || byte.is_ascii_alphanumeric()) + || (end == start + 1 && byte.is_ascii_digit()) + { + return None; + } + end += 1; + } + None +} + +fn push_postgres_statement(statements: &mut Vec, statement: &str) -> Result<(), AppError> { + let statement = statement.trim(); + if statement.is_empty() { + return Ok(()); + } + if statements.len() >= MAX_CONSOLE_STATEMENTS { + return Err(resource_error( + "postgres_console_too_many_statements", + format!( + "PostgreSQL Console scripts are limited to {MAX_CONSOLE_STATEMENTS} statements" + ), + )); + } + statements.push(statement.to_owned()); + Ok(()) +} + +fn invalid_console_script(detail: &str) -> AppError { + AppError::invalid( + "invalid_postgres_console_script", + format!("The PostgreSQL Console script contains an {detail}"), + ) +} + +#[cfg(test)] +mod tests { + use std::time::Duration; + + use chat2db_contract::DatasourceConnectionProperty; + + use super::*; + use crate::native_driver_types::{CreateSchemaSqlRequest, SchemaDefinition}; + + fn datasource_connection(jdbc_url: impl Into) -> DatasourceConnection { + DatasourceConnection { + jdbc_url: jdbc_url.into(), + properties: vec![ + DatasourceConnectionProperty { + key: "user".to_owned(), + value: "postgres".to_owned(), + sensitive: false, + }, + DatasourceConnectionProperty { + key: "password".to_owned(), + value: "postgres".to_owned(), + sensitive: true, + }, + ], + read_only: false, + ssh: None, + } + } + + #[test] + fn normalizes_jdbc_urls_and_rejects_unrelated_schemes() { + assert_eq!( + normalize_postgres_url("JDBC:POSTGRESQL://db.example/app").expect("JDBC URL"), + "postgresql://db.example/app" + ); + assert_eq!( + normalize_postgres_url("postgres://db.example/app").expect("Postgres URL"), + "postgresql://db.example/app" + ); + assert!(normalize_postgres_url("jdbc:mysql://db.example/app").is_err()); + } + + #[test] + fn connection_config_replaces_the_original_port_for_tunnels() { + let mut connection = datasource_connection( + "jdbc:postgresql://url_user@db.example:6432/app?sslmode=require¤tSchema=tenant", + ); + connection.properties.push(DatasourceConnectionProperty { + key: "user".to_owned(), + value: "property_user".to_owned(), + sensitive: false, + }); + let (config, tls_mode, target_host, target_port) = + connection_config(&connection, Some("selected_db"), Some(15_432)) + .expect("connection config"); + + assert_eq!(target_host, "db.example"); + assert_eq!(target_port, 6_432); + assert_eq!(tls_mode, PostgresTlsMode::Require); + assert_eq!(config.get_ports(), &[15_432]); + assert_eq!(config.get_user(), Some("property_user")); + assert_eq!(config.get_dbname(), Some("selected_db")); + assert_eq!(config.get_options(), Some("-c search_path=tenant")); + } + + #[test] + fn splits_dollar_quoted_scripts_and_nested_comments() { + let script = r" + CREATE FUNCTION public.answer() RETURNS integer AS $body$ + BEGIN + /* inner ; /* nested ; */ still inner */ + RETURN 42; + END; + $body$ LANGUAGE plpgsql; + SELECT ';' AS semicolon; + -- trailing comment ; + SELECT 2; + "; + let statements = split_postgres_script(script).expect("PostgreSQL script"); + + assert_eq!(statements.len(), 3); + assert!(statements[0].contains("RETURN 42;")); + assert!(statements[1].contains("SELECT ';'")); + assert!(statements[2].ends_with("SELECT 2")); + assert!(split_postgres_script("SELECT $broken$body").is_err()); + } + + #[test] + fn read_only_validation_blocks_writes_and_row_locks() { + assert!(validate_read_sql("WITH rows AS (SELECT 1) SELECT * FROM rows").is_ok()); + assert!(validate_read_sql("VALUES (1), (2)").is_ok()); + assert!(validate_read_sql("SELECT * FROM example FOR UPDATE").is_err()); + assert!( + validate_read_sql("WITH removed AS (DELETE FROM t RETURNING *) SELECT * FROM removed") + .is_err() + ); + assert!(validate_read_sql("SELECT 1; SELECT 2").is_err()); + assert!( + validate_forced_read_console(&[ + "SELECT 1".to_owned(), + "WITH row AS (SELECT 2) SELECT * FROM row".to_owned(), + ]) + .is_ok() + ); + assert!( + validate_forced_read_console(&["SELECT 1".to_owned(), "COMMIT".to_owned()]).is_err() + ); + assert!(validate_read_sql("SELECT * FROM example FOR KEY SHARE").is_err()); + } + + #[test] + fn console_dispatch_tracks_unknown_write_outcomes_fail_closed() { + for sql in ["SELECT 1", "VALUES (1)", "EXPLAIN SELECT 1"] { + assert_eq!( + ConsoleDispatchState::classify(sql).statement_kind, + ConsoleStatementKind::ReadOnly, + "{sql} must remain a read-only Console statement" + ); + } + for sql in [ + "INSERT INTO example VALUES (1)", + "UPDATE example SET value = 1", + "INSERT INTO example VALUES (1) RETURNING value", + ] { + assert_eq!( + ConsoleDispatchState::classify(sql).statement_kind, + ConsoleStatementKind::Write, + "{sql} must be treated as a Console write" + ); + } + + let mut write = ConsoleDispatchState::classify("INSERT INTO example VALUES (1)"); + assert!(!write.requires_unknown_outcome(ConsoleFailure::Cancelled)); + assert!(!write.requires_unknown_outcome(ConsoleFailure::TimedOut)); + write.mark_dispatched(); + assert!(write.requires_unknown_outcome(ConsoleFailure::Cancelled)); + assert!(write.requires_unknown_outcome(ConsoleFailure::TimedOut)); + assert!(write.requires_unknown_outcome(ConsoleFailure::ResultProcessing)); + assert!(write.requires_unknown_outcome(ConsoleFailure::Driver { + server_rejected: false, + })); + assert!(!write.requires_unknown_outcome(ConsoleFailure::Driver { + server_rejected: true, + })); + + let mut read = ConsoleDispatchState::classify("SELECT 1"); + read.mark_dispatched(); + assert!(!read.requires_unknown_outcome(ConsoleFailure::Cancelled)); + assert!(!read.requires_unknown_outcome(ConsoleFailure::TimedOut)); + assert!(!read.requires_unknown_outcome(ConsoleFailure::ResultProcessing)); + assert!(!read.requires_unknown_outcome(ConsoleFailure::Driver { + server_rejected: false, + })); + } + + #[test] + fn maps_postgres_type_names_without_losing_time_zone_qualifiers() { + assert_eq!( + postgres_jdbc_type_name("timestamp(6) with time zone"), + 2_014 + ); + assert_eq!(postgres_jdbc_type_name("timestamp without time zone"), 93); + assert_eq!(postgres_jdbc_type_name("time with time zone"), 2_013); + assert_eq!(postgres_jdbc_type_name("character varying(64)"), 12); + assert_eq!(postgres_jdbc_type_name("integer[]"), 2_003); + assert_eq!(postgres_jdbc_type(&Type::TIMETZ), 2_013); + } + + #[test] + fn decodes_binary_numeric_and_array_values() { + let mut numeric = Vec::new(); + numeric.extend_from_slice(&3_i16.to_be_bytes()); + numeric.extend_from_slice(&1_i16.to_be_bytes()); + numeric.extend_from_slice(&0_u16.to_be_bytes()); + numeric.extend_from_slice(&4_u16.to_be_bytes()); + for digit in [1_u16, 2_345, 6_789] { + numeric.extend_from_slice(&digit.to_be_bytes()); + } + assert_eq!( + decode_postgres_numeric(&numeric).expect("numeric value"), + "12345.6789" + ); + + let mut array = Vec::new(); + array.extend_from_slice(&1_i32.to_be_bytes()); + array.extend_from_slice(&1_i32.to_be_bytes()); + array.extend_from_slice(&Type::INT4.oid().to_be_bytes()); + array.extend_from_slice(&3_i32.to_be_bytes()); + array.extend_from_slice(&1_i32.to_be_bytes()); + for value in [Some(1_i32), Some(-2_i32), None] { + match value { + Some(value) => { + array.extend_from_slice(&4_i32.to_be_bytes()); + array.extend_from_slice(&value.to_be_bytes()); + } + None => array.extend_from_slice(&(-1_i32).to_be_bytes()), + } + } + assert_eq!( + decode_postgres_array(&Type::INT4_ARRAY, &array).expect("integer array"), + "{1,-2,NULL}" + ); + } + + #[test] + fn decodes_multidimensional_arrays_with_unambiguous_escaping() { + let mut array = Vec::new(); + array.extend_from_slice(&2_i32.to_be_bytes()); + array.extend_from_slice(&0_i32.to_be_bytes()); + array.extend_from_slice(&Type::TEXT.oid().to_be_bytes()); + for dimension in [2_i32, 2_i32] { + array.extend_from_slice(&dimension.to_be_bytes()); + array.extend_from_slice(&1_i32.to_be_bytes()); + } + for value in ["a,b", "NULL", "quote\"slash\\", "white space"] { + array.extend_from_slice( + &i32::try_from(value.len()) + .expect("small array value") + .to_be_bytes(), + ); + array.extend_from_slice(value.as_bytes()); + } + + assert_eq!( + decode_postgres_array(&Type::TEXT_ARRAY, &array).expect("text array"), + "{{\"a,b\",\"NULL\"},{\"quote\\\"slash\\\\\",\"white space\"}}" + ); + + let mut lower_bound = Vec::new(); + lower_bound.extend_from_slice(&1_i32.to_be_bytes()); + lower_bound.extend_from_slice(&0_i32.to_be_bytes()); + lower_bound.extend_from_slice(&Type::INT4.oid().to_be_bytes()); + lower_bound.extend_from_slice(&1_i32.to_be_bytes()); + lower_bound.extend_from_slice(&0_i32.to_be_bytes()); + lower_bound.extend_from_slice(&4_i32.to_be_bytes()); + lower_bound.extend_from_slice(&7_i32.to_be_bytes()); + assert_eq!( + decode_postgres_array(&Type::INT4_ARRAY, &lower_bound).expect("lower-bound array"), + "[0:0]={7}" + ); + + let mut excessive_dimensions = Vec::new(); + excessive_dimensions.extend_from_slice(&7_i32.to_be_bytes()); + excessive_dimensions.extend_from_slice(&0_i32.to_be_bytes()); + excessive_dimensions.extend_from_slice(&Type::INT4.oid().to_be_bytes()); + assert!(decode_postgres_array(&Type::INT4_ARRAY, &excessive_dimensions).is_err()); + } + + #[test] + fn network_and_money_binary_values_match_their_declared_schema() { + let inet = [2, 24, 0, 4, 192, 168, 4, 7]; + assert_eq!( + decode_postgres_network(&Type::INET, &inet).expect("IPv4 inet"), + "192.168.4.7/24" + ); + let mut cidr = vec![3, 64, 1, 16]; + cidr.extend_from_slice(&Ipv6Addr::from_str("2001:db8::").expect("IPv6").octets()); + assert_eq!( + decode_postgres_network(&Type::CIDR, &cidr).expect("IPv6 cidr"), + "2001:db8::/64" + ); + + assert_eq!( + postgres_value_type(&Type::MONEY), + wire::JdbcValueType::Opaque + ); + assert_eq!(postgres_jdbc_type(&Type::MONEY), 1_111); + let decoded = + decode_postgres_value(&Type::MONEY, &12_345_i64.to_be_bytes()).expect("money value"); + assert!(matches!( + decoded, + wire::jdbc_value::Value::OpaqueValue(value) + if value.type_name == "money" && value.display_value == "raw_units=12345" + )); + } + + #[test] + fn scalar_limits_are_checked_before_driver_owned_copies_and_hex_expansion() { + let oversized = vec![0_u8; MAX_SCALAR_BYTES + 1]; + let raw = RawPostgresValue::from_sql(&Type::BYTEA, &oversized) + .expect("oversized values are represented without cloning"); + assert!(matches!( + raw, + RawPostgresValue::TooLarge { byte_count } if byte_count == MAX_SCALAR_BYTES + 1 + )); + + let opaque = vec![0_u8; MAX_SCALAR_BYTES / 2 + 1]; + assert!(postgres_opaque_value("opaque", &opaque).is_err()); + } + + #[test] + fn quotes_identifiers_literals_and_builds_schema_sql() { + assert_eq!( + quote_identifier("Mixed\"Name", "name").expect("identifier"), + "\"Mixed\"\"Name\"" + ); + assert_eq!(quote_literal("owner's").expect("literal"), "E'owner''s'"); + assert_eq!( + quote_literal("path\\owner's").expect("escaped literal"), + "E'path\\\\owner''s'" + ); + assert!(quote_identifier(&"a".repeat(63), "name").is_ok()); + assert!(quote_identifier(&"a".repeat(64), "name").is_err()); + assert!(quote_identifier(&"é".repeat(31), "name").is_ok()); + assert!(quote_identifier(&"é".repeat(32), "name").is_err()); + let built = build_create_schema(CreateSchemaSqlRequest { + schema: SchemaDefinition { + database_name: "app".to_owned(), + name: "Team Data".to_owned(), + comment: "owner's schema".to_owned(), + owner: "postgres".to_owned(), + system: false, + }, + }) + .expect("schema SQL"); + assert_eq!( + built.sql, + "CREATE SCHEMA \"Team Data\" AUTHORIZATION \"postgres\";\nCOMMENT ON SCHEMA \"Team Data\" IS E'owner''s schema';" + ); + } + + #[test] + fn unknown_connection_properties_fail_closed() { + let mut connection = datasource_connection("jdbc:postgresql://db.example/app"); + connection.properties.push(DatasourceConnectionProperty { + key: "vendorMagic".to_owned(), + value: "enabled".to_owned(), + sensitive: false, + }); + assert!(connection_config(&connection, None, None).is_err()); + } + + #[test] + fn rustls_provider_initialization_is_panic_free_and_preserves_existing_provider() { + let before = rustls::crypto::CryptoProvider::get_default().map(std::sync::Arc::as_ptr); + let initialization = std::panic::catch_unwind(ensure_postgres_rustls_provider); + assert!( + initialization.is_ok(), + "provider initialization must not panic" + ); + initialization + .expect("caught provider initialization") + .expect("a Rustls provider must be available"); + let after = rustls::crypto::CryptoProvider::get_default().map(std::sync::Arc::as_ptr); + assert!(after.is_some()); + if before.is_some() { + assert_eq!(before, after, "an installed provider must not be replaced"); + } + } + + #[tokio::test] + async fn closed_cancellation_sender_does_not_spin() { + let (sender, mut receiver) = watch::channel(CancellationRequest::Waiting); + drop(sender); + let completed = tokio::time::timeout( + Duration::from_secs(1), + cancellable_postgres(async { Ok::(7) }, &mut receiver), + ) + .await + .expect("closed cancellation channels must not starve the database future"); + match completed { + Ok(value) => assert_eq!(value, 7), + Err(PostgresCancellableError::Cancelled(reason)) => { + panic!("unexpected cancellation after sender closed: {reason:?}"); + } + Err(PostgresCancellableError::Failed(error)) => { + panic!("unexpected database failure after sender closed: {error}"); + } + } + } + + fn local_smoke_connection() -> DatasourceConnection { + let mut connection = + datasource_connection(std::env::var("CHAT2DB_POSTGRES_URL").unwrap_or_else(|_| { + "jdbc:postgresql://127.0.0.1:5432/app?sslmode=disable".to_owned() + })); + connection.properties[0].value = + std::env::var("CHAT2DB_POSTGRES_USER").unwrap_or_else(|_| "postgres".to_owned()); + connection.properties[1].value = + std::env::var("CHAT2DB_POSTGRES_PASSWORD").unwrap_or_else(|_| "postgres".to_owned()); + connection + } + + #[tokio::test] + #[ignore = "requires a PostgreSQL server; run explicitly for native-driver verification"] + #[allow( + clippy::too_many_lines, + reason = "the real smoke intentionally verifies the complete PostgreSQL fixture lifecycle in one test" + )] + async fn real_postgres_driver_smoke() { + let connection = open_connection(&local_smoke_connection()) + .await + .expect("PostgreSQL connection"); + connection + .client() + .simple_query("SELECT 1") + .await + .expect("connection test"); + connection + .client() + .batch_execute("SET standard_conforming_strings = off") + .await + .expect("legacy string setting"); + let literal_value = "path\\segment's value"; + let literal_sql = format!( + "SELECT {}::text", + quote_literal(literal_value).expect("escaped PostgreSQL literal") + ); + let literal_round_trip: String = connection + .client() + .query_one(&literal_sql, &[]) + .await + .expect("escaped PostgreSQL literal query") + .try_get(0) + .expect("escaped PostgreSQL literal value"); + assert_eq!(literal_round_trip, literal_value); + let schema = format!("chat2db_native_smoke_{}", std::process::id()); + let quoted_schema = quote_identifier(&schema, "schemaName").expect("smoke schema"); + let fixture_sql = format!( + r" + CREATE SCHEMA {quoted_schema}; + CREATE TABLE {quoted_schema}.parent ( + id BIGSERIAL PRIMARY KEY, + label TEXT NOT NULL + ); + CREATE TABLE {quoted_schema}.child ( + id BIGSERIAL PRIMARY KEY, + parent_id BIGINT NOT NULL REFERENCES {quoted_schema}.parent(id), + payload NUMERIC(12, 4), + created_at TIMESTAMPTZ NOT NULL DEFAULT now() + ); + CREATE INDEX child_payload_idx ON {quoted_schema}.child(payload); + CREATE VIEW {quoted_schema}.child_view AS SELECT id, payload FROM {quoted_schema}.child; + CREATE FUNCTION {quoted_schema}.add_one(value integer) RETURNS integer + LANGUAGE SQL IMMUTABLE AS $function$ SELECT value + 1 $function$; + CREATE PROCEDURE {quoted_schema}.no_op() + LANGUAGE plpgsql AS $procedure$ BEGIN NULL; END $procedure$; + CREATE FUNCTION {quoted_schema}.touch_child() RETURNS trigger + LANGUAGE plpgsql AS $trigger_function$ + BEGIN NEW.created_at = clock_timestamp(); RETURN NEW; END + $trigger_function$; + CREATE TRIGGER child_touch BEFORE UPDATE ON {quoted_schema}.child + FOR EACH ROW EXECUTE FUNCTION {quoted_schema}.touch_child(); + INSERT INTO {quoted_schema}.parent(label) VALUES ('parent'); + INSERT INTO {quoted_schema}.child(parent_id, payload) VALUES (1, 12345.6789); + " + ); + connection + .client() + .batch_execute(&fixture_sql) + .await + .expect("PostgreSQL smoke fixture"); + + let verification = async { + let database_name: String = connection + .client() + .query_one("SELECT current_database()", &[]) + .await + .map_err(postgres_query_error)? + .try_get(0) + .map_err(postgres_query_error)?; + let databases: i64 = connection + .client() + .query_one("SELECT count(*) FROM pg_database", &[]) + .await + .map_err(postgres_query_error)? + .try_get(0) + .map_err(postgres_query_error)?; + let catalog_checks = [ + ("schema", "SELECT count(*) FROM pg_namespace WHERE nspname = $1"), + ("table", "SELECT count(*) FROM pg_class c JOIN pg_namespace n ON n.oid = c.relnamespace WHERE n.nspname = $1 AND c.relkind IN ('r', 'p')"), + ("column", "SELECT count(*) FROM pg_attribute a JOIN pg_class c ON c.oid = a.attrelid JOIN pg_namespace n ON n.oid = c.relnamespace WHERE n.nspname = $1 AND a.attnum > 0 AND NOT a.attisdropped"), + ("index", "SELECT count(*) FROM pg_index i JOIN pg_class c ON c.oid = i.indrelid JOIN pg_namespace n ON n.oid = c.relnamespace WHERE n.nspname = $1"), + ("view", "SELECT count(*) FROM pg_class c JOIN pg_namespace n ON n.oid = c.relnamespace WHERE n.nspname = $1 AND c.relkind IN ('v', 'm')"), + ("key", "SELECT count(*) FROM pg_constraint c JOIN pg_namespace n ON n.oid = c.connamespace WHERE n.nspname = $1 AND c.contype IN ('p', 'f')"), + ("function", "SELECT count(*) FROM pg_proc p JOIN pg_namespace n ON n.oid = p.pronamespace WHERE n.nspname = $1 AND p.prokind = 'f'"), + ("procedure", "SELECT count(*) FROM pg_proc p JOIN pg_namespace n ON n.oid = p.pronamespace WHERE n.nspname = $1 AND p.prokind = 'p'"), + ("trigger", "SELECT count(*) FROM pg_trigger t JOIN pg_class c ON c.oid = t.tgrelid JOIN pg_namespace n ON n.oid = c.relnamespace WHERE n.nspname = $1 AND NOT t.tgisinternal"), + ("ER foreign key", "SELECT count(*) FROM pg_constraint c JOIN pg_namespace n ON n.oid = c.connamespace WHERE n.nspname = $1 AND c.contype = 'f'"), + ]; + for (label, sql) in catalog_checks { + let count: i64 = connection + .client() + .query_one(sql, &[&schema]) + .await + .map_err(postgres_query_error)? + .try_get(0) + .map_err(postgres_query_error)?; + if count == 0 { + return Err(resource_error( + "postgres_smoke_metadata_missing", + format!("PostgreSQL {label} metadata returned no rows"), + )); + } + } + let ddl = build_table_ddl(&connection, &database_name, &schema, "child").await?; + let preview = connection + .client() + .query( + &format!("SELECT * FROM {quoted_schema}.child LIMIT 10"), + &[], + ) + .await + .map_err(postgres_query_error)?; + let preview_row = preview.first().ok_or_else(|| { + resource_error("postgres_smoke_preview_empty", "PostgreSQL preview returned no rows") + })?; + let decoded_preview = postgres_row(preview_row, preview_row.columns())?; + Ok::<_, AppError>((databases, ddl, decoded_preview.values.len())) + } + .await; + + let cleanup_result = connection + .client() + .batch_execute(&format!("DROP SCHEMA {quoted_schema} CASCADE")) + .await + .map_err(postgres_query_error); + finish_connection(connection, cleanup_result) + .await + .expect("PostgreSQL smoke cleanup"); + let (databases, ddl, preview_columns) = + verification.expect("PostgreSQL smoke verification"); + assert!(databases > 0); + assert!(ddl.contains("CREATE TABLE")); + assert!(ddl.contains("FOREIGN KEY")); + assert!(ddl.contains("CREATE INDEX")); + assert_eq!(preview_columns, 4); + } +} diff --git a/crates/chat2db-core/src/native_sqlserver.rs b/crates/chat2db-core/src/native_sqlserver.rs new file mode 100644 index 0000000..7f9dea9 --- /dev/null +++ b/crates/chat2db-core/src/native_sqlserver.rs @@ -0,0 +1,5594 @@ +use std::{ + collections::HashMap, + fmt::Write as _, + panic::{AssertUnwindSafe, catch_unwind}, + time::{Duration, Instant}, +}; + +use async_trait::async_trait; +use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64_STANDARD}; +use chat2db_contract::{ + ApiError, ColumnNullability, DatasourceConnection, JdbcValue, JdbcValueType, QueryLimits, + ResultColumn, ResultMetadata, ResultRow, StartQueryRequest, +}; +use chat2db_engine_protocol::wire; +use chat2db_storage::Storage; +use chrono::{DateTime, FixedOffset, NaiveDate, NaiveDateTime, NaiveTime}; +use futures_util::{FutureExt as _, stream::TryStreamExt}; +use prost::Message; +use sqlparser::{ + ast::{Query as SqlQuery, SetExpr, Statement}, + dialect::MsSqlDialect, + parser::Parser, +}; +use tiberius::{ + Client, Column, ColumnData, ColumnType, Config, Query, QueryItem, Row, + error::Error as TiberiusError, numeric::Numeric, +}; +use tokio::{net::TcpStream, sync::watch}; +use tokio_util::{ + compat::{Compat, TokioAsyncWriteCompatExt}, + sync::CancellationToken, +}; + +use crate::{ + AppError, AppErrorKind, Application, + datasource_session::{ResolvedDatasourceConnection, resolve_datasource_connection}, + native_driver::{ + NativeConnectionDriver, NativeDialectDriver, NativeDriver, NativeMetadataDriver, + NativeQueryDriver, NativeTableDriver, + }, + native_driver_types::{ + BuiltSql, ColumnList, ColumnMetadata, CreateSchemaSqlRequest, DatabaseDefinition, + DatabaseList, DatabaseMetadata, DmlAssignment, DmlColumn, DmlRow, DmlSqlRequest, + DmlStatement, DmlTarget, DmlTemporalKind, DmlValue, EntityRelationColumn, + EntityRelationForeignKey, EntityRelationTable, ForeignKeyList, ForeignKeyMetadata, + FunctionList, FunctionMetadata, FunctionParameterList, FunctionParameterMetadata, + IndexColumnMetadata, IndexList, IndexMetadata, ListColumnsRequest, ListDatabasesRequest, + ListIndexesRequest, ListRoutinesRequest, ListSchemasRequest, ListTableKeysRequest, + ListTablesRequest, ListTriggersRequest, ListViewsRequest, MetadataObjectRef, MetadataScope, + NamespaceSqlOperation, NamespaceSqlRequest, NativeDriverDescriptor, PrimaryKeyList, + PrimaryKeyMetadata, ProcedureList, ProcedureMetadata, ProcedureParameterList, + ProcedureParameterMetadata, SchemaDefinition, SchemaList, SchemaMetadata, TableList, + TableMetadata, TablePreviewAccepted, TablePreviewRequest, TableRef, TriggerList, + TriggerMetadata, ViewList, + }, + operation::CancellationRequest, + query::{ + DatabaseValue, DatabaseWriteError, NativeConsoleRequest, NativeConsoleResult, + PreparedQuery, QueryExecutionOptions, QueryParameter, QueryTaskError, RetainedWriter, + }, + ssh::{SshTunnel, SshTunnelIdentity}, +}; + +const CONNECT_TIMEOUT: Duration = Duration::from_secs(15); +const METADATA_TIMEOUT: Duration = Duration::from_secs(30); +const DEFAULT_BATCH_ROWS: u32 = 256; +const DEFAULT_BATCH_BYTES: u32 = 256 * 1024; +const DEFAULT_RESULT_BYTES: u64 = wire::JdbcResultByteLimit::DefaultResultBytes as u64; +const MAX_RESULT_BYTES: u64 = wire::JdbcResultByteLimit::MaxResultBytes as u64; +const MAX_BATCH_ROWS: u32 = wire::JdbcProtocolLimit::MaxBatchRows as u32; +const MAX_BATCH_BYTES: u32 = wire::JdbcProtocolLimit::MaxBatchBytes as u32; +const MAX_COLUMNS: usize = wire::JdbcProtocolLimit::MaxColumns as usize; +const MAX_PARAMETERS: usize = wire::JdbcProtocolLimit::MaxParameters as usize; +const MAX_SQL_BYTES: usize = wire::JdbcProtocolLimit::MaxSqlBytes as usize; +const MAX_SCALAR_BYTES: usize = wire::JdbcProtocolLimit::MaxScalarBytes as usize; +const MAX_CONSOLE_PAGE_SIZE: u32 = 10_000; +const MAX_CONSOLE_STATEMENTS: usize = 1_000; +const MAX_CONSOLE_RESULT_BYTES: u64 = DEFAULT_RESULT_BYTES; +const MAX_IDENTIFIER_BYTES: usize = 256; +const JDBC_SQLSERVER_PREFIX: &str = "jdbc:sqlserver://"; +const NATIVE_SQLSERVER_PREFIX: &str = "sqlserver://"; + +pub(crate) const SQLSERVER_DRIVER_DESCRIPTOR: NativeDriverDescriptor = NativeDriverDescriptor { + id: "sqlserver", + implementation: "tiberius", + database_types: &["SQLSERVER", "MSSQL", "SQL_SERVER"], + compatibility_aliases: &[ + "com.microsoft.sqlserver.jdbc.SQLServerDriver", + "microsoft-sql-server", + "sql-server", + ], +}; + +pub(crate) struct SqlServerNativeDriver; + +#[async_trait] +impl NativeConnectionDriver for SqlServerNativeDriver { + async fn test_connection(&self, connection: &DatasourceConnection) -> Result<(), AppError> { + let conn = open_connection(connection).await?; + finish_connection(conn, Ok(())).await + } + + async fn test_connection_with_local_port( + &self, + connection: &DatasourceConnection, + ) -> Result, AppError> { + let conn = open_connection(connection).await?; + let local_port = conn.tunnel.as_ref().map(SshTunnel::local_port); + finish_connection(conn, Ok(local_port)).await + } +} + +#[async_trait] +impl NativeQueryDriver for SqlServerNativeDriver { + fn is_read_candidate(&self, sql: &str) -> Result { + is_read_candidate(sql) + } + + fn validate_query(&self, query: &PreparedQuery) -> Result<(), AppError> { + validate_query(query) + } + + async fn execute_query_task( + &self, + application: &Application, + operation_id: &str, + cancellation: watch::Receiver, + query: PreparedQuery, + storage: Storage, + resolved: ResolvedDatasourceConnection, + ) -> Result { + match AssertUnwindSafe(execute_query_task( + application, + operation_id, + cancellation, + query, + storage, + resolved, + )) + .catch_unwind() + .await + { + Ok(result) => result, + Err(_) => Err(QueryTaskError::Failed(sqlserver_driver_failure())), + } + } + + async fn execute_update( + &self, + resolved: ResolvedDatasourceConnection, + sql: String, + cancellation: CancellationToken, + ) -> Result { + match AssertUnwindSafe(execute_update(resolved, sql, cancellation)) + .catch_unwind() + .await + { + Ok(result) => result, + Err(_) => Err(DatabaseWriteError::unknown(sqlserver_driver_failure())), + } + } + + async fn execute_console( + &self, + application: &Application, + request: NativeConsoleRequest, + cancellation: watch::Receiver, + force_read_only: bool, + ) -> Result, AppError> { + match AssertUnwindSafe(execute_console( + application, + request, + cancellation, + force_read_only, + )) + .catch_unwind() + .await + { + Ok(result) => result, + Err(_) => Err(sqlserver_driver_failure()), + } + } +} + +#[async_trait] +impl NativeMetadataDriver for SqlServerNativeDriver { + async fn list_schemas( + &self, + application: &Application, + request: ListSchemasRequest, + ) -> Result { + list_schemas(application, &request.datasource_id, &request.database_name).await + } + + async fn list_databases( + &self, + application: &Application, + request: ListDatabasesRequest, + ) -> Result { + list_databases(application, &request.datasource_id).await + } + + async fn list_tables( + &self, + application: &Application, + request: ListTablesRequest, + ) -> Result { + list_tables( + application, + &request.scope.datasource_id, + &request.scope.database_name, + &request.scope.schema_name, + &request.name_pattern, + ) + .await + } + + async fn list_columns( + &self, + application: &Application, + request: ListColumnsRequest, + ) -> Result { + list_columns( + application, + &request.table.scope.datasource_id, + &request.table.scope.database_name, + &request.table.scope.schema_name, + &request.table.table_name, + ) + .await + } + + async fn list_indexes( + &self, + application: &Application, + request: ListIndexesRequest, + ) -> Result { + list_indexes( + application, + &request.table.scope.datasource_id, + &request.table.scope.database_name, + &request.table.scope.schema_name, + &request.table.table_name, + ) + .await + } + + async fn list_views( + &self, + application: &Application, + request: ListViewsRequest, + ) -> Result { + list_views( + application, + &request.scope.datasource_id, + &request.scope.database_name, + &request.scope.schema_name, + &request.name_pattern, + ) + .await + } + + async fn get_view( + &self, + application: &Application, + request: MetadataObjectRef, + ) -> Result { + get_view( + application, + &request.scope.datasource_id, + &request.scope.database_name, + &request.scope.schema_name, + &request.object_name, + ) + .await + } + + async fn list_imported_keys( + &self, + application: &Application, + request: ListTableKeysRequest, + ) -> Result { + list_foreign_keys(application, &request, false).await + } + + async fn list_exported_keys( + &self, + application: &Application, + request: ListTableKeysRequest, + ) -> Result { + list_foreign_keys(application, &request, true).await + } + + async fn list_primary_keys( + &self, + application: &Application, + request: ListTableKeysRequest, + ) -> Result { + list_primary_keys(application, &request).await + } + + async fn list_functions( + &self, + application: &Application, + request: ListRoutinesRequest, + ) -> Result { + list_functions(application, &request).await + } + + async fn get_function( + &self, + application: &Application, + request: MetadataObjectRef, + ) -> Result { + get_function(application, &request).await + } + + async fn list_function_parameters( + &self, + application: &Application, + request: MetadataObjectRef, + ) -> Result { + list_function_parameters(application, &request).await + } + + async fn list_procedures( + &self, + application: &Application, + request: ListRoutinesRequest, + ) -> Result { + list_procedures(application, &request).await + } + + async fn get_procedure( + &self, + application: &Application, + request: MetadataObjectRef, + ) -> Result { + get_procedure(application, &request).await + } + + async fn list_procedure_parameters( + &self, + application: &Application, + request: MetadataObjectRef, + ) -> Result { + list_procedure_parameters(application, &request).await + } + + async fn list_triggers( + &self, + application: &Application, + request: ListTriggersRequest, + ) -> Result { + list_triggers(application, &request).await + } + + async fn get_trigger( + &self, + application: &Application, + request: MetadataObjectRef, + ) -> Result { + get_trigger(application, &request).await + } +} + +#[async_trait] +impl NativeTableDriver for SqlServerNativeDriver { + async fn load_er_tables( + &self, + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + ) -> Result, AppError> { + load_er_tables(application, datasource_id, database_name, schema_name).await + } + + async fn validate_column_reorder( + &self, + _application: &Application, + _datasource_id: &str, + _database_name: &str, + _table_name: &str, + column_names: &[String], + ) -> Result<(), AppError> { + if column_names.is_empty() { + return Ok(()); + } + Err(AppError::invalid( + "sqlserver_column_reorder_not_supported", + "SQL Server cannot reorder existing columns with an in-place ALTER TABLE operation", + )) + } + + async fn table_ddl( + &self, + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + table_name: &str, + ) -> Result { + table_ddl( + application, + datasource_id, + database_name, + schema_name, + table_name, + ) + .await + } + + async fn start_table_preview( + &self, + application: &Application, + request: TablePreviewRequest, + row_limit: u32, + ) -> Result { + start_table_preview(application, request, row_limit).await + } +} + +impl NativeDialectDriver for SqlServerNativeDriver { + fn build_create_schema(&self, request: CreateSchemaSqlRequest) -> Result { + build_create_schema(request) + } + + fn build_namespace_sql(&self, request: NamespaceSqlRequest) -> Result { + build_namespace_sql(request) + } + + fn build_dml(&self, request: DmlSqlRequest) -> Result { + build_dml(request) + } +} + +impl NativeDriver for SqlServerNativeDriver { + fn descriptor(&self) -> &'static NativeDriverDescriptor { + &SQLSERVER_DRIVER_DESCRIPTOR + } + + fn connection(&self) -> Option<&dyn NativeConnectionDriver> { + Some(self) + } + + fn query(&self) -> Option<&dyn NativeQueryDriver> { + Some(self) + } + + fn metadata(&self) -> Option<&dyn NativeMetadataDriver> { + Some(self) + } + + fn tables(&self) -> Option<&dyn NativeTableDriver> { + Some(self) + } + + fn dialect(&self) -> Option<&dyn NativeDialectDriver> { + Some(self) + } +} + +fn build_create_schema(request: CreateSchemaSqlRequest) -> Result { + let SchemaDefinition { + database_name, + name, + comment, + owner, + system: _, + } = request.schema; + let quoted_name = quote_dialect_identifier(&name, "schemaName")?; + let mut statements = Vec::new(); + if !database_name.trim().is_empty() { + statements.push(format!( + "USE {};", + quote_dialect_identifier(&database_name, "databaseName")? + )); + } + let authorization = if owner.trim().is_empty() { + String::new() + } else { + format!( + " AUTHORIZATION {}", + quote_dialect_identifier(&owner, "owner")? + ) + }; + let create_schema = format!("CREATE SCHEMA {quoted_name}{authorization}"); + statements.push(dynamic_statement(&create_schema)?); + if !comment.is_empty() { + statements.push(format!( + "EXEC sys.sp_addextendedproperty @name = N'MS_Description', @value = {}, @level0type = N'SCHEMA', @level0name = {};", + quote_dialect_literal(&comment, "schema comment")?, + quote_dialect_literal(&name, "schemaName")? + )); + } + Ok(BuiltSql { + sql: statements.join("\n"), + }) +} + +fn build_namespace_sql(request: NamespaceSqlRequest) -> Result { + let sql = match request.operation { + NamespaceSqlOperation::CreateDatabase { database } => build_create_database(&database)?, + NamespaceSqlOperation::AlterDatabase { + old_database, + new_database, + } => build_alter_database(&old_database, &new_database)?, + NamespaceSqlOperation::DropDatabase { database_name } => format!( + "DROP DATABASE {};", + quote_dialect_identifier(&database_name, "databaseName")? + ), + NamespaceSqlOperation::UseDatabase { database_name } => format!( + "USE {};", + quote_dialect_identifier(&database_name, "databaseName")? + ), + NamespaceSqlOperation::CreateSchema { schema } => { + return build_create_schema(CreateSchemaSqlRequest { schema }); + } + NamespaceSqlOperation::AlterSchema { + old_schema_name, + new_schema_name, + } => { + quote_dialect_identifier(&old_schema_name, "schemaName")?; + quote_dialect_identifier(&new_schema_name, "schemaName")?; + return Err(AppError::invalid( + "sqlserver_schema_rename_unsupported", + "SQL Server has no direct schema rename operation; create a new schema and transfer objects explicitly", + )); + } + NamespaceSqlOperation::DropSchema { schema_name } => format!( + "DROP SCHEMA {};", + quote_dialect_identifier(&schema_name, "schemaName")? + ), + }; + Ok(BuiltSql { sql }) +} + +fn build_create_database(database: &DatabaseDefinition) -> Result { + reject_database_charset(&database.charset)?; + let name = quote_dialect_identifier(&database.name, "databaseName")?; + let mut create = format!("CREATE DATABASE {name}"); + if !database.collation.trim().is_empty() { + write!( + &mut create, + " COLLATE {}", + validate_collation(&database.collation)? + ) + .map_err(|_| AppError::internal())?; + } + create.push(';'); + let mut statements = vec![create]; + if !database.owner.trim().is_empty() { + statements.push(format!( + "ALTER AUTHORIZATION ON DATABASE::{name} TO {};", + quote_dialect_identifier(&database.owner, "owner")? + )); + } + if !database.comment.is_empty() { + statements.push(database_comment_statement( + &database.name, + DatabaseCommentOperation::Add, + &database.comment, + )?); + } + Ok(statements.join("\n")) +} + +fn build_alter_database( + old_database: &DatabaseDefinition, + new_database: &DatabaseDefinition, +) -> Result { + if old_database.charset != new_database.charset { + return Err(database_charset_unsupported()); + } + reject_database_charset(&new_database.charset)?; + let old_name = quote_dialect_identifier(&old_database.name, "databaseName")?; + let new_name = quote_dialect_identifier(&new_database.name, "databaseName")?; + let mut statements = Vec::new(); + if old_database.name != new_database.name { + statements.push(format!( + "ALTER DATABASE {old_name} MODIFY NAME = {new_name};" + )); + } + let active_name = if old_database.name == new_database.name { + &old_database.name + } else { + &new_database.name + }; + let quoted_active_name = quote_dialect_identifier(active_name, "databaseName")?; + if old_database.collation != new_database.collation { + if new_database.collation.trim().is_empty() { + return Err(AppError::invalid( + "sqlserver_database_alter_unsupported", + "SQL Server database collation cannot be cleared", + )); + } + statements.push(format!( + "ALTER DATABASE {quoted_active_name} COLLATE {};", + validate_collation(&new_database.collation)? + )); + } + if old_database.owner != new_database.owner { + if new_database.owner.trim().is_empty() { + return Err(AppError::invalid( + "sqlserver_database_alter_unsupported", + "SQL Server database ownership cannot be cleared", + )); + } + statements.push(format!( + "ALTER AUTHORIZATION ON DATABASE::{quoted_active_name} TO {};", + quote_dialect_identifier(&new_database.owner, "owner")? + )); + } + if old_database.comment != new_database.comment { + let operation = match ( + old_database.comment.is_empty(), + new_database.comment.is_empty(), + ) { + (true, false) => DatabaseCommentOperation::Add, + (false, true) => DatabaseCommentOperation::Drop, + (false, false) => DatabaseCommentOperation::Update, + (true, true) => unreachable!("equal empty comments were filtered above"), + }; + statements.push(database_comment_statement( + active_name, + operation, + &new_database.comment, + )?); + } + if statements.is_empty() { + return Err(AppError::invalid( + "sqlserver_database_alter_empty", + "The SQL Server database definition has no supported changes", + )); + } + Ok(statements.join("\n")) +} + +#[derive(Clone, Copy)] +enum DatabaseCommentOperation { + Add, + Update, + Drop, +} + +fn database_comment_statement( + database_name: &str, + operation: DatabaseCommentOperation, + comment: &str, +) -> Result { + let database_name = quote_dialect_identifier(database_name, "databaseName")?; + let procedure = match operation { + DatabaseCommentOperation::Add => "sp_addextendedproperty", + DatabaseCommentOperation::Update => "sp_updateextendedproperty", + DatabaseCommentOperation::Drop => "sp_dropextendedproperty", + }; + let value = if matches!(operation, DatabaseCommentOperation::Drop) { + String::new() + } else { + format!( + ", @value = {}", + quote_dialect_literal(comment, "database comment")? + ) + }; + Ok(format!( + "EXEC {database_name}.sys.{procedure} @name = N'MS_Description'{value};" + )) +} + +fn dynamic_statement(sql: &str) -> Result { + Ok(format!( + "EXEC({});", + quote_dialect_literal(sql, "dynamic SQL statement")? + )) +} + +fn reject_database_charset(charset: &str) -> Result<(), AppError> { + if charset.trim().is_empty() { + Ok(()) + } else { + Err(database_charset_unsupported()) + } +} + +fn database_charset_unsupported() -> AppError { + AppError::invalid( + "sqlserver_database_charset_unsupported", + "SQL Server selects character encoding through collation and has no separate database charset option", + ) +} + +fn validate_collation(collation: &str) -> Result<&str, AppError> { + if collation.is_empty() + || collation.len() > MAX_IDENTIFIER_BYTES + || !collation + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || byte == b'_') + { + return Err(AppError::invalid( + "invalid_sqlserver_dialect_request", + "SQL Server collation must contain only ASCII letters, digits, and underscores", + )); + } + Ok(collation) +} + +fn build_dml(request: DmlSqlRequest) -> Result { + let target = dml_target(&request.target)?; + let sql = match request.statement { + DmlStatement::SingleInsert { columns, row } => { + insert_sql(&target, &columns, std::slice::from_ref(&row))? + } + DmlStatement::MultiInsert { columns, rows } => insert_sql(&target, &columns, &rows)?, + DmlStatement::Update { + assignments, + predicates, + } => update_sql(&target, &assignments, &predicates)?, + }; + Ok(BuiltSql { sql }) +} + +fn dml_target(target: &DmlTarget) -> Result { + let table = quote_dialect_identifier(&target.table_name, "tableName")?; + let database = target + .database_name + .as_deref() + .filter(|value| !value.trim().is_empty()); + let schema = target + .schema_name + .as_deref() + .filter(|value| !value.trim().is_empty()); + match (database, schema) { + (Some(database), Some(schema)) => Ok(format!( + "{}.{}.{}", + quote_dialect_identifier(database, "databaseName")?, + quote_dialect_identifier(schema, "schemaName")?, + table + )), + (Some(database), None) => Ok(format!( + "{}..{}", + quote_dialect_identifier(database, "databaseName")?, + table + )), + (None, Some(schema)) => Ok(format!( + "{}.{}", + quote_dialect_identifier(schema, "schemaName")?, + table + )), + (None, None) => Ok(table), + } +} + +fn insert_sql(target: &str, columns: &[DmlColumn], rows: &[DmlRow]) -> Result { + if columns.is_empty() || rows.is_empty() { + return Err(invalid_dml( + "SQL Server INSERT requires at least one column and row", + )); + } + let columns_sql = columns + .iter() + .map(|column| quote_dialect_identifier(&column.name, "columnName")) + .collect::, _>>()? + .join(", "); + let rows_sql = rows + .iter() + .map(|row| { + if row.values.len() != columns.len() { + return Err(invalid_dml( + "Each SQL Server INSERT row must match the selected column count", + )); + } + row.values + .iter() + .zip(columns) + .map(|(value, column)| dml_value(value, column)) + .collect::, _>>() + .map(|values| format!("({})", values.join(", "))) + }) + .collect::, _>>()? + .join(",\n"); + Ok(format!( + "INSERT INTO {target} ({columns_sql}) VALUES\n{rows_sql};" + )) +} + +fn update_sql( + target: &str, + assignments: &[DmlAssignment], + predicates: &[DmlAssignment], +) -> Result { + if assignments.is_empty() || predicates.is_empty() { + return Err(invalid_dml( + "SQL Server UPDATE requires assignments and key predicates", + )); + } + let assignments = assignments + .iter() + .map(|assignment| { + Ok(format!( + "{} = {}", + quote_dialect_identifier(&assignment.column.name, "columnName")?, + dml_value(&assignment.value, &assignment.column)? + )) + }) + .collect::, AppError>>()? + .join(", "); + let predicates = predicates + .iter() + .map(|predicate| { + let column = quote_dialect_identifier(&predicate.column.name, "columnName")?; + match predicate.value { + DmlValue::Null => Ok(format!("{column} IS NULL")), + _ => Ok(format!( + "{column} = {}", + dml_value(&predicate.value, &predicate.column)? + )), + } + }) + .collect::, AppError>>()? + .join(" AND "); + Ok(format!( + "UPDATE {target} SET {assignments} WHERE {predicates};" + )) +} + +fn dml_value(value: &DmlValue, column: &DmlColumn) -> Result { + match value { + DmlValue::Null => Ok("NULL".to_owned()), + DmlValue::String(value) => { + let literal = quote_dialect_literal(value, "DML string")?; + if column + .data_type_name + .eq_ignore_ascii_case("uniqueidentifier") + { + tiberius::Uuid::parse_str(value) + .map_err(|_| invalid_dml("The SQL Server uniqueidentifier value is invalid"))?; + Ok(format!("CAST({literal} AS uniqueidentifier)")) + } else { + Ok(literal) + } + } + DmlValue::Decimal(value) => parse_decimal(value) + .map(format_numeric) + .map_err(|_| invalid_dml("The SQL Server decimal value is invalid")), + DmlValue::Boolean(value) => Ok(if *value { "1" } else { "0" }.to_owned()), + DmlValue::Temporal { kind, iso8601 } => { + match kind { + DmlTemporalKind::Date => NaiveDate::parse_from_str(iso8601, "%Y-%m-%d") + .map(|_| ()) + .map_err(|_| invalid_dml("The SQL Server date value is invalid"))?, + DmlTemporalKind::Time => NaiveTime::parse_from_str(iso8601, "%H:%M:%S%.f") + .map(|_| ()) + .map_err(|_| invalid_dml("The SQL Server time value is invalid"))?, + DmlTemporalKind::LocalDatetime => parse_timestamp(iso8601) + .map(|_| ()) + .map_err(|_| invalid_dml("The SQL Server datetime2 value is invalid"))?, + DmlTemporalKind::OffsetDatetime => DateTime::parse_from_rfc3339(iso8601) + .map(|_| ()) + .map_err(|_| invalid_dml("The SQL Server datetimeoffset value is invalid"))?, + } + let data_type = match kind { + DmlTemporalKind::Date => "date", + DmlTemporalKind::Time => "time", + DmlTemporalKind::LocalDatetime => "datetime2", + DmlTemporalKind::OffsetDatetime => "datetimeoffset", + }; + Ok(format!( + "CAST({} AS {data_type})", + quote_dialect_literal(iso8601, "DML temporal value")? + )) + } + DmlValue::Binary(value) => { + if value.len() > MAX_SCALAR_BYTES { + return Err(invalid_dml( + "The SQL Server binary value exceeds the scalar limit", + )); + } + Ok(format!("0x{}", hex::encode(value))) + } + } +} + +fn quote_dialect_identifier(value: &str, field: &str) -> Result { + if value.trim().is_empty() || value.len() > MAX_IDENTIFIER_BYTES || value.contains('\0') { + return Err(AppError::invalid( + "invalid_sqlserver_dialect_request", + format!("{field} is invalid"), + )); + } + Ok(format!("[{}]", value.replace(']', "]]"))) +} + +fn quote_dialect_literal(value: &str, field: &str) -> Result { + if value.len() > MAX_SCALAR_BYTES || value.contains('\0') { + return Err(AppError::invalid( + "invalid_sqlserver_dialect_request", + format!("{field} is invalid"), + )); + } + Ok(format!("N'{}'", value.replace('\'', "''"))) +} + +fn invalid_dml(message: impl Into) -> AppError { + AppError::invalid("invalid_sqlserver_dml", message.into()) +} + +type SqlServerClient = Client>; + +struct PreparedSqlServerConnection { + config: Config, + connect_addr: String, + tunnel: Option, +} + +struct ManagedSqlServerConnection { + client: SqlServerClient, + tunnel: Option, +} + +enum QueryOpenError { + Cancelled(Option), + Failed(AppError), +} + +async fn resolve_native_connection( + application: &Application, + datasource_id: &str, +) -> Result { + let storage = application.require_storage()?; + let resolved = resolve_datasource_connection(&storage, datasource_id).await?; + if application + .native_driver_for_datasource_driver_id(&resolved.driver_id) + .is_none_or(|driver| driver.descriptor().id != SQLSERVER_DRIVER_DESCRIPTOR.id) + { + return Err(AppError::invalid( + "sqlserver_driver_mismatch", + "The datasource is not configured with the native SQL Server driver", + )); + } + Ok(resolved) +} + +async fn open_connection( + connection: &DatasourceConnection, +) -> Result { + open_prepared_connection(prepare_connection(connection, SshTunnelIdentity::Ephemeral).await?) + .await +} + +async fn open_resolved_connection( + resolved: &ResolvedDatasourceConnection, + database_name: Option<&str>, +) -> Result { + let identity = SshTunnelIdentity::Datasource { + datasource_id: &resolved.datasource_id, + revision: resolved.datasource_revision, + }; + let mut prepared = prepare_connection(&resolved.connection, identity).await?; + if let Some(database_name) = database_name.filter(|value| !value.trim().is_empty()) { + validate_identifier(database_name, "databaseName")?; + prepared.config.database(database_name); + } + open_prepared_connection(prepared).await +} + +async fn prepare_connection( + connection: &DatasourceConnection, + identity: SshTunnelIdentity<'_>, +) -> Result { + let config = connection_config(connection)?; + let direct_addr = config.get_addr(); + let Some(ssh) = connection.ssh.as_ref() else { + return Ok(PreparedSqlServerConnection { + config, + connect_addr: direct_addr, + tunnel: None, + }); + }; + let (target_host, target_port) = sqlserver_target(&connection.jdbc_url)?; + let tunnel = SshTunnel::open(identity, ssh, target_host, target_port).await?; + Ok(PreparedSqlServerConnection { + config, + connect_addr: format!("127.0.0.1:{}", tunnel.local_port()), + tunnel: Some(tunnel), + }) +} + +async fn open_prepared_connection( + mut prepared: PreparedSqlServerConnection, +) -> Result { + let config = prepared.config.clone(); + let connect_addr = prepared.connect_addr.clone(); + let open = async { + let tcp = TcpStream::connect(&connect_addr) + .await + .map_err(|_| sqlserver_connection_failed())?; + tcp.set_nodelay(true) + .map_err(|_| sqlserver_connection_failed())?; + Client::connect(config.clone(), tcp.compat_write()) + .await + .map_err(|_| sqlserver_connection_failed()) + }; + let client = match tokio::time::timeout(CONNECT_TIMEOUT, open).await { + Ok(Ok(client)) => client, + Ok(Err(error)) => { + close_tunnel_quietly(prepared.tunnel.take()).await; + return Err(error); + } + Err(_) => { + close_tunnel_quietly(prepared.tunnel.take()).await; + return Err(AppError::unavailable( + "sqlserver_connection_timeout", + "The SQL Server connection attempt timed out", + )); + } + }; + Ok(ManagedSqlServerConnection { + client, + tunnel: prepared.tunnel, + }) +} + +async fn finish_connection( + conn: ManagedSqlServerConnection, + result: Result, +) -> Result { + let ManagedSqlServerConnection { client, tunnel, .. } = conn; + drop(client); + let close = match tunnel { + Some(tunnel) => tunnel.close().await, + None => Ok(()), + }; + match result { + Ok(value) => close.map(|()| value), + Err(error) => { + if let Err(close_error) = close { + tracing::warn!(error = %close_error, "SQL Server SSH tunnel cleanup failed"); + } + Err(error) + } + } +} + +async fn close_tunnel_quietly(tunnel: Option) { + if let Some(tunnel) = tunnel + && let Err(error) = tunnel.close().await + { + tracing::warn!(error = %error, "SQL Server SSH tunnel cleanup failed"); + } +} + +async fn discard_connection(conn: ManagedSqlServerConnection) { + let ManagedSqlServerConnection { client, tunnel } = conn; + drop(client); + close_tunnel_quietly(tunnel).await; +} + +fn connection_config(connection: &DatasourceConnection) -> Result { + let mut jdbc_url = normalize_sqlserver_url(&connection.jdbc_url)?; + for property in &connection.properties { + let key = property.key.trim(); + if key.is_empty() || key.contains([';', '=', '\0']) { + return Err(AppError::invalid( + "invalid_sqlserver_connection", + "A SQL Server connection property name is invalid", + )); + } + if !jdbc_url.ends_with(';') { + jdbc_url.push(';'); + } + jdbc_url.push_str(key); + jdbc_url.push('='); + jdbc_url.push_str(&encode_jdbc_property(&property.value)); + } + validate_tls_connection_properties(&jdbc_url)?; + let Ok(Ok(mut config)) = catch_unwind(AssertUnwindSafe(|| Config::from_jdbc_string(&jdbc_url))) + else { + return Err(AppError::invalid( + "invalid_sqlserver_connection", + "A valid jdbc:sqlserver:// connection URL is required", + )); + }; + config.readonly(connection.read_only); + config.application_name("Chat2DB Rust"); + Ok(config) +} + +fn validate_tls_connection_properties(jdbc_url: &str) -> Result<(), AppError> { + let mut trust_server_certificate = false; + let mut trust_server_certificate_ca = false; + let mut start = jdbc_url.find(';').map_or(jdbc_url.len(), |index| index + 1); + let mut in_braces = false; + let mut chars = jdbc_url[start..].char_indices().peekable(); + while let Some((offset, ch)) = chars.next() { + match ch { + '{' if !in_braces => in_braces = true, + '}' if in_braces => { + if chars.peek().is_some_and(|(_, next)| *next == '}') { + chars.next(); + } else { + in_braces = false; + } + } + ';' if !in_braces => { + record_tls_property( + &jdbc_url[start..start + offset], + &mut trust_server_certificate, + &mut trust_server_certificate_ca, + ); + start += offset + 1; + chars = jdbc_url[start..].char_indices().peekable(); + } + _ => {} + } + } + record_tls_property( + &jdbc_url[start..], + &mut trust_server_certificate, + &mut trust_server_certificate_ca, + ); + if trust_server_certificate && trust_server_certificate_ca { + return Err(AppError::invalid( + "invalid_sqlserver_connection", + "trustServerCertificate and trustServerCertificateCA cannot be configured together", + )); + } + Ok(()) +} + +fn record_tls_property(segment: &str, trust_certificate: &mut bool, trust_ca: &mut bool) { + let key = segment + .split_once('=') + .map_or(segment, |(key, _)| key) + .trim(); + if key.eq_ignore_ascii_case("trustServerCertificate") { + *trust_certificate = true; + } else if key.eq_ignore_ascii_case("trustServerCertificateCA") { + *trust_ca = true; + } +} + +fn normalize_sqlserver_url(value: &str) -> Result { + let value = value.trim(); + if value + .get(..JDBC_SQLSERVER_PREFIX.len()) + .is_some_and(|prefix| prefix.eq_ignore_ascii_case(JDBC_SQLSERVER_PREFIX)) + { + return Ok(format!( + "{JDBC_SQLSERVER_PREFIX}{}", + &value[JDBC_SQLSERVER_PREFIX.len()..] + )); + } + if value + .get(..NATIVE_SQLSERVER_PREFIX.len()) + .is_some_and(|prefix| prefix.eq_ignore_ascii_case(NATIVE_SQLSERVER_PREFIX)) + { + return Ok(format!( + "{JDBC_SQLSERVER_PREFIX}{}", + &value[NATIVE_SQLSERVER_PREFIX.len()..] + )); + } + Err(AppError::invalid( + "invalid_sqlserver_connection", + "A valid jdbc:sqlserver:// or sqlserver:// connection URL is required", + )) +} + +fn encode_jdbc_property(value: &str) -> String { + if value.contains([';', '=', '{', '}']) || value.starts_with(' ') || value.ends_with(' ') { + format!("{{{}}}", value.replace('}', "}}")) + } else { + value.to_owned() + } +} + +fn sqlserver_target(value: &str) -> Result<(String, u16), AppError> { + let normalized = normalize_sqlserver_url(value)?; + let authority = normalized + .strip_prefix("jdbc:sqlserver://") + .expect("normalization establishes the SQL Server prefix") + .split(';') + .next() + .unwrap_or_default(); + if authority.is_empty() || authority.contains('\\') { + return Err(AppError::invalid( + "invalid_sqlserver_ssh_url", + "SQL Server SSH forwarding requires an explicit host and TCP port", + )); + } + if let Some(rest) = authority.strip_prefix('[') { + let (host, port) = rest + .split_once("]: ") + .or_else(|| rest.split_once("]:")) + .ok_or_else(invalid_sqlserver_ssh_url)?; + return Ok((host.to_owned(), parse_sqlserver_port(port)?)); + } + let (host, port) = authority + .rsplit_once(':') + .ok_or_else(invalid_sqlserver_ssh_url)?; + if host.trim().is_empty() || host.contains(':') { + return Err(invalid_sqlserver_ssh_url()); + } + Ok((host.to_owned(), parse_sqlserver_port(port)?)) +} + +fn parse_sqlserver_port(value: &str) -> Result { + value + .parse::() + .ok() + .filter(|port| *port != 0) + .ok_or_else(invalid_sqlserver_ssh_url) +} + +fn invalid_sqlserver_ssh_url() -> AppError { + AppError::invalid( + "invalid_sqlserver_ssh_url", + "SQL Server SSH forwarding requires an explicit host and TCP port", + ) +} + +fn sqlserver_connection_failed() -> AppError { + AppError::unavailable( + "sqlserver_connection_failed", + "The SQL Server instance could not be reached or rejected the connection", + ) +} + +fn sqlserver_driver_failure() -> AppError { + AppError::unavailable( + "sqlserver_driver_failure", + "The native SQL Server driver could not decode the server response safely", + ) +} + +fn sqlserver_query_error(error: impl std::fmt::Display) -> AppError { + AppError::invalid("sqlserver_query_failed", error.to_string()) +} + +fn metadata_timeout() -> AppError { + AppError::unavailable( + "sqlserver_metadata_timeout", + "The SQL Server metadata query did not finish in time", + ) +} + +fn validate_identifier(value: &str, field: &str) -> Result<(), AppError> { + if value.trim().is_empty() || value.len() > MAX_IDENTIFIER_BYTES || value.contains('\0') { + return Err(AppError::invalid( + "invalid_sqlserver_metadata_request", + format!("{field} is invalid"), + )); + } + Ok(()) +} + +fn quote_identifier(value: &str, field: &str) -> Result { + validate_identifier(value, field)?; + Ok(format!("[{}]", value.replace(']', "]]"))) +} + +fn qualified_table( + database_name: &str, + schema_name: &str, + table_name: &str, +) -> Result { + Ok(format!( + "{}.{}.{}", + quote_identifier(database_name, "databaseName")?, + quote_identifier(schema_name, "schemaName")?, + quote_identifier(table_name, "tableName")? + )) +} + +fn is_read_candidate(sql: &str) -> Result { + Ok(matches!( + sql_lexemes(sql)?.words.first().map(String::as_str), + Some("SELECT" | "WITH") + )) +} + +fn validate_query(query: &PreparedQuery) -> Result<(), AppError> { + if query.sql.len() > MAX_SQL_BYTES { + return Err(AppError::invalid( + "invalid_query_request", + format!("SQL cannot exceed {MAX_SQL_BYTES} UTF-8 bytes"), + )); + } + validate_query_options(query.options)?; + ordered_parameters(&query.parameters)?; + validate_read_sql(&query.sql) +} + +fn validate_read_sql(sql: &str) -> Result<(), AppError> { + let lexemes = sql_lexemes(sql)?; + if !matches!( + lexemes.words.first().map(String::as_str), + Some("SELECT" | "WITH") + ) { + return Err(AppError::invalid( + "sqlserver_native_query_unsupported", + "Native SQL Server read queries must be a single SELECT statement", + )); + } + if lexemes.words.iter().any(|word| { + matches!( + word.as_str(), + "INSERT" | "UPDATE" | "DELETE" | "MERGE" | "INTO" + ) + }) { + return Err(AppError::invalid( + "sqlserver_native_query_unsupported", + "Data-changing SQL is not allowed on the native SQL Server read path", + )); + } + let statements = Parser::parse_sql(&MsSqlDialect {}, sql).map_err(|_| { + AppError::invalid( + "sqlserver_native_query_unsupported", + "The SQL Server read query could not be parsed safely", + ) + })?; + let [Statement::Query(query)] = statements.as_slice() else { + return Err(AppError::invalid( + "sqlserver_native_query_unsupported", + "Native SQL Server read queries must contain exactly one SELECT statement", + )); + }; + if !query_tree_is_read_only(query) { + return Err(AppError::invalid( + "sqlserver_native_query_unsupported", + "Data-changing CTEs are not allowed on the native SQL Server read path", + )); + } + Ok(()) +} + +fn query_tree_is_read_only(query: &SqlQuery) -> bool { + query.with.as_ref().is_none_or(|with| { + with.cte_tables + .iter() + .all(|cte| query_tree_is_read_only(&cte.query)) + }) && set_expr_is_read_only(&query.body) +} + +fn set_expr_is_read_only(expression: &SetExpr) -> bool { + match expression { + SetExpr::Select(_) => true, + SetExpr::Query(query) => query_tree_is_read_only(query), + SetExpr::SetOperation { left, right, .. } => { + set_expr_is_read_only(left) && set_expr_is_read_only(right) + } + SetExpr::Insert(_) + | SetExpr::Update(_) + | SetExpr::Delete(_) + | SetExpr::Merge(_) + | SetExpr::Values(_) + | SetExpr::Table(_) => false, + } +} + +fn validate_query_options(options: QueryExecutionOptions) -> Result<(), AppError> { + if options.target_batch_rows > MAX_BATCH_ROWS { + return Err(AppError::invalid( + "invalid_query_limits", + format!("batchRows must be at most {MAX_BATCH_ROWS}"), + )); + } + if options.target_batch_bytes != 0 + && !(1024..=MAX_BATCH_BYTES).contains(&options.target_batch_bytes) + { + return Err(AppError::invalid( + "invalid_query_limits", + format!("batchBytes must be zero or between 1024 and {MAX_BATCH_BYTES}"), + )); + } + if options.max_result_bytes > MAX_RESULT_BYTES { + return Err(AppError::invalid( + "invalid_query_limits", + format!("maxResultBytes must be at most {MAX_RESULT_BYTES}"), + )); + } + Ok(()) +} + +fn ordered_parameters(parameters: &[QueryParameter]) -> Result, AppError> { + if parameters.len() > MAX_PARAMETERS { + return Err(AppError::invalid( + "invalid_query_parameter_count", + format!("SQL Server queries accept at most {MAX_PARAMETERS} parameters"), + )); + } + let mut ordered = parameters.iter().collect::>(); + ordered.sort_unstable_by_key(|parameter| parameter.position); + for (index, parameter) in ordered.iter().enumerate() { + let expected = u32::try_from(index + 1).map_err(|_| AppError::internal())?; + if parameter.position != expected { + return Err(AppError::invalid( + "invalid_query_parameter", + "SQL Server parameter positions must be unique and contiguous from 1", + )); + } + validate_parameter(¶meter.value)?; + } + Ok(ordered + .into_iter() + .map(|parameter| ¶meter.value) + .collect()) +} + +fn validate_parameter(value: &DatabaseValue) -> Result<(), AppError> { + let length = match value { + DatabaseValue::Decimal(value) + | DatabaseValue::Text(value) + | DatabaseValue::Date(value) + | DatabaseValue::Time(value) + | DatabaseValue::Timestamp(value) + | DatabaseValue::TimestampWithTimeZone(value) + | DatabaseValue::Json(value) + | DatabaseValue::Uuid(value) => value.len(), + DatabaseValue::Binary(value) => value.len(), + _ => 0, + }; + if length > MAX_SCALAR_BYTES { + return Err(AppError::invalid( + "invalid_query_parameter", + format!("A SQL Server parameter exceeds {MAX_SCALAR_BYTES} bytes"), + )); + } + match value { + DatabaseValue::UnsignedInteger(value) if *value > i64::MAX as u64 => { + Err(AppError::invalid( + "invalid_query_parameter", + "SQL Server does not support unsigned integers larger than BIGINT", + )) + } + DatabaseValue::Decimal(value) => parse_decimal(value).map(|_| ()), + DatabaseValue::Date(value) => NaiveDate::parse_from_str(value, "%Y-%m-%d") + .map(|_| ()) + .map_err(|_| invalid_temporal_parameter("date")), + DatabaseValue::Time(value) => NaiveTime::parse_from_str(value, "%H:%M:%S%.f") + .map(|_| ()) + .map_err(|_| invalid_temporal_parameter("time")), + DatabaseValue::Timestamp(value) => parse_timestamp(value).map(|_| ()), + DatabaseValue::TimestampWithTimeZone(value) => DateTime::parse_from_rfc3339(value) + .map(|_| ()) + .map_err(|_| invalid_temporal_parameter("timestamp with time zone")), + DatabaseValue::Uuid(value) => tiberius::Uuid::parse_str(value).map(|_| ()).map_err(|_| { + AppError::invalid( + "invalid_query_parameter", + "The SQL Server UUID parameter is invalid", + ) + }), + _ => Ok(()), + } +} + +fn parse_timestamp(value: &str) -> Result { + ["%Y-%m-%dT%H:%M:%S%.f", "%Y-%m-%d %H:%M:%S%.f"] + .into_iter() + .find_map(|format| NaiveDateTime::parse_from_str(value, format).ok()) + .ok_or_else(|| invalid_temporal_parameter("timestamp")) +} + +fn invalid_temporal_parameter(label: &str) -> AppError { + AppError::invalid( + "invalid_query_parameter", + format!("The SQL Server {label} parameter is invalid"), + ) +} + +fn parse_decimal(value: &str) -> Result { + let (negative, unsigned) = value + .strip_prefix('-') + .map_or((false, value), |value| (true, value)); + let unsigned = unsigned.strip_prefix('+').unwrap_or(unsigned); + let (whole, fraction) = unsigned + .split_once('.') + .map_or((unsigned, ""), |(whole, fraction)| (whole, fraction)); + if (whole.is_empty() && fraction.is_empty()) + || !whole.bytes().all(|byte| byte.is_ascii_digit()) + || !fraction.bytes().all(|byte| byte.is_ascii_digit()) + || fraction.len() >= 38 + { + return Err(AppError::invalid( + "invalid_query_parameter", + "The SQL Server decimal parameter is invalid", + )); + } + let digits = format!("{whole}{fraction}"); + if digits.len() > 38 { + return Err(AppError::invalid( + "invalid_query_parameter", + "SQL Server decimal parameters cannot exceed 38 digits", + )); + } + let mut integer = digits.parse::().map_err(|_| { + AppError::invalid( + "invalid_query_parameter", + "The SQL Server decimal parameter is invalid", + ) + })?; + if negative { + integer = -integer; + } + Ok(Numeric::new_with_scale( + integer, + u8::try_from(fraction.len()).map_err(|_| AppError::internal())?, + )) +} + +fn bind_query(sql: &str, parameters: &[QueryParameter]) -> Result, AppError> { + let ordered = ordered_parameters(parameters)?; + let sql = rewrite_positional_parameters(sql, ordered.len())?; + let mut query = Query::new(sql); + for value in ordered { + match value { + DatabaseValue::Null => query.bind(Option::::None), + DatabaseValue::Boolean(value) => query.bind(*value), + DatabaseValue::SignedInteger(value) => query.bind(*value), + DatabaseValue::UnsignedInteger(value) => { + query.bind(i64::try_from(*value).map_err(|_| { + AppError::invalid( + "invalid_query_parameter", + "SQL Server does not support unsigned integers larger than BIGINT", + ) + })?); + } + DatabaseValue::Float32(value) => query.bind(*value), + DatabaseValue::Float64(value) => query.bind(*value), + DatabaseValue::Decimal(value) => query.bind(parse_decimal(value)?), + DatabaseValue::Text(value) | DatabaseValue::Json(value) => query.bind(value.clone()), + DatabaseValue::Binary(value) => query.bind(value.clone()), + DatabaseValue::Date(value) => query.bind( + NaiveDate::parse_from_str(value, "%Y-%m-%d") + .map_err(|_| invalid_temporal_parameter("date"))?, + ), + DatabaseValue::Time(value) => query.bind( + NaiveTime::parse_from_str(value, "%H:%M:%S%.f") + .map_err(|_| invalid_temporal_parameter("time"))?, + ), + DatabaseValue::Timestamp(value) => query.bind(parse_timestamp(value)?), + DatabaseValue::TimestampWithTimeZone(value) => query.bind( + DateTime::::parse_from_rfc3339(value) + .map_err(|_| invalid_temporal_parameter("timestamp with time zone"))?, + ), + DatabaseValue::Uuid(value) => { + query.bind(tiberius::Uuid::parse_str(value).map_err(|_| { + AppError::invalid( + "invalid_query_parameter", + "The SQL Server UUID parameter is invalid", + ) + })?); + } + } + } + Ok(query) +} + +fn described_query( + sql: &str, + parameters: &[QueryParameter], +) -> Result<(String, Option), AppError> { + let ordered = ordered_parameters(parameters)?; + let sql = rewrite_positional_parameters(sql, ordered.len())?; + if ordered.is_empty() { + return Ok((sql, None)); + } + let declarations = ordered + .into_iter() + .enumerate() + .map(|(index, value)| { + Ok(format!( + "@P{} {}", + index + 1, + sqlserver_parameter_declaration(value)? + )) + }) + .collect::, AppError>>()? + .join(", "); + Ok((sql, Some(declarations))) +} + +fn sqlserver_parameter_declaration(value: &DatabaseValue) -> Result { + Ok(match value { + DatabaseValue::Null | DatabaseValue::Text(_) | DatabaseValue::Json(_) => { + "nvarchar(max)".to_owned() + } + DatabaseValue::Boolean(_) => "bit".to_owned(), + DatabaseValue::SignedInteger(_) | DatabaseValue::UnsignedInteger(_) => "bigint".to_owned(), + DatabaseValue::Float32(_) => "real".to_owned(), + DatabaseValue::Float64(_) => "float".to_owned(), + DatabaseValue::Decimal(value) => { + format!("decimal(38,{})", parse_decimal(value)?.scale()) + } + DatabaseValue::Binary(_) => "varbinary(max)".to_owned(), + DatabaseValue::Date(_) => "date".to_owned(), + DatabaseValue::Time(_) => "time(7)".to_owned(), + DatabaseValue::Timestamp(_) => "datetime2(7)".to_owned(), + DatabaseValue::TimestampWithTimeZone(_) => "datetimeoffset(7)".to_owned(), + DatabaseValue::Uuid(_) => "uniqueidentifier".to_owned(), + }) +} + +async fn validate_result_set( + client: &mut SqlServerClient, + sql: &str, + parameter_declarations: Option<&str>, +) -> Result<(), AppError> { + let mut query = Query::new( + "SELECT CONVERT(int, system_type_id) AS system_type_id, system_type_name, \ + user_type_name, CONVERT(bigint, max_length) AS max_length, error_number, error_message \ + FROM sys.dm_exec_describe_first_result_set(@P1, @P2, 0) \ + ORDER BY column_ordinal", + ); + query.bind(sql.to_owned()); + query.bind(parameter_declarations.map(str::to_owned)); + let stream = AssertUnwindSafe(query.query(client)) + .catch_unwind() + .await + .map_err(|_| sqlserver_driver_failure())? + .map_err(sqlserver_query_error)?; + let rows = AssertUnwindSafe(stream.into_first_result()) + .catch_unwind() + .await + .map_err(|_| sqlserver_driver_failure())? + .map_err(sqlserver_query_error)?; + for row in rows { + if row + .try_get::(4) + .map_err(sqlserver_query_error)? + .is_some() + { + return Err(AppError::invalid( + "sqlserver_result_description_failed", + "SQL Server could not describe the statement result safely", + )); + } + let system_type_id = row.try_get::(0).map_err(sqlserver_query_error)?; + let system_type_name = row + .try_get::<&str, _>(1) + .map_err(sqlserver_query_error)? + .unwrap_or_default(); + let user_type_name = row + .try_get::<&str, _>(2) + .map_err(sqlserver_query_error)? + .unwrap_or_default(); + validate_described_result_type(system_type_id, system_type_name, user_type_name)?; + if row + .try_get::(3) + .map_err(sqlserver_query_error)? + .is_some_and(|length| length > i64::try_from(MAX_SCALAR_BYTES).unwrap_or(i64::MAX)) + { + return Err(resource_error( + "sqlserver_scalar_too_large", + format!("A SQL Server scalar cannot exceed {MAX_SCALAR_BYTES} bytes"), + )); + } + } + Ok(()) +} + +enum ControlledResultSetValidationError { + Cancelled(Option), + TimedOut(AppError), + Failed(AppError), +} + +async fn validate_result_set_with_control( + client: &mut SqlServerClient, + sql: &str, + parameter_declarations: Option<&str>, + cancellation: &mut watch::Receiver, +) -> Result<(), ControlledResultSetValidationError> { + let validation = validate_result_set(client, sql, parameter_declarations); + tokio::pin!(validation); + let deadline = tokio::time::sleep(METADATA_TIMEOUT); + tokio::pin!(deadline); + let mut cancellation_open = true; + loop { + tokio::select! { + biased; + changed = cancellation.changed(), if cancellation_open => { + if changed.is_err() { + cancellation_open = false; + continue; + } + if let CancellationRequest::Requested { reason } = cancellation.borrow().clone() { + return Err(ControlledResultSetValidationError::Cancelled(reason)); + } + } + () = &mut deadline => { + return Err(ControlledResultSetValidationError::TimedOut( + AppError::unavailable( + "sqlserver_result_description_timeout", + "SQL Server did not describe the statement result before the deadline", + ), + )); + } + result = &mut validation => { + return result.map_err(ControlledResultSetValidationError::Failed); + } + } + } +} + +fn validate_described_result_type( + system_type_id: Option, + system_type_name: &str, + user_type_name: &str, +) -> Result<(), AppError> { + let normalized = system_type_name.trim().to_ascii_lowercase(); + if system_type_id == Some(240) + || matches!(normalized.as_str(), "money" | "smallmoney" | "sql_variant") + { + let type_name = if user_type_name.trim().is_empty() { + system_type_name + } else { + user_type_name + }; + return Err(unsupported_result_type(type_name)); + } + Ok(()) +} + +fn unsupported_result_type(type_name: &str) -> AppError { + AppError::invalid( + "sqlserver_result_type_unsupported", + format!("The native SQL Server driver cannot safely decode result type {type_name}"), + ) +} + +fn rewrite_positional_parameters(sql: &str, parameter_count: usize) -> Result { + let mut output = String::with_capacity(sql.len() + parameter_count.saturating_mul(3)); + let mut chars = sql.char_indices().peekable(); + let mut replaced = 0_usize; + while let Some((_, ch)) = chars.next() { + match ch { + '\'' => copy_quoted(&mut output, &mut chars, '\'', '\''), + '"' => copy_quoted(&mut output, &mut chars, '"', '"'), + '[' => copy_quoted(&mut output, &mut chars, '[', ']'), + '-' if chars.peek().is_some_and(|(_, next)| *next == '-') => { + output.push('-'); + output.push('-'); + chars.next(); + for (_, next) in chars.by_ref() { + output.push(next); + if next == '\n' { + break; + } + } + } + '/' if chars.peek().is_some_and(|(_, next)| *next == '*') => { + output.push('/'); + output.push('*'); + chars.next(); + let mut previous = '\0'; + for (_, next) in chars.by_ref() { + output.push(next); + if previous == '*' && next == '/' { + break; + } + previous = next; + } + } + '?' => { + replaced += 1; + output.push_str("@P"); + output.push_str(&replaced.to_string()); + } + _ => output.push(ch), + } + } + if replaced != 0 && replaced != parameter_count { + return Err(AppError::invalid( + "invalid_query_parameter_count", + format!( + "The SQL Server statement has {replaced} positional markers but {parameter_count} parameters were supplied" + ), + )); + } + Ok(output) +} + +fn copy_quoted( + output: &mut String, + chars: &mut std::iter::Peekable, + opening: char, + closing: char, +) where + I: Iterator, +{ + output.push(opening); + while let Some((_, ch)) = chars.next() { + output.push(ch); + if ch == closing { + if chars.peek().is_some_and(|(_, next)| *next == closing) { + output.push(closing); + chars.next(); + } else { + break; + } + } + } +} + +struct SqlLexemes { + words: Vec, +} + +fn sql_lexemes(sql: &str) -> Result { + if sql.trim().is_empty() { + return Err(AppError::invalid( + "invalid_query_request", + "SQL cannot be empty", + )); + } + let mut words = Vec::new(); + let mut word = String::new(); + let mut chars = sql.chars().peekable(); + let mut statement_terminators = 0_usize; + while let Some(ch) = chars.next() { + if ch.is_ascii_alphanumeric() || ch == '_' { + word.push(ch.to_ascii_uppercase()); + continue; + } + if !word.is_empty() { + words.push(std::mem::take(&mut word)); + } + match ch { + '\'' | '"' => skip_quoted_chars(&mut chars, ch), + '[' => skip_quoted_chars(&mut chars, ']'), + '-' if chars.peek().is_some_and(|next| *next == '-') => { + chars.next(); + for next in chars.by_ref() { + if next == '\n' { + break; + } + } + } + '/' if chars.peek().is_some_and(|next| *next == '*') => { + chars.next(); + let mut previous = '\0'; + let mut closed = false; + for next in chars.by_ref() { + if previous == '*' && next == '/' { + closed = true; + break; + } + previous = next; + } + if !closed { + return Err(AppError::invalid( + "sqlserver_native_query_unsupported", + "The SQL Server statement contains an unterminated comment", + )); + } + } + ';' => statement_terminators += 1, + _ => {} + } + } + if !word.is_empty() { + words.push(word); + } + if statement_terminators > 1 { + return Err(AppError::invalid( + "sqlserver_native_query_unsupported", + "Native SQL Server accepts exactly one statement", + )); + } + Ok(SqlLexemes { words }) +} + +fn skip_quoted_chars(chars: &mut std::iter::Peekable, closing: char) +where + I: Iterator, +{ + while let Some(ch) = chars.next() { + if ch == closing { + if chars.peek().is_some_and(|next| *next == closing) { + chars.next(); + } else { + break; + } + } + } +} + +#[allow( + clippy::too_many_lines, + reason = "the query stream owns one connection, cancellation session, and retained writer lifecycle" +)] +async fn execute_query_task( + application: &Application, + operation_id: &str, + mut cancellation: watch::Receiver, + query: PreparedQuery, + storage: Storage, + resolved: ResolvedDatasourceConnection, +) -> Result { + if let CancellationRequest::Requested { reason } = cancellation.borrow().clone() { + return Err(QueryTaskError::Cancelled(reason)); + } + validate_query(&query)?; + let (described_sql, parameter_declarations) = described_query(&query.sql, &query.parameters)?; + let bound = bind_query(&query.sql, &query.parameters)?; + let open = open_resolved_connection(&resolved, None); + tokio::pin!(open); + let mut conn = loop { + tokio::select! { + biased; + changed = cancellation.changed() => { + if changed.is_ok() + && let CancellationRequest::Requested { reason } = cancellation.borrow().clone() + { + return Err(QueryTaskError::Cancelled(reason)); + } + } + result = &mut open => break result?, + } + }; + let validation = validate_result_set_with_control( + &mut conn.client, + &described_sql, + parameter_declarations.as_deref(), + &mut cancellation, + ) + .await; + match validation { + Ok(()) => {} + Err(ControlledResultSetValidationError::Cancelled(reason)) => { + discard_connection(conn).await; + return Err(QueryTaskError::Cancelled(reason)); + } + Err( + ControlledResultSetValidationError::TimedOut(error) + | ControlledResultSetValidationError::Failed(error), + ) => { + discard_connection(conn).await; + return Err(error.into()); + } + } + let cancellation_request = cancellation.borrow().clone(); + if let CancellationRequest::Requested { reason } = cancellation_request { + discard_connection(conn).await; + return Err(QueryTaskError::Cancelled(reason)); + } + let opened: Result<_, QueryOpenError> = { + let query_open = bound.query(&mut conn.client); + tokio::pin!(query_open); + loop { + tokio::select! { + biased; + changed = cancellation.changed() => { + if changed.is_ok() + && let CancellationRequest::Requested { reason } = cancellation.borrow().clone() + { + break Err(QueryOpenError::Cancelled(reason)); + } + } + result = &mut query_open => { + break result.map_err(sqlserver_query_error).map_err(QueryOpenError::Failed); + } + } + } + }; + let mut stream = opened.map_err(|error| match error { + QueryOpenError::Failed(error) => QueryTaskError::Failed(error), + QueryOpenError::Cancelled(reason) => QueryTaskError::Cancelled(reason), + })?; + + let first = next_query_item(&mut stream, &mut cancellation).await; + let metadata = match first { + Ok(Some(QueryItem::Metadata(metadata))) => metadata, + Ok(Some(QueryItem::Row(_)) | None) => { + drop(stream); + discard_connection(conn).await; + return Err(AppError::invalid( + "sqlserver_query_has_no_result_set", + "The SQL Server read query did not return a tabular result set", + ) + .into()); + } + Err(QueryTaskError::Cancelled(reason)) => { + drop(stream); + discard_connection(conn).await; + return Err(QueryTaskError::Cancelled(reason)); + } + Err(error) => { + drop(stream); + discard_connection(conn).await; + return Err(error); + } + }; + let columns = metadata.columns().to_vec(); + if columns.len() > MAX_COLUMNS { + drop(stream); + discard_connection(conn).await; + return Err(resource_error( + "sqlserver_result_too_wide", + format!("SQL Server returned more than {MAX_COLUMNS} columns"), + ) + .into()); + } + let schema_columns = columns + .iter() + .enumerate() + .map(|(index, column)| wire_column(index, column)) + .collect::, _>>(); + let schema = match schema_columns { + Ok(columns) => wire::QueryStarted { columns }, + Err(error) => { + drop(stream); + discard_connection(conn).await; + return Err(error.into()); + } + }; + let mut writer = match RetainedWriter::begin(storage, schema, query.retention).await { + Ok(writer) => writer, + Err(error) => { + drop(stream); + discard_connection(conn).await; + return Err(error.into()); + } + }; + if let Err(error) = application.inner.operations.started(operation_id).await { + drop(stream); + abort_writer(&mut writer).await; + discard_connection(conn).await; + return Err(error.into()); + } + + let max_rows = query.options.max_rows; + let max_result_bytes = if query.options.max_result_bytes == 0 { + DEFAULT_RESULT_BYTES + } else { + query.options.max_result_bytes + }; + let batch_rows = if query.options.target_batch_rows == 0 { + DEFAULT_BATCH_ROWS + } else { + query.options.target_batch_rows + }; + let batch_bytes = if query.options.target_batch_bytes == 0 { + DEFAULT_BATCH_BYTES + } else { + query.options.target_batch_bytes + }; + let mut pending_rows = Vec::new(); + let mut pending_bytes = 0_u64; + let mut row_count = 0_u64; + let mut result_bytes = 0_u64; + let mut requires_discard = false; + let consumption: Result = async { + loop { + let item = next_query_item(&mut stream, &mut cancellation).await?; + let Some(item) = item else { + return Ok(wire::QueryCompleted { + row_count, + truncated_by_max_rows: false, + truncated_by_max_result_bytes: false, + }); + }; + let QueryItem::Row(row) = item else { + return Err(AppError::invalid( + "sqlserver_query_multiple_results", + "Native SQL Server retained queries accept one result set", + ) + .into()); + }; + if max_rows != 0 && row_count >= max_rows { + requires_discard = true; + return Ok(wire::QueryCompleted { + row_count, + truncated_by_max_rows: true, + truncated_by_max_result_bytes: false, + }); + } + let row = wire_row(row)?; + let row_bytes = u64::try_from(row.encoded_len()) + .map_err(|_| QueryTaskError::Failed(AppError::internal()))?; + if result_bytes.saturating_add(row_bytes) > max_result_bytes { + requires_discard = true; + return Ok(wire::QueryCompleted { + row_count, + truncated_by_max_rows: false, + truncated_by_max_result_bytes: true, + }); + } + let entry_bytes = row_batch_entry_bytes(&row)?; + let candidate_bytes = pending_bytes + .saturating_add(if pending_rows.is_empty() { + row_batch_prefix_bytes(row_count) + } else { + 0 + }) + .saturating_add(entry_bytes); + if !pending_rows.is_empty() + && (pending_rows.len() >= usize::try_from(batch_rows).unwrap_or(usize::MAX) + || candidate_bytes > u64::from(batch_bytes)) + { + flush_rows( + application, + operation_id, + &mut writer, + &mut pending_rows, + row_count, + ) + .await?; + pending_bytes = 0; + } + if pending_rows.is_empty() { + pending_bytes = row_batch_prefix_bytes(row_count); + } + pending_rows.push(row); + pending_bytes = pending_bytes.saturating_add(entry_bytes); + row_count = row_count + .checked_add(1) + .ok_or_else(|| QueryTaskError::Failed(AppError::internal()))?; + result_bytes = result_bytes + .checked_add(row_bytes) + .ok_or_else(|| QueryTaskError::Failed(AppError::internal()))?; + } + } + .await; + drop(stream); + + let completion = match consumption { + Ok(completion) => completion, + Err(error) => { + abort_writer(&mut writer).await; + discard_connection(conn).await; + return Err(error); + } + }; + if requires_discard { + discard_connection(conn).await; + } else if let Err(error) = finish_connection(conn, Ok(())).await { + abort_writer(&mut writer).await; + return Err(error.into()); + } + if let Err(error) = flush_rows( + application, + operation_id, + &mut writer, + &mut pending_rows, + row_count, + ) + .await + { + abort_writer(&mut writer).await; + return Err(error); + } + let metadata = match writer.finish(completion).await { + Ok(metadata) => metadata, + Err(error) => { + abort_writer(&mut writer).await; + return Err(error.into()); + } + }; + Ok(metadata) +} + +async fn next_query_item( + stream: &mut tiberius::QueryStream<'_>, + cancellation: &mut watch::Receiver, +) -> Result, QueryTaskError> { + let mut cancellation_open = true; + loop { + tokio::select! { + biased; + changed = cancellation.changed(), if cancellation_open => { + if changed.is_err() { + cancellation_open = false; + continue; + } + if let CancellationRequest::Requested { reason } = cancellation.borrow().clone() { + return Err(QueryTaskError::Cancelled(reason)); + } + } + item = AssertUnwindSafe(stream.try_next()).catch_unwind() => { + return match item { + Ok(item) => item.map_err(sqlserver_query_error).map_err(QueryTaskError::from), + Err(_) => Err(QueryTaskError::Failed(sqlserver_driver_failure())), + }; + } + } + } +} + +fn wire_column(index: usize, column: &Column) -> Result { + validate_supported_column_type(column.column_type())?; + let ordinal = u32::try_from(index) + .ok() + .and_then(|index| index.checked_add(1)) + .ok_or_else(AppError::internal)?; + let value_type = sqlserver_value_type(column.column_type()); + Ok(wire::JdbcColumn { + ordinal, + name: column.name().to_owned(), + label: column.name().to_owned(), + jdbc_type: sqlserver_jdbc_type(column.column_type()), + jdbc_type_name: sqlserver_type_name(column.column_type()).to_owned(), + value_type: value_type as i32, + nullability: wire::ColumnNullability::Unknown as i32, + precision: None, + scale: None, + display_size: None, + signed: sqlserver_numeric_type(column.column_type()).then_some(true), + catalog_name: None, + schema_name: None, + table_name: None, + }) +} + +fn wire_row(row: Row) -> Result { + Ok(wire::JdbcRow { + values: row + .into_iter() + .map(wire_value) + .collect::, _>>()?, + }) +} + +fn wire_value(value: ColumnData<'static>) -> Result { + use wire::jdbc_value::Value as WireValue; + let value = match value { + ColumnData::U8(None) + | ColumnData::I16(None) + | ColumnData::I32(None) + | ColumnData::I64(None) + | ColumnData::F32(None) + | ColumnData::F64(None) + | ColumnData::Bit(None) + | ColumnData::String(None) + | ColumnData::Guid(None) + | ColumnData::Binary(None) + | ColumnData::Numeric(None) + | ColumnData::Xml(None) + | ColumnData::DateTime(None) + | ColumnData::SmallDateTime(None) + | ColumnData::Time(None) + | ColumnData::Date(None) + | ColumnData::DateTime2(None) + | ColumnData::DateTimeOffset(None) => WireValue::NullValue(wire::JdbcNull {}), + ColumnData::U8(Some(value)) => WireValue::UnsignedIntegerValue(u64::from(value)), + ColumnData::I16(Some(value)) => WireValue::SignedIntegerValue(i64::from(value)), + ColumnData::I32(Some(value)) => WireValue::SignedIntegerValue(i64::from(value)), + ColumnData::I64(Some(value)) => WireValue::SignedIntegerValue(value), + ColumnData::F32(Some(value)) => WireValue::Float32Value(value), + ColumnData::F64(Some(value)) => WireValue::Float64Value(value), + ColumnData::Bit(Some(value)) => WireValue::BooleanValue(value), + ColumnData::String(Some(value)) => { + let value = value.into_owned(); + validate_scalar_bytes(value.len())?; + WireValue::TextValue(value) + } + ColumnData::Guid(Some(value)) => WireValue::UuidValue(value.to_string()), + ColumnData::Binary(Some(value)) => { + let value = value.into_owned(); + validate_scalar_bytes(value.len())?; + WireValue::BinaryValue(value) + } + ColumnData::Numeric(Some(value)) => WireValue::DecimalValue(format_numeric(value)), + ColumnData::Xml(Some(value)) => { + let display_value = value.to_string(); + validate_scalar_bytes(display_value.len())?; + WireValue::OpaqueValue(wire::OpaqueValue { + type_name: "xml".to_owned(), + display_value, + }) + } + ColumnData::DateTime(Some(value)) => { + WireValue::TimestampValue(format_datetime(value.days(), value.seconds_fragments())?) + } + ColumnData::SmallDateTime(Some(value)) => WireValue::TimestampValue(format_small_datetime( + value.days(), + value.seconds_fragments(), + )?), + ColumnData::Time(Some(value)) => WireValue::TimeValue(format_time(value)?), + ColumnData::Date(Some(value)) => WireValue::DateValue(format_date(value)?), + ColumnData::DateTime2(Some(value)) => WireValue::TimestampValue(format_datetime2(value)?), + ColumnData::DateTimeOffset(Some(value)) => { + WireValue::TimestampWithTimeZoneValue(format_datetime_offset(value)?) + } + }; + Ok(wire::JdbcValue { value: Some(value) }) +} + +fn validate_scalar_bytes(length: usize) -> Result<(), AppError> { + if length > MAX_SCALAR_BYTES { + return Err(resource_error( + "sqlserver_scalar_too_large", + format!("A SQL Server scalar cannot exceed {MAX_SCALAR_BYTES} bytes"), + )); + } + Ok(()) +} + +fn validate_supported_column_type(column_type: ColumnType) -> Result<(), AppError> { + if matches!( + column_type, + ColumnType::Money | ColumnType::Money4 | ColumnType::Udt | ColumnType::SSVariant + ) { + return Err(unsupported_result_type(sqlserver_type_name(column_type))); + } + Ok(()) +} + +fn sqlserver_value_type(column_type: ColumnType) -> wire::JdbcValueType { + match column_type { + ColumnType::Bit | ColumnType::Bitn => wire::JdbcValueType::Boolean, + ColumnType::Int1 => wire::JdbcValueType::UnsignedInteger, + ColumnType::Int2 | ColumnType::Int4 | ColumnType::Int8 | ColumnType::Intn => { + wire::JdbcValueType::SignedInteger + } + ColumnType::Float4 | ColumnType::Money4 => wire::JdbcValueType::Float32, + ColumnType::Float8 | ColumnType::Floatn | ColumnType::Money => wire::JdbcValueType::Float64, + ColumnType::Decimaln | ColumnType::Numericn => wire::JdbcValueType::Decimal, + ColumnType::Guid => wire::JdbcValueType::Uuid, + ColumnType::BigVarBin | ColumnType::BigBinary | ColumnType::Image | ColumnType::Udt => { + wire::JdbcValueType::Binary + } + ColumnType::Daten => wire::JdbcValueType::Date, + ColumnType::Timen => wire::JdbcValueType::Time, + ColumnType::Datetime4 + | ColumnType::Datetime + | ColumnType::Datetimen + | ColumnType::Datetime2 => wire::JdbcValueType::Timestamp, + ColumnType::DatetimeOffsetn => wire::JdbcValueType::TimestampWithTimeZone, + ColumnType::Null => wire::JdbcValueType::Unspecified, + ColumnType::Xml | ColumnType::SSVariant => wire::JdbcValueType::Opaque, + ColumnType::BigVarChar + | ColumnType::BigChar + | ColumnType::NVarchar + | ColumnType::NChar + | ColumnType::Text + | ColumnType::NText => wire::JdbcValueType::Text, + } +} + +fn sqlserver_jdbc_type(column_type: ColumnType) -> i32 { + match column_type { + ColumnType::Bit | ColumnType::Bitn => -7, + ColumnType::Int1 => -6, + ColumnType::Int2 => 5, + ColumnType::Int4 | ColumnType::Intn => 4, + ColumnType::Int8 => -5, + ColumnType::Float4 => 7, + ColumnType::Float8 | ColumnType::Floatn => 8, + ColumnType::Money | ColumnType::Money4 | ColumnType::Decimaln | ColumnType::Numericn => 3, + ColumnType::Guid => -11, + ColumnType::BigVarBin => -3, + ColumnType::BigBinary => -2, + ColumnType::Image => -4, + ColumnType::Daten => 91, + ColumnType::Timen => 92, + ColumnType::Datetime4 + | ColumnType::Datetime + | ColumnType::Datetimen + | ColumnType::Datetime2 => 93, + ColumnType::DatetimeOffsetn => 2014, + ColumnType::BigVarChar => 12, + ColumnType::BigChar => 1, + ColumnType::NVarchar => -9, + ColumnType::NChar => -15, + ColumnType::Text => -1, + ColumnType::NText => -16, + ColumnType::Xml => 2009, + ColumnType::Udt | ColumnType::SSVariant | ColumnType::Null => 1111, + } +} + +fn sqlserver_type_name(column_type: ColumnType) -> &'static str { + match column_type { + ColumnType::Null => "null", + ColumnType::Bit | ColumnType::Bitn => "bit", + ColumnType::Int1 => "tinyint", + ColumnType::Int2 => "smallint", + ColumnType::Int4 | ColumnType::Intn => "int", + ColumnType::Int8 => "bigint", + ColumnType::Datetime4 => "smalldatetime", + ColumnType::Float4 => "real", + ColumnType::Float8 | ColumnType::Floatn => "float", + ColumnType::Money => "money", + ColumnType::Datetime | ColumnType::Datetimen => "datetime", + ColumnType::Money4 => "smallmoney", + ColumnType::Guid => "uniqueidentifier", + ColumnType::Decimaln => "decimal", + ColumnType::Numericn => "numeric", + ColumnType::Daten => "date", + ColumnType::Timen => "time", + ColumnType::Datetime2 => "datetime2", + ColumnType::DatetimeOffsetn => "datetimeoffset", + ColumnType::BigVarBin => "varbinary", + ColumnType::BigVarChar => "varchar", + ColumnType::BigBinary => "binary", + ColumnType::BigChar => "char", + ColumnType::NVarchar => "nvarchar", + ColumnType::NChar => "nchar", + ColumnType::Xml => "xml", + ColumnType::Udt => "udt", + ColumnType::Text => "text", + ColumnType::Image => "image", + ColumnType::NText => "ntext", + ColumnType::SSVariant => "sql_variant", + } +} + +fn sqlserver_numeric_type(column_type: ColumnType) -> bool { + matches!( + column_type, + ColumnType::Int1 + | ColumnType::Int2 + | ColumnType::Int4 + | ColumnType::Int8 + | ColumnType::Intn + | ColumnType::Float4 + | ColumnType::Float8 + | ColumnType::Floatn + | ColumnType::Money + | ColumnType::Money4 + | ColumnType::Decimaln + | ColumnType::Numericn + ) +} + +fn format_numeric(value: Numeric) -> String { + let scale = usize::from(value.scale()); + if scale == 0 { + return value.value().to_string(); + } + let negative = value.value().is_negative(); + let digits = value.value().unsigned_abs().to_string(); + let padded = if digits.len() <= scale { + format!("{}{}", "0".repeat(scale + 1 - digits.len()), digits) + } else { + digits + }; + let split = padded.len() - scale; + format!( + "{}{}.{}", + if negative { "-" } else { "" }, + &padded[..split], + &padded[split..] + ) +} + +fn format_datetime(days: i32, fragments: u32) -> Result { + let date = NaiveDate::from_ymd_opt(1900, 1, 1) + .and_then(|date| date.checked_add_signed(chrono::Duration::days(i64::from(days)))) + .ok_or_else(AppError::internal)?; + let nanos = u64::from(fragments) * 1_000_000_000 / 300; + let time = NaiveTime::from_num_seconds_from_midnight_opt( + u32::try_from(nanos / 1_000_000_000).map_err(|_| AppError::internal())?, + u32::try_from(nanos % 1_000_000_000).map_err(|_| AppError::internal())?, + ) + .ok_or_else(AppError::internal)?; + Ok(format_naive_datetime(NaiveDateTime::new(date, time))) +} + +fn format_small_datetime(days: u16, minutes: u16) -> Result { + let date = NaiveDate::from_ymd_opt(1900, 1, 1) + .and_then(|date| date.checked_add_signed(chrono::Duration::days(i64::from(days)))) + .ok_or_else(AppError::internal)?; + let time = NaiveTime::from_num_seconds_from_midnight_opt(u32::from(minutes) * 60, 0) + .ok_or_else(AppError::internal)?; + Ok(format_naive_datetime(NaiveDateTime::new(date, time))) +} + +fn format_time(value: tiberius::time::Time) -> Result { + let nanos = value + .increments() + .checked_mul(10_u64.pow(u32::from(9_u8.saturating_sub(value.scale())))) + .ok_or_else(AppError::internal)?; + let seconds = nanos / 1_000_000_000; + let nanos = u32::try_from(nanos % 1_000_000_000).map_err(|_| AppError::internal())?; + NaiveTime::from_num_seconds_from_midnight_opt( + u32::try_from(seconds).map_err(|_| AppError::internal())?, + nanos, + ) + .map(|time| time.format("%H:%M:%S%.f").to_string()) + .ok_or_else(AppError::internal) +} + +fn format_date(value: tiberius::time::Date) -> Result { + NaiveDate::from_ymd_opt(1, 1, 1) + .and_then(|date| date.checked_add_signed(chrono::Duration::days(i64::from(value.days())))) + .map(|date| date.format("%Y-%m-%d").to_string()) + .ok_or_else(AppError::internal) +} + +fn datetime2_as_naive(value: tiberius::time::DateTime2) -> Result { + let date = NaiveDate::from_ymd_opt(1, 1, 1) + .and_then(|date| { + date.checked_add_signed(chrono::Duration::days(i64::from(value.date().days()))) + }) + .ok_or_else(AppError::internal)?; + if date > NaiveDate::from_ymd_opt(9999, 12, 31).ok_or_else(AppError::internal)? { + return Err(AppError::internal()); + } + let time = value.time(); + let precision = 9_u32 + .checked_sub(u32::from(time.scale())) + .ok_or_else(AppError::internal)?; + let nanos = time + .increments() + .checked_mul(10_u64.pow(precision)) + .ok_or_else(AppError::internal)?; + let seconds = u32::try_from(nanos / 1_000_000_000).map_err(|_| AppError::internal())?; + let nanos = u32::try_from(nanos % 1_000_000_000).map_err(|_| AppError::internal())?; + let time = NaiveTime::from_num_seconds_from_midnight_opt(seconds, nanos) + .ok_or_else(AppError::internal)?; + Ok(NaiveDateTime::new(date, time)) +} + +fn format_datetime2(value: tiberius::time::DateTime2) -> Result { + datetime2_as_naive(value).map(format_naive_datetime) +} + +fn format_datetime_offset(value: tiberius::time::DateTimeOffset) -> Result { + let offset = value.offset(); + if !(-840..=840).contains(&offset) { + return Err(AppError::internal()); + } + let datetime = datetime2_as_naive(value.datetime2())? + .checked_add_signed(chrono::Duration::minutes(i64::from(offset))) + .ok_or_else(AppError::internal)?; + let minimum_date = NaiveDate::from_ymd_opt(1, 1, 1).ok_or_else(AppError::internal)?; + let maximum_date = NaiveDate::from_ymd_opt(9999, 12, 31).ok_or_else(AppError::internal)?; + if !(minimum_date..=maximum_date).contains(&datetime.date()) { + return Err(AppError::internal()); + } + let datetime = format_naive_datetime(datetime); + let sign = if offset < 0 { '-' } else { '+' }; + let minutes = offset.unsigned_abs(); + Ok(format!( + "{datetime}{sign}{:02}:{:02}", + minutes / 60, + minutes % 60 + )) +} + +fn format_naive_datetime(value: NaiveDateTime) -> String { + value.format("%Y-%m-%dT%H:%M:%S%.f").to_string() +} + +fn row_batch_prefix_bytes(start_row_offset: u64) -> u64 { + if start_row_offset == 0 { + 0 + } else { + 1_u64.saturating_add( + u64::try_from(prost::encoding::encoded_len_varint(start_row_offset)) + .unwrap_or(u64::MAX), + ) + } +} + +fn row_batch_entry_bytes(row: &wire::JdbcRow) -> Result { + let row_bytes = row.encoded_len(); + let length_bytes = prost::encoding::length_delimiter_len(row_bytes); + u64::try_from( + 1_usize + .saturating_add(length_bytes) + .saturating_add(row_bytes), + ) + .map_err(|_| QueryTaskError::Failed(AppError::internal())) +} + +async fn flush_rows( + application: &Application, + operation_id: &str, + writer: &mut RetainedWriter, + rows: &mut Vec, + row_count: u64, +) -> Result<(), QueryTaskError> { + if rows.is_empty() { + return Ok(()); + } + let row_len = u64::try_from(rows.len()).map_err(|_| AppError::internal())?; + let start_row_offset = row_count + .checked_sub(row_len) + .ok_or_else(AppError::internal)?; + let batch = wire::RowBatch { + start_row_offset, + rows: std::mem::take(rows), + }; + if batch.encoded_len() > usize::try_from(MAX_BATCH_BYTES).unwrap_or(usize::MAX) { + return Err(resource_error( + "sqlserver_result_batch_too_large", + "One SQL Server result row exceeds the retained-result batch limit", + ) + .into()); + } + let byte_count = writer.append(batch).await?; + application + .inner + .operations + .progress(operation_id, row_count, byte_count) + .await?; + Ok(()) +} + +async fn abort_writer(writer: &mut RetainedWriter) { + if let Err(error) = writer.abort().await { + tracing::warn!(error = %error, "native SQL Server retained-result cleanup failed"); + } +} + +fn resource_error(code: impl Into, message: impl Into) -> AppError { + AppError::new( + AppErrorKind::ResourceExhausted, + ApiError::new(code, message), + ) +} + +async fn execute_update( + resolved: ResolvedDatasourceConnection, + sql: String, + cancellation: CancellationToken, +) -> Result { + if cancellation.is_cancelled() { + return Err(DatabaseWriteError::not_started( + write_cancelled_before_dispatch(), + )); + } + let sql = validate_single_statement(&sql).map_err(DatabaseWriteError::not_started)?; + if resolved.connection.read_only { + return Err(DatabaseWriteError::not_started(AppError::new( + AppErrorKind::Conflict, + ApiError::new( + "datasource_read_only", + "The datasource connection is configured as read-only", + ), + ))); + } + let open = open_resolved_connection(&resolved, None); + tokio::pin!(open); + let mut conn = tokio::select! { + biased; + () = cancellation.cancelled() => { + return Err(DatabaseWriteError::not_started(write_cancelled_before_dispatch())); + } + result = &mut open => result.map_err(DatabaseWriteError::not_started)?, + }; + if cancellation.is_cancelled() { + discard_connection(conn).await; + return Err(DatabaseWriteError::not_started( + write_cancelled_before_dispatch(), + )); + } + let result = { + let execution = conn.client.execute(sql, &[]); + tokio::pin!(execution); + tokio::select! { + biased; + () = cancellation.cancelled() => None, + result = &mut execution => Some(result), + } + }; + let Some(result) = result else { + discard_connection(conn).await; + return Err(DatabaseWriteError::unknown(write_outcome_unknown( + "The SQL Server write was interrupted after dispatch; do not retry it blindly", + ))); + }; + match result { + Ok(result) => { + let affected = result.total(); + finish_connection(conn, Ok(())) + .await + .map_err(DatabaseWriteError::unknown)?; + Ok(affected) + } + Err(error) => { + tracing::warn!(error = %error, "SQL Server rejected a dispatched write"); + discard_connection(conn).await; + Err(DatabaseWriteError::unknown(write_outcome_unknown( + "SQL Server reported an error after write dispatch; partial effects cannot be excluded, so do not retry it blindly", + ))) + } + } +} + +fn validate_single_statement(sql: &str) -> Result { + if sql.len() > MAX_SQL_BYTES || sql.trim().is_empty() { + return Err(AppError::invalid( + "invalid_database_write", + "SQL must be non-empty and within the configured size limit", + )); + } + let statements = split_sqlserver_script(sql)?; + if statements.len() != 1 { + return Err(AppError::invalid( + "invalid_database_write", + "Exactly one SQL Server statement is required", + )); + } + Ok(statements.into_iter().next().expect("one statement")) +} + +fn write_cancelled_before_dispatch() -> AppError { + AppError::new( + AppErrorKind::Conflict, + ApiError::new( + "database_write_cancelled", + "The database write was cancelled before dispatch", + ), + ) +} + +fn write_outcome_unknown(message: &'static str) -> AppError { + AppError::new( + AppErrorKind::Unavailable, + ApiError::new("database_write_outcome_unknown", message), + ) +} + +enum ConsoleExecutionError { + Cancelled(Option), + ConnectionUnusable(AppError), + WriteOutcomeUnknown, + Statement(AppError), +} + +struct ConsolePending { + id: u32, + started: Instant, + columns: Vec, + rows: Vec, + row_count: u64, + retain: bool, + page_end: u64, +} + +async fn execute_console( + application: &Application, + request: NativeConsoleRequest, + mut cancellation: watch::Receiver, + force_read_only: bool, +) -> Result, AppError> { + let (mut statements, page_offset, page_end) = prepare_console_request(&request)?; + if let CancellationRequest::Requested { reason } = cancellation.borrow().clone() { + return Err(console_cancelled(reason)); + } + let resolved = resolve_native_connection(application, &request.datasource_id).await?; + if force_read_only || resolved.connection.read_only { + for statement in &statements { + validate_read_sql(statement).map_err(|_| { + AppError::new( + AppErrorKind::Conflict, + ApiError::new( + "datasource_read_only", + "The SQL Server datasource accepts read-only Console statements", + ), + ) + })?; + } + } + if request.explain { + statements = statements + .into_iter() + .map(|statement| format!("SET SHOWPLAN_ALL ON; {statement}; SET SHOWPLAN_ALL OFF")) + .collect(); + } + let open = open_resolved_connection( + &resolved, + (!request.database_name.trim().is_empty()).then_some(request.database_name.as_str()), + ); + tokio::pin!(open); + let mut conn = loop { + tokio::select! { + biased; + changed = cancellation.changed() => { + if changed.is_ok() + && let CancellationRequest::Requested { reason } = cancellation.borrow().clone() + { + return Err(console_cancelled(reason)); + } + } + result = &mut open => break result?, + } + }; + let mut results = Vec::new(); + let mut retained_bytes = 0_u64; + for (index, statement) in statements.into_iter().enumerate() { + let sequence = u32::try_from(index) + .ok() + .and_then(|value| value.checked_add(1)) + .ok_or_else(AppError::internal)?; + let started = Instant::now(); + let execution = execute_console_statement( + &mut conn.client, + &statement, + sequence, + page_offset, + page_end, + request.result_set_id, + &mut retained_bytes, + &mut cancellation, + request.explain, + ) + .await; + match execution { + Ok(mut statement_results) => results.append(&mut statement_results), + Err(ConsoleExecutionError::Statement(error)) => { + results.push(console_failure_result( + sequence, + statement, + &error, + elapsed_millis(started), + )); + if !request.error_continue { + break; + } + } + Err(ConsoleExecutionError::Cancelled(reason)) => { + discard_connection(conn).await; + return Err(console_cancelled(reason)); + } + Err(ConsoleExecutionError::ConnectionUnusable(error)) => { + discard_connection(conn).await; + return Err(error); + } + Err(ConsoleExecutionError::WriteOutcomeUnknown) => { + discard_connection(conn).await; + return Err(write_outcome_unknown( + "The SQL Server Console write was interrupted after dispatch; do not retry it blindly", + )); + } + } + } + finish_connection(conn, Ok(())).await?; + Ok(results) +} + +fn prepare_console_request( + request: &NativeConsoleRequest, +) -> Result<(Vec, u64, u64), AppError> { + if request.sql.trim().is_empty() || request.sql.len() > MAX_SQL_BYTES { + return Err(AppError::invalid( + "invalid_sqlserver_console_request", + "sql must be non-empty and within the configured size limit", + )); + } + if request.page_no == 0 || request.page_size == 0 || request.page_size > MAX_CONSOLE_PAGE_SIZE { + return Err(AppError::invalid( + "invalid_sqlserver_console_request", + format!( + "pageNo and pageSize must be positive, and pageSize cannot exceed {MAX_CONSOLE_PAGE_SIZE}" + ), + )); + } + let statements = if request.single { + vec![request.sql.trim().to_owned()] + } else { + split_sqlserver_script(&request.sql)? + }; + if statements.is_empty() || statements.len() > MAX_CONSOLE_STATEMENTS { + return Err(AppError::invalid( + "invalid_sqlserver_console_request", + format!( + "A Console request must contain between 1 and {MAX_CONSOLE_STATEMENTS} statements" + ), + )); + } + let page_size = if request.page_size_all { + u64::from(MAX_CONSOLE_PAGE_SIZE) + } else { + u64::from(request.page_size) + }; + let page_offset = if request.page_size_all { + 0 + } else { + u64::from(request.page_no - 1) + .checked_mul(page_size) + .ok_or_else(AppError::internal)? + }; + let page_end = page_offset + .checked_add(page_size) + .ok_or_else(AppError::internal)?; + Ok((statements, page_offset, page_end)) +} + +#[allow(clippy::too_many_arguments)] +async fn execute_console_statement( + client: &mut SqlServerClient, + statement: &str, + statement_sequence: u32, + page_offset: u64, + page_end: u64, + selected_result_set_id: Option, + retained_bytes: &mut u64, + cancellation: &mut watch::Receiver, + force_tabular: bool, +) -> Result, ConsoleExecutionError> { + if let CancellationRequest::Requested { reason } = cancellation.borrow().clone() { + return Err(ConsoleExecutionError::Cancelled(reason)); + } + let read_only = validate_read_sql(statement).is_ok(); + let creates_local_temp_table = creates_local_temp_table(statement); + if force_tabular || is_tabular_statement(statement) || creates_local_temp_table { + return execute_console_query( + client, + statement, + statement_sequence, + page_offset, + page_end, + selected_result_set_id, + retained_bytes, + cancellation, + read_only, + !force_tabular && !creates_local_temp_table, + ) + .await; + } + let started = Instant::now(); + let execution = AssertUnwindSafe(client.execute(statement, &[])).catch_unwind(); + tokio::pin!(execution); + let result = loop { + tokio::select! { + biased; + changed = cancellation.changed() => { + if changed.is_ok() + && let CancellationRequest::Requested { reason } = cancellation.borrow().clone() + { + return Err(console_execution_interrupted(read_only, reason)); + } + } + result = &mut execution => { + break result.map_err(|_| console_driver_interrupted(read_only))?; + } + } + } + .map_err(|error| console_dispatched_driver_error(read_only, error))?; + Ok(vec![NativeConsoleResult { + statement_sequence, + result_set_id: None, + sql: statement.to_owned(), + success: true, + message: "Statement executed successfully".to_owned(), + update_count: result.total(), + columns: Vec::new(), + rows: Vec::new(), + row_count: 0, + has_more: false, + duration_ms: elapsed_millis(started), + error: None, + }]) +} + +#[allow(clippy::too_many_arguments, clippy::too_many_lines)] +async fn execute_console_query( + client: &mut SqlServerClient, + statement: &str, + statement_sequence: u32, + page_offset: u64, + page_end: u64, + selected_result_set_id: Option, + retained_bytes: &mut u64, + cancellation: &mut watch::Receiver, + read_only: bool, + preflight_result: bool, +) -> Result, ConsoleExecutionError> { + let statement_started = Instant::now(); + if preflight_result { + match validate_result_set_with_control(client, statement, None, cancellation).await { + Ok(()) => {} + Err(ControlledResultSetValidationError::Cancelled(reason)) => { + return Err(ConsoleExecutionError::Cancelled(reason)); + } + Err(ControlledResultSetValidationError::TimedOut(error)) => { + return Err(ConsoleExecutionError::ConnectionUnusable(error)); + } + Err(ControlledResultSetValidationError::Failed(error)) => { + return Err(ConsoleExecutionError::Statement(error)); + } + } + } + if let CancellationRequest::Requested { reason } = cancellation.borrow().clone() { + return Err(ConsoleExecutionError::Cancelled(reason)); + } + let open = AssertUnwindSafe(client.simple_query(statement)).catch_unwind(); + tokio::pin!(open); + let mut stream = loop { + tokio::select! { + biased; + changed = cancellation.changed() => { + if changed.is_ok() + && let CancellationRequest::Requested { reason } = cancellation.borrow().clone() + { + return Err(console_execution_interrupted(read_only, reason)); + } + } + result = &mut open => { + break result.map_err(|_| console_driver_interrupted(read_only))?; + } + } + } + .map_err(|error| console_dispatched_driver_error(read_only, error))?; + + let mut results = Vec::new(); + let mut pending: Option = None; + let mut result_set_id = 0_u32; + loop { + let item = loop { + tokio::select! { + biased; + changed = cancellation.changed() => { + if changed.is_ok() + && let CancellationRequest::Requested { reason } = cancellation.borrow().clone() + { + return Err(console_execution_interrupted(read_only, reason)); + } + } + item = AssertUnwindSafe(stream.try_next()).catch_unwind() => { + break item.map_err(|_| console_driver_interrupted(read_only))?; + }, + } + } + .map_err(|error| console_dispatched_driver_error(read_only, error))?; + let Some(item) = item else { + break; + }; + match item { + QueryItem::Metadata(metadata) => { + if let Some(previous) = pending.take() { + push_console_result(&mut results, statement_sequence, statement, previous); + } + result_set_id = result_set_id + .checked_add(1) + .ok_or_else(|| console_post_dispatch_error(read_only, AppError::internal()))?; + let retain = + selected_result_set_id.is_none_or(|selected| selected == result_set_id); + let columns = if retain { + if metadata.columns().len() > MAX_COLUMNS { + return Err(console_post_dispatch_error( + read_only, + resource_error( + "sqlserver_result_too_wide", + format!("SQL Server returned more than {MAX_COLUMNS} columns"), + ), + )); + } + metadata + .columns() + .iter() + .enumerate() + .map(|(index, column)| portable_column(index, column)) + .collect::, _>>() + .map_err(|error| console_post_dispatch_error(read_only, error))? + } else { + Vec::new() + }; + pending = Some(ConsolePending { + id: result_set_id, + started: Instant::now(), + columns, + rows: Vec::new(), + row_count: 0, + retain, + page_end, + }); + } + QueryItem::Row(row) => { + let current = pending + .as_mut() + .ok_or_else(|| console_post_dispatch_error(read_only, AppError::internal()))?; + if current.retain && (page_offset..page_end).contains(¤t.row_count) { + let row = portable_row(row) + .map_err(|error| console_post_dispatch_error(read_only, error))?; + reserve_console_bytes(retained_bytes, &row) + .map_err(|error| console_post_dispatch_error(read_only, error))?; + current.rows.push(row); + } + current.row_count = current + .row_count + .checked_add(1) + .ok_or_else(|| console_post_dispatch_error(read_only, AppError::internal()))?; + } + } + } + if let Some(previous) = pending { + push_console_result(&mut results, statement_sequence, statement, previous); + } + if results.is_empty() && selected_result_set_id.is_none() { + results.push(NativeConsoleResult { + statement_sequence, + result_set_id: None, + sql: statement.to_owned(), + success: true, + message: "Statement executed successfully".to_owned(), + update_count: 0, + columns: Vec::new(), + rows: Vec::new(), + row_count: 0, + has_more: false, + duration_ms: elapsed_millis(statement_started), + error: None, + }); + } + Ok(results) +} + +fn console_execution_interrupted(read_only: bool, reason: Option) -> ConsoleExecutionError { + if read_only { + ConsoleExecutionError::Cancelled(reason) + } else { + ConsoleExecutionError::WriteOutcomeUnknown + } +} + +fn console_driver_interrupted(read_only: bool) -> ConsoleExecutionError { + if read_only { + ConsoleExecutionError::Statement(sqlserver_driver_failure()) + } else { + ConsoleExecutionError::WriteOutcomeUnknown + } +} + +fn console_dispatched_driver_error(read_only: bool, error: TiberiusError) -> ConsoleExecutionError { + if read_only || matches!(error, TiberiusError::Server(_)) { + ConsoleExecutionError::Statement(sqlserver_query_error(error)) + } else { + ConsoleExecutionError::WriteOutcomeUnknown + } +} + +fn console_post_dispatch_error(read_only: bool, error: AppError) -> ConsoleExecutionError { + if read_only { + ConsoleExecutionError::Statement(error) + } else { + ConsoleExecutionError::WriteOutcomeUnknown + } +} + +fn push_console_result( + output: &mut Vec, + statement_sequence: u32, + statement: &str, + pending: ConsolePending, +) { + if !pending.retain { + return; + } + output.push(NativeConsoleResult { + statement_sequence, + result_set_id: Some(pending.id), + sql: statement.to_owned(), + success: true, + message: "Statement executed successfully".to_owned(), + update_count: 0, + columns: pending.columns, + rows: pending.rows, + row_count: pending.row_count, + has_more: pending.row_count > pending.page_end, + duration_ms: elapsed_millis(pending.started), + error: None, + }); +} + +fn portable_column(index: usize, column: &Column) -> Result { + validate_supported_column_type(column.column_type())?; + let ordinal = u32::try_from(index) + .ok() + .and_then(|value| value.checked_add(1)) + .ok_or_else(AppError::internal)?; + Ok(ResultColumn { + ordinal, + label: column.name().to_owned(), + name: column.name().to_owned(), + jdbc_type: sqlserver_jdbc_type(column.column_type()), + jdbc_type_name: sqlserver_type_name(column.column_type()).to_owned(), + value_type: portable_value_type(column.column_type()), + nullability: ColumnNullability::Unknown, + precision: None, + scale: None, + display_size: None, + signed: sqlserver_numeric_type(column.column_type()).then_some(true), + catalog_name: None, + schema_name: None, + table_name: None, + }) +} + +fn portable_value_type(column_type: ColumnType) -> JdbcValueType { + match sqlserver_value_type(column_type) { + wire::JdbcValueType::Boolean => JdbcValueType::Boolean, + wire::JdbcValueType::SignedInteger => JdbcValueType::SignedInteger, + wire::JdbcValueType::UnsignedInteger => JdbcValueType::UnsignedInteger, + wire::JdbcValueType::Float32 => JdbcValueType::Float32, + wire::JdbcValueType::Float64 => JdbcValueType::Float64, + wire::JdbcValueType::Decimal => JdbcValueType::Decimal, + wire::JdbcValueType::Text => JdbcValueType::Text, + wire::JdbcValueType::Binary => JdbcValueType::Binary, + wire::JdbcValueType::Date => JdbcValueType::Date, + wire::JdbcValueType::Time => JdbcValueType::Time, + wire::JdbcValueType::Timestamp => JdbcValueType::Timestamp, + wire::JdbcValueType::TimestampWithTimeZone => JdbcValueType::TimestampWithTimeZone, + wire::JdbcValueType::Json => JdbcValueType::Json, + wire::JdbcValueType::Uuid => JdbcValueType::Uuid, + wire::JdbcValueType::Opaque | wire::JdbcValueType::Unspecified => JdbcValueType::Opaque, + } +} + +fn portable_row(row: Row) -> Result { + Ok(ResultRow { + values: row + .into_iter() + .map(portable_value) + .collect::, _>>()?, + }) +} + +fn portable_value(value: ColumnData<'static>) -> Result { + use wire::jdbc_value::Value as WireValue; + + let value = wire_value(value)?; + let Some(value) = value.value else { + return Err(AppError::internal()); + }; + Ok(match value { + WireValue::NullValue(_) => JdbcValue::Null, + WireValue::BooleanValue(value) => JdbcValue::Boolean { value }, + WireValue::SignedIntegerValue(value) => JdbcValue::SignedInteger { + value: value.to_string(), + }, + WireValue::UnsignedIntegerValue(value) => JdbcValue::UnsignedInteger { + value: value.to_string(), + }, + WireValue::Float32Value(value) => JdbcValue::Float32 { + value: display_f32(value), + }, + WireValue::Float64Value(value) => JdbcValue::Float64 { + value: display_f64(value), + }, + WireValue::DecimalValue(value) => JdbcValue::Decimal { value }, + WireValue::TextValue(value) => JdbcValue::Text { value }, + WireValue::BinaryValue(value) => JdbcValue::Binary { + value: BASE64_STANDARD.encode(value), + }, + WireValue::DateValue(value) => JdbcValue::Date { value }, + WireValue::TimeValue(value) => JdbcValue::Time { value }, + WireValue::TimestampValue(value) => JdbcValue::Timestamp { value }, + WireValue::TimestampWithTimeZoneValue(value) => JdbcValue::TimestampWithTimeZone { value }, + WireValue::JsonValue(value) => JdbcValue::Json { value }, + WireValue::UuidValue(value) => JdbcValue::Uuid { value }, + WireValue::OpaqueValue(value) => JdbcValue::Opaque { + type_name: value.type_name, + display_value: value.display_value, + }, + }) +} + +fn reserve_console_bytes(total: &mut u64, row: &ResultRow) -> Result<(), AppError> { + let bytes = u64::try_from( + serde_json::to_vec(row) + .map_err(|_| AppError::internal())? + .len(), + ) + .map_err(|_| AppError::internal())?; + *total = total.checked_add(bytes).ok_or_else(AppError::internal)?; + if *total > MAX_CONSOLE_RESULT_BYTES { + return Err(resource_error( + "sqlserver_console_result_too_large", + "The retained SQL Server Console result exceeded the configured byte limit", + )); + } + Ok(()) +} + +fn display_f32(value: f32) -> String { + if value.is_nan() { + "NaN".to_owned() + } else if value == f32::INFINITY { + "Infinity".to_owned() + } else if value == f32::NEG_INFINITY { + "-Infinity".to_owned() + } else { + value.to_string() + } +} + +fn display_f64(value: f64) -> String { + if value.is_nan() { + "NaN".to_owned() + } else if value == f64::INFINITY { + "Infinity".to_owned() + } else if value == f64::NEG_INFINITY { + "-Infinity".to_owned() + } else { + value.to_string() + } +} + +fn console_failure_result( + statement_sequence: u32, + sql: String, + error: &AppError, + duration_ms: u64, +) -> NativeConsoleResult { + NativeConsoleResult { + statement_sequence, + result_set_id: None, + sql, + success: false, + message: error.api_error().message.clone(), + update_count: 0, + columns: Vec::new(), + rows: Vec::new(), + row_count: 0, + has_more: false, + duration_ms, + error: Some(error.api_error()), + } +} + +fn console_cancelled(reason: Option) -> AppError { + AppError::new( + AppErrorKind::Conflict, + ApiError::new( + "sqlserver_console_cancelled", + reason.unwrap_or_else(|| "The SQL Server Console execution was cancelled".to_owned()), + ), + ) +} + +fn elapsed_millis(started: Instant) -> u64 { + u64::try_from(started.elapsed().as_millis()).unwrap_or(u64::MAX) +} + +fn is_tabular_statement(statement: &str) -> bool { + sql_lexemes(statement).is_ok_and(|lexemes| { + if is_extended_property_procedure(&lexemes.words) { + return false; + } + matches!( + lexemes.words.first().map(String::as_str), + Some("SELECT" | "WITH" | "EXEC" | "EXECUTE" | "DBCC") + ) + }) +} + +fn is_extended_property_procedure(words: &[String]) -> bool { + let [command, procedure @ ..] = words else { + return false; + }; + if !matches!(command.as_str(), "EXEC" | "EXECUTE") { + return false; + } + let procedure = match procedure { + [system_schema, procedure, ..] if system_schema == "SYS" => procedure.as_str(), + [procedure, ..] => procedure.as_str(), + [] => return false, + }; + matches!( + procedure, + "SP_ADDEXTENDEDPROPERTY" | "SP_UPDATEEXTENDEDPROPERTY" | "SP_DROPEXTENDEDPROPERTY" + ) +} + +fn creates_local_temp_table(statement: &str) -> bool { + let Ok(statements) = Parser::parse_sql(&MsSqlDialect {}, statement) else { + return false; + }; + let [Statement::CreateTable(table)] = statements.as_slice() else { + return false; + }; + table + .name + .0 + .last() + .and_then(sqlparser::ast::ObjectNamePart::as_ident) + .is_some_and(|name| name.value.starts_with('#') && !name.value.starts_with("##")) +} + +fn split_sqlserver_script(script: &str) -> Result, AppError> { + let mut statements = Vec::new(); + let mut current = String::new(); + let mut chars = script.chars().peekable(); + while let Some(ch) = chars.next() { + match ch { + '\'' | '"' => { + current.push(ch); + copy_quoted_chars(&mut current, &mut chars, ch)?; + } + '[' => { + current.push(ch); + copy_quoted_chars(&mut current, &mut chars, ']')?; + } + '-' if chars.peek().is_some_and(|next| *next == '-') => { + current.push('-'); + current.push('-'); + chars.next(); + for next in chars.by_ref() { + current.push(next); + if next == '\n' { + break; + } + } + } + '/' if chars.peek().is_some_and(|next| *next == '*') => { + current.push('/'); + current.push('*'); + chars.next(); + let mut previous = '\0'; + let mut closed = false; + for next in chars.by_ref() { + current.push(next); + if previous == '*' && next == '/' { + closed = true; + break; + } + previous = next; + } + if !closed { + return Err(AppError::invalid( + "invalid_sqlserver_console_request", + "The SQL Server script contains an unterminated comment", + )); + } + } + ';' => push_statement(&mut statements, &mut current), + _ => current.push(ch), + } + } + push_statement(&mut statements, &mut current); + Ok(statements) +} + +fn copy_quoted_chars( + output: &mut String, + chars: &mut std::iter::Peekable, + closing: char, +) -> Result<(), AppError> +where + I: Iterator, +{ + while let Some(ch) = chars.next() { + output.push(ch); + if ch == closing { + if chars.peek().is_some_and(|next| *next == closing) { + output.push(closing); + chars.next(); + } else { + return Ok(()); + } + } + } + Err(AppError::invalid( + "invalid_sqlserver_console_request", + "The SQL Server script contains an unterminated quoted value", + )) +} + +fn push_statement(statements: &mut Vec, current: &mut String) { + let statement = current.trim(); + if !statement.is_empty() { + statements.push(statement.to_owned()); + } + current.clear(); +} + +async fn metadata_rows( + application: &Application, + datasource_id: &str, + database_name: Option<&str>, + sql: &str, + values: Vec, +) -> Result, AppError> { + let resolved = resolve_native_connection(application, datasource_id).await?; + let mut conn = open_resolved_connection(&resolved, database_name).await?; + let parameters = values + .into_iter() + .enumerate() + .map(|(index, value)| QueryParameter { + position: u32::try_from(index + 1).unwrap_or(u32::MAX), + value, + }) + .collect::>(); + let query = bind_query(sql, ¶meters)?; + let result = tokio::time::timeout(METADATA_TIMEOUT, async { + query + .query(&mut conn.client) + .await + .map_err(sqlserver_query_error)? + .into_first_result() + .await + .map_err(sqlserver_query_error) + }) + .await + .map_err(|_| metadata_timeout())?; + finish_connection(conn, result).await +} + +fn row_string(row: &Row, index: usize) -> Result { + Ok(row + .try_get::<&str, _>(index) + .map_err(sqlserver_query_error)? + .unwrap_or_default() + .to_owned()) +} + +fn row_optional_string(row: &Row, index: usize) -> Result, AppError> { + row.try_get::<&str, _>(index) + .map(|value| value.map(ToOwned::to_owned)) + .map_err(sqlserver_query_error) +} + +fn row_i32(row: &Row, index: usize) -> Result { + row.try_get::(index) + .map_err(sqlserver_query_error)? + .ok_or_else(AppError::internal) +} + +fn row_bool(row: &Row, index: usize) -> Result { + row.try_get::(index) + .map_err(sqlserver_query_error)? + .ok_or_else(AppError::internal) +} + +fn row_optional_i32(row: &Row, index: usize) -> Result, AppError> { + row.try_get::(index).map_err(sqlserver_query_error) +} + +async fn list_databases( + application: &Application, + datasource_id: &str, +) -> Result { + let rows = metadata_rows( + application, + datasource_id, + None, + "SELECT d.name, COALESCE(CONVERT(nvarchar(128), DATABASEPROPERTYEX(d.name, 'Collation')), ''), \ + COALESCE(SUSER_SNAME(d.owner_sid), ''), CONVERT(int, d.database_id) \ + FROM sys.databases d WHERE d.state = 0 ORDER BY d.name", + Vec::new(), + ) + .await?; + rows.into_iter() + .map(|row| { + let name = row_string(&row, 0)?; + Ok(DatabaseMetadata { + name, + collation: row_string(&row, 1)?, + owner: row_string(&row, 2)?, + system: row_i32(&row, 3)? <= 4, + ..DatabaseMetadata::default() + }) + }) + .collect::, _>>() + .map(|items| DatabaseList { items }) +} + +async fn list_schemas( + application: &Application, + datasource_id: &str, + database_name: &str, +) -> Result { + validate_identifier(database_name, "databaseName")?; + let rows = metadata_rows( + application, + datasource_id, + Some(database_name), + "SELECT s.name, COALESCE(p.name, ''), CONVERT(int, s.schema_id) \ + FROM sys.schemas s LEFT JOIN sys.database_principals p ON p.principal_id = s.principal_id \ + ORDER BY s.name", + Vec::new(), + ) + .await?; + rows.into_iter() + .map(|row| { + let name = row_string(&row, 0)?; + let schema_id = row_i32(&row, 2)?; + Ok(SchemaMetadata { + database_name: database_name.to_owned(), + system: schema_id <= 4 + || matches!( + name.as_str(), + "sys" | "INFORMATION_SCHEMA" | "db_owner" | "db_accessadmin" + ), + name, + owner: row_string(&row, 1)?, + ..SchemaMetadata::default() + }) + }) + .collect::, _>>() + .map(|items| SchemaList { items }) +} + +fn effective_schema(schema_name: &str) -> &str { + if schema_name.trim().is_empty() { + "dbo" + } else { + schema_name + } +} + +async fn list_tables( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + name_pattern: &str, +) -> Result { + validate_identifier(database_name, "databaseName")?; + let schema_name = effective_schema(schema_name); + validate_identifier(schema_name, "schemaName")?; + let rows = metadata_rows( + application, + datasource_id, + Some(database_name), + "SELECT s.name, t.name, COALESCE(CONVERT(nvarchar(max), ep.value), ''), \ + COALESCE(CONVERT(nvarchar(40), SUM(CASE WHEN p.index_id IN (0,1) THEN p.rows ELSE 0 END)), '0'), \ + COALESCE(CONVERT(nvarchar(40), SUM(CASE WHEN p.index_id IN (0,1) THEN a.total_pages * 8192 ELSE 0 END)), '0'), \ + CONVERT(nvarchar(33), t.create_date, 126), CONVERT(nvarchar(33), t.modify_date, 126) \ + FROM sys.tables t JOIN sys.schemas s ON s.schema_id=t.schema_id \ + LEFT JOIN sys.partitions p ON p.object_id=t.object_id \ + LEFT JOIN sys.allocation_units a ON a.container_id=p.partition_id \ + LEFT JOIN sys.extended_properties ep ON ep.major_id=t.object_id AND ep.minor_id=0 AND ep.name='MS_Description' \ + WHERE s.name=@P1 AND (@P2='' OR t.name LIKE @P2) \ + GROUP BY s.name,t.name,ep.value,t.create_date,t.modify_date ORDER BY t.name", + vec![ + DatabaseValue::Text(schema_name.to_owned()), + DatabaseValue::Text(name_pattern.trim().to_owned()), + ], + ) + .await?; + rows.into_iter() + .map(|row| { + Ok(TableMetadata { + database_name: database_name.to_owned(), + schema_name: row_string(&row, 0)?, + name: row_string(&row, 1)?, + table_type: "TABLE".to_owned(), + comment: row_string(&row, 2)?, + database_type: "SQLSERVER".to_owned(), + rows: Some(row_string(&row, 3)?), + data_length: Some(row_string(&row, 4)?), + create_time: row_string(&row, 5)?, + update_time: row_string(&row, 6)?, + ..TableMetadata::default() + }) + }) + .collect::, _>>() + .map(|items| TableList { items }) +} + +async fn list_columns( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + table_name: &str, +) -> Result { + validate_identifier(database_name, "databaseName")?; + let schema_name = effective_schema(schema_name); + validate_identifier(schema_name, "schemaName")?; + validate_identifier(table_name, "tableName")?; + let rows = metadata_rows( + application, + datasource_id, + Some(database_name), + "SELECT c.name, ty.name, \ + ty.name + CASE WHEN ty.name IN ('varchar','char','varbinary','binary') THEN '(' + CASE WHEN c.max_length=-1 THEN 'max' ELSE CONVERT(varchar(10),c.max_length) END + ')' \ + WHEN ty.name IN ('nvarchar','nchar') THEN '(' + CASE WHEN c.max_length=-1 THEN 'max' ELSE CONVERT(varchar(10),c.max_length/2) END + ')' \ + WHEN ty.name IN ('decimal','numeric') THEN '('+CONVERT(varchar(10),c.precision)+','+CONVERT(varchar(10),c.scale)+')' \ + WHEN ty.name IN ('datetime2','datetimeoffset','time') THEN '('+CONVERT(varchar(10),c.scale)+')' ELSE '' END, \ + dc.definition, COALESCE(CONVERT(nvarchar(max),ep.value),''), c.is_nullable, CONVERT(int,c.column_id), \ + CONVERT(int,c.max_length), CONVERT(int,c.precision), CONVERT(int,c.scale), c.is_identity, \ + TRY_CONVERT(int,ic.seed_value), TRY_CONVERT(int,ic.increment_value), c.is_computed, cc.definition, c.is_sparse, \ + COALESCE(dc.name,''), COALESCE(c.collation_name,''), COALESCE(kc.name,''), COALESCE(CONVERT(int,icx.key_ordinal),0) \ + FROM sys.columns c JOIN sys.tables t ON t.object_id=c.object_id \ + JOIN sys.schemas s ON s.schema_id=t.schema_id JOIN sys.types ty ON ty.user_type_id=c.user_type_id \ + LEFT JOIN sys.default_constraints dc ON dc.object_id=c.default_object_id \ + LEFT JOIN sys.identity_columns ic ON ic.object_id=c.object_id AND ic.column_id=c.column_id \ + LEFT JOIN sys.computed_columns cc ON cc.object_id=c.object_id AND cc.column_id=c.column_id \ + LEFT JOIN sys.extended_properties ep ON ep.major_id=c.object_id AND ep.minor_id=c.column_id AND ep.name='MS_Description' \ + LEFT JOIN sys.indexes pix ON pix.object_id=t.object_id AND pix.is_primary_key=1 \ + LEFT JOIN sys.index_columns icx ON icx.object_id=t.object_id AND icx.index_id=pix.index_id AND icx.column_id=c.column_id \ + LEFT JOIN sys.key_constraints kc ON kc.parent_object_id=t.object_id AND kc.unique_index_id=pix.index_id AND kc.type='PK' \ + WHERE s.name=@P1 AND t.name=@P2 ORDER BY c.column_id", + vec![ + DatabaseValue::Text(schema_name.to_owned()), + DatabaseValue::Text(table_name.to_owned()), + ], + ) + .await?; + rows.into_iter() + .map(|row| { + let type_name = row_string(&row, 1)?; + let generated = row_bool(&row, 13)?; + let primary_key_order = row_i32(&row, 19)?; + Ok(ColumnMetadata { + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + table_name: table_name.to_owned(), + name: row_string(&row, 0)?, + column_type: row_string(&row, 2)?, + data_type: Some(sqlserver_metadata_jdbc_type(&type_name)), + default_value: row_optional_string(&row, 3)?, + auto_increment: Some(row_bool(&row, 10)?), + comment: row_string(&row, 4)?, + primary_key: Some(primary_key_order > 0), + primary_key_name: row_string(&row, 18)?, + primary_key_order, + column_size: Some(row_i32(&row, 7)?), + buffer_length: Some(row_i32(&row, 7)?), + decimal_digits: Some(row_i32(&row, 9)?), + num_prec_radix: sqlserver_numeric_name(&type_name).then_some(10), + char_octet_length: Some(row_i32(&row, 7)?), + ordinal_position: Some(row_i32(&row, 6)?), + nullable: Some(i32::from(row_bool(&row, 5)?)), + generated_column: Some(generated), + extent: if generated { + row_optional_string(&row, 14)?.unwrap_or_default() + } else { + String::new() + }, + collation: row_string(&row, 17)?, + sparse: Some(row_bool(&row, 15)?), + default_constraint_name: row_string(&row, 16)?, + seed: row_optional_i32(&row, 11)?, + increment: row_optional_i32(&row, 12)?, + ..ColumnMetadata::default() + }) + }) + .collect::, _>>() + .map(|items| ColumnList { items }) +} + +async fn list_indexes( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + table_name: &str, +) -> Result { + validate_identifier(database_name, "databaseName")?; + let schema_name = effective_schema(schema_name); + validate_identifier(schema_name, "schemaName")?; + validate_identifier(table_name, "tableName")?; + let rows = metadata_rows( + application, + datasource_id, + Some(database_name), + "SELECT i.name, i.is_unique, i.type_desc, COALESCE(i.filter_definition,''), \ + CONVERT(int,ic.key_ordinal), ic.is_descending_key, c.name, ic.is_included_column, \ + COALESCE(CONVERT(nvarchar(max),ep.value),'') \ + FROM sys.indexes i JOIN sys.tables t ON t.object_id=i.object_id \ + JOIN sys.schemas s ON s.schema_id=t.schema_id \ + JOIN sys.index_columns ic ON ic.object_id=i.object_id AND ic.index_id=i.index_id \ + JOIN sys.columns c ON c.object_id=ic.object_id AND c.column_id=ic.column_id \ + LEFT JOIN sys.extended_properties ep ON ep.major_id=i.object_id AND ep.minor_id=i.index_id AND ep.name='MS_Description' \ + WHERE s.name=@P1 AND t.name=@P2 AND i.index_id>0 ORDER BY i.index_id,ic.index_column_id", + vec![ + DatabaseValue::Text(schema_name.to_owned()), + DatabaseValue::Text(table_name.to_owned()), + ], + ) + .await?; + let mut items = Vec::::new(); + let mut positions = HashMap::::new(); + for row in rows { + let name = row_string(&row, 0)?; + let index = if let Some(index) = positions.get(&name).copied() { + index + } else { + let index = items.len(); + positions.insert(name.clone(), index); + items.push(IndexMetadata { + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + table_name: table_name.to_owned(), + name: name.clone(), + index_type: row_string(&row, 2)?, + unique: Some(row_bool(&row, 1)?), + comment: row_string(&row, 8)?, + method: row_string(&row, 2)?, + ..IndexMetadata::default() + }); + index + }; + items[index].columns.push(IndexColumnMetadata { + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + table_name: table_name.to_owned(), + index_name: name, + column_name: row_string(&row, 6)?, + column_type: if row_bool(&row, 7)? { + "INCLUDED".to_owned() + } else { + "KEY".to_owned() + }, + ordinal_position: Some(row_i32(&row, 4)?), + non_unique: Some(!row_bool(&row, 1)?), + sort_order: if row_bool(&row, 5)? { "D" } else { "A" }.to_owned(), + filter_condition: row_string(&row, 3)?, + ..IndexColumnMetadata::default() + }); + } + Ok(IndexList { items }) +} + +async fn list_views( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + name_pattern: &str, +) -> Result { + let schema_name = effective_schema(schema_name); + validate_identifier(database_name, "databaseName")?; + validate_identifier(schema_name, "schemaName")?; + let rows = metadata_rows( + application, + datasource_id, + Some(database_name), + "SELECT s.name,v.name,COALESCE(m.definition,''),CONVERT(nvarchar(33),v.create_date,126), \ + CONVERT(nvarchar(33),v.modify_date,126),COALESCE(CONVERT(nvarchar(max),ep.value),'') \ + FROM sys.views v JOIN sys.schemas s ON s.schema_id=v.schema_id \ + LEFT JOIN sys.sql_modules m ON m.object_id=v.object_id \ + LEFT JOIN sys.extended_properties ep ON ep.major_id=v.object_id AND ep.minor_id=0 AND ep.name='MS_Description' \ + WHERE s.name=@P1 AND (@P2='' OR v.name LIKE @P2) ORDER BY v.name", + vec![ + DatabaseValue::Text(schema_name.to_owned()), + DatabaseValue::Text(name_pattern.trim().to_owned()), + ], + ) + .await?; + rows.into_iter() + .map(|row| view_metadata(database_name, &row)) + .collect::, _>>() + .map(|items| ViewList { items }) +} + +async fn get_view( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + view_name: &str, +) -> Result { + let schema_name = effective_schema(schema_name); + validate_identifier(view_name, "viewName")?; + let rows = metadata_rows( + application, + datasource_id, + Some(database_name), + "SELECT s.name,v.name,COALESCE(m.definition,''),CONVERT(nvarchar(33),v.create_date,126), \ + CONVERT(nvarchar(33),v.modify_date,126),COALESCE(CONVERT(nvarchar(max),ep.value),'') \ + FROM sys.views v JOIN sys.schemas s ON s.schema_id=v.schema_id \ + LEFT JOIN sys.sql_modules m ON m.object_id=v.object_id \ + LEFT JOIN sys.extended_properties ep ON ep.major_id=v.object_id AND ep.minor_id=0 AND ep.name='MS_Description' \ + WHERE s.name=@P1 AND v.name=@P2", + vec![ + DatabaseValue::Text(schema_name.to_owned()), + DatabaseValue::Text(view_name.to_owned()), + ], + ) + .await?; + rows.into_iter() + .next() + .ok_or_else(|| metadata_not_found("view", database_name, schema_name, view_name)) + .and_then(|row| view_metadata(database_name, &row)) +} + +fn view_metadata(database_name: &str, row: &Row) -> Result { + Ok(TableMetadata { + database_name: database_name.to_owned(), + schema_name: row_string(row, 0)?, + name: row_string(row, 1)?, + table_type: "VIEW".to_owned(), + database_type: "SQLSERVER".to_owned(), + ddl: row_string(row, 2)?, + create_time: row_string(row, 3)?, + update_time: row_string(row, 4)?, + comment: row_string(row, 5)?, + ..TableMetadata::default() + }) +} + +async fn list_foreign_keys( + application: &Application, + request: &ListTableKeysRequest, + exported: bool, +) -> Result { + let database_name = &request.table.scope.database_name; + let schema_name = effective_schema(&request.table.scope.schema_name); + let table_name = &request.table.table_name; + validate_identifier(database_name, "databaseName")?; + validate_identifier(schema_name, "schemaName")?; + validate_identifier(table_name, "tableName")?; + let predicate = if exported { + "rs.name=@P1 AND rt.name=@P2" + } else { + "fs.name=@P1 AND ft.name=@P2" + }; + let sql = format!( + "SELECT rs.name,rt.name,rc.name,fs.name,ft.name,fc.name,CONVERT(int,fkc.constraint_column_id), \ + fk.update_referential_action_desc,fk.delete_referential_action_desc,fk.name,COALESCE(pk.name,'') \ + FROM sys.foreign_keys fk JOIN sys.foreign_key_columns fkc ON fkc.constraint_object_id=fk.object_id \ + JOIN sys.tables ft ON ft.object_id=fk.parent_object_id JOIN sys.schemas fs ON fs.schema_id=ft.schema_id \ + JOIN sys.columns fc ON fc.object_id=ft.object_id AND fc.column_id=fkc.parent_column_id \ + JOIN sys.tables rt ON rt.object_id=fk.referenced_object_id JOIN sys.schemas rs ON rs.schema_id=rt.schema_id \ + JOIN sys.columns rc ON rc.object_id=rt.object_id AND rc.column_id=fkc.referenced_column_id \ + LEFT JOIN sys.key_constraints pk ON pk.parent_object_id=rt.object_id AND pk.type='PK' \ + WHERE {predicate} ORDER BY fk.name,fkc.constraint_column_id" + ); + let rows = metadata_rows( + application, + &request.table.scope.datasource_id, + Some(database_name), + &sql, + vec![ + DatabaseValue::Text(schema_name.to_owned()), + DatabaseValue::Text(table_name.to_owned()), + ], + ) + .await?; + rows.into_iter() + .map(|row| { + Ok(ForeignKeyMetadata { + primary_table_database: database_name.to_owned(), + primary_table_schema: row_string(&row, 0)?, + primary_table_name: row_string(&row, 1)?, + primary_column_name: row_string(&row, 2)?, + foreign_table_database: database_name.to_owned(), + foreign_table_schema: row_string(&row, 3)?, + foreign_table_name: row_string(&row, 4)?, + foreign_column_name: row_string(&row, 5)?, + key_sequence: row_i32(&row, 6)?, + update_rule: referential_rule(&row_string(&row, 7)?), + delete_rule: referential_rule(&row_string(&row, 8)?), + foreign_key_name: row_string(&row, 9)?, + primary_key_name: row_string(&row, 10)?, + deferrability: 7, + }) + }) + .collect::, _>>() + .map(|items| ForeignKeyList { items }) +} + +async fn list_primary_keys( + application: &Application, + request: &ListTableKeysRequest, +) -> Result { + let database_name = &request.table.scope.database_name; + let schema_name = effective_schema(&request.table.scope.schema_name); + let table_name = &request.table.table_name; + let rows = metadata_rows( + application, + &request.table.scope.datasource_id, + Some(database_name), + "SELECT c.name,kc.name FROM sys.key_constraints kc \ + JOIN sys.tables t ON t.object_id=kc.parent_object_id JOIN sys.schemas s ON s.schema_id=t.schema_id \ + JOIN sys.index_columns ic ON ic.object_id=t.object_id AND ic.index_id=kc.unique_index_id \ + JOIN sys.columns c ON c.object_id=t.object_id AND c.column_id=ic.column_id \ + WHERE kc.type='PK' AND s.name=@P1 AND t.name=@P2 ORDER BY ic.key_ordinal", + vec![ + DatabaseValue::Text(schema_name.to_owned()), + DatabaseValue::Text(table_name.to_owned()), + ], + ) + .await?; + rows.into_iter() + .map(|row| { + Ok(PrimaryKeyMetadata { + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + table_name: table_name.to_owned(), + column_name: row_string(&row, 0)?, + name: row_string(&row, 1)?, + }) + }) + .collect::, _>>() + .map(|items| PrimaryKeyList { items }) +} + +fn referential_rule(value: &str) -> i32 { + match value.to_ascii_uppercase().as_str() { + "CASCADE" => 0, + "SET_NULL" => 2, + "SET_DEFAULT" => 4, + _ => 3, + } +} + +fn metadata_not_found( + kind: &str, + database_name: &str, + schema_name: &str, + object_name: &str, +) -> AppError { + AppError::not_found( + "sqlserver_metadata_not_found", + format!("SQL Server {kind} {database_name}.{schema_name}.{object_name} does not exist"), + ) +} + +fn sqlserver_numeric_name(type_name: &str) -> bool { + matches!( + type_name.to_ascii_lowercase().as_str(), + "tinyint" + | "smallint" + | "int" + | "bigint" + | "real" + | "float" + | "decimal" + | "numeric" + | "money" + | "smallmoney" + ) +} + +fn sqlserver_metadata_jdbc_type(type_name: &str) -> i32 { + match type_name.to_ascii_lowercase().as_str() { + "bit" => -7, + "tinyint" => -6, + "smallint" => 5, + "int" => 4, + "bigint" => -5, + "real" => 7, + "float" => 8, + "decimal" | "numeric" | "money" | "smallmoney" => 3, + "date" => 91, + "time" => 92, + "datetime" | "datetime2" | "smalldatetime" => 93, + "datetimeoffset" => 2014, + "uniqueidentifier" => -11, + "binary" => -2, + "varbinary" | "image" | "timestamp" | "rowversion" => -3, + "char" => 1, + "varchar" => 12, + "nchar" => -15, + "nvarchar" => -9, + "text" => -1, + "ntext" => -16, + "xml" => 2009, + _ => 1111, + } +} + +async fn list_functions( + application: &Application, + request: &ListRoutinesRequest, +) -> Result { + let database_name = &request.scope.database_name; + let schema_name = effective_schema(&request.scope.schema_name); + let rows = metadata_rows( + application, + &request.scope.datasource_id, + Some(database_name), + "SELECT s.name,o.name,o.type,COALESCE(CONVERT(nvarchar(max),ep.value),''),COALESCE(m.definition,'') \ + FROM sys.objects o JOIN sys.schemas s ON s.schema_id=o.schema_id \ + LEFT JOIN sys.sql_modules m ON m.object_id=o.object_id \ + LEFT JOIN sys.extended_properties ep ON ep.major_id=o.object_id AND ep.minor_id=0 AND ep.name='MS_Description' \ + WHERE s.name=@P1 AND o.type IN ('FN','IF','TF','FS','FT') ORDER BY o.name", + vec![DatabaseValue::Text(schema_name.to_owned())], + ) + .await?; + rows.into_iter() + .map(|row| function_metadata(database_name, &row)) + .collect::, _>>() + .map(|items| FunctionList { items }) +} + +async fn get_function( + application: &Application, + request: &MetadataObjectRef, +) -> Result { + let database_name = &request.scope.database_name; + let schema_name = effective_schema(&request.scope.schema_name); + validate_identifier(&request.object_name, "functionName")?; + let rows = metadata_rows( + application, + &request.scope.datasource_id, + Some(database_name), + "SELECT s.name,o.name,o.type,COALESCE(CONVERT(nvarchar(max),ep.value),''),COALESCE(m.definition,'') \ + FROM sys.objects o JOIN sys.schemas s ON s.schema_id=o.schema_id \ + LEFT JOIN sys.sql_modules m ON m.object_id=o.object_id \ + LEFT JOIN sys.extended_properties ep ON ep.major_id=o.object_id AND ep.minor_id=0 AND ep.name='MS_Description' \ + WHERE s.name=@P1 AND o.name=@P2 AND o.type IN ('FN','IF','TF','FS','FT')", + vec![ + DatabaseValue::Text(schema_name.to_owned()), + DatabaseValue::Text(request.object_name.clone()), + ], + ) + .await?; + rows.into_iter() + .next() + .ok_or_else(|| { + metadata_not_found("function", database_name, schema_name, &request.object_name) + }) + .and_then(|row| function_metadata(database_name, &row)) +} + +async fn list_function_parameters( + application: &Application, + request: &MetadataObjectRef, +) -> Result { + let rows = routine_parameter_rows(application, request, true).await?; + rows.into_iter() + .map(|row| { + let type_name = row_string(&row, 4)?; + let parameter_id = row_i32(&row, 3)?; + Ok(FunctionParameterMetadata { + function_database: request.scope.database_name.clone(), + function_schema: row_string(&row, 0)?, + function_name: row_string(&row, 1)?, + column_name: row_string(&row, 2)?, + column_type: Some(if parameter_id == 0 { + 4 + } else if row_bool(&row, 5)? { + 3 + } else { + 1 + }), + data_type: Some(sqlserver_metadata_jdbc_type(&type_name)), + type_name, + precision: Some(row_i32(&row, 7)?), + length: Some(row_i32(&row, 6)?), + scale: Some(row_i32(&row, 8)?), + radix: Some(10), + nullable: Some(1), + char_octet_length: Some(row_i32(&row, 6)?), + ordinal_position: Some(parameter_id), + is_nullable: "YES".to_owned(), + specific_name: request.object_name.clone(), + ..FunctionParameterMetadata::default() + }) + }) + .collect::, _>>() + .map(|items| FunctionParameterList { items }) +} + +async fn list_procedures( + application: &Application, + request: &ListRoutinesRequest, +) -> Result { + let database_name = &request.scope.database_name; + let schema_name = effective_schema(&request.scope.schema_name); + let rows = metadata_rows( + application, + &request.scope.datasource_id, + Some(database_name), + "SELECT s.name,o.name,COALESCE(CONVERT(nvarchar(max),ep.value),''),COALESCE(m.definition,'') \ + FROM sys.objects o JOIN sys.schemas s ON s.schema_id=o.schema_id \ + LEFT JOIN sys.sql_modules m ON m.object_id=o.object_id \ + LEFT JOIN sys.extended_properties ep ON ep.major_id=o.object_id AND ep.minor_id=0 AND ep.name='MS_Description' \ + WHERE s.name=@P1 AND o.type IN ('P','PC','X') ORDER BY o.name", + vec![DatabaseValue::Text(schema_name.to_owned())], + ) + .await?; + rows.into_iter() + .map(|row| procedure_metadata(database_name, &row)) + .collect::, _>>() + .map(|items| ProcedureList { items }) +} + +async fn get_procedure( + application: &Application, + request: &MetadataObjectRef, +) -> Result { + let database_name = &request.scope.database_name; + let schema_name = effective_schema(&request.scope.schema_name); + validate_identifier(&request.object_name, "procedureName")?; + let rows = metadata_rows( + application, + &request.scope.datasource_id, + Some(database_name), + "SELECT s.name,o.name,COALESCE(CONVERT(nvarchar(max),ep.value),''),COALESCE(m.definition,'') \ + FROM sys.objects o JOIN sys.schemas s ON s.schema_id=o.schema_id \ + LEFT JOIN sys.sql_modules m ON m.object_id=o.object_id \ + LEFT JOIN sys.extended_properties ep ON ep.major_id=o.object_id AND ep.minor_id=0 AND ep.name='MS_Description' \ + WHERE s.name=@P1 AND o.name=@P2 AND o.type IN ('P','PC','X')", + vec![ + DatabaseValue::Text(schema_name.to_owned()), + DatabaseValue::Text(request.object_name.clone()), + ], + ) + .await?; + rows.into_iter() + .next() + .ok_or_else(|| { + metadata_not_found( + "procedure", + database_name, + schema_name, + &request.object_name, + ) + }) + .and_then(|row| procedure_metadata(database_name, &row)) +} + +async fn list_procedure_parameters( + application: &Application, + request: &MetadataObjectRef, +) -> Result { + let rows = routine_parameter_rows(application, request, false).await?; + rows.into_iter() + .map(|row| { + let type_name = row_string(&row, 4)?; + let output = row_bool(&row, 5)?; + let parameter_id = row_i32(&row, 3)?; + Ok(ProcedureParameterMetadata { + procedure_database: request.scope.database_name.clone(), + procedure_schema: row_string(&row, 0)?, + procedure_name: row_string(&row, 1)?, + column_name: row_string(&row, 2)?, + column_type: Some(if parameter_id == 0 { + 5 + } else if output { + 4 + } else { + 1 + }), + data_type: Some(sqlserver_metadata_jdbc_type(&type_name)), + type_name, + precision: Some(row_i32(&row, 7)?), + length: Some(row_i32(&row, 6)?), + scale: Some(row_i32(&row, 8)?), + radix: Some(10), + nullable: Some(1), + column_default: row_string(&row, 9)?, + char_octet_length: Some(row_i32(&row, 6)?), + ordinal_position: Some(parameter_id), + is_nullable: "YES".to_owned(), + specific_name: request.object_name.clone(), + ..ProcedureParameterMetadata::default() + }) + }) + .collect::, _>>() + .map(|items| ProcedureParameterList { items }) +} + +async fn list_triggers( + application: &Application, + request: &ListTriggersRequest, +) -> Result { + let rows = trigger_rows(application, request, None).await?; + rows.into_iter() + .map(|row| trigger_metadata(&request.scope.database_name, &row)) + .collect::, _>>() + .map(|items| TriggerList { items }) +} + +async fn get_trigger( + application: &Application, + request: &MetadataObjectRef, +) -> Result { + let trigger_request = ListTriggersRequest { + scope: request.scope.clone(), + }; + let rows = trigger_rows(application, &trigger_request, Some(&request.object_name)).await?; + rows.into_iter() + .next() + .ok_or_else(|| { + metadata_not_found( + "trigger", + &request.scope.database_name, + effective_schema(&request.scope.schema_name), + &request.object_name, + ) + }) + .and_then(|row| trigger_metadata(&request.scope.database_name, &row)) +} + +fn function_metadata(database_name: &str, row: &Row) -> Result { + let body = row_string(row, 4)?; + Ok(FunctionMetadata { + database_name: database_name.to_owned(), + schema_name: row_string(row, 0)?, + name: row_string(row, 1)?, + function_type: Some( + if matches!(row_string(row, 2)?.as_str(), "IF" | "TF" | "FT") { + 2 + } else { + 1 + }, + ), + remarks: row_string(row, 3)?, + specific_name: row_string(row, 1)?, + template: body.clone(), + body, + }) +} + +fn procedure_metadata(database_name: &str, row: &Row) -> Result { + Ok(ProcedureMetadata { + database_name: database_name.to_owned(), + schema_name: row_string(row, 0)?, + name: row_string(row, 1)?, + remarks: row_string(row, 2)?, + procedure_type: Some(2), + specific_name: row_string(row, 1)?, + body: row_string(row, 3)?, + }) +} + +async fn routine_parameter_rows( + application: &Application, + request: &MetadataObjectRef, + function: bool, +) -> Result, AppError> { + let schema_name = effective_schema(&request.scope.schema_name); + let object_types = if function { + "'FN','IF','TF','FS','FT'" + } else { + "'P','PC','X'" + }; + let sql = format!( + "SELECT s.name,o.name,COALESCE(p.name,''),CONVERT(int,p.parameter_id),ty.name,p.is_output, \ + CONVERT(int,p.max_length),CONVERT(int,p.precision),CONVERT(int,p.scale), \ + COALESCE(TRY_CONVERT(nvarchar(max),p.default_value),'') \ + FROM sys.objects o JOIN sys.schemas s ON s.schema_id=o.schema_id \ + JOIN sys.parameters p ON p.object_id=o.object_id JOIN sys.types ty ON ty.user_type_id=p.user_type_id \ + WHERE s.name=@P1 AND o.name=@P2 AND o.type IN ({object_types}) ORDER BY p.parameter_id" + ); + metadata_rows( + application, + &request.scope.datasource_id, + Some(&request.scope.database_name), + &sql, + vec![ + DatabaseValue::Text(schema_name.to_owned()), + DatabaseValue::Text(request.object_name.clone()), + ], + ) + .await +} + +async fn trigger_rows( + application: &Application, + request: &ListTriggersRequest, + trigger_name: Option<&str>, +) -> Result, AppError> { + let schema_name = effective_schema(&request.scope.schema_name); + let name = trigger_name.unwrap_or_default(); + metadata_rows( + application, + &request.scope.datasource_id, + Some(&request.scope.database_name), + "SELECT s.name,tr.name,COALESCE(STRING_AGG(te.type_desc,','),''),COALESCE(m.definition,'') \ + FROM sys.triggers tr JOIN sys.tables t ON t.object_id=tr.parent_id JOIN sys.schemas s ON s.schema_id=t.schema_id \ + LEFT JOIN sys.trigger_events te ON te.object_id=tr.object_id \ + LEFT JOIN sys.sql_modules m ON m.object_id=tr.object_id \ + WHERE s.name=@P1 AND (@P2='' OR tr.name=@P2) GROUP BY s.name,tr.name,m.definition ORDER BY tr.name", + vec![ + DatabaseValue::Text(schema_name.to_owned()), + DatabaseValue::Text(name.to_owned()), + ], + ) + .await +} + +fn trigger_metadata(database_name: &str, row: &Row) -> Result { + Ok(TriggerMetadata { + database_name: database_name.to_owned(), + schema_name: row_string(row, 0)?, + name: row_string(row, 1)?, + event_manipulation: row_string(row, 2)?, + body: row_string(row, 3)?, + }) +} + +async fn load_er_tables( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, +) -> Result, AppError> { + let schema_name = effective_schema(schema_name).to_owned(); + let tables = list_tables(application, datasource_id, database_name, &schema_name, "").await?; + let mut result = Vec::with_capacity(tables.items.len()); + for table in tables.items { + let columns = list_columns( + application, + datasource_id, + database_name, + &schema_name, + &table.name, + ) + .await?; + let keys_request = + table_keys_request(datasource_id, database_name, &schema_name, &table.name); + let foreign_keys = list_foreign_keys(application, &keys_request, false).await?; + result.push(EntityRelationTable { + name: table.name, + comment: table.comment, + columns: columns + .items + .into_iter() + .map(|column| EntityRelationColumn { + name: column.name, + column_type: column.column_type, + primary_key: column.primary_key.unwrap_or(false), + comment: column.comment, + }) + .collect(), + foreign_keys: foreign_keys + .items + .into_iter() + .map(|key| EntityRelationForeignKey { + primary_table: key.primary_table_name, + primary_column: key.primary_column_name, + foreign_table: key.foreign_table_name, + foreign_column: key.foreign_column_name, + }) + .collect(), + }); + } + Ok(result) +} + +async fn table_ddl( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + table_name: &str, +) -> Result { + let schema_name = effective_schema(schema_name); + let columns = list_columns( + application, + datasource_id, + database_name, + schema_name, + table_name, + ) + .await?; + if columns.items.is_empty() { + return Err(metadata_not_found( + "table", + database_name, + schema_name, + table_name, + )); + } + let keys_request = table_keys_request(datasource_id, database_name, schema_name, table_name); + let primary_keys = list_primary_keys(application, &keys_request).await?; + let foreign_keys = list_foreign_keys(application, &keys_request, false).await?; + let indexes = list_indexes( + application, + datasource_id, + database_name, + schema_name, + table_name, + ) + .await?; + render_table_ddl( + database_name, + schema_name, + table_name, + columns, + &primary_keys, + foreign_keys, + indexes, + ) +} + +fn table_keys_request( + datasource_id: &str, + database_name: &str, + schema_name: &str, + table_name: &str, +) -> ListTableKeysRequest { + ListTableKeysRequest { + table: TableRef { + scope: MetadataScope { + datasource_id: datasource_id.to_owned(), + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + }, + table_name: table_name.to_owned(), + }, + } +} + +#[allow( + clippy::too_many_lines, + reason = "SQL Server table DDL is rendered in dependency order from one metadata snapshot" +)] +fn render_table_ddl( + database_name: &str, + schema_name: &str, + table_name: &str, + columns: ColumnList, + primary_keys: &PrimaryKeyList, + foreign_keys: ForeignKeyList, + indexes: IndexList, +) -> Result { + let target = qualified_table(database_name, schema_name, table_name)?; + let mut definitions = Vec::new(); + for column in columns.items { + let name = quote_identifier(&column.name, "columnName")?; + if column.generated_column.unwrap_or(false) { + if column.extent.trim().is_empty() { + return Err(AppError::invalid( + "sqlserver_ddl_unavailable", + format!( + "Computed column {} has no recoverable definition", + column.name + ), + )); + } + definitions.push(format!(" {name} AS {}", column.extent)); + continue; + } + let mut definition = format!(" {name} {}", column.column_type); + if !column.collation.is_empty() && sqlserver_textual_jdbc_type(column.data_type) { + definition.push_str(" COLLATE "); + definition.push_str("e_identifier(&column.collation, "collation")?); + } + if column.sparse.unwrap_or(false) { + definition.push_str(" SPARSE"); + } + if column.auto_increment.unwrap_or(false) { + let _ = write!( + definition, + " IDENTITY({},{})", + column.seed.unwrap_or(1), + column.increment.unwrap_or(1) + ); + } + if let Some(default_value) = column + .default_value + .filter(|value| !value.trim().is_empty()) + { + if !column.default_constraint_name.is_empty() { + definition.push_str(" CONSTRAINT "); + definition.push_str("e_identifier( + &column.default_constraint_name, + "defaultConstraintName", + )?); + } + definition.push_str(" DEFAULT "); + definition.push_str(&default_value); + } + definition.push_str(if column.nullable == Some(1) { + " NULL" + } else { + " NOT NULL" + }); + definitions.push(definition); + } + + if !primary_keys.items.is_empty() { + let constraint_name = primary_keys.items[0].name.clone(); + let columns = primary_keys + .items + .iter() + .map(|key| quote_identifier(&key.column_name, "primaryKeyColumn")) + .collect::, _>>()? + .join(", "); + definitions.push(format!( + " CONSTRAINT {} PRIMARY KEY ({columns})", + quote_identifier(&constraint_name, "primaryKeyName")? + )); + } + + let mut foreign_key_order = Vec::::new(); + let mut grouped_foreign_keys = HashMap::>::new(); + for key in foreign_keys.items { + if !grouped_foreign_keys.contains_key(&key.foreign_key_name) { + foreign_key_order.push(key.foreign_key_name.clone()); + } + grouped_foreign_keys + .entry(key.foreign_key_name.clone()) + .or_default() + .push(key); + } + for name in foreign_key_order { + let mut keys = grouped_foreign_keys.remove(&name).unwrap_or_default(); + keys.sort_by_key(|key| key.key_sequence); + let Some(first) = keys.first() else { + continue; + }; + let local_columns = keys + .iter() + .map(|key| quote_identifier(&key.foreign_column_name, "foreignColumn")) + .collect::, _>>()? + .join(", "); + let referenced_columns = keys + .iter() + .map(|key| quote_identifier(&key.primary_column_name, "primaryColumn")) + .collect::, _>>()? + .join(", "); + let referenced_table = format!( + "{}.{}.{}", + quote_identifier(&first.primary_table_database, "primaryDatabase")?, + quote_identifier(&first.primary_table_schema, "primarySchema")?, + quote_identifier(&first.primary_table_name, "primaryTable")? + ); + let mut definition = format!( + " CONSTRAINT {} FOREIGN KEY ({local_columns}) REFERENCES {referenced_table} ({referenced_columns})", + quote_identifier(&name, "foreignKeyName")? + ); + append_referential_action(&mut definition, "UPDATE", first.update_rule); + append_referential_action(&mut definition, "DELETE", first.delete_rule); + definitions.push(definition); + } + + let mut ddl = format!("CREATE TABLE {target} (\n{}\n);", definitions.join(",\n")); + let primary_name = primary_keys.items.first().map(|key| key.name.as_str()); + for index in indexes.items { + if primary_name.is_some_and(|name| name.eq_ignore_ascii_case(&index.name)) { + continue; + } + let key_columns = index + .columns + .iter() + .filter(|column| column.column_type == "KEY") + .map(|column| { + quote_identifier(&column.column_name, "indexColumn").map(|name| { + format!( + "{name} {}", + if column.sort_order == "D" { + "DESC" + } else { + "ASC" + } + ) + }) + }) + .collect::, _>>()?; + if key_columns.is_empty() { + continue; + } + ddl.push_str("\n\nCREATE "); + if index.unique.unwrap_or(false) { + ddl.push_str("UNIQUE "); + } + if index.index_type.contains("CLUSTERED") && !index.index_type.contains("NONCLUSTERED") { + ddl.push_str("CLUSTERED "); + } else { + ddl.push_str("NONCLUSTERED "); + } + ddl.push_str("INDEX "); + ddl.push_str("e_identifier(&index.name, "indexName")?); + ddl.push_str(" ON "); + ddl.push_str(&target); + ddl.push_str(" ("); + ddl.push_str(&key_columns.join(", ")); + ddl.push(')'); + let included = index + .columns + .iter() + .filter(|column| column.column_type == "INCLUDED") + .map(|column| quote_identifier(&column.column_name, "includedColumn")) + .collect::, _>>()?; + if !included.is_empty() { + ddl.push_str(" INCLUDE ("); + ddl.push_str(&included.join(", ")); + ddl.push(')'); + } + if let Some(filter) = index + .columns + .first() + .map(|column| column.filter_condition.trim()) + .filter(|filter| !filter.is_empty()) + { + ddl.push_str(" WHERE "); + ddl.push_str(filter); + } + ddl.push(';'); + } + Ok(ddl) +} + +fn sqlserver_textual_jdbc_type(jdbc_type: Option) -> bool { + matches!(jdbc_type, Some(1 | 12 | -1 | -9 | -15 | -16)) +} + +fn append_referential_action(sql: &mut String, action: &str, rule: i32) { + let rule = match rule { + 0 => Some("CASCADE"), + 2 => Some("SET NULL"), + 4 => Some("SET DEFAULT"), + _ => None, + }; + if let Some(rule) = rule { + sql.push_str(" ON "); + sql.push_str(action); + sql.push(' '); + sql.push_str(rule); + } +} + +async fn start_table_preview( + application: &Application, + request: TablePreviewRequest, + row_limit: u32, +) -> Result { + let table = qualified_table( + &request.table.scope.database_name, + effective_schema(&request.table.scope.schema_name), + &request.table.table_name, + )?; + let sql = format!("SELECT TOP ({row_limit}) * FROM {table}"); + let accepted = application + .start_read_query(StartQueryRequest { + datasource_id: request.table.scope.datasource_id, + sql: sql.clone(), + parameters: Vec::new(), + limits: QueryLimits { + max_rows: row_limit.to_string(), + max_result_bytes: (8 * 1024 * 1024_u64).to_string(), + batch_rows: row_limit.min(200), + batch_bytes: 1024 * 1024, + result_ttl_seconds: 60 * 60, + }, + }) + .await?; + Ok(TablePreviewAccepted { + operation_id: accepted.operation_id, + sql, + row_limit, + }) +} + +#[cfg(test)] +mod tests { + use chat2db_contract::DatasourceConnectionProperty; + + use super::*; + + #[test] + fn descriptor_exposes_only_implemented_native_capabilities() { + let driver = SqlServerNativeDriver; + + assert_eq!(driver.descriptor(), &SQLSERVER_DRIVER_DESCRIPTOR); + assert_eq!(driver.descriptor().implementation, "tiberius"); + assert!(driver.connection().is_some()); + assert!(driver.query().is_some()); + assert!(driver.metadata().is_some()); + assert!(driver.tables().is_some()); + assert!(driver.routines().is_none()); + assert!(driver.transfer().is_none()); + assert!(driver.dialect().is_some()); + assert!(driver.administration().is_none()); + assert!(driver.schema_diff().is_none()); + } + + #[test] + fn namespace_builders_render_safe_sqlserver_batches() { + let driver = SqlServerNativeDriver; + let schema = driver + .build_create_schema(CreateSchemaSqlRequest { + schema: SchemaDefinition { + database_name: "sales]archive".to_owned(), + name: "reporting]daily".to_owned(), + comment: "owner's notes".to_owned(), + owner: "db_owner".to_owned(), + system: false, + }, + }) + .expect("SQL Server schema SQL should render"); + assert_eq!( + schema.sql, + "USE [sales]]archive];\nEXEC(N'CREATE SCHEMA [reporting]]daily] AUTHORIZATION [db_owner]');\nEXEC sys.sp_addextendedproperty @name = N'MS_Description', @value = N'owner''s notes', @level0type = N'SCHEMA', @level0name = N'reporting]daily';" + ); + + let database = driver + .build_namespace_sql(NamespaceSqlRequest { + operation: NamespaceSqlOperation::CreateDatabase { + database: database_definition("inventory]2026", "initial owner's note"), + }, + }) + .expect("SQL Server database SQL should render"); + assert_eq!( + database.sql, + "CREATE DATABASE [inventory]]2026] COLLATE Latin1_General_100_CI_AS_SC_UTF8;\nALTER AUTHORIZATION ON DATABASE::[inventory]]2026] TO [sa];\nEXEC [inventory]]2026].sys.sp_addextendedproperty @name = N'MS_Description', @value = N'initial owner''s note';" + ); + + let use_database = driver + .build_namespace_sql(NamespaceSqlRequest { + operation: NamespaceSqlOperation::UseDatabase { + database_name: "inventory]2026".to_owned(), + }, + }) + .expect("SQL Server USE SQL should render"); + assert_eq!(use_database.sql, "USE [inventory]]2026];"); + } + + #[test] + fn database_alter_builder_covers_rename_collation_owner_and_comment() { + let driver = SqlServerNativeDriver; + let old_database = database_definition("inventory", "old note"); + let mut new_database = database_definition("inventory_archive", "new owner's note"); + new_database.collation = "SQL_Latin1_General_CP1_CI_AS".to_owned(); + new_database.owner = "archive_owner".to_owned(); + let built = driver + .build_namespace_sql(NamespaceSqlRequest { + operation: NamespaceSqlOperation::AlterDatabase { + old_database, + new_database, + }, + }) + .expect("supported SQL Server database changes should render"); + assert_eq!( + built.sql, + "ALTER DATABASE [inventory] MODIFY NAME = [inventory_archive];\nALTER DATABASE [inventory_archive] COLLATE SQL_Latin1_General_CP1_CI_AS;\nALTER AUTHORIZATION ON DATABASE::[inventory_archive] TO [archive_owner];\nEXEC [inventory_archive].sys.sp_updateextendedproperty @name = N'MS_Description', @value = N'new owner''s note';" + ); + + let mut clear_comment = database_definition("inventory", "old note"); + clear_comment.comment.clear(); + let built = driver + .build_namespace_sql(NamespaceSqlRequest { + operation: NamespaceSqlOperation::AlterDatabase { + old_database: database_definition("inventory", "old note"), + new_database: clear_comment, + }, + }) + .expect("SQL Server database comments should be removable"); + assert_eq!( + built.sql, + "EXEC [inventory].sys.sp_dropextendedproperty @name = N'MS_Description';" + ); + } + + #[test] + fn namespace_builder_rejects_unsupported_or_unsafe_changes() { + let driver = SqlServerNativeDriver; + let charset = driver + .build_namespace_sql(NamespaceSqlRequest { + operation: NamespaceSqlOperation::CreateDatabase { + database: DatabaseDefinition { + charset: "UTF-8".to_owned(), + ..database_definition("inventory", "") + }, + }, + }) + .expect_err("SQL Server has no independent database charset"); + assert_eq!( + charset.api_error().code, + "sqlserver_database_charset_unsupported" + ); + + let collation = driver + .build_namespace_sql(NamespaceSqlRequest { + operation: NamespaceSqlOperation::CreateDatabase { + database: DatabaseDefinition { + collation: "Latin1_General_CI_AS; DROP DATABASE master".to_owned(), + ..database_definition("inventory", "") + }, + }, + }) + .expect_err("unsafe SQL Server collations must be rejected"); + assert_eq!( + collation.api_error().code, + "invalid_sqlserver_dialect_request" + ); + + let rename = driver + .build_namespace_sql(NamespaceSqlRequest { + operation: NamespaceSqlOperation::AlterSchema { + old_schema_name: "old_schema".to_owned(), + new_schema_name: "new_schema".to_owned(), + }, + }) + .expect_err("SQL Server schema rename must not silently degrade"); + assert_eq!( + rename.api_error().code, + "sqlserver_schema_rename_unsupported" + ); + } + + #[test] + fn typed_dml_builder_renders_three_part_unicode_and_binary_values() { + let driver = SqlServerNativeDriver; + let columns = vec![ + dml_column("label", "nvarchar"), + dml_column("amount", "decimal"), + dml_column("active", "bit"), + dml_column("created_at", "datetimeoffset"), + dml_column("payload", "varbinary"), + ]; + let built = driver + .build_dml(DmlSqlRequest { + target: DmlTarget { + database_name: Some("sales]archive".to_owned()), + schema_name: Some("dbo".to_owned()), + table_name: "order]items".to_owned(), + }, + statement: DmlStatement::SingleInsert { + columns, + row: DmlRow { + values: vec![ + DmlValue::String("O'Brien".to_owned()), + DmlValue::Decimal("+001.20".to_owned()), + DmlValue::Boolean(true), + DmlValue::Temporal { + kind: DmlTemporalKind::OffsetDatetime, + iso8601: "2026-08-07T12:30:45.1234567+08:00".to_owned(), + }, + DmlValue::Binary(vec![0, 255]), + ], + }, + }, + }) + .expect("typed SQL Server INSERT should render"); + assert_eq!( + built.sql, + "INSERT INTO [sales]]archive].[dbo].[order]]items] ([label], [amount], [active], [created_at], [payload]) VALUES\n(N'O''Brien', 1.20, 1, CAST(N'2026-08-07T12:30:45.1234567+08:00' AS datetimeoffset), 0x00ff);" + ); + + let update = driver + .build_dml(DmlSqlRequest { + target: DmlTarget { + database_name: Some("sales".to_owned()), + schema_name: None, + table_name: "items".to_owned(), + }, + statement: DmlStatement::Update { + assignments: vec![DmlAssignment { + column: dml_column("label", "nvarchar"), + value: DmlValue::String("updated".to_owned()), + }], + predicates: vec![DmlAssignment { + column: dml_column("deleted_at", "datetime2"), + value: DmlValue::Null, + }], + }, + }) + .expect("typed SQL Server UPDATE should render"); + assert_eq!( + update.sql, + "UPDATE [sales]..[items] SET [label] = N'updated' WHERE [deleted_at] IS NULL;" + ); + } + + #[test] + fn typed_dml_builder_fails_closed_for_invalid_values_and_shapes() { + let driver = SqlServerNativeDriver; + let target = DmlTarget { + database_name: None, + schema_name: Some("dbo".to_owned()), + table_name: "items".to_owned(), + }; + let decimal = driver + .build_dml(DmlSqlRequest { + target: target.clone(), + statement: DmlStatement::SingleInsert { + columns: vec![dml_column("amount", "decimal")], + row: DmlRow { + values: vec![DmlValue::Decimal("1; DROP TABLE items".to_owned())], + }, + }, + }) + .expect_err("invalid SQL Server decimals must be rejected"); + assert_eq!(decimal.api_error().code, "invalid_sqlserver_dml"); + + let shape = driver + .build_dml(DmlSqlRequest { + target: target.clone(), + statement: DmlStatement::MultiInsert { + columns: vec![dml_column("id", "int")], + rows: vec![DmlRow { values: Vec::new() }], + }, + }) + .expect_err("SQL Server row shape mismatches must be rejected"); + assert_eq!(shape.api_error().code, "invalid_sqlserver_dml"); + + let update = driver + .build_dml(DmlSqlRequest { + target, + statement: DmlStatement::Update { + assignments: vec![DmlAssignment { + column: dml_column("label", "nvarchar"), + value: DmlValue::String("unsafe".to_owned()), + }], + predicates: Vec::new(), + }, + }) + .expect_err("unbounded SQL Server updates must be rejected"); + assert_eq!(update.api_error().code, "invalid_sqlserver_dml"); + } + + fn database_definition(name: &str, comment: &str) -> DatabaseDefinition { + DatabaseDefinition { + name: name.to_owned(), + comment: comment.to_owned(), + charset: String::new(), + collation: "Latin1_General_100_CI_AS_SC_UTF8".to_owned(), + owner: "sa".to_owned(), + system: false, + } + } + + fn dml_column(name: &str, data_type_name: &str) -> DmlColumn { + DmlColumn { + name: name.to_owned(), + data_type_name: data_type_name.to_owned(), + precision: None, + scale: None, + } + } + + #[test] + fn connection_url_properties_and_ssh_target_are_normalized_safely() { + assert_eq!( + normalize_sqlserver_url(" JDBC:SQLSERVER://db.example:1433;databaseName=inventory ") + .expect("JDBC URL should normalize"), + "jdbc:sqlserver://db.example:1433;databaseName=inventory" + ); + assert_eq!( + normalize_sqlserver_url("sqlserver://db.example:1433") + .expect("native URL should normalize"), + "jdbc:sqlserver://db.example:1433" + ); + assert!(normalize_sqlserver_url("postgres://db.example:5432").is_err()); + + assert_eq!(encode_jdbc_property("plain"), "plain"); + assert_eq!(encode_jdbc_property(" a;b=c}"), "{ a;b=c}}}"); + assert_eq!( + sqlserver_target("jdbc:sqlserver://db.example:15433;databaseName=master") + .expect("IPv4 target should parse"), + ("db.example".to_owned(), 15_433) + ); + assert_eq!( + sqlserver_target("sqlserver://[2001:db8::10]:1433").expect("IPv6 target should parse"), + ("2001:db8::10".to_owned(), 1433) + ); + assert!(sqlserver_target("sqlserver://named-host\\instance").is_err()); + + let config = connection_config(&DatasourceConnection { + jdbc_url: "sqlserver://localhost:1433;encrypt=false".to_owned(), + properties: vec![ + DatasourceConnectionProperty { + key: "user".to_owned(), + value: "sa".to_owned(), + sensitive: false, + }, + DatasourceConnectionProperty { + key: "password".to_owned(), + value: "secret;with=delimiters".to_owned(), + sensitive: true, + }, + ], + read_only: true, + ssh: None, + }); + assert!(config.is_ok()); + + let tls_conflict = connection_config(&DatasourceConnection { + jdbc_url: "sqlserver://localhost:1433;encrypt=true;trustServerCertificate=true" + .to_owned(), + properties: vec![DatasourceConnectionProperty { + key: "trustServerCertificateCA".to_owned(), + value: "/tmp/sqlserver-ca.pem".to_owned(), + sensitive: false, + }], + read_only: false, + ssh: None, + }); + let Err(error) = tls_conflict else { + panic!("conflicting SQL Server trust properties must be rejected"); + }; + assert_eq!(error.api_error().code, "invalid_sqlserver_connection"); + } + + #[test] + fn positional_parameters_are_ordered_and_rewritten_outside_literals() { + let parameters = vec![ + QueryParameter { + position: 2, + value: DatabaseValue::Text("second".to_owned()), + }, + QueryParameter { + position: 1, + value: DatabaseValue::SignedInteger(1), + }, + ]; + let ordered = ordered_parameters(¶meters).expect("parameters should order"); + assert_eq!(ordered[0], &DatabaseValue::SignedInteger(1)); + assert_eq!(ordered[1], &DatabaseValue::Text("second".to_owned())); + + let sql = "SELECT ?, '?', \"?\", [?] -- ?\n, ? /* ? */"; + assert_eq!( + rewrite_positional_parameters(sql, 2).expect("markers should rewrite"), + "SELECT @P1, '?', \"?\", [?] -- ?\n, @P2 /* ? */" + ); + assert!(rewrite_positional_parameters("SELECT ?, ?", 1).is_err()); + assert!( + ordered_parameters(&[ + QueryParameter { + position: 1, + value: DatabaseValue::Null, + }, + QueryParameter { + position: 3, + value: DatabaseValue::Null, + }, + ]) + .is_err() + ); + } + + #[test] + fn read_policy_and_console_splitter_respect_sql_lexical_boundaries() { + assert!(validate_read_sql("SELECT 'INTO' AS keyword_value;").is_ok()); + assert!(validate_read_sql("WITH source AS (SELECT 1 AS id) SELECT id FROM source").is_ok()); + assert!( + validate_read_sql( + "SELECT SUM(CONVERT(bigint, a.object_id % 2)) FROM sys.all_objects AS a CROSS JOIN sys.all_objects AS b" + ) + .is_ok() + ); + assert!(validate_read_sql("SELECT id INTO copied FROM source").is_err()); + assert!(validate_read_sql("UPDATE source SET id = 2").is_err()); + assert!( + validate_read_sql( + "WITH target AS (SELECT id FROM source) DELETE FROM target WHERE id = 1" + ) + .is_err() + ); + assert!( + validate_read_sql("WITH target AS (SELECT id FROM source) UPDATE target SET id = 2") + .is_err() + ); + assert!( + validate_read_sql( + "WITH target AS (SELECT id FROM source) INSERT INTO copied SELECT id FROM target" + ) + .is_err() + ); + assert!( + validate_read_sql( + "WITH target AS (SELECT id FROM source) MERGE copied AS c USING target AS t ON c.id=t.id WHEN MATCHED THEN DELETE" + ) + .is_err() + ); + assert!(validate_read_sql("SELECT 1; SELECT 2").is_err()); + assert!(validate_read_sql("SELECT 1 /* unterminated").is_err()); + + let statements = split_sqlserver_script( + "SELECT ';' AS literal; -- keep ; here\nSELECT [semi;colon] FROM [table]; /* ; */ SELECT 3", + ) + .expect("console script should split"); + assert_eq!(statements.len(), 3); + assert_eq!(statements[0], "SELECT ';' AS literal"); + assert!(statements[1].contains("SELECT [semi;colon] FROM [table]")); + assert!(statements[2].contains("SELECT 3")); + assert!(split_sqlserver_script("SELECT 'unterminated").is_err()); + assert!(creates_local_temp_table( + "CREATE TABLE #native_temp(id int NOT NULL)" + )); + assert!(creates_local_temp_table( + "CREATE TABLE [#native temp](id int NOT NULL)" + )); + assert!(!creates_local_temp_table( + "CREATE TABLE dbo.native_temp(id int NOT NULL)" + )); + assert!(!creates_local_temp_table( + "CREATE TABLE ##native_global_temp(id int NOT NULL)" + )); + for statement in [ + "EXEC sys.sp_addextendedproperty @name=N'MS_Description', @value=N'note'", + "EXECUTE sys.sp_updateextendedproperty @name=N'MS_Description', @value=N'note'", + "EXEC [inventory].sys.sp_dropextendedproperty @name=N'MS_Description'", + ] { + assert!(!is_tabular_statement(statement)); + } + assert!(is_tabular_statement("EXEC dbo.procedure_that_selects")); + assert!(is_tabular_statement( + "EXEC dbo.procedure_that_selects @sp_addextendedproperty = 1" + )); + + for statement in [ + "SELECT id INTO copied FROM source", + "WITH target AS (SELECT id FROM source) DELETE FROM target OUTPUT deleted.id", + ] { + let read_only = validate_read_sql(statement).is_ok(); + assert!(!read_only); + assert!(matches!( + console_execution_interrupted(read_only, Some("cancelled".to_owned())), + ConsoleExecutionError::WriteOutcomeUnknown + )); + assert!(matches!( + console_driver_interrupted(read_only), + ConsoleExecutionError::WriteOutcomeUnknown + )); + } + } + + #[test] + fn dispatched_console_failures_preserve_read_errors_and_fence_write_retries() { + let transport = TiberiusError::Io { + kind: std::io::ErrorKind::ConnectionReset, + message: "connection reset after dispatch".to_owned(), + }; + assert!(matches!( + console_dispatched_driver_error(false, transport), + ConsoleExecutionError::WriteOutcomeUnknown + )); + assert!(matches!( + console_dispatched_driver_error( + false, + TiberiusError::Protocol("invalid response after dispatch".into()) + ), + ConsoleExecutionError::WriteOutcomeUnknown + )); + + let read_error = console_dispatched_driver_error( + true, + TiberiusError::Io { + kind: std::io::ErrorKind::ConnectionReset, + message: "connection reset".to_owned(), + }, + ); + assert!(matches!( + read_error, + ConsoleExecutionError::Statement(error) + if error.api_error().code == "sqlserver_query_failed" + )); + assert!(matches!( + console_post_dispatch_error(false, AppError::internal()), + ConsoleExecutionError::WriteOutcomeUnknown + )); + assert!(matches!( + console_post_dispatch_error(true, AppError::internal()), + ConsoleExecutionError::Statement(_) + )); + } + + #[test] + fn unsafe_tiberius_result_types_and_oversized_scalars_fail_closed() { + for (type_id, type_name, user_type_name) in [ + (Some(60), "money", ""), + (Some(122), "smallmoney", ""), + (Some(98), "sql_variant", ""), + (Some(240), "hierarchyid", "hierarchyid"), + ] { + let error = validate_described_result_type(type_id, type_name, user_type_name) + .expect_err("unsafe Tiberius result types must be rejected"); + assert_eq!(error.api_error().code, "sqlserver_result_type_unsupported"); + } + assert_eq!( + sqlserver_value_type(ColumnType::Money), + wire::JdbcValueType::Float64 + ); + assert_eq!( + sqlserver_value_type(ColumnType::Money4), + wire::JdbcValueType::Float32 + ); + let oversized = ColumnData::String(Some(std::borrow::Cow::Owned( + "x".repeat(MAX_SCALAR_BYTES + 1), + ))); + let error = wire_value(oversized).expect_err("oversized scalar must be rejected"); + assert_eq!(error.api_error().code, "sqlserver_scalar_too_large"); + } + + #[test] + fn datetimeoffset_values_restore_local_time_and_reject_overflow() { + fn date(year: i32, month: u32, day: u32) -> tiberius::time::Date { + let epoch = NaiveDate::from_ymd_opt(1, 1, 1).expect("valid TDS date epoch"); + let value = NaiveDate::from_ymd_opt(year, month, day).expect("valid test date"); + let days = value.signed_duration_since(epoch).num_days(); + tiberius::time::Date::new(u32::try_from(days).expect("positive TDS day count")) + } + + let exact = tiberius::time::DateTimeOffset::new( + tiberius::time::DateTime2::new( + date(2026, 8, 7), + tiberius::time::Time::new(16_496_123_456, 6), + ), + 480, + ); + assert_eq!( + format_datetime_offset(exact).expect("valid datetimeoffset"), + "2026-08-07T12:34:56.123456+08:00" + ); + + let next_day = tiberius::time::DateTimeOffset::new( + tiberius::time::DateTime2::new( + date(2026, 12, 31), + tiberius::time::Time::new(73_800, 0), + ), + 330, + ); + assert_eq!( + format_datetime_offset(next_day).expect("valid cross-day datetimeoffset"), + "2027-01-01T02:00:00+05:30" + ); + + let overflow = tiberius::time::DateTimeOffset::new( + tiberius::time::DateTime2::new( + date(9999, 12, 31), + tiberius::time::Time::new(86_399, 0), + ), + 60, + ); + assert!(format_datetime_offset(overflow).is_err()); + } + + #[test] + fn decimal_and_temporal_parameters_enforce_sqlserver_limits() { + assert_eq!( + format_numeric(parse_decimal("-123.450").expect("decimal should parse")), + "-123.450" + ); + assert_eq!( + format_numeric(parse_decimal(".5").expect("fraction should parse")), + "0.5" + ); + assert!(parse_decimal("1.2.3").is_err()); + assert!(parse_decimal("123456789012345678901234567890123456789").is_err()); + assert!(parse_timestamp("2026-08-07T11:22:33.1234567").is_ok()); + assert!(parse_timestamp("2026/08/07 11:22:33").is_err()); + assert!(validate_parameter(&DatabaseValue::UnsignedInteger(i64::MAX as u64 + 1)).is_err()); + } + + #[test] + fn table_ddl_preserves_identity_constraints_foreign_keys_and_indexes() { + let columns = ColumnList { + items: vec![ + ColumnMetadata { + name: "id".to_owned(), + column_type: "int".to_owned(), + data_type: Some(4), + nullable: Some(0), + auto_increment: Some(true), + seed: Some(10), + increment: Some(5), + ..ColumnMetadata::default() + }, + ColumnMetadata { + name: "name".to_owned(), + column_type: "nvarchar(80)".to_owned(), + data_type: Some(-9), + nullable: Some(1), + default_value: Some("(N'')".to_owned()), + default_constraint_name: "DF_child_name".to_owned(), + collation: "Latin1_General_100_CI_AS".to_owned(), + ..ColumnMetadata::default() + }, + ColumnMetadata { + name: "slug".to_owned(), + generated_column: Some(true), + extent: "LOWER([name])".to_owned(), + ..ColumnMetadata::default() + }, + ], + }; + let primary_keys = PrimaryKeyList { + items: vec![PrimaryKeyMetadata { + column_name: "id".to_owned(), + name: "PK_child".to_owned(), + ..PrimaryKeyMetadata::default() + }], + }; + let foreign_keys = ForeignKeyList { + items: vec![ForeignKeyMetadata { + primary_table_database: "master".to_owned(), + primary_table_schema: "dbo".to_owned(), + primary_table_name: "parent".to_owned(), + primary_column_name: "id".to_owned(), + foreign_column_name: "id".to_owned(), + key_sequence: 1, + update_rule: 0, + delete_rule: 2, + foreign_key_name: "FK_child_parent".to_owned(), + ..ForeignKeyMetadata::default() + }], + }; + let indexes = IndexList { + items: vec![IndexMetadata { + name: "IX_child_name".to_owned(), + index_type: "NONCLUSTERED".to_owned(), + columns: vec![ + IndexColumnMetadata { + column_name: "name".to_owned(), + column_type: "KEY".to_owned(), + sort_order: "D".to_owned(), + filter_condition: "[name] <> N''".to_owned(), + ..IndexColumnMetadata::default() + }, + IndexColumnMetadata { + column_name: "slug".to_owned(), + column_type: "INCLUDED".to_owned(), + ..IndexColumnMetadata::default() + }, + ], + ..IndexMetadata::default() + }], + }; + + let ddl = render_table_ddl( + "master", + "dbo", + "child", + columns, + &primary_keys, + foreign_keys, + indexes, + ) + .expect("DDL should render"); + assert!(ddl.contains("CREATE TABLE [master].[dbo].[child]")); + assert!(ddl.contains("[id] int IDENTITY(10,5) NOT NULL")); + assert!(ddl.contains("[name] nvarchar(80) COLLATE [Latin1_General_100_CI_AS]")); + assert!(ddl.contains("CONSTRAINT [PK_child] PRIMARY KEY ([id])")); + assert!(ddl.contains("ON UPDATE CASCADE ON DELETE SET NULL")); + assert!(ddl.contains("[slug] AS LOWER([name])")); + assert!(ddl.contains( + "CREATE NONCLUSTERED INDEX [IX_child_name] ON [master].[dbo].[child] ([name] DESC) INCLUDE ([slug]) WHERE [name] <> N'';" + )); + } + + async fn live_execute(client: &mut SqlServerClient, sql: &str) -> Result<(), AppError> { + let mut stream = client + .simple_query(sql) + .await + .map_err(sqlserver_query_error)?; + while stream + .try_next() + .await + .map_err(sqlserver_query_error)? + .is_some() + {} + Ok(()) + } + + async fn live_cleanup(client: &mut SqlServerClient) { + let cleanup = "DROP VIEW IF EXISTS dbo.chat2db_native_view; \ + DROP FUNCTION IF EXISTS dbo.chat2db_native_function; \ + DROP PROCEDURE IF EXISTS dbo.chat2db_native_procedure; \ + DROP TABLE IF EXISTS dbo.chat2db_native_child; \ + DROP TABLE IF EXISTS dbo.chat2db_native_parent;"; + if let Err(error) = live_execute(client, cleanup).await { + tracing::warn!(error = %error, "SQL Server smoke cleanup failed"); + } + } + + #[tokio::test] + #[ignore = "requires SQLSERVER_TEST_HOST, SQLSERVER_TEST_PORT, SQLSERVER_TEST_USER, and SQLSERVER_TEST_PASSWORD"] + #[allow( + clippy::too_many_lines, + reason = "the product smoke validates one complete SQL Server object lifecycle" + )] + async fn live_sqlserver_connection_query_and_catalog_smoke() { + let host = std::env::var("SQLSERVER_TEST_HOST").expect("SQLSERVER_TEST_HOST is required"); + let port = std::env::var("SQLSERVER_TEST_PORT") + .expect("SQLSERVER_TEST_PORT is required") + .parse::() + .expect("SQLSERVER_TEST_PORT must be a TCP port"); + let user = std::env::var("SQLSERVER_TEST_USER").expect("SQLSERVER_TEST_USER is required"); + let password = + std::env::var("SQLSERVER_TEST_PASSWORD").expect("SQLSERVER_TEST_PASSWORD is required"); + let connection = DatasourceConnection { + jdbc_url: format!( + "jdbc:sqlserver://{host}:{port};databaseName=master;encrypt=false;trustServerCertificate=true" + ), + properties: vec![ + DatasourceConnectionProperty { + key: "user".to_owned(), + value: user, + sensitive: false, + }, + DatasourceConnectionProperty { + key: "password".to_owned(), + value: password, + sensitive: true, + }, + ], + read_only: false, + ssh: None, + }; + let mut conn = open_connection(&connection) + .await + .expect("native SQL Server connection should open"); + live_cleanup(&mut conn.client).await; + + let smoke_result = async { + live_execute( + &mut conn.client, + "CREATE TABLE dbo.chat2db_native_parent (id int NOT NULL CONSTRAINT PK_chat2db_native_parent PRIMARY KEY, label nvarchar(80) NULL);", + ) + .await?; + live_execute( + &mut conn.client, + "CREATE TABLE dbo.chat2db_native_child (id int NOT NULL CONSTRAINT PK_chat2db_native_child PRIMARY KEY, parent_id int NOT NULL, note nvarchar(80) NULL, CONSTRAINT FK_chat2db_native_child_parent FOREIGN KEY (parent_id) REFERENCES dbo.chat2db_native_parent(id)); CREATE INDEX IX_chat2db_native_child_note ON dbo.chat2db_native_child(note);", + ) + .await?; + live_execute( + &mut conn.client, + "CREATE VIEW dbo.chat2db_native_view AS SELECT id, label FROM dbo.chat2db_native_parent;", + ) + .await?; + live_execute( + &mut conn.client, + "CREATE FUNCTION dbo.chat2db_native_function(@value int) RETURNS int AS BEGIN RETURN @value + 1; END;", + ) + .await?; + live_execute( + &mut conn.client, + "CREATE PROCEDURE dbo.chat2db_native_procedure @value int AS SELECT @value AS value;", + ) + .await?; + live_execute( + &mut conn.client, + "CREATE TRIGGER dbo.chat2db_native_trigger ON dbo.chat2db_native_child AFTER INSERT AS BEGIN SET NOCOUNT ON; END;", + ) + .await?; + + bind_query( + "INSERT INTO dbo.chat2db_native_parent(id, label) VALUES (?, ?)", + &[ + QueryParameter { + position: 1, + value: DatabaseValue::SignedInteger(7), + }, + QueryParameter { + position: 2, + value: DatabaseValue::Text("native-rust".to_owned()), + }, + ], + )? + .execute(&mut conn.client) + .await + .map_err(sqlserver_query_error)?; + bind_query( + "INSERT INTO dbo.chat2db_native_child(id, parent_id, note) VALUES (?, ?, ?)", + &[ + QueryParameter { + position: 1, + value: DatabaseValue::SignedInteger(11), + }, + QueryParameter { + position: 2, + value: DatabaseValue::SignedInteger(7), + }, + QueryParameter { + position: 3, + value: DatabaseValue::Text("catalog".to_owned()), + }, + ], + )? + .execute(&mut conn.client) + .await + .map_err(sqlserver_query_error)?; + + let rows = bind_query( + "SELECT p.id, p.label, c.note FROM dbo.chat2db_native_parent p JOIN dbo.chat2db_native_child c ON c.parent_id=p.id WHERE p.id=?", + &[QueryParameter { + position: 1, + value: DatabaseValue::SignedInteger(7), + }], + )? + .query(&mut conn.client) + .await + .map_err(sqlserver_query_error)? + .into_first_result() + .await + .map_err(sqlserver_query_error)?; + assert_eq!(rows.len(), 1); + assert_eq!(row_i32(&rows[0], 0)?, 7); + assert_eq!(row_string(&rows[0], 1)?, "native-rust"); + assert_eq!(row_string(&rows[0], 2)?, "catalog"); + + let catalog = conn + .client + .simple_query( + "SELECT CONVERT(int, COUNT(*)) FROM sys.objects WHERE name IN ('chat2db_native_parent','chat2db_native_child','chat2db_native_view','chat2db_native_function','chat2db_native_procedure','chat2db_native_trigger')", + ) + .await + .map_err(sqlserver_query_error)? + .into_first_result() + .await + .map_err(sqlserver_query_error)?; + assert_eq!(row_i32(&catalog[0], 0)?, 6); + + let metadata = conn + .client + .simple_query( + "SELECT CONVERT(int, COUNT(DISTINCT o.object_id)) FROM sys.objects o JOIN sys.schemas s ON s.schema_id=o.schema_id LEFT JOIN sys.columns c ON c.object_id=o.object_id LEFT JOIN sys.indexes i ON i.object_id=o.object_id LEFT JOIN sys.foreign_keys fk ON fk.parent_object_id=o.object_id WHERE s.name='dbo' AND o.name IN ('chat2db_native_parent','chat2db_native_child') AND c.column_id IS NOT NULL AND i.index_id IS NOT NULL", + ) + .await + .map_err(sqlserver_query_error)? + .into_first_result() + .await + .map_err(sqlserver_query_error)?; + assert_eq!(row_i32(&metadata[0], 0)?, 2); + Ok::<(), AppError>(()) + } + .await; + + live_cleanup(&mut conn.client).await; + finish_connection(conn, smoke_result) + .await + .expect("native SQL Server product smoke should pass"); + } +} diff --git a/crates/chat2db-core/src/ssh.rs b/crates/chat2db-core/src/ssh.rs index 9fc0e1a..4ad0aef 100644 --- a/crates/chat2db-core/src/ssh.rs +++ b/crates/chat2db-core/src/ssh.rs @@ -906,8 +906,9 @@ mod tests { .port(); let (shutdown, receiver) = oneshot::channel(); let task = tokio::spawn(async move { - let _listener = listener; + let listener = listener; let _ = receiver.await; + drop(listener); shutdowns.fetch_add(1, Ordering::SeqCst); Ok(()) }); diff --git a/crates/chat2db-core/tests/java_h2_product.rs b/crates/chat2db-core/tests/java_h2_product.rs index d025c85..1d73e7e 100644 --- a/crates/chat2db-core/tests/java_h2_product.rs +++ b/crates/chat2db-core/tests/java_h2_product.rs @@ -135,10 +135,10 @@ async fn assert_community_disabled(application: &Application) { database_type: "H2".to_owned(), }) .await - .expect_err("unconfigured Community database metadata must stay disabled"); + .expect_err("unknown datasource must fail before Community engine acquisition"); assert_eq!( disabled_database_error.api_error().code, - "community_compatibility_disabled" + "datasource_not_found" ); let disabled_table_error = application .list_community_tables(ListCommunityTablesRequest { @@ -149,10 +149,10 @@ async fn assert_community_disabled(application: &Application) { table_name_pattern: "%".to_owned(), }) .await - .expect_err("unconfigured Community table metadata must stay disabled"); + .expect_err("unknown datasource must fail before Community engine acquisition"); assert_eq!( disabled_table_error.api_error().code, - "community_compatibility_disabled" + "datasource_not_found" ); let disabled_column_error = application .list_community_columns(ListCommunityColumnsRequest { @@ -163,10 +163,10 @@ async fn assert_community_disabled(application: &Application) { table_name: "items".to_owned(), }) .await - .expect_err("unconfigured Community column metadata must stay disabled"); + .expect_err("unknown datasource must fail before Community engine acquisition"); assert_eq!( disabled_column_error.api_error().code, - "community_compatibility_disabled" + "datasource_not_found" ); let disabled_index_error = application .list_community_indexes(ListCommunityIndexesRequest { @@ -177,10 +177,10 @@ async fn assert_community_disabled(application: &Application) { table_name: "items".to_owned(), }) .await - .expect_err("unconfigured Community index metadata must stay disabled"); + .expect_err("unknown datasource must fail before Community engine acquisition"); assert_eq!( disabled_index_error.api_error().code, - "community_compatibility_disabled" + "datasource_not_found" ); } @@ -198,7 +198,7 @@ async fn runtime_host_open_keeps_java_dormant() { .expect("opening storage must not spawn the missing Java executable"); assert_engine_available_on_demand(&host.application()); let inventory = host.application().list_drivers(); - assert_eq!(inventory.items.len(), 1); + assert_eq!(inventory.items.len(), 4); assert_native_mysql_driver(&inventory.items); host.shutdown() .await @@ -233,7 +233,7 @@ async fn managed_h2_starts_on_demand_and_reloads_after_idle_shutdown() { let application = host.application(); assert_engine_available_on_demand(&application); let inventory = application.list_drivers(); - assert_eq!(inventory.items.len(), 2); + assert_eq!(inventory.items.len(), 5); assert_native_mysql_driver(&inventory.items); let installed = managed_driver(&inventory.items, "h2"); assert_eq!(installed.pack_id, "h2"); @@ -429,7 +429,7 @@ async fn partial_managed_driver_preload_cleans_generation_and_releases_storage() )) .await .expect("driver discovery must not start Java"); - assert_eq!(host.application().list_drivers().items.len(), 3); + assert_eq!(host.application().list_drivers().items.len(), 6); let error = host .acquire_engine() .await @@ -451,7 +451,7 @@ async fn partial_managed_driver_preload_cleans_generation_and_releases_storage() )) .await .expect("storage and driver discovery must reopen immediately"); - assert_eq!(host.application().list_drivers().items.len(), 2); + assert_eq!(host.application().list_drivers().items.len(), 5); let lease = host .acquire_engine() .await diff --git a/crates/chat2db-core/tests/native_oracle_smoke.rs b/crates/chat2db-core/tests/native_oracle_smoke.rs new file mode 100644 index 0000000..a07b101 --- /dev/null +++ b/crates/chat2db-core/tests/native_oracle_smoke.rs @@ -0,0 +1,1158 @@ +use std::{panic::AssertUnwindSafe, time::Duration}; + +use base64::{Engine as _, engine::general_purpose::STANDARD}; +use chat2db_contract::{ + ComponentState, CreateDatasourceRequest, DatasourceConnection, DatasourceConnectionProperty, + GetCommunityFunctionRequest, GetCommunityProcedureRequest, GetCommunityTriggerRequest, + JdbcValue, ListCommunityColumnsRequest, ListCommunityDatabasesRequest, + ListCommunityFunctionsRequest, ListCommunityIndexesRequest, ListCommunityProceduresRequest, + ListCommunitySchemasRequest, ListCommunityTableKeysRequest, ListCommunityTablesRequest, + ListCommunityTriggersRequest, ListCommunityViewsRequest, OperationEvent, QueryLimits, + QueryParameter, ResultMetadata, ResultPageRequest, StartCommunityTablePreviewRequest, + StartQueryRequest, +}; +use chat2db_core::{ + Application, NativeConsoleCancellation, NativeConsoleRequest, RuntimeConfig, RuntimeHost, +}; +use chat2db_java_bridge::{EngineCommand, EngineConfig}; +use futures_util::FutureExt as _; +use oracle_rs::{Config, Connection, Error as OracleError, Value}; +use tempfile::TempDir; +use uuid::Uuid; + +const ORACLE_DATABASE_TYPE: &str = "ORACLE"; +const EVENT_TIMEOUT: Duration = Duration::from_secs(15); + +struct OracleTestConfig { + host: String, + port: u16, + service: String, + username: String, + password: String, +} + +impl OracleTestConfig { + fn from_environment() -> Self { + let host = required_env("CHAT2DB_ORACLE_HOST"); + assert!( + !host.trim().is_empty() + && !host.chars().any(char::is_control) + && !host.contains(['/', '?', '#', ':']), + "CHAT2DB_ORACLE_HOST must be a valid IPv4 address or hostname" + ); + let port = std::env::var("CHAT2DB_ORACLE_PORT") + .unwrap_or_else(|_| "1521".to_owned()) + .parse::() + .expect("CHAT2DB_ORACLE_PORT must be a TCP port"); + assert_ne!(port, 0, "CHAT2DB_ORACLE_PORT cannot be zero"); + let service = required_env("CHAT2DB_ORACLE_SERVICE"); + assert!(!service.trim().is_empty(), "Oracle service cannot be empty"); + let username = required_env("CHAT2DB_ORACLE_USERNAME"); + assert!(!username.is_empty(), "Oracle username cannot be empty"); + Self { + host, + port, + service, + username, + password: required_env("CHAT2DB_ORACLE_PASSWORD"), + } + } + + fn driver_config(&self) -> Config { + Config::new( + self.host.clone(), + self.port, + self.service.clone(), + self.username.clone(), + self.password.clone(), + ) + } + + fn connection(&self) -> DatasourceConnection { + DatasourceConnection { + jdbc_url: format!( + "jdbc:oracle:thin:@{}:{}/{}", + self.host, self.port, self.service + ), + properties: vec![ + DatasourceConnectionProperty { + key: "user".to_owned(), + value: self.username.clone(), + sensitive: false, + }, + DatasourceConnectionProperty { + key: "password".to_owned(), + value: self.password.clone(), + sensitive: true, + }, + ], + read_only: false, + ssh: None, + } + } + + fn read_only_connection(&self) -> DatasourceConnection { + let mut connection = self.connection(); + connection.read_only = true; + connection + } +} + +struct OracleFixture { + parent_table: String, + table: String, + index: String, + foreign_key: String, + view: String, + function: String, + side_effect_function: String, + procedure: String, + trigger: String, +} + +impl OracleFixture { + fn unique() -> Self { + let suffix = Uuid::new_v4().simple().to_string()[..8].to_ascii_uppercase(); + Self { + parent_table: format!("C2P_{suffix}"), + table: format!("C2T_{suffix}"), + index: format!("C2I_{suffix}"), + foreign_key: format!("C2K_{suffix}"), + view: format!("C2V_{suffix}"), + function: format!("C2F_{suffix}"), + side_effect_function: format!("C2S_{suffix}"), + procedure: format!("C2R_{suffix}"), + trigger: format!("C2G_{suffix}"), + } + } +} + +#[tokio::test] +#[ignore = "requires a reachable Oracle 12.1+ database"] +async fn connects_and_selects_one() { + let config = OracleTestConfig::from_environment(); + let connection = Connection::connect_with_config(config.driver_config()) + .await + .expect("Oracle connection should succeed"); + let result = connection + .query("SELECT 1 FROM DUAL", &[]) + .await + .expect("SELECT 1 should succeed"); + assert_eq!(result.rows.len(), 1); + match result.rows[0].get(0) { + Some(Value::Integer(1)) => {} + Some(Value::String(value)) if value == "1" => {} + other => panic!("SELECT 1 returned an unexpected value: {other:?}"), + } + connection.close().await.expect("connection should close"); +} + +#[tokio::test] +#[ignore = "requires CHAT2DB_ORACLE_* variables and a reachable Oracle 12.1+ database"] +async fn native_oracle_product_paths_keep_java_dormant() { + let config = OracleTestConfig::from_environment(); + let fixture = OracleFixture::unique(); + let verification = AssertUnwindSafe(async { + provision_fixture(&config, &fixture).await; + verify_native_product(&config, &fixture).await; + }) + .catch_unwind() + .await; + let cleanup = cleanup_fixture(&config, &fixture).await; + if let Err(payload) = verification { + if let Err(error) = cleanup { + eprintln!("native Oracle cleanup also failed: {error}"); + } + std::panic::resume_unwind(payload); + } + cleanup.expect("native Oracle fixture must be removed"); +} + +async fn verify_native_product(config: &OracleTestConfig, fixture: &OracleFixture) { + let directory = TempDir::new().expect("temporary native Oracle runtime"); + let missing_java = directory.path().join("missing-java"); + let runtime = RuntimeConfig::new(EngineConfig::new(EngineCommand::new(missing_java))) + .with_data_dir(directory.path().join("data")) + .with_vault_master_key_base64(STANDARD.encode([0x72; 32])); + let mut host = RuntimeHost::open(runtime) + .await + .expect("native Oracle runtime must open without Java"); + let application = host.application(); + assert_java_dormant(&application); + + let drivers = application.list_drivers(); + let oracle_driver = drivers + .items + .iter() + .find(|driver| driver.driver_id == "oracle") + .expect("native Oracle driver must be present in the driver inventory"); + assert_eq!(oracle_driver.driver_class, "rust:oracle-rs"); + assert_eq!(oracle_driver.artifact_count, 0); + + application + .test_datasource_connection("oracle", config.connection()) + .await + .expect("native Oracle connection test must succeed without a JDBC pack"); + assert_java_dormant(&application); + + let datasource = application + .create_datasource(CreateDatasourceRequest { + name: "Native Oracle".to_owned(), + driver_id: "oracle".to_owned(), + connection: Some(config.connection()), + }) + .await + .expect("native Oracle datasource must persist without a JDBC pack"); + + let (database_name, schema_name) = + verify_database_and_schema_metadata(&application, &datasource.id).await; + verify_table_metadata( + &application, + &datasource.id, + &database_name, + &schema_name, + fixture, + ) + .await; + verify_view_routine_and_trigger_metadata( + &application, + &datasource.id, + &database_name, + &schema_name, + fixture, + ) + .await; + verify_query_console_and_preview( + &application, + &datasource.id, + &database_name, + &schema_name, + fixture, + ) + .await; + verify_type_matrix( + &application, + &datasource.id, + &database_name, + &schema_name, + fixture, + ) + .await; + verify_read_only_side_effect(&application, config, &datasource.id, &schema_name, fixture).await; + assert_java_dormant(&application); + + host.shutdown() + .await + .expect("native-only Oracle runtime must shut down cleanly"); +} + +async fn verify_database_and_schema_metadata( + application: &Application, + datasource_id: &str, +) -> (String, String) { + let databases = application + .list_community_databases(ListCommunityDatabasesRequest { + datasource_id: datasource_id.to_owned(), + database_type: ORACLE_DATABASE_TYPE.to_owned(), + }) + .await + .expect("native Oracle databases must list"); + let database = databases + .items + .first() + .expect("Oracle must expose its current database"); + assert!(!database.name.is_empty()); + assert!(!database.owner.is_empty()); + + let schemas = application + .list_community_schemas(ListCommunitySchemasRequest { + datasource_id: datasource_id.to_owned(), + database_type: ORACLE_DATABASE_TYPE.to_owned(), + database_name: database.name.clone(), + }) + .await + .expect("native Oracle schemas must list"); + assert!( + schemas + .items + .iter() + .any(|schema| schema.name.eq_ignore_ascii_case(&database.owner)) + ); + (database.name.clone(), database.owner.to_ascii_uppercase()) +} + +async fn verify_table_metadata( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + fixture: &OracleFixture, +) { + let tables = application + .list_community_tables(ListCommunityTablesRequest { + datasource_id: datasource_id.to_owned(), + database_type: ORACLE_DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + table_name_pattern: fixture.table.clone(), + }) + .await + .expect("native Oracle tables must list"); + assert!( + tables + .items + .iter() + .any(|table| table.name == fixture.table && table.table_type == "TABLE") + ); + + verify_column_metadata( + application, + datasource_id, + database_name, + schema_name, + fixture, + ) + .await; + + let indexes = application + .list_community_indexes(ListCommunityIndexesRequest { + datasource_id: datasource_id.to_owned(), + database_type: ORACLE_DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + table_name: fixture.table.clone(), + }) + .await + .expect("native Oracle indexes must list"); + assert!(indexes.items.iter().any(|index| { + index.name == fixture.index + && index + .columns + .iter() + .any(|column| column.column_name == "LABEL") + })); + + let child_keys = table_keys(datasource_id, database_name, schema_name, &fixture.table); + let primary = application + .list_community_primary_keys(child_keys.clone()) + .await + .expect("native Oracle primary keys must list"); + assert!(primary.items.iter().any(|key| key.column_name == "ID")); + let imported = application + .list_community_imported_keys(child_keys) + .await + .expect("native Oracle imported keys must list"); + assert!(imported.items.iter().any(|key| { + key.foreign_key_name == fixture.foreign_key + && key.primary_table_name == fixture.parent_table + && key.foreign_table_name == fixture.table + })); + let exported = application + .list_community_exported_keys(table_keys( + datasource_id, + database_name, + schema_name, + &fixture.parent_table, + )) + .await + .expect("native Oracle exported keys must list"); + assert!( + exported + .items + .iter() + .any(|key| key.foreign_key_name == fixture.foreign_key) + ); + + let ddl = application + .table_ddl(datasource_id, database_name, schema_name, &fixture.table) + .await + .expect("native Oracle table DDL must load"); + assert!(ddl.to_ascii_uppercase().contains(&fixture.table)); + assert!(ddl.to_ascii_uppercase().contains("CREATE TABLE")); +} + +async fn verify_column_metadata( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + fixture: &OracleFixture, +) { + let columns = application + .list_community_columns(ListCommunityColumnsRequest { + datasource_id: datasource_id.to_owned(), + database_type: ORACLE_DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + table_name: fixture.table.clone(), + }) + .await + .expect("native Oracle columns must list"); + assert!(columns.items.iter().any(|column| { + column.name == "ID" && column.column_type == "NUMBER" && column.primary_key == Some(true) + })); + assert!( + columns + .items + .iter() + .any(|column| column.name == "LABEL" && column.column_type == "VARCHAR2") + ); + assert!(columns.items.iter().any(|column| { + column.name == "SCORE_X2" + && column.column_type == "NUMBER" + && column.generated_column == Some(true) + })); + assert!( + columns + .items + .iter() + .all(|column| column.name != "HIDDEN_VALUE"), + "Oracle hidden columns must not leak through native metadata" + ); +} + +async fn verify_view_routine_and_trigger_metadata( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + fixture: &OracleFixture, +) { + verify_view_metadata( + application, + datasource_id, + database_name, + schema_name, + fixture, + ) + .await; + verify_function_metadata( + application, + datasource_id, + database_name, + schema_name, + fixture, + ) + .await; + verify_procedure_metadata( + application, + datasource_id, + database_name, + schema_name, + fixture, + ) + .await; + verify_trigger_metadata( + application, + datasource_id, + database_name, + schema_name, + fixture, + ) + .await; +} + +async fn verify_view_metadata( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + fixture: &OracleFixture, +) { + let views_request = ListCommunityViewsRequest { + datasource_id: datasource_id.to_owned(), + database_type: ORACLE_DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + view_name_pattern: fixture.view.clone(), + }; + let views = application + .list_community_views(views_request.clone()) + .await + .expect("native Oracle views must list"); + assert!(views.items.iter().any(|view| view.name == fixture.view)); + let view = application + .get_community_view(views_request) + .await + .expect("native Oracle view detail must load"); + assert!(view.ddl.to_ascii_uppercase().contains(&fixture.view)); +} + +async fn verify_function_metadata( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + fixture: &OracleFixture, +) { + let functions = application + .list_community_functions(ListCommunityFunctionsRequest { + datasource_id: datasource_id.to_owned(), + database_type: ORACLE_DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + }) + .await + .expect("native Oracle functions must list"); + assert!( + functions + .items + .iter() + .any(|function| function.name == fixture.function) + ); + let function_request = GetCommunityFunctionRequest { + datasource_id: datasource_id.to_owned(), + database_type: ORACLE_DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + function_name: fixture.function.clone(), + }; + let function = application + .get_community_function(function_request.clone()) + .await + .expect("native Oracle function detail must load"); + assert!( + function + .body + .to_ascii_uppercase() + .contains(&fixture.function) + ); + let function_parameters = application + .list_community_function_parameters(function_request) + .await + .expect("native Oracle function parameters must list"); + assert!( + function_parameters + .items + .iter() + .any(|parameter| parameter.column_name == "P_VALUE") + ); +} + +async fn verify_procedure_metadata( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + fixture: &OracleFixture, +) { + let procedures = application + .list_community_procedures(ListCommunityProceduresRequest { + datasource_id: datasource_id.to_owned(), + database_type: ORACLE_DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + }) + .await + .expect("native Oracle procedures must list"); + assert!( + procedures + .items + .iter() + .any(|procedure| procedure.name == fixture.procedure) + ); + let procedure_request = GetCommunityProcedureRequest { + datasource_id: datasource_id.to_owned(), + database_type: ORACLE_DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + procedure_name: fixture.procedure.clone(), + }; + let procedure = application + .get_community_procedure(procedure_request.clone()) + .await + .expect("native Oracle procedure detail must load"); + assert!( + procedure + .body + .to_ascii_uppercase() + .contains(&fixture.procedure) + ); + let procedure_parameters = application + .list_community_procedure_parameters(procedure_request) + .await + .expect("native Oracle procedure parameters must list"); + assert!( + procedure_parameters + .items + .iter() + .any(|parameter| parameter.column_name == "P_OUTPUT") + ); +} + +async fn verify_trigger_metadata( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + fixture: &OracleFixture, +) { + let triggers = application + .list_community_triggers(ListCommunityTriggersRequest { + datasource_id: datasource_id.to_owned(), + database_type: ORACLE_DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + }) + .await + .expect("native Oracle triggers must list"); + assert!( + triggers + .items + .iter() + .any(|trigger| trigger.name == fixture.trigger) + ); + let trigger = application + .get_community_trigger(GetCommunityTriggerRequest { + datasource_id: datasource_id.to_owned(), + database_type: ORACLE_DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + trigger_name: fixture.trigger.clone(), + }) + .await + .expect("native Oracle trigger detail must load"); + assert!(trigger.event_manipulation.contains("INSERT")); + assert!(trigger.body.to_ascii_uppercase().contains(&fixture.trigger)); +} + +async fn verify_query_console_and_preview( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + fixture: &OracleFixture, +) { + let query = application + .start_query(StartQueryRequest { + datasource_id: datasource_id.to_owned(), + sql: format!( + "SELECT LABEL, SCORE FROM \"{schema_name}\".\"{}\" WHERE ID = :1", + fixture.table + ), + parameters: vec![QueryParameter { + position: 1, + value: JdbcValue::SignedInteger { + value: "1".to_owned(), + }, + }], + limits: query_limits("10"), + }) + .await + .expect("native Oracle query must be accepted"); + let result = wait_for_result(application, &query.operation_id).await; + assert_eq!(result.row_count, "1"); + let page = result_page(application, &result).await; + assert!(matches!( + &page.rows[0].values[0], + JdbcValue::Text { value } if value == "alpha" + )); + + let console = application + .execute_native_console( + NativeConsoleRequest { + datasource_id: datasource_id.to_owned(), + database_name: database_name.to_owned(), + sql: format!( + "SELECT LABEL FROM \"{schema_name}\".\"{}\" ORDER BY ID", + fixture.table + ), + page_no: 1, + page_size: 10, + result_set_id: None, + single: false, + page_size_all: false, + explain: false, + error_continue: false, + }, + NativeConsoleCancellation::new(), + ) + .await + .expect("native Oracle Console query must execute"); + assert_eq!(console.len(), 1); + assert!(console[0].success); + assert_eq!(console[0].row_count, 2); + assert!(matches!( + &console[0].rows[0].values[0], + JdbcValue::Text { value } if value == "alpha" + )); + + let preview = application + .start_community_table_preview(StartCommunityTablePreviewRequest { + datasource_id: datasource_id.to_owned(), + database_type: ORACLE_DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + table_name: fixture.table.clone(), + row_limit: Some(2), + }) + .await + .expect("native Oracle table preview must be accepted"); + assert!(preview.sql.contains(&fixture.table)); + let preview_result = wait_for_result(application, &preview.operation_id).await; + assert_eq!(preview_result.row_count, "2"); + assert_eq!( + result_page(application, &preview_result).await.rows.len(), + 2 + ); +} + +async fn verify_type_matrix( + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + fixture: &OracleFixture, +) { + let sql = oracle_type_matrix_sql(schema_name, fixture); + let query = application + .start_query(StartQueryRequest { + datasource_id: datasource_id.to_owned(), + sql: sql.clone(), + parameters: Vec::new(), + limits: query_limits("10"), + }) + .await + .expect("native Oracle type matrix query must be accepted"); + let result = wait_for_result(application, &query.operation_id).await; + let page = result_page(application, &result).await; + assert_eq!(page.rows.len(), 1); + assert_oracle_type_matrix(&page.rows[0].values); + + let console = application + .execute_native_console( + NativeConsoleRequest { + datasource_id: datasource_id.to_owned(), + database_name: database_name.to_owned(), + sql, + page_no: 1, + page_size: 10, + result_set_id: None, + single: false, + page_size_all: false, + explain: false, + error_continue: false, + }, + NativeConsoleCancellation::new(), + ) + .await + .expect("native Oracle Console type matrix must execute"); + assert_eq!(console.len(), 1); + assert_eq!(console[0].rows.len(), 1); + assert_oracle_type_matrix(&console[0].rows[0].values); + + verify_unsupported_type_result( + application, + datasource_id, + database_name, + "SELECT CAST(1.25 AS BINARY_FLOAT), CAST(2.5 AS BINARY_DOUBLE) FROM DUAL", + ) + .await; + verify_unsupported_type_result( + application, + datasource_id, + database_name, + &format!( + "SELECT TO_CLOB('probe'), ROWID FROM \"{schema_name}\".\"{}\" WHERE ID = 1", + fixture.table + ), + ) + .await; +} + +async fn verify_unsupported_type_result( + application: &Application, + datasource_id: &str, + database_name: &str, + sql: &str, +) { + let query = application + .start_query(StartQueryRequest { + datasource_id: datasource_id.to_owned(), + sql: sql.to_owned(), + parameters: Vec::new(), + limits: query_limits("10"), + }) + .await + .expect("unsupported Oracle type query must be accepted before column describe"); + let error = wait_for_failure(application, &query.operation_id).await; + assert_eq!( + error.code, "oracle_result_type_not_supported", + "unsupported SQL must fail after describe: {sql}" + ); + let error = application + .execute_native_console( + NativeConsoleRequest { + datasource_id: datasource_id.to_owned(), + database_name: database_name.to_owned(), + sql: sql.to_owned(), + page_no: 1, + page_size: 10, + result_set_id: None, + single: false, + page_size_all: false, + explain: false, + error_continue: false, + }, + NativeConsoleCancellation::new(), + ) + .await + .expect_err("Oracle Console must reject lossy native result types"); + assert_eq!(error.api_error().code, "oracle_result_type_not_supported"); +} + +fn oracle_type_matrix_sql(schema_name: &str, fixture: &OracleFixture) -> String { + format!( + "SELECT CAST(123.45 AS NUMBER(10,2)) AS NUMBER_VALUE, \ + HEXTORAW('00FF') AS RAW_VALUE, \ + TO_BLOB(HEXTORAW('0102')) AS BLOB_VALUE, \ + TO_CLOB('clob-value') AS CLOB_VALUE, \ + CAST(DATE '2026-08-07' AS DATE) AS DATE_VALUE, \ + TIMESTAMP '2026-08-07 12:34:56.123456' AS TIMESTAMP_VALUE, \ + TIMESTAMP '2026-08-07 12:34:56.123456 +08:00' AS TIMESTAMP_TZ_VALUE, \ + JSON('{{\"ready\":true}}') AS JSON_VALUE, \ + TRUE AS BOOLEAN_VALUE \ + FROM \"{schema_name}\".\"{}\" WHERE ID = 1", + fixture.table + ) +} + +fn assert_oracle_type_matrix(values: &[JdbcValue]) { + assert_eq!(values.len(), 9); + assert!(matches!( + &values[0], + JdbcValue::Decimal { value } + if value.parse::().is_ok_and(|value| (value - 123.45).abs() < f64::EPSILON) + )); + assert!(matches!(&values[1], JdbcValue::Binary { value } if value == "AP8=")); + assert!(matches!(&values[2], JdbcValue::Binary { value } if value == "AQI=")); + assert!(matches!(&values[3], JdbcValue::Text { value } if value == "clob-value")); + assert!( + matches!( + &values[4], + JdbcValue::Timestamp { value } if value == "2026-08-07T00:00:00" + ), + "unexpected Oracle type matrix: {values:?}" + ); + assert!(matches!( + &values[5], + JdbcValue::Timestamp { value } if value == "2026-08-07T12:34:56.123456" + )); + assert!(matches!( + &values[6], + JdbcValue::TimestampWithTimeZone { value } + if value == "2026-08-07T12:34:56.123456+08:00" + )); + assert!(matches!( + &values[7], + JdbcValue::Json { value } if value == "{\"ready\":true}" + )); + assert!(matches!(&values[8], JdbcValue::Boolean { value: true })); +} + +async fn verify_read_only_side_effect( + application: &Application, + config: &OracleTestConfig, + verification_datasource_id: &str, + schema_name: &str, + fixture: &OracleFixture, +) { + let read_only = application + .create_datasource(CreateDatasourceRequest { + name: "Native Oracle read only".to_owned(), + driver_id: "oracle".to_owned(), + connection: Some(config.read_only_connection()), + }) + .await + .expect("read-only native Oracle datasource must persist"); + let query = application + .start_query(StartQueryRequest { + datasource_id: read_only.id, + sql: format!( + "SELECT \"{schema_name}\".\"{}\"() FROM DUAL", + fixture.side_effect_function + ), + parameters: Vec::new(), + limits: query_limits("10"), + }) + .await + .expect("read-only Oracle side-effect query must be accepted before dispatch"); + let error = wait_for_failure(application, &query.operation_id).await; + assert!( + matches!( + error.code.as_str(), + "oracle_query_failed" | "oracle_connection_failed" + ), + "read-only Oracle side-effect rejection returned an unexpected error: {error:?}" + ); + + let verification = application + .start_query(StartQueryRequest { + datasource_id: verification_datasource_id.to_owned(), + sql: format!( + "SELECT SCORE FROM \"{schema_name}\".\"{}\" WHERE ID = 1", + fixture.table + ), + parameters: Vec::new(), + limits: query_limits("10"), + }) + .await + .expect("Oracle side-effect verification query must be accepted"); + let result = wait_for_result(application, &verification.operation_id).await; + let page = result_page(application, &result).await; + assert!(matches!( + &page.rows[0].values[0], + JdbcValue::Decimal { value } + if value.parse::().is_ok_and(|value| (value - 10.5).abs() < f64::EPSILON) + )); +} + +fn table_keys( + datasource_id: &str, + database_name: &str, + schema_name: &str, + table_name: &str, +) -> ListCommunityTableKeysRequest { + ListCommunityTableKeysRequest { + datasource_id: datasource_id.to_owned(), + database_type: ORACLE_DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: schema_name.to_owned(), + table_name: table_name.to_owned(), + } +} + +fn query_limits(max_rows: &str) -> QueryLimits { + QueryLimits { + max_rows: max_rows.to_owned(), + max_result_bytes: (8_u64 * 1024 * 1024).to_string(), + batch_rows: 2, + batch_bytes: 1024 * 1024, + result_ttl_seconds: 60, + } +} + +async fn wait_for_result(application: &Application, operation_id: &str) -> ResultMetadata { + let mut subscription = application + .subscribe_operation(operation_id, None) + .await + .expect("native Oracle query operation must be subscribable"); + tokio::time::timeout(EVENT_TIMEOUT, async { + while let Some(envelope) = subscription + .next_event() + .await + .expect("native Oracle operation event must decode") + { + match envelope.event { + OperationEvent::Completed { result } => return result, + OperationEvent::Failed { error } => { + panic!("native Oracle query failed: {error:?}") + } + OperationEvent::Cancelled { reason } => { + panic!("native Oracle query was cancelled: {reason:?}") + } + OperationEvent::Started | OperationEvent::Progress { .. } => {} + } + } + panic!("native Oracle operation ended without a terminal event") + }) + .await + .expect("native Oracle query must finish before timeout") +} + +async fn wait_for_failure( + application: &Application, + operation_id: &str, +) -> chat2db_contract::ApiError { + let mut subscription = application + .subscribe_operation(operation_id, None) + .await + .expect("failed native Oracle operation must be subscribable"); + tokio::time::timeout(EVENT_TIMEOUT, async { + while let Some(envelope) = subscription + .next_event() + .await + .expect("native Oracle operation event must decode") + { + match envelope.event { + OperationEvent::Failed { error } => return error, + OperationEvent::Completed { result } => { + panic!("native Oracle query unexpectedly completed: {result:?}") + } + OperationEvent::Cancelled { reason } => { + panic!("native Oracle query was cancelled: {reason:?}") + } + OperationEvent::Started | OperationEvent::Progress { .. } => {} + } + } + panic!("native Oracle operation ended without a failure event") + }) + .await + .expect("native Oracle query must fail before timeout") +} + +async fn result_page( + application: &Application, + result: &ResultMetadata, +) -> chat2db_contract::ResultPage { + application + .result_page( + &result.id, + ResultPageRequest { + offset: "0".to_owned(), + max_rows: "20".to_owned(), + max_bytes: (8_u64 * 1024 * 1024).to_string(), + }, + ) + .await + .expect("native Oracle result page must be retained") +} + +async fn provision_fixture(config: &OracleTestConfig, fixture: &OracleFixture) { + let connection = Connection::connect_with_config(config.driver_config()) + .await + .expect("native Oracle fixture connection must open"); + for sql in fixture_ddl(fixture) { + connection + .execute(&sql, &[]) + .await + .unwrap_or_else(|error| panic!("Oracle fixture SQL failed: {sql}: {error}")); + } + connection + .execute( + &format!("INSERT INTO {} (ID) VALUES (1)", fixture.parent_table), + &[], + ) + .await + .expect("Oracle parent fixture row must insert"); + connection + .execute( + &format!( + "INSERT ALL \ + INTO {} (ID, PARENT_ID, LABEL, SCORE) VALUES (1, 1, ' alpha ', 10.50) \ + INTO {} (ID, PARENT_ID, LABEL, SCORE) VALUES (2, 1, 'beta', 20.25) \ + SELECT 1 FROM DUAL", + fixture.table, fixture.table + ), + &[], + ) + .await + .expect("Oracle child fixture rows must insert"); + connection + .commit() + .await + .expect("Oracle fixture transaction must commit"); + connection + .close() + .await + .expect("Oracle fixture connection must close"); +} + +fn fixture_ddl(fixture: &OracleFixture) -> Vec { + vec![ + format!( + "CREATE TABLE {} (ID NUMBER(10) CONSTRAINT C2PP_{} PRIMARY KEY)", + fixture.parent_table, + fixture_suffix(&fixture.parent_table) + ), + format!( + "CREATE TABLE {} (\ + ID NUMBER(10) CONSTRAINT C2TP_{} PRIMARY KEY, \ + PARENT_ID NUMBER(10) NOT NULL, LABEL VARCHAR2(64) NOT NULL, \ + SCORE NUMBER(10,2), \ + SCORE_X2 NUMBER GENERATED ALWAYS AS (SCORE * 2) VIRTUAL, \ + HIDDEN_VALUE NUMBER INVISIBLE, \ + CREATED_AT TIMESTAMP DEFAULT CURRENT_TIMESTAMP, \ + CONSTRAINT {} FOREIGN KEY (PARENT_ID) REFERENCES {} (ID))", + fixture.table, + fixture_suffix(&fixture.table), + fixture.foreign_key, + fixture.parent_table + ), + format!( + "CREATE INDEX {} ON {} (LABEL)", + fixture.index, fixture.table + ), + format!( + "CREATE VIEW {} AS SELECT ID, LABEL FROM {}", + fixture.view, fixture.table + ), + format!( + "CREATE FUNCTION {} (P_VALUE IN NUMBER) RETURN NUMBER IS \ + BEGIN RETURN P_VALUE + 1; END;", + fixture.function + ), + format!( + "CREATE FUNCTION {} RETURN NUMBER IS \ + BEGIN UPDATE {} SET SCORE = SCORE + 100 WHERE ID = 1; \ + RETURN SQL%ROWCOUNT; END;", + fixture.side_effect_function, fixture.table + ), + format!( + "CREATE PROCEDURE {} (P_INPUT IN NUMBER, P_OUTPUT OUT NUMBER) IS \ + BEGIN P_OUTPUT := P_INPUT + 1; END;", + fixture.procedure + ), + format!( + "CREATE TRIGGER {} BEFORE INSERT ON {} FOR EACH ROW \ + BEGIN :NEW.LABEL := TRIM(:NEW.LABEL); END;", + fixture.trigger, fixture.table + ), + ] +} + +async fn cleanup_fixture(config: &OracleTestConfig, fixture: &OracleFixture) -> Result<(), String> { + let connection = Connection::connect_with_config(config.driver_config()) + .await + .map_err(|error| error.to_string())?; + let drops = [ + format!("DROP TRIGGER {}", fixture.trigger), + format!("DROP FUNCTION {}", fixture.side_effect_function), + format!("DROP FUNCTION {}", fixture.function), + format!("DROP PROCEDURE {}", fixture.procedure), + format!("DROP VIEW {}", fixture.view), + format!("DROP TABLE {} PURGE", fixture.table), + format!("DROP TABLE {} PURGE", fixture.parent_table), + ]; + let mut cleanup_error = None; + for sql in drops { + if let Err(error) = connection.execute(&sql, &[]).await + && !missing_fixture_object(&error) + && cleanup_error.is_none() + { + cleanup_error = Some(format!("{sql}: {error}")); + } + } + if let Err(error) = connection.close().await + && cleanup_error.is_none() + { + cleanup_error = Some(error.to_string()); + } + cleanup_error.map_or(Ok(()), Err) +} + +fn missing_fixture_object(error: &OracleError) -> bool { + let message = error.to_string(); + ["ORA-00942", "ORA-04043", "ORA-04080"] + .iter() + .any(|code| message.contains(code)) +} + +fn fixture_suffix(name: &str) -> &str { + name.rsplit_once('_').map_or(name, |(_, suffix)| suffix) +} + +fn assert_java_dormant(application: &Application) { + let engine = application + .health() + .components + .into_iter() + .find(|component| component.id == "database-engine") + .expect("database engine health must be present"); + assert_eq!(engine.state, ComponentState::Ready); + assert_eq!(engine.detail, "Available on demand; Java is not running"); +} + +fn required_env(name: &str) -> String { + std::env::var(name).unwrap_or_else(|_| panic!("{name} must be configured")) +} diff --git a/crates/chat2db-core/tests/native_postgres_smoke.rs b/crates/chat2db-core/tests/native_postgres_smoke.rs new file mode 100644 index 0000000..bf2a37b --- /dev/null +++ b/crates/chat2db-core/tests/native_postgres_smoke.rs @@ -0,0 +1,1087 @@ +use std::{error::Error, fs, panic::AssertUnwindSafe, path::Path, time::Duration}; + +use base64::{Engine as _, engine::general_purpose::STANDARD}; +use chat2db_contract::{ + CommunityErQueryRequest, ComponentState, CreateDatasourceRequest, DatasourceConnection, + DatasourceConnectionProperty, GetCommunityFunctionRequest, GetCommunityProcedureRequest, + GetCommunityTriggerRequest, JdbcValue, ListCommunityColumnsRequest, + ListCommunityDatabasesRequest, ListCommunityFunctionsRequest, ListCommunityIndexesRequest, + ListCommunityProceduresRequest, ListCommunitySchemasRequest, ListCommunityTableKeysRequest, + ListCommunityTablesRequest, ListCommunityTriggersRequest, ListCommunityViewsRequest, + OperationEvent, QueryLimits, QueryParameter, ResultMetadata, ResultPageRequest, + StartCommunityTablePreviewRequest, StartQueryRequest, +}; +use chat2db_core::{ + Application, NativeConsoleCancellation, NativeConsoleRequest, RuntimeConfig, RuntimeHost, +}; +use chat2db_engine_protocol::wire; +use chat2db_java_bridge::{EngineCommand, EngineConfig}; +use futures_util::FutureExt as _; +use tempfile::TempDir; +use tokio_postgres::{Config, NoTls}; +use uuid::Uuid; + +const POSTGRES_DATABASE_TYPE: &str = "POSTGRESQL"; +const EVENT_TIMEOUT: Duration = Duration::from_secs(15); + +struct PostgresTestConfig { + host: String, + port: u16, + database: String, + username: String, + password: String, +} + +impl PostgresTestConfig { + fn from_environment() -> Self { + Self { + host: std::env::var("CHAT2DB_POSTGRES_HOST").unwrap_or_else(|_| "127.0.0.1".to_owned()), + port: std::env::var("CHAT2DB_POSTGRES_PORT") + .unwrap_or_else(|_| "5432".to_owned()) + .parse() + .expect("CHAT2DB_POSTGRES_PORT must be a TCP port"), + database: std::env::var("CHAT2DB_POSTGRES_DATABASE") + .unwrap_or_else(|_| "app".to_owned()), + username: std::env::var("CHAT2DB_POSTGRES_USER") + .unwrap_or_else(|_| "postgres".to_owned()), + password: std::env::var("CHAT2DB_POSTGRES_PASSWORD") + .unwrap_or_else(|_| "postgres".to_owned()), + } + } + + fn driver_config(&self) -> Config { + let mut config = Config::new(); + config + .host(&self.host) + .port(self.port) + .dbname(&self.database) + .user(&self.username) + .password(&self.password); + config + } + + fn connection(&self) -> DatasourceConnection { + DatasourceConnection { + jdbc_url: format!( + "jdbc:postgresql://{}:{}/{}?sslmode=disable", + self.host, self.port, self.database + ), + properties: vec![ + DatasourceConnectionProperty { + key: "user".to_owned(), + value: self.username.clone(), + sensitive: false, + }, + DatasourceConnectionProperty { + key: "password".to_owned(), + value: self.password.clone(), + sensitive: true, + }, + ], + read_only: false, + ssh: None, + } + } + + fn read_only_connection(&self) -> DatasourceConnection { + let mut connection = self.connection(); + connection.read_only = true; + connection + } +} + +struct PostgresFixture { + schema: String, + parent_table: String, + table: String, + index: String, + foreign_key: String, + view: String, + function: String, + procedure: String, + trigger: String, + trigger_function: String, +} + +impl PostgresFixture { + fn unique() -> Self { + let suffix = Uuid::new_v4().simple().to_string(); + let suffix = &suffix[..8]; + Self { + schema: format!("c2s_{suffix}"), + parent_table: format!("c2p_{suffix}"), + table: format!("c2t_{suffix}"), + index: format!("c2i_{suffix}"), + foreign_key: format!("c2k_{suffix}"), + view: format!("c2v_{suffix}"), + function: format!("c2f_{suffix}"), + procedure: format!("c2r_{suffix}"), + trigger: format!("c2g_{suffix}"), + trigger_function: format!("c2tf_{suffix}"), + } + } +} + +#[tokio::test] +#[ignore = "requires a reachable PostgreSQL database"] +async fn native_postgres_product_paths_keep_java_dormant() { + let config = PostgresTestConfig::from_environment(); + let fixture = PostgresFixture::unique(); + provision_fixture(&config, &fixture).await; + + let verification = AssertUnwindSafe(verify_native_product(&config, &fixture)) + .catch_unwind() + .await; + let cleanup = cleanup_fixture(&config, &fixture).await; + if let Err(payload) = verification { + if let Err(error) = cleanup { + eprintln!("native PostgreSQL cleanup also failed: {error}"); + } + std::panic::resume_unwind(payload); + } + cleanup.expect("native PostgreSQL fixture must be removed"); + assert_fixture_residue_zero(&config).await; +} + +async fn verify_native_product(config: &PostgresTestConfig, fixture: &PostgresFixture) { + let directory = TempDir::new().expect("temporary native PostgreSQL runtime"); + let missing_java = directory.path().join("missing-java"); + let data_dir = directory.path().join("data"); + let runtime = RuntimeConfig::new(EngineConfig::new(EngineCommand::new(missing_java))) + .with_data_dir(data_dir.clone()) + .with_vault_master_key_base64(STANDARD.encode([0x70; 32])); + let mut host = RuntimeHost::open(runtime) + .await + .expect("native PostgreSQL runtime must open without Java"); + let application = host.application(); + assert_java_dormant(&application); + + let drivers = application.list_drivers(); + let driver = drivers + .items + .iter() + .find(|driver| driver.driver_id == "postgresql") + .expect("native PostgreSQL driver must be present"); + assert_eq!(driver.driver_class, "rust:tokio-postgres"); + assert_eq!(driver.artifact_count, 0); + + application + .test_datasource_connection("postgresql", config.connection()) + .await + .expect("native PostgreSQL connection test must avoid Java"); + assert_java_dormant(&application); + + let datasource = application + .create_datasource(CreateDatasourceRequest { + name: "Native PostgreSQL".to_owned(), + driver_id: "postgresql".to_owned(), + connection: Some(config.connection()), + }) + .await + .expect("native PostgreSQL datasource must persist without a JDBC pack"); + + verify_oversized_scalar_cleanup(&application, &datasource.id, &data_dir).await; + verify_read_only_console(&application, config, fixture).await; + verify_database_schema_and_tables(&application, &datasource.id, config, fixture).await; + verify_views_routines_and_triggers(&application, &datasource.id, config, fixture).await; + verify_query_console_preview_and_er(&application, &datasource.id, config, fixture).await; + verify_console_write_outcomes(&application, &datasource.id, config, fixture).await; + assert_java_dormant(&application); + + host.shutdown() + .await + .expect("native-only PostgreSQL runtime must shut down cleanly"); +} + +async fn verify_oversized_scalar_cleanup( + application: &Application, + datasource_id: &str, + data_dir: &Path, +) { + assert!(retained_result_files(data_dir).is_empty()); + let scalar_bytes = wire::JdbcProtocolLimit::MaxScalarBytes as usize + 1; + let query = application + .start_query(StartQueryRequest { + datasource_id: datasource_id.to_owned(), + sql: format!("SELECT repeat('x', {scalar_bytes})"), + parameters: Vec::new(), + limits: query_limits("10"), + }) + .await + .expect("oversized PostgreSQL scalar query must be accepted before row decoding"); + let error = wait_for_failure(application, &query.operation_id).await; + assert_eq!(error.code, "postgres_scalar_too_large"); + assert!( + retained_result_files(data_dir).is_empty(), + "an aborted oversized result must not leave a retained result file" + ); + assert_java_dormant(application); +} + +async fn verify_read_only_console( + application: &Application, + config: &PostgresTestConfig, + fixture: &PostgresFixture, +) { + let datasource = application + .create_datasource(CreateDatasourceRequest { + name: "Native PostgreSQL read only".to_owned(), + driver_id: "postgresql".to_owned(), + connection: Some(config.read_only_connection()), + }) + .await + .expect("read-only native PostgreSQL datasource must persist"); + let request = |sql: String| NativeConsoleRequest { + datasource_id: datasource.id.clone(), + database_name: config.database.clone(), + sql, + page_no: 1, + page_size: 10, + result_set_id: None, + single: false, + page_size_all: false, + explain: false, + error_continue: false, + }; + let results = application + .execute_native_console( + request("SELECT 1 AS first_value; VALUES (2)".to_owned()), + NativeConsoleCancellation::new(), + ) + .await + .expect("a read-only Console must allow multiple read statements"); + assert_eq!(results.len(), 2); + assert!(results.iter().all(|result| result.success)); + + let error = application + .execute_native_console( + request(format!( + "SELECT 1; UPDATE \"{}\".\"{}\" SET label = 'changed' WHERE id = 1", + fixture.schema, fixture.table + )), + NativeConsoleCancellation::new(), + ) + .await + .expect_err("read-only validation must reject the entire script before dispatch"); + assert_eq!(error.api_error().code, "postgres_console_must_be_read_only"); + + let label = query_single_text( + config, + &format!( + "SELECT label FROM \"{}\".\"{}\" WHERE id = 1", + fixture.schema, fixture.table + ), + ) + .await; + assert_eq!(label, "alpha"); + assert_java_dormant(application); +} + +fn retained_result_files(data_dir: &Path) -> Vec { + fs::read_dir(data_dir.join("results")) + .expect("retained result directory") + .map(|entry| { + entry + .expect("retained result directory entry") + .file_name() + .to_string_lossy() + .into_owned() + }) + .collect() +} + +async fn verify_database_schema_and_tables( + application: &Application, + datasource_id: &str, + config: &PostgresTestConfig, + fixture: &PostgresFixture, +) { + let databases = application + .list_community_databases(ListCommunityDatabasesRequest { + datasource_id: datasource_id.to_owned(), + database_type: POSTGRES_DATABASE_TYPE.to_owned(), + }) + .await + .expect("native PostgreSQL databases must list"); + assert!( + databases + .items + .iter() + .any(|item| item.name == config.database) + ); + + let schemas = application + .list_community_schemas(ListCommunitySchemasRequest { + datasource_id: datasource_id.to_owned(), + database_type: POSTGRES_DATABASE_TYPE.to_owned(), + database_name: config.database.clone(), + }) + .await + .expect("native PostgreSQL schemas must list"); + assert!(schemas.items.iter().any(|item| item.name == fixture.schema)); + + let tables = application + .list_community_tables(ListCommunityTablesRequest { + datasource_id: datasource_id.to_owned(), + database_type: POSTGRES_DATABASE_TYPE.to_owned(), + database_name: config.database.clone(), + schema_name: fixture.schema.clone(), + table_name_pattern: fixture.table.clone(), + }) + .await + .expect("native PostgreSQL tables must list"); + assert!(tables.items.iter().any(|item| item.name == fixture.table)); + + let columns = application + .list_community_columns(ListCommunityColumnsRequest { + datasource_id: datasource_id.to_owned(), + database_type: POSTGRES_DATABASE_TYPE.to_owned(), + database_name: config.database.clone(), + schema_name: fixture.schema.clone(), + table_name: fixture.table.clone(), + }) + .await + .expect("native PostgreSQL columns must list"); + assert!(columns.items.iter().any(|column| { + column.name == "id" && column.column_type == "bigint" && column.primary_key == Some(true) + })); + assert!(columns.items.iter().any(|column| column.name == "label")); + + let indexes = application + .list_community_indexes(ListCommunityIndexesRequest { + datasource_id: datasource_id.to_owned(), + database_type: POSTGRES_DATABASE_TYPE.to_owned(), + database_name: config.database.clone(), + schema_name: fixture.schema.clone(), + table_name: fixture.table.clone(), + }) + .await + .expect("native PostgreSQL indexes must list"); + assert!( + indexes + .items + .iter() + .any(|index| index.name == fixture.index) + ); + + verify_table_keys_and_ddl(application, datasource_id, config, fixture).await; +} + +async fn verify_table_keys_and_ddl( + application: &Application, + datasource_id: &str, + config: &PostgresTestConfig, + fixture: &PostgresFixture, +) { + let keys = table_keys(datasource_id, config, fixture, &fixture.table); + let primary = application + .list_community_primary_keys(keys.clone()) + .await + .expect("native PostgreSQL primary keys must list"); + assert!(primary.items.iter().any(|key| key.column_name == "id")); + let imported = application + .list_community_imported_keys(keys) + .await + .expect("native PostgreSQL imported keys must list"); + assert!(imported.items.iter().any(|key| { + key.foreign_key_name == fixture.foreign_key + && key.primary_table_name == fixture.parent_table + && key.foreign_table_name == fixture.table + })); + let exported = application + .list_community_exported_keys(table_keys( + datasource_id, + config, + fixture, + &fixture.parent_table, + )) + .await + .expect("native PostgreSQL exported keys must list"); + assert!( + exported + .items + .iter() + .any(|key| key.foreign_key_name == fixture.foreign_key) + ); + + let ddl = application + .table_ddl( + datasource_id, + &config.database, + &fixture.schema, + &fixture.table, + ) + .await + .expect("native PostgreSQL table DDL must load"); + assert!(ddl.contains("CREATE TABLE")); + assert!(ddl.contains("FOREIGN KEY")); + assert!(ddl.contains(&fixture.index)); +} + +async fn verify_views_routines_and_triggers( + application: &Application, + datasource_id: &str, + config: &PostgresTestConfig, + fixture: &PostgresFixture, +) { + let views_request = ListCommunityViewsRequest { + datasource_id: datasource_id.to_owned(), + database_type: POSTGRES_DATABASE_TYPE.to_owned(), + database_name: config.database.clone(), + schema_name: fixture.schema.clone(), + view_name_pattern: fixture.view.clone(), + }; + let views = application + .list_community_views(views_request.clone()) + .await + .expect("native PostgreSQL views must list"); + assert!(views.items.iter().any(|view| view.name == fixture.view)); + let view = application + .get_community_view(views_request) + .await + .expect("native PostgreSQL view detail must load"); + assert!(view.ddl.contains(&fixture.view)); + + let functions = application + .list_community_functions(ListCommunityFunctionsRequest { + datasource_id: datasource_id.to_owned(), + database_type: POSTGRES_DATABASE_TYPE.to_owned(), + database_name: config.database.clone(), + schema_name: fixture.schema.clone(), + }) + .await + .expect("native PostgreSQL functions must list"); + assert!( + functions + .items + .iter() + .any(|item| item.name == fixture.function) + ); + let function_request = GetCommunityFunctionRequest { + datasource_id: datasource_id.to_owned(), + database_type: POSTGRES_DATABASE_TYPE.to_owned(), + database_name: config.database.clone(), + schema_name: fixture.schema.clone(), + function_name: fixture.function.clone(), + }; + let function = application + .get_community_function(function_request.clone()) + .await + .expect("native PostgreSQL function detail must load"); + assert!(function.body.contains(&fixture.function)); + let parameters = application + .list_community_function_parameters(function_request) + .await + .expect("native PostgreSQL function parameters must list"); + assert!( + parameters + .items + .iter() + .any(|parameter| parameter.column_name == "p_value") + ); + + verify_procedures_and_triggers(application, datasource_id, config, fixture).await; +} + +async fn verify_procedures_and_triggers( + application: &Application, + datasource_id: &str, + config: &PostgresTestConfig, + fixture: &PostgresFixture, +) { + let procedures = application + .list_community_procedures(ListCommunityProceduresRequest { + datasource_id: datasource_id.to_owned(), + database_type: POSTGRES_DATABASE_TYPE.to_owned(), + database_name: config.database.clone(), + schema_name: fixture.schema.clone(), + }) + .await + .expect("native PostgreSQL procedures must list"); + assert!( + procedures + .items + .iter() + .any(|item| item.name == fixture.procedure) + ); + let procedure_request = GetCommunityProcedureRequest { + datasource_id: datasource_id.to_owned(), + database_type: POSTGRES_DATABASE_TYPE.to_owned(), + database_name: config.database.clone(), + schema_name: fixture.schema.clone(), + procedure_name: fixture.procedure.clone(), + }; + let procedure = application + .get_community_procedure(procedure_request.clone()) + .await + .expect("native PostgreSQL procedure detail must load"); + assert!(procedure.body.contains(&fixture.procedure)); + let parameters = application + .list_community_procedure_parameters(procedure_request) + .await + .expect("native PostgreSQL procedure parameters must list"); + assert!( + parameters + .items + .iter() + .any(|parameter| parameter.column_name == "p_output") + ); + + let triggers = application + .list_community_triggers(ListCommunityTriggersRequest { + datasource_id: datasource_id.to_owned(), + database_type: POSTGRES_DATABASE_TYPE.to_owned(), + database_name: config.database.clone(), + schema_name: fixture.schema.clone(), + }) + .await + .expect("native PostgreSQL triggers must list"); + assert!( + triggers + .items + .iter() + .any(|item| item.name == fixture.trigger) + ); + let trigger = application + .get_community_trigger(GetCommunityTriggerRequest { + datasource_id: datasource_id.to_owned(), + database_type: POSTGRES_DATABASE_TYPE.to_owned(), + database_name: config.database.clone(), + schema_name: fixture.schema.clone(), + trigger_name: fixture.trigger.clone(), + }) + .await + .expect("native PostgreSQL trigger detail must load"); + assert!(trigger.event_manipulation.contains("INSERT")); + assert!(trigger.body.contains(&fixture.trigger)); +} + +async fn verify_query_console_preview_and_er( + application: &Application, + datasource_id: &str, + config: &PostgresTestConfig, + fixture: &PostgresFixture, +) { + let query = application + .start_query(StartQueryRequest { + datasource_id: datasource_id.to_owned(), + sql: format!( + "SELECT label, score FROM \"{}\".\"{}\" WHERE id = $1", + fixture.schema, fixture.table + ), + parameters: vec![QueryParameter { + position: 1, + value: JdbcValue::SignedInteger { + value: "1".to_owned(), + }, + }], + limits: query_limits("10"), + }) + .await + .expect("native PostgreSQL retained query must be accepted"); + let result = wait_for_result(application, &query.operation_id).await; + assert_eq!(result.row_count, "1"); + let page = result_page(application, &result).await; + assert!(matches!( + &page.rows[0].values[0], + JdbcValue::Text { value } if value == "alpha" + )); + + let console = application + .execute_native_console( + NativeConsoleRequest { + datasource_id: datasource_id.to_owned(), + database_name: config.database.clone(), + sql: format!( + "SELECT label FROM \"{}\".\"{}\" ORDER BY id", + fixture.schema, fixture.table + ), + page_no: 1, + page_size: 10, + result_set_id: None, + single: false, + page_size_all: false, + explain: false, + error_continue: false, + }, + NativeConsoleCancellation::new(), + ) + .await + .expect("native PostgreSQL Console must execute"); + assert_eq!(console.len(), 1); + assert!(console[0].success); + assert_eq!(console[0].row_count, 2); + + verify_extended_binary_types(application, datasource_id, config).await; + + let preview = application + .start_community_table_preview(StartCommunityTablePreviewRequest { + datasource_id: datasource_id.to_owned(), + database_type: POSTGRES_DATABASE_TYPE.to_owned(), + database_name: config.database.clone(), + schema_name: fixture.schema.clone(), + table_name: fixture.table.clone(), + row_limit: Some(2), + }) + .await + .expect("native PostgreSQL table preview must be accepted"); + let preview_result = wait_for_result(application, &preview.operation_id).await; + assert_eq!(preview_result.row_count, "2"); + assert_eq!( + result_page(application, &preview_result).await.rows.len(), + 2 + ); + + let er = application + .community_mysql_er_model(CommunityErQueryRequest { + data_source_id: datasource_id.to_owned(), + database_name: config.database.clone(), + schema_name: fixture.schema.clone(), + }) + .await + .expect("native PostgreSQL ER metadata must load through the table SPI"); + let child = er + .tables + .iter() + .find(|table| table.name == fixture.table) + .expect("ER metadata must include the child table"); + assert!(child.column_list.iter().any(|column| column.name == "id")); + assert!(child.foreign_key_list.iter().any(|key| { + key.pk_table_name == fixture.parent_table && key.fk_table_name == fixture.table + })); +} + +async fn verify_extended_binary_types( + application: &Application, + datasource_id: &str, + config: &PostgresTestConfig, +) { + let types = application + .execute_native_console( + NativeConsoleRequest { + datasource_id: datasource_id.to_owned(), + database_name: config.database.clone(), + sql: concat!( + "SELECT inet '192.168.4.7/24', ", + "cidr '2001:db8::/64', ", + "12.34::money, ", + "ARRAY[['a,b', 'NULL'], ['brace{', 'white space']]::text[][]" + ) + .to_owned(), + page_no: 1, + page_size: 10, + result_set_id: None, + single: true, + page_size_all: false, + explain: false, + error_continue: false, + }, + NativeConsoleCancellation::new(), + ) + .await + .expect("native PostgreSQL extended binary types must decode"); + let values = &types[0].rows[0].values; + assert!(matches!( + &values[0], + JdbcValue::Text { value } if value == "192.168.4.7/24" + )); + assert!(matches!( + &values[1], + JdbcValue::Text { value } if value == "2001:db8::/64" + )); + assert!(matches!( + &values[2], + JdbcValue::Opaque { type_name, display_value } + if type_name == "money" && display_value == "raw_units=1234" + )); + assert!(matches!( + &values[3], + JdbcValue::Opaque { display_value, .. } + if display_value == "{{\"a,b\",\"NULL\"},{\"brace{\",\"white space\"}}" + )); +} + +async fn verify_console_write_outcomes( + application: &Application, + datasource_id: &str, + config: &PostgresTestConfig, + fixture: &PostgresFixture, +) { + let request = |sql: String| NativeConsoleRequest { + datasource_id: datasource_id.to_owned(), + database_name: config.database.clone(), + sql, + page_no: 1, + page_size: 10, + result_set_id: None, + single: true, + page_size_all: false, + explain: false, + error_continue: false, + }; + + let rejected = application + .execute_native_console( + request(format!( + "INSERT INTO \"{}\".\"{}\" (id, parent_id, label) VALUES (1, 1, 'duplicate') RETURNING id", + fixture.schema, fixture.table + )), + NativeConsoleCancellation::new(), + ) + .await + .expect("an explicit PostgreSQL rejection must remain a statement failure"); + assert_eq!(rejected.len(), 1); + assert!(!rejected[0].success); + assert_eq!( + rejected[0] + .error + .as_ref() + .expect("the rejected write must include a safe error") + .code, + "postgres_query_rejected" + ); + + let cancellation = NativeConsoleCancellation::new(); + let cancellation_control = cancellation.clone(); + let cancellation_marker = format!("chat2db_cancel_{}", Uuid::new_v4().simple()); + let cancelled = tokio::time::timeout(EVENT_TIMEOUT, async { + tokio::join!( + application.execute_native_console( + request(long_running_insert( + fixture, + &cancellation_marker, + "cancelled" + )), + cancellation, + ), + async { + wait_for_active_statement(config, &cancellation_marker).await; + assert!( + cancellation_control + .cancel(Some("cancel after PostgreSQL dispatch".to_owned())) + ); + } + ) + }) + .await + .expect("the cancelled PostgreSQL Console write must terminate") + .0 + .expect_err("a write cancelled after dispatch must not report a definite outcome"); + assert_eq!(cancelled.api_error().code, "database_write_outcome_unknown"); + + let termination_marker = format!("chat2db_terminate_{}", Uuid::new_v4().simple()); + let terminated = tokio::time::timeout(EVENT_TIMEOUT, async { + tokio::join!( + application.execute_native_console( + request(long_running_insert( + fixture, + &termination_marker, + "terminated" + )), + NativeConsoleCancellation::new(), + ), + async { + let backend_pid = wait_for_active_statement(config, &termination_marker).await; + terminate_backend(config, backend_pid).await; + } + ) + }) + .await + .expect("the terminated PostgreSQL Console write must finish") + .0 + .expect_err("a transport failure after write dispatch must have an unknown outcome"); + assert_eq!( + terminated.api_error().code, + "database_write_outcome_unknown" + ); +} + +fn long_running_insert(fixture: &PostgresFixture, marker: &str, label: &str) -> String { + format!( + "/* {marker} */ INSERT INTO \"{}\".\"{}\" (parent_id, label, score) SELECT 1, '{label}', 1 FROM pg_sleep(30) RETURNING id", + fixture.schema, fixture.table + ) +} + +async fn wait_for_active_statement(config: &PostgresTestConfig, marker: &str) -> i32 { + let (client, connection) = config + .driver_config() + .connect(NoTls) + .await + .expect("PostgreSQL activity monitor connection"); + let task = tokio::spawn(connection); + let query_pattern = format!("%{marker}%"); + let backend_pid = tokio::time::timeout(Duration::from_secs(10), async { + loop { + if let Some(row) = client + .query_opt( + "SELECT pid FROM pg_stat_activity WHERE pid <> pg_backend_pid() AND state = 'active' AND wait_event_type = 'Timeout' AND wait_event = 'PgSleep' AND query LIKE $1 ORDER BY query_start DESC LIMIT 1", + &[&query_pattern], + ) + .await + .expect("PostgreSQL activity lookup") + { + return row.try_get(0).expect("PostgreSQL backend pid"); + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + }) + .await + .expect("the marked PostgreSQL statement must reach the server"); + drop(client); + task.await + .expect("PostgreSQL activity monitor task") + .expect("PostgreSQL activity monitor shutdown"); + backend_pid +} + +async fn terminate_backend(config: &PostgresTestConfig, backend_pid: i32) { + let (client, connection) = config + .driver_config() + .connect(NoTls) + .await + .expect("PostgreSQL termination connection"); + let task = tokio::spawn(connection); + let terminated: bool = client + .query_one("SELECT pg_terminate_backend($1)", &[&backend_pid]) + .await + .expect("PostgreSQL backend termination") + .try_get(0) + .expect("PostgreSQL backend termination result"); + assert!(terminated, "the marked PostgreSQL backend must terminate"); + drop(client); + task.await + .expect("PostgreSQL termination task") + .expect("PostgreSQL termination connection shutdown"); +} + +fn table_keys( + datasource_id: &str, + config: &PostgresTestConfig, + fixture: &PostgresFixture, + table_name: &str, +) -> ListCommunityTableKeysRequest { + ListCommunityTableKeysRequest { + datasource_id: datasource_id.to_owned(), + database_type: POSTGRES_DATABASE_TYPE.to_owned(), + database_name: config.database.clone(), + schema_name: fixture.schema.clone(), + table_name: table_name.to_owned(), + } +} + +fn query_limits(max_rows: &str) -> QueryLimits { + QueryLimits { + max_rows: max_rows.to_owned(), + max_result_bytes: (8_u64 * 1024 * 1024).to_string(), + batch_rows: 2, + batch_bytes: 1024 * 1024, + result_ttl_seconds: 60, + } +} + +async fn wait_for_result(application: &Application, operation_id: &str) -> ResultMetadata { + let mut subscription = application + .subscribe_operation(operation_id, None) + .await + .expect("native PostgreSQL operation must be subscribable"); + tokio::time::timeout(EVENT_TIMEOUT, async { + while let Some(envelope) = subscription + .next_event() + .await + .expect("native PostgreSQL operation event must decode") + { + match envelope.event { + OperationEvent::Completed { result } => return result, + OperationEvent::Failed { error } => { + panic!("native PostgreSQL query failed: {error:?}") + } + OperationEvent::Cancelled { reason } => { + panic!("native PostgreSQL query was cancelled: {reason:?}") + } + OperationEvent::Started | OperationEvent::Progress { .. } => {} + } + } + panic!("native PostgreSQL operation ended without a terminal event") + }) + .await + .expect("native PostgreSQL query must finish before timeout") +} + +async fn wait_for_failure( + application: &Application, + operation_id: &str, +) -> chat2db_contract::ApiError { + let mut subscription = application + .subscribe_operation(operation_id, None) + .await + .expect("failed native PostgreSQL operation must be subscribable"); + tokio::time::timeout(EVENT_TIMEOUT, async { + while let Some(envelope) = subscription + .next_event() + .await + .expect("native PostgreSQL operation event must decode") + { + match envelope.event { + OperationEvent::Failed { error } => return error, + OperationEvent::Completed { result } => { + panic!("native PostgreSQL query unexpectedly completed: {result:?}") + } + OperationEvent::Cancelled { reason } => { + panic!("native PostgreSQL query was cancelled: {reason:?}") + } + OperationEvent::Started | OperationEvent::Progress { .. } => {} + } + } + panic!("native PostgreSQL operation ended without a failure event") + }) + .await + .expect("native PostgreSQL query must fail before timeout") +} + +async fn result_page( + application: &Application, + result: &ResultMetadata, +) -> chat2db_contract::ResultPage { + application + .result_page( + &result.id, + ResultPageRequest { + offset: "0".to_owned(), + max_rows: "20".to_owned(), + max_bytes: (8_u64 * 1024 * 1024).to_string(), + }, + ) + .await + .expect("native PostgreSQL result page must be retained") +} + +fn assert_java_dormant(application: &Application) { + let engine = application + .health() + .components + .into_iter() + .find(|component| component.id == "database-engine") + .expect("database engine health must be present"); + assert_eq!(engine.state, ComponentState::Ready); + assert_eq!(engine.detail, "Available on demand; Java is not running"); +} + +async fn provision_fixture(config: &PostgresTestConfig, fixture: &PostgresFixture) { + let sql = format!( + r#" + CREATE SCHEMA "{schema}"; + CREATE TABLE "{schema}"."{parent}" ( + id BIGSERIAL PRIMARY KEY, + label TEXT NOT NULL + ); + CREATE TABLE "{schema}"."{table}" ( + id BIGSERIAL PRIMARY KEY, + parent_id BIGINT NOT NULL, + label VARCHAR(64) NOT NULL, + score NUMERIC(12, 4), + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + CONSTRAINT "{foreign_key}" FOREIGN KEY (parent_id) + REFERENCES "{schema}"."{parent}"(id) + ); + CREATE INDEX "{index}" ON "{schema}"."{table}"(label); + CREATE VIEW "{schema}"."{view}" AS + SELECT id, label FROM "{schema}"."{table}"; + CREATE FUNCTION "{schema}"."{function}"(p_value integer) RETURNS integer + LANGUAGE SQL IMMUTABLE AS $function_body$ SELECT p_value + 1 $function_body$; + CREATE PROCEDURE "{schema}"."{procedure}"(IN p_value integer, OUT p_output integer) + LANGUAGE plpgsql AS $procedure_body$ BEGIN p_output := p_value + 1; END $procedure_body$; + CREATE FUNCTION "{schema}"."{trigger_function}"() RETURNS trigger + LANGUAGE plpgsql AS $trigger_body$ + BEGIN NEW.created_at := clock_timestamp(); RETURN NEW; END + $trigger_body$; + CREATE TRIGGER "{trigger}" BEFORE INSERT ON "{schema}"."{table}" + FOR EACH ROW EXECUTE FUNCTION "{schema}"."{trigger_function}"(); + INSERT INTO "{schema}"."{parent}"(label) VALUES ('parent'); + INSERT INTO "{schema}"."{table}"(parent_id, label, score) + VALUES (1, 'alpha', 12.3400), (1, 'beta', 56.7800); + "#, + schema = fixture.schema, + parent = fixture.parent_table, + table = fixture.table, + foreign_key = fixture.foreign_key, + index = fixture.index, + view = fixture.view, + function = fixture.function, + procedure = fixture.procedure, + trigger_function = fixture.trigger_function, + trigger = fixture.trigger, + ); + execute_fixture_sql(config, &sql) + .await + .expect("native PostgreSQL fixture must be provisioned"); +} + +async fn cleanup_fixture( + config: &PostgresTestConfig, + fixture: &PostgresFixture, +) -> Result<(), Box> { + execute_fixture_sql( + config, + &format!("DROP SCHEMA IF EXISTS \"{}\" CASCADE", fixture.schema), + ) + .await +} + +async fn assert_fixture_residue_zero(config: &PostgresTestConfig) { + let (client, connection) = config + .driver_config() + .connect(NoTls) + .await + .expect("PostgreSQL residue verification connection"); + let task = tokio::spawn(connection); + let count: i64 = client + .query_one( + "SELECT count(*) FROM pg_namespace \ + WHERE nspname LIKE 'c2s\\_%' ESCAPE '\\' \ + OR nspname LIKE 'chat2db\\_native\\_smoke\\_%' ESCAPE '\\'", + &[], + ) + .await + .expect("PostgreSQL residue verification query") + .try_get(0) + .expect("PostgreSQL residue count"); + drop(client); + task.await + .expect("PostgreSQL residue connection task") + .expect("PostgreSQL residue connection shutdown"); + assert_eq!(count, 0, "native PostgreSQL smoke schemas must be removed"); +} + +async fn execute_fixture_sql( + config: &PostgresTestConfig, + sql: &str, +) -> Result<(), Box> { + let (client, connection) = config.driver_config().connect(NoTls).await?; + let task = tokio::spawn(connection); + let execution = client.batch_execute(sql).await; + drop(client); + task.await??; + execution?; + Ok(()) +} + +async fn query_single_text(config: &PostgresTestConfig, sql: &str) -> String { + let (client, connection) = config + .driver_config() + .connect(NoTls) + .await + .expect("PostgreSQL verification connection"); + let task = tokio::spawn(connection); + let value = client + .query_one(sql, &[]) + .await + .expect("PostgreSQL verification query") + .try_get(0) + .expect("PostgreSQL text value"); + drop(client); + task.await + .expect("PostgreSQL verification connection task") + .expect("PostgreSQL verification connection shutdown"); + value +} diff --git a/crates/chat2db-core/tests/native_sqlserver_product.rs b/crates/chat2db-core/tests/native_sqlserver_product.rs new file mode 100644 index 0000000..332a409 --- /dev/null +++ b/crates/chat2db-core/tests/native_sqlserver_product.rs @@ -0,0 +1,1387 @@ +use std::{panic::AssertUnwindSafe, time::Duration}; + +use base64::{Engine as _, engine::general_purpose::STANDARD}; +use chat2db_contract::{ + BuildCommunityCreateSchemaRequest, BuildCommunityDmlRequest, BuildCommunityNamespaceSqlRequest, + CommunityDmlAssignment, CommunityDmlColumn, CommunityDmlRow, CommunityDmlStatement, + CommunityDmlTarget, CommunityDmlTemporalKind, CommunityDmlValue, + CommunityNamespaceSqlOperation, CommunitySchema, ComponentState, CreateDatasourceRequest, + DatasourceConnection, DatasourceConnectionProperty, GetCommunityFunctionRequest, + GetCommunityProcedureRequest, GetCommunityTriggerRequest, JdbcValue, + ListCommunityColumnsRequest, ListCommunityDatabasesRequest, ListCommunityFunctionsRequest, + ListCommunityIndexesRequest, ListCommunityProceduresRequest, ListCommunitySchemasRequest, + ListCommunityTableKeysRequest, ListCommunityTablesRequest, ListCommunityTriggersRequest, + ListCommunityViewsRequest, OperationEvent, QueryLimits, QueryParameter, ResultMetadata, + ResultPageRequest, StartCommunityTablePreviewRequest, StartQueryRequest, +}; +use chat2db_core::{ + Application, NativeConsoleCancellation, NativeConsoleRequest, RuntimeConfig, RuntimeHost, +}; +use chat2db_java_bridge::{EngineCommand, EngineConfig}; +use futures_util::{FutureExt as _, TryStreamExt as _}; +use tempfile::TempDir; +use tiberius::{AuthMethod, Client, Config, EncryptionLevel}; +use tokio::net::TcpStream; +use tokio_util::compat::{Compat, TokioAsyncWriteCompatExt as _}; +use uuid::Uuid; + +const DATABASE_TYPE: &str = "SQLSERVER"; +const EVENT_TIMEOUT: Duration = Duration::from_secs(20); + +type DirectClient = Client>; + +struct SqlServerTestConfig { + host: String, + port: u16, + user: String, + password: String, +} + +impl SqlServerTestConfig { + fn from_environment() -> Self { + let host = required_env("SQLSERVER_TEST_HOST"); + assert!( + !host.trim().is_empty(), + "SQLSERVER_TEST_HOST cannot be empty" + ); + let port = required_env("SQLSERVER_TEST_PORT") + .parse::() + .expect("SQLSERVER_TEST_PORT must be a TCP port"); + assert_ne!(port, 0, "SQLSERVER_TEST_PORT cannot be zero"); + Self { + host, + port, + user: required_env("SQLSERVER_TEST_USER"), + password: required_env("SQLSERVER_TEST_PASSWORD"), + } + } + + fn connection(&self) -> DatasourceConnection { + let host = if self.host.contains(':') + && !(self.host.starts_with('[') && self.host.ends_with(']')) + { + format!("[{}]", self.host) + } else { + self.host.clone() + }; + DatasourceConnection { + jdbc_url: format!( + "jdbc:sqlserver://{host}:{};databaseName=master;encrypt=false;trustServerCertificate=true", + self.port + ), + properties: vec![ + DatasourceConnectionProperty { + key: "user".to_owned(), + value: self.user.clone(), + sensitive: false, + }, + DatasourceConnectionProperty { + key: "password".to_owned(), + value: self.password.clone(), + sensitive: true, + }, + ], + read_only: false, + ssh: None, + } + } + + async fn direct_client(&self) -> Result { + let mut config = Config::new(); + config.host(&self.host); + config.port(self.port); + config.database("master"); + config.authentication(AuthMethod::sql_server(&self.user, &self.password)); + config.encryption(EncryptionLevel::NotSupported); + config.trust_cert(); + let address = config.get_addr(); + let tcp = TcpStream::connect(address) + .await + .map_err(|error| error.to_string())?; + tcp.set_nodelay(true).map_err(|error| error.to_string())?; + Client::connect(config, tcp.compat_write()) + .await + .map_err(|error| error.to_string()) + } +} + +#[tokio::test] +#[ignore = "requires SQLSERVER_TEST_HOST, SQLSERVER_TEST_PORT, SQLSERVER_TEST_USER, and SQLSERVER_TEST_PASSWORD"] +async fn native_sqlserver_product_paths_keep_java_dormant() { + let config = SqlServerTestConfig::from_environment(); + let database_name = format!("chat2db_native_it_{}", Uuid::new_v4().simple()); + provision_database(&config, &database_name).await; + + let verification = AssertUnwindSafe(verify_native_product(&config, &database_name)) + .catch_unwind() + .await; + let cleanup = cleanup_database(&config, &database_name).await; + if let Err(payload) = verification { + if let Err(error) = cleanup { + eprintln!("native SQL Server cleanup also failed: {error}"); + } + std::panic::resume_unwind(payload); + } + cleanup.expect("native SQL Server fixture database must be removed"); +} + +async fn verify_native_product(config: &SqlServerTestConfig, database_name: &str) { + let directory = TempDir::new().expect("temporary native SQL Server runtime"); + let missing_java = directory.path().join("missing-java"); + let runtime = RuntimeConfig::new(EngineConfig::new(EngineCommand::new(missing_java))) + .with_data_dir(directory.path().join("data")) + .with_vault_master_key_base64(STANDARD.encode([0x73; 32])); + let mut host = RuntimeHost::open(runtime) + .await + .expect("native SQL Server runtime must open without Java"); + let application = host.application(); + assert_java_dormant(&application); + + let driver = application + .list_drivers() + .items + .into_iter() + .find(|driver| driver.driver_id == "sqlserver") + .expect("native SQL Server driver must be present"); + assert_eq!(driver.driver_class, "rust:tiberius"); + assert_eq!(driver.artifact_count, 0); + application + .test_datasource_connection("sqlserver", config.connection()) + .await + .expect("native SQL Server connection test must succeed"); + let mut tls_conflict = config.connection(); + tls_conflict.properties.push(DatasourceConnectionProperty { + key: "trustServerCertificateCA".to_owned(), + value: "/tmp/sqlserver-ca.pem".to_owned(), + sensitive: false, + }); + let conflict = application + .test_datasource_connection("sqlserver", tls_conflict) + .await + .expect_err("conflicting SQL Server trust settings must fail without panicking"); + assert_eq!(conflict.api_error().code, "invalid_sqlserver_connection"); + assert_java_dormant(&application); + + let datasource = application + .create_datasource(CreateDatasourceRequest { + name: "Native SQL Server".to_owned(), + driver_id: "sqlserver".to_owned(), + connection: Some(config.connection()), + }) + .await + .expect("native SQL Server datasource must persist"); + + verify_native_dialect_builders(&application, &datasource.id, database_name).await; + verify_database_and_table_metadata(&application, &datasource.id, database_name).await; + verify_object_metadata(&application, &datasource.id, database_name).await; + verify_console_query_and_preview(&application, &datasource.id, database_name).await; + verify_console_preflight_compatibility(&application, &datasource.id, database_name).await; + verify_fail_closed_query_safety(&application, &datasource.id, database_name).await; + verify_cancellation_and_recovery(&application, &datasource.id, database_name).await; + assert_java_dormant(&application); + host.shutdown() + .await + .expect("native-only SQL Server runtime must shut down cleanly"); +} + +async fn verify_native_dialect_builders( + application: &Application, + datasource_id: &str, + database_name: &str, +) { + let schema = application + .build_community_create_schema(BuildCommunityCreateSchemaRequest { + database_type: DATABASE_TYPE.to_owned(), + schema: CommunitySchema { + database_name: database_name.to_owned(), + name: "native_builder".to_owned(), + comment: "native builder's schema".to_owned(), + owner: "dbo".to_owned(), + system: false, + }, + }) + .await + .expect("native SQL Server CREATE SCHEMA must build without Java"); + assert!(schema.sql.contains("CREATE SCHEMA [native_builder]")); + execute_product_console(application, datasource_id, database_name, &schema.sql).await; + + let use_database = application + .build_community_namespace_sql(BuildCommunityNamespaceSqlRequest { + database_type: DATABASE_TYPE.to_owned(), + operation: CommunityNamespaceSqlOperation::UseDatabase { + database_name: database_name.to_owned(), + }, + }) + .await + .expect("native SQL Server USE must build without Java"); + assert_eq!(use_database.sql, format!("USE [{database_name}];")); + let rename = application + .build_community_namespace_sql(BuildCommunityNamespaceSqlRequest { + database_type: DATABASE_TYPE.to_owned(), + operation: CommunityNamespaceSqlOperation::AlterSchema { + old_schema_name: "native_builder".to_owned(), + new_schema_name: "native_builder_renamed".to_owned(), + }, + }) + .await + .expect_err("SQL Server schema rename must fail explicitly without Java fallback"); + assert_eq!( + rename.api_error().code, + "sqlserver_schema_rename_unsupported" + ); + + execute_product_console( + application, + datasource_id, + database_name, + "CREATE TABLE native_builder.items (id int NOT NULL PRIMARY KEY, label nvarchar(80) NOT NULL, amount decimal(12,2) NOT NULL, active bit NOT NULL, created_at datetimeoffset NOT NULL, payload varbinary(16) NOT NULL)", + ) + .await; + verify_native_dml_builders(application, datasource_id, database_name).await; + assert_java_dormant(application); +} + +async fn verify_native_dml_builders( + application: &Application, + datasource_id: &str, + database_name: &str, +) { + let columns = vec![ + dml_column("id", "int"), + dml_column("label", "nvarchar"), + dml_column("amount", "decimal"), + dml_column("active", "bit"), + dml_column("created_at", "datetimeoffset"), + dml_column("payload", "varbinary"), + ]; + let insert = application + .build_community_dml(BuildCommunityDmlRequest { + database_type: DATABASE_TYPE.to_owned(), + target: dml_target(database_name), + statement: CommunityDmlStatement::SingleInsert { + columns, + row: CommunityDmlRow { + values: vec![ + CommunityDmlValue::Decimal { + value: "1".to_owned(), + }, + CommunityDmlValue::String { + value: "O'Brien".to_owned(), + }, + CommunityDmlValue::Decimal { + value: "12.50".to_owned(), + }, + CommunityDmlValue::Boolean { value: true }, + CommunityDmlValue::Temporal { + temporal_kind: CommunityDmlTemporalKind::OffsetDatetime, + value: "2026-08-07T12:30:45.1234567+08:00".to_owned(), + }, + CommunityDmlValue::Binary { + base64: STANDARD.encode([0, 255]), + }, + ], + }, + }, + }) + .await + .expect("native SQL Server INSERT must build without Java"); + assert!(insert.sql.contains("N'O''Brien'")); + assert!(insert.sql.contains("0x00ff")); + execute_product_console(application, datasource_id, database_name, &insert.sql).await; + + let update = application + .build_community_dml(BuildCommunityDmlRequest { + database_type: DATABASE_TYPE.to_owned(), + target: dml_target(database_name), + statement: CommunityDmlStatement::Update { + assignments: vec![CommunityDmlAssignment { + column: dml_column("label", "nvarchar"), + value: CommunityDmlValue::String { + value: "updated".to_owned(), + }, + }], + predicates: vec![CommunityDmlAssignment { + column: dml_column("id", "int"), + value: CommunityDmlValue::Decimal { + value: "1".to_owned(), + }, + }], + }, + }) + .await + .expect("native SQL Server UPDATE must build without Java"); + execute_product_console(application, datasource_id, database_name, &update.sql).await; + let selected = application + .execute_native_console( + NativeConsoleRequest { + datasource_id: datasource_id.to_owned(), + database_name: database_name.to_owned(), + sql: "SELECT label FROM native_builder.items WHERE id = 1".to_owned(), + page_no: 1, + page_size: 20, + result_set_id: None, + single: true, + page_size_all: false, + explain: false, + error_continue: false, + }, + NativeConsoleCancellation::new(), + ) + .await + .expect("native SQL Server builder result must be queryable"); + assert!(matches!( + selected + .iter() + .find_map(|result| result.rows.first()) + .and_then(|row| row.values.first()), + Some(JdbcValue::Text { value }) if value == "updated" + )); +} + +async fn execute_product_console( + application: &Application, + datasource_id: &str, + database_name: &str, + sql: &str, +) { + let results = application + .execute_native_console( + NativeConsoleRequest { + datasource_id: datasource_id.to_owned(), + database_name: database_name.to_owned(), + sql: sql.to_owned(), + page_no: 1, + page_size: 20, + result_set_id: None, + single: false, + page_size_all: false, + explain: false, + error_continue: false, + }, + NativeConsoleCancellation::new(), + ) + .await + .unwrap_or_else(|error| panic!("native SQL Server Console failed for {sql}: {error}")); + assert!( + results.iter().all(|result| result.success), + "native SQL Server Console returned failed results for {sql}: {results:#?}" + ); +} + +fn dml_target(database_name: &str) -> CommunityDmlTarget { + CommunityDmlTarget { + database_name: Some(database_name.to_owned()), + schema_name: Some("native_builder".to_owned()), + table_name: "items".to_owned(), + } +} + +fn dml_column(name: &str, data_type_name: &str) -> CommunityDmlColumn { + CommunityDmlColumn { + name: name.to_owned(), + data_type_name: data_type_name.to_owned(), + precision: None, + scale: None, + } +} + +#[allow( + clippy::too_many_lines, + reason = "the product smoke checks every SQL Server table metadata projection together" +)] +async fn verify_database_and_table_metadata( + application: &Application, + datasource_id: &str, + database_name: &str, +) { + let databases = application + .list_community_databases(ListCommunityDatabasesRequest { + datasource_id: datasource_id.to_owned(), + database_type: DATABASE_TYPE.to_owned(), + }) + .await + .expect("native SQL Server databases must list"); + assert!( + databases + .items + .iter() + .any(|item| item.name == database_name) + ); + assert!( + databases + .items + .iter() + .find(|item| item.name == "master") + .is_some_and(|item| item.system) + ); + + let schemas = application + .list_community_schemas(ListCommunitySchemasRequest { + datasource_id: datasource_id.to_owned(), + database_type: DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + }) + .await + .expect("native SQL Server schemas must list"); + assert!(schemas.items.iter().any(|item| item.name == "dbo")); + + let tables = application + .list_community_tables(ListCommunityTablesRequest { + datasource_id: datasource_id.to_owned(), + database_type: DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: "dbo".to_owned(), + table_name_pattern: String::new(), + }) + .await + .expect("native SQL Server tables must list"); + assert!(tables.items.iter().any(|item| item.name == "native_parent")); + assert!(tables.items.iter().any(|item| item.name == "native_child")); + assert!(tables.items.iter().all(|item| item.name != "native_view")); + + let columns = application + .list_community_columns(ListCommunityColumnsRequest { + datasource_id: datasource_id.to_owned(), + database_type: DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: "dbo".to_owned(), + table_name: "native_child".to_owned(), + }) + .await + .expect("native SQL Server columns must list"); + let identity = columns + .items + .iter() + .find(|column| column.name == "id") + .expect("identity column must exist"); + assert_eq!(identity.auto_increment, Some(true)); + assert_eq!(identity.seed, Some(10)); + assert_eq!(identity.increment, Some(5)); + assert!( + columns + .items + .iter() + .find(|column| column.name == "note_lower") + .is_some_and(|column| column.generated_column == Some(true)) + ); + + let indexes = application + .list_community_indexes(ListCommunityIndexesRequest { + datasource_id: datasource_id.to_owned(), + database_type: DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: "dbo".to_owned(), + table_name: "native_child".to_owned(), + }) + .await + .expect("native SQL Server indexes must list"); + assert!( + indexes + .items + .iter() + .any(|index| index.name == "PK_native_child") + ); + let filtered = indexes + .items + .iter() + .find(|index| index.name == "IX_native_child_note") + .expect("filtered index must exist"); + assert!( + filtered + .columns + .iter() + .any(|column| column.column_name == "note") + ); + assert!( + filtered + .columns + .iter() + .any(|column| column.column_name == "parent_id" && column.column_type == "INCLUDED") + ); + + let child_keys = table_keys_request(datasource_id, database_name, "native_child"); + let imported = application + .list_community_imported_keys(child_keys.clone()) + .await + .expect("native SQL Server imported keys must list"); + assert!(imported.items.iter().any(|key| { + key.foreign_key_name == "FK_native_child_parent" + && key.primary_table_name == "native_parent" + && key.foreign_table_name == "native_child" + })); + let primary = application + .list_community_primary_keys(child_keys) + .await + .expect("native SQL Server primary keys must list"); + assert!( + primary + .items + .iter() + .any(|key| key.name == "PK_native_child" && key.column_name == "id") + ); + let exported = application + .list_community_exported_keys(table_keys_request( + datasource_id, + database_name, + "native_parent", + )) + .await + .expect("native SQL Server exported keys must list"); + assert!( + exported + .items + .iter() + .any(|key| key.foreign_key_name == "FK_native_child_parent") + ); + + let ddl = application + .table_ddl(datasource_id, database_name, "dbo", "native_child") + .await + .expect("native SQL Server table DDL must render"); + assert!(ddl.starts_with(&format!( + "CREATE TABLE [{database_name}].[dbo].[native_child]" + ))); + assert!(ddl.contains("IDENTITY(10,5)")); + assert!(ddl.contains("CONSTRAINT [FK_native_child_parent] FOREIGN KEY")); + assert!(ddl.contains("CREATE NONCLUSTERED INDEX [IX_native_child_note]")); + let invalid = application + .table_ddl(datasource_id, database_name, "dbo", "") + .await + .expect_err("empty SQL Server table names must be rejected"); + assert_eq!( + invalid.api_error().code, + "invalid_sqlserver_metadata_request" + ); + assert_java_dormant(application); +} + +#[allow( + clippy::too_many_lines, + reason = "the product smoke checks every SQL Server programmable-object projection together" +)] +async fn verify_object_metadata( + application: &Application, + datasource_id: &str, + database_name: &str, +) { + let views_request = ListCommunityViewsRequest { + datasource_id: datasource_id.to_owned(), + database_type: DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: "dbo".to_owned(), + view_name_pattern: "native_view".to_owned(), + }; + let views = application + .list_community_views(views_request.clone()) + .await + .expect("native SQL Server views must list"); + assert!(views.items.iter().any(|view| view.name == "native_view")); + let view = application + .get_community_view(views_request) + .await + .expect("native SQL Server view detail must load"); + assert!(view.ddl.contains("native_view")); + + let functions = application + .list_community_functions(ListCommunityFunctionsRequest { + datasource_id: datasource_id.to_owned(), + database_type: DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: "dbo".to_owned(), + }) + .await + .expect("native SQL Server functions must list"); + assert!( + functions + .items + .iter() + .any(|item| item.name == "native_function") + ); + let function_request = GetCommunityFunctionRequest { + datasource_id: datasource_id.to_owned(), + database_type: DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: "dbo".to_owned(), + function_name: "native_function".to_owned(), + }; + let function = application + .get_community_function(function_request.clone()) + .await + .expect("native SQL Server function detail must load"); + assert!(function.body.contains("native_function")); + let function_parameters = application + .list_community_function_parameters(function_request) + .await + .expect("native SQL Server function parameters must list"); + assert!( + function_parameters + .items + .iter() + .any(|parameter| parameter.column_name == "@value") + ); + + let procedures = application + .list_community_procedures(ListCommunityProceduresRequest { + datasource_id: datasource_id.to_owned(), + database_type: DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: "dbo".to_owned(), + }) + .await + .expect("native SQL Server procedures must list"); + assert!( + procedures + .items + .iter() + .any(|item| item.name == "native_procedure") + ); + let procedure_request = GetCommunityProcedureRequest { + datasource_id: datasource_id.to_owned(), + database_type: DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: "dbo".to_owned(), + procedure_name: "native_procedure".to_owned(), + }; + let procedure = application + .get_community_procedure(procedure_request.clone()) + .await + .expect("native SQL Server procedure detail must load"); + assert!(procedure.body.contains("native_procedure")); + let procedure_parameters = application + .list_community_procedure_parameters(procedure_request) + .await + .expect("native SQL Server procedure parameters must list"); + assert!( + procedure_parameters + .items + .iter() + .any(|parameter| parameter.column_name == "@value") + ); + + let triggers = application + .list_community_triggers(ListCommunityTriggersRequest { + datasource_id: datasource_id.to_owned(), + database_type: DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: "dbo".to_owned(), + }) + .await + .expect("native SQL Server triggers must list"); + assert!( + triggers + .items + .iter() + .any(|item| item.name == "native_trigger") + ); + let trigger = application + .get_community_trigger(GetCommunityTriggerRequest { + datasource_id: datasource_id.to_owned(), + database_type: DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: "dbo".to_owned(), + trigger_name: "native_trigger".to_owned(), + }) + .await + .expect("native SQL Server trigger detail must load"); + assert!(trigger.event_manipulation.contains("INSERT")); + assert!(trigger.body.contains("native_trigger")); + assert_java_dormant(application); +} + +async fn verify_console_query_and_preview( + application: &Application, + datasource_id: &str, + database_name: &str, +) { + verify_datetimeoffset_projection(application, datasource_id, database_name).await; + + let console = application + .execute_native_console( + NativeConsoleRequest { + datasource_id: datasource_id.to_owned(), + database_name: database_name.to_owned(), + sql: "SELECT parent_id, note FROM dbo.native_child ORDER BY id".to_owned(), + page_no: 1, + page_size: 20, + result_set_id: None, + single: true, + page_size_all: false, + explain: false, + error_continue: false, + }, + NativeConsoleCancellation::new(), + ) + .await + .expect("native SQL Server Console SELECT must execute"); + let tabular = console + .iter() + .find(|result| !result.columns.is_empty()) + .expect("Console must return one tabular result"); + assert!(tabular.success); + assert_eq!(tabular.rows.len(), 2); + assert!(matches!( + tabular.rows[0].values.as_slice(), + [JdbcValue::SignedInteger { value: parent_id }, JdbcValue::Text { value: note }] + if parent_id == "7" && note == "alpha" + )); + + let query = application + .start_query(StartQueryRequest { + datasource_id: datasource_id.to_owned(), + sql: format!( + "SELECT parent_id, note FROM [{database_name}].[dbo].[native_child] WHERE parent_id = ? ORDER BY id" + ), + parameters: vec![QueryParameter { + position: 1, + value: JdbcValue::SignedInteger { + value: "7".to_owned(), + }, + }], + limits: query_limits("10"), + }) + .await + .expect("native SQL Server retained query must be accepted"); + let query_result = wait_for_result(application, &query.operation_id).await; + assert_eq!(query_result.row_count, "2"); + assert_eq!(result_page(application, &query_result).await.rows.len(), 2); + + let preview = application + .start_community_table_preview(StartCommunityTablePreviewRequest { + datasource_id: datasource_id.to_owned(), + database_type: DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: "dbo".to_owned(), + table_name: "native_child".to_owned(), + row_limit: Some(1), + }) + .await + .expect("native SQL Server table preview must be accepted"); + assert_eq!( + preview.sql, + format!("SELECT TOP (1) * FROM [{database_name}].[dbo].[native_child]") + ); + let preview_result = wait_for_result(application, &preview.operation_id).await; + assert_eq!(preview_result.row_count, "1"); + assert_eq!( + result_page(application, &preview_result).await.rows.len(), + 1 + ); + assert_java_dormant(application); +} + +async fn verify_datetimeoffset_projection( + application: &Application, + datasource_id: &str, + database_name: &str, +) { + let datetime_offset = application + .execute_native_console( + NativeConsoleRequest { + datasource_id: datasource_id.to_owned(), + database_name: database_name.to_owned(), + sql: "SELECT CAST('2026-08-07T12:34:56.123456+08:00' AS datetimeoffset(6)) AS exact_offset".to_owned(), + page_no: 1, + page_size: 20, + result_set_id: None, + single: true, + page_size_all: false, + explain: false, + error_continue: false, + }, + NativeConsoleCancellation::new(), + ) + .await + .expect("native SQL Server datetimeoffset SELECT must execute"); + let value = datetime_offset + .iter() + .find_map(|result| result.rows.first()) + .and_then(|row| row.values.first()) + .expect("datetimeoffset SELECT must return one value"); + assert_eq!( + value, + &JdbcValue::TimestampWithTimeZone { + value: "2026-08-07T12:34:56.123456+08:00".to_owned(), + } + ); +} + +async fn verify_console_preflight_compatibility( + application: &Application, + datasource_id: &str, + database_name: &str, +) { + verify_temporary_table_preflight(application, datasource_id, database_name).await; + verify_select_into_preflight(application, datasource_id, database_name).await; + verify_limited_reader_preflight(application, datasource_id, database_name).await; + assert_java_dormant(application); +} + +async fn verify_temporary_table_preflight( + application: &Application, + datasource_id: &str, + database_name: &str, +) { + let temporary_table = application + .execute_native_console( + NativeConsoleRequest { + datasource_id: datasource_id.to_owned(), + database_name: database_name.to_owned(), + sql: "CREATE TABLE #chat2db_native_temp(id int NOT NULL);\n\ + INSERT INTO #chat2db_native_temp(id) VALUES (1), (2);\n\ + WITH target AS (SELECT id FROM #chat2db_native_temp WHERE id = 2)\n\ + DELETE FROM target OUTPUT deleted.id;\n\ + SELECT id FROM #chat2db_native_temp ORDER BY id;\n\ + DROP TABLE #chat2db_native_temp" + .to_owned(), + page_no: 1, + page_size: 20, + result_set_id: None, + single: false, + page_size_all: false, + explain: false, + error_continue: false, + }, + NativeConsoleCancellation::new(), + ) + .await + .expect("temporary-table and CTE DML Console statements must execute"); + assert!( + temporary_table.iter().all(|result| result.success), + "temporary-table Console statements failed: {temporary_table:#?}" + ); + let returned_values = temporary_table + .iter() + .filter_map(|result| result.rows.first()) + .filter_map(|row| row.values.first()) + .filter_map(|value| match value { + JdbcValue::SignedInteger { value } => Some(value.as_str()), + _ => None, + }) + .collect::>(); + assert_eq!(returned_values, ["2", "1"]); +} + +async fn verify_select_into_preflight( + application: &Application, + datasource_id: &str, + database_name: &str, +) { + let select_into = application + .execute_native_console( + NativeConsoleRequest { + datasource_id: datasource_id.to_owned(), + database_name: database_name.to_owned(), + sql: "SELECT id INTO dbo.native_select_into FROM dbo.native_child;\n\ + SELECT COUNT_BIG(*) AS copied FROM dbo.native_select_into;\n\ + DROP TABLE dbo.native_select_into" + .to_owned(), + page_no: 1, + page_size: 20, + result_set_id: None, + single: false, + page_size_all: false, + explain: false, + error_continue: false, + }, + NativeConsoleCancellation::new(), + ) + .await + .expect("SELECT INTO must execute as a Console write"); + assert!( + select_into.iter().all(|result| result.success), + "SELECT INTO Console statements failed: {select_into:#?}" + ); + assert!(select_into.iter().any(|result| { + matches!( + result.rows.first().map(|row| row.values.as_slice()), + Some([JdbcValue::SignedInteger { value }]) if value == "2" + ) + })); +} + +async fn verify_limited_reader_preflight( + application: &Application, + datasource_id: &str, + database_name: &str, +) { + let limited_reader = application + .execute_native_console( + NativeConsoleRequest { + datasource_id: datasource_id.to_owned(), + database_name: database_name.to_owned(), + sql: "CREATE USER chat2db_native_reader WITHOUT LOGIN;\n\ + GRANT SELECT ON dbo.native_child TO chat2db_native_reader;\n\ + EXECUTE AS USER = 'chat2db_native_reader';\n\ + SELECT COUNT_BIG(*) AS visible_rows FROM dbo.native_child;\n\ + REVERT;\n\ + DROP USER chat2db_native_reader" + .to_owned(), + page_no: 1, + page_size: 20, + result_set_id: None, + single: false, + page_size_all: false, + explain: false, + error_continue: false, + }, + NativeConsoleCancellation::new(), + ) + .await + .expect("a SELECT-only database user must pass safe result preflight"); + assert!( + limited_reader.iter().all(|result| result.success), + "limited-reader Console statements failed: {limited_reader:#?}" + ); + assert!(limited_reader.iter().any(|result| { + matches!( + result.rows.first().map(|row| row.values.as_slice()), + Some([JdbcValue::SignedInteger { value }]) if value == "2" + ) + })); +} + +async fn verify_cancellation_and_recovery( + application: &Application, + datasource_id: &str, + database_name: &str, +) { + let cancellation = NativeConsoleCancellation::new(); + let task_application = application.clone(); + let task_cancellation = cancellation.clone(); + let request = NativeConsoleRequest { + datasource_id: datasource_id.to_owned(), + database_name: database_name.to_owned(), + sql: "SELECT SUM(CONVERT(bigint, a.object_id % 2)) AS total FROM sys.all_objects AS a CROSS JOIN sys.all_objects AS b CROSS JOIN sys.all_objects AS c CROSS JOIN sys.all_objects AS d".to_owned(), + page_no: 1, + page_size: 20, + result_set_id: None, + single: true, + page_size_all: false, + explain: false, + error_continue: false, + }; + let task = tokio::spawn(async move { + task_application + .execute_native_console(request, task_cancellation) + .await + }); + tokio::time::sleep(Duration::from_millis(300)).await; + assert!(cancellation.cancel(Some("SQL Server product smoke".to_owned()))); + let error = tokio::time::timeout(Duration::from_secs(10), task) + .await + .expect("cancelled SQL Server Console must stop before timeout") + .expect("cancelled SQL Server Console task must join") + .expect_err("cancelled SQL Server Console must return an error"); + assert_eq!(error.api_error().code, "sqlserver_console_cancelled"); + + let write_cancellation = NativeConsoleCancellation::new(); + let write_application = application.clone(); + let task_cancellation = write_cancellation.clone(); + let write_request = NativeConsoleRequest { + datasource_id: datasource_id.to_owned(), + database_name: database_name.to_owned(), + sql: "WAITFOR DELAY '00:00:03'; INSERT INTO dbo.native_parent(id,label) VALUES (99,N'write-cancelled')".to_owned(), + page_no: 1, + page_size: 20, + result_set_id: None, + single: true, + page_size_all: false, + explain: false, + error_continue: false, + }; + let write_task = tokio::spawn(async move { + write_application + .execute_native_console(write_request, task_cancellation) + .await + }); + tokio::time::sleep(Duration::from_millis(300)).await; + assert!(write_cancellation.cancel(Some("SQL Server write cancellation".to_owned()))); + let write_error = tokio::time::timeout(Duration::from_secs(10), write_task) + .await + .expect("cancelled SQL Server Console write must stop before timeout") + .expect("cancelled SQL Server Console write task must join") + .expect_err("a dispatched Console write cancellation must be outcome-unknown"); + assert_eq!( + write_error.api_error().code, + "database_write_outcome_unknown" + ); + verify_authoritative_server_rejections(application, datasource_id, database_name).await; + + let recovered = application + .execute_native_console( + NativeConsoleRequest { + datasource_id: datasource_id.to_owned(), + database_name: database_name.to_owned(), + sql: "SELECT 1 AS recovered".to_owned(), + page_no: 1, + page_size: 20, + result_set_id: None, + single: true, + page_size_all: false, + explain: false, + error_continue: false, + }, + NativeConsoleCancellation::new(), + ) + .await + .expect("SQL Server Console must recover after session cancellation"); + assert!( + recovered + .iter() + .any(|result| result.success && result.row_count == 1) + ); + assert_java_dormant(application); +} + +async fn verify_authoritative_server_rejections( + application: &Application, + datasource_id: &str, + database_name: &str, +) { + for sql in [ + "INSERT INTO dbo.native_parent(id,label) VALUES (7,N'duplicate')", + "WITH duplicate AS (SELECT 7 AS id, N'duplicate' AS label) \ + INSERT INTO dbo.native_parent(id,label) OUTPUT inserted.id \ + SELECT id,label FROM duplicate", + ] { + let results = application + .execute_native_console( + NativeConsoleRequest { + datasource_id: datasource_id.to_owned(), + database_name: database_name.to_owned(), + sql: sql.to_owned(), + page_no: 1, + page_size: 20, + result_set_id: None, + single: true, + page_size_all: false, + explain: false, + error_continue: false, + }, + NativeConsoleCancellation::new(), + ) + .await + .expect("an authoritative SQL Server rejection must remain a statement result"); + let error = results + .first() + .and_then(|result| result.error.as_ref()) + .expect("the rejected SQL Server write must expose its server error"); + assert_eq!(error.code, "sqlserver_query_failed"); + } +} + +async fn verify_fail_closed_query_safety( + application: &Application, + datasource_id: &str, + database_name: &str, +) { + verify_retained_query_fail_closed(application, datasource_id, database_name).await; + verify_console_results_fail_closed(application, datasource_id, database_name).await; + + let rows = application + .execute_native_console( + NativeConsoleRequest { + datasource_id: datasource_id.to_owned(), + database_name: database_name.to_owned(), + sql: "SELECT COUNT_BIG(*) AS row_count FROM dbo.native_child".to_owned(), + page_no: 1, + page_size: 20, + result_set_id: None, + single: true, + page_size_all: false, + explain: false, + error_continue: false, + }, + NativeConsoleCancellation::new(), + ) + .await + .expect("SQL Server fixture row count must remain readable"); + assert!(matches!( + rows.first() + .and_then(|result| result.rows.first()) + .map(|row| row.values.as_slice()), + Some([JdbcValue::SignedInteger { value }]) if value == "2" + )); + assert_java_dormant(application); +} + +async fn verify_retained_query_fail_closed( + application: &Application, + datasource_id: &str, + database_name: &str, +) { + let cte_write = application + .start_query(StartQueryRequest { + datasource_id: datasource_id.to_owned(), + sql: format!( + "WITH target AS (SELECT id FROM [{database_name}].[dbo].[native_child]) DELETE FROM target WHERE id = 10" + ), + parameters: Vec::new(), + limits: query_limits("10"), + }) + .await + .expect_err("a data-changing CTE must be rejected by the native read path"); + assert_eq!( + cte_write.api_error().code, + "sqlserver_native_query_unsupported" + ); + + let unsafe_query = application + .start_query(StartQueryRequest { + datasource_id: datasource_id.to_owned(), + sql: "SELECT CONVERT(sql_variant, 7) AS unsafe_variant".to_owned(), + parameters: Vec::new(), + limits: query_limits("10"), + }) + .await + .expect("unsafe SQL Server result query must be accepted as an operation"); + assert_eq!( + wait_for_failure_code(application, &unsafe_query.operation_id).await, + "sqlserver_result_type_unsupported" + ); +} + +async fn verify_console_results_fail_closed( + application: &Application, + datasource_id: &str, + database_name: &str, +) { + for (sql, expected_codes) in [ + ( + "SELECT CONVERT(money, 1.23) AS unsafe_money", + &["sqlserver_result_type_unsupported"][..], + ), + ( + "SELECT CONVERT(sql_variant, 7) AS unsafe_variant", + &["sqlserver_result_type_unsupported"][..], + ), + ( + "SELECT hierarchyid::Parse('/1/') AS unsafe_udt", + &[ + "sqlserver_result_type_unsupported", + "sqlserver_result_description_failed", + ][..], + ), + ( + "SELECT REPLICATE(CONVERT(varchar(max), 'x'), 4194305) AS oversized_value", + &["sqlserver_scalar_too_large"][..], + ), + ] { + let results = application + .execute_native_console( + NativeConsoleRequest { + datasource_id: datasource_id.to_owned(), + database_name: database_name.to_owned(), + sql: sql.to_owned(), + page_no: 1, + page_size: 20, + result_set_id: None, + single: true, + page_size_all: false, + explain: false, + error_continue: false, + }, + NativeConsoleCancellation::new(), + ) + .await + .expect("fail-closed SQL Server Console checks must terminate normally"); + let error = results + .first() + .and_then(|result| result.error.as_ref()) + .expect("unsafe SQL Server result must return a statement error"); + assert!( + expected_codes.contains(&error.code.as_str()), + "unexpected error {} for SQL: {sql}", + error.code + ); + } +} + +fn table_keys_request( + datasource_id: &str, + database_name: &str, + table_name: &str, +) -> ListCommunityTableKeysRequest { + ListCommunityTableKeysRequest { + datasource_id: datasource_id.to_owned(), + database_type: DATABASE_TYPE.to_owned(), + database_name: database_name.to_owned(), + schema_name: "dbo".to_owned(), + table_name: table_name.to_owned(), + } +} + +fn query_limits(max_rows: &str) -> QueryLimits { + QueryLimits { + max_rows: max_rows.to_owned(), + max_result_bytes: (8_u64 * 1024 * 1024).to_string(), + batch_rows: 2, + batch_bytes: 1024 * 1024, + result_ttl_seconds: 60, + } +} + +async fn wait_for_result(application: &Application, operation_id: &str) -> ResultMetadata { + let mut subscription = application + .subscribe_operation(operation_id, None) + .await + .expect("SQL Server query operation must be subscribable"); + tokio::time::timeout(EVENT_TIMEOUT, async { + while let Some(envelope) = subscription + .next_event() + .await + .expect("SQL Server operation event must decode") + { + match envelope.event { + OperationEvent::Completed { result } => return result, + OperationEvent::Failed { error } => { + panic!("native SQL Server query failed: {error:?}") + } + OperationEvent::Cancelled { reason } => { + panic!("native SQL Server query was cancelled: {reason:?}") + } + OperationEvent::Started | OperationEvent::Progress { .. } => {} + } + } + panic!("native SQL Server operation ended without a terminal event") + }) + .await + .expect("native SQL Server query must finish before timeout") +} + +async fn wait_for_failure_code(application: &Application, operation_id: &str) -> String { + let mut subscription = application + .subscribe_operation(operation_id, None) + .await + .expect("SQL Server query operation must be subscribable"); + tokio::time::timeout(EVENT_TIMEOUT, async { + while let Some(envelope) = subscription + .next_event() + .await + .expect("SQL Server operation event must decode") + { + match envelope.event { + OperationEvent::Failed { error } => return error.code, + OperationEvent::Completed { result } => { + panic!("unsafe SQL Server query unexpectedly completed: {result:?}") + } + OperationEvent::Cancelled { reason } => { + panic!("unsafe SQL Server query was cancelled: {reason:?}") + } + OperationEvent::Started | OperationEvent::Progress { .. } => {} + } + } + panic!("unsafe SQL Server query ended without a terminal event") + }) + .await + .expect("unsafe SQL Server query must fail before timeout") +} + +async fn result_page( + application: &Application, + result: &ResultMetadata, +) -> chat2db_contract::ResultPage { + application + .result_page( + &result.id, + ResultPageRequest { + offset: "0".to_owned(), + max_rows: "20".to_owned(), + max_bytes: (8_u64 * 1024 * 1024).to_string(), + }, + ) + .await + .expect("native SQL Server result page must be retained") +} + +async fn provision_database(config: &SqlServerTestConfig, database_name: &str) { + let mut client = config + .direct_client() + .await + .expect("SQL Server fixture connection must open"); + execute_batch(&mut client, &format!("CREATE DATABASE [{database_name}]")).await; + execute_batch(&mut client, &format!("USE [{database_name}]")).await; + execute_batch( + &mut client, + "CREATE TABLE dbo.native_parent (id int NOT NULL, label nvarchar(80) NULL, CONSTRAINT PK_native_parent PRIMARY KEY (id))", + ) + .await; + execute_batch( + &mut client, + "CREATE TABLE dbo.native_child (id int IDENTITY(10,5) NOT NULL, parent_id int NOT NULL, note nvarchar(80) NULL, note_lower AS LOWER(note), CONSTRAINT PK_native_child PRIMARY KEY (id), CONSTRAINT FK_native_child_parent FOREIGN KEY (parent_id) REFERENCES dbo.native_parent(id) ON UPDATE CASCADE ON DELETE CASCADE)", + ) + .await; + execute_batch( + &mut client, + "CREATE INDEX IX_native_child_note ON dbo.native_child(note DESC) INCLUDE(parent_id) WHERE note IS NOT NULL", + ) + .await; + execute_batch( + &mut client, + "CREATE VIEW dbo.native_view AS SELECT id, label FROM dbo.native_parent", + ) + .await; + execute_batch( + &mut client, + "CREATE FUNCTION dbo.native_function(@value int) RETURNS int AS BEGIN RETURN @value + 1; END", + ) + .await; + execute_batch( + &mut client, + "CREATE PROCEDURE dbo.native_procedure @value int AS SELECT @value AS value", + ) + .await; + execute_batch( + &mut client, + "CREATE TRIGGER dbo.native_trigger ON dbo.native_child AFTER INSERT AS BEGIN SET NOCOUNT ON; END", + ) + .await; + execute_batch( + &mut client, + "INSERT INTO dbo.native_parent(id,label) VALUES (7,N'native-rust'); INSERT INTO dbo.native_child(parent_id,note) VALUES (7,N'alpha'),(7,N'beta')", + ) + .await; +} + +async fn cleanup_database(config: &SqlServerTestConfig, database_name: &str) -> Result<(), String> { + let mut client = config.direct_client().await?; + let sql = format!( + "IF DB_ID(N'{database_name}') IS NOT NULL BEGIN ALTER DATABASE [{database_name}] SET SINGLE_USER WITH ROLLBACK IMMEDIATE; DROP DATABASE [{database_name}]; END" + ); + execute_batch_result(&mut client, &sql).await +} + +async fn execute_batch(client: &mut DirectClient, sql: &str) { + execute_batch_result(client, sql) + .await + .unwrap_or_else(|error| panic!("SQL Server fixture statement failed: {error}; SQL: {sql}")); +} + +async fn execute_batch_result(client: &mut DirectClient, sql: &str) -> Result<(), String> { + let mut stream = client + .simple_query(sql) + .await + .map_err(|error| error.to_string())?; + while stream + .try_next() + .await + .map_err(|error| error.to_string())? + .is_some() + {} + Ok(()) +} + +fn assert_java_dormant(application: &Application) { + let engine = application + .health() + .components + .into_iter() + .find(|component| component.id == "database-engine") + .expect("database engine health must be present"); + assert_eq!(engine.state, ComponentState::Ready); + assert_eq!(engine.detail, "Available on demand; Java is not running"); +} + +fn required_env(name: &str) -> String { + std::env::var(name).unwrap_or_else(|_| panic!("{name} must be configured")) +} diff --git a/docs/architecture.md b/docs/architecture.md index 9dd5614..ff8b790 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -95,6 +95,9 @@ real-MySQL rerun pass with this Dashboard/Chart increment included. | AI agent | Rust | Provider adapters, tool loop, limits, compaction, and cancellation | | MCP and CLI | Rust | Adapters around the same product services and policy | | Native MySQL product slice | Rust / `mysql_async` | Connection and SSH, datasource lifecycle/portability, object metadata, typed SELECT binds, editable DML/DDL, Console, Dashboard/Chart refresh, routines/migration, transfer and class generation, accounts, schema diff, workspace state, Agent/CLI/MCP writes, cancellation, large values, and historical HTTP/IPC envelopes | +| Native PostgreSQL product slice | Rust / `tokio-postgres` | Connection and SSH, retained query and typed binds, Console, relational/programmability metadata, DDL, preview, ER, cancellation, limits, and native schema/namespace/DML builders | +| Native SQL Server product slice | Rust / `tiberius` | Connection and SSH over TDS, retained query and typed binds, Console with direct-batch semantics, relational/programmability metadata, DDL, preview, ER, cancellation, limits, and native schema/namespace/DML builders | +| Native Oracle product slice | Rust / `oracle-rs` | Pure-Rust Oracle protocol without OCI/ODPI-C/JDBC; connection and SSH, retained query and typed binds, Console, relational/programmability metadata, DDL, preview, ER, cancellation, limits, and native schema/namespace/DML builders | | Hybrid DM product slice | Rust SPI plus generic Java JDBC | Rust owns DM metadata and preview behavior; the Java engine loads only the official DM JDBC JAR and streams typed JDBC results | | Remaining compatibility databases and exact Community helpers | Java 17 | Existing SPI/plugins for databases without a Rust-owned adapter plus Community parsing, formatting, completion, SQL builders, and plugin-specific behavior | | SQL parsing, formatting, and completion | Java 17 | Existing Java ANTLR grammars, parser behavior, formatter behavior, and completion | @@ -115,6 +118,9 @@ React in system WebView React in browser -> AI agent runtime -> Rust Driver SPI -> native MySQL / mysql_async + -> native PostgreSQL / tokio-postgres + -> native SQL Server / tiberius + -> native Oracle / oracle-rs -> DM metadata and preview adapter -> generic JDBC session bridge -> official dmJdbcDriver JAR @@ -152,12 +158,17 @@ cross-language acceptance gates pass. ## Database boundary -For DM, the Rust Driver SPI owns database-specific behavior and the project-owned -Java process is only a generic JDBC transport for the official vendor JAR. It -does not load Community's DM or Oracle plugins. The fixed Community runtime -remains the compatibility implementation only for database types that do not -have a registered Rust-owned driver, and for explicitly requested Community -parser, formatter, completion, builder, and plugin behavior. +For MySQL, PostgreSQL, SQL Server, and Oracle, the Rust Driver SPI owns the +implemented database wire protocol and workbench behavior. For DM, the SPI owns +database-specific behavior while the project-owned Java process is only a +generic JDBC transport for the official vendor JAR; it does not load +Community's DM or Oracle plugins. Existing persisted managed JDBC datasources +do not switch engines merely because a native driver is registered. A native +route is selected only by its explicit persisted driver id, except for MySQL +and DM whose adapters deliberately opt into the existing managed-JDBC aliases. +The fixed Community runtime remains the compatibility implementation for +managed JDBC and unregistered database types, and for explicitly requested +Community parser, formatter, completion, and plugin behavior. The native route uses upstream `mysql_async 0.37.0` for the complete MySQL product data plane: connection and SSH, metadata, editable DML and DDL, Console, typed SELECT bind parameters, Dashboard/Chart refresh, routines, @@ -202,6 +213,33 @@ The native MySQL baseline implements: Community-shaped response-only metadata, records `CHART` history, and rejects writes, multiple statements, locking reads, and server-file output. +The additional native relational adapters implement the common connection, +query/Console, metadata, table, routine, and dialect capabilities through the +same registry and neutral Core models: + +- PostgreSQL uses `tokio-postgres 0.7.18` and Rustls. It supports JDBC/native + URL normalization, SSH, positional typed parameters, explicit read-only + transactions, retained results, script splitting, database/schema/table and + relation metadata, functions/procedures/triggers, table DDL/preview, ER, and + schema/namespace/DML builders. +- SQL Server uses `tiberius 0.12.3` and Rustls over TDS. Direct unparameterized + batches preserve local temporary tables across Console statements, while + parameterized work uses the prepared path. Result-bearing writes, dispatch + failures, cancellation, and driver panics preserve conservative unknown-write + outcomes. Generated schema comments use the narrowly classified non-tabular + extended-property procedures; ordinary `EXEC` can still return rows. +- Oracle uses `oracle-rs 0.1.7`, a pure-Rust protocol client with no OCI, + ODPI-C, JDBC, or Java dependency. The adapter supports service/SID URLs, + TCPS, SSH with the remote TLS identity, forced read-only transactions, + metadata, query/Console, LOBs, DDL/preview, ER, and dialect builders. + `oracle-rs` fixes the first query prefetch at 100 rows and its legacy first + batch decoder cannot safely expose `BINARY_FLOAT`, `BINARY_DOUBLE`, `ROWID`, + or `UROWID`. The adapter rejects those described result types with + `oracle_result_type_not_supported`; it neither guesses values nor vendors a + patched driver. Oracle Free 23 is runtime-tested, while the advertised + connection floor remains Oracle 12.1+ and is not claimed as exhaustively + runtime-tested. + The JDBC baseline implements: - verified external JAR snapshots and per-driver classloader isolation;