diff --git a/Cargo.lock b/Cargo.lock index 32ad0774c..fc82ab6e9 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -622,9 +622,9 @@ checksum = "1fd0f2584146f6f2ef48085050886acf353beff7305ebd1ae69500e27c67f64b" [[package]] name = "bytes" -version = "1.11.1" +version = "1.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1e748733b7cbc798e1434b6ac524f0c1ff2ab456fe201501e6497c8417a4fc33" +checksum = "fc652a48c352aef3ea3aed32080501cf3ef6ed5da78602a020c991775b0aff04" dependencies = [ "serde", ] @@ -751,7 +751,7 @@ dependencies = [ "dirs", "futures", "hex", - "hmac", + "hmac 0.12.1", "itertools 0.12.1", "lazy_static", "log", @@ -771,7 +771,7 @@ dependencies = [ "serde_cbor", "serde_json", "serdect", - "sha2", + "sha2 0.10.8", "stack-auth", "stack-profile", "static_assertions", @@ -807,12 +807,12 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1946988e0b7f9de259d85b10c9c1fd7e1327751103b808c2c6e930b9a87c25c9" dependencies = [ "getrandom 0.2.15", - "hmac", + "hmac 0.12.1", "lazy_static", "num-bigint", "rand 0.8.6", "regex", - "sha2", + "sha2 0.10.8", "thiserror 1.0.69", ] @@ -822,7 +822,6 @@ version = "2.2.4" dependencies = [ "arc-swap", "async-trait", - "aws-lc-rs", "bigdecimal", "blake3", "bytes", @@ -836,12 +835,11 @@ dependencies = [ "eql-mapper", "exitcode", "hex", - "md-5", + "md-5 0.10.6", "metrics", "metrics-exporter-prometheus", "moka", - "oid-registry", - "pg_escape", + "pg-proto", "postgres-protocol", "postgres-types", "rand 0.9.2", @@ -853,20 +851,17 @@ dependencies = [ "rustls-platform-verifier 0.5.1", "serde", "serde_json", - "socket2 0.5.8", "sqltk", "temp-env", "thiserror 2.0.18", "tokio", "tokio-postgres", "tokio-postgres-rustls", - "tokio-rustls", "tokio-util", "tracing", "tracing-subscriber", "uuid", "vitaminc-protected 0.1.0-pre4.2", - "x509-parser", ] [[package]] @@ -972,6 +967,12 @@ dependencies = [ "cc", ] +[[package]] +name = "cmov" +version = "0.5.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0c9ea0ac24bc397ab3c98583a3c9ba74fa56b09a4449bbe172b9b1ddb016027a" + [[package]] name = "colorchoice" version = "1.0.3" @@ -1031,6 +1032,12 @@ version = "0.9.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c2459377285ad874054d797f3ccebf984978aa39129f6eafde5cdc8315b612f8" +[[package]] +name = "const-oid" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a6ef517f0926dd24a1582492c791b6a4818a4d94e789a334894aa15b0d12f55c" + [[package]] name = "constant_time_eq" version = "0.3.1" @@ -1192,6 +1199,15 @@ dependencies = [ "vitaminc", ] +[[package]] +name = "ctutils" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7d5515a3834141de9eafb9717ad39eea8247b5674e6066c404e8c4b365d2a29e" +dependencies = [ + "cmov", +] + [[package]] name = "darling" version = "0.20.10" @@ -1248,7 +1264,7 @@ version = "0.7.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f55bf8e7b65898637379c1b74eb1551107c8294ed26d855ceb9fd1a09cfc9bc0" dependencies = [ - "const-oid", + "const-oid 0.9.6", "der_derive", "flagset", "zeroize", @@ -1393,7 +1409,9 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f1dd6dbb5841937940781866fa1281a1ff7bd3bf827091440879f9994983d5c2" dependencies = [ "block-buffer 0.12.1", + "const-oid 0.10.2", "crypto-common 0.2.2", + "ctutils", ] [[package]] @@ -1569,7 +1587,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "33d852cb9b869c2a9b3df2f71a3074817f01e1844f839a144f5fcef059a4eb5d" dependencies = [ "libc", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] @@ -2006,6 +2024,15 @@ dependencies = [ "digest 0.10.7", ] +[[package]] +name = "hmac" +version = "0.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6303bc9732ae41b04cb554b844a762b4115a61bfaa81e3e83050991eeb56863f" +dependencies = [ + "digest 0.11.3", +] + [[package]] name = "http" version = "1.3.1" @@ -2120,7 +2147,7 @@ dependencies = [ "libc", "percent-encoding", "pin-project-lite", - "socket2 0.6.1", + "socket2 0.6.5", "system-configuration", "tokio", "tower-service", @@ -2500,9 +2527,9 @@ checksum = "09edd9e8b54e49e587e4f6295a7d29c3ea94d469cb40ab8ca70b288248a81db2" [[package]] name = "libc" -version = "0.2.177" +version = "0.2.189" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2874a2af47a2325c2001a6e6fad9b16a53b802102b528163885171cf92b15976" +checksum = "3eaf3ede3fee6db1a4c2ee091bf8a8b4dccdc6d17f656fb07896ee72867612f2" [[package]] name = "libm" @@ -2592,6 +2619,16 @@ dependencies = [ "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 = "memchr" version = "2.7.4" @@ -2654,7 +2691,7 @@ dependencies = [ "cfg-if", "miette-derive", "thiserror 1.0.69", - "unicode-width", + "unicode-width 0.1.14", ] [[package]] @@ -2691,13 +2728,13 @@ dependencies = [ [[package]] name = "mio" -version = "1.0.3" +version = "1.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2886843bf800fba2e3377cff24abf6379b4c4d5c6681eaf9ea5b0d15090450bd" +checksum = "30d65c71f1ce40ab09135ce117d742b9f8a19ff91a41a8b57ed50bc2de59c427" dependencies = [ "libc", "wasi 0.11.0+wasi-snapshot-preview1", - "windows-sys 0.52.0", + "windows-sys 0.61.2", ] [[package]] @@ -2983,45 +3020,47 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e3148f5046208a5d56bcfc03053e3ca6334e51da8dfb19b6cdc8b306fae3283e" [[package]] -name = "pg_escape" -version = "0.1.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "44c7bc82ccbe2c7ef7ceed38dcac90d7ff46681e061e9d7310cbcd409113e303" -dependencies = [ - "phf", -] - -[[package]] -name = "phf" -version = "0.11.3" +name = "pg-proto" +version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1fd6780a80ae0c52cc120a26a1a42c1ae51b247a253e4e06113d23d2c2edd078" +checksum = "5275df13ca49c9ef28ad5d63006b6d93a38d6f1ab9f9655475b305be0a8c4e58" dependencies = [ - "phf_macros", - "phf_shared", + "base64", + "bytes", + "hmac 0.13.0", + "pg-proto-fsm", + "postgres-protocol", + "rand 0.10.2", + "rustls", + "sha2 0.11.0", + "socket2 0.6.5", + "stringprep", + "subtle", + "tokio", + "tokio-rustls", + "tokio-util", + "x509-parser", ] [[package]] -name = "phf_generator" -version = "0.11.3" +name = "pg-proto-fsm" +version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3c80231409c20246a13fddb31776fb942c38553c51e871f8cbd687a4cfb5843d" +checksum = "3574ed461ac3cbc508cb2c0d44e8280c346a46856e3760bb544c0bc7ec4bcc64" dependencies = [ - "phf_shared", - "rand 0.8.6", + "proc-macro2", + "quote", + "railroad", + "syn 3.0.3", ] [[package]] -name = "phf_macros" +name = "phf" version = "0.11.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f84ac04429c13a7ff43785d75ad27569f2951ce0ffd30a3321230db2fc727216" +checksum = "1fd6780a80ae0c52cc120a26a1a42c1ae51b247a253e4e06113d23d2c2edd078" dependencies = [ - "phf_generator", "phf_shared", - "proc-macro2", - "quote", - "syn 2.0.117", ] [[package]] @@ -3083,19 +3122,19 @@ dependencies = [ [[package]] name = "postgres-protocol" -version = "0.6.8" +version = "0.6.12" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "76ff0abab4a9b844b93ef7b81f1efc0a366062aaef2cd702c76256b5dc075c54" +checksum = "08808e3c483c46e999108051c78334f473d5adb59d78bb80a1268c7e6aa6c514" dependencies = [ "base64", "byteorder", "bytes", "fallible-iterator", - "hmac", - "md-5", + "hmac 0.13.0", + "md-5 0.11.0", "memchr", - "rand 0.9.2", - "sha2", + "rand 0.10.2", + "sha2 0.11.0", "stringprep", ] @@ -3205,9 +3244,9 @@ dependencies = [ [[package]] name = "proc-macro2" -version = "1.0.95" +version = "1.0.107" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "02b3e5e68a3a1a02aad3ec490a98007cbc13c37cbe84a3cd7b8e406d76e7f778" +checksum = "985e7ec9bb745e6ce6535b544d84d6cd6f7ad8bd711c398938ae983b91a766d9" dependencies = [ "unicode-ident", ] @@ -3325,14 +3364,14 @@ dependencies = [ "once_cell", "socket2 0.5.8", "tracing", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] name = "quote" -version = "1.0.40" +version = "1.0.47" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1885c039570dc00dcb4ff087a89e185fd56bae234ddc7f056a945bf36467248d" +checksum = "1fbf4db142a473a8d80c26bbf18454ed458bf8d26c8219c331daecfdbd079001" dependencies = [ "proc-macro2", ] @@ -3355,6 +3394,15 @@ version = "0.7.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "dc33ff2d4973d518d823d61aa239014831e521c75da58e3df4840d3f47749d09" +[[package]] +name = "railroad" +version = "0.3.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4bf842ad92d09c4dd1e68be1507189b01898c46fe01916c0467fb405b1f9ee0a" +dependencies = [ + "unicode-width 0.2.2", +] + [[package]] name = "rand" version = "0.8.6" @@ -3474,7 +3522,7 @@ dependencies = [ "rand_chacha 0.3.1", "serde", "serde_cbor", - "sha2", + "sha2 0.10.8", "thiserror 1.0.69", "zeroize", ] @@ -3495,7 +3543,7 @@ dependencies = [ "rand_chacha 0.3.1", "serde", "serde_cbor", - "sha2", + "sha2 0.10.8", "thiserror 1.0.69", "zeroize", ] @@ -3557,7 +3605,7 @@ checksum = "2c9283685feec7d69af75fb0e858d5e7378f33fe4fc699383b2916ab9273e03c" dependencies = [ "proc-macro2", "quote", - "syn 3.0.2", + "syn 3.0.3", ] [[package]] @@ -3775,14 +3823,14 @@ dependencies = [ "errno", "libc", "linux-raw-sys", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] name = "rustls" -version = "0.23.28" +version = "0.23.43" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7160e3e10bf4535308537f3c4e1641468cd0e485175d6163087c0393c7d46643" +checksum = "0283386ce02abc0151e1761d08802dfe86c173b0b494af5cbc086574e453da06" dependencies = [ "aws-lc-rs", "log", @@ -3807,11 +3855,12 @@ dependencies = [ [[package]] name = "rustls-pki-types" -version = "1.11.0" +version = "1.15.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "917ce264624a4b4db1c364dcc35bfca9ded014d0a958cd47ad3e960e988ea51c" +checksum = "2f4925028c7eb5d1fcdaf196971378ed9d2c1c4efc7dc5d011256f76c99c0a96" dependencies = [ "web-time", + "zeroize", ] [[package]] @@ -3832,7 +3881,7 @@ dependencies = [ "security-framework", "security-framework-sys", "webpki-root-certs 0.26.8", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] @@ -3864,9 +3913,9 @@ checksum = "f87165f0995f63a9fbeea62b64d10b4d9d8e78ec6d7d51fb2125fda7bb36788f" [[package]] name = "rustls-webpki" -version = "0.103.3" +version = "0.103.13" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e4a72fe2bcf7a6ac6fd7d0b9e5cb68aeb7d4c0a0271730218b3e92d43b4eb435" +checksum = "61c429a8649f110dddef65e2a5ad240f747e85f7758a6bccc7e5777bd33f756e" dependencies = [ "aws-lc-rs", "ring", @@ -4161,6 +4210,17 @@ dependencies = [ "digest 0.10.7", ] +[[package]] +name = "sha2" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "446ba717509524cb3f22f17ecc096f10f4822d76ab5c0b9822c5f9c284e825f4" +dependencies = [ + "cfg-if", + "cpufeatures 0.3.0", + "digest 0.11.3", +] + [[package]] name = "sharded-slab" version = "0.1.7" @@ -4258,12 +4318,12 @@ dependencies = [ [[package]] name = "socket2" -version = "0.6.1" +version = "0.6.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "17129e116933cf371d018bb80ae557e889637989d8638274fb25622827b03881" +checksum = "c3d1e2c7f27f8d4cb10542a02c49005dbd6e93095799d6f3be745fae9f8fedd4" dependencies = [ "libc", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -4355,7 +4415,7 @@ dependencies = [ "cfg-if", "libc", "psm", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] @@ -4417,9 +4477,9 @@ dependencies = [ [[package]] name = "syn" -version = "3.0.2" +version = "3.0.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a207d6d6a2b7fc470b80443726053f18a2481b7e1eee970597051596567987a3" +checksum = "53e9bae58849f64dfa4f5d5ae372c8341f7305f82a3868709269343628b659a3" dependencies = [ "proc-macro2", "quote", @@ -4626,9 +4686,9 @@ dependencies = [ [[package]] name = "tokio" -version = "1.48.0" +version = "1.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ff360e02eab121e0bc37a2d3b4d4dc622e6eda3a8e5253d5435ecf5bd4c68408" +checksum = "202caea871b69668250d242070849eb495be178ed697a3e98aebce5bc81a0bed" dependencies = [ "bytes", "libc", @@ -4636,20 +4696,20 @@ dependencies = [ "parking_lot", "pin-project-lite", "signal-hook-registry", - "socket2 0.6.1", + "socket2 0.6.5", "tokio-macros", "windows-sys 0.61.2", ] [[package]] name = "tokio-macros" -version = "2.6.0" +version = "2.7.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "af407857209536a95c8e56f8231ef2c2e2aff839b22e07a1ffcbc617e9db9fa5" +checksum = "78773a2a397f451582ce068015985c33193cf6dea8b74d2a639fe457b2f07b0e" dependencies = [ "proc-macro2", "quote", - "syn 2.0.117", + "syn 3.0.3", ] [[package]] @@ -4684,7 +4744,7 @@ version = "0.13.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "27d684bad428a0f2481f42241f821db42c54e2dc81d8c00db8536c506b0a0144" dependencies = [ - "const-oid", + "const-oid 0.9.6", "ring", "rustls", "tokio", @@ -4695,9 +4755,9 @@ dependencies = [ [[package]] name = "tokio-rustls" -version = "0.26.2" +version = "0.26.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8e727b36a1a0e8b74c376ac2211e40c2c8af09fb4013c60d910495810f008e9b" +checksum = "1729aa945f29d91ba541258c8df89027d5792d85a8841fb65e8bf0f4ede4ef61" dependencies = [ "rustls", "tokio", @@ -4705,15 +4765,15 @@ dependencies = [ [[package]] name = "tokio-util" -version = "0.7.14" +version = "0.7.19" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6b9590b93e6fcc1739458317cccd391ad3955e2bde8913edf6f95f9e65a8f034" +checksum = "494815d09bf52b5548659851081238f0ca39ff638363907596da739561c62c52" dependencies = [ "bytes", "futures-core", "futures-sink", "futures-util", - "hashbrown 0.14.5", + "libc", "pin-project-lite", "tokio", ] @@ -4964,6 +5024,12 @@ version = "0.1.14" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7dd6e30e90baa6f72411720665d41d89b9a3d039dc45b8faea1ddd07f617f6af" +[[package]] +name = "unicode-width" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b4ac048d71ede7ee76d585517add45da530660ef4390e49b098733c6e897f254" + [[package]] name = "unicode-xid" version = "0.2.6" @@ -5056,7 +5122,7 @@ dependencies = [ "atomic", "getrandom 0.3.2", "js-sys", - "md-5", + "md-5 0.10.6", "serde", "sha1_smol", "wasm-bindgen", @@ -5496,7 +5562,7 @@ version = "0.1.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cf221c93e13a30d793f7645a0e7762c55d169dbb0a49671918a2319d289b10bb" dependencies = [ - "windows-sys 0.48.0", + "windows-sys 0.59.0", ] [[package]] @@ -5678,15 +5744,6 @@ dependencies = [ "windows-targets 0.52.6", ] -[[package]] -name = "windows-sys" -version = "0.60.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f2f500e4d28234f72040990ec9d39e3a6b950f9f22d3dba18416c35882612bcb" -dependencies = [ - "windows-targets 0.53.5", -] - [[package]] name = "windows-sys" version = "0.61.2" @@ -6100,7 +6157,7 @@ version = "0.2.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1301e935010a701ae5f8655edc0ad17c44bad3ac5ce8c39185f75453b720ae94" dependencies = [ - "const-oid", + "const-oid 0.9.6", "der", "spki", "tls_codec", @@ -6108,9 +6165,9 @@ dependencies = [ [[package]] name = "x509-parser" -version = "0.17.0" +version = "0.18.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4569f339c0c402346d4a75a9e39cf8dad310e287eef1ff56d4c68e5067f53460" +checksum = "d43b0f71ce057da06bc0851b23ee24f3f86190b07203dd8f567d0b706a185202" dependencies = [ "asn1-rs", "data-encoding", diff --git a/PG_PROTO_FOLLOWUPS.md b/PG_PROTO_FOLLOWUPS.md new file mode 100644 index 000000000..d41fab3af --- /dev/null +++ b/PG_PROTO_FOLLOWUPS.md @@ -0,0 +1,40 @@ +# pg-proto follow-ups + +The proxy migration delegates framing, typed messages, startup/authentication +state, typed startup middleware, async runtime middleware, demultiplexing, and +compile-time checked bounded pipeline dispatch to `pg-proto` 0.5.0. The +remaining transport adapters below would be better eliminated in `pg-proto` +itself. + +## Preserve buffered transport state across a split + +`Buffered::into_inner()` cannot return retained inbound bytes, pending outbound +bytes, or demultiplexer state. The proxy must finish startup through +`ReadyForQuery` before splitting the bidirectional stream and then construct new +`Buffered` values for concurrent frontend/backend processing. + +An `into_parts`/`from_parts` API, or a buffer-preserving split API, would let a +proxy change transport ownership without risking loss of bytes already read past +a message boundary. It should preserve both codec buffers and the demultiplexer. + +## Transfer a demultiplexer between transport owners + +The proxy routes backend messages through pg-proto's demultiplexer, but startup +and concurrent runtime currently use separate `Buffered` owners. A supported way +to extract and restore `Demux` state would retain startup parameter status, +cancellation-key, and readiness state without application bookkeeping. + +This may naturally be solved by the buffer-preserving transport-parts API above. + +## Publish the example proxy driver as a library API + +`pg-proto` demonstrates a clean `Buffered` + `Middleware` forwarding loop in +`examples/proxy_support`, but does not expose a configurable proxy driver from +the crate. CipherStash therefore still owns connection orchestration, concurrent +forwarding, and the small amount of glue that invokes middleware. + +A library-level proxy builder should accept downstream/upstream transports, +startup and authentication policy, typed frontend/backend middleware, timeout +policy, and an output strategy. It should own framing, phase transitions, +bounded pipeline dispatch, demultiplexing, and shutdown. Applications would then +only supply policy and message transformations. diff --git a/PG_PROTO_MIGRATION_PLAN.md b/PG_PROTO_MIGRATION_PLAN.md new file mode 100644 index 000000000..09c2ba7c7 --- /dev/null +++ b/PG_PROTO_MIGRATION_PLAN.md @@ -0,0 +1,38 @@ +# pg-proto Migration + +## Summary + +- Create `/Users/jamessadler/cipherstash/proxy-pg-proto` from current `main` (`15b7f996`) on branch `refactor/pg-proto`. +- Save this migration plan as `PG_PROTO_MIGRATION_PLAN.md` in that worktree’s repository root. +- Limit the initial deliverable to the worktree, branch, and plan document; implementation follows separately. +- Target a full migration to published [`pg-proto` 0.2.1](https://crates.io/crates/pg-proto/0.2.1), covering codecs, startup/authentication, runtime protocol-state validation, and bounded pipeline orchestration. + +## Implementation Changes + +- Replace handwritten framing, startup packet parsing, message codes, and message serialization with direction-specific `pg-proto` frontend/backend codecs. +- Convert CipherStash-specific behavior into adapters over `pg-proto` messages: + - Preserve Parse/Query SQL rewriting and parameter OID mapping. + - Preserve Bind format-code semantics, nulls, parameter reshaping, and encryption. + - Preserve ParameterDescription, RowDescription, and DataRow rewriting and batched decryption. + - Retain diagnostic-response factories while emitting `pg-proto` response types. +- Use `pg-proto` pre-startup and authentication APIs for SSL negotiation, startup, cancellation, client-facing MD5 authentication, and upstream cleartext/MD5/SCRAM authentication. Continue using existing TLS configuration and certificate policy. +- Pair downstream server-role and upstream client-role runtime FSMs through `Intermediary`. Advance both sides for forwarded messages and only the affected side for locally intercepted or synthesised messages. +- Preserve concurrent client-to-server and server-to-client processing, connection timeouts, response ordering, metrics, logging, schema reloads, and row buffering. +- Track protocol state even when encryption mapping is disabled. Preserve one-to-one cancellation forwarding; do not introduce pooling or cancellation-key translation. +- Reject unknown message tags as protocol errors, matching `pg-proto`’s fail-closed behavior. +- Remove obsolete handwritten protocol modules and direct low-level dependencies once unused. Preserve the public configuration and CLI surfaces; retain existing `ProtocolError` variants for source compatibility even where `pg-proto` supersedes them. + +## Test Plan + +- Port existing message round-trip and rewrite tests to `pg-proto` message fixtures. +- Add coverage for partial/oversized frames, malformed messages, unknown-tag rejection, SSL/TLS startup, cancellation, and all supported authentication modes. +- Exercise simple queries and extended Parse/Bind/Describe/Execute/Close/Sync pipelines, including pipelining and error draining through Sync. +- Verify text/binary formats, nulls, reshaped parameters, prepared statements, portals, COPY messages, asynchronous backend messages, and buffered DataRow decryption. +- Run formatting, clippy, proxy unit tests, and TCP/TLS integration suites. Unset `CS_PROMETHEUS__ENABLED` for the baseline unit suite; its current environment value causes the otherwise unrelated Prometheus test to fail. + +## Assumptions + +- The existing untracked `.claude/worktrees/` directory remains untouched. +- The worktree path and branch are currently available. +- No compatibility feature flag or dual protocol implementation is required. +- The plan document is left as an uncommitted worktree change unless a commit is requested separately. diff --git a/packages/cipherstash-proxy/Cargo.toml b/packages/cipherstash-proxy/Cargo.toml index 73dce1b9b..e7a0e925b 100644 --- a/packages/cipherstash-proxy/Cargo.toml +++ b/packages/cipherstash-proxy/Cargo.toml @@ -5,7 +5,6 @@ edition = "2021" [dependencies] async-trait = "0.1" -aws-lc-rs = "1.13.3" bigdecimal = { version = "0.4.6", features = ["serde-json"] } blake3 = "1" arc-swap = "1.7.1" @@ -29,8 +28,7 @@ md-5 = "0.10.6" metrics = "0.24.3" metrics-exporter-prometheus = "0.17" moka = { version = "0.12", features = ["future"] } -oid-registry = "0.8" -pg_escape = "0.1.1" +pg-proto = "0.5.0" postgres-protocol = "0.6.7" postgres-types = { version = "0.2.8", features = ["with-serde_json-1"] } rand = "0.9" @@ -43,7 +41,6 @@ rustls-platform-verifier = "0.5.0" rustls-pki-types = "1.10.0" serde = "1.0" serde_json = "1.0" -socket2 = "0.5.7" sqltk = { workspace = true } thiserror = { workspace = true } tokio = { workspace = true } @@ -52,13 +49,11 @@ tokio-postgres = { version = "0.7", features = [ "with-serde_json-1", ] } tokio-postgres-rustls = "0.13.0" -tokio-rustls = "0.26.0" tokio-util = { version = "0.7.13", features = ["rt"] } tracing = { workspace = true } tracing-subscriber = { workspace = true } uuid = { version = "1.11.0", features = ["serde", "v4"] } vitaminc-protected = "0.1.0-pre4.2" -x509-parser = "0.17.0" [dev-dependencies] diff --git a/packages/cipherstash-proxy/src/connect/async_stream.rs b/packages/cipherstash-proxy/src/connect/async_stream.rs deleted file mode 100644 index 83436a05b..000000000 --- a/packages/cipherstash-proxy/src/connect/async_stream.rs +++ /dev/null @@ -1,148 +0,0 @@ -use super::{configure, connect_with_retry}; -use crate::{error::Error, log::AUTHENTICATION}; -use aws_lc_rs::digest; -use core::str; -use oid_registry::{ - Oid, OID_HASH_SHA1, OID_NIST_HASH_SHA256, OID_NIST_HASH_SHA384, OID_NIST_HASH_SHA512, - OID_PKCS1_SHA1WITHRSA, OID_PKCS1_SHA256WITHRSA, OID_PKCS1_SHA384WITHRSA, - OID_PKCS1_SHA512WITHRSA, OID_SIG_ECDSA_WITH_SHA256, OID_SIG_ECDSA_WITH_SHA384, OID_SIG_ED25519, -}; -use postgres_protocol::authentication::sasl::ChannelBinding; - -use std::{ - pin::Pin, - task::{Context, Poll}, -}; -use tokio::{ - io::{split, AsyncRead, AsyncWrite, ReadBuf}, - net::{TcpListener, TcpStream}, -}; -use tokio_rustls::TlsStream; -use tracing::debug; -use x509_parser::prelude::{FromDer, X509Certificate}; - -#[derive(Debug)] -pub enum AsyncStream { - Tcp(TcpStream), - Tls(Box>), -} - -impl AsyncStream { - pub async fn accept(listener: &TcpListener) -> Result { - let (stream, _) = listener.accept().await?; - configure(&stream); - Ok(AsyncStream::Tcp(stream)) - } - - pub async fn connect(addr: &str) -> Result { - let stream = connect_with_retry(addr).await?; - configure(&stream); - Ok(AsyncStream::Tcp(stream)) - } - - pub fn split( - self, - ) -> ( - tokio::io::ReadHalf, - tokio::io::WriteHalf, - ) { - split(self) - } - - pub fn is_tls(&self) -> bool { - matches!(self, AsyncStream::Tls(_)) - } - - pub fn is_tcp(&self) -> bool { - !self.is_tls() - } - - pub fn channel_binding(&self) -> ChannelBinding { - match self { - AsyncStream::Tcp(_) => ChannelBinding::unsupported(), - AsyncStream::Tls(stream) => { - let (_, session) = stream.get_ref(); - let certs = session.peer_certificates(); - match certs { - Some(certs) if !certs.is_empty() => { - let cert_der = &certs[0]; - X509Certificate::from_der(cert_der) - .ok() - .map(|(_, cert)| get_digest(&cert.signature_algorithm.algorithm)) - .map_or_else(ChannelBinding::unsupported, |algorithm| { - let hash = digest::digest(algorithm, certs[0].as_ref()); - ChannelBinding::tls_server_end_point(hash.as_ref().into()) - }) - } - _ => { - debug!( - target: AUTHENTICATION, - msg = "Missing certificates, ChannelBinding is unsupported" - ); - ChannelBinding::unsupported() - } - } - } - } - } -} - -/// -/// Note: SHA1 is upgraded to SHA256 as per https://datatracker.ietf.org/doc/html/rfc5929#section-4.1 -/// -fn get_digest(oid: &Oid) -> &'static digest::Algorithm { - match oid { - oid if oid == &OID_HASH_SHA1 => &digest::SHA256, - oid if oid == &OID_NIST_HASH_SHA256 => &digest::SHA256, - oid if oid == &OID_PKCS1_SHA1WITHRSA => &digest::SHA256, - oid if oid == &OID_PKCS1_SHA256WITHRSA => &digest::SHA256, - oid if oid == &OID_SIG_ECDSA_WITH_SHA256 => &digest::SHA256, - oid if oid == &OID_NIST_HASH_SHA384 => &digest::SHA384, - oid if oid == &OID_PKCS1_SHA384WITHRSA => &digest::SHA384, - oid if oid == &OID_SIG_ECDSA_WITH_SHA384 => &digest::SHA384, - oid if oid == &OID_NIST_HASH_SHA512 => &digest::SHA512, - oid if oid == &OID_PKCS1_SHA512WITHRSA => &digest::SHA512, - oid if oid == &OID_SIG_ED25519 => &digest::SHA512, - _ => panic!("Unsupported OID"), - } -} - -impl AsyncRead for AsyncStream { - fn poll_read( - mut self: Pin<&mut Self>, - cx: &mut Context<'_>, - buf: &mut ReadBuf<'_>, - ) -> Poll> { - match *self { - AsyncStream::Tcp(ref mut stream) => Pin::new(stream).poll_read(cx, buf), - AsyncStream::Tls(ref mut stream) => Pin::new(stream).poll_read(cx, buf), - } - } -} - -impl AsyncWrite for AsyncStream { - fn poll_write( - mut self: Pin<&mut Self>, - cx: &mut Context<'_>, - buf: &[u8], - ) -> Poll> { - match *self { - AsyncStream::Tcp(ref mut stream) => Pin::new(stream).poll_write(cx, buf), - AsyncStream::Tls(ref mut stream) => Pin::new(stream).poll_write(cx, buf), - } - } - - fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { - match *self { - AsyncStream::Tcp(ref mut stream) => Pin::new(stream).poll_flush(cx), - AsyncStream::Tls(ref mut stream) => Pin::new(stream).poll_flush(cx), - } - } - - fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { - match *self { - AsyncStream::Tcp(ref mut stream) => Pin::new(stream).poll_shutdown(cx), - AsyncStream::Tls(ref mut stream) => Pin::new(stream).poll_shutdown(cx), - } - } -} diff --git a/packages/cipherstash-proxy/src/connect/channel_writer.rs b/packages/cipherstash-proxy/src/connect/channel_writer.rs index a6dc68ffb..bd72ba1f2 100644 --- a/packages/cipherstash-proxy/src/connect/channel_writer.rs +++ b/packages/cipherstash-proxy/src/connect/channel_writer.rs @@ -1,4 +1,7 @@ -use bytes::BytesMut; +use pg_proto::{ + codec::{BackendMessage, Frontend}, + transport::Buffered, +}; use tokio::{ io::{AsyncWrite, AsyncWriteExt}, sync::mpsc::{self, UnboundedReceiver, UnboundedSender}, @@ -7,15 +10,15 @@ use tracing::{debug, error}; use crate::log::PROTOCOL; -pub type Receiver = UnboundedReceiver; -pub type Sender = UnboundedSender; +pub type Receiver = UnboundedReceiver; +pub type Sender = UnboundedSender; #[derive(Debug)] pub struct ChannelWriter where W: AsyncWrite + Unpin, { - writer: W, + writer: Buffered, receiver: Receiver, sender: Sender, client_id: i32, @@ -26,11 +29,13 @@ where W: AsyncWrite + Unpin, { pub fn new(writer: W, client_id: i32) -> Self { - let (sender, receiver): (UnboundedSender, UnboundedReceiver) = - mpsc::unbounded_channel(); + let (sender, receiver): ( + UnboundedSender, + UnboundedReceiver, + ) = mpsc::unbounded_channel(); ChannelWriter { - writer, + writer: Buffered::new_frontend(writer), receiver, sender, client_id, @@ -48,14 +53,19 @@ where // but we're holding one of them ourselves! drop(self.sender); - while let Some(bytes) = self.receiver.recv().await { + while let Some(message) = self.receiver.recv().await { debug!(target: PROTOCOL, client_id = self.client_id, msg = "Writing", - ?bytes + ?message ); - match self.writer.write_all(&bytes).await { + let result = message.to_frame().and_then(|frame| self.writer.push(frame)); + let result = match result { + Ok(()) => self.writer.flush().await, + Err(error) => Err(error), + }; + match result { Ok(_) => { debug!(target: PROTOCOL, client_id = self.client_id, @@ -89,7 +99,7 @@ where } // Shutdown the write half to send FIN and properly close the connection - if let Err(err) = self.writer.shutdown().await { + if let Err(err) = self.writer.into_inner().shutdown().await { error!(target: PROTOCOL, client_id = self.client_id, msg = "Error shutting down writer", diff --git a/packages/cipherstash-proxy/src/connect/mod.rs b/packages/cipherstash-proxy/src/connect/mod.rs index d17d97969..9fa25cd4e 100644 --- a/packages/cipherstash-proxy/src/connect/mod.rs +++ b/packages/cipherstash-proxy/src/connect/mod.rs @@ -1,11 +1,9 @@ -mod async_stream; mod channel_writer; -pub use async_stream::AsyncStream; pub use channel_writer::{ChannelWriter, Sender}; -use crate::{config::ServerConfig, error::Error, log::DEVELOPMENT, tls, DatabaseConfig}; -use socket2::TcpKeepalive; +use crate::{config::ServerConfig, error::Error, tls, DatabaseConfig}; +use pg_proto::net::{ConnectRetry, NetworkStream, TcpSettings}; use std::time::Duration; use tokio::{ net::{TcpListener, TcpStream}, @@ -14,13 +12,48 @@ use tokio::{ use tokio_postgres::Client; use tracing::{debug, error, info, warn}; +const MAX_RETRY_DELAY: Duration = Duration::from_secs(2); +const MAX_RETRY_COUNT: u32 = 3; const TCP_USER_TIMEOUT: Duration = Duration::from_secs(10); const TCP_KEEPALIVE_INTERVAL: Duration = Duration::from_secs(5); const TCP_KEEPALIVE_TIME: Duration = Duration::from_secs(5); const TCP_KEEPALIVE_RETRIES: u32 = 5; -const MAX_RETRY_DELAY: Duration = Duration::from_secs(2); -const MAX_RETRY_COUNT: u32 = 3; +fn configure_tcp(stream: &TcpStream) { + let settings = TcpSettings { + no_delay: true, + user_timeout: Some(TCP_USER_TIMEOUT), + keepalive_time: Some(TCP_KEEPALIVE_TIME), + keepalive_interval: Some(TCP_KEEPALIVE_INTERVAL), + keepalive_retries: Some(TCP_KEEPALIVE_RETRIES), + }; + for error in pg_proto::net::configure_tcp(stream, settings) { + warn!(msg = "Error configuring connection", error = %error); + } +} + +pub async fn accept(listener: &TcpListener) -> Result, Error> { + let (stream, _) = listener.accept().await?; + configure_tcp(&stream); + Ok(NetworkStream::plain(stream)) +} + +pub async fn connect(address: &str) -> Result, Error> { + debug!(msg = "Connecting to database"); + let retry = ConnectRetry { + max_retries: MAX_RETRY_COUNT, + initial_delay: Duration::from_millis(100), + max_delay: MAX_RETRY_DELAY, + }; + let stream = pg_proto::net::connect_with_retry(address, retry) + .await + .map_err(|error| { + error!(msg = "Could not connect to database", error = %error); + Error::DatabaseConnection + })?; + configure_tcp(&stream); + Ok(NetworkStream::plain(stream)) +} pub async fn database(config: &DatabaseConfig) -> Result { let connection_config = config.to_connection_config(); @@ -79,81 +112,3 @@ pub async fn bind_with_retry(server: &ServerConfig) -> TcpListener { retry_count += 1; } } - -pub async fn connect_with_retry(addr: &str) -> Result { - let mut retry_count = 0; - - loop { - debug!(target: DEVELOPMENT, msg = "Connecting to database"); - match TcpStream::connect(&addr).await { - Ok(stream) => { - return Ok(stream); - } - Err(err) => { - if retry_count > MAX_RETRY_COUNT { - error!(msg = "Could not connect to database", retries = ?retry_count, error = err.to_string()); - return Err(Error::DatabaseConnection); - } - } - }; - let sleep_duration_ms = - (100 * 2_u64.pow(retry_count)).min(MAX_RETRY_DELAY.as_millis() as _); - time::sleep(Duration::from_millis(sleep_duration_ms)).await; - - retry_count += 1; - } -} - -/// -/// Configure the tcp socket -/// set_nodelay -/// set_keepalive -/// -/// Keepalive is not as important without connection pooling timeouts to deal with -/// -pub fn configure(stream: &TcpStream) { - let sock_ref = socket2::SockRef::from(&stream); - - stream.set_nodelay(true).unwrap_or_else(|err| { - warn!( - msg = "Error configuring nodelay for connection", - error = err.to_string() - ); - }); - - #[cfg(target_os = "linux")] - match sock_ref.set_tcp_user_timeout(Some(TCP_USER_TIMEOUT)) { - Ok(_) => (), - Err(err) => { - warn!( - msg = "Error configuring tcp_user_timeout for connection", - error = err.to_string() - ); - } - } - - match sock_ref.set_keepalive(true) { - Ok(_) => { - let params = &TcpKeepalive::new() - .with_interval(TCP_KEEPALIVE_INTERVAL) - .with_retries(TCP_KEEPALIVE_RETRIES) - .with_time(TCP_KEEPALIVE_TIME); - - match sock_ref.set_tcp_keepalive(params) { - Ok(_) => (), - Err(err) => { - warn!( - msg = "Error configuring keepalive for connection", - error = err.to_string() - ); - } - } - } - Err(err) => { - warn!( - msg = "Error configuring connection", - error = err.to_string() - ); - } - } -} diff --git a/packages/cipherstash-proxy/src/error.rs b/packages/cipherstash-proxy/src/error.rs index fb3982b4f..1b12e7f8c 100644 --- a/packages/cipherstash-proxy/src/error.rs +++ b/packages/cipherstash-proxy/src/error.rs @@ -1,5 +1,4 @@ use crate::{postgresql::Column, Identifier}; -use bytes::BytesMut; use cipherstash_client::{encryption, schema::ColumnType}; use eql_mapper::{EqlMapperError, EqlTermVariant}; use metrics_exporter_prometheus::BuildError; @@ -55,7 +54,7 @@ pub enum Error { Unknown, #[error(transparent)] - SendError(#[from] tokio::sync::mpsc::error::SendError), + SendError(#[from] tokio::sync::mpsc::error::SendError), } #[derive(Error, Debug)] diff --git a/packages/cipherstash-proxy/src/main.rs b/packages/cipherstash-proxy/src/main.rs index a11da143d..39f0555fc 100644 --- a/packages/cipherstash-proxy/src/main.rs +++ b/packages/cipherstash-proxy/src/main.rs @@ -1,5 +1,5 @@ use cipherstash_proxy::config::TandemConfig; -use cipherstash_proxy::connect::{self, AsyncStream}; +use cipherstash_proxy::connect; use cipherstash_proxy::error::{ConfigError, Error}; use cipherstash_proxy::prometheus::CLIENTS_ACTIVE_CONNECTIONS; use cipherstash_proxy::proxy::Proxy; @@ -87,7 +87,7 @@ fn main() -> Result<(), Box> { info!(msg = "Received SIGTERM"); break; }, - Ok(client_stream) = AsyncStream::accept(&listener) => { + Ok(client_stream) = connect::accept(&listener) => { client_id += 1; diff --git a/packages/cipherstash-proxy/src/postgresql/context/mod.rs b/packages/cipherstash-proxy/src/postgresql/context/mod.rs index d42e015cf..7382f2c90 100644 --- a/packages/cipherstash-proxy/src/postgresql/context/mod.rs +++ b/packages/cipherstash-proxy/src/postgresql/context/mod.rs @@ -4,11 +4,7 @@ pub mod portal; pub mod statement; pub mod statement_metadata; pub use self::{phase_timing::PhaseTiming, portal::Portal, statement::Statement}; -use super::{ - column_mapper::ColumnMapper, - messages::{describe::Describe, Name, Target}, - Column, -}; +use super::{column_mapper::ColumnMapper, rewrite::Name, Column}; use crate::{ config::TandemConfig, error::{EncryptError, Error}, @@ -22,6 +18,15 @@ use crate::{ use cipherstash_client::IdentifiedBy; use eql_mapper::{Schema, TableResolver}; use metrics::{counter, histogram}; +use pg_proto::{ + codec::{BackendMessage, Describe, DescribeTarget, FrontendMessage}, + intermediary::Intermediary, + middleware::{Identity, Middleware}, + pipeline::{ + BackendAction, BoundedPipeline, FrontendAction, FrontendHandling, FrontendProjectionError, + OperationId, PipelineMiddlewareError, + }, +}; use serde_json::json; use sqltk::parser::ast::{Expr, Ident, ObjectName, ObjectNamePart, Set, Value, ValueWithSpan}; pub use statement_metadata::StatementMetadata; @@ -33,7 +38,7 @@ use std::{ }, time::{Duration, Instant}, }; -use tokio::sync::oneshot; +use tokio::sync::{oneshot, Mutex, Notify}; use tracing::{debug, error, warn}; use uuid::Uuid; @@ -42,6 +47,25 @@ type ExecuteQueue = Queue; type SessionMetricsQueue = Queue; type PortalQueue = Queue>; +fn protocol_transition_error(error: impl std::fmt::Debug) -> Error { + std::io::Error::new(std::io::ErrorKind::InvalidData, format!("{error:?}")).into() +} + +#[derive(Debug)] +struct ProtocolState { + sides: Intermediary<(), (), BoundedPipeline>, +} + +impl ProtocolState { + fn new() -> Self { + Self { + sides: Intermediary::new((), ()).with_pipeline( + BoundedPipeline::new(256).expect("non-zero PostgreSQL pipeline limit"), + ), + } + } +} + #[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] pub struct SessionId(u64); @@ -76,6 +100,8 @@ where unsafe_disable_mapping: bool, keyset_id: Arc>>, session_id_counter: Arc, + protocol: Arc>, + protocol_changed: Arc, } /// Context for tracking an in-flight Execute operation. @@ -197,9 +223,122 @@ where unsafe_disable_mapping: false, keyset_id: Arc::new(RwLock::new(None)), session_id_counter: Arc::new(AtomicU64::new(1)), + protocol: Arc::new(Mutex::new(ProtocolState::new())), + protocol_changed: Arc::new(Notify::new()), + } + } + + /// Records a message accepted from the downstream client. The downstream + /// server-role state advances even when proxy policy intercepts the message. + pub async fn protocol_frontend_received( + &self, + mut message: FrontendMessage, + handling: FrontendHandling, + ) -> Result { + loop { + let notified = self.protocol_changed.notified(); + tokio::pin!(notified); + notified.as_mut().enable(); + let admission = { + let mut protocol = self.protocol.lock().await; + let mut middleware = Middleware::new((), Identity); + match protocol + .sides + .pipeline_mut() + .accept_frontend_typed(&mut middleware, message, handling) + .await + { + Ok(admission) => Ok(admission.into_action()), + Err(PipelineMiddlewareError::Projection( + FrontendProjectionError::Capacity(returned), + )) => Err(*returned), + Err(error) => return Err(protocol_transition_error(error)), + } + }; + match admission { + Ok(FrontendAction::Forward { id, .. } | FrontendAction::Discard { id }) => { + return Ok(id); + } + Err(returned) => { + message = returned; + notified.await; + } + } } } + pub async fn protocol_in_extended_error(&self) -> bool { + let protocol = self.protocol.lock().await; + matches!( + protocol.sides.pipeline().state(), + pg_proto::pipeline::PipelineState::ExtendedError + ) + } + + /// Records a message accepted from the upstream database. + /// Records a backend response actually emitted to the downstream client. + pub async fn protocol_backend_forwarded( + &self, + mut message: BackendMessage, + ) -> Result<(), Error> { + loop { + let notified = self.protocol_changed.notified(); + tokio::pin!(notified); + notified.as_mut().enable(); + let action = { + let mut protocol = self.protocol.lock().await; + let mut middleware = Middleware::new((), Identity); + protocol + .sides + .pipeline_mut() + .accept_backend_typed(&mut middleware, message) + .await + .map_err(protocol_transition_error)? + }; + match action { + BackendAction::Emit(_) => return Ok(()), + BackendAction::Deferred(returned) => { + message = returned; + notified.await; + } + } + } + } + + /// Emits a response synthesized for one locally handled frontend operation. + pub async fn protocol_backend_local( + &self, + id: OperationId, + mut message: BackendMessage, + ) -> Result<(), Error> { + loop { + let notified = self.protocol_changed.notified(); + tokio::pin!(notified); + notified.as_mut().enable(); + let action = { + let mut protocol = self.protocol.lock().await; + let mut middleware = Middleware::new((), Identity); + protocol + .sides + .pipeline_mut() + .try_emit_local_typed(&mut middleware, id, message) + .await + .map_err(protocol_transition_error)? + }; + match action { + BackendAction::Emit(_) => return Ok(()), + BackendAction::Deferred(returned) => { + message = returned; + notified.await; + } + } + } + } + + pub fn protocol_backend_sent(&self) { + self.protocol_changed.notify_waiters(); + } + pub fn set_describe(&mut self, describe: Describe) { debug!(target: CONTEXT, client_id = self.client_id, describe = ?describe); let _ = self.describe.write().map(|mut queue| queue.add(describe)); @@ -387,7 +526,7 @@ where ) .record(execute.duration()); - if execute.name.is_unnamed() { + if execute.name.is_empty() { self.close_portal(&execute.name); } } @@ -462,7 +601,7 @@ where warn!( target: CONTEXT, client_id = self.client_id, - prepared_statement = %name.as_str(), + prepared_statement = %String::from_utf8_lossy(name), msg = "Session lookup failed for prepared statement, using latest session" ); } @@ -532,11 +671,11 @@ where match describe { Describe { ref name, - target: Target::Portal, + target: DescribeTarget::Portal, } => self.get_portal_statement(name), Describe { ref name, - target: Target::Statement, + target: DescribeTarget::Statement, } => self.get_statement(name), } } @@ -1056,25 +1195,30 @@ impl Queue { #[cfg(test)] mod tests { - use super::{Context, Describe, KeysetIdentifier, Portal, Statement}; + use super::{Context, Describe, KeysetIdentifier, Portal, ProtocolState, Statement}; use crate::{ config::LogConfig, error::Error, log, - postgresql::{ - messages::{Name, Target}, - Column, - }, + postgresql::{rewrite::Name, Column}, proxy::{EncryptConfig, EncryptionService}, TandemConfig, }; + use bytes::Bytes; use cipherstash_client::IdentifiedBy; use eql_mapper::Schema; + use pg_proto::codec::{ + BackendMessage, Bind, DescribeTarget, DiagnosticResponse, Execute, FrontendMessage, Parse, + TransactionStatus, + }; + use pg_proto::middleware::{Identity, Middleware}; + use pg_proto::pipeline::{BackendAction, FrontendAction, FrontendHandling, PipelineState}; use sqltk::parser::{dialect::PostgreSqlDialect, parser::Parser}; use std::sync::Arc; use tokio::sync::mpsc; use uuid::Uuid; + #[derive(Clone)] struct TestService {} #[async_trait::async_trait] @@ -1117,6 +1261,129 @@ mod tests { ) } + #[tokio::test] + async fn pg_proto_ledger_tracks_pipelined_extended_messages_in_processing_order() { + let mut protocol = ProtocolState::new(); + let mut middleware = Middleware::new((), Identity); + let messages = [ + FrontendMessage::Parse(Parse { + statement: Bytes::new(), + query: Bytes::from_static(b"select $1"), + parameter_types: vec![23], + }), + FrontendMessage::Bind(Bind { + portal: Bytes::new(), + statement: Bytes::new(), + parameter_formats: vec![0], + parameters: vec![Some(Bytes::from_static(b"1"))], + result_formats: vec![0], + }), + FrontendMessage::Execute(Execute { + portal: Bytes::new(), + max_rows: 0, + }), + FrontendMessage::Sync, + ]; + for message in messages { + assert!(matches!( + protocol + .sides + .pipeline_mut() + .accept_frontend_typed(&mut middleware, message, FrontendHandling::Forward,) + .await + .unwrap(), + pg_proto::pipeline::FrontendAdmission::Immediate(FrontendAction::Forward { .. }) + | pg_proto::pipeline::FrontendAdmission::Waiting( + FrontendAction::Forward { .. } + ) + )); + } + + for response in [ + BackendMessage::ParseComplete, + BackendMessage::BindComplete, + BackendMessage::CommandComplete(Bytes::from_static(b"SELECT 1")), + BackendMessage::ReadyForQuery(TransactionStatus::Idle), + ] { + assert!(matches!( + protocol + .sides + .pipeline_mut() + .accept_backend_typed(&mut middleware, response) + .await + .unwrap(), + BackendAction::Emit(_) + )); + } + + assert!(protocol.sides.pipeline().is_empty()); + assert_eq!(protocol.sides.pipeline().state(), PipelineState::Ready); + } + + #[tokio::test] + async fn local_extended_error_waits_for_earlier_forwarded_response() { + let context = create_context(); + context + .protocol_frontend_received( + FrontendMessage::Parse(Parse { + statement: Bytes::new(), + query: Bytes::from_static(b"select $1"), + parameter_types: vec![23], + }), + FrontendHandling::Forward, + ) + .await + .unwrap(); + let bind_id = context + .protocol_frontend_received( + FrontendMessage::Bind(Bind { + portal: Bytes::new(), + statement: Bytes::new(), + parameter_formats: vec![0], + parameters: vec![Some(Bytes::from_static(b"1"))], + result_formats: vec![0], + }), + FrontendHandling::Local, + ) + .await + .unwrap(); + + let local_context = context.clone(); + let local_error = tokio::spawn(async move { + local_context + .protocol_backend_local( + bind_id, + BackendMessage::ErrorResponse(DiagnosticResponse { fields: vec![] }), + ) + .await + }); + tokio::task::yield_now().await; + assert!(!local_error.is_finished()); + + context + .protocol_backend_forwarded(BackendMessage::ParseComplete) + .await + .unwrap(); + context.protocol_backend_sent(); + local_error.await.unwrap().unwrap(); + + let sync_id = context + .protocol_frontend_received(FrontendMessage::Sync, FrontendHandling::Local) + .await + .unwrap(); + context + .protocol_backend_local( + sync_id, + BackendMessage::ReadyForQuery(TransactionStatus::Idle), + ) + .await + .unwrap(); + assert_eq!( + context.protocol.lock().await.sides.pipeline().state(), + PipelineState::Ready + ); + } + fn statement() -> Statement { Statement { param_columns: vec![], @@ -1154,7 +1421,7 @@ mod tests { let describe = Describe { name, - target: Target::Statement, + target: DescribeTarget::Statement, }; context.set_describe(describe); @@ -1223,7 +1490,7 @@ mod tests { for _ in 0..1000 { // Frontend: a session + execute are enqueued for every statement. let session_id = context.start_session(); - context.set_execute(Name::unnamed(), Some(session_id)); + context.set_execute(Name::new(), Some(session_id)); // Drain primitives, normally called by the backend on an // execute-terminating message (CommandComplete / ErrorResponse / …). @@ -1290,10 +1557,10 @@ mod tests { let mut context = create_context(); let statement_name_1 = Name::from("statement_1"); - let portal_name_1 = Name::unnamed(); + let portal_name_1 = Name::new(); let statement_name_2 = Name::from("statement_2"); - let portal_name_2 = Name::unnamed(); + let portal_name_2 = Name::new(); let statement_name_3 = Name::from("statement_3"); let portal_name_3 = Name::from("portal_3"); diff --git a/packages/cipherstash-proxy/src/postgresql/data/from_sql.rs b/packages/cipherstash-proxy/src/postgresql/data/from_sql.rs index 1deff1b4d..f426b4876 100644 --- a/packages/cipherstash-proxy/src/postgresql/data/from_sql.rs +++ b/packages/cipherstash-proxy/src/postgresql/data/from_sql.rs @@ -1,7 +1,7 @@ use crate::{ error::{Error, MappingError}, log::ENCODING, - postgresql::{format_code::FormatCode, messages::bind::BindParam}, + postgresql::{format_code::FormatCode, rewrite::bind::BindParam}, }; use bigdecimal::BigDecimal; use bytes::BytesMut; @@ -566,7 +566,7 @@ fn decimal_from_sql( #[cfg(test)] mod binary_json_value_tests { use super::*; - use crate::postgresql::{format_code::FormatCode, messages::bind::BindParam}; + use crate::postgresql::{format_code::FormatCode, rewrite::bind::BindParam}; use bytes::BytesMut; fn binary_param(bytes: &[u8]) -> BindParam { @@ -635,7 +635,7 @@ mod tests { config::LogConfig, log, postgresql::{ - data::bind_param_from_sql, format_code::FormatCode, messages::bind::BindParam, Column, + data::bind_param_from_sql, format_code::FormatCode, rewrite::bind::BindParam, Column, }, Identifier, }; diff --git a/packages/cipherstash-proxy/src/postgresql/messages/error_response.rs b/packages/cipherstash-proxy/src/postgresql/diagnostics.rs similarity index 87% rename from packages/cipherstash-proxy/src/postgresql/messages/error_response.rs rename to packages/cipherstash-proxy/src/postgresql/diagnostics.rs index 7013fe1f4..29558a4be 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/error_response.rs +++ b/packages/cipherstash-proxy/src/postgresql/diagnostics.rs @@ -1,13 +1,9 @@ -use super::BackendCode; -use crate::error::{Error, ProtocolError}; -use crate::postgresql::protocol::BytesMutReadString; -use crate::SIZE_I32; -use bytes::{Buf, BufMut, BytesMut}; +//! CipherStash diagnostic response factories. +use bytes::Bytes; use core::fmt; +use pg_proto::codec::{BackendMessage, DiagnosticField, DiagnosticResponse}; use regex::Regex; -use std::io::Cursor; use std::sync::LazyLock; -use std::{convert::TryFrom, ffi::CString}; /// /// Postgres Error Codes /// https://www.postgresql.org/docs/current/errcodes-appendix.html @@ -60,6 +56,9 @@ pub enum ErrorResponseCode { } impl ErrorResponse { + pub fn into_backend_message(self) -> BackendMessage { + BackendMessage::ErrorResponse(self.into()) + } /// Create a FATAL error response for connection timeout. /// /// Uses PostgreSQL error code 57P05 (idle_session_timeout). While this code @@ -293,6 +292,36 @@ impl ErrorResponse { } } +impl From<&DiagnosticResponse> for ErrorResponse { + fn from(response: &DiagnosticResponse) -> Self { + Self { + fields: response + .fields + .iter() + .map(|field| Field { + code: field.code.into(), + value: String::from_utf8_lossy(&field.value).into_owned(), + }) + .collect(), + } + } +} + +impl From for DiagnosticResponse { + fn from(response: ErrorResponse) -> Self { + Self { + fields: response + .fields + .into_iter() + .map(|field| DiagnosticField { + code: field.code.into(), + value: Bytes::from(field.value), + }) + .collect(), + } + } +} + /// /// Extracts line (if present) from a SQL Parser error message /// @@ -313,73 +342,6 @@ fn extract_position_from_parse_error(error_message: &str) -> Option { .and_then(|c| c.get(1)?.as_str().parse::().ok()) } -impl TryFrom<&BytesMut> for ErrorResponse { - type Error = Error; - - fn try_from(buf: &BytesMut) -> Result { - let mut cursor = Cursor::new(buf); - let code = cursor.get_u8(); - - if BackendCode::from(code) != BackendCode::ErrorResponse { - return Err(ProtocolError::UnexpectedMessageCode { - expected: BackendCode::ErrorResponse.into(), - received: code as char, - } - .into()); - } - - let _len = cursor.get_i32(); - - // The message body consists of one or more identified fields, followed by a zero byte as a terminator. - let mut fields = Vec::new(); - - loop { - let code = cursor.get_u8(); - - // zero byte is terminator - if code == 0 { - break; - } - - let value = cursor.read_string()?; - let field = Field { - code: code.into(), - value, - }; - fields.push(field); - } - - Ok(ErrorResponse { fields }) - } -} - -impl TryFrom for BytesMut { - type Error = Error; - - fn try_from(error_response: ErrorResponse) -> Result { - let mut field_bytes = BytesMut::new(); - - for field in error_response.fields { - let value = CString::new(field.value)?; - let value = value.as_bytes_with_nul(); - - field_bytes.put_u8(field.code.into()); - field_bytes.put_slice(value); - } - field_bytes.put_u8(0); // field terminator - - let mut bytes = BytesMut::new(); - - let len = SIZE_I32 + field_bytes.len(); // len + fields - - bytes.put_u8(BackendCode::ErrorResponse.into()); - bytes.put_i32(len as i32); - bytes.put_slice(&field_bytes); - - Ok(bytes) - } -} - impl fmt::Display for ErrorResponse { fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result { for field in self.fields.iter() { @@ -492,8 +454,10 @@ impl From for ErrorResponseCode { #[cfg(test)] mod tests { use super::ErrorResponseCode; - use crate::postgresql::messages::error_response::ErrorResponse; + use crate::postgresql::diagnostics::ErrorResponse; + use crate::postgresql::test_codec::{decode_backend_frame, encode_backend_message}; use bytes::BytesMut; + use pg_proto::codec::BackendMessage; fn to_message(s: &[u8]) -> BytesMut { BytesMut::from(s) @@ -503,13 +467,17 @@ mod tests { pub fn parse_error_response_message() { let message = to_message(b"E\0\0\0kSERROR\0VERROR\0C26000\0Mprepared statement \"a37\" does not exist\0Fprepare.c\0L454\0RFetchPreparedStatement\0\0Z\0\0\0\x05I"); - let error_response = ErrorResponse::try_from(&message).unwrap(); + let BackendMessage::ErrorResponse(response) = decode_backend_frame(&message).unwrap() + else { + panic!("expected ErrorResponse") + }; + let error_response = ErrorResponse::from(&response); assert_eq!(error_response.fields.len(), 7); // let next = cursor.get_u8() as char; // assert_eq!(next, 'Z'); - let bytes = BytesMut::try_from(error_response).unwrap(); + let bytes = encode_backend_message(&error_response.into_backend_message()).unwrap(); let message = to_message(b"E\0\0\0kSERROR\0VERROR\0C26000\0Mprepared statement \"a37\" does not exist\0Fprepare.c\0L454\0RFetchPreparedStatement\0\0"); assert_eq!(bytes, message); } diff --git a/packages/cipherstash-proxy/src/postgresql/driver.rs b/packages/cipherstash-proxy/src/postgresql/driver.rs new file mode 100644 index 000000000..7476dd0eb --- /dev/null +++ b/packages/cipherstash-proxy/src/postgresql/driver.rs @@ -0,0 +1,770 @@ +use super::middleware::{Backend, BackendDisposition, Frontend, FrontendDisposition}; +use crate::connect::{self, ChannelWriter}; +use crate::error::ConfigError; +use crate::log::AUTHENTICATION; +use crate::postgresql::diagnostics::ErrorResponse; +use crate::prometheus::{ + CLIENTS_BYTES_RECEIVED_TOTAL, SERVER_BYTES_RECEIVED_TOTAL, SERVER_BYTES_SENT_TOTAL, +}; +use crate::proxy::ZeroKms; +use crate::{ + error::{Error, ProtocolError}, + postgresql::context::Context, + tls, TandemConfig, +}; +use bytes::Bytes; +use md5::{Digest, Md5}; +use metrics::counter; +use pg_proto::pre_startup::PreStartupMessage; +use pg_proto::{ + auth::{AuthCompletion, AuthEvent, AuthOffer, AwaitingStartupReady, SaslEvent}, + codec::{ + Backend as BackendDirection, BackendMessage, Frontend as FrontendDirection, FrontendMessage, + }, + middleware::{Identity, Inbound, Middleware, PhaseAssociation, ServerRole, TypedReceiveError}, + net::NetworkStream, + pre_startup::{Negotiation, PreStartup, PreStartupOffer}, + server_auth::{ServerPassword, ServerProtocolOffer}, + startup::ProtocolVersion, + transport::Buffered, + Conn, +}; +use postgres_protocol::authentication::sasl::{ChannelBinding, ScramSha256}; +use rand::Rng; +use std::{ + sync::Arc, + time::{Duration, Instant}, +}; +use tokio::{ + io::{AsyncRead, AsyncWrite, AsyncWriteExt}, + net::TcpStream, + time::timeout, +}; +use tracing::{debug, error, info, warn}; + +const SCRAM_SHA_256_PLUS: &[u8] = b"SCRAM-SHA-256-PLUS"; +const SCRAM_SHA_256: &[u8] = b"SCRAM-SHA-256"; +type PgStream = NetworkStream; +type ProtocolMiddleware = Middleware<(), Identity>; + +fn typed_receive_error( + error: TypedReceiveError, +) -> Error +where + Message: std::fmt::Debug, +{ + match error { + TypedReceiveError::Io(error) => error.into(), + TypedReceiveError::Illegal(message) => std::io::Error::new( + std::io::ErrorKind::InvalidData, + format!("message is illegal in the current PostgreSQL phase: {message:?}"), + ) + .into(), + TypedReceiveError::Middleware(never) => match never {}, + TypedReceiveError::InvalidWire(message) => std::io::Error::new( + std::io::ErrorKind::InvalidData, + format!("middleware produced an invalid PostgreSQL message: {message:?}"), + ) + .into(), + } +} + +async fn receive_pre_startup( + mut conn: Conn, PreStartup>, + middleware: &mut ProtocolMiddleware, + connection_timeout: Option, +) -> Result< + ( + Conn, PreStartup>, + PreStartupMessage, + ), + (Conn, PreStartup>, Error), +> { + let received = match connection_timeout { + Some(duration) => match timeout(duration, conn.receive_pre_startup_typed(middleware)).await + { + Ok(received) => received, + Err(_) => return Err((conn, Error::ConnectionTimeout { duration })), + }, + None => conn.receive_pre_startup_typed(middleware).await, + }; + match received { + Ok(message) => Ok((conn, message.into_wire())), + Err(error) => Err((conn, typed_receive_error(error))), + } +} + +async fn receive_frontend_auth( + mut conn: Conn, ServerPassword>, + middleware: &mut ProtocolMiddleware, + connection_timeout: Option, +) -> Result< + ( + Conn, ServerPassword>, + FrontendMessage, + ), + (Conn, ServerPassword>, Error), +> { + let received = match connection_timeout { + Some(duration) => match timeout(duration, conn.receive_frontend_typed(middleware)).await { + Ok(received) => received, + Err(_) => return Err((conn, Error::ConnectionTimeout { duration })), + }, + None => conn.receive_frontend_typed(middleware).await, + }; + match received { + Ok(message) => Ok((conn, message.into_wire())), + Err(error) => Err((conn, typed_receive_error(error))), + } +} + +async fn receive_backend_conn( + mut conn: Conn, Phase>, + middleware: &mut ProtocolMiddleware, +) -> Result<(Conn, Phase>, BackendMessage), Error> +where + S: AsyncRead + Unpin, + Phase: PhaseAssociation, + >::Message: Into, +{ + let duration = Duration::from_secs(10); + let message = timeout(duration, conn.receive_backend_typed(middleware)) + .await + .map_err(|_| Error::ConnectionTimeout { duration })? + .map_err(typed_receive_error)?; + Ok((conn, message.into())) +} + +async fn authenticate_upstream( + startup: Conn, pg_proto::pre_startup::Startup>, + context: &Context, + middleware: &mut ProtocolMiddleware, +) -> Result, AwaitingStartupReady>, Error> { + let mut auth = startup.authentication(); + let offer = loop { + let (current, message) = receive_backend_conn(auth, middleware).await?; + match current.offer_backend(message) { + Ok(AuthEvent::Authentication(offer)) => break offer, + Ok(AuthEvent::Negotiate { conn, .. }) => auth = conn, + Ok(AuthEvent::Error { conn, .. }) => { + conn.into_transport(); + return Err(ProtocolError::AuthenticationFailed.into()); + } + Err((conn, _message, source)) => { + conn.into_transport(); + return Err(source + .unwrap_or_else(|| { + std::io::Error::new( + std::io::ErrorKind::InvalidData, + "unexpected message during upstream authentication", + ) + }) + .into()); + } + } + }; + + match offer { + AuthOffer::Ok(conn) => Ok(conn), + AuthOffer::Cleartext(conn) => { + let (mut awaiting, frame) = conn.password(context.database_password().as_bytes())?; + awaiting.push_frame(frame)?; + awaiting.flush().await?; + complete_upstream_auth(awaiting, middleware).await + } + AuthOffer::Md5 { conn, salt } => { + let hash = md5_hash( + context.database_username().as_bytes(), + context.database_password().as_bytes(), + &salt, + ); + let (mut awaiting, frame) = conn.password(hash.as_bytes())?; + awaiting.push_frame(frame)?; + awaiting.flush().await?; + complete_upstream_auth(awaiting, middleware).await + } + AuthOffer::Sasl { conn, mechanisms } => { + let mechanism = sasl_mechanism(&mechanisms)?; + let mut scram = ScramSha256::new( + context.database_password().as_bytes(), + match mechanism { + SaslMechanism::ScramSha256 => ChannelBinding::unsupported(), + SaslMechanism::ScramSha256Plus => { + ChannelBinding::tls_server_end_point(conn.tls_server_end_point().to_vec()) + } + }, + ); + let (mut sasl, frame) = match mechanism { + SaslMechanism::ScramSha256 => conn.scram_sha_256(scram.message())?, + SaslMechanism::ScramSha256Plus => conn.scram_sha_256_plus(scram.message())?, + }; + sasl.push_frame(frame)?; + sasl.flush().await?; + + let (sasl, message) = receive_backend_conn(sasl, middleware).await?; + let BackendMessage::Authentication(authentication) = message else { + sasl.into_transport(); + return Err(ProtocolError::UnexpectedStartupMessage.into()); + }; + let SaslEvent::Continue { + conn: challenge, + challenge: server_first, + } = sasl.offer(authentication).map_err(|(conn, _)| { + conn.into_transport(); + ProtocolError::AuthenticationFailed + })? + else { + return Err(ProtocolError::AuthenticationFailed.into()); + }; + scram.update(&server_first)?; + let (mut sasl, frame) = challenge.respond(Bytes::copy_from_slice(scram.message())); + sasl.push_frame(frame)?; + sasl.flush().await?; + + let (sasl, message) = receive_backend_conn(sasl, middleware).await?; + let BackendMessage::Authentication(authentication) = message else { + sasl.into_transport(); + return Err(ProtocolError::UnexpectedStartupMessage.into()); + }; + let SaslEvent::Final { + conn: final_state, + server_final, + } = sasl.offer(authentication).map_err(|(conn, _)| { + conn.into_transport(); + ProtocolError::AuthenticationFailed + })? + else { + return Err(ProtocolError::AuthenticationFailed.into()); + }; + scram.finish(&server_final)?; + complete_upstream_auth(final_state.verified(), middleware).await + } + AuthOffer::Gss(conn) | AuthOffer::Sspi(conn) | AuthOffer::KerberosV5(conn) => { + conn.into_transport(); + Err(ProtocolError::UnsupportedAuthentication { method_code: -1 }.into()) + } + } +} + +async fn complete_upstream_auth( + awaiting: Conn, pg_proto::auth::AwaitingAuthOk>, + middleware: &mut ProtocolMiddleware, +) -> Result, AwaitingStartupReady>, Error> { + let (awaiting, message) = receive_backend_conn(awaiting, middleware).await?; + match awaiting.offer(message) { + Ok(AuthCompletion::Ok(conn)) => Ok(conn), + Ok(AuthCompletion::Error { conn, .. }) => { + conn.into_transport(); + Err(ProtocolError::AuthenticationFailed.into()) + } + Err((conn, _)) => { + conn.into_transport(); + Err(ProtocolError::AuthenticationFailed.into()) + } + } +} + +async fn drain_upstream_startup( + mut database: Conn, AwaitingStartupReady>, + client: &mut Buffered, + middleware: &mut ProtocolMiddleware, +) -> Result, Error> { + loop { + let typed = database + .receive_backend_typed(middleware) + .await + .map_err(typed_receive_error)?; + let message: BackendMessage = typed.into(); + let session_item = database.project_backend(message.clone()); + send_backend_message(client, message).await?; + if let Some(item) = session_item { + match database.offer_ready(item) { + Ok(ready) => return Ok(ready.into_transport()), + Err((conn, _)) => database = conn, + } + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq)] +enum SaslMechanism { + ScramSha256, + ScramSha256Plus, +} +/// Handles one downstream PostgreSQL connection and its paired upstream connection. +/// +/// Negotiation and message validation are delegated to `pg-proto`; this function +/// retains the proxy-specific TLS policy, authentication policy, and forwarding. +pub async fn handler(client_stream: PgStream, context: Context) -> Result<(), Error> { + let mut downstream_middleware = Middleware::new((), Identity); + let mut upstream_middleware = Middleware::new((), Identity); + let mut client_is_tls = client_stream.is_tls(); + let mut client = Conn::new(Buffered::<_, FrontendDirection>::new_frontend( + client_stream, + )); + let client_id = context.client_id; + + // Connect to the database server, using TLS if configured + let stream = connect::connect(&context.database_socket_address()).await?; + let mut database_stream = connect_upstream_tls(stream, context.config()).await?; + info!( + msg = "Client connected", + database = context.database_socket_address(), + client_id = client_id, + ); + + let (client_startup, startup_message) = loop { + let (pre_startup, startup_message) = match receive_pre_startup( + client, + &mut downstream_middleware, + context.connection_timeout(), + ) + .await + { + Ok(result) => result, + Err((conn, err @ Error::ConnectionTimeout { .. })) => { + let mut transport = conn.into_transport(); + send_timeout_error(&mut transport, &err).await; + return Err(err); + } + Err((conn, err)) => { + conn.into_transport(); + return Err(err); + } + }; + + match pre_startup.offer_pre_startup(startup_message) { + PreStartupOffer::Ssl(decision) => { + let mut stream = if context.use_tls() { + let (accepted, reply) = decision.accept_ssl(); + let mut stream = accepted.into_transport().into_inner(); + stream.write_all(&[reply]).await?; + stream + } else { + let (rejected, reply) = decision.reject_ssl(); + let mut stream = rejected.into_transport().into_inner(); + stream.write_all(&[reply]).await?; + stream + }; + if let Some(ref tls) = context.tls_config() { + let tcp_stream = stream + .into_plain() + .map_err(|_| ProtocolError::UnexpectedStartupMessage)?; + let (server_config, leaf) = tls::configure_server_with_leaf(tls)?; + let tls_stream = pg_proto::tls::accept( + tcp_stream, + std::sync::Arc::new(server_config), + &leaf, + ) + .await?; + client_is_tls = true; + stream = NetworkStream::server_tls(tls_stream); + } + client = Conn::new(Buffered::new_frontend(stream)); + } + PreStartupOffer::Cancel { + conn, + process_id, + secret_key, + } => { + conn.into_transport(); + let startup_message = PreStartupMessage::CancelRequest { + process_id, + secret_key, + }; + database_stream + .write_all(&startup_message.to_packet()?) + .await?; + return Err(Error::CancelRequest); + } + PreStartupOffer::Startup { conn, message } => break (conn, message), + PreStartupOffer::Gss(conn) => { + conn.into_transport(); + return Err(ProtocolError::UnexpectedStartupMessage.into()); + } + } + }; + + let (mut database_startup, startup_packet) = + Conn::new(Buffered::<_, BackendDirection>::new(database_stream)) + .startup(&startup_message)?; + database_startup.push_startup_packet(&startup_packet); + database_startup.flush().await?; + let database_startup = + authenticate_upstream(database_startup, &context, &mut upstream_middleware).await?; + + // Proxy -> Client Authentication + // Uses MD5 + // SASL is not supported because I need to RTFM https://datatracker.ietf.org/doc/html/rfc5802 + // + // Proxy -> Send AuthenticationMD5Password + // Client -> Send PasswordMessage + // + let client_validated = + match client_startup.validate_protocol(startup_message, ProtocolVersion::V3_2) { + ServerProtocolOffer::Supported { conn, .. } => conn, + ServerProtocolOffer::Rejected { conn, .. } => { + conn.into_transport(); + return Err(ProtocolError::UnexpectedStartupMessage.into()); + } + }; + + let mut client_stream = { + let salt = generate_md5_password_salt(); + + let username = context.database_username().as_bytes(); + let password = context.database_password(); + + let password = password.as_bytes(); + + let hash = md5_hash(username, password, &salt); + + let (mut password_state, frame) = client_validated.begin_server_auth().request_md5(salt)?; + password_state.push_frame(frame)?; + password_state.flush().await?; + + let connection_timeout = context.connection_timeout(); + let (password_state, message) = match receive_frontend_auth( + password_state, + &mut downstream_middleware, + connection_timeout, + ) + .await + { + Ok(result) => result, + Err((conn, err @ Error::ConnectionTimeout { .. })) => { + let mut transport = conn.into_transport(); + send_timeout_error(&mut transport, &err).await; + return Err(err); + } + Err((conn, err)) => { + conn.into_transport(); + return Err(err); + } + }; + + let (auth_state, password) = + password_state + .receive_password(message) + .map_err(|rejected| { + let (conn, _message) = *rejected; + conn.into_transport(); + ProtocolError::UnexpectedAuthenticationResponse { + expected: "PasswordResponse".into(), + received: -1, + } + })?; + let password = + std::str::from_utf8(&password).map_err(|_| ProtocolError::AuthenticationFailed)?; + + if hash != password { + auth_state.into_transport(); + return Err(ProtocolError::ClientAuthenticationFailed.into()); + } + + debug!(target: AUTHENTICATION, msg = "Client AuthenticationOk"); + let (mut startup_ready, frame) = auth_state.authentication_ok()?; + startup_ready.push_frame(frame)?; + startup_ready.flush().await?; + startup_ready.into_transport() + }; + + if context.require_tls() && !client_is_tls { + let message = ErrorResponse::tls_required(); + send_backend_message(&mut client_stream, message.into_backend_message()).await?; + + error!(msg = "Client must connect with Transport Layer Security (TLS)"); + return Err(ConfigError::TlsRequired.into()); + } + + let database_stream = drain_upstream_startup( + database_startup, + &mut client_stream, + &mut upstream_middleware, + ) + .await?; + + let (client_reader, client_writer) = client_stream.into_inner().split(); + let (server_reader, server_writer) = database_stream.into_inner().split(); + + let channel_writer = ChannelWriter::new(client_writer, client_id); + + let mut client_reader: Buffered<_, FrontendDirection> = Buffered::new_frontend(client_reader); + let mut server_writer: Buffered<_, BackendDirection> = Buffered::new(server_writer); + let mut server_reader: Buffered<_, BackendDirection> = Buffered::new(server_reader); + let mut frontend = Middleware::new( + FrontendDisposition::Forward, + Frontend::new(channel_writer.sender(), context.clone()), + ); + let mut backend = Middleware::new( + BackendDisposition::Emit, + Backend::new(channel_writer.sender(), context.clone()), + ); + + if context.is_passthrough() { + if context.use_structured_logging() { + warn!(msg = "RUNNING IN PASSTHROUGH MODE"); + warn!(msg = "DATA IS NOT PROTECTED WITH ENCRYPTION"); + } else { + warn!(msg = "========================================"); + warn!(msg = "RUNNING IN PASSTHROUGH MODE"); + warn!(msg = "DATA IS NOT PROTECTED WITH ENCRYPTION"); + warn!(msg = "========================================"); + } + } + + let timeout_sender = channel_writer.sender(); + let channel_writer_task = tokio::spawn(channel_writer.receive()); + let client_context = context.clone(); + let mut backend_context = context.clone(); + let mut server_write_context = context.clone(); + + let client_to_server = async { + loop { + let message = match receive_frontend_runtime( + &mut client_reader, + client_context.connection_timeout(), + ) + .await + { + Ok(message) => message, + Err(Error::ConnectionClosed) => { + write_to_server( + &mut server_writer, + &mut server_write_context, + FrontendMessage::Terminate, + ) + .await?; + return Ok::<(), Error>(()); + } + Err(error) => return Err(error), + }; + + let frame = message.to_frame()?; + counter!(CLIENTS_BYTES_RECEIVED_TOTAL).increment((frame.body.len() + 5) as u64); + + let tracking_message = message.clone(); + *frontend.state_mut() = FrontendDisposition::Forward; + let outbound = frontend.intercept(message).await?; + if *frontend.state() == FrontendDisposition::Forward { + client_context + .protocol_frontend_received( + tracking_message, + pg_proto::pipeline::FrontendHandling::Forward, + ) + .await?; + write_to_server(&mut server_writer, &mut server_write_context, outbound).await?; + } + } + }; + + let server_to_client = async { + loop { + let read_start = Instant::now(); + let message = + receive_backend_runtime(&mut server_reader, backend_context.connection_timeout()) + .await?; + if server_reader.project_backend(message.clone()).is_none() { + let _ = server_reader.demux_mut().pop_async_event(); + } + let read_duration = read_start.elapsed(); + backend_context.record_execute_server_timing(read_duration); + let frame = message.to_frame()?; + counter!(SERVER_BYTES_RECEIVED_TOTAL).increment((frame.body.len() + 5) as u64); + if read_duration > backend_context.slow_db_response_min_duration() { + warn!( + client_id = backend_context.client_id, + msg = "Slow database response", + duration_ms = read_duration.as_millis(), + message = ?message, + ); + } + + *backend.state_mut() = BackendDisposition::Emit; + let outbound = backend.intercept(message).await?; + if *backend.state() == BackendDisposition::Emit { + backend.handler_mut().write_with_flush(outbound).await?; + } + } + #[allow(unreachable_code)] + Ok::<(), Error>(()) + }; + + // Run frontend and backend tasks + let result = tokio::try_join!(client_to_server, server_to_client); + + if let Err(ref err @ Error::ConnectionTimeout { .. }) = &result { + let error_response = ErrorResponse::connection_timeout(err.to_string()); + let _ = timeout_sender.send(error_response.into_backend_message()); + // Best-effort yield to allow ChannelWriter to flush the error response + // before the connection tears down. Not guaranteed — if the runtime doesn't + // schedule the writer task before teardown, the client may see a connection + // reset instead of the ErrorResponse. + tokio::task::yield_now().await; + } + + // Drop frontend and backend to drop their senders and close the channel + // The async blocks above captured frontend/backend by reference, so they're still alive + drop(frontend); + drop(backend); + + // Wait for channel writer to finish shutdown sequence + // The senders are now dropped, which closes the channel and allows + // the writer task to complete its shutdown + if let Err(err) = channel_writer_task.await { + error!( + client_id, + msg = "Channel writer task panicked", + error = ?err + ); + } + + result?; + Ok(()) +} + +/// Applies CipherStash's upstream TLS policy while pg-proto owns negotiation. +async fn connect_upstream_tls( + stream: NetworkStream, + config: &TandemConfig, +) -> Result, Error> { + if config.database_tls_disabled() { + warn!(msg = "Connecting to database without Transport Layer Security (TLS)"); + return Ok(stream); + } + + let mut request = Conn::new(Buffered::<_, BackendDirection>::new(stream)).request_ssl(); + request.flush().await?; + let mut middleware = Middleware::new((), Identity); + let reply = request + .receive_encryption_reply_typed(&mut middleware) + .await + .map_err(|error| match error { + TypedReceiveError::Io(error) => Error::from(error), + TypedReceiveError::Illegal(reply) | TypedReceiveError::InvalidWire(reply) => { + Error::from(std::io::Error::new( + std::io::ErrorKind::InvalidData, + format!("invalid PostgreSQL encryption reply: {reply:?}"), + )) + } + TypedReceiveError::Middleware(never) => match never {}, + })?; + match request.receive_reply(reply.into_wire()) { + Negotiation::Accepted(conn) => { + let stream = conn + .into_transport() + .into_inner() + .into_plain() + .map_err(|_| ProtocolError::UnexpectedStartupMessage)?; + let tls = pg_proto::tls::connect( + stream, + config.database.server_name()?.to_owned(), + Arc::new(tls::configure_client(&config.database)), + ) + .await?; + Ok(NetworkStream::client_tls(tls)) + } + Negotiation::Rejected(conn) => { + warn!(msg = "Connecting to database without Transport Layer Security (TLS)"); + Ok(conn.into_transport().into_inner()) + } + Negotiation::LegacyError(conn) => { + conn.into_transport(); + Err(ProtocolError::UnexpectedStartupMessage.into()) + } + } +} + +async fn receive_frontend_runtime( + reader: &mut Buffered, + connection_timeout: Option, +) -> Result { + match connection_timeout { + Some(duration) => timeout(duration, reader.receive_wire()) + .await + .map_err(|_| Error::ConnectionTimeout { duration })? + .map_err(Into::into), + None => reader.receive_wire().await.map_err(Into::into), + } +} + +async fn receive_backend_runtime( + reader: &mut Buffered, + connection_timeout: Option, +) -> Result { + match connection_timeout { + Some(duration) => timeout(duration, reader.receive_backend()) + .await + .map_err(|_| Error::ConnectionTimeout { duration })? + .map_err(Into::into), + None => reader.receive_backend().await.map_err(Into::into), + } +} + +async fn write_to_server( + writer: &mut Buffered, + context: &mut Context, + message: FrontendMessage, +) -> Result<(), Error> +where + E: crate::proxy::EncryptionService, +{ + debug!(target: crate::log::PROTOCOL, msg = "Write to server", ?message); + let frame = message.to_frame()?; + counter!(SERVER_BYTES_SENT_TOTAL).increment((frame.body.len() + 5) as u64); + let start = Instant::now(); + writer.push(frame)?; + writer.flush().await?; + if let Some(session_id) = context.latest_session_id() { + context.add_server_write_duration(session_id, start.elapsed()); + } + Ok(()) +} + +fn sasl_mechanism(mechanisms: &[Bytes]) -> Result { + match mechanisms.first().map(Bytes::as_ref) { + Some(SCRAM_SHA_256) => Ok(SaslMechanism::ScramSha256), + Some(SCRAM_SHA_256_PLUS) => Ok(SaslMechanism::ScramSha256Plus), + Some(mechanism) => Err(ProtocolError::UnexpectedSaslAuthenticationMethod( + String::from_utf8_lossy(mechanism).into_owned(), + ) + .into()), + None => Err(ProtocolError::UnexpectedSaslAuthenticationMethod("None".to_string()).into()), + } +} + +pub fn md5_hash(username: &[u8], password: &[u8], salt: &[u8; 4]) -> String { + let mut md5 = Md5::new(); + md5.update(password); + md5.update(username); + let output = md5.finalize_reset(); + md5.update(format!("{output:x}")); + md5.update(salt); + format!("md5{:x}", md5.finalize()) +} + +fn generate_md5_password_salt() -> [u8; 4] { + let mut rng = rand::rng(); + let mut bytes = [0u8; 4]; + rng.fill(&mut bytes); + bytes +} + +/// Best-effort send of a connection timeout ErrorResponse directly to a client stream. +/// Used for pre-split timeout sites where no ChannelWriter exists yet. +async fn send_timeout_error( + stream: &mut Buffered, + err: &Error, +) { + let error_response = ErrorResponse::connection_timeout(err.to_string()); + let _ = send_backend_message(stream, error_response.into_backend_message()).await; +} + +async fn send_backend_message( + stream: &mut Buffered, + message: BackendMessage, +) -> Result<(), Error> { + stream.push(message.to_frame()?)?; + stream.flush().await?; + Ok(()) +} diff --git a/packages/cipherstash-proxy/src/postgresql/error_handler.rs b/packages/cipherstash-proxy/src/postgresql/error_handler.rs index 4de24e03f..367835f38 100644 --- a/packages/cipherstash-proxy/src/postgresql/error_handler.rs +++ b/packages/cipherstash-proxy/src/postgresql/error_handler.rs @@ -6,7 +6,7 @@ use crate::{ connect::Sender, error::{EncryptError, Error, MappingError}, - postgresql::messages::error_response::ErrorResponse, + postgresql::diagnostics::ErrorResponse, }; /// Trait for components that can send PostgreSQL error responses to clients. @@ -57,22 +57,12 @@ pub trait PostgreSqlErrorHandler { _ => ErrorResponse::system_error(err.to_string()), } } - - /// Send an ErrorResponse message to the client. - /// - /// Converts the error to a PostgreSQL ErrorResponse and sends it - /// to the client via the component's sender channel. - /// - /// # Arguments - /// - /// * `error_response` - The ErrorResponse to send to the client - fn send_error_response(&mut self, err: Error) -> Result<(), Error>; } #[cfg(test)] mod tests { use super::*; - use crate::postgresql::messages::error_response::{ + use crate::postgresql::diagnostics::{ ErrorResponseCode, CODE_IDLE_SESSION_TIMEOUT, CODE_SYSTEM_ERROR, }; use std::time::Duration; @@ -88,10 +78,6 @@ mod tests { fn client_id(&self) -> i32 { 0 } - - fn send_error_response(&mut self, _err: Error) -> Result<(), Error> { - unimplemented!("not needed for error_to_response tests") - } } fn error_code(response: &ErrorResponse) -> Option<&str> { diff --git a/packages/cipherstash-proxy/src/postgresql/handler.rs b/packages/cipherstash-proxy/src/postgresql/handler.rs deleted file mode 100644 index c1fe8e50e..000000000 --- a/packages/cipherstash-proxy/src/postgresql/handler.rs +++ /dev/null @@ -1,394 +0,0 @@ -use super::backend::Backend; -use super::frontend::Frontend; -use super::protocol::StartupCode; -use crate::connect::ChannelWriter; -use crate::error::ConfigError; -use crate::log::{AUTHENTICATION, PROTOCOL}; -use crate::postgresql::messages::authentication::auth::{AuthenticationMethod, SaslMechanism}; -use crate::postgresql::messages::authentication::sasl::SASLResponse; -use crate::postgresql::messages::authentication::{ - Authentication, PasswordMessage, SASLInitialResponse, -}; -use crate::postgresql::messages::error_response::ErrorResponse; -use crate::postgresql::{protocol, startup}; -use crate::proxy::ZeroKms; -use crate::{ - connect::AsyncStream, - error::{Error, ProtocolError}, - postgresql::context::Context, - tls, -}; -use bytes::BytesMut; -use md5::{Digest, Md5}; -use postgres_protocol::authentication::sasl::{ChannelBinding, ScramSha256}; -use rand::Rng; -use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt}; -use tracing::{debug, error, info, warn}; -/// -/// -/// Entry point for handling postgres protocol connections -/// Each inbound client connection is mapped to a database connection -/// Hilarity ensues -/// -/// Startup flow -/// -/// Connect to database with TLS if required -/// First message is either: -/// - SSLRequest -/// - ProtocolVersionNumber -/// - CancelRequest -/// -/// On SSLRequest -/// Send SSLResponse -/// Connect with TLS if configured -/// -/// On TLS Connect -/// Expect message containing ProtocolVersionNumber is sent -/// -/// On CancelRequest -/// Propagate and disconnect -/// -/// On ProtocolVersionNumber -/// Propagate and continue -/// -/// -pub async fn handler(client_stream: AsyncStream, context: Context) -> Result<(), Error> { - let mut client_stream = client_stream; - let client_id = context.client_id; - - // Connect to the database server, using TLS if configured - let stream = AsyncStream::connect(&context.database_socket_address()).await?; - let mut database_stream = startup::with_tls(stream, context.config()).await?; - info!( - msg = "Client connected", - database = context.database_socket_address(), - client_id = client_id, - ); - - loop { - let startup_message = - match startup::read_message(&mut client_stream, context.connection_timeout()).await { - Ok(msg) => msg, - Err(err @ Error::ConnectionTimeout { .. }) => { - send_timeout_error(&mut client_stream, &err).await; - return Err(err); - } - Err(err) => return Err(err), - }; - - match &startup_message.code { - StartupCode::SSLRequest => { - startup::send_ssl_response(&mut client_stream, context.use_tls()).await?; - if let Some(ref tls) = context.tls_config() { - match client_stream { - AsyncStream::Tcp(stream) => { - // The Client is connecting to our Server - let tls_stream = tls::server(stream, tls).await?; - client_stream = AsyncStream::Tls(Box::new(tls_stream)); - } - AsyncStream::Tls(_) => { - unreachable!(); - } - } - } - } - StartupCode::CancelRequest => { - database_stream.write_all(&startup_message.bytes).await?; - return Err(Error::CancelRequest); - } - StartupCode::ProtocolVersionNumber => { - database_stream.write_all(&startup_message.bytes).await?; - break; - } - } - } - - // Proxy -> Client Authentication - // Uses MD5 - // SASL is not supported because I need to RTFM https://datatracker.ietf.org/doc/html/rfc5802 - // - // Proxy -> Send AuthenticationMD5Password - // Client -> Send PasswordMessage - // - { - let salt = generate_md5_password_salt(); - - let username = context.database_username().as_bytes(); - let password = context.database_password(); - - let password = password.as_bytes(); - - let hash = md5_hash(username, password, &salt); - - let message = Authentication::md5_password(salt); - let bytes = BytesMut::try_from(message)?; - client_stream.write_all(&bytes).await?; - - let connection_timeout = context.connection_timeout(); - let (_code, bytes) = - match protocol::read_message(&mut client_stream, client_id, connection_timeout).await { - Ok(result) => result, - Err(err @ Error::ConnectionTimeout { .. }) => { - send_timeout_error(&mut client_stream, &err).await; - return Err(err); - } - Err(err) => return Err(err), - }; - - let password_message = PasswordMessage::try_from(&bytes)?; - - if hash == password_message.password { - let message = Authentication::authentication_ok(); - debug!(target: AUTHENTICATION, msg = "Client AuthenticationOk"); - let bytes = BytesMut::try_from(message)?; - client_stream.write_all(&bytes).await?; - } else { - let message = ProtocolError::ClientAuthenticationFailed.to_string(); - error!(msg = message); - - let message = ErrorResponse::invalid_password(message); - let bytes = BytesMut::try_from(message)?; - client_stream.write_all(&bytes).await?; - } - } - - // Database authentication flow - // 1. Database -> Authentication message (SASL) - // -> Proxy -> Auth Reponse flow with SASL - // - // 2. Proxy -> Auth message to the client Md5, SASL etc - // -> Client -> Auth response - // - - // First message should always be Auth - let auth = protocol::read_auth_message(&mut database_stream, client_id).await?; - - match &auth.method { - AuthenticationMethod::AuthenticationOk => { - debug!(target: AUTHENTICATION, msg = "AuthenticationOk"); - } - AuthenticationMethod::AuthenticationCleartextPassword => { - debug!(target: AUTHENTICATION, msg = "AuthenticationCleartextPassword"); - let password = context.database_password(); - let message = PasswordMessage::new(password); - let bytes = BytesMut::try_from(message)?; - database_stream.write_all(&bytes).await?; - } - AuthenticationMethod::Md5Password { salt } => { - debug!(target: AUTHENTICATION, msg = "Md5Password"); - let username = context.database_username().as_bytes(); - let password = context.database_password(); - let password = password.as_bytes(); - - let hash = md5_hash(username, password, salt); - let message = PasswordMessage::new(hash); - let bytes = BytesMut::try_from(message)?; - database_stream.write_all(&bytes).await?; - } - AuthenticationMethod::Sasl { .. } => { - debug!(target: AUTHENTICATION, msg = "Sasl"); - let mechanism = auth.sasl_mechanism()?; - sanity_check_sasl_mechanism(&mechanism, &client_stream); - - // Toby: I don't think we need to do anything here - // If we are connected via TLS, we can support SCRAM-SHA-256-PLUS - // If we are not connected via TLS, the database won't ask for SCRAM-SHA-256-PLUS - let channel_binding = database_stream.channel_binding(); - let password = context.database_password(); - let password = password.as_bytes(); - scram_sha_256_plus_handler(&mut database_stream, mechanism, password, channel_binding) - .await?; - } - AuthenticationMethod::Other { method_code, .. } => { - debug!(target: AUTHENTICATION, msg = "UnsupportedAuthentication"); - return Err(ProtocolError::UnsupportedAuthentication { - method_code: *method_code, - } - .into()); - } - method => { - debug!(target: AUTHENTICATION, msg = "UnexpectedStartupMessage", authentication_method = ?method); - return Err(ProtocolError::UnexpectedStartupMessage.into()); - } - } - - if context.require_tls() && !client_stream.is_tls() { - let message = ErrorResponse::tls_required(); - let bytes = BytesMut::try_from(message)?; - client_stream.write_all(&bytes).await?; - - error!(msg = "Client must connect with Transport Layer Security (TLS)"); - return Err(ConfigError::TlsRequired.into()); - } - - let (client_reader, client_writer) = client_stream.split(); - let (server_reader, server_writer) = database_stream.split(); - - let channel_writer = ChannelWriter::new(client_writer, client_id); - - let mut frontend = Frontend::new( - client_reader, - channel_writer.sender(), - server_writer, - context.clone(), - ); - let mut backend = Backend::new(channel_writer.sender(), server_reader, context.clone()); - - if context.is_passthrough() { - if context.use_structured_logging() { - warn!(msg = "RUNNING IN PASSTHROUGH MODE"); - warn!(msg = "DATA IS NOT PROTECTED WITH ENCRYPTION"); - } else { - warn!(msg = "========================================"); - warn!(msg = "RUNNING IN PASSTHROUGH MODE"); - warn!(msg = "DATA IS NOT PROTECTED WITH ENCRYPTION"); - warn!(msg = "========================================"); - } - } - - let timeout_sender = channel_writer.sender(); - let channel_writer_task = tokio::spawn(channel_writer.receive()); - - let client_to_server = async { - loop { - let result = frontend.rewrite().await; - // Ensure the connection is terminated if the client closes the connection - // The client ConnectionClosed error is triggered before the terminate message is passed through - if matches!(result, Err(Error::ConnectionClosed)) { - frontend.terminate().await? - } - result?; - } - // Unreachable, but helps the compiler understand the return type - // TODO: extract into a function or something with type - #[allow(unreachable_code)] - Ok::<(), Error>(()) - }; - - let server_to_client = async { - loop { - backend.rewrite().await?; - } - #[allow(unreachable_code)] - Ok::<(), Error>(()) - }; - - // Run frontend and backend tasks - let result = tokio::try_join!(client_to_server, server_to_client); - - if let Err(ref err @ Error::ConnectionTimeout { .. }) = &result { - let error_response = ErrorResponse::connection_timeout(err.to_string()); - if let Ok(bytes) = BytesMut::try_from(error_response) { - let _ = timeout_sender.send(bytes); - } - // Best-effort yield to allow ChannelWriter to flush the error response - // before the connection tears down. Not guaranteed — if the runtime doesn't - // schedule the writer task before teardown, the client may see a connection - // reset instead of the ErrorResponse. - tokio::task::yield_now().await; - } - - // Drop frontend and backend to drop their senders and close the channel - // The async blocks above captured frontend/backend by reference, so they're still alive - drop(frontend); - drop(backend); - - // Wait for channel writer to finish shutdown sequence - // The senders are now dropped, which closes the channel and allows - // the writer task to complete its shutdown - if let Err(err) = channel_writer_task.await { - error!( - client_id, - msg = "Channel writer task panicked", - error = ?err - ); - } - - result?; - Ok(()) -} - -// Keep for debugging -fn sanity_check_sasl_mechanism(mechanism: &SaslMechanism, client_stream: &AsyncStream) { - match mechanism { - SaslMechanism::ScramSha256 => { - if client_stream.is_tls() { - debug!( - PROTOCOL, - msg = "Database requested SCRAM-SHA-256, but Proxy has a TLS connection" - ); - } - } - SaslMechanism::ScramSha256Plus => { - if client_stream.is_tcp() { - debug!( - PROTOCOL, - msg = "Database requested SCRAM-SHA-256-PLUS, but Proxy has a TCP connection" - ); - } - } - } -} - -pub fn md5_hash(username: &[u8], password: &[u8], salt: &[u8; 4]) -> String { - let mut md5 = Md5::new(); - md5.update(password); - md5.update(username); - let output = md5.finalize_reset(); - md5.update(format!("{output:x}")); - md5.update(salt); - format!("md5{:x}", md5.finalize()) -} - -fn generate_md5_password_salt() -> [u8; 4] { - let mut rng = rand::rng(); - let mut bytes = [0u8; 4]; - rng.fill(&mut bytes); - bytes -} - -async fn scram_sha_256_plus_handler( - mut stream: S, - mechanism: SaslMechanism, - password: &[u8], - channel_binding: ChannelBinding, -) -> Result<(), Error> { - let mut scram = ScramSha256::new(password, channel_binding); - let bytes = scram.message().to_vec(); - - let sasl_initial_response = SASLInitialResponse::new(mechanism, bytes); - let bytes = BytesMut::try_from(sasl_initial_response)?; - stream.write_all(&bytes).await?; - - let auth = protocol::read_auth_message(&mut stream, 1).await?; - - let bytes = auth.sasl_continue()?; - scram.update(bytes)?; - - let sasl_response = SASLResponse::new(scram.message().to_vec()); - - let bytes = BytesMut::try_from(sasl_response)?; - stream.write_all(&bytes).await?; - - let auth = protocol::read_auth_message(&mut stream, 1).await?; - let bytes = auth.sasl_final()?; - scram.finish(bytes)?; - - let auth = protocol::read_auth_message(&mut stream, 1).await?; - - if auth.is_ok() { - debug!(target: AUTHENTICATION, msg = "SASL authentication successful"); - Ok(()) - } else { - Err(ProtocolError::AuthenticationFailed.into()) - } -} - -/// Best-effort send of a connection timeout ErrorResponse directly to a client stream. -/// Used for pre-split timeout sites where no ChannelWriter exists yet. -async fn send_timeout_error(stream: &mut S, err: &Error) { - let error_response = ErrorResponse::connection_timeout(err.to_string()); - if let Ok(bytes) = BytesMut::try_from(error_response) { - let _ = stream.write_all(&bytes).await; - } -} diff --git a/packages/cipherstash-proxy/src/postgresql/message_buffer.rs b/packages/cipherstash-proxy/src/postgresql/message_buffer.rs deleted file mode 100644 index 750b79b2c..000000000 --- a/packages/cipherstash-proxy/src/postgresql/message_buffer.rs +++ /dev/null @@ -1,42 +0,0 @@ -use super::messages::data_row::DataRow; - -pub struct MessageBuffer { - // buffer: RwLock>, - buffer: Vec, -} - -impl MessageBuffer { - /// Default number of rows to keep in the buffer. - /// Larger rows will require more memory. - const DEFAULT_RESPONSE_BUFFER_SIZE: usize = 4096; - - pub fn new() -> Self { - Self { - buffer: Vec::with_capacity(Self::DEFAULT_RESPONSE_BUFFER_SIZE), - } - } - - pub fn push(&mut self, row: DataRow) { - self.buffer.push(row); - } - - pub fn drain(&mut self) -> Vec { - self.buffer.drain(..).collect() - } - - pub fn clear(&mut self) { - self.buffer.clear(); - } - - pub fn len(&self) -> usize { - self.buffer.len() - } - - pub fn is_empty(&self) -> bool { - self.buffer.is_empty() - } - - pub fn at_capacity(&self) -> bool { - self.buffer.len() >= Self::DEFAULT_RESPONSE_BUFFER_SIZE - } -} diff --git a/packages/cipherstash-proxy/src/postgresql/messages/authentication/auth.rs b/packages/cipherstash-proxy/src/postgresql/messages/authentication/auth.rs deleted file mode 100644 index c96f9e119..000000000 --- a/packages/cipherstash-proxy/src/postgresql/messages/authentication/auth.rs +++ /dev/null @@ -1,358 +0,0 @@ -use crate::error::{Error, ProtocolError}; -use crate::postgresql::messages::{BackendCode, FrontendCode}; -use crate::postgresql::protocol::BytesMutReadString; -use crate::SIZE_I32; -use bytes::{Buf, BufMut, BytesMut}; - -use std::convert::TryFrom; -use std::ffi::CString; -use std::fmt::{self, Display, Formatter}; -use std::io::{Cursor, Read}; - -const SIZE_NULL_BYTE: usize = 1; - -pub const SCRAM_SHA_256_PLUS: &str = "SCRAM-SHA-256-PLUS"; -pub const SCRAM_SHA_256: &str = "SCRAM-SHA-256"; - -#[derive(Debug, Clone, Copy, PartialEq)] -pub enum SaslMechanism { - ScramSha256, - ScramSha256Plus, -} - -#[derive(Debug, Clone)] -pub struct Authentication { - #[allow(dead_code)] - code: u8, - pub method: AuthenticationMethod, -} - -#[derive(Clone, Debug)] -#[repr(i32)] -pub enum AuthenticationMethod { - AuthenticationOk = 0, - AuthenticationCleartextPassword = 3, - Md5Password { salt: [u8; 4] } = 5, - Sasl { mechanisms: Vec } = 10, - AuthenticationSASLContinue { bytes: Vec } = 11, - AuthenticationSASLFinal { bytes: Vec } = 12, - Other { method_code: i32, bytes: Vec }, -} - -#[derive(Clone, Debug)] -pub struct PasswordMessage { - code: u8, - pub password: String, -} - -impl Authentication { - pub fn is_ok(&self) -> bool { - matches!(self.method, AuthenticationMethod::AuthenticationOk) - } - - pub fn is_sasl(&self) -> bool { - matches!(self.method, AuthenticationMethod::Sasl { .. }) - } - - pub fn is_scram_sha_256_plus(&self) -> bool { - match self.method { - AuthenticationMethod::Sasl { ref mechanisms } => { - mechanisms.contains(&SaslMechanism::ScramSha256Plus) - } - _ => false, - } - } - - /// - /// Returns the first mechanism in the list of mechanisms - /// If the method is not SASL, it will return an error - /// If the method is SASL and there are no mechanisms, it will return an error - /// - This should never happen as the server should always return at least one mechanism - /// - If it does, it is a protocol error and the message parse should already have returned an error - /// - This is a safety check to ensure that the server is behaving as expected - /// - pub fn sasl_mechanism(&self) -> Result { - let mechanism = match self.method { - AuthenticationMethod::Sasl { ref mechanisms } => mechanisms.first(), - _ => None, - }; - - match mechanism { - Some(m) => Ok(*m), - None => { - Err(ProtocolError::UnexpectedSaslAuthenticationMethod("None".to_string()).into()) - } - } - } - - pub fn sasl_continue(&self) -> Result<&Vec, Error> { - match self.method { - AuthenticationMethod::AuthenticationSASLContinue { ref bytes } => Ok(bytes), - _ => Err(ProtocolError::UnexpectedAuthenticationResponse { - expected: "SASLContinue".into(), - received: (&self.method).into(), - } - .into()), - } - } - - pub fn sasl_final(&self) -> Result<&Vec, Error> { - match self.method { - AuthenticationMethod::AuthenticationSASLFinal { ref bytes } => Ok(bytes), - _ => Err(ProtocolError::UnexpectedAuthenticationResponse { - expected: "SASLFinal".into(), - received: (&self.method).into(), - } - .into()), - } - } - - pub fn md5_password(salt: [u8; 4]) -> Authentication { - Authentication { - code: BackendCode::Authentication.into(), - method: AuthenticationMethod::Md5Password { salt }, - } - } - - pub fn authentication_ok() -> Authentication { - Authentication { - code: BackendCode::Authentication.into(), - method: AuthenticationMethod::AuthenticationOk, - } - } -} - -impl PasswordMessage { - pub fn new(password: String) -> PasswordMessage { - PasswordMessage { - code: FrontendCode::PasswordMessage.into(), - password, - } - } -} - -impl TryFrom<&BytesMut> for Authentication { - type Error = Error; - - fn try_from(bytes: &BytesMut) -> Result { - let mut cursor = Cursor::new(bytes); - let code = cursor.get_u8(); - - if BackendCode::from(code) != BackendCode::Authentication { - return Err(ProtocolError::UnexpectedMessageCode { - expected: BackendCode::Authentication.into(), - received: code as char, - } - .into()); - } - - let len = cursor.get_i32(); // read and progress cursor - let method_code = cursor.get_i32(); - - let method = match method_code { - 0 => AuthenticationMethod::AuthenticationOk, - 5 => { - let mut salt = [0; 4]; - cursor.read_exact(&mut salt)?; - AuthenticationMethod::Md5Password { salt } - } - 10 => { - let mut mechanisms = Vec::new(); - let mut count = SIZE_I32 // message len - + SIZE_I32 // method_code - + SIZE_NULL_BYTE; // terminating null byte; - - while count < (len as usize) { - let m = cursor.read_string()?; - count += m.len() + SIZE_NULL_BYTE; - mechanisms.push(SaslMechanism::try_from(m)?); - } - AuthenticationMethod::Sasl { mechanisms } - } - 11 => { - let mut bytes = Vec::new(); - cursor.read_to_end(&mut bytes)?; - AuthenticationMethod::AuthenticationSASLContinue { bytes } - } - 12 => { - let mut bytes = Vec::new(); - cursor.read_to_end(&mut bytes)?; - AuthenticationMethod::AuthenticationSASLFinal { bytes } - } - _ => { - // Get any remaining bytes from the cursor - let mut bytes = Vec::new(); - cursor.read_to_end(&mut bytes)?; - AuthenticationMethod::Other { method_code, bytes } - } - }; - - Ok(Authentication { code, method }) - } -} - -impl TryFrom for BytesMut { - type Error = Error; - - fn try_from(auth: Authentication) -> Result { - let mut method_bytes = BytesMut::new(); - - let method_code = (&auth.method).into(); - method_bytes.put_i32(method_code); - - match auth.method { - AuthenticationMethod::AuthenticationOk => {} - AuthenticationMethod::AuthenticationCleartextPassword => {} - AuthenticationMethod::Md5Password { salt } => { - method_bytes.put_slice(&salt); - } - AuthenticationMethod::Sasl { mechanisms } => { - for m in mechanisms { - let s = m.to_string(); - let c = CString::new(s)?; - let s = c.as_bytes_with_nul(); - method_bytes.put_slice(s); - } - method_bytes.put_u8(0); // null byte - } - AuthenticationMethod::AuthenticationSASLContinue { bytes, .. } => { - method_bytes.put_slice(&bytes); - } - AuthenticationMethod::AuthenticationSASLFinal { bytes, .. } => { - method_bytes.put_slice(&bytes); - } - AuthenticationMethod::Other { bytes, .. } => { - method_bytes.put_slice(&bytes); - } - } - - let mut bytes = BytesMut::new(); - - let len = SIZE_I32 // message len - + method_bytes.len(); // method_code - - bytes.put_u8(BackendCode::Authentication.into()); - bytes.put_i32(len as i32); - bytes.put_slice(&method_bytes); - - Ok(bytes) - } -} - -impl TryFrom<&BytesMut> for PasswordMessage { - type Error = Error; - - fn try_from(bytes: &BytesMut) -> Result { - let mut cursor = Cursor::new(bytes); - let code = cursor.get_u8(); - - if FrontendCode::from(code) != FrontendCode::PasswordMessage { - return Err(ProtocolError::UnexpectedMessageCode { - expected: FrontendCode::PasswordMessage.into(), - received: code as char, - } - .into()); - } - - let _len = cursor.get_i32(); // read and progress cursor - let password = cursor.read_string()?; - - Ok(PasswordMessage { code, password }) - } -} - -impl TryFrom for BytesMut { - type Error = Error; - - fn try_from(password_message: PasswordMessage) -> Result { - let mut bytes = BytesMut::new(); - - let password = CString::new(password_message.password)?; - let password = password.as_bytes_with_nul(); - - let len = SIZE_I32 // message len - + password.len(); // password - - bytes.put_u8(FrontendCode::PasswordMessage.into()); - bytes.put_i32(len as i32); - bytes.put_slice(password); - - Ok(bytes) - } -} - -impl TryFrom for SaslMechanism { - type Error = Error; - fn try_from(s: String) -> Result { - match s.as_str() { - SCRAM_SHA_256 => Ok(SaslMechanism::ScramSha256), - SCRAM_SHA_256_PLUS => Ok(SaslMechanism::ScramSha256Plus), - s => Err(ProtocolError::UnexpectedSaslAuthenticationMethod(s.to_owned()).into()), - } - } -} - -impl Display for SaslMechanism { - fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { - let s = match self { - SaslMechanism::ScramSha256 => SCRAM_SHA_256.to_owned(), - SaslMechanism::ScramSha256Plus => SCRAM_SHA_256_PLUS.to_owned(), - }; - write!(f, "{s}") - } -} - -impl From<&AuthenticationMethod> for i32 { - fn from(method: &AuthenticationMethod) -> Self { - match method { - AuthenticationMethod::AuthenticationOk => 0, - AuthenticationMethod::AuthenticationCleartextPassword => 3, - AuthenticationMethod::Md5Password { .. } => 5, - AuthenticationMethod::Sasl { .. } => 10, - AuthenticationMethod::AuthenticationSASLContinue { .. } => 11, - AuthenticationMethod::AuthenticationSASLFinal { .. } => 12, - AuthenticationMethod::Other { method_code, .. } => *method_code, - } - } -} - -#[cfg(test)] -mod tests { - use bytes::BytesMut; - - use crate::{config::LogConfig, log}; - - use super::Authentication; - - fn to_message(s: &[u8]) -> BytesMut { - BytesMut::from(s) - } - - #[test] - pub fn parse_auth_message() { - log::init(LogConfig::default()); - - let bytes = to_message(b"R\0\0\0*\0\0\0\nSCRAM-SHA-256-PLUS\0SCRAM-SHA-256\0\0"); - - let auth = Authentication::try_from(&bytes).unwrap(); - - assert!(matches!( - auth.method, - super::AuthenticationMethod::Sasl { .. } - )); - - let auth_bytes = BytesMut::try_from(auth).unwrap(); - - assert_eq!(bytes, auth_bytes); - } - - #[test] - pub fn is_scram_sha_256_plus() { - log::init(LogConfig::default()); - - let bytes = to_message(b"R\0\0\0*\0\0\0\nSCRAM-SHA-256-PLUS\0SCRAM-SHA-256\0\0"); - let auth = Authentication::try_from(&bytes).unwrap(); - - assert!(auth.is_scram_sha_256_plus()); - } -} diff --git a/packages/cipherstash-proxy/src/postgresql/messages/authentication/mod.rs b/packages/cipherstash-proxy/src/postgresql/messages/authentication/mod.rs deleted file mode 100644 index 40ca7809a..000000000 --- a/packages/cipherstash-proxy/src/postgresql/messages/authentication/mod.rs +++ /dev/null @@ -1,6 +0,0 @@ -pub mod auth; -pub mod sasl; - -pub use auth::Authentication; -pub use auth::PasswordMessage; -pub use sasl::SASLInitialResponse; diff --git a/packages/cipherstash-proxy/src/postgresql/messages/authentication/sasl.rs b/packages/cipherstash-proxy/src/postgresql/messages/authentication/sasl.rs deleted file mode 100644 index b5806619f..000000000 --- a/packages/cipherstash-proxy/src/postgresql/messages/authentication/sasl.rs +++ /dev/null @@ -1,133 +0,0 @@ -use std::{ - ffi::CString, - io::{Cursor, Read}, -}; - -use bytes::{Buf, BufMut, BytesMut}; - -use crate::{ - error::{Error, ProtocolError}, - postgresql::{messages::FrontendCode, protocol::BytesMutReadString}, - SIZE_I32, -}; - -use super::auth::{self, SaslMechanism}; - -#[derive(Clone, Debug)] -pub struct SASLInitialResponse { - #[allow(dead_code)] - code: u8, - pub mechanism: String, - pub response: Vec, -} - -#[derive(Clone, Debug)] -pub struct SASLResponse { - #[allow(dead_code)] - code: u8, - response: Vec, -} - -impl SASLInitialResponse { - pub fn new(mechanism: SaslMechanism, response: Vec) -> Self { - let mechanism = mechanism.to_string(); - - SASLInitialResponse { - code: FrontendCode::SASLInitialResponse.into(), - mechanism, - - response, - } - } - - pub fn is_scram_sha_256(&self) -> bool { - self.mechanism == auth::SCRAM_SHA_256 - } - - pub fn is_scram_sha_256_plus(&self) -> bool { - self.mechanism == auth::SCRAM_SHA_256_PLUS - } -} - -impl SASLResponse { - pub fn new(response: Vec) -> Self { - SASLResponse { - code: FrontendCode::SASLResponse.into(), - response, - } - } -} - -impl TryFrom<&BytesMut> for SASLInitialResponse { - type Error = Error; - - fn try_from(bytes: &BytesMut) -> Result { - let mut cursor = Cursor::new(bytes); - let code = cursor.get_u8(); - - // Note: all password messages use the 'p' code - if code != b'p' { - return Err(ProtocolError::UnexpectedMessageCode { - expected: FrontendCode::SASLInitialResponse.into(), - received: code as char, - } - .into()); - } - let _len = cursor.get_i32(); - let mechanism = cursor.read_string()?; - let _response_len = cursor.get_i32(); - let mut bytes = Vec::new(); - cursor.read_to_end(&mut bytes)?; - - Ok(SASLInitialResponse { - code, - mechanism, - response: bytes, - }) - } -} - -impl TryFrom for BytesMut { - type Error = Error; - - fn try_from(response: SASLInitialResponse) -> Result { - let mut bytes = BytesMut::new(); - - let mechanism = CString::new(response.mechanism)?; - let mechanism = mechanism.as_bytes_with_nul(); - - let response_len = response.response.len(); - - let len = SIZE_I32 // len length - + mechanism.len() - + SIZE_I32 // response_len - + response_len; - - bytes.put_u8(FrontendCode::SASLInitialResponse.into()); - bytes.put_i32(len as i32); - bytes.put_slice(mechanism); - bytes.put_i32(response_len as i32); - bytes.put_slice(&response.response); - - Ok(bytes) - } -} - -impl TryFrom for BytesMut { - type Error = Error; - - fn try_from(response: SASLResponse) -> Result { - let mut bytes = BytesMut::new(); - - let response_len = response.response.len(); - - let len = SIZE_I32 // len length - + response_len; - - bytes.put_u8(FrontendCode::SASLResponse.into()); - bytes.put_i32(len as i32); - bytes.put_slice(&response.response); - - Ok(bytes) - } -} diff --git a/packages/cipherstash-proxy/src/postgresql/messages/close.rs b/packages/cipherstash-proxy/src/postgresql/messages/close.rs deleted file mode 100644 index 1a8fb11d2..000000000 --- a/packages/cipherstash-proxy/src/postgresql/messages/close.rs +++ /dev/null @@ -1,144 +0,0 @@ -use crate::error::{Error, ProtocolError}; -use crate::postgresql::protocol::BytesMutReadString; -use crate::{SIZE_I32, SIZE_U8}; - -use bytes::{Buf, BufMut, BytesMut}; -use std::convert::TryFrom; -use std::ffi::CString; -use std::io::Cursor; - -use super::target::Target; -use super::{FrontendCode, Name}; - -/// -/// Close b'C' (Frontend) message. -/// -/// See: -/// -/// Byte1('C') -/// Identifies the message as a Close command. -/// -/// Int32 -/// Length of message contents in bytes, including self. -/// -/// Byte1 -/// 'S' to close a prepared statement; or 'P' to close a portal. -/// -/// String -/// The name of the prepared statement or portal to close (an empty string selects the unnamed prepared statement or portal). - -#[derive(Debug, Clone)] -pub(crate) struct Close { - pub target: Target, - pub name: Name, -} - -impl TryFrom<&BytesMut> for Close { - type Error = Error; - - fn try_from(bytes: &BytesMut) -> Result { - let mut cursor = Cursor::new(bytes); - let code = cursor.get_u8(); - - if FrontendCode::from(code) != FrontendCode::Close { - return Err(ProtocolError::UnexpectedMessageCode { - expected: FrontendCode::Close.into(), - received: code as char, - } - .into()); - } - - let _len = cursor.get_i32(); // read and progress cursor - let target = cursor.get_u8(); - let target = Target::try_from(target)?; - let name = cursor.read_string()?; - let name = Name::from(name); - - Ok(Close { target, name }) - } -} - -impl TryFrom for BytesMut { - type Error = Error; - - fn try_from(close: Close) -> Result { - let mut bytes = BytesMut::new(); - - let name = CString::new(close.name.as_str())?; - let name = name.as_bytes_with_nul(); - - let len = SIZE_I32 + SIZE_U8 + name.len(); - - bytes.put_u8(FrontendCode::Close.into()); - bytes.put_i32(len as i32); - bytes.put_u8(close.target.into()); - bytes.put_slice(name); - - Ok(bytes) - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::{config::LogConfig, log, postgresql::messages::Name}; - use bytes::BytesMut; - use std::convert::TryFrom; - - fn to_message(s: &[u8]) -> BytesMut { - BytesMut::from(s) - } - - #[test] - pub fn test_close_statement() { - log::init(LogConfig::default()); - - // Close unnamed prepared statement: C\0\0\0\x06S\0 - let bytes = to_message(b"C\0\0\0\x06S\0"); - let close = Close::try_from(&bytes).unwrap(); - - assert!(matches!(close.target, Target::Statement)); - assert!(close.name.is_unnamed()); - } - - #[test] - pub fn test_close_portal() { - log::init(LogConfig::default()); - - // Close unnamed portal: C\0\0\0\x06P\0 - let bytes = to_message(b"C\0\0\0\x06P\0"); - let close = Close::try_from(&bytes).unwrap(); - - assert!(matches!(close.target, Target::Portal)); - assert!(close.name.is_unnamed()); - } - - #[test] - pub fn test_close_named_statement() { - log::init(LogConfig::default()); - - // Close named prepared statement "stmt1": C\0\0\0\x0bSstmt1\0 - let bytes = to_message(b"C\0\0\0\x0bSstmt1\0"); - let close = Close::try_from(&bytes).unwrap(); - - assert!(matches!(close.target, Target::Statement)); - assert_eq!(close.name.as_str(), "stmt1"); - assert!(!close.name.is_unnamed()); - } - - #[test] - pub fn test_close_to_bytes() { - log::init(LogConfig::default()); - - let close = Close { - target: Target::Portal, - name: Name::from("portal1"), - }; - - let bytes = BytesMut::try_from(close).unwrap(); - let parsed = Close::try_from(&bytes).unwrap(); - - assert!(matches!(parsed.target, Target::Portal)); - assert_eq!(parsed.name.as_str(), "portal1"); - } -} diff --git a/packages/cipherstash-proxy/src/postgresql/messages/describe.rs b/packages/cipherstash-proxy/src/postgresql/messages/describe.rs deleted file mode 100644 index 6f0bbd8e3..000000000 --- a/packages/cipherstash-proxy/src/postgresql/messages/describe.rs +++ /dev/null @@ -1,79 +0,0 @@ -use crate::error::{Error, ProtocolError}; -use crate::postgresql::protocol::BytesMutReadString; -use crate::{SIZE_I32, SIZE_U8}; - -use bytes::{Buf, BufMut, BytesMut}; -use std::convert::TryFrom; -use std::ffi::CString; -use std::io::Cursor; - -use super::target::Target; -use super::{FrontendCode, Name}; - -/// -/// Describe b'D' (Frontend) message. -/// -/// See: -/// -/// Byte1('D') -/// Identifies the message as a Describe command. -/// -/// Int32 -/// Length of message contents in bytes, including self. -/// -/// Byte1 -/// 'S' to describe a prepared statement; or 'P' to describe a portal. -/// -/// String -/// The name of the prepared statement or portal to describe (an empty string selects the unnamed prepared statement or portal). - -#[derive(Debug, Clone)] -pub struct Describe { - pub target: Target, - pub name: Name, -} - -impl TryFrom<&BytesMut> for Describe { - type Error = Error; - - fn try_from(bytes: &BytesMut) -> Result { - let mut cursor = Cursor::new(bytes); - let code = cursor.get_u8(); - - if FrontendCode::from(code) != FrontendCode::Describe { - return Err(ProtocolError::UnexpectedMessageCode { - expected: FrontendCode::Describe.into(), - received: code as char, - } - .into()); - } - - let _len = cursor.get_i32(); // read and progress cursor - let target = cursor.get_u8(); - let target = Target::try_from(target)?; - let name = cursor.read_string()?; - let name = Name::from(name); - - Ok(Describe { target, name }) - } -} - -impl TryFrom for BytesMut { - type Error = Error; - - fn try_from(describe: Describe) -> Result { - let mut bytes = BytesMut::new(); - - let name = CString::new(describe.name.as_str())?; - let name = name.as_bytes_with_nul(); - - let len = SIZE_I32 + SIZE_U8 + name.len(); - - bytes.put_u8(FrontendCode::Describe.into()); - bytes.put_i32(len as i32); - bytes.put_u8(describe.target.into()); - bytes.put_slice(name); - - Ok(bytes) - } -} diff --git a/packages/cipherstash-proxy/src/postgresql/messages/execute.rs b/packages/cipherstash-proxy/src/postgresql/messages/execute.rs deleted file mode 100644 index f5e17ae29..000000000 --- a/packages/cipherstash-proxy/src/postgresql/messages/execute.rs +++ /dev/null @@ -1,37 +0,0 @@ -use super::{FrontendCode, Name}; -use crate::error::{Error, ProtocolError}; -use crate::postgresql::protocol::BytesMutReadString; -use bytes::{Buf, BytesMut}; -use std::convert::TryFrom; -use std::io::Cursor; - -#[derive(Debug, Clone)] -pub(crate) struct Execute { - pub portal: Name, - pub max_rows: i32, -} - -impl TryFrom<&BytesMut> for Execute { - type Error = Error; - - fn try_from(bytes: &BytesMut) -> Result { - let mut cursor = Cursor::new(bytes); - let code = cursor.get_u8(); - - if FrontendCode::from(code) != FrontendCode::Execute { - return Err(ProtocolError::UnexpectedMessageCode { - expected: FrontendCode::Execute.into(), - received: code as char, - } - .into()); - } - - let _len = cursor.get_i32(); // read and progress cursor - - let portal = cursor.read_string()?; - let portal = Name::from(portal); - let max_rows = cursor.get_i32(); - - Ok(Execute { portal, max_rows }) - } -} diff --git a/packages/cipherstash-proxy/src/postgresql/messages/mod.rs b/packages/cipherstash-proxy/src/postgresql/messages/mod.rs deleted file mode 100644 index b271c001c..000000000 --- a/packages/cipherstash-proxy/src/postgresql/messages/mod.rs +++ /dev/null @@ -1,304 +0,0 @@ -use std::fmt; - -use bytes::BytesMut; - -pub mod authentication; -pub mod bind; -pub mod close; -pub mod data_row; -pub mod describe; -pub mod error_response; -pub mod execute; -pub mod name; -pub mod param_description; -pub mod parse; -pub mod query; -pub mod ready_for_query; -pub mod row_description; -pub mod target; -pub mod terminate; - -// Re-export commonly used types -pub use name::Name; -pub use target::Target; - -pub const NULL: i32 = -1; - -/// PostgreSQL's "unspecified type, infer it" param OID, used in `Parse` and -/// when a param's type is not known to the proxy. -pub const UNSPECIFIED_TYPE_OID: i32 = 0; - -#[derive(Clone, Copy, Debug, PartialEq)] -pub enum FrontendCode { - Bind, - Close, - Describe, - Execute, - Flush, - Parse, - PasswordMessage, - Query, - SASLInitialResponse, - SASLResponse, - Sync, - Terminate, - Unknown(char), -} - -#[derive(Clone, Copy, Debug, PartialEq)] -pub enum BackendCode { - Authentication, - BindComplete, - BackendKeyData, - CloseComplete, - CommandComplete, - CopyBothResponse, - CopyInResponse, - CopyOutResponse, - DataRow, - EmptyQueryResponse, - ErrorResponse, - NoData, - NoticeResponse, - NotificationResponse, - ParameterDescription, - ParameterStatus, - ParseComplete, - PortalSuspended, - ReadyForQuery, - RowDescription, - Unknown(char), -} - -impl From for FrontendCode { - fn from(code: u8) -> Self { - (code as char).into() - } -} - -impl From for FrontendCode { - fn from(code: char) -> Self { - match code { - 'B' => FrontendCode::Bind, - 'C' => FrontendCode::Close, - 'D' => FrontendCode::Describe, - 'E' => FrontendCode::Execute, - 'H' => FrontendCode::Flush, - 'p' => FrontendCode::PasswordMessage, - 'P' => FrontendCode::Parse, - 'Q' => FrontendCode::Query, - #[allow(unreachable_patterns)] - 'p' => FrontendCode::SASLInitialResponse, // Uses same char, here for completeness - #[allow(unreachable_patterns)] - 'p' => FrontendCode::SASLResponse, // Uses same char, here for completeness - 'S' => FrontendCode::Sync, - 'X' => FrontendCode::Terminate, - _ => FrontendCode::Unknown(code), - } - } -} - -impl From for u8 { - fn from(code: FrontendCode) -> Self { - match code { - FrontendCode::Bind => b'B', - FrontendCode::Close => b'C', - FrontendCode::Describe => b'D', - FrontendCode::Execute => b'E', - FrontendCode::Flush => b'F', - FrontendCode::Parse => b'P', - FrontendCode::PasswordMessage => b'p', - FrontendCode::Query => b'Q', - FrontendCode::SASLInitialResponse => b'p', - FrontendCode::SASLResponse => b'p', - FrontendCode::Sync => b'S', - FrontendCode::Terminate => b'X', - FrontendCode::Unknown(c) => c as u8, - } - } -} - -impl From for char { - fn from(code: FrontendCode) -> Self { - match code { - FrontendCode::Bind => 'B', - FrontendCode::Close => 'C', - FrontendCode::Describe => 'D', - FrontendCode::Execute => 'E', - FrontendCode::Flush => 'F', - FrontendCode::Parse => 'P', - FrontendCode::PasswordMessage => 'p', - FrontendCode::Query => 'Q', - FrontendCode::SASLInitialResponse => 'p', - FrontendCode::SASLResponse => 'p', - FrontendCode::Sync => 'S', - FrontendCode::Terminate => 'X', - FrontendCode::Unknown(c) => c, - } - } -} - -impl From for BackendCode { - fn from(code: u8) -> Self { - match code as char { - 'R' => BackendCode::Authentication, - 'K' => BackendCode::BackendKeyData, - '2' => BackendCode::BindComplete, - '3' => BackendCode::CloseComplete, - 'C' => BackendCode::CommandComplete, - 'W' => BackendCode::CopyBothResponse, - 'G' => BackendCode::CopyInResponse, - 'H' => BackendCode::CopyOutResponse, - 'D' => BackendCode::DataRow, - 'I' => BackendCode::EmptyQueryResponse, - 'E' => BackendCode::ErrorResponse, - 'n' => BackendCode::NoData, - 'N' => BackendCode::NoticeResponse, - 'A' => BackendCode::NotificationResponse, - 't' => BackendCode::ParameterDescription, - 'S' => BackendCode::ParameterStatus, - '1' => BackendCode::ParseComplete, - 's' => BackendCode::PortalSuspended, - 'Z' => BackendCode::ReadyForQuery, - 'T' => BackendCode::RowDescription, - _ => BackendCode::Unknown(code as char), - } - } -} - -impl From for u8 { - fn from(code: BackendCode) -> Self { - match code { - BackendCode::Authentication => b'R', - BackendCode::BackendKeyData => b'K', - BackendCode::BindComplete => b'2', - BackendCode::CloseComplete => b'3', - BackendCode::CommandComplete => b'C', - BackendCode::CopyBothResponse => b'W', - BackendCode::CopyInResponse => b'G', - BackendCode::CopyOutResponse => b'H', - BackendCode::DataRow => b'D', - BackendCode::EmptyQueryResponse => b'I', - BackendCode::ErrorResponse => b'E', - BackendCode::NoData => b'n', - BackendCode::NoticeResponse => b'N', - BackendCode::NotificationResponse => b'A', - BackendCode::ParameterDescription => b't', - BackendCode::ParameterStatus => b'S', - BackendCode::ParseComplete => b'1', - BackendCode::PortalSuspended => b's', - BackendCode::ReadyForQuery => b'Z', - BackendCode::RowDescription => b'T', - BackendCode::Unknown(c) => c as u8, - } - } -} - -impl From for char { - fn from(code: BackendCode) -> Self { - match code { - BackendCode::Authentication => 'R', - BackendCode::BackendKeyData => 'K', - BackendCode::BindComplete => '2', - BackendCode::CloseComplete => '3', - BackendCode::CommandComplete => 'C', - BackendCode::CopyBothResponse => 'W', - BackendCode::CopyInResponse => 'G', - BackendCode::CopyOutResponse => 'H', - BackendCode::DataRow => 'D', - BackendCode::EmptyQueryResponse => 'I', - BackendCode::ErrorResponse => 'E', - BackendCode::NoData => 'n', - BackendCode::NoticeResponse => 'N', - BackendCode::NotificationResponse => 'A', - BackendCode::ParameterDescription => 't', - BackendCode::ParameterStatus => 'S', - BackendCode::ParseComplete => '1', - BackendCode::PortalSuspended => 's', - BackendCode::ReadyForQuery => 'Z', - BackendCode::RowDescription => 'T', - BackendCode::Unknown(c) => c, - } - } -} - -impl fmt::Display for BackendCode { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - match self { - BackendCode::Authentication => write!(f, "BackendCode::Authentication"), - BackendCode::BackendKeyData => write!(f, "BackendCode::BackendKeyData"), - BackendCode::BindComplete => write!(f, "BackendCode::BindComplete"), - BackendCode::CloseComplete => write!(f, "BackendCode::CloseComplete"), - BackendCode::CommandComplete => write!(f, "BackendCode::CommandComplete"), - BackendCode::CopyBothResponse => write!(f, "BackendCode::CopyBothResponse"), - BackendCode::CopyInResponse => write!(f, "BackendCode::CopyInResponse"), - BackendCode::CopyOutResponse => write!(f, "BackendCode::CopyOutResponse"), - BackendCode::DataRow => write!(f, "BackendCode::DataRow"), - BackendCode::EmptyQueryResponse => write!(f, "BackendCode::EmptyQueryResponse"), - BackendCode::ErrorResponse => write!(f, "BackendCode::ErrorResponse"), - BackendCode::NoData => write!(f, "BackendCode::NoData"), - BackendCode::NoticeResponse => write!(f, "BackendCode::NoticeResponse"), - BackendCode::NotificationResponse => write!(f, "BackendCode::NotificationResponse"), - BackendCode::ParameterDescription => write!(f, "BackendCode::ParameterDescription"), - BackendCode::ParameterStatus => write!(f, "BackendCode::ParameterStatus"), - BackendCode::ParseComplete => write!(f, "BackendCode::ParseComplete"), - BackendCode::PortalSuspended => write!(f, "BackendCode::PortalSuspended"), - BackendCode::ReadyForQuery => write!(f, "BackendCode::ReadyForQuery"), - BackendCode::RowDescription => write!(f, "BackendCode::RowDescription"), - BackendCode::Unknown(c) => write!(f, "BackendCode::Unknown('{}')", c), - } - } -} - -impl fmt::Display for FrontendCode { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - match self { - FrontendCode::Bind => write!(f, "FrontendCode::Bind"), - FrontendCode::Close => write!(f, "FrontendCode::Close"), - FrontendCode::Describe => write!(f, "FrontendCode::Describe"), - FrontendCode::Execute => write!(f, "FrontendCode::Execute"), - FrontendCode::Flush => write!(f, "FrontendCode::Flush"), - FrontendCode::Parse => write!(f, "FrontendCode::Parse"), - FrontendCode::PasswordMessage => write!(f, "FrontendCode::PasswordMessage"), - FrontendCode::Query => write!(f, "FrontendCode::Query"), - FrontendCode::SASLInitialResponse => write!(f, "FrontendCode::SASLInitialResponse"), - FrontendCode::SASLResponse => write!(f, "FrontendCode::SASLResponse"), - FrontendCode::Sync => write!(f, "FrontendCode::Sync"), - FrontendCode::Terminate => write!(f, "FrontendCode::Terminate"), - FrontendCode::Unknown(c) => write!(f, "FrontendCode::Unknown('{}')", c), - } - } -} - -/// -/// Peaks at the first byte char. -/// Assumes that a leading `{` may be a JSON value -/// The Plaintext Payload is always a JSON object so this is a pretty naive approach -/// We are not worried about an exhaustive check here -/// -pub fn maybe_json(bytes: &BytesMut) -> bool { - if bytes.is_empty() { - return false; - } - - let b = bytes.as_ref()[0]; - b == b'{' -} - -/// -/// Postgres binary json is regular json with a leading header byte -/// The header byte is always 1 -/// -pub fn maybe_jsonb(bytes: &BytesMut) -> bool { - // Empty JSONB is at least 3 bytes - // `1{}`` - if bytes.len() <= 3 { - return false; - } - - let b = bytes.as_ref(); - - let header = b[0]; - let first = b[1]; - header == 1 && first == b'{' -} diff --git a/packages/cipherstash-proxy/src/postgresql/messages/name.rs b/packages/cipherstash-proxy/src/postgresql/messages/name.rs deleted file mode 100644 index f4f1b9acc..000000000 --- a/packages/cipherstash-proxy/src/postgresql/messages/name.rs +++ /dev/null @@ -1,50 +0,0 @@ -#[derive(Debug, Clone, Hash, Eq, PartialEq)] -pub enum Name { - Named(String), - Unnamed, -} - -impl Name { - pub fn unnamed() -> Name { - Name::Unnamed - } - - pub fn is_unnamed(&self) -> bool { - matches!(self, Name::Unnamed) - } - - pub fn as_str(&self) -> &str { - match self { - Name::Named(s) => s, - Name::Unnamed => "", - } - } -} - -impl std::ops::Deref for Name { - type Target = str; - - fn deref(&self) -> &str { - self.as_str() - } -} - -impl From for Name { - fn from(s: String) -> Self { - if s.is_empty() { - Name::Unnamed - } else { - Name::Named(s) - } - } -} - -impl From<&str> for Name { - fn from(s: &str) -> Self { - if s.is_empty() { - Name::Unnamed - } else { - Name::Named(s.to_string()) - } - } -} diff --git a/packages/cipherstash-proxy/src/postgresql/messages/ready_for_query.rs b/packages/cipherstash-proxy/src/postgresql/messages/ready_for_query.rs deleted file mode 100644 index 324a03d25..000000000 --- a/packages/cipherstash-proxy/src/postgresql/messages/ready_for_query.rs +++ /dev/null @@ -1,21 +0,0 @@ -use crate::{postgresql::messages::BackendCode, SIZE_I32, SIZE_U8}; -use bytes::{BufMut, BytesMut}; - -/// Bind (Z) message. -/// See: -#[derive(Clone, Debug)] -pub struct ReadyForQuery; - -impl From for BytesMut { - fn from(_: ReadyForQuery) -> BytesMut { - let mut bytes = BytesMut::new(); - - let len = SIZE_I32 + SIZE_U8; - - bytes.put_u8(BackendCode::ReadyForQuery.into()); - bytes.put_i32(len as i32); - bytes.put_u8(b'I'); - - bytes - } -} diff --git a/packages/cipherstash-proxy/src/postgresql/messages/target.rs b/packages/cipherstash-proxy/src/postgresql/messages/target.rs deleted file mode 100644 index cf9bfacd2..000000000 --- a/packages/cipherstash-proxy/src/postgresql/messages/target.rs +++ /dev/null @@ -1,46 +0,0 @@ -use crate::error::{Error, ProtocolError}; -use std::convert::TryFrom; - -/// -/// The target of describe or close messages. -/// -/// Valid values are PreparedStatement or Portal -/// -/// A Portal is a parsed statement PLUS any bound parameters -/// Describe with `Target::Portal` returns the RowDescription describing the result set. -/// The assumption is that the parameters are already bound to the portal, so the Describe message is not required to include any parameter information. -/// -/// Calls to Execute are made on a Portal (not a prepared statement) as execute requires any bound parameters -/// -/// A Statement is the parsed statement -/// Describe with `Target::Statement` returns a ParameterDescription followed by the RowDescription. -/// -/// -/// See https://www.postgresql.org/docs/current/protocol-flow.html#PROTOCOL-FLOW-EXT-QUERY -/// -#[derive(Debug, Clone)] -pub enum Target { - Portal, - Statement, -} - -impl TryFrom for Target { - type Error = Error; - - fn try_from(t: u8) -> Result { - match t as char { - 'S' => Ok(Target::Statement), - 'P' => Ok(Target::Portal), - t => Err(ProtocolError::UnexpectedDescribeTarget(t).into()), - } - } -} - -impl From for u8 { - fn from(target: Target) -> u8 { - match target { - Target::Statement => b'S', - Target::Portal => b'P', - } - } -} diff --git a/packages/cipherstash-proxy/src/postgresql/messages/terminate.rs b/packages/cipherstash-proxy/src/postgresql/messages/terminate.rs deleted file mode 100644 index 0024c10ef..000000000 --- a/packages/cipherstash-proxy/src/postgresql/messages/terminate.rs +++ /dev/null @@ -1,15 +0,0 @@ -use super::FrontendCode; -use bytes::{BufMut, BytesMut}; - -pub struct Terminate; - -impl Terminate { - pub fn message() -> BytesMut { - let mut bytes = BytesMut::new(); - - bytes.put_u8(FrontendCode::Terminate.into()); - bytes.put_i32(4); - - bytes - } -} diff --git a/packages/cipherstash-proxy/src/postgresql/backend.rs b/packages/cipherstash-proxy/src/postgresql/middleware/backend.rs similarity index 78% rename from packages/cipherstash-proxy/src/postgresql/backend.rs rename to packages/cipherstash-proxy/src/postgresql/middleware/backend.rs index d22730092..66ac91330 100644 --- a/packages/cipherstash-proxy/src/postgresql/backend.rs +++ b/packages/cipherstash-proxy/src/postgresql/middleware/backend.rs @@ -1,29 +1,26 @@ -use super::context::Context; -use super::data::to_sql; -use super::error_handler::PostgreSqlErrorHandler; -use super::message_buffer::MessageBuffer; -use super::messages::error_response::ErrorResponse; -use super::messages::row_description::RowDescription; -use super::messages::{BackendCode, UNSPECIFIED_TYPE_OID}; -use super::Column; +use super::super::context::Context; +use super::super::data::to_sql; +use super::super::diagnostics::ErrorResponse; +use super::super::error_handler::PostgreSqlErrorHandler; +use super::super::rewrite::row_description::RowDescription; +use super::super::rewrite::UNSPECIFIED_TYPE_OID; +use super::super::Column; use crate::connect::Sender; use crate::error::{EncryptError, Error}; use crate::log::{CONTEXT, DEVELOPMENT, MAPPER, PROTOCOL}; use crate::postgresql::context::Portal; -use crate::postgresql::messages::data_row::DataRow; -use crate::postgresql::messages::param_description::ParamDescription; -use crate::postgresql::protocol::{self}; +use crate::postgresql::rewrite::data_row::DataRow; +use crate::postgresql::rewrite::param_description::ParamDescription; use crate::prometheus::{ CLIENTS_BYTES_SENT_TOTAL, DECRYPTED_VALUES_TOTAL, DECRYPTION_DURATION_SECONDS, DECRYPTION_ERROR_TOTAL, DECRYPTION_REQUESTS_TOTAL, ROWS_ENCRYPTED_TOTAL, - ROWS_PASSTHROUGH_TOTAL, ROWS_TOTAL, SERVER_BYTES_RECEIVED_TOTAL, + ROWS_PASSTHROUGH_TOTAL, ROWS_TOTAL, }; use crate::proxy::EncryptionService; use crate::EqlCiphertext; -use bytes::BytesMut; use metrics::{counter, histogram}; +use pg_proto::{codec::BackendMessage, middleware::MessageMiddleware}; use std::time::Instant; -use tokio::io::AsyncRead; use tracing::{debug, error, info, warn}; /// The PostgreSQL proxy backend that handles server-to-client message processing. @@ -70,26 +67,37 @@ use tracing::{debug, error, info, warn}; /// - `RowDescription`: Result column metadata (modified for encrypted columns) /// - `ParameterDescription`: Parameter metadata (modified for encrypted parameters) /// - `ReadyForQuery`: Session ready state (triggers schema reload if needed) -pub struct Backend -where - R: AsyncRead + Unpin, - S: EncryptionService, -{ +pub struct Backend { /// Sender for outgoing messages to client client_sender: Sender, - /// Reader for incoming messages from server - server_reader: R, /// Session context with portal and statement metadata context: Context, /// Buffer for batching DataRow messages before decryption - buffer: MessageBuffer, + buffer: Vec, } -impl Backend -where - R: AsyncRead + Unpin, - S: EncryptionService, -{ +#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)] +pub(crate) enum BackendDisposition { + #[default] + Emit, + Suppress, +} + +impl MessageMiddleware for Backend { + type Error = Error; + + async fn intercept( + &mut self, + disposition: &mut BackendDisposition, + message: BackendMessage, + ) -> Result { + self.intercept_backend(disposition, message).await + } +} + +impl Backend { + const RESPONSE_BUFFER_SIZE: usize = 4096; + /// Creates a new Backend instance. /// /// # Arguments @@ -98,87 +106,27 @@ where /// * `server_reader` - Stream for reading messages from the PostgreSQL server /// * `encrypt` - Encryption service for handling column decryption /// * `context` - Session context shared with the frontend - pub fn new(client_sender: Sender, server_reader: R, context: Context) -> Self { - let buffer = MessageBuffer::new(); + pub fn new(client_sender: Sender, context: Context) -> Self { + let buffer = Vec::with_capacity(Self::RESPONSE_BUFFER_SIZE); Backend { client_sender, - server_reader, context, buffer, } } - /// Main message processing loop for handling server messages. - /// - /// Reads messages from the PostgreSQL server, processes them based on message type, - /// performs decryption for encrypted result data, and forwards messages to the client. - /// - /// # PostgreSQL Protocol Phases - /// - /// ## Execute Phase - /// Execute operations produce a stream of DataRow messages followed by exactly one of: - /// - `CommandComplete` - Successful completion - /// - `EmptyQueryResponse` - Empty query completed - /// - `ErrorResponse` - Error occurred - /// - `PortalSuspended` - Portal execution suspended (LIMIT reached) - /// - /// ## Describe Phase - /// Describe operations return metadata about statements or portals: - /// - `ParameterDescription` - Parameter metadata (for statements) - /// - `RowDescription` - Result column metadata - /// - `NoData` - No result columns - /// - /// # Message Processing Flow - /// - /// 1. **Read Message**: Read and parse PostgreSQL wire protocol message - /// 2. **Check Passthrough**: Skip processing if encryption is disabled - /// 3. **Handle by Type**: Route to appropriate handler based on message code - /// 4. **Buffer Management**: Buffer DataRows, flush on completion/errors - /// 5. **Forward**: Send processed message to PostgreSQL client - /// - /// # Buffering Behavior - /// - /// DataRow messages are buffered for batch decryption to improve performance. - /// The buffer is automatically flushed when: - /// - Buffer reaches capacity - /// - Execute phase completes (CommandComplete, ErrorResponse, etc.) - /// - Non-DataRow message is encountered - /// - /// # Returns - /// - /// Returns `Ok(())` on successful message processing, or an `Error` if a fatal - /// error occurs that should terminate the connection. - pub async fn rewrite(&mut self) -> Result<(), Error> { - let read_start = Instant::now(); - let (code, mut bytes) = protocol::read_message( - &mut self.server_reader, - self.context.client_id, - self.context.connection_timeout(), - ) - .await?; - let read_duration = read_start.elapsed(); - self.context.record_execute_server_timing(read_duration); - - let sent: u64 = bytes.len() as u64; - counter!(SERVER_BYTES_RECEIVED_TOTAL).increment(sent); - - // Log slow database responses (configurable threshold, default 100ms) - if read_duration > self.context.slow_db_response_min_duration() { - warn!( - client_id = self.context.client_id, - msg = "Slow database response", - duration_ms = read_duration.as_millis(), - message_code = ?code, - ); - } + async fn intercept_backend( + &mut self, + disposition: &mut BackendDisposition, + protocol_message: BackendMessage, + ) -> Result { + let mut outbound_message = protocol_message.clone(); if self.context.is_passthrough() { debug!(target: DEVELOPMENT, client_id = self.context.client_id, msg = "Passthrough enabled" ); - self.write_with_flush(bytes).await?; - // The frontend starts a session and enqueues an execute for every // statement (start_session / set_execute), regardless of whether // the statement is mapped. Those per-connection queues are only @@ -188,60 +136,59 @@ where // otherwise the execute and session_metrics queues grow by one // entry per statement and never shrink, leaking memory until the // process is OOM-killed. See BUG-300. - match code.into() { - BackendCode::CommandComplete - | BackendCode::EmptyQueryResponse - | BackendCode::PortalSuspended - | BackendCode::ErrorResponse => { + match protocol_message { + BackendMessage::CommandComplete(_) + | BackendMessage::EmptyQueryResponse + | BackendMessage::PortalSuspended + | BackendMessage::ErrorResponse(_) => { self.context.complete_execution(); self.context.finish_session(); } _ => {} } - return Ok(()); + return Ok(outbound_message); } let keyset_id = self.context.keyset_identifier(); debug!(target: CONTEXT, client_id = ?self.context.client_id, ?keyset_id); - match code.into() { - BackendCode::DataRow => { + match protocol_message { + BackendMessage::DataRow(row) => { // Encrypted DataRows are added to the buffer and we return early // Otherwise, continue and write - if self.data_row_handler(&bytes).await? { - return Ok(()); + if self.data_row_handler(DataRow::from(row)).await? { + *disposition = BackendDisposition::Suppress; + return Ok(outbound_message); } } // Execute phase is always terminated by the appearance of exactly one of these messages: // CommandComplete, EmptyQueryResponse (if the portal was created from an empty query string), ErrorResponse, or PortalSuspended. - BackendCode::CommandComplete - | BackendCode::EmptyQueryResponse - | BackendCode::PortalSuspended => { + BackendMessage::CommandComplete(_) + | BackendMessage::EmptyQueryResponse + | BackendMessage::PortalSuspended => { debug!(target: PROTOCOL, client_id = self.context.client_id, msg = "CommandComplete | EmptyQueryResponse | PortalSuspended"); match self.flush().await { Ok(_) => (), Err(err) => { warn!(client_id = self.client_id(), error = err.to_string()); - self.send_error_response(err)?; + self.send_error_response(err).await?; } } self.context.complete_execution(); self.context.finish_session(); } - BackendCode::ErrorResponse => { - if let Some(b) = self.error_response_handler(&bytes)? { - bytes = b - } + BackendMessage::ErrorResponse(ref response) => { + self.error_response_handler(response); match self.flush().await { Ok(_) => (), Err(err) => { warn!(client_id = self.client_id(), error = err.to_string()); - self.send_error_response(err)?; + self.send_error_response(err).await?; } } @@ -251,9 +198,12 @@ where // Describe with Target:Statement // Returns a ParameterDescription followed by RowDescription // The Describe is complete after the RowDescription - BackendCode::ParameterDescription => { - if let Some(b) = self.parameter_description_handler(&bytes).await? { - bytes = b + BackendMessage::ParameterDescription(types) => { + if let Some(message) = self + .parameter_description_handler(ParamDescription::from(types)) + .await? + { + outbound_message = message; } } // Describe with Target:Statement or Target::Portal @@ -261,21 +211,24 @@ where // Target::Portal returns a RowDescription // If no rows are returned, NoData is returned instead of a RowDescription // Complete the Describe - BackendCode::RowDescription => { - if let Some(b) = self.row_description_handler(&bytes).await? { - bytes = b + BackendMessage::RowDescription(description) => { + if let Some(message) = self + .row_description_handler(RowDescription::from(description)) + .await? + { + outbound_message = message; } self.context.complete_describe(); } // Describe with Target:Statement or Target::Portal // If the statement returns no rows, NoData is returned instead of a RowDescription - BackendCode::NoData => { + BackendMessage::NoData => { self.context.complete_describe(); } // Reload for SompleQuery flow // Reload is potentially triggered by a FrontEnd Sync message. // However, the SimpleQuery flow does not use Sync so we check here as well - BackendCode::ReadyForQuery => { + BackendMessage::ReadyForQuery(_) => { debug!(target: PROTOCOL, client_id = self.context.client_id, msg = "ReadyForQuery" @@ -285,18 +238,16 @@ where } } - code => { + _ => { debug!(target: PROTOCOL, client_id = self.context.client_id, msg = "Passthrough", - ?code, + message = ?protocol_message, ); } } - self.write_with_flush(bytes).await?; - - Ok(()) + Ok(outbound_message) } /// Handles PostgreSQL ErrorResponse messages from the server. @@ -335,11 +286,10 @@ where /// /// Always returns `Some(bytes)` containing the original error response /// to forward to the client unchanged. - fn error_response_handler(&mut self, bytes: &BytesMut) -> Result, Error> { - let error_response = ErrorResponse::try_from(bytes)?; + fn error_response_handler(&mut self, response: &pg_proto::codec::DiagnosticResponse) { + let error_response = ErrorResponse::from(response); error!(msg = "PostgreSQL Error", error = ?error_response); info!(msg = "PostgreSQL Errors originate in the database"); - Ok(Some(bytes.to_owned())) } /// @@ -351,7 +301,7 @@ where /// async fn buffer(&mut self, data_row: DataRow) -> Result<(), Error> { self.buffer.push(data_row); - if self.buffer.at_capacity() { + if self.buffer.len() >= Self::RESPONSE_BUFFER_SIZE { debug!(target: DEVELOPMENT, client_id = self.context.client_id, msg = "Flush message buffer"); self.flush().await?; } @@ -362,30 +312,35 @@ where /// Write a message to the client /// Flushes all messages in the buffer before writing the message /// - pub async fn write_with_flush(&mut self, bytes: BytesMut) -> Result<(), Error> { + pub async fn write_with_flush(&mut self, message: BackendMessage) -> Result<(), Error> { debug!(target: DEVELOPMENT, client_id = self.context.client_id, msg = "Write"); match self.flush().await { Ok(_) => (), Err(err) => { warn!(client_id = self.client_id(), error = err.to_string()); - self.send_error_response(err)?; + self.send_error_response(err).await?; } } - self.write(bytes).await?; + self.write(message).await?; Ok(()) } /// /// Write a message to the client /// - pub async fn write(&mut self, bytes: BytesMut) -> Result<(), Error> { - let sent: u64 = bytes.len() as u64; + pub async fn write(&mut self, message: BackendMessage) -> Result<(), Error> { + self.context + .protocol_backend_forwarded(message.clone()) + .await?; + let frame = message.to_frame()?; + let sent: u64 = (frame.body.len() + 5) as u64; counter!(CLIENTS_BYTES_SENT_TOTAL).increment(sent); let start = Instant::now(); - self.client_sender.send(bytes)?; + self.client_sender.send(message)?; + self.context.protocol_backend_sent(); let duration = start.elapsed(); self.context.add_client_write_duration_for_execute(duration); @@ -451,7 +406,7 @@ where } }; - let mut rows: Vec = self.buffer.drain().into_iter().collect(); + let mut rows: Vec = self.buffer.drain(..).collect(); debug!(target: DEVELOPMENT, client_id = self.context.client_id, rows = rows.len()); let result_column_count = match rows.first() { @@ -523,8 +478,7 @@ where row.rewrite(&data)?; - let bytes = BytesMut::try_from(row)?; - self.write(bytes).await?; + self.write(BackendMessage::from(row)).await?; } Ok(()) @@ -565,10 +519,8 @@ where async fn parameter_description_handler( &self, - bytes: &BytesMut, - ) -> Result, Error> { - let mut description = ParamDescription::try_from(bytes)?; - + mut description: ParamDescription, + ) -> Result, Error> { debug!(target: PROTOCOL, client_id = self.context.client_id, ParamDescription = ?description); if let Some(statement) = self.context.get_statement_from_describe() { @@ -603,9 +555,9 @@ where } if description.requires_rewrite() { - let bytes = BytesMut::try_from(description)?; - debug!(target: MAPPER, client_id = self.context.client_id, msg = "Rewrite ParamDescription", bytes = ?bytes); - Ok(Some(bytes)) + let message = BackendMessage::from(description); + debug!(target: MAPPER, client_id = self.context.client_id, msg = "Rewrite ParamDescription", ?message); + Ok(Some(message)) } else { Ok(None) } @@ -620,10 +572,8 @@ where /// async fn row_description_handler( &mut self, - bytes: &BytesMut, - ) -> Result, Error> { - let mut description = RowDescription::try_from(bytes)?; - + mut description: RowDescription, + ) -> Result, Error> { debug!(target: PROTOCOL, client_id = self.context.client_id, RowDescription = ?description); if let Some(statement) = self.context.get_statement_for_row_decription() { @@ -639,9 +589,9 @@ where } if description.requires_rewrite() { - let bytes = BytesMut::try_from(description)?; - debug!(target: MAPPER, client_id = self.context.client_id, msg = "Rewrite RowDescription", bytes = ?bytes); - Ok(Some(bytes)) + let message = BackendMessage::from(description); + debug!(target: MAPPER, client_id = self.context.client_id, msg = "Rewrite RowDescription", ?message); + Ok(Some(message)) } else { Ok(None) } @@ -681,13 +631,12 @@ where /// /// Records metrics for both encrypted and passthrough row processing to /// track proxy performance and encryption usage patterns. - async fn data_row_handler(&mut self, bytes: &BytesMut) -> Result { + async fn data_row_handler(&mut self, data_row: DataRow) -> Result { counter!(ROWS_TOTAL).increment(1); match self.context.get_portal_from_execute().as_deref() { Some(Portal::Encrypted { .. }) => { debug!(target: MAPPER, client_id = self.context.client_id, msg = "Encrypted"); - let data_row = DataRow::try_from(bytes)?; self.buffer(data_row).await?; counter!(ROWS_ENCRYPTED_TOTAL).increment(1); @@ -703,11 +652,7 @@ where } /// Implementation of PostgreSQL error handling for the Backend component. -impl PostgreSqlErrorHandler for Backend -where - R: AsyncRead + Unpin, - S: EncryptionService, -{ +impl PostgreSqlErrorHandler for Backend { fn client_sender(&mut self) -> &mut Sender { &mut self.client_sender } @@ -715,18 +660,15 @@ where fn client_id(&self) -> i32 { self.context.client_id } +} - /// Backend-specific error response handling. - /// - /// Unlike the frontend, the backend doesn't need to set an error state - /// since errors during result processing should immediately terminate - /// the current query execution. - fn send_error_response(&mut self, err: Error) -> Result<(), Error> { +impl Backend { + async fn send_error_response(&mut self, err: Error) -> Result<(), Error> { let error_response = self.error_to_response(err); // Ensure any buffered data is cleared before sending error self.buffer.clear(); - let message = BytesMut::try_from(error_response)?; + let message = error_response.into_backend_message(); debug!( target: "PROTOCOL", @@ -735,7 +677,11 @@ where ?message, ); + self.context + .protocol_backend_forwarded(message.clone()) + .await?; self.client_sender.send(message)?; + self.context.protocol_backend_sent(); Ok(()) } @@ -747,10 +693,12 @@ mod tests { use crate::config::{LogConfig, TandemConfig}; use crate::log; use crate::postgresql::context::KeysetIdentifier; - use crate::postgresql::messages::Name; + use crate::postgresql::test_codec::decode_backend_frame; use crate::proxy::{EncryptConfig, EncryptionService}; + use bytes::Bytes as Name; + use bytes::{Bytes, BytesMut}; use eql_mapper::Schema; - use std::io::Cursor; + use pg_proto::codec::FrontendMessage; use std::sync::Arc; use tokio::sync::mpsc; @@ -866,29 +814,55 @@ mod tests { "test context must be in passthrough mode" ); - // A stream of identical terminating messages — one per statement — - // that the backend reads from the "server". - let message = encode(); - let mut server_bytes = BytesMut::new(); - for _ in 0..STATEMENTS { - server_bytes.extend_from_slice(&message); - } + let message = decode_backend_frame(&encode()).unwrap(); // Keep the client receiver alive so write_with_flush succeeds. let (client_sender, _client_receiver) = mpsc::unbounded_channel(); - let reader = Cursor::new(server_bytes.to_vec()); - let mut backend = Backend::new(client_sender, reader, context); + let mut backend = Backend::new(client_sender, context); for i in 0..STATEMENTS { // Frontend: enqueue a session + execute for the statement. let session_id = backend.context.start_session(); + backend.context.set_execute(Name::new(), Some(session_id)); backend .context - .set_execute(Name::unnamed(), Some(session_id)); + .protocol_frontend_received( + FrontendMessage::Execute(pg_proto::codec::Execute { + portal: Bytes::new(), + max_rows: 0, + }), + pg_proto::pipeline::FrontendHandling::Forward, + ) + .await + .unwrap(); + + let mut disposition = BackendDisposition::Emit; + backend + .intercept_backend(&mut disposition, message.clone()) + .await + .unwrap(); + if disposition == BackendDisposition::Emit { + backend.write_with_flush(message.clone()).await.unwrap(); + } - // Backend: process the terminating message via the passthrough - // path, which must drain the queues. - backend.rewrite().await.unwrap(); + if label == "ErrorResponse" { + backend + .context + .protocol_frontend_received( + FrontendMessage::Sync, + pg_proto::pipeline::FrontendHandling::Forward, + ) + .await + .unwrap(); + let ready = + BackendMessage::ReadyForQuery(pg_proto::codec::TransactionStatus::Idle); + backend + .context + .protocol_backend_forwarded(ready) + .await + .unwrap(); + backend.context.protocol_backend_sent(); + } // The queues must be drained every iteration — not grow by one // per statement (the BUG-300 leak). diff --git a/packages/cipherstash-proxy/src/postgresql/frontend.rs b/packages/cipherstash-proxy/src/postgresql/middleware/frontend.rs similarity index 81% rename from packages/cipherstash-proxy/src/postgresql/frontend.rs rename to packages/cipherstash-proxy/src/postgresql/middleware/frontend.rs index 88a3e1f0c..196593265 100644 --- a/packages/cipherstash-proxy/src/postgresql/frontend.rs +++ b/packages/cipherstash-proxy/src/postgresql/middleware/frontend.rs @@ -1,14 +1,10 @@ -use super::context::phase_timing::PhaseTimer; -use super::context::{Context, SessionId, Statement}; -use super::error_handler::PostgreSqlErrorHandler; -use super::messages::bind::Bind; -use super::messages::describe::Describe; -use super::messages::execute::Execute; -use super::messages::parse::Parse; -use super::messages::query::Query; -use super::messages::FrontendCode as Code; -use super::parser::SqlParser; -use super::protocol::{self}; +use super::super::context::phase_timing::PhaseTimer; +use super::super::context::{Context, SessionId, Statement}; +use super::super::error_handler::PostgreSqlErrorHandler; +use super::super::parser::SqlParser; +use super::super::rewrite::bind::Bind; +use super::super::rewrite::parse::Parse; +use super::super::rewrite::query::Query; use crate::connect::Sender; use crate::error::{EncryptError, Error, MappingError}; use crate::log::{MAPPER, PROTOCOL}; @@ -21,31 +17,33 @@ use crate::postgresql::context::Portal; use crate::postgresql::data::{ json_value_selector_plaintext, literal_from_sql, literal_json_value, }; -use crate::postgresql::messages::close::Close; -use crate::postgresql::messages::ready_for_query::ReadyForQuery; -use crate::postgresql::messages::terminate::Terminate; -use crate::postgresql::messages::{Name, Target}; +use crate::postgresql::rewrite::Name; use crate::prometheus::{ - CLIENTS_BYTES_RECEIVED_TOTAL, ENCRYPTED_VALUES_TOTAL, ENCRYPTION_DURATION_SECONDS, - ENCRYPTION_ERROR_TOTAL, ENCRYPTION_REQUESTS_TOTAL, SERVER_BYTES_SENT_TOTAL, - STATEMENTS_ENCRYPTED_TOTAL, STATEMENTS_PASSTHROUGH_MAPPING_DISABLED_TOTAL, - STATEMENTS_PASSTHROUGH_TOTAL, STATEMENTS_UNMAPPABLE_TOTAL, + ENCRYPTED_VALUES_TOTAL, ENCRYPTION_DURATION_SECONDS, ENCRYPTION_ERROR_TOTAL, + ENCRYPTION_REQUESTS_TOTAL, STATEMENTS_ENCRYPTED_TOTAL, + STATEMENTS_PASSTHROUGH_MAPPING_DISABLED_TOTAL, STATEMENTS_PASSTHROUGH_TOTAL, + STATEMENTS_UNMAPPABLE_TOTAL, }; use crate::proxy::EncryptionService; use crate::{EqlOutput, EqlQueryPayload}; -use bytes::BytesMut; use cipherstash_client::encryption::Plaintext; use eql_mapper::{self, EqlMapperError, EqlTermVariant, JsonSelectorSource, TypeCheckedStatement}; use metrics::{counter, histogram}; -use pg_escape::quote_literal; +use pg_proto::{ + codec::{ + BackendMessage, Close, Describe, DescribeTarget, Execute, FrontendMessage, + TransactionStatus, + }, + middleware::MessageMiddleware, + pipeline::{FrontendHandling, OperationId}, +}; use serde::Serialize; use sqltk::parser::ast::{self, Value}; use sqltk::NodeKey; use std::collections::HashMap; use std::sync::Arc; use std::time::Instant; -use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt}; -use tracing::{debug, error, info, warn}; +use tracing::{debug, info, warn}; /// The PostgreSQL proxy frontend that handles client-to-server message processing. /// @@ -90,33 +88,33 @@ use tracing::{debug, error, info, warn}; /// Encryption and mapping errors are converted to appropriate PostgreSQL error responses /// and sent back to the client. The frontend maintains error state to properly handle /// the PostgreSQL extended query error recovery protocol. -pub struct Frontend -where - R: AsyncRead + Unpin, - W: AsyncWrite + Unpin, - S: EncryptionService, -{ - /// Reader for incoming client messages - client_reader: R, +pub struct Frontend { /// Sender for outgoing messages to client client_sender: Sender, - /// Writer for forwarding messages to server - server_writer: W, /// Session context tracking statements, portals, and keyset IDs context: Context, - /// Error state flag for extended query protocol error handling - error_state: Option, } -#[derive(Debug)] -struct ErrorState; +#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)] +pub(crate) enum FrontendDisposition { + #[default] + Forward, + Local, +} -impl Frontend -where - R: AsyncRead + Unpin, - W: AsyncWrite + Unpin, - S: EncryptionService, -{ +impl MessageMiddleware for Frontend { + type Error = Error; + + async fn intercept( + &mut self, + disposition: &mut FrontendDisposition, + message: FrontendMessage, + ) -> Result { + self.intercept_frontend(disposition, message).await + } +} + +impl Frontend { /// Creates a new Frontend instance. /// /// # Arguments @@ -125,79 +123,47 @@ where /// * `client_sender` - Channel sender for sending messages back to client /// * `server_writer` - Stream for writing messages to the PostgreSQL server /// * `context` - Session context for tracking statements and portals with service access - pub fn new( - client_reader: R, - client_sender: Sender, - server_writer: W, - context: Context, - ) -> Self { + pub fn new(client_sender: Sender, context: Context) -> Self { Frontend { - client_reader, client_sender, - server_writer, context, - error_state: None, } } - /// Main message processing loop for handling client messages. - /// - /// Reads a message from the client, processes it based on the PostgreSQL message type, - /// performs any necessary encryption/transformation, and forwards it to the server. - /// - /// # Message Processing Flow - /// - /// 1. **Read Message**: Read and parse the PostgreSQL wire protocol message - /// 2. **Check Mapping**: Skip processing if mapping is disabled - /// 3. **Handle by Type**: Route to appropriate handler based on message type - /// 4. **Error Recovery**: Handle extended query protocol error states - /// 5. **Forward**: Send processed message to PostgreSQL server - /// - /// # Error States - /// - /// When an error occurs during extended query processing, the frontend enters - /// error state and discards messages until a Sync message is received, following - /// the PostgreSQL protocol specification. - /// - /// # Returns - /// - /// Returns `Ok(())` on successful message processing, or an `Error` if a fatal - /// error occurs that should terminate the connection. - pub async fn rewrite(&mut self) -> Result<(), Error> { - let (code, mut bytes) = protocol::read_message( - &mut self.client_reader, - self.context.client_id, - self.context.connection_timeout(), - ) - .await?; - - let sent: u64 = bytes.len() as u64; - counter!(CLIENTS_BYTES_RECEIVED_TOTAL).increment(sent); + async fn intercept_frontend( + &mut self, + disposition: &mut FrontendDisposition, + protocol_message: FrontendMessage, + ) -> Result { + let mut outbound_message = protocol_message.clone(); if self.context.mapping_disabled() { - self.write_to_server(bytes).await?; - return Ok(()); + return Ok(outbound_message); } - let code = Code::from(code); + let recovering_from_extended_error = self.context.protocol_in_extended_error().await; // When an error is detected while processing any extended-query message, the backend issues ErrorResponse, then reads and discards messages until a Sync is reached, // https://www.postgresql.org/docs/current/protocol-flow.html#PROTOCOL-FLOW-EXT-QUERY - if self.error_state.is_some() { + if recovering_from_extended_error { warn!(target: PROTOCOL, client_id = self.context.client_id, - error_state = ?self.error_state, - ?code, + message = ?protocol_message, ); - if code != Code::Sync { - return Ok(()); + if !matches!(protocol_message, FrontendMessage::Sync) { + self.context + .protocol_frontend_received(protocol_message, FrontendHandling::Local) + .await?; + *disposition = FrontendDisposition::Local; + return Ok(outbound_message); } } - match code { - Code::Query => { - match self.query_handler(&bytes).await { - Ok(Some(mapped)) => bytes = mapped, + let tracking_message = protocol_message.clone(); + match protocol_message { + FrontendMessage::Query(query) => { + match self.query_handler(Query::from(query)).await { + Ok(Some(mapped)) => outbound_message = mapped, // No mapping needed, don't change the bytes Ok(None) => (), Err(err) => { @@ -206,21 +172,26 @@ where msg = "Query Handler Error", error = ?err.to_string(), ); - self.send_error_response(err)?; - self.send_ready_for_query()?; - return Ok(()); + let id = self + .context + .protocol_frontend_received(tracking_message, FrontendHandling::Local) + .await?; + self.send_error_response(id, err).await?; + self.send_ready_for_query(id).await?; + *disposition = FrontendDisposition::Local; + return Ok(outbound_message); } } } - Code::Describe => { - self.describe_handler(&bytes).await?; + FrontendMessage::Describe(describe) => { + self.describe_handler(describe).await?; } - Code::Execute => { - self.execute_handler(&bytes).await?; + FrontendMessage::Execute(execute) => { + self.execute_handler(execute).await?; } - Code::Parse => { - match self.parse_handler(&bytes).await { - Ok(Some(mapped)) => bytes = mapped, + FrontendMessage::Parse(parse) => { + match self.parse_handler(Parse::from(parse)).await { + Ok(Some(mapped)) => outbound_message = mapped, // No mapping needed, don't change the bytes Ok(None) => (), Err(err) => { @@ -229,119 +200,108 @@ where msg = "Parse Handler Error", error = ?err.to_string(), ); - self.send_error_response(err)?; - return Ok(()); + let id = self + .context + .protocol_frontend_received(tracking_message, FrontendHandling::Local) + .await?; + self.send_error_response(id, err).await?; + *disposition = FrontendDisposition::Local; + return Ok(outbound_message); } } } - Code::Bind => { - match self.bind_handler(&bytes).await { - Ok(Some(mapped)) => bytes = mapped, + FrontendMessage::Bind(bind) => { + match self.bind_handler(Bind::try_from(bind)?).await { + Ok(Some(mapped)) => outbound_message = mapped, // No mapping needed, don't change the bytes Ok(None) => (), - Err(err) => match err { - Error::Mapping(MappingError::InvalidParameter(_)) => { - warn!(target: PROTOCOL, - client_id = self.context.client_id, - msg = "EncryptError::InvalidParameter", - ); - self.send_error_response(err)?; - return Ok(()); - } - Error::Encrypt(EncryptError::UnknownKeysetIdentifier { .. }) => { - warn!(target: PROTOCOL, - client_id = self.context.client_id, - msg = "EncryptError::UnknownKeysetIdentifier", - ); - self.send_error_response(err)?; - return Ok(()); - } - _ => { - warn!(target: PROTOCOL, - client_id = self.context.client_id, - msg = "Bind Error", - err = err.to_string() - ); - self.send_error_response(err)?; - return Ok(()); + Err(err) => { + let id = self + .context + .protocol_frontend_received(tracking_message, FrontendHandling::Local) + .await?; + match err { + Error::Mapping(MappingError::InvalidParameter(_)) => { + warn!(target: PROTOCOL, + client_id = self.context.client_id, + msg = "EncryptError::InvalidParameter", + ); + self.send_error_response(id, err).await?; + } + Error::Encrypt(EncryptError::UnknownKeysetIdentifier { .. }) => { + warn!(target: PROTOCOL, + client_id = self.context.client_id, + msg = "EncryptError::UnknownKeysetIdentifier", + ); + self.send_error_response(id, err).await?; + } + _ => { + warn!(target: PROTOCOL, + client_id = self.context.client_id, + msg = "Bind Error", + err = err.to_string() + ); + self.send_error_response(id, err).await?; + } } - }, + *disposition = FrontendDisposition::Local; + return Ok(outbound_message); + } } } - Code::Sync => { + FrontendMessage::Sync => { debug!(target: PROTOCOL, client_id = self.context.client_id, - ?code, + message = ?protocol_message, ); self.context.reload_schema_if_changed().await; - if self.error_state.is_some() { + if recovering_from_extended_error { debug!(target: PROTOCOL, client_id = self.context.client_id, msg = "Ready for Query", ); - self.send_ready_for_query()?; - return Ok(()); + let id = self + .context + .protocol_frontend_received(tracking_message, FrontendHandling::Local) + .await?; + self.send_ready_for_query(id).await?; + *disposition = FrontendDisposition::Local; + return Ok(outbound_message); } } - Code::Close => { - self.close_handler(&bytes).await?; + FrontendMessage::Close(close) => { + self.close_handler(close).await?; } - code => { + _ => { debug!(target: PROTOCOL, client_id = self.context.client_id, msg = "Passthrough", - ?code, + message = ?protocol_message, ); } } - self.write_to_server(bytes).await?; - Ok(()) - } - - pub async fn write_to_server(&mut self, bytes: BytesMut) -> Result<(), Error> { - debug!(target: PROTOCOL, msg = "Write to server", ?bytes); - let sent: u64 = bytes.len() as u64; - counter!(SERVER_BYTES_SENT_TOTAL).increment(sent); - - let start = Instant::now(); - self.server_writer.write_all(&bytes).await?; - let duration = start.elapsed(); - if let Some(session_id) = self.context.latest_session_id() { - self.context.add_server_write_duration(session_id, duration); - } - - Ok(()) + Ok(outbound_message) } - pub async fn terminate(&mut self) -> Result<(), Error> { - debug!(target: PROTOCOL, msg = "Terminate server connection"); - let bytes = Terminate::message(); - self.write_to_server(bytes).await?; - Ok(()) - } - - async fn describe_handler(&mut self, bytes: &BytesMut) -> Result<(), Error> { - let describe = Describe::try_from(bytes)?; + async fn describe_handler(&mut self, describe: Describe) -> Result<(), Error> { debug!(target: PROTOCOL, client_id = self.context.client_id, ?describe); self.context.set_describe(describe); Ok(()) } - async fn close_handler(&mut self, bytes: &BytesMut) -> Result<(), Error> { - let close = Close::try_from(bytes)?; + async fn close_handler(&mut self, close: Close) -> Result<(), Error> { debug!(target: PROTOCOL, client_id = self.context.client_id, ?close); match close.target { - Target::Portal => self.context.close_portal(&close.name), - Target::Statement => self.context.close_statement_and_portal(&close.name), + DescribeTarget::Portal => self.context.close_portal(&close.name), + DescribeTarget::Statement => self.context.close_statement_and_portal(&close.name), } Ok(()) } - async fn execute_handler(&mut self, bytes: &BytesMut) -> Result<(), Error> { - let execute = Execute::try_from(bytes)?; + async fn execute_handler(&mut self, execute: Execute) -> Result<(), Error> { debug!(target: PROTOCOL, client_id = self.context.client_id, ?execute); self.context .set_execute_for_portal(execute.portal.to_owned()); @@ -380,7 +340,7 @@ where /// - `Ok(Some(bytes))` - Transformed query that should replace the original /// - `Ok(None)` - No transformation needed, forward original query /// - `Err(error)` - Processing failed, error should be sent to client - async fn query_handler(&mut self, bytes: &BytesMut) -> Result, Error> { + async fn query_handler(&mut self, mut query: Query) -> Result, Error> { let handler_start = Instant::now(); let session_id = self.context.start_session(); @@ -391,8 +351,6 @@ where let parse_timer = PhaseTimer::start(); - let mut query = Query::try_from(bytes)?; - // Simple Query may contain many statements let parsed_statements = SqlParser::parse_statements(&query.statement)?; let mut transformed_statements = vec![]; @@ -525,8 +483,8 @@ where m.set_query_fingerprint(&query.statement); }); - self.context.add_portal(Name::unnamed(), portal); - self.context.set_execute(Name::unnamed(), Some(session_id)); + self.context.add_portal(Name::new(), portal); + self.context.set_execute(Name::new(), Some(session_id)); if encrypted { let transformed_statement = transformed_statements @@ -537,7 +495,7 @@ where query.rewrite(transformed_statement.to_string()); - let bytes = BytesMut::try_from(query)?; + let message = FrontendMessage::from(query); let handler_duration = handler_start.elapsed(); debug!( target: MAPPER, @@ -545,7 +503,7 @@ where msg = "Rewrite Query", transformed_statement = transformed_statement.to_string(), duration_ms = handler_duration.as_millis(), - bytes = ?bytes, + ?message, ); if handler_duration.as_millis() > 100 { warn!( @@ -554,7 +512,7 @@ where duration_ms = handler_duration.as_millis(), ); } - Ok(Some(bytes)) + Ok(Some(message)) } else { let handler_duration = handler_start.elapsed(); if handler_duration.as_millis() > 50 { @@ -740,7 +698,10 @@ where /// - `Ok(Some(bytes))` - Modified Parse message with transformed SQL/parameters /// - `Ok(None)` - No transformation needed, forward original message /// - `Err(error)` - Processing failed, error should be sent to client - async fn parse_handler(&mut self, bytes: &BytesMut) -> Result, Error> { + async fn parse_handler( + &mut self, + mut message: Parse, + ) -> Result, Error> { let session_id = self.context.start_session(); // Set protocol type @@ -750,7 +711,6 @@ where let parse_timer = PhaseTimer::start(); - let mut message = Parse::try_from(bytes)?; self.context .set_statement_session(message.name.to_owned(), session_id); @@ -882,14 +842,14 @@ where }); if message.requires_rewrite() { - let bytes = BytesMut::try_from(message)?; + let message = FrontendMessage::from(message); debug!(target: MAPPER, client_id = self.context.client_id, msg = "Rewrite Parse", - bytes = ?bytes); + ?message); - Ok(Some(bytes)) + Ok(Some(message)) } else { Ok(None) } @@ -1034,7 +994,7 @@ where /// - `Ok(Some(bytes))` - Modified Bind message with encrypted parameter values /// - `Ok(None)` - No parameter encryption needed, forward original message /// - `Err(error)` - Processing failed, error should be sent to client - async fn bind_handler(&mut self, bytes: &BytesMut) -> Result, Error> { + async fn bind_handler(&mut self, mut bind: Bind) -> Result, Error> { if self.context.unsafe_disable_mapping() { warn!(msg = "Encrypted statement mapping is not enabled"); counter!(STATEMENTS_PASSTHROUGH_MAPPING_DISABLED_TOTAL).increment(1); @@ -1042,8 +1002,6 @@ where return Ok(None); } - let mut bind = Bind::try_from(bytes)?; - let session_id = self .context .get_statement_session_or_latest(&bind.prepared_statement); @@ -1079,15 +1037,15 @@ where self.context.add_portal(bind.portal.to_owned(), portal); if bind.requires_rewrite() { - let bytes = BytesMut::try_from(bind)?; + let message = FrontendMessage::from(bind); debug!( target: MAPPER, client_id = self.context.client_id, msg = "Rewrite Bind", - bytes = ?bytes + ?message ); - Ok(Some(bytes)) + Ok(Some(message)) } else { Ok(None) } @@ -1194,8 +1152,8 @@ where /// /// Send an ReadyForQuery to the client and remove error state. /// - fn send_ready_for_query(&mut self) -> Result<(), Error> { - let message = BytesMut::from(ReadyForQuery); + async fn send_ready_for_query(&mut self, id: OperationId) -> Result<(), Error> { + let message = BackendMessage::ReadyForQuery(TransactionStatus::Idle); debug!(target: PROTOCOL, client_id = self.context.client_id, @@ -1203,34 +1161,13 @@ where ?message, ); + self.context + .protocol_backend_local(id, message.clone()) + .await?; self.client_sender.send(message)?; - self.error_state = None; - + self.context.protocol_backend_sent(); Ok(()) } - - /// TODO output err as structured data. - /// err can carry any additional context from caller - fn to_database_exception(&self, err: Error) -> Result { - error!(client_id = self.context.client_id, msg = err.to_string(), error = ?err); - - // This *should* be sufficient for escaping error messages as we're only - // using the string literal, and not identifiers - let quoted_error = quote_literal(format!("{err}").as_str()); - let content = format!("DO $$ BEGIN RAISE EXCEPTION {quoted_error}; END; $$;"); - - debug!( - target: MAPPER, - client_id = self.context.client_id, - msg = "Frontend exception", - error = err.to_string() - ); - - let query = Query::new(content); - let bytes = BytesMut::try_from(query)?; - - Ok(bytes) - } } /// Projects a stored payload into its query operand when the value is bound in @@ -1329,12 +1266,7 @@ where } /// Implementation of PostgreSQL error handling for the Frontend component. -impl PostgreSqlErrorHandler for Frontend -where - R: AsyncRead + Unpin, - W: AsyncWrite + Unpin, - S: EncryptionService, -{ +impl PostgreSqlErrorHandler for Frontend { fn client_sender(&mut self) -> &mut Sender { &mut self.client_sender } @@ -1342,10 +1274,12 @@ where fn client_id(&self) -> i32 { self.context.client_id } +} - fn send_error_response(&mut self, err: Error) -> Result<(), Error> { +impl Frontend { + async fn send_error_response(&mut self, id: OperationId, err: Error) -> Result<(), Error> { let error_response = self.error_to_response(err); - let message = BytesMut::try_from(error_response)?; + let message = error_response.into_backend_message(); debug!(target: PROTOCOL, client_id = self.context.client_id, @@ -1353,9 +1287,11 @@ where ?message, ); + self.context + .protocol_backend_local(id, message.clone()) + .await?; self.client_sender.send(message)?; - self.error_state = Some(ErrorState); // Frontend-specific: set error state for extended query protocol - + self.context.protocol_backend_sent(); Ok(()) } } diff --git a/packages/cipherstash-proxy/src/postgresql/middleware/mod.rs b/packages/cipherstash-proxy/src/postgresql/middleware/mod.rs new file mode 100644 index 000000000..621a6e7fa --- /dev/null +++ b/packages/cipherstash-proxy/src/postgresql/middleware/mod.rs @@ -0,0 +1,5 @@ +mod backend; +mod frontend; + +pub(crate) use backend::{Backend, BackendDisposition}; +pub(crate) use frontend::{Frontend, FrontendDisposition}; diff --git a/packages/cipherstash-proxy/src/postgresql/mod.rs b/packages/cipherstash-proxy/src/postgresql/mod.rs index 71c8df249..12e21feb9 100644 --- a/packages/cipherstash-proxy/src/postgresql/mod.rs +++ b/packages/cipherstash-proxy/src/postgresql/mod.rs @@ -1,28 +1,17 @@ -mod backend; mod column_mapper; mod context; mod data; +mod diagnostics; +mod driver; mod error_handler; mod format_code; -mod frontend; -mod handler; -mod message_buffer; -mod messages; +mod middleware; mod parser; -mod protocol; -mod startup; +mod rewrite; +#[cfg(test)] +mod test_codec; pub use context::column::Column; pub use context::Context; pub use context::KeysetIdentifier; -pub use handler::handler; - -pub const PROTOCOL_VERSION_NUMBER: i32 = 196608; - -pub const SSL_REQUEST: i32 = 80877103; - -pub const CANCEL_REQUEST: i32 = 80877102; - -pub const SSL_RESPONSE_NO: u8 = b'N'; - -pub const SSL_RESPONSE_YES: u8 = b'S'; +pub use driver::handler; diff --git a/packages/cipherstash-proxy/src/postgresql/protocol.rs b/packages/cipherstash-proxy/src/postgresql/protocol.rs deleted file mode 100644 index 9c1afb79b..000000000 --- a/packages/cipherstash-proxy/src/postgresql/protocol.rs +++ /dev/null @@ -1,163 +0,0 @@ -use super::{messages::authentication::Authentication, CANCEL_REQUEST, SSL_REQUEST}; -use crate::{ - error::{Error, ProtocolError}, - log::PROTOCOL, - postgresql::PROTOCOL_VERSION_NUMBER, - SIZE_I32, SIZE_U8, -}; -use bytes::{BufMut, BytesMut}; -use std::{ - io::{BufRead, Cursor}, - time::Duration, -}; -use tokio::{ - io::{AsyncRead, AsyncReadExt}, - time::timeout, -}; -use tracing::{debug, error}; - -type Code = u8; - -#[derive(Clone, Debug, PartialEq)] -pub enum StartupCode { - ProtocolVersionNumber, - CancelRequest, - SSLRequest, -} - -#[derive(Clone, Debug)] -pub struct StartupMessage { - pub code: StartupCode, - pub bytes: BytesMut, -} - -#[derive(Clone, Debug)] -pub struct Message { - pub code: u8, - pub bytes: BytesMut, -} - -impl From for StartupCode { - fn from(code: i32) -> Self { - match code { - PROTOCOL_VERSION_NUMBER => StartupCode::ProtocolVersionNumber, - SSL_REQUEST => StartupCode::SSLRequest, - CANCEL_REQUEST => StartupCode::CancelRequest, - _ => panic!("Unexpected startup code {code}"), - } - } -} - -pub trait BytesMutReadString { - fn read_string(&mut self) -> Result; -} - -impl BytesMutReadString for Cursor<&BytesMut> { - /// Should only be used when reading strings from the message protocol. - /// Can be used to read multiple strings from the same message which are separated by the null byte - fn read_string(&mut self) -> Result { - let mut buf = Vec::with_capacity(512); - match self.read_until(b'\0', &mut buf) { - Ok(_) => Ok(String::from_utf8_lossy(&buf[..buf.len() - 1]).to_string()), - Err(err) => Err(err.into()), - } - } -} - -/// -/// Reads an Auth Message from Stream -/// -/// Does not use the default connection timeout as the auth message is expected to be sent immediately -/// 10 seconds is a reasonable timeout for the auth message -/// -/// -pub async fn read_auth_message( - mut stream: S, - client_id: i32, -) -> Result { - let connection_timeout = Duration::from_millis(1000 * 10); - let (_code, bytes) = - read_message_with_timeout(&mut stream, client_id, connection_timeout).await?; - Authentication::try_from(&bytes) -} - -/// -/// Reads a Postgres message from client with an optional timeout -/// -/// Timeout values are in config -/// -/// -pub async fn read_message( - mut stream: S, - client_id: i32, - connection_timeout: Option, -) -> Result<(Code, BytesMut), Error> { - match connection_timeout { - Some(duration) => read_message_with_timeout(stream, client_id, duration).await, - None => read(&mut stream, client_id).await, - } -} - -/// -/// Reads a Postgres message from client with a timeout -/// -/// Timeout values are in config -/// -/// -async fn read_message_with_timeout( - mut stream: S, - client_id: i32, - duration: Duration, -) -> Result<(Code, BytesMut), Error> { - timeout(duration, read(&mut stream, client_id)) - .await - .map_err(|_| Error::ConnectionTimeout { duration })? -} - -/// -/// Reads a Postgres message from client -/// -/// The SSLRequest/Response sequence requires the Backend to inspect the first byte of the message -/// Byte is then passed as `code` to this function to preserve the message structure -/// -/// -async fn read( - mut stream: S, - client_id: i32, -) -> Result<(Code, BytesMut), Error> { - let code = stream.read_u8().await?; - let len = stream.read_i32().await?; - - // Detect unexpected message len and avoid panic on read_exact - // Len must be at least 4 bytes (4 bytes for len/i32) - if (len as usize) < SIZE_I32 { - error!( - msg = "Unexpected PostgreSQL message length", - code = code, - len = len - ); - return Err(ProtocolError::UnexpectedMessageLength { - code, - len: len as usize, - } - .into()); - } - - let capacity = len as usize + SIZE_U8; //len plus len of code - let mut bytes = BytesMut::with_capacity(capacity); - - bytes.put_u8(code); - bytes.put_i32(len); - - let slice_start = bytes.len(); - - // Capacity and len are not the same!! - // resize populates the buffer with 0s - bytes.resize(capacity, 0); - - stream.read_exact(&mut bytes[slice_start..]).await?; - - debug!(target: PROTOCOL, client_id, code = ?(code as char), ?bytes); - - Ok((code, bytes)) -} diff --git a/packages/cipherstash-proxy/src/postgresql/messages/bind.rs b/packages/cipherstash-proxy/src/postgresql/rewrite/bind.rs similarity index 71% rename from packages/cipherstash-proxy/src/postgresql/messages/bind.rs rename to packages/cipherstash-proxy/src/postgresql/rewrite/bind.rs index 410436953..60609872f 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/bind.rs +++ b/packages/cipherstash-proxy/src/postgresql/rewrite/bind.rs @@ -1,3 +1,4 @@ +//! CipherStash Bind parameter rewriting. use super::{maybe_json, maybe_jsonb, Name, NULL}; use crate::error::{Error, MappingError, ProtocolError}; use crate::log::MAPPER; @@ -9,29 +10,25 @@ use crate::postgresql::data::{ bind_param_from_sql, bind_param_json_value, json_value_selector_plaintext, }; use crate::postgresql::format_code::FormatCode; -use crate::postgresql::protocol::BytesMutReadString; +#[cfg(test)] +use crate::postgresql::test_codec::{decode_frontend_frame, encode_frontend_message}; use crate::{EqlOutput, EqlQueryPayload}; -use crate::{SIZE_I16, SIZE_I32}; -use bytes::{Buf, BufMut, BytesMut}; +use bytes::{BufMut, BytesMut}; use cipherstash_client::encryption::Plaintext; +use pg_proto::codec::{Bind as PgBind, FrontendMessage}; use postgres_types::Type; +use std::convert::TryFrom; use std::fmt::{self, Display, Formatter}; -use std::io::Cursor; -use std::{convert::TryFrom, ffi::CString}; use tracing::debug; /// Bind (B) message. /// See: #[derive(Clone, Debug)] pub struct Bind { - pub code: char, pub portal: Name, pub prepared_statement: Name, - pub num_param_format_codes: i16, pub param_format_codes: Vec, - pub num_param_values: i16, pub param_values: Vec, - pub num_result_column_format_codes: i16, pub result_columns_format_codes: Vec, /// Set when the param list was rebuilt because the rewrite reshaped the /// params. The message must then be re-sent even if no individual param was @@ -43,6 +40,7 @@ pub struct Bind { pub struct BindParam { pub format_code: FormatCode, pub bytes: BytesMut, + null: bool, dirty: bool, } @@ -184,8 +182,6 @@ impl Bind { } self.param_format_codes = param_values.iter().map(|param| param.format_code).collect(); - self.num_param_format_codes = self.param_format_codes.len() as i16; - self.num_param_values = param_values.len() as i16; self.param_values = param_values; self.reshaped = true; @@ -217,6 +213,55 @@ impl Bind { } } +impl TryFrom for Bind { + type Error = Error; + + fn try_from(bind: PgBind) -> Result { + let portal = bind.portal; + let prepared_statement = bind.statement; + let param_format_codes = bind + .parameter_formats + .iter() + .copied() + .map(FormatCode::from) + .collect::>(); + let num_param_values = bind.parameters.len(); + let mut param_values = Vec::with_capacity(num_param_values); + for (idx, parameter) in bind.parameters.into_iter().enumerate() { + let format_code = match param_format_codes.len() { + 0 => FormatCode::Text, + 1 => param_format_codes[0], + len if len == num_param_values => param_format_codes[idx], + _ => { + return Err(ProtocolError::ParameterFormatCodesMismatch { + expected: num_param_values, + received: param_format_codes.len(), + } + .into()); + } + }; + match parameter { + None => param_values.push(BindParam::null_with_format(format_code)), + Some(bytes) => { + param_values.push(BindParam::new(format_code, BytesMut::from(&bytes[..]))); + } + } + } + Ok(Self { + portal, + prepared_statement, + param_format_codes, + param_values, + result_columns_format_codes: bind + .result_formats + .into_iter() + .map(FormatCode::from) + .collect(), + reshaped: false, + }) + } +} + /// /// Param type is either provided with Parse message or the column type /// Column type is the cast of the encrypted column @@ -233,6 +278,7 @@ impl BindParam { Self { format_code, bytes, + null: false, dirty: false, } } @@ -241,6 +287,16 @@ impl BindParam { Self { format_code: FormatCode::Text, bytes: BytesMut::new(), + null: true, + dirty: false, + } + } + + fn null_with_format(format_code: FormatCode) -> Self { + Self { + format_code, + bytes: BytesMut::new(), + null: true, dirty: false, } } @@ -268,6 +324,7 @@ impl BindParam { pub fn rewrite(&mut self, bytes: &[u8]) { self.bytes.clear(); + self.null = false; if self.is_binary() { self.bytes.put_u8(1); @@ -287,6 +344,7 @@ impl BindParam { /// stop `->` from matching any stored entry. pub fn rewrite_text(&mut self, bytes: Vec) { self.bytes.clear(); + self.null = false; self.bytes.extend_from_slice(&bytes); self.dirty = true; } @@ -311,7 +369,7 @@ impl BindParam { } pub fn is_null(&self) -> bool { - self.bytes.is_empty() + self.null } pub fn is_text(&self) -> bool { @@ -330,144 +388,105 @@ impl Display for BindParam { } } +#[cfg(test)] impl TryFrom<&BytesMut> for Bind { type Error = Error; fn try_from(buf: &BytesMut) -> Result { - let mut cursor = Cursor::new(buf); - let code = cursor.get_u8() as char; - let _len = cursor.get_i32(); - - let portal = cursor.read_string()?; - let portal = Name::from(portal); - - let prepared_statement = cursor.read_string()?; - let prepared_statement = Name::from(prepared_statement); - - let num_param_format_codes = cursor.get_i16(); - let mut param_format_codes = Vec::new(); - - for _ in 0..num_param_format_codes { - param_format_codes.push(cursor.get_i16().into()); - } - - let num_param_values = cursor.get_i16(); - let mut param_values = Vec::new(); - - for idx in 0..num_param_values as usize { - let param_len = cursor.get_i32(); - - let format_code = match num_param_format_codes { + let FrontendMessage::Bind(bind) = decode_frontend_frame(buf)? else { + return Err(ProtocolError::UnexpectedMessageCode { + expected: 'B', + received: buf.first().copied().unwrap_or_default() as char, + } + .into()); + }; + let portal = bind.portal; + let prepared_statement = bind.statement; + let param_format_codes = bind + .parameter_formats + .iter() + .copied() + .map(FormatCode::from) + .collect::>(); + let num_param_values = bind.parameters.len(); + let mut param_values = Vec::with_capacity(bind.parameters.len()); + for (idx, parameter) in bind.parameters.into_iter().enumerate() { + let format_code = match param_format_codes.len() { 0 => FormatCode::Text, 1 => param_format_codes[0], - _ => param_format_codes[idx], - }; - - // NULL parameters have a length of -1 and no bytes - match param_len { - NULL => { - param_values.push(BindParam::null()); - } + len if len == num_param_values => param_format_codes[idx], _ => { - let mut bytes = BytesMut::with_capacity(param_len as usize); - bytes.resize(param_len as usize, b'0'); - cursor.copy_to_slice(&mut bytes); - param_values.push(BindParam::new(format_code, bytes)); + return Err(ProtocolError::ParameterFormatCodesMismatch { + expected: num_param_values, + received: param_format_codes.len(), + } + .into()) + } + }; + match parameter { + None => param_values.push(BindParam::null_with_format(format_code)), + Some(bytes) => { + param_values.push(BindParam::new(format_code, BytesMut::from(&bytes[..]))) } } } - - let num_result_column_format_codes = cursor.get_i16(); - let mut result_columns_format_codes = Vec::new(); - - for _ in 0..num_result_column_format_codes { - result_columns_format_codes.push(cursor.get_i16().into()); - } - + let result_columns_format_codes = bind + .result_formats + .iter() + .copied() + .map(FormatCode::from) + .collect::>(); Ok(Bind { - code, portal, prepared_statement, - num_param_format_codes, param_format_codes, - num_param_values, param_values, - num_result_column_format_codes, result_columns_format_codes, reshaped: false, }) } } +#[cfg(test)] impl TryFrom for BytesMut { type Error = Error; fn try_from(bind: Bind) -> Result { - let mut bytes = BytesMut::new(); - - let portal_binding = CString::new(&*bind.portal)?; - let portal = portal_binding.as_bytes_with_nul(); - - let prepared_statement_binding = CString::new(&*bind.prepared_statement)?; - let prepared_statement = prepared_statement_binding.as_bytes_with_nul(); - - if bind.num_param_format_codes != bind.param_format_codes.len() as i16 { - let err = ProtocolError::ParameterFormatCodesMismatch { - expected: bind.num_param_format_codes as usize, - received: bind.param_format_codes.len(), - }; - return Err(err.into()); - } - - if bind.num_result_column_format_codes != bind.result_columns_format_codes.len() as i16 { - let err = ProtocolError::ParameterResultFormatCodesMismatch { - expected: bind.num_result_column_format_codes as usize, - received: bind.result_columns_format_codes.len(), - }; - return Err(err.into()); - } - - // sum of param byte_lens (the *actual* byte lengths of the parameters) - let param_byte_len = &bind - .param_values - .iter() - .fold(0, |acc, param| acc + SIZE_I32 + param.byte_len()); - - let len = SIZE_I32 // self/len of len - + portal.len() - + prepared_statement.len() - + SIZE_I16 // num_param_format_codes - + SIZE_I16 * bind.num_param_format_codes as usize // num_param_format_codes - + SIZE_I16 // num_param_values - + param_byte_len // parameter bytes - + SIZE_I16 // num_result_column_format_codes - + SIZE_I16 * bind.num_result_column_format_codes as usize; - - bytes.put_u8(bind.code as u8); - bytes.put_i32(len as i32); - bytes.put_slice(portal); - bytes.put_slice(prepared_statement); - bytes.put_i16(bind.num_param_format_codes); - for param_format_code in bind.param_format_codes { - bytes.put_i16(param_format_code.into()); - } - - let num_param_values = bind.param_values.len() as i16; - bytes.put_i16(num_param_values); - - for p in bind.param_values { - // len is not the same as byte_len - // A NULL param len is -1 - bytes.put_i32(p.len()); - bytes.put_slice(&p.bytes); - } - - bytes.put_i16(bind.num_result_column_format_codes); - for result_column_format_code in bind.result_columns_format_codes { - bytes.put_i16(result_column_format_code.into()); - } + encode_frontend_message(&FrontendMessage::Bind(PgBind { + portal: bind.portal, + statement: bind.prepared_statement, + parameter_formats: bind.param_format_codes.into_iter().map(i16::from).collect(), + parameters: bind + .param_values + .into_iter() + .map(|param| (!param.null).then(|| param.bytes.freeze())) + .collect(), + result_formats: bind + .result_columns_format_codes + .into_iter() + .map(i16::from) + .collect(), + })) + } +} - Ok(bytes) +impl From for FrontendMessage { + fn from(bind: Bind) -> Self { + Self::Bind(PgBind { + portal: bind.portal, + statement: bind.prepared_statement, + parameter_formats: bind.param_format_codes.into_iter().map(i16::from).collect(), + parameters: bind + .param_values + .into_iter() + .map(|param| (!param.null).then(|| param.bytes.freeze())) + .collect(), + result_formats: bind + .result_columns_format_codes + .into_iter() + .map(i16::from) + .collect(), + }) } } @@ -477,7 +496,7 @@ mod tests { use crate::{ config::LogConfig, log, - postgresql::{format_code::FormatCode, messages::bind::Bind}, + postgresql::{format_code::FormatCode, rewrite::bind::Bind}, }; use bytes::BytesMut; @@ -520,6 +539,19 @@ mod tests { assert_eq!(bytes, expected); } + #[test] + pub fn preserves_empty_and_null_params_distinctly() { + let bytes = to_message(b"B\0\0\0\x14\0\0\0\0\0\x02\0\0\0\0\xff\xff\xff\xff\0\0"); + let expected = bytes.clone(); + + let bind = Bind::try_from(&bytes).unwrap(); + + assert!(!bind.param_values[0].is_null()); + assert_eq!(bind.param_values[0].byte_len(), 0); + assert!(bind.param_values[1].is_null()); + assert_eq!(BytesMut::try_from(bind).unwrap(), expected); + } + #[test] fn bind_should_rewrite() { log::init(LogConfig::default()); diff --git a/packages/cipherstash-proxy/src/postgresql/messages/data_row.rs b/packages/cipherstash-proxy/src/postgresql/rewrite/data_row.rs similarity index 89% rename from packages/cipherstash-proxy/src/postgresql/messages/data_row.rs rename to packages/cipherstash-proxy/src/postgresql/rewrite/data_row.rs index 96d5cb8af..54239e757 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/data_row.rs +++ b/packages/cipherstash-proxy/src/postgresql/rewrite/data_row.rs @@ -1,12 +1,17 @@ -use super::{BackendCode, NULL}; +//! CipherStash DataRow rewriting. use crate::EqlCiphertext; +#[cfg(test)] +use crate::{ + error::ProtocolError, + postgresql::test_codec::{decode_backend_frame, encode_backend_message}, +}; use crate::{ - error::{EncryptError, Error, ProtocolError}, + error::{EncryptError, Error}, log::DECRYPT, postgresql::Column, }; -use bytes::{Buf, BufMut, BytesMut}; -use std::io::Cursor; +use bytes::BytesMut; +use pg_proto::codec::{BackendMessage, DataRow as PgDataRow}; use tracing::{debug, error}; /// Leading byte of `jsonb`'s binary wire format. PostgreSQL has only ever @@ -63,11 +68,9 @@ impl DataRow { } fn len_of_columns(&self) -> usize { - let column_len_size = size_of::(); // len of column len - self.columns .iter() - .map(|col| column_len_size + col.bytes.as_ref().map(|b| b.len()).unwrap_or(0)) + .map(|column| size_of::() + column.bytes.as_ref().map_or(0, BytesMut::len)) .sum() } @@ -98,84 +101,68 @@ impl DataColumn { } } +#[cfg(test)] impl TryFrom<&BytesMut> for DataRow { type Error = Error; fn try_from(buf: &BytesMut) -> Result { - let mut cursor = Cursor::new(buf); - - let code = cursor.get_u8(); - - if BackendCode::from(code) != BackendCode::DataRow { + let BackendMessage::DataRow(row) = decode_backend_frame(buf)? else { return Err(ProtocolError::UnexpectedMessageCode { - expected: BackendCode::DataRow.into(), - received: code as char, + expected: 'D', + received: buf.first().copied().unwrap_or_default() as char, } .into()); - } - - let _len = cursor.get_i32(); - - let num_columns = cursor.get_i16(); - - let mut columns = Vec::new(); - for _ in 0..num_columns { - let len = cursor.get_i32(); - - if len == NULL { - columns.push(DataColumn { bytes: None }); - } else { - let len = len as usize; + }; + let columns = row + .columns + .into_iter() + .map(|bytes| DataColumn { + bytes: bytes.map(BytesMut::from), + }) + .collect(); - let mut bytes = BytesMut::with_capacity(len); - bytes.resize(len, 0); - cursor.copy_to_slice(&mut bytes); + Ok(DataRow { columns }) + } +} - columns.push(DataColumn { bytes: Some(bytes) }); - } +impl From for DataRow { + fn from(row: PgDataRow) -> Self { + Self { + columns: row + .columns + .into_iter() + .map(|bytes| DataColumn { + bytes: bytes.map(BytesMut::from), + }) + .collect(), } - - Ok(DataRow { columns }) } } +#[cfg(test)] impl TryFrom for BytesMut { type Error = Error; fn try_from(data_row: DataRow) -> Result { - let mut bytes = BytesMut::new(); - - let len = size_of::() // len of len - + size_of::() // num columns - + data_row.len_of_columns(); // len data columns - - bytes.put_u8(BackendCode::DataRow.into()); - bytes.put_i32(len as i32); - bytes.put_i16(data_row.columns.len() as i16); - - for col in data_row.columns.into_iter() { - let b = BytesMut::try_from(col)?; - bytes.put_slice(&b); - } - - Ok(bytes) + encode_backend_message(&BackendMessage::DataRow(PgDataRow { + columns: data_row + .columns + .into_iter() + .map(|column| column.bytes.map(|bytes| bytes.freeze())) + .collect(), + })) } } -impl TryFrom for BytesMut { - type Error = Error; - - fn try_from(data_column: DataColumn) -> Result { - let mut bytes = BytesMut::new(); - - if let Some(data) = data_column.bytes { - bytes.put_i32(data.len() as i32); - bytes.put_slice(&data); - } else { - bytes.put_i32(NULL); - } - - Ok(bytes) +impl From for BackendMessage { + fn from(data_row: DataRow) -> Self { + Self::DataRow(PgDataRow { + columns: data_row + .columns + .into_iter() + .map(|column| column.bytes.map(|bytes| bytes.freeze())) + .collect(), + }) } } @@ -279,7 +266,7 @@ mod tests { use crate::{ config::{LogConfig, LogLevel}, log, - postgresql::{messages::data_row::DataColumn, Column}, + postgresql::{rewrite::data_row::DataColumn, Column}, }; use bytes::BytesMut; use cipherstash_client::schema::{ColumnConfig, ColumnType}; diff --git a/packages/cipherstash-proxy/src/postgresql/rewrite/mod.rs b/packages/cipherstash-proxy/src/postgresql/rewrite/mod.rs new file mode 100644 index 000000000..3637cdbee --- /dev/null +++ b/packages/cipherstash-proxy/src/postgresql/rewrite/mod.rs @@ -0,0 +1,20 @@ +use bytes::BytesMut; + +pub mod bind; +pub mod data_row; +pub mod param_description; +pub mod parse; +pub mod query; +pub mod row_description; + +pub type Name = bytes::Bytes; +pub const NULL: i32 = -1; +pub const UNSPECIFIED_TYPE_OID: i32 = 0; + +pub fn maybe_json(bytes: &BytesMut) -> bool { + bytes.first() == Some(&b'{') +} + +pub fn maybe_jsonb(bytes: &BytesMut) -> bool { + bytes.len() > 3 && bytes[0] == 1 && bytes[1] == b'{' +} diff --git a/packages/cipherstash-proxy/src/postgresql/messages/param_description.rs b/packages/cipherstash-proxy/src/postgresql/rewrite/param_description.rs similarity index 76% rename from packages/cipherstash-proxy/src/postgresql/messages/param_description.rs rename to packages/cipherstash-proxy/src/postgresql/rewrite/param_description.rs index 6a9b4c19b..10e6f6df5 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/param_description.rs +++ b/packages/cipherstash-proxy/src/postgresql/rewrite/param_description.rs @@ -1,12 +1,14 @@ -use super::BackendCode; +//! CipherStash ParameterDescription rewriting. +use crate::log::MAPPER; +#[cfg(test)] use crate::{ error::{Error, ProtocolError}, - log::MAPPER, - SIZE_I16, SIZE_I32, + postgresql::test_codec::{decode_backend_frame, encode_backend_message}, }; -use bytes::{Buf, BufMut, BytesMut}; +#[cfg(test)] +use bytes::BytesMut; +use pg_proto::codec::BackendMessage; use postgres_types::Type; -use std::io::Cursor; use tracing::debug; /// @@ -66,58 +68,58 @@ impl ParamDescription { } } +#[cfg(test)] impl TryFrom<&BytesMut> for ParamDescription { type Error = Error; fn try_from(bytes: &BytesMut) -> Result { - let mut cursor = Cursor::new(bytes); - - let code = cursor.get_u8(); - - if BackendCode::from(code) != BackendCode::ParameterDescription { + let BackendMessage::ParameterDescription(types) = decode_backend_frame(bytes)? else { return Err(ProtocolError::UnexpectedMessageCode { - expected: BackendCode::ParameterDescription.into(), - received: code as char, + expected: 't', + received: bytes.first().copied().unwrap_or_default() as char, } .into()); - } - - let _len = cursor.get_i32(); // move the cursor - let count = cursor.get_i16() as usize; - - let mut types = vec![]; - for _idx in 0..count { - let type_oid = cursor.get_i32(); - types.push(type_oid) - } + }; Ok(ParamDescription { - types, + types: types.into_iter().map(|oid| oid as i32).collect(), dirty: false, }) } } +impl From> for ParamDescription { + fn from(types: Vec) -> Self { + Self { + types: types.into_iter().map(|oid| oid as i32).collect(), + dirty: false, + } + } +} + +#[cfg(test)] impl TryFrom for BytesMut { type Error = Error; fn try_from(parameter_description: ParamDescription) -> Result { - let mut bytes = BytesMut::new(); - - let count = parameter_description.types.len(); - let size_of_types = count * SIZE_I32; - - let len = SIZE_I32 + SIZE_I16 + size_of_types; - - bytes.put_u8(BackendCode::ParameterDescription.into()); - bytes.put_i32(len as i32); - bytes.put_i16(count as i16); - - for type_oid in parameter_description.types.into_iter() { - bytes.put_i32(type_oid); - } + let types = parameter_description + .types + .into_iter() + .map(|oid| oid as u32) + .collect(); + encode_backend_message(&BackendMessage::ParameterDescription(types)) + } +} - Ok(bytes) +impl From for BackendMessage { + fn from(parameter_description: ParamDescription) -> Self { + Self::ParameterDescription( + parameter_description + .types + .into_iter() + .map(|oid| oid as u32) + .collect(), + ) } } diff --git a/packages/cipherstash-proxy/src/postgresql/messages/parse.rs b/packages/cipherstash-proxy/src/postgresql/rewrite/parse.rs similarity index 76% rename from packages/cipherstash-proxy/src/postgresql/messages/parse.rs rename to packages/cipherstash-proxy/src/postgresql/rewrite/parse.rs index 8f7c9666d..bb5cc451a 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/parse.rs +++ b/packages/cipherstash-proxy/src/postgresql/rewrite/parse.rs @@ -1,20 +1,22 @@ -use super::{FrontendCode, Name, UNSPECIFIED_TYPE_OID}; +//! CipherStash Parse rewriting. +use super::{Name, UNSPECIFIED_TYPE_OID}; +use crate::postgresql::context::statement::OutputParam; +#[cfg(test)] use crate::{ error::{Error, ProtocolError}, - postgresql::{context::statement::OutputParam, protocol::BytesMutReadString}, - SIZE_I16, SIZE_I32, + postgresql::test_codec::{decode_frontend_frame, encode_frontend_message}, }; -use bytes::{Buf, BufMut, BytesMut}; +use bytes::Bytes; +#[cfg(test)] +use bytes::BytesMut; use eql_mapper::EqlTermVariant; +use pg_proto::codec::{FrontendMessage, Parse as PgParse}; use postgres_types::Type; -use std::{ffi::CString, io::Cursor}; #[derive(Debug, Clone)] pub struct Parse { - pub code: char, pub name: Name, pub statement: String, - pub num_params: i16, pub param_types: Vec, dirty: bool, } @@ -76,7 +78,6 @@ impl Parse { .collect::>(); if param_types != self.param_types { - self.num_params = param_types.len() as i16; self.param_types = param_types; self.dirty = true; } @@ -88,72 +89,78 @@ impl Parse { } } +impl From for Parse { + fn from(parse: PgParse) -> Self { + Self { + name: parse.statement, + statement: String::from_utf8_lossy(&parse.query).into_owned(), + param_types: parse + .parameter_types + .into_iter() + .map(|oid| oid as i32) + .collect(), + dirty: false, + } + } +} + +#[cfg(test)] impl TryFrom<&BytesMut> for Parse { type Error = Error; fn try_from(buf: &BytesMut) -> Result { - let mut cursor = Cursor::new(buf); - let code = cursor.get_u8() as char; - - if FrontendCode::from(code) != FrontendCode::Parse { + let FrontendMessage::Parse(parse) = decode_frontend_frame(buf)? else { return Err(ProtocolError::UnexpectedMessageCode { - expected: FrontendCode::Parse.into(), - received: code, + expected: 'P', + received: buf.first().copied().unwrap_or_default() as char, } .into()); - } - - let _len = cursor.get_i32(); - let name = cursor.read_string()?; - let name = Name::from(name); - - let statement = cursor.read_string()?; - let num_params = cursor.get_i16(); - let mut param_types = Vec::new(); - - for _ in 0..num_params { - param_types.push(cursor.get_i32()); - } + }; + let name = parse.statement; + let statement = String::from_utf8_lossy(&parse.query).into_owned(); + let param_types = parse + .parameter_types + .iter() + .map(|oid| *oid as i32) + .collect::>(); Ok(Parse { - code, name, statement, - num_params, param_types, dirty: false, }) } } +#[cfg(test)] impl TryFrom for BytesMut { type Error = Error; fn try_from(parse: Parse) -> Result { - let mut bytes = BytesMut::new(); - - let name = CString::new(parse.name.as_str())?; - let name = name.as_bytes_with_nul(); - - let statement = CString::new(parse.statement)?; - let statement = statement.as_bytes_with_nul(); - - let len = SIZE_I32 // len - + name.len() - + statement.len() - + SIZE_I16 // num_params - + SIZE_I32 * parse.param_types.len(); - - bytes.put_u8(FrontendCode::Parse.into()); - bytes.put_i32(len as i32); - bytes.put_slice(name); - bytes.put_slice(statement); - bytes.put_i16(parse.num_params); - for param in parse.param_types { - bytes.put_i32(param); - } + encode_frontend_message(&FrontendMessage::Parse(PgParse { + statement: parse.name, + query: Bytes::from(parse.statement), + parameter_types: parse + .param_types + .into_iter() + .map(|oid| oid as u32) + .collect(), + })) + } +} - Ok(bytes) +impl From for FrontendMessage { + fn from(parse: Parse) -> Self { + Self::Parse(PgParse { + statement: parse.name, + query: Bytes::from(parse.statement), + parameter_types: parse + .param_types + .into_iter() + .map(|oid| oid as u32) + .collect(), + }) } } @@ -164,7 +171,7 @@ mod tests { log, postgresql::{ context::statement::{OutputParam, OutputParamSource}, - messages::parse::Parse, + rewrite::parse::Parse, Column, }, Identifier, @@ -250,7 +257,7 @@ mod tests { parse.rewrite_param_types(&output_params); assert!(parse.requires_rewrite()); - assert_eq!(parse.num_params, 1); + assert_eq!(parse.param_types.len(), 1); assert_eq!( parse.param_types, vec![postgres_types::Type::INT2.oid() as i32] diff --git a/packages/cipherstash-proxy/src/postgresql/messages/query.rs b/packages/cipherstash-proxy/src/postgresql/rewrite/query.rs similarity index 50% rename from packages/cipherstash-proxy/src/postgresql/messages/query.rs rename to packages/cipherstash-proxy/src/postgresql/rewrite/query.rs index 88a5cd570..ea311d005 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/query.rs +++ b/packages/cipherstash-proxy/src/postgresql/rewrite/query.rs @@ -1,13 +1,15 @@ +//! CipherStash simple Query rewriting. +#[cfg(test)] use crate::error::{Error, ProtocolError}; -use crate::postgresql::protocol::BytesMutReadString; -use crate::SIZE_I32; - -use bytes::{Buf, BufMut, BytesMut}; +#[cfg(test)] +use crate::postgresql::test_codec::{decode_frontend_frame, encode_frontend_message}; + +use bytes::Bytes; +#[cfg(test)] +use bytes::BytesMut; +use pg_proto::codec::FrontendMessage; +#[cfg(test)] use std::convert::TryFrom; -use std::ffi::CString; -use std::io::Cursor; - -use super::FrontendCode; #[derive(Debug, Clone)] pub struct Query { @@ -34,46 +36,46 @@ impl Query { } } +impl From for Query { + fn from(query: Bytes) -> Self { + Self { + statement: String::from_utf8_lossy(&query).into_owned(), + dirty: false, + } + } +} + +#[cfg(test)] impl TryFrom<&BytesMut> for Query { type Error = Error; fn try_from(bytes: &BytesMut) -> Result { - let mut cursor = Cursor::new(bytes); - let code = cursor.get_u8(); - - if FrontendCode::from(code) != FrontendCode::Query { + let FrontendMessage::Query(query) = decode_frontend_frame(bytes)? else { return Err(ProtocolError::UnexpectedMessageCode { - expected: FrontendCode::Query.into(), - received: code as char, + expected: 'Q', + received: bytes.first().copied().unwrap_or_default() as char, } .into()); - } - - let _len = cursor.get_i32(); // read and progress cursor - let query = cursor.read_string()?; + }; Ok(Query { - statement: query, + statement: String::from_utf8_lossy(&query).into_owned(), dirty: false, }) } } +#[cfg(test)] impl TryFrom for BytesMut { type Error = Error; fn try_from(query: Query) -> Result { - let mut bytes = BytesMut::new(); - - let statement = CString::new(query.statement).map_err(|_| ProtocolError::UnexpectedNull)?; - let statement_bytes = statement.as_bytes_with_nul(); - - let len = SIZE_I32 + statement_bytes.len(); // len of query - - bytes.put_u8(FrontendCode::Query.into()); - bytes.put_i32(len as i32); - bytes.put_slice(statement_bytes); + encode_frontend_message(&FrontendMessage::Query(Bytes::from(query.statement))) + } +} - Ok(bytes) +impl From for FrontendMessage { + fn from(query: Query) -> Self { + Self::Query(Bytes::from(query.statement)) } } diff --git a/packages/cipherstash-proxy/src/postgresql/messages/row_description.rs b/packages/cipherstash-proxy/src/postgresql/rewrite/row_description.rs similarity index 55% rename from packages/cipherstash-proxy/src/postgresql/messages/row_description.rs rename to packages/cipherstash-proxy/src/postgresql/rewrite/row_description.rs index 7fd8aea5e..c60db234d 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/row_description.rs +++ b/packages/cipherstash-proxy/src/postgresql/rewrite/row_description.rs @@ -1,16 +1,19 @@ -use std::{ffi::CString, io::Cursor}; - -use bytes::{Buf, BufMut, BytesMut}; +//! CipherStash RowDescription rewriting. +use bytes::Bytes; +#[cfg(test)] +use bytes::BytesMut; +use pg_proto::codec::{ + BackendMessage, FieldDescription as PgFieldDescription, RowDescription as PgRowDescription, +}; use postgres_types::Type; +use crate::postgresql::format_code::FormatCode; +#[cfg(test)] use crate::{ error::{Error, ProtocolError}, - postgresql::{format_code::FormatCode, protocol::BytesMutReadString}, - SIZE_I16, SIZE_I32, + postgresql::test_codec::{decode_backend_frame, encode_backend_message}, }; -use super::BackendCode; - #[derive(Debug)] pub struct RowDescription { pub fields: Vec, @@ -56,116 +59,106 @@ impl RowDescriptionField { } } +#[cfg(test)] impl TryFrom<&BytesMut> for RowDescription { type Error = Error; fn try_from(bytes: &BytesMut) -> Result { - let mut cursor = Cursor::new(bytes); - - let code = cursor.get_u8(); - - if BackendCode::from(code) != BackendCode::RowDescription { + let BackendMessage::RowDescription(description) = decode_backend_frame(bytes)? else { return Err(ProtocolError::UnexpectedMessageCode { - expected: BackendCode::RowDescription.into(), - received: code as char, + expected: 'T', + received: bytes.first().copied().unwrap_or_default() as char, } .into()); - } + }; - let _len = cursor.get_i32(); // move the cursor - let num_fields = cursor.get_i16() as usize; - - let fields = std::iter::repeat_with(|| RowDescriptionField::try_from(&mut cursor)) - .take(num_fields) - .collect::>()?; + let fields = description + .fields + .into_iter() + .map(|field| RowDescriptionField { + name: String::from_utf8_lossy(&field.name).into_owned(), + table_oid: field.table_oid as i32, + table_column: field.column, + type_oid: field.type_oid as i32, + type_size: field.type_size, + type_modifier: field.type_modifier, + format_code: field.format.into(), + dirty: false, + }) + .collect(); Ok(RowDescription { fields }) } } +impl From for RowDescription { + fn from(description: PgRowDescription) -> Self { + Self { + fields: description + .fields + .into_iter() + .map(|field| RowDescriptionField { + name: String::from_utf8_lossy(&field.name).into_owned(), + table_oid: field.table_oid as i32, + table_column: field.column, + type_oid: field.type_oid as i32, + type_size: field.type_size, + type_modifier: field.type_modifier, + format_code: field.format.into(), + dirty: false, + }) + .collect(), + } + } +} + +#[cfg(test)] impl TryFrom for BytesMut { type Error = Error; fn try_from(row_description: RowDescription) -> Result { - let mut bytes = BytesMut::new(); - - // Convert each field to bytes let fields = row_description .fields .into_iter() - .map(BytesMut::try_from) - .collect::, _>>()?; - - let field_count = fields.len(); - let field_size = fields.iter().map(|x| x.len()).sum::(); - - let len = SIZE_I32 + SIZE_I16 + field_size; - - bytes.put_u8(BackendCode::RowDescription.into()); - bytes.put_i32(len as i32); - bytes.put_i16(field_count as i16); - - for field in fields.into_iter() { - bytes.put_slice(&field); - } - - Ok(bytes) + .map(|field| PgFieldDescription { + name: Bytes::from(field.name), + table_oid: field.table_oid as u32, + column: field.table_column, + type_oid: field.type_oid as u32, + type_size: field.type_size, + type_modifier: field.type_modifier, + format: field.format_code.into(), + }) + .collect(); + + encode_backend_message(&BackendMessage::RowDescription(PgRowDescription { fields })) } } -// impl TryFrom<&BytesMut> for RowDescriptionField { -impl TryFrom<&mut Cursor<&BytesMut>> for RowDescriptionField { - type Error = Error; - - fn try_from(cursor: &mut Cursor<&BytesMut>) -> Result { - let name = cursor.read_string()?; - - let table_oid = cursor.get_i32(); - let table_column = cursor.get_i16(); - let type_oid = cursor.get_i32(); - - let type_size = cursor.get_i16(); - let type_modifier = cursor.get_i32(); - let format_code = cursor.get_i16().into(); - - Ok(Self { - name, - table_oid, - table_column, - type_oid, - type_size, - type_modifier, - format_code, - dirty: false, +impl From for BackendMessage { + fn from(row_description: RowDescription) -> Self { + Self::RowDescription(PgRowDescription { + fields: row_description + .fields + .into_iter() + .map(|field| PgFieldDescription { + name: Bytes::from(field.name), + table_oid: field.table_oid as u32, + column: field.table_column, + type_oid: field.type_oid as u32, + type_size: field.type_size, + type_modifier: field.type_modifier, + format: field.format_code.into(), + }) + .collect(), }) } } -impl TryFrom for BytesMut { - type Error = Error; - - fn try_from(field: RowDescriptionField) -> Result { - let mut bytes = BytesMut::new(); - - let name = CString::new(field.name)?; - let name = name.as_bytes_with_nul(); - - bytes.put_slice(name); - bytes.put_i32(field.table_oid); - bytes.put_i16(field.table_column); - bytes.put_i32(field.type_oid); - bytes.put_i16(field.type_size); - bytes.put_i32(field.type_modifier); - bytes.put_i16(field.format_code.into()); - - Ok(bytes) - } -} - #[cfg(test)] mod tests { - use crate::{config::LogConfig, log, postgresql::messages::row_description::RowDescription}; + use crate::{config::LogConfig, log, postgresql::rewrite::row_description::RowDescription}; use bytes::BytesMut; use tracing::info; diff --git a/packages/cipherstash-proxy/src/postgresql/startup.rs b/packages/cipherstash-proxy/src/postgresql/startup.rs deleted file mode 100644 index a21e64f95..000000000 --- a/packages/cipherstash-proxy/src/postgresql/startup.rs +++ /dev/null @@ -1,163 +0,0 @@ -use std::time::Duration; - -use bytes::{BufMut, BytesMut}; -use tokio::{ - io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}, - time::timeout, -}; -use tracing::{debug, error, warn}; - -use crate::{ - connect::AsyncStream, - error::{Error, ProtocolError}, - log::PROTOCOL, - postgresql::{SSL_REQUEST, SSL_RESPONSE_NO, SSL_RESPONSE_YES}, - tls, TandemConfig, SIZE_I32, -}; - -use super::protocol::StartupMessage; - -pub async fn with_tls(stream: AsyncStream, config: &TandemConfig) -> Result { - if config.database_tls_disabled() { - warn!(msg = "Connecting to database without Transport Layer Security (TLS)"); - return Ok(stream); - } - match stream { - AsyncStream::Tcp(mut tcp_stream) => { - let server_supports_ssl = send_ssl_request(&mut tcp_stream).await?; - - match server_supports_ssl { - true => { - let tls_stream = tls::client(tcp_stream, config).await?; - Ok(AsyncStream::Tls(Box::new(tls_stream))) - } - false => { - warn!(msg = "Connecting to database without Transport Layer Security (TLS)"); - Ok(AsyncStream::Tcp(tcp_stream)) - } - } - } - AsyncStream::Tls(_) => { - // Technically unreachable unless the server is misbehaving - warn!(msg = "Database already connected over Transport Layer Security (TLS)"); - Ok(stream) - } - } -} - -/// -/// Reads a Postgres startup message from client with an optional timeout -/// -/// Timeout values are in config -/// -/// -pub async fn read_message( - mut stream: S, - connection_timeout: Option, -) -> Result { - match connection_timeout { - Some(duration) => read_message_with_timeout(stream, duration).await, - None => read(&mut stream).await, - } -} - -/// -/// Reads a Postgres message from client with a timeout -/// -/// Timeout values are in config -/// -/// -async fn read_message_with_timeout( - mut stream: S, - duration: Duration, -) -> Result { - timeout(duration, read(&mut stream)) - .await - .map_err(|_| Error::ConnectionTimeout { duration })? -} - -/// -/// Read the start up message from the client -/// Startup messages are sent by the client to the server to initiate a connection -/// -/// -/// -async fn read(client: &mut C) -> Result -where - C: AsyncRead + Unpin, -{ - let len = client.read_i32().await?; - - let capacity = len as usize; - - let mut bytes = BytesMut::with_capacity(capacity); - bytes.put_i32(len); - bytes.resize(capacity, b'0'); - - let slice_start = SIZE_I32; - client.read_exact(&mut bytes[slice_start..]).await?; - - // code is the first 4 bytes after len - let code_bytes: [u8; 4] = [ - bytes.as_ref()[4], - bytes.as_ref()[5], - bytes.as_ref()[6], - bytes.as_ref()[7], - ]; - - let code = i32::from_be_bytes(code_bytes); - - let message = StartupMessage { - code: code.into(), - bytes, - }; - debug!(target: PROTOCOL, StartupMessage = ?message); - - Ok(message) -} - -/// -/// Send SSLRequest to the stream and return the response -/// Returns true if the server indicates support for TLS -/// -pub async fn send_ssl_request( - stream: &mut T, -) -> Result { - let mut bytes = BytesMut::with_capacity(12); - bytes.put_i32(8); - bytes.put_i32(SSL_REQUEST); - - stream.write_all(&bytes).await?; - - // Server supports TLS - let response = match stream.read_u8().await? { - SSL_RESPONSE_YES => true, - SSL_RESPONSE_NO => false, - code => { - error!(msg = "Unexpected startup message", code = ?(code as char)); - return Err(ProtocolError::UnexpectedStartupMessage.into()); - } - }; - - debug!(target: PROTOCOL, msg = "Database SSLResponse", SSLResponse = ?response); - Ok(response) -} - -/// -/// Send SSLRequest to the stream -/// Returns true if the server indicates support for TLS -/// N for no, S for yeS or tlS -/// The SSLResponse MUST come before the TLS handshake -/// -pub async fn send_ssl_response( - stream: &mut T, - tls: bool, -) -> Result<(), Error> { - let response = if tls { b'S' } else { b'N' }; - - debug!(target: PROTOCOL, msg = "SSLResponse to Client", SSLResponse = ?response); - - stream.write_all(&[response]).await?; - - Ok(()) -} diff --git a/packages/cipherstash-proxy/src/postgresql/test_codec.rs b/packages/cipherstash-proxy/src/postgresql/test_codec.rs new file mode 100644 index 000000000..148e0a74e --- /dev/null +++ b/packages/cipherstash-proxy/src/postgresql/test_codec.rs @@ -0,0 +1,31 @@ +use bytes::BytesMut; +use pg_proto::codec::{Backend, BackendMessage, Frontend, FrontendMessage, PgCodec}; +use tokio_util::codec::{Decoder, Encoder}; + +use crate::error::Error; + +pub fn decode_frontend_frame(bytes: &BytesMut) -> Result { + let mut bytes = bytes.clone(); + PgCodec::::default() + .decode(&mut bytes)? + .ok_or_else(|| std::io::Error::from(std::io::ErrorKind::UnexpectedEof).into()) +} + +pub fn decode_backend_frame(bytes: &BytesMut) -> Result { + let mut bytes = bytes.clone(); + PgCodec::::default() + .decode(&mut bytes)? + .ok_or_else(|| std::io::Error::from(std::io::ErrorKind::UnexpectedEof).into()) +} + +pub fn encode_frontend_message(message: &FrontendMessage) -> Result { + let mut bytes = BytesMut::new(); + PgCodec::::default().encode(message.to_frame()?, &mut bytes)?; + Ok(bytes) +} + +pub fn encode_backend_message(message: &BackendMessage) -> Result { + let mut bytes = BytesMut::new(); + PgCodec::::default().encode(message.to_frame()?, &mut bytes)?; + Ok(bytes) +} diff --git a/packages/cipherstash-proxy/src/tls/mod.rs b/packages/cipherstash-proxy/src/tls/mod.rs index 82c70cea1..8c1fab46b 100644 --- a/packages/cipherstash-proxy/src/tls/mod.rs +++ b/packages/cipherstash-proxy/src/tls/mod.rs @@ -1,40 +1,10 @@ +use crate::DatabaseConfig; use crate::{config::TlsConfig, error::Error}; -use crate::{DatabaseConfig, TandemConfig}; use rustls::client::danger::ServerCertVerifier; use rustls::ClientConfig; use rustls_pki_types::{pem::PemObject, CertificateDer, PrivateKeyDer, ServerName}; use rustls_platform_verifier::ConfigVerifierExt; use std::sync::Arc; -use tokio::net::TcpStream; -use tokio_rustls::{TlsAcceptor, TlsConnector, TlsStream}; - -/// -/// Create a Server TLS connection -/// The returned type is the higher-level TlsStream that wraps both Client & Server variants -/// -pub async fn client( - stream: TcpStream, - config: &TandemConfig, -) -> Result, Error> { - let tls_config = configure_client(&config.database); - let connector = TlsConnector::from(Arc::new(tls_config)); - let domain = config.database.server_name()?.to_owned(); - let tls_stream = connector.connect(domain, stream).await?; - - Ok(tls_stream.into()) -} - -/// -/// Create a Server TLS connection -/// The returned type is the higher-level TlsStream that wraps both Client & Server variants -/// -pub async fn server(stream: TcpStream, config: &TlsConfig) -> Result, Error> { - let server_config = configure_server(config)?; - let acceptor = TlsAcceptor::from(Arc::new(server_config)); - let tls_stream = acceptor.accept(stream).await?; - - Ok(tls_stream.into()) -} /// /// Configure the server TLS settings @@ -44,6 +14,14 @@ pub async fn server(stream: TcpStream, config: &TlsConfig) -> Result Result { + configure_server_with_leaf(config).map(|(config, _)| config) +} + +/// Builds the server TLS configuration and returns its leaf certificate for +/// pg-proto's RFC 5929 channel-binding transport. +pub fn configure_server_with_leaf( + config: &TlsConfig, +) -> Result<(rustls::ServerConfig, CertificateDer<'static>), Error> { let certs = match config { TlsConfig::Pem { certificate_pem: certificate, @@ -68,11 +46,15 @@ pub fn configure_server(config: &TlsConfig) -> Result PrivateKeyDer::from_pem_file(private_key), }?; + let leaf = certs + .first() + .cloned() + .ok_or(rustls::Error::NoCertificatesPresented)?; let server_config = rustls::ServerConfig::builder() .with_no_client_auth() .with_single_cert(certs, key)?; - Ok(server_config) + Ok((server_config, leaf)) } ///