From a7a61a3c06647552b3a3cb33c8ebaf8523750e70 Mon Sep 17 00:00:00 2001 From: James Sadler Date: Tue, 4 Aug 2026 16:39:11 +1000 Subject: [PATCH 01/16] refactor(proxy): migrate PostgreSQL protocol to pg-proto --- Cargo.lock | 250 +++++++++++++----- PG_PROTO_MIGRATION_PLAN.md | 38 +++ packages/cipherstash-proxy/Cargo.toml | 1 + .../src/postgresql/backend.rs | 40 +-- .../src/postgresql/context/mod.rs | 233 +++++++++++++++- .../src/postgresql/frontend.rs | 28 +- .../src/postgresql/handler.rs | 26 +- .../src/postgresql/messages/bind.rs | 186 ++++++------- .../src/postgresql/messages/close.rs | 53 ++-- .../src/postgresql/messages/data_row.rs | 94 ++----- .../src/postgresql/messages/describe.rs | 53 ++-- .../src/postgresql/messages/execute.rs | 27 +- .../postgresql/messages/param_description.rs | 52 +--- .../src/postgresql/messages/parse.rs | 84 +++--- .../src/postgresql/messages/query.rs | 39 +-- .../postgresql/messages/row_description.rs | 128 +++------ .../cipherstash-proxy/src/postgresql/mod.rs | 10 - .../src/postgresql/protocol.rs | 188 +++++++++++-- .../src/postgresql/startup.rs | 106 ++++++-- 19 files changed, 1009 insertions(+), 627 deletions(-) create mode 100644 PG_PROTO_MIGRATION_PLAN.md diff --git a/Cargo.lock b/Cargo.lock index 32ad0774c..ff6f77ab3 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", ] @@ -836,11 +836,12 @@ dependencies = [ "eql-mapper", "exitcode", "hex", - "md-5", + "md-5 0.10.6", "metrics", "metrics-exporter-prometheus", "moka", "oid-registry", + "pg-proto", "pg_escape", "postgres-protocol", "postgres-types", @@ -866,7 +867,7 @@ dependencies = [ "tracing-subscriber", "uuid", "vitaminc-protected 0.1.0-pre4.2", - "x509-parser", + "x509-parser 0.17.0", ] [[package]] @@ -972,6 +973,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 +1038,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 +1205,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 +1270,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 +1415,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 +1593,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 +2030,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 +2153,7 @@ dependencies = [ "libc", "percent-encoding", "pin-project-lite", - "socket2 0.6.1", + "socket2 0.6.5", "system-configuration", "tokio", "tower-service", @@ -2500,9 +2533,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 +2625,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 +2697,7 @@ dependencies = [ "cfg-if", "miette-derive", "thiserror 1.0.69", - "unicode-width", + "unicode-width 0.1.14", ] [[package]] @@ -2691,13 +2734,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]] @@ -2982,6 +3025,40 @@ version = "2.3.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e3148f5046208a5d56bcfc03053e3ca6334e51da8dfb19b6cdc8b306fae3283e" +[[package]] +name = "pg-proto" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7a51a1b6becd8bc3571cc353b5f61ae8726f46ecacd78ab8222535cf710b92dd" +dependencies = [ + "base64", + "bytes", + "hmac 0.13.0", + "pg-proto-fsm", + "postgres-protocol", + "rand 0.10.2", + "rustls", + "sha2 0.11.0", + "stringprep", + "subtle", + "tokio", + "tokio-rustls", + "tokio-util", + "x509-parser 0.18.1", +] + +[[package]] +name = "pg-proto-fsm" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ba10fdad1b08db94d8d087940da00427599fe49e1e1d2762c8d95e7c7a7d8509" +dependencies = [ + "proc-macro2", + "quote", + "railroad", + "syn 3.0.3", +] + [[package]] name = "pg_escape" version = "0.1.1" @@ -3083,19 +3160,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 +3282,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 +3402,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 +3432,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 +3560,7 @@ dependencies = [ "rand_chacha 0.3.1", "serde", "serde_cbor", - "sha2", + "sha2 0.10.8", "thiserror 1.0.69", "zeroize", ] @@ -3495,7 +3581,7 @@ dependencies = [ "rand_chacha 0.3.1", "serde", "serde_cbor", - "sha2", + "sha2 0.10.8", "thiserror 1.0.69", "zeroize", ] @@ -3557,7 +3643,7 @@ checksum = "2c9283685feec7d69af75fb0e858d5e7378f33fe4fc699383b2916ab9273e03c" dependencies = [ "proc-macro2", "quote", - "syn 3.0.2", + "syn 3.0.3", ] [[package]] @@ -3775,14 +3861,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 +3893,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 +3919,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 +3951,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 +4248,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 +4356,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 +4453,7 @@ dependencies = [ "cfg-if", "libc", "psm", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] @@ -4417,9 +4515,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 +4724,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 +4734,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 +4782,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 +4793,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 +4803,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 +5062,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 +5160,7 @@ dependencies = [ "atomic", "getrandom 0.3.2", "js-sys", - "md-5", + "md-5 0.10.6", "serde", "sha1_smol", "wasm-bindgen", @@ -5496,7 +5600,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 +5782,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 +6195,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", @@ -6123,6 +6218,23 @@ dependencies = [ "time", ] +[[package]] +name = "x509-parser" +version = "0.18.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d43b0f71ce057da06bc0851b23ee24f3f86190b07203dd8f567d0b706a185202" +dependencies = [ + "asn1-rs", + "data-encoding", + "der-parser", + "lazy_static", + "nom 7.1.3", + "oid-registry", + "rusticata-macros", + "thiserror 2.0.18", + "time", +] + [[package]] name = "yansi" version = "1.0.1" diff --git a/PG_PROTO_MIGRATION_PLAN.md b/PG_PROTO_MIGRATION_PLAN.md new file mode 100644 index 000000000..c83a76a30 --- /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.1.0](https://crates.io/crates/pg-proto/0.1.0), covering codecs, startup/authentication, and runtime protocol-state validation. + +## 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..bbca02cb7 100644 --- a/packages/cipherstash-proxy/Cargo.toml +++ b/packages/cipherstash-proxy/Cargo.toml @@ -31,6 +31,7 @@ metrics-exporter-prometheus = "0.17" moka = { version = "0.12", features = ["future"] } oid-registry = "0.8" pg_escape = "0.1.1" +pg-proto = "0.1.0" postgres-protocol = "0.6.7" postgres-types = { version = "0.2.8", features = ["with-serde_json-1"] } rand = "0.9" diff --git a/packages/cipherstash-proxy/src/postgresql/backend.rs b/packages/cipherstash-proxy/src/postgresql/backend.rs index d22730092..36aa9c58d 100644 --- a/packages/cipherstash-proxy/src/postgresql/backend.rs +++ b/packages/cipherstash-proxy/src/postgresql/backend.rs @@ -4,7 +4,7 @@ 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::messages::UNSPECIFIED_TYPE_OID; use super::Column; use crate::connect::Sender; use crate::error::{EncryptError, Error}; @@ -22,6 +22,7 @@ use crate::proxy::EncryptionService; use crate::EqlCiphertext; use bytes::BytesMut; use metrics::{counter, histogram}; +use pg_proto::codec::BackendMessage; use std::time::Instant; use tokio::io::AsyncRead; use tracing::{debug, error, info, warn}; @@ -150,12 +151,14 @@ where /// 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( + let (code, mut bytes, protocol_message) = protocol::read_backend_message( &mut self.server_reader, self.context.client_id, self.context.connection_timeout(), ) .await?; + + self.context.protocol_backend_received(&protocol_message)?; let read_duration = read_start.elapsed(); self.context.record_execute_server_timing(read_duration); @@ -188,11 +191,11 @@ 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(); } @@ -205,8 +208,8 @@ where 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(_) => { // Encrypted DataRows are added to the buffer and we return early // Otherwise, continue and write if self.data_row_handler(&bytes).await? { @@ -216,9 +219,9 @@ where // 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 { @@ -232,7 +235,7 @@ where self.context.complete_execution(); self.context.finish_session(); } - BackendCode::ErrorResponse => { + BackendMessage::ErrorResponse(_) => { if let Some(b) = self.error_response_handler(&bytes)? { bytes = b } @@ -251,7 +254,7 @@ where // Describe with Target:Statement // Returns a ParameterDescription followed by RowDescription // The Describe is complete after the RowDescription - BackendCode::ParameterDescription => { + BackendMessage::ParameterDescription(_) => { if let Some(b) = self.parameter_description_handler(&bytes).await? { bytes = b } @@ -261,7 +264,7 @@ where // Target::Portal returns a RowDescription // If no rows are returned, NoData is returned instead of a RowDescription // Complete the Describe - BackendCode::RowDescription => { + BackendMessage::RowDescription(_) => { if let Some(b) = self.row_description_handler(&bytes).await? { bytes = b } @@ -269,13 +272,13 @@ where } // 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,7 +288,7 @@ where } } - code => { + _ => { debug!(target: PROTOCOL, client_id = self.context.client_id, msg = "Passthrough", @@ -381,6 +384,7 @@ where /// Write a message to the client /// pub async fn write(&mut self, bytes: BytesMut) -> Result<(), Error> { + self.context.protocol_backend_forwarded(&bytes)?; let sent: u64 = bytes.len() as u64; counter!(CLIENTS_BYTES_SENT_TOTAL).increment(sent); diff --git a/packages/cipherstash-proxy/src/postgresql/context/mod.rs b/packages/cipherstash-proxy/src/postgresql/context/mod.rs index d42e015cf..475333ec5 100644 --- a/packages/cipherstash-proxy/src/postgresql/context/mod.rs +++ b/packages/cipherstash-proxy/src/postgresql/context/mod.rs @@ -7,6 +7,7 @@ pub use self::{phase_timing::PhaseTiming, portal::Portal, statement::Statement}; use super::{ column_mapper::ColumnMapper, messages::{describe::Describe, Name, Target}, + protocol::{decode_backend_frame, decode_frontend_frame}, Column, }; use crate::{ @@ -19,9 +20,15 @@ use crate::{ }, proxy::{EncryptConfig, EncryptionService, ReloadCommand, ReloadSender}, }; +use bytes::BytesMut; use cipherstash_client::IdentifiedBy; use eql_mapper::{Schema, TableResolver}; use metrics::{counter, histogram}; +use pg_proto::{ + codec::{BackendMessage, FrontendMessage}, + grammar::{backend as server_role, frontend as client_role}, + intermediary::Intermediary, +}; use serde_json::json; use sqltk::parser::ast::{Expr, Ident, ObjectName, ObjectNamePart, Set, Value, ValueWithSpan}; pub use statement_metadata::StatementMetadata; @@ -29,7 +36,7 @@ use std::{ collections::{HashMap, VecDeque}, sync::{ atomic::{AtomicU64, Ordering}, - Arc, LazyLock, RwLock, + Arc, LazyLock, Mutex, RwLock, }, time::{Duration, Instant}, }; @@ -42,6 +49,83 @@ type ExecuteQueue = Queue; type SessionMetricsQueue = Queue; type PortalQueue = Queue>; +fn protocol_lock_error(_: std::sync::PoisonError) -> Error { + std::io::Error::other("PostgreSQL protocol state lock poisoned").into() +} + +fn protocol_transition_error(error: impl std::fmt::Debug) -> Error { + std::io::Error::new(std::io::ErrorKind::InvalidData, format!("{error:?}")).into() +} + +fn protocol_neutral_backend_message(message: &BackendMessage) -> bool { + matches!( + message, + BackendMessage::ParameterStatus { .. } + | BackendMessage::NoticeResponse(_) + | BackendMessage::NotificationResponse { .. } + | BackendMessage::BackendKeyData { .. } + | BackendMessage::NegotiateProtocolVersion(_) + ) +} + +#[derive(Debug)] +struct ProtocolState { + sides: Intermediary, + upstream_startup_ready: bool, + downstream_startup_ready: bool, + runtime_started: bool, + downstream_frontend_queue: VecDeque, + upstream_backend_queue: VecDeque, +} + +impl ProtocolState { + fn new() -> Self { + Self { + sides: Intermediary::new( + server_role::RuntimeFsm::new(), + client_role::RuntimeFsm::new(), + ), + upstream_startup_ready: false, + downstream_startup_ready: false, + runtime_started: false, + downstream_frontend_queue: VecDeque::new(), + upstream_backend_queue: VecDeque::new(), + } + } + + fn advance_downstream_frontend(&mut self, message: &FrontendMessage) -> bool { + let (downstream, _) = self.sides.sides_mut(); + downstream + .step_projected(message, server_role::project_external) + .is_ok() + } + + fn drain_downstream_frontend(&mut self) { + while let Some(message) = self.downstream_frontend_queue.front().cloned() { + if !self.advance_downstream_frontend(&message) { + break; + } + self.downstream_frontend_queue.pop_front(); + } + } + + fn advance_upstream_backend(&mut self, message: &BackendMessage) -> bool { + let (_, upstream) = self.sides.sides_mut(); + upstream + .step_projected(message, client_role::project_external) + .is_ok() + } + + fn drain_upstream_backend(&mut self) { + while let Some(message) = self.upstream_backend_queue.front().cloned() { + if !self.advance_upstream_backend(&message) { + break; + } + self.upstream_backend_queue.pop_front(); + } + } +} + #[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] pub struct SessionId(u64); @@ -76,6 +160,7 @@ where unsafe_disable_mapping: bool, keyset_id: Arc>>, session_id_counter: Arc, + protocol: Arc>, } /// Context for tracking an in-flight Execute operation. @@ -197,7 +282,95 @@ 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())), + } + } + + /// Records a message accepted from the downstream client. The downstream + /// server-role state advances even when proxy policy intercepts the message. + pub fn protocol_frontend_received(&self, message: &FrontendMessage) -> Result<(), Error> { + let mut protocol = self.protocol.lock().map_err(protocol_lock_error)?; + protocol.runtime_started = true; + if !protocol.downstream_frontend_queue.is_empty() + || !protocol.advance_downstream_frontend(message) + { + protocol + .downstream_frontend_queue + .push_back(message.clone()); } + Ok(()) + } + + /// Records a frontend message actually forwarded upstream after rewriting. + pub fn protocol_frontend_forwarded(&self, bytes: &BytesMut) -> Result<(), Error> { + let message = decode_frontend_frame(bytes)?; + let mut protocol = self.protocol.lock().map_err(protocol_lock_error)?; + let (_, upstream) = protocol.sides.sides_mut(); + if upstream.state() == client_role::RuntimeState::Ready + && matches!( + message, + FrontendMessage::Parse(_) + | FrontendMessage::Bind(_) + | FrontendMessage::Describe(_) + | FrontendMessage::Execute(_) + | FrontendMessage::Close(_) + | FrontendMessage::Flush + | FrontendMessage::Sync + ) + { + upstream + .step(client_role::Event::BeginExtended) + .map_err(protocol_transition_error)?; + } + upstream + .step_projected(&message, client_role::project_internal) + .map_err(protocol_transition_error)?; + protocol.drain_upstream_backend(); + Ok(()) + } + + /// Records a message accepted from the upstream database. + pub fn protocol_backend_received(&self, message: &BackendMessage) -> Result<(), Error> { + if protocol_neutral_backend_message(message) { + return Ok(()); + } + let mut protocol = self.protocol.lock().map_err(protocol_lock_error)?; + if matches!(message, BackendMessage::ReadyForQuery(_)) && !protocol.upstream_startup_ready { + protocol.upstream_startup_ready = true; + return Ok(()); + } + if !protocol.runtime_started { + return Ok(()); + } + if !protocol.upstream_backend_queue.is_empty() + || !protocol.advance_upstream_backend(message) + { + protocol.upstream_backend_queue.push_back(message.clone()); + } + Ok(()) + } + + /// Records a backend response actually emitted to the downstream client. + pub fn protocol_backend_forwarded(&self, bytes: &BytesMut) -> Result<(), Error> { + let message = decode_backend_frame(bytes)?; + if protocol_neutral_backend_message(&message) { + return Ok(()); + } + let mut protocol = self.protocol.lock().map_err(protocol_lock_error)?; + if matches!(message, BackendMessage::ReadyForQuery(_)) && !protocol.downstream_startup_ready + { + protocol.downstream_startup_ready = true; + return Ok(()); + } + if !protocol.runtime_started { + return Ok(()); + } + let (downstream, _) = protocol.sides.sides_mut(); + downstream + .step_projected(&message, server_role::project_internal) + .map_err(protocol_transition_error)?; + protocol.drain_downstream_frontend(); + Ok(()) } pub fn set_describe(&mut self, describe: Describe) { @@ -1056,7 +1229,9 @@ impl Queue { #[cfg(test)] mod tests { - use super::{Context, Describe, KeysetIdentifier, Portal, Statement}; + use super::{ + server_role, Context, Describe, KeysetIdentifier, Portal, ProtocolState, Statement, + }; use crate::{ config::LogConfig, error::Error, @@ -1068,8 +1243,12 @@ mod tests { proxy::{EncryptConfig, EncryptionService}, TandemConfig, }; + use bytes::Bytes; use cipherstash_client::IdentifiedBy; use eql_mapper::Schema; + use pg_proto::codec::{ + BackendMessage, Bind, Execute, FrontendMessage, Parse, TransactionStatus, + }; use sqltk::parser::{dialect::PostgreSqlDialect, parser::Parser}; use std::sync::Arc; use tokio::sync::mpsc; @@ -1117,6 +1296,56 @@ mod tests { ) } + #[test] + fn server_role_fsm_tracks_pipelined_extended_messages_in_processing_order() { + let mut protocol = ProtocolState::new(); + 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 { + if !protocol.downstream_frontend_queue.is_empty() + || !protocol.advance_downstream_frontend(&message) + { + protocol.downstream_frontend_queue.push_back(message); + } + } + + for response in [ + BackendMessage::ParseComplete, + BackendMessage::BindComplete, + BackendMessage::CommandComplete(Bytes::from_static(b"SELECT 1")), + BackendMessage::ReadyForQuery(TransactionStatus::Idle), + ] { + let (downstream, _) = protocol.sides.sides_mut(); + downstream + .step_projected(&response, server_role::project_internal) + .unwrap(); + protocol.drain_downstream_frontend(); + } + + assert!(protocol.downstream_frontend_queue.is_empty()); + assert_eq!( + protocol.sides.downstream().state(), + server_role::RuntimeState::Ready + ); + } + fn statement() -> Statement { Statement { param_columns: vec![], diff --git a/packages/cipherstash-proxy/src/postgresql/frontend.rs b/packages/cipherstash-proxy/src/postgresql/frontend.rs index 88a3e1f0c..24d35fa2c 100644 --- a/packages/cipherstash-proxy/src/postgresql/frontend.rs +++ b/packages/cipherstash-proxy/src/postgresql/frontend.rs @@ -38,6 +38,7 @@ 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::FrontendMessage; use serde::Serialize; use sqltk::parser::ast::{self, Value}; use sqltk::NodeKey; @@ -164,13 +165,15 @@ where /// 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( + let (code, mut bytes, protocol_message) = protocol::read_frontend_message( &mut self.client_reader, self.context.client_id, self.context.connection_timeout(), ) .await?; + self.context.protocol_frontend_received(&protocol_message)?; + let sent: u64 = bytes.len() as u64; counter!(CLIENTS_BYTES_RECEIVED_TOTAL).increment(sent); @@ -189,13 +192,13 @@ where error_state = ?self.error_state, ?code, ); - if code != Code::Sync { + if !matches!(protocol_message, FrontendMessage::Sync) { return Ok(()); } } - match code { - Code::Query => { + match protocol_message { + FrontendMessage::Query(_) => { match self.query_handler(&bytes).await { Ok(Some(mapped)) => bytes = mapped, // No mapping needed, don't change the bytes @@ -212,13 +215,13 @@ where } } } - Code::Describe => { + FrontendMessage::Describe(_) => { self.describe_handler(&bytes).await?; } - Code::Execute => { + FrontendMessage::Execute(_) => { self.execute_handler(&bytes).await?; } - Code::Parse => { + FrontendMessage::Parse(_) => { match self.parse_handler(&bytes).await { Ok(Some(mapped)) => bytes = mapped, // No mapping needed, don't change the bytes @@ -234,7 +237,7 @@ where } } } - Code::Bind => { + FrontendMessage::Bind(_) => { match self.bind_handler(&bytes).await { Ok(Some(mapped)) => bytes = mapped, // No mapping needed, don't change the bytes @@ -268,7 +271,7 @@ where }, } } - Code::Sync => { + FrontendMessage::Sync => { debug!(target: PROTOCOL, client_id = self.context.client_id, ?code, @@ -285,10 +288,10 @@ where return Ok(()); } } - Code::Close => { + FrontendMessage::Close(_) => { self.close_handler(&bytes).await?; } - code => { + _ => { debug!(target: PROTOCOL, client_id = self.context.client_id, msg = "Passthrough", @@ -302,6 +305,7 @@ where } pub async fn write_to_server(&mut self, bytes: BytesMut) -> Result<(), Error> { + self.context.protocol_frontend_forwarded(&bytes)?; debug!(target: PROTOCOL, msg = "Write to server", ?bytes); let sent: u64 = bytes.len() as u64; counter!(SERVER_BYTES_SENT_TOTAL).increment(sent); @@ -1203,6 +1207,7 @@ where ?message, ); + self.context.protocol_backend_forwarded(&message)?; self.client_sender.send(message)?; self.error_state = None; @@ -1353,6 +1358,7 @@ where ?message, ); + self.context.protocol_backend_forwarded(&message)?; self.client_sender.send(message)?; self.error_state = Some(ErrorState); // Frontend-specific: set error state for extended query protocol diff --git a/packages/cipherstash-proxy/src/postgresql/handler.rs b/packages/cipherstash-proxy/src/postgresql/handler.rs index c1fe8e50e..b8c593454 100644 --- a/packages/cipherstash-proxy/src/postgresql/handler.rs +++ b/packages/cipherstash-proxy/src/postgresql/handler.rs @@ -100,6 +100,9 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R database_stream.write_all(&startup_message.bytes).await?; break; } + StartupCode::GSSENCRequest => { + return Err(ProtocolError::UnexpectedStartupMessage.into()); + } } } @@ -125,15 +128,20 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R 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 (_code, bytes, _message) = match protocol::read_frontend_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)?; diff --git a/packages/cipherstash-proxy/src/postgresql/messages/bind.rs b/packages/cipherstash-proxy/src/postgresql/messages/bind.rs index 410436953..7509101d9 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/bind.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/bind.rs @@ -9,15 +9,14 @@ 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; +use crate::postgresql::protocol::{decode_frontend_frame, encode_frontend_message}; use crate::{EqlOutput, EqlQueryPayload}; -use crate::{SIZE_I16, SIZE_I32}; -use bytes::{Buf, BufMut, BytesMut}; +use bytes::{BufMut, Bytes, 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. @@ -43,6 +42,7 @@ pub struct Bind { pub struct BindParam { pub format_code: FormatCode, pub bytes: BytesMut, + null: bool, dirty: bool, } @@ -233,6 +233,7 @@ impl BindParam { Self { format_code, bytes, + null: false, dirty: false, } } @@ -241,6 +242,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 +279,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 +299,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 +324,7 @@ impl BindParam { } pub fn is_null(&self) -> bool { - self.bytes.is_empty() + self.null } pub fn is_text(&self) -> bool { @@ -334,58 +347,54 @@ 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 = Name::from(String::from_utf8_lossy(&bind.portal).into_owned()); + let prepared_statement = Name::from(String::from_utf8_lossy(&bind.statement).into_owned()); + let param_format_codes = bind + .parameter_formats + .iter() + .copied() + .map(FormatCode::from) + .collect::>(); + let num_param_format_codes = param_format_codes.len() as i16; + let num_param_values = bind.parameters.len() as i16; + 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 as usize => 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 as usize, + 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::>(); + let num_result_column_format_codes = result_columns_format_codes.len() as i16; Ok(Bind { - code, + code: 'B', portal, prepared_statement, num_param_format_codes, @@ -403,14 +412,6 @@ 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, @@ -427,47 +428,21 @@ impl TryFrom for BytesMut { 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()); - } - - Ok(bytes) + encode_frontend_message(&FrontendMessage::Bind(PgBind { + portal: Bytes::copy_from_slice(bind.portal.as_str().as_bytes()), + statement: Bytes::copy_from_slice(bind.prepared_statement.as_str().as_bytes()), + 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(), + })) } } @@ -520,6 +495,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/close.rs b/packages/cipherstash-proxy/src/postgresql/messages/close.rs index 1a8fb11d2..acd91f163 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/close.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/close.rs @@ -1,14 +1,12 @@ use crate::error::{Error, ProtocolError}; -use crate::postgresql::protocol::BytesMutReadString; -use crate::{SIZE_I32, SIZE_U8}; +use crate::postgresql::protocol::{decode_frontend_frame, encode_frontend_message}; -use bytes::{Buf, BufMut, BytesMut}; +use bytes::{Bytes, BytesMut}; +use pg_proto::codec::{Close as PgClose, DescribeTarget, FrontendMessage}; use std::convert::TryFrom; -use std::ffi::CString; -use std::io::Cursor; use super::target::Target; -use super::{FrontendCode, Name}; +use super::Name; /// /// Close b'C' (Frontend) message. @@ -37,22 +35,18 @@ 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 { + let FrontendMessage::Close(close) = decode_frontend_frame(bytes)? else { return Err(ProtocolError::UnexpectedMessageCode { - expected: FrontendCode::Close.into(), - received: code as char, + expected: 'C', + received: bytes.first().copied().unwrap_or_default() 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); + }; + let target = match close.target { + DescribeTarget::Statement => Target::Statement, + DescribeTarget::Portal => Target::Portal, + }; + let name = Name::from(String::from_utf8_lossy(&close.name).into_owned()); Ok(Close { target, name }) } @@ -62,19 +56,14 @@ 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) + let target = match close.target { + Target::Statement => DescribeTarget::Statement, + Target::Portal => DescribeTarget::Portal, + }; + encode_frontend_message(&FrontendMessage::Close(PgClose { + target, + name: Bytes::copy_from_slice(close.name.as_str().as_bytes()), + })) } } diff --git a/packages/cipherstash-proxy/src/postgresql/messages/data_row.rs b/packages/cipherstash-proxy/src/postgresql/messages/data_row.rs index 96d5cb8af..843a657af 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/data_row.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/data_row.rs @@ -1,12 +1,14 @@ -use super::{BackendCode, NULL}; use crate::EqlCiphertext; use crate::{ error::{EncryptError, Error, ProtocolError}, log::DECRYPT, - postgresql::Column, + postgresql::{ + protocol::{decode_backend_frame, encode_backend_message}, + 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 +65,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() } @@ -102,38 +102,20 @@ 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 mut bytes = BytesMut::with_capacity(len); - bytes.resize(len, 0); - cursor.copy_to_slice(&mut bytes); - - columns.push(DataColumn { bytes: Some(bytes) }); - } - } + }; + let columns = row + .columns + .into_iter() + .map(|bytes| DataColumn { + bytes: bytes.map(BytesMut::from), + }) + .collect(); Ok(DataRow { columns }) } @@ -143,39 +125,13 @@ 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) - } -} - -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) + encode_backend_message(&BackendMessage::DataRow(PgDataRow { + columns: data_row + .columns + .into_iter() + .map(|column| column.bytes.map(|bytes| bytes.freeze())) + .collect(), + })) } } diff --git a/packages/cipherstash-proxy/src/postgresql/messages/describe.rs b/packages/cipherstash-proxy/src/postgresql/messages/describe.rs index 6f0bbd8e3..88ad1107e 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/describe.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/describe.rs @@ -1,14 +1,12 @@ use crate::error::{Error, ProtocolError}; -use crate::postgresql::protocol::BytesMutReadString; -use crate::{SIZE_I32, SIZE_U8}; +use crate::postgresql::protocol::{decode_frontend_frame, encode_frontend_message}; -use bytes::{Buf, BufMut, BytesMut}; +use bytes::{Bytes, BytesMut}; +use pg_proto::codec::{Describe as PgDescribe, DescribeTarget, FrontendMessage}; use std::convert::TryFrom; -use std::ffi::CString; -use std::io::Cursor; use super::target::Target; -use super::{FrontendCode, Name}; +use super::Name; /// /// Describe b'D' (Frontend) message. @@ -37,22 +35,18 @@ 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 { + let FrontendMessage::Describe(description) = decode_frontend_frame(bytes)? else { return Err(ProtocolError::UnexpectedMessageCode { - expected: FrontendCode::Describe.into(), - received: code as char, + expected: 'D', + received: bytes.first().copied().unwrap_or_default() 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); + }; + let target = match description.target { + DescribeTarget::Statement => Target::Statement, + DescribeTarget::Portal => Target::Portal, + }; + let name = Name::from(String::from_utf8_lossy(&description.name).into_owned()); Ok(Describe { target, name }) } @@ -62,18 +56,13 @@ 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) + let target = match describe.target { + Target::Statement => DescribeTarget::Statement, + Target::Portal => DescribeTarget::Portal, + }; + encode_frontend_message(&FrontendMessage::Describe(PgDescribe { + target, + name: Bytes::copy_from_slice(describe.name.as_str().as_bytes()), + })) } } diff --git a/packages/cipherstash-proxy/src/postgresql/messages/execute.rs b/packages/cipherstash-proxy/src/postgresql/messages/execute.rs index f5e17ae29..dc6d2aacd 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/execute.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/execute.rs @@ -1,9 +1,9 @@ -use super::{FrontendCode, Name}; +use super::Name; use crate::error::{Error, ProtocolError}; -use crate::postgresql::protocol::BytesMutReadString; -use bytes::{Buf, BytesMut}; +use crate::postgresql::protocol::decode_frontend_frame; +use bytes::BytesMut; +use pg_proto::codec::FrontendMessage; use std::convert::TryFrom; -use std::io::Cursor; #[derive(Debug, Clone)] pub(crate) struct Execute { @@ -15,22 +15,15 @@ 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 { + let FrontendMessage::Execute(execute) = decode_frontend_frame(bytes)? else { return Err(ProtocolError::UnexpectedMessageCode { - expected: FrontendCode::Execute.into(), - received: code as char, + expected: 'E', + received: bytes.first().copied().unwrap_or_default() 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(); + }; + let portal = Name::from(String::from_utf8_lossy(&execute.portal).into_owned()); + let max_rows = execute.max_rows; Ok(Execute { portal, max_rows }) } diff --git a/packages/cipherstash-proxy/src/postgresql/messages/param_description.rs b/packages/cipherstash-proxy/src/postgresql/messages/param_description.rs index 6a9b4c19b..50841a226 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/param_description.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/param_description.rs @@ -1,12 +1,11 @@ -use super::BackendCode; use crate::{ error::{Error, ProtocolError}, log::MAPPER, - SIZE_I16, SIZE_I32, + postgresql::protocol::{decode_backend_frame, encode_backend_message}, }; -use bytes::{Buf, BufMut, BytesMut}; +use bytes::BytesMut; +use pg_proto::codec::BackendMessage; use postgres_types::Type; -use std::io::Cursor; use tracing::debug; /// @@ -70,29 +69,16 @@ 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, }) } @@ -102,22 +88,12 @@ 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); - } - - Ok(bytes) + let types = parameter_description + .types + .into_iter() + .map(|oid| oid as u32) + .collect(); + encode_backend_message(&BackendMessage::ParameterDescription(types)) } } diff --git a/packages/cipherstash-proxy/src/postgresql/messages/parse.rs b/packages/cipherstash-proxy/src/postgresql/messages/parse.rs index 8f7c9666d..9f3962ba9 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/parse.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/parse.rs @@ -1,13 +1,15 @@ -use super::{FrontendCode, Name, UNSPECIFIED_TYPE_OID}; +use super::{Name, UNSPECIFIED_TYPE_OID}; use crate::{ error::{Error, ProtocolError}, - postgresql::{context::statement::OutputParam, protocol::BytesMutReadString}, - SIZE_I16, SIZE_I32, + postgresql::{ + context::statement::OutputParam, + protocol::{decode_frontend_frame, encode_frontend_message}, + }, }; -use bytes::{Buf, BufMut, BytesMut}; +use bytes::{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 { @@ -92,34 +94,26 @@ 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 = Name::from(String::from_utf8_lossy(&parse.statement).into_owned()); + 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, + code: 'P', name, statement, - num_params, + num_params: param_types.len() as i16, param_types, dirty: false, }) @@ -130,30 +124,22 @@ 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); + if parse.num_params as usize != parse.param_types.len() { + return Err(ProtocolError::UnexpectedMessageLength { + code: b'P', + len: parse.param_types.len(), + } + .into()); } - - Ok(bytes) + encode_frontend_message(&FrontendMessage::Parse(PgParse { + statement: Bytes::copy_from_slice(parse.name.as_str().as_bytes()), + query: Bytes::from(parse.statement), + parameter_types: parse + .param_types + .into_iter() + .map(|oid| oid as u32) + .collect(), + })) } } diff --git a/packages/cipherstash-proxy/src/postgresql/messages/query.rs b/packages/cipherstash-proxy/src/postgresql/messages/query.rs index 88a5cd570..919de8af8 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/query.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/query.rs @@ -1,13 +1,9 @@ use crate::error::{Error, ProtocolError}; -use crate::postgresql::protocol::BytesMutReadString; -use crate::SIZE_I32; +use crate::postgresql::protocol::{decode_frontend_frame, encode_frontend_message}; -use bytes::{Buf, BufMut, BytesMut}; +use bytes::{Bytes, BytesMut}; +use pg_proto::codec::FrontendMessage; use std::convert::TryFrom; -use std::ffi::CString; -use std::io::Cursor; - -use super::FrontendCode; #[derive(Debug, Clone)] pub struct Query { @@ -38,22 +34,16 @@ 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, }) } @@ -63,17 +53,6 @@ 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); - - Ok(bytes) + encode_frontend_message(&FrontendMessage::Query(Bytes::from(query.statement))) } } diff --git a/packages/cipherstash-proxy/src/postgresql/messages/row_description.rs b/packages/cipherstash-proxy/src/postgresql/messages/row_description.rs index 7fd8aea5e..06145de98 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/row_description.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/row_description.rs @@ -1,16 +1,17 @@ -use std::{ffi::CString, io::Cursor}; - -use bytes::{Buf, BufMut, BytesMut}; +use bytes::{Bytes, BytesMut}; +use pg_proto::codec::{ + BackendMessage, FieldDescription as PgFieldDescription, RowDescription as PgRowDescription, +}; use postgres_types::Type; use crate::{ error::{Error, ProtocolError}, - postgresql::{format_code::FormatCode, protocol::BytesMutReadString}, - SIZE_I16, SIZE_I32, + postgresql::{ + format_code::FormatCode, + protocol::{decode_backend_frame, encode_backend_message}, + }, }; -use super::BackendCode; - #[derive(Debug)] pub struct RowDescription { pub fields: Vec, @@ -60,24 +61,28 @@ 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 }) } @@ -87,78 +92,21 @@ 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) - } -} - -// 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 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) + .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 })) } } diff --git a/packages/cipherstash-proxy/src/postgresql/mod.rs b/packages/cipherstash-proxy/src/postgresql/mod.rs index 71c8df249..d09a7a2ee 100644 --- a/packages/cipherstash-proxy/src/postgresql/mod.rs +++ b/packages/cipherstash-proxy/src/postgresql/mod.rs @@ -16,13 +16,3 @@ 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'; diff --git a/packages/cipherstash-proxy/src/postgresql/protocol.rs b/packages/cipherstash-proxy/src/postgresql/protocol.rs index 9c1afb79b..3f683691e 100644 --- a/packages/cipherstash-proxy/src/postgresql/protocol.rs +++ b/packages/cipherstash-proxy/src/postgresql/protocol.rs @@ -1,11 +1,13 @@ -use super::{messages::authentication::Authentication, CANCEL_REQUEST, SSL_REQUEST}; +use super::messages::authentication::Authentication; use crate::{ error::{Error, ProtocolError}, log::PROTOCOL, - postgresql::PROTOCOL_VERSION_NUMBER, SIZE_I32, SIZE_U8, }; use bytes::{BufMut, BytesMut}; +use pg_proto::codec::{ + Backend, BackendMessage, Direction, Frontend, FrontendMessage, PgCodec, DEFAULT_MAX_FRAME_LEN, +}; use std::{ io::{BufRead, Cursor}, time::Duration, @@ -14,15 +16,48 @@ use tokio::{ io::{AsyncRead, AsyncReadExt}, time::timeout, }; +use tokio_util::codec::Decoder; +use tokio_util::codec::Encoder; use tracing::{debug, error}; type Code = u8; +pub fn decode_frontend_frame(bytes: &BytesMut) -> Result { + let mut bytes = bytes.clone(); + PgCodec::::default() + .decode(&mut bytes)? + .ok_or_else(|| { + std::io::Error::new(std::io::ErrorKind::UnexpectedEof, "partial frontend frame").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::new(std::io::ErrorKind::UnexpectedEof, "partial backend frame").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) +} + #[derive(Clone, Debug, PartialEq)] pub enum StartupCode { ProtocolVersionNumber, CancelRequest, SSLRequest, + GSSENCRequest, } #[derive(Clone, Debug)] @@ -37,17 +72,6 @@ pub struct Message { 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; } @@ -76,8 +100,8 @@ pub async fn read_auth_message( 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?; + let (_code, bytes, _message) = + read_backend_message_with_timeout(&mut stream, client_id, connection_timeout).await?; Authentication::try_from(&bytes) } @@ -87,14 +111,25 @@ pub async fn read_auth_message( /// Timeout values are in config /// /// -pub async fn read_message( +pub async fn read_frontend_message( mut stream: S, client_id: i32, connection_timeout: Option, -) -> Result<(Code, BytesMut), Error> { +) -> Result<(Code, BytesMut, FrontendMessage), Error> { match connection_timeout { - Some(duration) => read_message_with_timeout(stream, client_id, duration).await, - None => read(&mut stream, client_id).await, + Some(duration) => read_frontend_message_with_timeout(stream, client_id, duration).await, + None => read::(&mut stream, client_id).await, + } +} + +pub async fn read_backend_message( + mut stream: S, + client_id: i32, + connection_timeout: Option, +) -> Result<(Code, BytesMut, BackendMessage), Error> { + match connection_timeout { + Some(duration) => read_backend_message_with_timeout(stream, client_id, duration).await, + None => read::(&mut stream, client_id).await, } } @@ -104,12 +139,22 @@ pub async fn read_message( /// Timeout values are in config /// /// -async fn read_message_with_timeout( +async fn read_frontend_message_with_timeout( mut stream: S, client_id: i32, duration: Duration, -) -> Result<(Code, BytesMut), Error> { - timeout(duration, read(&mut stream, client_id)) +) -> Result<(Code, BytesMut, FrontendMessage), Error> { + timeout(duration, read::(&mut stream, client_id)) + .await + .map_err(|_| Error::ConnectionTimeout { duration })? +} + +async fn read_backend_message_with_timeout( + mut stream: S, + client_id: i32, + duration: Duration, +) -> Result<(Code, BytesMut, BackendMessage), Error> { + timeout(duration, read::(&mut stream, client_id)) .await .map_err(|_| Error::ConnectionTimeout { duration })? } @@ -121,16 +166,16 @@ async fn read_message_with_timeout( /// Byte is then passed as `code` to this function to preserve the message structure /// /// -async fn read( +async fn read( mut stream: S, client_id: i32, -) -> Result<(Code, BytesMut), Error> { +) -> Result<(Code, BytesMut, D::Message), 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 { + if len < SIZE_I32 as i32 || len as usize + SIZE_U8 > DEFAULT_MAX_FRAME_LEN { error!( msg = "Unexpected PostgreSQL message length", code = code, @@ -138,7 +183,7 @@ async fn read( ); return Err(ProtocolError::UnexpectedMessageLength { code, - len: len as usize, + len: len.max(0) as usize, } .into()); } @@ -157,7 +202,96 @@ async fn read( stream.read_exact(&mut bytes[slice_start..]).await?; + // Direction-specific pg-proto decoding validates both the frame and message body. + // In particular, unknown tags fail closed instead of being passed through. + let mut validated = bytes.clone(); + let message = PgCodec::::default() + .decode(&mut validated)? + .ok_or_else(|| { + std::io::Error::new( + std::io::ErrorKind::UnexpectedEof, + "partial PostgreSQL frame", + ) + })?; + debug!(target: PROTOCOL, client_id, code = ?(code as char), ?bytes); - Ok((code, bytes)) + Ok((code, bytes, message)) +} + +#[cfg(test)] +mod tests { + use super::*; + use pg_proto::codec::{Authentication as PgAuthentication, BackendMessage}; + use tokio::io::{duplex, AsyncWriteExt}; + use tokio_util::codec::Encoder; + + fn encode_backend(message: BackendMessage) -> BytesMut { + let mut bytes = BytesMut::new(); + PgCodec::::default() + .encode(message.to_frame().unwrap(), &mut bytes) + .unwrap(); + bytes + } + + #[tokio::test] + async fn frontend_frame_can_arrive_in_partial_writes() { + let (mut writer, mut reader) = duplex(64); + let task = tokio::spawn(async move { + writer.write_all(b"Q\0\0").await.unwrap(); + writer.write_all(b"\0\x0dselect 1\0").await.unwrap(); + }); + + let (tag, bytes, _) = read_frontend_message(&mut reader, 1, None).await.unwrap(); + assert_eq!(tag, b'Q'); + assert_eq!(&bytes[..], b"Q\0\0\0\x0dselect 1\0"); + task.await.unwrap(); + } + + #[tokio::test] + async fn unknown_frontend_tag_is_rejected() { + let (mut writer, mut reader) = duplex(16); + writer.write_all(b"?\0\0\0\x04").await.unwrap(); + + let error = read_frontend_message(&mut reader, 1, None) + .await + .unwrap_err(); + assert!(error.to_string().contains("unknown frontend message tag")); + } + + #[tokio::test] + async fn malformed_and_oversized_frames_are_rejected_before_body_allocation() { + let (mut writer, mut reader) = duplex(16); + writer.write_all(b"Q\0\0\0\x03").await.unwrap(); + assert!(read_frontend_message(&mut reader, 1, None).await.is_err()); + + let (mut writer, mut reader) = duplex(16); + let oversized = (DEFAULT_MAX_FRAME_LEN as u32).to_be_bytes(); + writer.write_all(b"Q").await.unwrap(); + writer.write_all(&oversized).await.unwrap(); + assert!(read_frontend_message(&mut reader, 1, None).await.is_err()); + } + + #[tokio::test] + async fn authentication_modes_are_validated_by_the_backend_codec() { + let messages = [ + PgAuthentication::Ok, + PgAuthentication::CleartextPassword, + PgAuthentication::Md5Password { salt: *b"salt" }, + PgAuthentication::Sasl { + mechanisms: vec![bytes::Bytes::from_static(b"SCRAM-SHA-256")], + }, + ]; + + for authentication in messages { + let (mut writer, mut reader) = duplex(128); + writer + .write_all(&encode_backend(BackendMessage::Authentication( + authentication, + ))) + .await + .unwrap(); + read_auth_message(&mut reader, 1).await.unwrap(); + } + } } diff --git a/packages/cipherstash-proxy/src/postgresql/startup.rs b/packages/cipherstash-proxy/src/postgresql/startup.rs index a21e64f95..3f6822f28 100644 --- a/packages/cipherstash-proxy/src/postgresql/startup.rs +++ b/packages/cipherstash-proxy/src/postgresql/startup.rs @@ -1,6 +1,9 @@ use std::time::Duration; use bytes::{BufMut, BytesMut}; +use pg_proto::pre_startup::{ + decode_pre_startup, EncryptionReply, PreStartupMessage, DEFAULT_MAX_PRE_STARTUP_PACKET_LEN, +}; use tokio::{ io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}, time::timeout, @@ -11,7 +14,6 @@ use crate::{ connect::AsyncStream, error::{Error, ProtocolError}, log::PROTOCOL, - postgresql::{SSL_REQUEST, SSL_RESPONSE_NO, SSL_RESPONSE_YES}, tls, TandemConfig, SIZE_I32, }; @@ -88,6 +90,14 @@ where { let len = client.read_i32().await?; + if len < 8 || len as usize > DEFAULT_MAX_PRE_STARTUP_PACKET_LEN { + return Err(ProtocolError::UnexpectedMessageLength { + code: 0, + len: len.max(0) as usize, + } + .into()); + } + let capacity = len as usize; let mut bytes = BytesMut::with_capacity(capacity); @@ -97,20 +107,18 @@ where 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, + let mut decode_buffer = bytes.clone(); + let decoded = decode_pre_startup(&mut decode_buffer)?.ok_or_else(|| { + std::io::Error::new(std::io::ErrorKind::UnexpectedEof, "partial startup packet") + })?; + let code = match decoded { + PreStartupMessage::SslRequest => super::protocol::StartupCode::SSLRequest, + PreStartupMessage::GssEncRequest => super::protocol::StartupCode::GSSENCRequest, + PreStartupMessage::CancelRequest { .. } => super::protocol::StartupCode::CancelRequest, + PreStartupMessage::Startup(_) => super::protocol::StartupCode::ProtocolVersionNumber, }; + + let message = StartupMessage { code, bytes }; debug!(target: PROTOCOL, StartupMessage = ?message); Ok(message) @@ -123,17 +131,19 @@ where 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?; + stream + .write_all(&PreStartupMessage::SslRequest.to_packet()?) + .await?; // Server supports TLS - let response = match stream.read_u8().await? { - SSL_RESPONSE_YES => true, - SSL_RESPONSE_NO => false, - code => { + let response = match EncryptionReply::try_from(stream.read_u8().await?) { + Ok(EncryptionReply::Accepted) => true, + Ok(EncryptionReply::Rejected) => false, + Ok(EncryptionReply::LegacyError) => { + return Err(ProtocolError::UnexpectedStartupMessage.into()); + } + Err(err) => { + let code = err.0; error!(msg = "Unexpected startup message", code = ?(code as char)); return Err(ProtocolError::UnexpectedStartupMessage.into()); } @@ -153,11 +163,57 @@ pub async fn send_ssl_response( stream: &mut T, tls: bool, ) -> Result<(), Error> { - let response = if tls { b'S' } else { b'N' }; + let response = if tls { + EncryptionReply::Accepted + } else { + EncryptionReply::Rejected + }; debug!(target: PROTOCOL, msg = "SSLResponse to Client", SSLResponse = ?response); - stream.write_all(&[response]).await?; + stream.write_all(&[response.as_byte()]).await?; Ok(()) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::postgresql::protocol::StartupCode; + use tokio::io::{duplex, AsyncReadExt, AsyncWriteExt}; + + #[tokio::test] + async fn ssl_and_cancellation_packets_use_pg_proto_pre_startup_decoding() { + let (mut writer, mut reader) = duplex(64); + writer + .write_all(&PreStartupMessage::SslRequest.to_packet().unwrap()) + .await + .unwrap(); + let ssl = read_message(&mut reader, None).await.unwrap(); + assert_eq!(ssl.code, StartupCode::SSLRequest); + + let cancel = PreStartupMessage::CancelRequest { + process_id: 42, + secret_key: bytes::Bytes::from_static(b"key!"), + } + .to_packet() + .unwrap(); + writer.write_all(&cancel).await.unwrap(); + let decoded = read_message(&mut reader, None).await.unwrap(); + assert_eq!(decoded.code, StartupCode::CancelRequest); + assert_eq!(&decoded.bytes[..], &cancel[..]); + } + + #[tokio::test] + async fn ssl_reply_rejects_unknown_bytes() { + let (mut writer, mut reader) = duplex(8); + writer.write_all(b"?").await.unwrap(); + assert!(send_ssl_request(&mut reader).await.is_err()); + + let (mut client, mut server) = duplex(8); + send_ssl_response(&mut server, true).await.unwrap(); + let mut response = [0]; + client.read_exact(&mut response).await.unwrap(); + assert_eq!(response, [EncryptionReply::Accepted.as_byte()]); + } +} From a5c106492a7a3a2531185630ddfcfcd2771f4308 Mon Sep 17 00:00:00 2001 From: James Sadler Date: Tue, 4 Aug 2026 16:56:41 +1000 Subject: [PATCH 02/16] refactor(proxy): remove legacy protocol plumbing --- .../src/postgresql/backend.rs | 23 +- .../src/postgresql/frontend.rs | 24 +- .../src/postgresql/handler.rs | 239 +++++++----- .../messages/authentication/auth.rs | 358 ------------------ .../postgresql/messages/authentication/mod.rs | 6 - .../messages/authentication/sasl.rs | 133 ------- .../src/postgresql/messages/bind.rs | 35 +- .../src/postgresql/messages/close.rs | 18 +- .../src/postgresql/messages/describe.rs | 18 +- .../src/postgresql/messages/error_response.rs | 119 +++--- .../src/postgresql/messages/mod.rs | 279 +------------- .../src/postgresql/messages/parse.rs | 14 +- .../postgresql/messages/ready_for_query.rs | 21 - .../src/postgresql/messages/terminate.rs | 15 - .../src/postgresql/protocol.rs | 182 +++------ .../src/postgresql/startup.rs | 63 +-- 16 files changed, 296 insertions(+), 1251 deletions(-) delete mode 100644 packages/cipherstash-proxy/src/postgresql/messages/authentication/auth.rs delete mode 100644 packages/cipherstash-proxy/src/postgresql/messages/authentication/mod.rs delete mode 100644 packages/cipherstash-proxy/src/postgresql/messages/authentication/sasl.rs delete mode 100644 packages/cipherstash-proxy/src/postgresql/messages/ready_for_query.rs delete mode 100644 packages/cipherstash-proxy/src/postgresql/messages/terminate.rs diff --git a/packages/cipherstash-proxy/src/postgresql/backend.rs b/packages/cipherstash-proxy/src/postgresql/backend.rs index 36aa9c58d..b68400dcf 100644 --- a/packages/cipherstash-proxy/src/postgresql/backend.rs +++ b/packages/cipherstash-proxy/src/postgresql/backend.rs @@ -151,9 +151,8 @@ where /// 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_message) = protocol::read_backend_message( + let (mut bytes, protocol_message) = protocol::read_backend_message( &mut self.server_reader, - self.context.client_id, self.context.connection_timeout(), ) .await?; @@ -171,7 +170,7 @@ where client_id = self.context.client_id, msg = "Slow database response", duration_ms = read_duration.as_millis(), - message_code = ?code, + message = ?protocol_message, ); } @@ -235,8 +234,8 @@ where self.context.complete_execution(); self.context.finish_session(); } - BackendMessage::ErrorResponse(_) => { - if let Some(b) = self.error_response_handler(&bytes)? { + BackendMessage::ErrorResponse(ref response) => { + if let Some(b) = self.error_response_handler(response, &bytes) { bytes = b } @@ -292,7 +291,7 @@ where debug!(target: PROTOCOL, client_id = self.context.client_id, msg = "Passthrough", - ?code, + message = ?protocol_message, ); } } @@ -338,11 +337,15 @@ 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, + bytes: &BytesMut, + ) -> Option { + 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())) + Some(bytes.to_owned()) } /// @@ -730,7 +733,7 @@ where // Ensure any buffered data is cleared before sending error self.buffer.clear(); - let message = BytesMut::try_from(error_response)?; + let message = protocol::encode_backend_message(&error_response.into_backend_message())?; debug!( target: "PROTOCOL", diff --git a/packages/cipherstash-proxy/src/postgresql/frontend.rs b/packages/cipherstash-proxy/src/postgresql/frontend.rs index 24d35fa2c..f04f58024 100644 --- a/packages/cipherstash-proxy/src/postgresql/frontend.rs +++ b/packages/cipherstash-proxy/src/postgresql/frontend.rs @@ -6,7 +6,6 @@ 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 crate::connect::Sender; @@ -22,8 +21,6 @@ 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::prometheus::{ CLIENTS_BYTES_RECEIVED_TOTAL, ENCRYPTED_VALUES_TOTAL, ENCRYPTION_DURATION_SECONDS, @@ -38,7 +35,7 @@ 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::FrontendMessage; +use pg_proto::codec::{BackendMessage, FrontendMessage, TransactionStatus}; use serde::Serialize; use sqltk::parser::ast::{self, Value}; use sqltk::NodeKey; @@ -165,9 +162,8 @@ where /// 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_message) = protocol::read_frontend_message( + let (mut bytes, protocol_message) = protocol::read_frontend_message( &mut self.client_reader, - self.context.client_id, self.context.connection_timeout(), ) .await?; @@ -182,15 +178,13 @@ where return Ok(()); } - let code = Code::from(code); - // 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() { warn!(target: PROTOCOL, client_id = self.context.client_id, error_state = ?self.error_state, - ?code, + message = ?protocol_message, ); if !matches!(protocol_message, FrontendMessage::Sync) { return Ok(()); @@ -274,7 +268,7 @@ where FrontendMessage::Sync => { debug!(target: PROTOCOL, client_id = self.context.client_id, - ?code, + message = ?protocol_message, ); self.context.reload_schema_if_changed().await; @@ -295,7 +289,7 @@ where debug!(target: PROTOCOL, client_id = self.context.client_id, msg = "Passthrough", - ?code, + message = ?protocol_message, ); } } @@ -322,7 +316,7 @@ where pub async fn terminate(&mut self) -> Result<(), Error> { debug!(target: PROTOCOL, msg = "Terminate server connection"); - let bytes = Terminate::message(); + let bytes = protocol::encode_frontend_message(&FrontendMessage::Terminate)?; self.write_to_server(bytes).await?; Ok(()) } @@ -1199,7 +1193,9 @@ 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); + let message = protocol::encode_backend_message(&BackendMessage::ReadyForQuery( + TransactionStatus::Idle, + ))?; debug!(target: PROTOCOL, client_id = self.context.client_id, @@ -1350,7 +1346,7 @@ where fn send_error_response(&mut self, err: Error) -> Result<(), Error> { let error_response = self.error_to_response(err); - let message = BytesMut::try_from(error_response)?; + let message = protocol::encode_backend_message(&error_response.into_backend_message())?; debug!(target: PROTOCOL, client_id = self.context.client_id, diff --git a/packages/cipherstash-proxy/src/postgresql/handler.rs b/packages/cipherstash-proxy/src/postgresql/handler.rs index b8c593454..f91781dff 100644 --- a/packages/cipherstash-proxy/src/postgresql/handler.rs +++ b/packages/cipherstash-proxy/src/postgresql/handler.rs @@ -1,14 +1,8 @@ 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; @@ -18,40 +12,27 @@ use crate::{ postgresql::context::Context, tls, }; -use bytes::BytesMut; +use bytes::{BufMut, Bytes, BytesMut}; use md5::{Digest, Md5}; +use pg_proto::codec::{Authentication, BackendMessage, FrontendMessage}; +use pg_proto::pre_startup::PreStartupMessage; use postgres_protocol::authentication::sasl::{ChannelBinding, ScramSha256}; use rand::Rng; use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt}; 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"; + +#[derive(Debug, Clone, Copy, PartialEq)] +enum SaslMechanism { + ScramSha256, + ScramSha256Plus, +} +/// Handles one downstream PostgreSQL connection and its paired upstream connection. /// -/// -/// 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 -/// -/// +/// 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: AsyncStream, context: Context) -> Result<(), Error> { let mut client_stream = client_stream; let client_id = context.client_id; @@ -76,8 +57,8 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R Err(err) => return Err(err), }; - match &startup_message.code { - StartupCode::SSLRequest => { + match &startup_message { + PreStartupMessage::SslRequest => { startup::send_ssl_response(&mut client_stream, context.use_tls()).await?; if let Some(ref tls) = context.tls_config() { match client_stream { @@ -92,15 +73,19 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R } } } - StartupCode::CancelRequest => { - database_stream.write_all(&startup_message.bytes).await?; + PreStartupMessage::CancelRequest { .. } => { + database_stream + .write_all(&startup_message.to_packet()?) + .await?; return Err(Error::CancelRequest); } - StartupCode::ProtocolVersionNumber => { - database_stream.write_all(&startup_message.bytes).await?; + PreStartupMessage::Startup(_) => { + database_stream + .write_all(&startup_message.to_packet()?) + .await?; break; } - StartupCode::GSSENCRequest => { + PreStartupMessage::GssEncRequest => { return Err(ProtocolError::UnexpectedStartupMessage.into()); } } @@ -123,39 +108,47 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R let hash = md5_hash(username, password, &salt); - let message = Authentication::md5_password(salt); - let bytes = BytesMut::try_from(message)?; + let bytes = protocol::encode_backend_message(&BackendMessage::Authentication( + Authentication::Md5Password { salt }, + ))?; client_stream.write_all(&bytes).await?; let connection_timeout = context.connection_timeout(); - let (_code, bytes, _message) = match protocol::read_frontend_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); + let (_bytes, message) = + match protocol::read_frontend_message(&mut client_stream, 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 FrontendMessage::PasswordResponse(password) = message else { + return Err(ProtocolError::UnexpectedAuthenticationResponse { + expected: "PasswordResponse".into(), + received: -1, } - Err(err) => return Err(err), + .into()); }; + let password = password + .strip_suffix(&[0]) + .ok_or(ProtocolError::UnexpectedStartupMessage)?; + let password = + std::str::from_utf8(password).map_err(|_| ProtocolError::AuthenticationFailed)?; - let password_message = PasswordMessage::try_from(&bytes)?; - - if hash == password_message.password { - let message = Authentication::authentication_ok(); + if hash == password { debug!(target: AUTHENTICATION, msg = "Client AuthenticationOk"); - let bytes = BytesMut::try_from(message)?; + let bytes = protocol::encode_backend_message(&BackendMessage::Authentication( + Authentication::Ok, + ))?; 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)?; + let bytes = protocol::encode_backend_message(&message.into_backend_message())?; client_stream.write_all(&bytes).await?; } } @@ -169,33 +162,31 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R // // First message should always be Auth - let auth = protocol::read_auth_message(&mut database_stream, client_id).await?; + let auth = protocol::read_auth_message(&mut database_stream).await?; - match &auth.method { - AuthenticationMethod::AuthenticationOk => { + match &auth { + Authentication::Ok => { debug!(target: AUTHENTICATION, msg = "AuthenticationOk"); } - AuthenticationMethod::AuthenticationCleartextPassword => { + Authentication::CleartextPassword => { debug!(target: AUTHENTICATION, msg = "AuthenticationCleartextPassword"); let password = context.database_password(); - let message = PasswordMessage::new(password); - let bytes = BytesMut::try_from(message)?; + let bytes = password_message(password)?; database_stream.write_all(&bytes).await?; } - AuthenticationMethod::Md5Password { salt } => { + Authentication::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)?; + let bytes = password_message(hash)?; database_stream.write_all(&bytes).await?; } - AuthenticationMethod::Sasl { .. } => { + Authentication::Sasl { mechanisms } => { debug!(target: AUTHENTICATION, msg = "Sasl"); - let mechanism = auth.sasl_mechanism()?; + let mechanism = sasl_mechanism(mechanisms)?; sanity_check_sasl_mechanism(&mechanism, &client_stream); // Toby: I don't think we need to do anything here @@ -207,22 +198,25 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R scram_sha_256_plus_handler(&mut database_stream, mechanism, password, channel_binding) .await?; } - AuthenticationMethod::Other { method_code, .. } => { + Authentication::KerberosV5 + | Authentication::Gss + | Authentication::GssContinue(_) + | Authentication::Sspi => { debug!(target: AUTHENTICATION, msg = "UnsupportedAuthentication"); return Err(ProtocolError::UnsupportedAuthentication { - method_code: *method_code, + method_code: authentication_method_code(&auth), } .into()); } - method => { - debug!(target: AUTHENTICATION, msg = "UnexpectedStartupMessage", authentication_method = ?method); + Authentication::SaslContinue(_) | Authentication::SaslFinal(_) => { + debug!(target: AUTHENTICATION, msg = "UnexpectedStartupMessage", authentication_method = ?auth); return Err(ProtocolError::UnexpectedStartupMessage.into()); } } if context.require_tls() && !client_stream.is_tls() { let message = ErrorResponse::tls_required(); - let bytes = BytesMut::try_from(message)?; + let bytes = protocol::encode_backend_message(&message.into_backend_message())?; client_stream.write_all(&bytes).await?; error!(msg = "Client must connect with Transport Layer Security (TLS)"); @@ -286,7 +280,8 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R 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) { + if let Ok(bytes) = protocol::encode_backend_message(&error_response.into_backend_message()) + { let _ = timeout_sender.send(bytes); } // Best-effort yield to allow ChannelWriter to flush the error response @@ -338,6 +333,40 @@ fn sanity_check_sasl_mechanism(mechanism: &SaslMechanism, client_stream: &AsyncS } } +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()), + } +} + +fn authentication_method_code(authentication: &Authentication) -> i32 { + match authentication { + Authentication::Ok => 0, + Authentication::KerberosV5 => 2, + Authentication::CleartextPassword => 3, + Authentication::Md5Password { .. } => 5, + Authentication::Gss => 7, + Authentication::GssContinue(_) => 8, + Authentication::Sspi => 9, + Authentication::Sasl { .. } => 10, + Authentication::SaslContinue(_) => 11, + Authentication::SaslFinal(_) => 12, + } +} + +fn password_message(password: String) -> Result { + let password = std::ffi::CString::new(password)?; + protocol::encode_frontend_message(&FrontendMessage::PasswordResponse(Bytes::copy_from_slice( + password.as_bytes_with_nul(), + ))) +} + pub fn md5_hash(username: &[u8], password: &[u8], salt: &[u8; 4]) -> String { let mut md5 = Md5::new(); md5.update(password); @@ -364,27 +393,53 @@ async fn scram_sha_256_plus_handler( 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)?; + let mechanism = match mechanism { + SaslMechanism::ScramSha256 => SCRAM_SHA_256, + SaslMechanism::ScramSha256Plus => SCRAM_SHA_256_PLUS, + }; + let mut initial = BytesMut::new(); + initial.extend_from_slice(mechanism); + initial.put_u8(0); + initial.put_i32(bytes.len().try_into().map_err(|_| { + std::io::Error::new( + std::io::ErrorKind::InvalidInput, + "SASL response is too large", + ) + })?); + initial.extend_from_slice(&bytes); + let bytes = + protocol::encode_frontend_message(&FrontendMessage::PasswordResponse(initial.freeze()))?; stream.write_all(&bytes).await?; - let auth = protocol::read_auth_message(&mut stream, 1).await?; + let auth = protocol::read_auth_message(&mut stream).await?; - let bytes = auth.sasl_continue()?; - scram.update(bytes)?; - - let sasl_response = SASLResponse::new(scram.message().to_vec()); + let Authentication::SaslContinue(bytes) = auth else { + return Err(ProtocolError::UnexpectedAuthenticationResponse { + expected: "SaslContinue".into(), + received: authentication_method_code(&auth), + } + .into()); + }; + scram.update(&bytes)?; - let bytes = BytesMut::try_from(sasl_response)?; + let bytes = protocol::encode_frontend_message(&FrontendMessage::PasswordResponse( + Bytes::copy_from_slice(scram.message()), + ))?; 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).await?; + let Authentication::SaslFinal(bytes) = auth else { + return Err(ProtocolError::UnexpectedAuthenticationResponse { + expected: "SaslFinal".into(), + received: authentication_method_code(&auth), + } + .into()); + }; + scram.finish(&bytes)?; - let auth = protocol::read_auth_message(&mut stream, 1).await?; + let auth = protocol::read_auth_message(&mut stream).await?; - if auth.is_ok() { + if matches!(auth, Authentication::Ok) { debug!(target: AUTHENTICATION, msg = "SASL authentication successful"); Ok(()) } else { @@ -396,7 +451,7 @@ async fn scram_sha_256_plus_handler( /// 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) { + if let Ok(bytes) = protocol::encode_backend_message(&error_response.into_backend_message()) { let _ = stream.write_all(&bytes).await; } } 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/bind.rs b/packages/cipherstash-proxy/src/postgresql/messages/bind.rs index 7509101d9..94563c28c 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/bind.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/bind.rs @@ -23,14 +23,10 @@ use tracing::debug; /// 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 @@ -184,8 +180,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; @@ -362,17 +356,16 @@ impl TryFrom<&BytesMut> for Bind { .copied() .map(FormatCode::from) .collect::>(); - let num_param_format_codes = param_format_codes.len() as i16; - let num_param_values = bind.parameters.len() as i16; + 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], - len if len == num_param_values as usize => param_format_codes[idx], + len if len == num_param_values => param_format_codes[idx], _ => { return Err(ProtocolError::ParameterFormatCodesMismatch { - expected: num_param_values as usize, + expected: num_param_values, received: param_format_codes.len(), } .into()) @@ -391,17 +384,11 @@ impl TryFrom<&BytesMut> for Bind { .copied() .map(FormatCode::from) .collect::>(); - let num_result_column_format_codes = result_columns_format_codes.len() as i16; - Ok(Bind { - code: 'B', 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, }) @@ -412,22 +399,6 @@ impl TryFrom for BytesMut { type Error = Error; fn try_from(bind: Bind) -> Result { - 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()); - } - encode_frontend_message(&FrontendMessage::Bind(PgBind { portal: Bytes::copy_from_slice(bind.portal.as_str().as_bytes()), statement: Bytes::copy_from_slice(bind.prepared_statement.as_str().as_bytes()), diff --git a/packages/cipherstash-proxy/src/postgresql/messages/close.rs b/packages/cipherstash-proxy/src/postgresql/messages/close.rs index acd91f163..b3ee2c55c 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/close.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/close.rs @@ -8,23 +8,7 @@ use std::convert::TryFrom; use super::target::Target; use super::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). - +/// Proxy state extracted from a typed frontend `Close` message. #[derive(Debug, Clone)] pub(crate) struct Close { pub target: Target, diff --git a/packages/cipherstash-proxy/src/postgresql/messages/describe.rs b/packages/cipherstash-proxy/src/postgresql/messages/describe.rs index 88ad1107e..20af83058 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/describe.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/describe.rs @@ -8,23 +8,7 @@ use std::convert::TryFrom; use super::target::Target; use super::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). - +/// Proxy state extracted from a typed frontend `Describe` message. #[derive(Debug, Clone)] pub struct Describe { pub target: Target, diff --git a/packages/cipherstash-proxy/src/postgresql/messages/error_response.rs b/packages/cipherstash-proxy/src/postgresql/messages/error_response.rs index 7013fe1f4..c54bd9764 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/error_response.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/error_response.rs @@ -1,13 +1,8 @@ -use super::BackendCode; -use crate::error::{Error, ProtocolError}; -use crate::postgresql::protocol::BytesMutReadString; -use crate::SIZE_I32; -use bytes::{Buf, BufMut, BytesMut}; +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 +55,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 +291,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 +341,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() { @@ -493,7 +454,9 @@ impl From for ErrorResponseCode { mod tests { use super::ErrorResponseCode; use crate::postgresql::messages::error_response::ErrorResponse; + use crate::postgresql::protocol::{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 +466,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/messages/mod.rs b/packages/cipherstash-proxy/src/postgresql/messages/mod.rs index b271c001c..df7d98661 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/mod.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/mod.rs @@ -1,8 +1,5 @@ -use std::fmt; - use bytes::BytesMut; -pub mod authentication; pub mod bind; pub mod close; pub mod data_row; @@ -13,12 +10,9 @@ 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; @@ -28,277 +22,12 @@ pub const NULL: i32 = -1; /// 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 -/// +/// Returns whether a text value may contain a JSON object. pub fn maybe_json(bytes: &BytesMut) -> bool { - if bytes.is_empty() { - return false; - } - - let b = bytes.as_ref()[0]; - b == b'{' + bytes.first() == Some(&b'{') } -/// -/// Postgres binary json is regular json with a leading header byte -/// The header byte is always 1 -/// +/// Returns whether a binary value may contain a JSONB object. 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'{' + bytes.len() > 3 && bytes[0] == 1 && bytes[1] == b'{' } diff --git a/packages/cipherstash-proxy/src/postgresql/messages/parse.rs b/packages/cipherstash-proxy/src/postgresql/messages/parse.rs index 9f3962ba9..aef0d6b4d 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/parse.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/parse.rs @@ -13,10 +13,8 @@ use postgres_types::Type; #[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, } @@ -78,7 +76,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; } @@ -110,10 +107,8 @@ impl TryFrom<&BytesMut> for Parse { .collect::>(); Ok(Parse { - code: 'P', name, statement, - num_params: param_types.len() as i16, param_types, dirty: false, }) @@ -124,13 +119,6 @@ impl TryFrom for BytesMut { type Error = Error; fn try_from(parse: Parse) -> Result { - if parse.num_params as usize != parse.param_types.len() { - return Err(ProtocolError::UnexpectedMessageLength { - code: b'P', - len: parse.param_types.len(), - } - .into()); - } encode_frontend_message(&FrontendMessage::Parse(PgParse { statement: Bytes::copy_from_slice(parse.name.as_str().as_bytes()), query: Bytes::from(parse.statement), @@ -236,7 +224,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/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/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/protocol.rs b/packages/cipherstash-proxy/src/postgresql/protocol.rs index 3f683691e..f35173b73 100644 --- a/packages/cipherstash-proxy/src/postgresql/protocol.rs +++ b/packages/cipherstash-proxy/src/postgresql/protocol.rs @@ -1,26 +1,15 @@ -use super::messages::authentication::Authentication; -use crate::{ - error::{Error, ProtocolError}, - log::PROTOCOL, - SIZE_I32, SIZE_U8, -}; -use bytes::{BufMut, BytesMut}; +use crate::error::{Error, ProtocolError}; +use bytes::BytesMut; use pg_proto::codec::{ - Backend, BackendMessage, Direction, Frontend, FrontendMessage, PgCodec, DEFAULT_MAX_FRAME_LEN, -}; -use std::{ - io::{BufRead, Cursor}, - time::Duration, + Authentication, Backend, BackendMessage, Direction, Frontend, FrontendMessage, PgCodec, }; +use std::time::Duration; use tokio::{ io::{AsyncRead, AsyncReadExt}, time::timeout, }; use tokio_util::codec::Decoder; use tokio_util::codec::Encoder; -use tracing::{debug, error}; - -type Code = u8; pub fn decode_frontend_frame(bytes: &BytesMut) -> Result { let mut bytes = bytes.clone(); @@ -52,42 +41,6 @@ pub fn encode_backend_message(message: &BackendMessage) -> Result 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 /// @@ -97,12 +50,18 @@ impl BytesMutReadString for Cursor<&BytesMut> { /// pub async fn read_auth_message( mut stream: S, - client_id: i32, ) -> Result { let connection_timeout = Duration::from_millis(1000 * 10); - let (_code, bytes, _message) = - read_backend_message_with_timeout(&mut stream, client_id, connection_timeout).await?; - Authentication::try_from(&bytes) + let (_bytes, message) = + read_backend_message_with_timeout(&mut stream, connection_timeout).await?; + match message { + BackendMessage::Authentication(authentication) => Ok(authentication), + _ => Err(ProtocolError::UnexpectedAuthenticationResponse { + expected: "Authentication".into(), + received: -1, + } + .into()), + } } /// @@ -113,23 +72,21 @@ pub async fn read_auth_message( /// pub async fn read_frontend_message( mut stream: S, - client_id: i32, connection_timeout: Option, -) -> Result<(Code, BytesMut, FrontendMessage), Error> { +) -> Result<(BytesMut, FrontendMessage), Error> { match connection_timeout { - Some(duration) => read_frontend_message_with_timeout(stream, client_id, duration).await, - None => read::(&mut stream, client_id).await, + Some(duration) => read_frontend_message_with_timeout(stream, duration).await, + None => read_frontend(&mut stream).await, } } pub async fn read_backend_message( mut stream: S, - client_id: i32, connection_timeout: Option, -) -> Result<(Code, BytesMut, BackendMessage), Error> { +) -> Result<(BytesMut, BackendMessage), Error> { match connection_timeout { - Some(duration) => read_backend_message_with_timeout(stream, client_id, duration).await, - None => read::(&mut stream, client_id).await, + Some(duration) => read_backend_message_with_timeout(stream, duration).await, + None => read_backend(&mut stream).await, } } @@ -141,20 +98,18 @@ pub async fn read_backend_message( /// async fn read_frontend_message_with_timeout( mut stream: S, - client_id: i32, duration: Duration, -) -> Result<(Code, BytesMut, FrontendMessage), Error> { - timeout(duration, read::(&mut stream, client_id)) +) -> Result<(BytesMut, FrontendMessage), Error> { + timeout(duration, read_frontend(&mut stream)) .await .map_err(|_| Error::ConnectionTimeout { duration })? } async fn read_backend_message_with_timeout( mut stream: S, - client_id: i32, duration: Duration, -) -> Result<(Code, BytesMut, BackendMessage), Error> { - timeout(duration, read::(&mut stream, client_id)) +) -> Result<(BytesMut, BackendMessage), Error> { + timeout(duration, read_backend(&mut stream)) .await .map_err(|_| Error::ConnectionTimeout { duration })? } @@ -166,63 +121,41 @@ async fn read_backend_message_with_timeout( /// 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, D::Message), Error> { - let code = stream.read_u8().await?; - let len = stream.read_i32().await?; +async fn read_frontend( + stream: &mut S, +) -> Result<(BytesMut, FrontendMessage), Error> { + let message = read::(stream).await?; + let bytes = encode_frontend_message(&message)?; + Ok((bytes, message)) +} + +async fn read_backend( + stream: &mut S, +) -> Result<(BytesMut, BackendMessage), Error> { + let message = read::(stream).await?; + let bytes = encode_backend_message(&message)?; + Ok((bytes, message)) +} - // Detect unexpected message len and avoid panic on read_exact - // Len must be at least 4 bytes (4 bytes for len/i32) - if len < SIZE_I32 as i32 || len as usize + SIZE_U8 > DEFAULT_MAX_FRAME_LEN { - error!( - msg = "Unexpected PostgreSQL message length", - code = code, - len = len - ); - return Err(ProtocolError::UnexpectedMessageLength { - code, - len: len.max(0) as usize, +async fn read(stream: &mut S) -> Result { + let mut codec = PgCodec::::default(); + let mut bytes = BytesMut::with_capacity(5); + loop { + if let Some(message) = codec.decode(&mut bytes)? { + return Ok(message); + } + if stream.read_buf(&mut bytes).await? == 0 { + return Err(Error::ConnectionClosed); } - .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?; - - // Direction-specific pg-proto decoding validates both the frame and message body. - // In particular, unknown tags fail closed instead of being passed through. - let mut validated = bytes.clone(); - let message = PgCodec::::default() - .decode(&mut validated)? - .ok_or_else(|| { - std::io::Error::new( - std::io::ErrorKind::UnexpectedEof, - "partial PostgreSQL frame", - ) - })?; - - debug!(target: PROTOCOL, client_id, code = ?(code as char), ?bytes); - - Ok((code, bytes, message)) } #[cfg(test)] mod tests { use super::*; - use pg_proto::codec::{Authentication as PgAuthentication, BackendMessage}; + use pg_proto::codec::{ + Authentication as PgAuthentication, BackendMessage, DEFAULT_MAX_FRAME_LEN, + }; use tokio::io::{duplex, AsyncWriteExt}; use tokio_util::codec::Encoder; @@ -242,8 +175,7 @@ mod tests { writer.write_all(b"\0\x0dselect 1\0").await.unwrap(); }); - let (tag, bytes, _) = read_frontend_message(&mut reader, 1, None).await.unwrap(); - assert_eq!(tag, b'Q'); + let (bytes, _) = read_frontend_message(&mut reader, None).await.unwrap(); assert_eq!(&bytes[..], b"Q\0\0\0\x0dselect 1\0"); task.await.unwrap(); } @@ -253,9 +185,7 @@ mod tests { let (mut writer, mut reader) = duplex(16); writer.write_all(b"?\0\0\0\x04").await.unwrap(); - let error = read_frontend_message(&mut reader, 1, None) - .await - .unwrap_err(); + let error = read_frontend_message(&mut reader, None).await.unwrap_err(); assert!(error.to_string().contains("unknown frontend message tag")); } @@ -263,13 +193,13 @@ mod tests { async fn malformed_and_oversized_frames_are_rejected_before_body_allocation() { let (mut writer, mut reader) = duplex(16); writer.write_all(b"Q\0\0\0\x03").await.unwrap(); - assert!(read_frontend_message(&mut reader, 1, None).await.is_err()); + assert!(read_frontend_message(&mut reader, None).await.is_err()); let (mut writer, mut reader) = duplex(16); let oversized = (DEFAULT_MAX_FRAME_LEN as u32).to_be_bytes(); writer.write_all(b"Q").await.unwrap(); writer.write_all(&oversized).await.unwrap(); - assert!(read_frontend_message(&mut reader, 1, None).await.is_err()); + assert!(read_frontend_message(&mut reader, None).await.is_err()); } #[tokio::test] @@ -291,7 +221,7 @@ mod tests { ))) .await .unwrap(); - read_auth_message(&mut reader, 1).await.unwrap(); + read_auth_message(&mut reader).await.unwrap(); } } } diff --git a/packages/cipherstash-proxy/src/postgresql/startup.rs b/packages/cipherstash-proxy/src/postgresql/startup.rs index 3f6822f28..eabb6261f 100644 --- a/packages/cipherstash-proxy/src/postgresql/startup.rs +++ b/packages/cipherstash-proxy/src/postgresql/startup.rs @@ -1,9 +1,7 @@ use std::time::Duration; -use bytes::{BufMut, BytesMut}; -use pg_proto::pre_startup::{ - decode_pre_startup, EncryptionReply, PreStartupMessage, DEFAULT_MAX_PRE_STARTUP_PACKET_LEN, -}; +use bytes::BytesMut; +use pg_proto::pre_startup::{decode_pre_startup, EncryptionReply, PreStartupMessage}; use tokio::{ io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}, time::timeout, @@ -14,11 +12,9 @@ use crate::{ connect::AsyncStream, error::{Error, ProtocolError}, log::PROTOCOL, - tls, TandemConfig, SIZE_I32, + tls, TandemConfig, }; -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)"); @@ -56,7 +52,7 @@ pub async fn with_tls(stream: AsyncStream, config: &TandemConfig) -> Result( mut stream: S, connection_timeout: Option, -) -> Result { +) -> Result { match connection_timeout { Some(duration) => read_message_with_timeout(stream, duration).await, None => read(&mut stream).await, @@ -72,7 +68,7 @@ pub async fn read_message( async fn read_message_with_timeout( mut stream: S, duration: Duration, -) -> Result { +) -> Result { timeout(duration, read(&mut stream)) .await .map_err(|_| Error::ConnectionTimeout { duration })? @@ -84,44 +80,20 @@ async fn read_message_with_timeout( /// /// /// -async fn read(client: &mut C) -> Result +async fn read(client: &mut C) -> Result where C: AsyncRead + Unpin, { - let len = client.read_i32().await?; - - if len < 8 || len as usize > DEFAULT_MAX_PRE_STARTUP_PACKET_LEN { - return Err(ProtocolError::UnexpectedMessageLength { - code: 0, - len: len.max(0) as usize, + let mut bytes = BytesMut::with_capacity(4); + loop { + if let Some(message) = decode_pre_startup(&mut bytes)? { + debug!(target: PROTOCOL, pre_startup = ?message); + return Ok(message); + } + if client.read_buf(&mut bytes).await? == 0 { + return Err(Error::ConnectionClosed); } - .into()); } - - 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?; - - let mut decode_buffer = bytes.clone(); - let decoded = decode_pre_startup(&mut decode_buffer)?.ok_or_else(|| { - std::io::Error::new(std::io::ErrorKind::UnexpectedEof, "partial startup packet") - })?; - let code = match decoded { - PreStartupMessage::SslRequest => super::protocol::StartupCode::SSLRequest, - PreStartupMessage::GssEncRequest => super::protocol::StartupCode::GSSENCRequest, - PreStartupMessage::CancelRequest { .. } => super::protocol::StartupCode::CancelRequest, - PreStartupMessage::Startup(_) => super::protocol::StartupCode::ProtocolVersionNumber, - }; - - let message = StartupMessage { code, bytes }; - debug!(target: PROTOCOL, StartupMessage = ?message); - - Ok(message) } /// @@ -179,7 +151,6 @@ pub async fn send_ssl_response( #[cfg(test)] mod tests { use super::*; - use crate::postgresql::protocol::StartupCode; use tokio::io::{duplex, AsyncReadExt, AsyncWriteExt}; #[tokio::test] @@ -190,7 +161,7 @@ mod tests { .await .unwrap(); let ssl = read_message(&mut reader, None).await.unwrap(); - assert_eq!(ssl.code, StartupCode::SSLRequest); + assert!(matches!(ssl, PreStartupMessage::SslRequest)); let cancel = PreStartupMessage::CancelRequest { process_id: 42, @@ -200,8 +171,8 @@ mod tests { .unwrap(); writer.write_all(&cancel).await.unwrap(); let decoded = read_message(&mut reader, None).await.unwrap(); - assert_eq!(decoded.code, StartupCode::CancelRequest); - assert_eq!(&decoded.bytes[..], &cancel[..]); + assert!(matches!(decoded, PreStartupMessage::CancelRequest { .. })); + assert_eq!(decoded.to_packet().unwrap(), cancel); } #[tokio::test] From 75bd1f3db59817dfff62c4beef0b4240183a8309 Mon Sep 17 00:00:00 2001 From: James Sadler Date: Tue, 4 Aug 2026 23:15:52 +1000 Subject: [PATCH 03/16] refactor(proxy): delegate pipelining to pg-proto --- Cargo.lock | 8 +- PG_PROTO_MIGRATION_PLAN.md | 2 +- packages/cipherstash-proxy/Cargo.toml | 2 +- .../src/postgresql/backend.rs | 55 +++- .../src/postgresql/context/mod.rs | 241 ++++++------------ .../src/postgresql/error_handler.rs | 14 - .../src/postgresql/frontend.rs | 52 ++-- 7 files changed, 161 insertions(+), 213 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index ff6f77ab3..d2fae3c95 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3027,9 +3027,9 @@ checksum = "e3148f5046208a5d56bcfc03053e3ca6334e51da8dfb19b6cdc8b306fae3283e" [[package]] name = "pg-proto" -version = "0.1.0" +version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7a51a1b6becd8bc3571cc353b5f61ae8726f46ecacd78ab8222535cf710b92dd" +checksum = "6e850c5837b91e4dfd30bd4fb4923e4173f175b48d6d421bdfd3f5316b42a198" dependencies = [ "base64", "bytes", @@ -3049,9 +3049,9 @@ dependencies = [ [[package]] name = "pg-proto-fsm" -version = "0.1.0" +version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ba10fdad1b08db94d8d087940da00427599fe49e1e1d2762c8d95e7c7a7d8509" +checksum = "4460a8f63f7626bd5c0a446bd5f7d0d3b72bcd41d67efea1eeee9e6ca69396ec" dependencies = [ "proc-macro2", "quote", diff --git a/PG_PROTO_MIGRATION_PLAN.md b/PG_PROTO_MIGRATION_PLAN.md index c83a76a30..09c2ba7c7 100644 --- a/PG_PROTO_MIGRATION_PLAN.md +++ b/PG_PROTO_MIGRATION_PLAN.md @@ -5,7 +5,7 @@ - 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.1.0](https://crates.io/crates/pg-proto/0.1.0), covering codecs, startup/authentication, and runtime protocol-state validation. +- 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 diff --git a/packages/cipherstash-proxy/Cargo.toml b/packages/cipherstash-proxy/Cargo.toml index bbca02cb7..c2bb7ffcf 100644 --- a/packages/cipherstash-proxy/Cargo.toml +++ b/packages/cipherstash-proxy/Cargo.toml @@ -31,7 +31,7 @@ metrics-exporter-prometheus = "0.17" moka = { version = "0.12", features = ["future"] } oid-registry = "0.8" pg_escape = "0.1.1" -pg-proto = "0.1.0" +pg-proto = "0.2.1" postgres-protocol = "0.6.7" postgres-types = { version = "0.2.8", features = ["with-serde_json-1"] } rand = "0.9" diff --git a/packages/cipherstash-proxy/src/postgresql/backend.rs b/packages/cipherstash-proxy/src/postgresql/backend.rs index b68400dcf..51007a013 100644 --- a/packages/cipherstash-proxy/src/postgresql/backend.rs +++ b/packages/cipherstash-proxy/src/postgresql/backend.rs @@ -157,7 +157,6 @@ where ) .await?; - self.context.protocol_backend_received(&protocol_message)?; let read_duration = read_start.elapsed(); self.context.record_execute_server_timing(read_duration); @@ -227,7 +226,7 @@ where Ok(_) => (), Err(err) => { warn!(client_id = self.client_id(), error = err.to_string()); - self.send_error_response(err)?; + self.send_error_response(err).await?; } } @@ -243,7 +242,7 @@ where Ok(_) => (), Err(err) => { warn!(client_id = self.client_id(), error = err.to_string()); - self.send_error_response(err)?; + self.send_error_response(err).await?; } } @@ -375,7 +374,7 @@ where Ok(_) => (), Err(err) => { warn!(client_id = self.client_id(), error = err.to_string()); - self.send_error_response(err)?; + self.send_error_response(err).await?; } } @@ -387,12 +386,13 @@ where /// Write a message to the client /// pub async fn write(&mut self, bytes: BytesMut) -> Result<(), Error> { - self.context.protocol_backend_forwarded(&bytes)?; + self.context.protocol_backend_forwarded(&bytes).await?; let sent: u64 = bytes.len() as u64; counter!(CLIENTS_BYTES_SENT_TOTAL).increment(sent); let start = Instant::now(); self.client_sender.send(bytes)?; + self.context.protocol_backend_sent(); let duration = start.elapsed(); self.context.add_client_write_duration_for_execute(duration); @@ -722,13 +722,14 @@ 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 +where + R: AsyncRead + Unpin, + S: EncryptionService, +{ + 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(); @@ -742,7 +743,9 @@ where ?message, ); + self.context.protocol_backend_forwarded(&message).await?; self.client_sender.send(message)?; + self.context.protocol_backend_sent(); Ok(()) } @@ -756,7 +759,9 @@ mod tests { use crate::postgresql::context::KeysetIdentifier; use crate::postgresql::messages::Name; use crate::proxy::{EncryptConfig, EncryptionService}; + use bytes::Bytes; use eql_mapper::Schema; + use pg_proto::codec::FrontendMessage; use std::io::Cursor; use std::sync::Arc; use tokio::sync::mpsc; @@ -892,11 +897,39 @@ mod tests { backend .context .set_execute(Name::unnamed(), Some(session_id)); + backend + .context + .protocol_frontend_received(FrontendMessage::Execute( + pg_proto::codec::Execute { + portal: Bytes::new(), + max_rows: 0, + }, + )) + .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) + .await + .unwrap(); + let ready = protocol::encode_backend_message(&BackendMessage::ReadyForQuery( + pg_proto::codec::TransactionStatus::Idle, + )) + .unwrap(); + 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). assert_eq!( diff --git a/packages/cipherstash-proxy/src/postgresql/context/mod.rs b/packages/cipherstash-proxy/src/postgresql/context/mod.rs index 475333ec5..38345cb2c 100644 --- a/packages/cipherstash-proxy/src/postgresql/context/mod.rs +++ b/packages/cipherstash-proxy/src/postgresql/context/mod.rs @@ -7,7 +7,7 @@ pub use self::{phase_timing::PhaseTiming, portal::Portal, statement::Statement}; use super::{ column_mapper::ColumnMapper, messages::{describe::Describe, Name, Target}, - protocol::{decode_backend_frame, decode_frontend_frame}, + protocol::decode_backend_frame, Column, }; use crate::{ @@ -25,9 +25,9 @@ use cipherstash_client::IdentifiedBy; use eql_mapper::{Schema, TableResolver}; use metrics::{counter, histogram}; use pg_proto::{ - codec::{BackendMessage, FrontendMessage}, - grammar::{backend as server_role, frontend as client_role}, + codec::FrontendMessage, intermediary::Intermediary, + pipeline::{BackendAction, BoundedPipeline, FrontendAction, FrontendHandling}, }; use serde_json::json; use sqltk::parser::ast::{Expr, Ident, ObjectName, ObjectNamePart, Set, Value, ValueWithSpan}; @@ -40,7 +40,7 @@ use std::{ }, time::{Duration, Instant}, }; -use tokio::sync::oneshot; +use tokio::sync::{oneshot, Notify}; use tracing::{debug, error, warn}; use uuid::Uuid; @@ -57,71 +57,17 @@ fn protocol_transition_error(error: impl std::fmt::Debug) -> Error { std::io::Error::new(std::io::ErrorKind::InvalidData, format!("{error:?}")).into() } -fn protocol_neutral_backend_message(message: &BackendMessage) -> bool { - matches!( - message, - BackendMessage::ParameterStatus { .. } - | BackendMessage::NoticeResponse(_) - | BackendMessage::NotificationResponse { .. } - | BackendMessage::BackendKeyData { .. } - | BackendMessage::NegotiateProtocolVersion(_) - ) -} - #[derive(Debug)] struct ProtocolState { - sides: Intermediary, - upstream_startup_ready: bool, - downstream_startup_ready: bool, - runtime_started: bool, - downstream_frontend_queue: VecDeque, - upstream_backend_queue: VecDeque, + sides: Intermediary<(), (), BoundedPipeline>, } impl ProtocolState { fn new() -> Self { Self { - sides: Intermediary::new( - server_role::RuntimeFsm::new(), - client_role::RuntimeFsm::new(), + sides: Intermediary::new((), ()).with_pipeline( + BoundedPipeline::new(256).expect("non-zero PostgreSQL pipeline limit"), ), - upstream_startup_ready: false, - downstream_startup_ready: false, - runtime_started: false, - downstream_frontend_queue: VecDeque::new(), - upstream_backend_queue: VecDeque::new(), - } - } - - fn advance_downstream_frontend(&mut self, message: &FrontendMessage) -> bool { - let (downstream, _) = self.sides.sides_mut(); - downstream - .step_projected(message, server_role::project_external) - .is_ok() - } - - fn drain_downstream_frontend(&mut self) { - while let Some(message) = self.downstream_frontend_queue.front().cloned() { - if !self.advance_downstream_frontend(&message) { - break; - } - self.downstream_frontend_queue.pop_front(); - } - } - - fn advance_upstream_backend(&mut self, message: &BackendMessage) -> bool { - let (_, upstream) = self.sides.sides_mut(); - upstream - .step_projected(message, client_role::project_external) - .is_ok() - } - - fn drain_upstream_backend(&mut self) { - while let Some(message) = self.upstream_backend_queue.front().cloned() { - if !self.advance_upstream_backend(&message) { - break; - } - self.upstream_backend_queue.pop_front(); } } } @@ -161,6 +107,7 @@ where keyset_id: Arc>>, session_id_counter: Arc, protocol: Arc>, + protocol_changed: Arc, } /// Context for tracking an in-flight Execute operation. @@ -283,94 +230,74 @@ where 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 fn protocol_frontend_received(&self, message: &FrontendMessage) -> Result<(), Error> { - let mut protocol = self.protocol.lock().map_err(protocol_lock_error)?; - protocol.runtime_started = true; - if !protocol.downstream_frontend_queue.is_empty() - || !protocol.advance_downstream_frontend(message) - { - protocol - .downstream_frontend_queue - .push_back(message.clone()); - } - Ok(()) - } - - /// Records a frontend message actually forwarded upstream after rewriting. - pub fn protocol_frontend_forwarded(&self, bytes: &BytesMut) -> Result<(), Error> { - let message = decode_frontend_frame(bytes)?; - let mut protocol = self.protocol.lock().map_err(protocol_lock_error)?; - let (_, upstream) = protocol.sides.sides_mut(); - if upstream.state() == client_role::RuntimeState::Ready - && matches!( - message, - FrontendMessage::Parse(_) - | FrontendMessage::Bind(_) - | FrontendMessage::Describe(_) - | FrontendMessage::Execute(_) - | FrontendMessage::Close(_) - | FrontendMessage::Flush - | FrontendMessage::Sync - ) - { - upstream - .step(client_role::Event::BeginExtended) - .map_err(protocol_transition_error)?; + pub async fn protocol_frontend_received( + &self, + mut message: FrontendMessage, + ) -> Result<(), Error> { + loop { + let notified = self.protocol_changed.notified(); + tokio::pin!(notified); + notified.as_mut().enable(); + let action = { + let mut protocol = self.protocol.lock().map_err(protocol_lock_error)?; + protocol + .sides + .pipeline_mut() + .frontend_action(message, FrontendHandling::Forward) + .map_err(protocol_transition_error)? + }; + match action { + FrontendAction::Forward { .. } | FrontendAction::Discard { .. } => return Ok(()), + FrontendAction::Backpressure(returned) => { + message = returned; + notified.await; + } + } } - upstream - .step_projected(&message, client_role::project_internal) - .map_err(protocol_transition_error)?; - protocol.drain_upstream_backend(); - Ok(()) } - /// Records a message accepted from the upstream database. - pub fn protocol_backend_received(&self, message: &BackendMessage) -> Result<(), Error> { - if protocol_neutral_backend_message(message) { - return Ok(()); - } - let mut protocol = self.protocol.lock().map_err(protocol_lock_error)?; - if matches!(message, BackendMessage::ReadyForQuery(_)) && !protocol.upstream_startup_ready { - protocol.upstream_startup_ready = true; - return Ok(()); - } - if !protocol.runtime_started { - return Ok(()); - } - if !protocol.upstream_backend_queue.is_empty() - || !protocol.advance_upstream_backend(message) - { - protocol.upstream_backend_queue.push_back(message.clone()); - } - Ok(()) + pub fn protocol_in_extended_error(&self) -> Result { + let protocol = self.protocol.lock().map_err(protocol_lock_error)?; + Ok(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 fn protocol_backend_forwarded(&self, bytes: &BytesMut) -> Result<(), Error> { - let message = decode_backend_frame(bytes)?; - if protocol_neutral_backend_message(&message) { - return Ok(()); - } - let mut protocol = self.protocol.lock().map_err(protocol_lock_error)?; - if matches!(message, BackendMessage::ReadyForQuery(_)) && !protocol.downstream_startup_ready - { - protocol.downstream_startup_ready = true; - return Ok(()); - } - if !protocol.runtime_started { - return Ok(()); + pub async fn protocol_backend_forwarded(&self, bytes: &BytesMut) -> Result<(), Error> { + let mut message = decode_backend_frame(bytes)?; + loop { + let notified = self.protocol_changed.notified(); + tokio::pin!(notified); + notified.as_mut().enable(); + let action = { + let mut protocol = self.protocol.lock().map_err(protocol_lock_error)?; + protocol + .sides + .pipeline_mut() + .accept_backend(message) + .map_err(protocol_transition_error)? + }; + match action { + BackendAction::Emit(_) => return Ok(()), + BackendAction::Deferred(returned) => { + message = returned; + notified.await; + } + } } - let (downstream, _) = protocol.sides.sides_mut(); - downstream - .step_projected(&message, server_role::project_internal) - .map_err(protocol_transition_error)?; - protocol.drain_downstream_frontend(); - Ok(()) + } + + pub fn protocol_backend_sent(&self) { + self.protocol_changed.notify_waiters(); } pub fn set_describe(&mut self, describe: Describe) { @@ -1229,9 +1156,7 @@ impl Queue { #[cfg(test)] mod tests { - use super::{ - server_role, Context, Describe, KeysetIdentifier, Portal, ProtocolState, Statement, - }; + use super::{Context, Describe, KeysetIdentifier, Portal, ProtocolState, Statement}; use crate::{ config::LogConfig, error::Error, @@ -1249,6 +1174,7 @@ mod tests { use pg_proto::codec::{ BackendMessage, Bind, Execute, FrontendMessage, Parse, TransactionStatus, }; + use pg_proto::pipeline::{BackendAction, FrontendAction, FrontendHandling, PipelineState}; use sqltk::parser::{dialect::PostgreSqlDialect, parser::Parser}; use std::sync::Arc; use tokio::sync::mpsc; @@ -1297,7 +1223,7 @@ mod tests { } #[test] - fn server_role_fsm_tracks_pipelined_extended_messages_in_processing_order() { + fn pg_proto_ledger_tracks_pipelined_extended_messages_in_processing_order() { let mut protocol = ProtocolState::new(); let messages = [ FrontendMessage::Parse(Parse { @@ -1319,11 +1245,14 @@ mod tests { FrontendMessage::Sync, ]; for message in messages { - if !protocol.downstream_frontend_queue.is_empty() - || !protocol.advance_downstream_frontend(&message) - { - protocol.downstream_frontend_queue.push_back(message); - } + assert!(matches!( + protocol + .sides + .pipeline_mut() + .frontend_action(message, FrontendHandling::Forward) + .unwrap(), + FrontendAction::Forward { .. } + )); } for response in [ @@ -1332,18 +1261,18 @@ mod tests { BackendMessage::CommandComplete(Bytes::from_static(b"SELECT 1")), BackendMessage::ReadyForQuery(TransactionStatus::Idle), ] { - let (downstream, _) = protocol.sides.sides_mut(); - downstream - .step_projected(&response, server_role::project_internal) - .unwrap(); - protocol.drain_downstream_frontend(); + assert!(matches!( + protocol + .sides + .pipeline_mut() + .accept_backend(response) + .unwrap(), + BackendAction::Emit(_) + )); } - assert!(protocol.downstream_frontend_queue.is_empty()); - assert_eq!( - protocol.sides.downstream().state(), - server_role::RuntimeState::Ready - ); + assert!(protocol.sides.pipeline().is_empty()); + assert_eq!(protocol.sides.pipeline().state(), PipelineState::Ready); } fn statement() -> Statement { diff --git a/packages/cipherstash-proxy/src/postgresql/error_handler.rs b/packages/cipherstash-proxy/src/postgresql/error_handler.rs index 4de24e03f..bde5f1e0d 100644 --- a/packages/cipherstash-proxy/src/postgresql/error_handler.rs +++ b/packages/cipherstash-proxy/src/postgresql/error_handler.rs @@ -57,16 +57,6 @@ 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)] @@ -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/frontend.rs b/packages/cipherstash-proxy/src/postgresql/frontend.rs index f04f58024..b4c6824b7 100644 --- a/packages/cipherstash-proxy/src/postgresql/frontend.rs +++ b/packages/cipherstash-proxy/src/postgresql/frontend.rs @@ -102,13 +102,8 @@ where 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; - impl Frontend where R: AsyncRead + Unpin, @@ -134,7 +129,6 @@ where client_sender, server_writer, context, - error_state: None, } } @@ -168,7 +162,10 @@ where ) .await?; - self.context.protocol_frontend_received(&protocol_message)?; + let recovering_from_extended_error = self.context.protocol_in_extended_error()?; + self.context + .protocol_frontend_received(protocol_message.clone()) + .await?; let sent: u64 = bytes.len() as u64; counter!(CLIENTS_BYTES_RECEIVED_TOTAL).increment(sent); @@ -180,10 +177,9 @@ where // 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, message = ?protocol_message, ); if !matches!(protocol_message, FrontendMessage::Sync) { @@ -203,8 +199,8 @@ where msg = "Query Handler Error", error = ?err.to_string(), ); - self.send_error_response(err)?; - self.send_ready_for_query()?; + self.send_error_response(err).await?; + self.send_ready_for_query().await?; return Ok(()); } } @@ -226,7 +222,7 @@ where msg = "Parse Handler Error", error = ?err.to_string(), ); - self.send_error_response(err)?; + self.send_error_response(err).await?; return Ok(()); } } @@ -242,7 +238,7 @@ where client_id = self.context.client_id, msg = "EncryptError::InvalidParameter", ); - self.send_error_response(err)?; + self.send_error_response(err).await?; return Ok(()); } Error::Encrypt(EncryptError::UnknownKeysetIdentifier { .. }) => { @@ -250,7 +246,7 @@ where client_id = self.context.client_id, msg = "EncryptError::UnknownKeysetIdentifier", ); - self.send_error_response(err)?; + self.send_error_response(err).await?; return Ok(()); } _ => { @@ -259,7 +255,7 @@ where msg = "Bind Error", err = err.to_string() ); - self.send_error_response(err)?; + self.send_error_response(err).await?; return Ok(()); } }, @@ -273,12 +269,12 @@ where 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()?; + self.send_ready_for_query().await?; return Ok(()); } } @@ -299,7 +295,6 @@ where } pub async fn write_to_server(&mut self, bytes: BytesMut) -> Result<(), Error> { - self.context.protocol_frontend_forwarded(&bytes)?; debug!(target: PROTOCOL, msg = "Write to server", ?bytes); let sent: u64 = bytes.len() as u64; counter!(SERVER_BYTES_SENT_TOTAL).increment(sent); @@ -1192,7 +1187,7 @@ where /// /// Send an ReadyForQuery to the client and remove error state. /// - fn send_ready_for_query(&mut self) -> Result<(), Error> { + async fn send_ready_for_query(&mut self) -> Result<(), Error> { let message = protocol::encode_backend_message(&BackendMessage::ReadyForQuery( TransactionStatus::Idle, ))?; @@ -1203,10 +1198,9 @@ where ?message, ); - self.context.protocol_backend_forwarded(&message)?; + self.context.protocol_backend_forwarded(&message).await?; self.client_sender.send(message)?; - self.error_state = None; - + self.context.protocol_backend_sent(); Ok(()) } @@ -1343,8 +1337,15 @@ where fn client_id(&self) -> i32 { self.context.client_id } +} - fn send_error_response(&mut self, err: Error) -> Result<(), Error> { +impl Frontend +where + R: AsyncRead + Unpin, + W: AsyncWrite + Unpin, + S: EncryptionService, +{ + async fn send_error_response(&mut self, err: Error) -> Result<(), Error> { let error_response = self.error_to_response(err); let message = protocol::encode_backend_message(&error_response.into_backend_message())?; @@ -1354,10 +1355,9 @@ where ?message, ); - self.context.protocol_backend_forwarded(&message)?; + self.context.protocol_backend_forwarded(&message).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(()) } } From dcf55b74f118a662b2dbf19273601d360ea4081d Mon Sep 17 00:00:00 2001 From: James Sadler Date: Tue, 4 Aug 2026 23:25:59 +1000 Subject: [PATCH 04/16] refactor(proxy): use pg-proto buffered transport --- .../src/postgresql/backend.rs | 9 +- .../src/postgresql/frontend.rs | 9 +- .../src/postgresql/handler.rs | 315 +++++++++++------- .../src/postgresql/protocol.rs | 65 ++-- .../src/postgresql/startup.rs | 33 +- 5 files changed, 254 insertions(+), 177 deletions(-) diff --git a/packages/cipherstash-proxy/src/postgresql/backend.rs b/packages/cipherstash-proxy/src/postgresql/backend.rs index 51007a013..a029af61f 100644 --- a/packages/cipherstash-proxy/src/postgresql/backend.rs +++ b/packages/cipherstash-proxy/src/postgresql/backend.rs @@ -22,7 +22,10 @@ use crate::proxy::EncryptionService; use crate::EqlCiphertext; use bytes::BytesMut; use metrics::{counter, histogram}; -use pg_proto::codec::BackendMessage; +use pg_proto::{ + codec::{Backend as BackendDirection, BackendMessage}, + transport::Buffered, +}; use std::time::Instant; use tokio::io::AsyncRead; use tracing::{debug, error, info, warn}; @@ -79,7 +82,7 @@ where /// Sender for outgoing messages to client client_sender: Sender, /// Reader for incoming messages from server - server_reader: R, + server_reader: Buffered, /// Session context with portal and statement metadata context: Context, /// Buffer for batching DataRow messages before decryption @@ -103,7 +106,7 @@ where let buffer = MessageBuffer::new(); Backend { client_sender, - server_reader, + server_reader: Buffered::new(server_reader), context, buffer, } diff --git a/packages/cipherstash-proxy/src/postgresql/frontend.rs b/packages/cipherstash-proxy/src/postgresql/frontend.rs index b4c6824b7..ed8e9ab65 100644 --- a/packages/cipherstash-proxy/src/postgresql/frontend.rs +++ b/packages/cipherstash-proxy/src/postgresql/frontend.rs @@ -35,7 +35,10 @@ 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, FrontendMessage, TransactionStatus}; +use pg_proto::{ + codec::{BackendMessage, Frontend as FrontendDirection, FrontendMessage, TransactionStatus}, + transport::Buffered, +}; use serde::Serialize; use sqltk::parser::ast::{self, Value}; use sqltk::NodeKey; @@ -95,7 +98,7 @@ where S: EncryptionService, { /// Reader for incoming client messages - client_reader: R, + client_reader: Buffered, /// Sender for outgoing messages to client client_sender: Sender, /// Writer for forwarding messages to server @@ -125,7 +128,7 @@ where context: Context, ) -> Self { Frontend { - client_reader, + client_reader: Buffered::new_frontend(client_reader), client_sender, server_writer, context, diff --git a/packages/cipherstash-proxy/src/postgresql/handler.rs b/packages/cipherstash-proxy/src/postgresql/handler.rs index f91781dff..dc83cca92 100644 --- a/packages/cipherstash-proxy/src/postgresql/handler.rs +++ b/packages/cipherstash-proxy/src/postgresql/handler.rs @@ -2,7 +2,7 @@ use super::backend::Backend; use super::frontend::Frontend; use crate::connect::ChannelWriter; use crate::error::ConfigError; -use crate::log::{AUTHENTICATION, PROTOCOL}; +use crate::log::AUTHENTICATION; use crate::postgresql::messages::error_response::ErrorResponse; use crate::postgresql::{protocol, startup}; use crate::proxy::ZeroKms; @@ -14,16 +14,76 @@ use crate::{ }; use bytes::{BufMut, Bytes, BytesMut}; use md5::{Digest, Md5}; -use pg_proto::codec::{Authentication, BackendMessage, FrontendMessage}; use pg_proto::pre_startup::PreStartupMessage; +use pg_proto::{ + codec::{ + Authentication, Backend as BackendDirection, BackendMessage, Frontend as FrontendDirection, + FrontendMessage, + }, + pre_startup::{PreStartup, PreStartupOffer}, + server_auth::{ServerPassword, ServerProtocolOffer}, + startup::ProtocolVersion, + transport::Buffered, + Conn, +}; use postgres_protocol::authentication::sasl::{ChannelBinding, ScramSha256}; use rand::Rng; -use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt}; +use std::time::Duration; +use tokio::{ + io::{AsyncRead, AsyncWrite, AsyncWriteExt}, + 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"; +async fn receive_pre_startup( + mut conn: Conn, PreStartup>, + connection_timeout: Option, +) -> Result< + ( + Conn, PreStartup>, + PreStartupMessage, + ), + (Conn, PreStartup>, Error), +> { + let received = match connection_timeout { + Some(duration) => match timeout(duration, conn.receive_pre_startup_wire()).await { + Ok(received) => received, + Err(_) => return Err((conn, Error::ConnectionTimeout { duration })), + }, + None => conn.receive_pre_startup_wire().await, + }; + match received { + Ok(message) => Ok((conn, message)), + Err(error) => Err((conn, error.into())), + } +} + +async fn receive_frontend_auth( + mut conn: Conn, ServerPassword>, + connection_timeout: Option, +) -> Result< + ( + Conn, ServerPassword>, + FrontendMessage, + ), + (Conn, ServerPassword>, Error), +> { + let received = match connection_timeout { + Some(duration) => match timeout(duration, conn.receive_frontend_wire()).await { + Ok(received) => received, + Err(_) => return Err((conn, Error::ConnectionTimeout { duration })), + }, + None => conn.receive_frontend_wire().await, + }; + match received { + Ok(message) => Ok((conn, message)), + Err(error) => Err((conn, error.into())), + } +} + #[derive(Debug, Clone, Copy, PartialEq)] enum SaslMechanism { ScramSha256, @@ -34,62 +94,86 @@ enum SaslMechanism { /// 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: AsyncStream, context: Context) -> Result<(), Error> { - let mut client_stream = client_stream; + 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 = AsyncStream::connect(&context.database_socket_address()).await?; let mut database_stream = startup::with_tls(stream, context.config()).await?; + let database_channel_binding = database_stream.channel_binding(); 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; + let (client_startup, startup_message) = loop { + let (pre_startup, startup_message) = + match receive_pre_startup(client, 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); } - Err(err) => return Err(err), }; - match &startup_message { - PreStartupMessage::SslRequest => { - startup::send_ssl_response(&mut client_stream, context.use_tls()).await?; + match pre_startup.offer_pre_startup(startup_message) { + PreStartupOffer::Ssl(decision) => { + let mut stream = decision.into_transport().into_inner(); + startup::send_ssl_response(&mut stream, context.use_tls()).await?; if let Some(ref tls) = context.tls_config() { - match client_stream { - AsyncStream::Tcp(stream) => { + stream = match stream { + AsyncStream::Tcp(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)); + let tls_stream = tls::server(tcp_stream, tls).await?; + client_is_tls = true; + AsyncStream::Tls(Box::new(tls_stream)) } AsyncStream::Tls(_) => { unreachable!(); } - } + }; } + client = Conn::new(Buffered::new_frontend(stream)); } - PreStartupMessage::CancelRequest { .. } => { + 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); } - PreStartupMessage::Startup(_) => { + PreStartupOffer::Startup { conn, message } => { + let startup_message = PreStartupMessage::Startup(message.clone()); database_stream .write_all(&startup_message.to_packet()?) .await?; - break; + break (conn, message); } - PreStartupMessage::GssEncRequest => { + PreStartupOffer::Gss(conn) => { + conn.into_transport(); return Err(ProtocolError::UnexpectedStartupMessage.into()); } } - } + }; + + let mut database_stream = Buffered::<_, BackendDirection>::new(database_stream); // Proxy -> Client Authentication // Uses MD5 @@ -98,7 +182,16 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R // 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(); @@ -108,50 +201,50 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R let hash = md5_hash(username, password, &salt); - let bytes = protocol::encode_backend_message(&BackendMessage::Authentication( - Authentication::Md5Password { salt }, - ))?; - client_stream.write_all(&bytes).await?; + 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 (_bytes, message) = - match protocol::read_frontend_message(&mut client_stream, connection_timeout).await { + let (password_state, message) = + match receive_frontend_auth(password_state, connection_timeout).await { Ok(result) => result, - Err(err @ Error::ConnectionTimeout { .. }) => { - send_timeout_error(&mut client_stream, &err).await; + 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); } - Err(err) => return Err(err), }; - let FrontendMessage::PasswordResponse(password) = message else { - return Err(ProtocolError::UnexpectedAuthenticationResponse { - expected: "PasswordResponse".into(), - received: -1, - } - .into()); - }; - let password = password - .strip_suffix(&[0]) - .ok_or(ProtocolError::UnexpectedStartupMessage)?; + 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 { - debug!(target: AUTHENTICATION, msg = "Client AuthenticationOk"); - let bytes = protocol::encode_backend_message(&BackendMessage::Authentication( - Authentication::Ok, - ))?; - client_stream.write_all(&bytes).await?; - } else { - let message = ProtocolError::ClientAuthenticationFailed.to_string(); - error!(msg = message); + std::str::from_utf8(&password).map_err(|_| ProtocolError::AuthenticationFailed)?; - let message = ErrorResponse::invalid_password(message); - let bytes = protocol::encode_backend_message(&message.into_backend_message())?; - client_stream.write_all(&bytes).await?; + 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() + }; // Database authentication flow // 1. Database -> Authentication message (SASL) @@ -171,8 +264,7 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R Authentication::CleartextPassword => { debug!(target: AUTHENTICATION, msg = "AuthenticationCleartextPassword"); let password = context.database_password(); - let bytes = password_message(password)?; - database_stream.write_all(&bytes).await?; + send_frontend_message(&mut database_stream, password_message(password)?).await?; } Authentication::Md5Password { salt } => { debug!(target: AUTHENTICATION, msg = "Md5Password"); @@ -181,22 +273,20 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R let password = password.as_bytes(); let hash = md5_hash(username, password, salt); - let bytes = password_message(hash)?; - database_stream.write_all(&bytes).await?; + send_frontend_message(&mut database_stream, password_message(hash)?).await?; } Authentication::Sasl { mechanisms } => { debug!(target: AUTHENTICATION, msg = "Sasl"); let mechanism = sasl_mechanism(mechanisms)?; - 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?; + scram_sha_256_plus_handler( + &mut database_stream, + mechanism, + password, + database_channel_binding, + ) + .await?; } Authentication::KerberosV5 | Authentication::Gss @@ -214,17 +304,16 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R } } - if context.require_tls() && !client_stream.is_tls() { + if context.require_tls() && !client_is_tls { let message = ErrorResponse::tls_required(); - let bytes = protocol::encode_backend_message(&message.into_backend_message())?; - client_stream.write_all(&bytes).await?; + 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 (client_reader, client_writer) = client_stream.split(); - let (server_reader, server_writer) = database_stream.split(); + 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); @@ -311,28 +400,6 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R 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" - ); - } - } - } -} - fn sasl_mechanism(mechanisms: &[Bytes]) -> Result { match mechanisms.first().map(Bytes::as_ref) { Some(SCRAM_SHA_256) => Ok(SaslMechanism::ScramSha256), @@ -360,9 +427,9 @@ fn authentication_method_code(authentication: &Authentication) -> i32 { } } -fn password_message(password: String) -> Result { +fn password_message(password: String) -> Result { let password = std::ffi::CString::new(password)?; - protocol::encode_frontend_message(&FrontendMessage::PasswordResponse(Bytes::copy_from_slice( + Ok(FrontendMessage::PasswordResponse(Bytes::copy_from_slice( password.as_bytes_with_nul(), ))) } @@ -385,7 +452,7 @@ fn generate_md5_password_salt() -> [u8; 4] { } async fn scram_sha_256_plus_handler( - mut stream: S, + stream: &mut Buffered, mechanism: SaslMechanism, password: &[u8], channel_binding: ChannelBinding, @@ -407,11 +474,9 @@ async fn scram_sha_256_plus_handler( ) })?); initial.extend_from_slice(&bytes); - let bytes = - protocol::encode_frontend_message(&FrontendMessage::PasswordResponse(initial.freeze()))?; - stream.write_all(&bytes).await?; + send_frontend_message(stream, FrontendMessage::PasswordResponse(initial.freeze())).await?; - let auth = protocol::read_auth_message(&mut stream).await?; + let auth = protocol::read_auth_message(stream).await?; let Authentication::SaslContinue(bytes) = auth else { return Err(ProtocolError::UnexpectedAuthenticationResponse { @@ -422,12 +487,13 @@ async fn scram_sha_256_plus_handler( }; scram.update(&bytes)?; - let bytes = protocol::encode_frontend_message(&FrontendMessage::PasswordResponse( - Bytes::copy_from_slice(scram.message()), - ))?; - stream.write_all(&bytes).await?; + send_frontend_message( + stream, + FrontendMessage::PasswordResponse(Bytes::copy_from_slice(scram.message())), + ) + .await?; - let auth = protocol::read_auth_message(&mut stream).await?; + let auth = protocol::read_auth_message(stream).await?; let Authentication::SaslFinal(bytes) = auth else { return Err(ProtocolError::UnexpectedAuthenticationResponse { expected: "SaslFinal".into(), @@ -437,7 +503,7 @@ async fn scram_sha_256_plus_handler( }; scram.finish(&bytes)?; - let auth = protocol::read_auth_message(&mut stream).await?; + let auth = protocol::read_auth_message(stream).await?; if matches!(auth, Authentication::Ok) { debug!(target: AUTHENTICATION, msg = "SASL authentication successful"); @@ -449,9 +515,28 @@ async fn scram_sha_256_plus_handler( /// 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) { +async fn send_timeout_error( + stream: &mut Buffered, + err: &Error, +) { let error_response = ErrorResponse::connection_timeout(err.to_string()); - if let Ok(bytes) = protocol::encode_backend_message(&error_response.into_backend_message()) { - let _ = stream.write_all(&bytes).await; - } + 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(()) +} + +async fn send_frontend_message( + stream: &mut Buffered, + message: FrontendMessage, +) -> Result<(), Error> { + stream.push(message.to_frame()?)?; + stream.flush().await?; + Ok(()) } diff --git a/packages/cipherstash-proxy/src/postgresql/protocol.rs b/packages/cipherstash-proxy/src/postgresql/protocol.rs index f35173b73..472402281 100644 --- a/packages/cipherstash-proxy/src/postgresql/protocol.rs +++ b/packages/cipherstash-proxy/src/postgresql/protocol.rs @@ -1,13 +1,11 @@ use crate::error::{Error, ProtocolError}; use bytes::BytesMut; use pg_proto::codec::{ - Authentication, Backend, BackendMessage, Direction, Frontend, FrontendMessage, PgCodec, + Authentication, Backend, BackendMessage, Frontend, FrontendMessage, PgCodec, }; +use pg_proto::transport::Buffered; use std::time::Duration; -use tokio::{ - io::{AsyncRead, AsyncReadExt}, - time::timeout, -}; +use tokio::{io::AsyncRead, time::timeout}; use tokio_util::codec::Decoder; use tokio_util::codec::Encoder; @@ -49,11 +47,10 @@ pub fn encode_backend_message(message: &BackendMessage) -> Result( - mut stream: S, + stream: &mut Buffered, ) -> Result { let connection_timeout = Duration::from_millis(1000 * 10); - let (_bytes, message) = - read_backend_message_with_timeout(&mut stream, connection_timeout).await?; + let (_bytes, message) = read_backend_message_with_timeout(stream, connection_timeout).await?; match message { BackendMessage::Authentication(authentication) => Ok(authentication), _ => Err(ProtocolError::UnexpectedAuthenticationResponse { @@ -71,22 +68,22 @@ pub async fn read_auth_message( /// /// pub async fn read_frontend_message( - mut stream: S, + stream: &mut Buffered, connection_timeout: Option, ) -> Result<(BytesMut, FrontendMessage), Error> { match connection_timeout { Some(duration) => read_frontend_message_with_timeout(stream, duration).await, - None => read_frontend(&mut stream).await, + None => read_frontend(stream).await, } } pub async fn read_backend_message( - mut stream: S, + stream: &mut Buffered, connection_timeout: Option, ) -> Result<(BytesMut, BackendMessage), Error> { match connection_timeout { Some(duration) => read_backend_message_with_timeout(stream, duration).await, - None => read_backend(&mut stream).await, + None => read_backend(stream).await, } } @@ -97,19 +94,19 @@ pub async fn read_backend_message( /// /// async fn read_frontend_message_with_timeout( - mut stream: S, + stream: &mut Buffered, duration: Duration, ) -> Result<(BytesMut, FrontendMessage), Error> { - timeout(duration, read_frontend(&mut stream)) + timeout(duration, read_frontend(stream)) .await .map_err(|_| Error::ConnectionTimeout { duration })? } async fn read_backend_message_with_timeout( - mut stream: S, + stream: &mut Buffered, duration: Duration, ) -> Result<(BytesMut, BackendMessage), Error> { - timeout(duration, read_backend(&mut stream)) + timeout(duration, read_backend(stream)) .await .map_err(|_| Error::ConnectionTimeout { duration })? } @@ -122,34 +119,21 @@ async fn read_backend_message_with_timeout( /// /// async fn read_frontend( - stream: &mut S, + stream: &mut Buffered, ) -> Result<(BytesMut, FrontendMessage), Error> { - let message = read::(stream).await?; + let message = stream.receive_wire().await?; let bytes = encode_frontend_message(&message)?; Ok((bytes, message)) } async fn read_backend( - stream: &mut S, + stream: &mut Buffered, ) -> Result<(BytesMut, BackendMessage), Error> { - let message = read::(stream).await?; + let message = stream.receive_backend().await?; let bytes = encode_backend_message(&message)?; Ok((bytes, message)) } -async fn read(stream: &mut S) -> Result { - let mut codec = PgCodec::::default(); - let mut bytes = BytesMut::with_capacity(5); - loop { - if let Some(message) = codec.decode(&mut bytes)? { - return Ok(message); - } - if stream.read_buf(&mut bytes).await? == 0 { - return Err(Error::ConnectionClosed); - } - } -} - #[cfg(test)] mod tests { use super::*; @@ -169,11 +153,12 @@ mod tests { #[tokio::test] async fn frontend_frame_can_arrive_in_partial_writes() { - let (mut writer, mut reader) = duplex(64); + let (mut writer, reader) = duplex(64); let task = tokio::spawn(async move { writer.write_all(b"Q\0\0").await.unwrap(); writer.write_all(b"\0\x0dselect 1\0").await.unwrap(); }); + let mut reader = Buffered::new_frontend(reader); let (bytes, _) = read_frontend_message(&mut reader, None).await.unwrap(); assert_eq!(&bytes[..], b"Q\0\0\0\x0dselect 1\0"); @@ -182,8 +167,9 @@ mod tests { #[tokio::test] async fn unknown_frontend_tag_is_rejected() { - let (mut writer, mut reader) = duplex(16); + let (mut writer, reader) = duplex(16); writer.write_all(b"?\0\0\0\x04").await.unwrap(); + let mut reader = Buffered::new_frontend(reader); let error = read_frontend_message(&mut reader, None).await.unwrap_err(); assert!(error.to_string().contains("unknown frontend message tag")); @@ -191,14 +177,16 @@ mod tests { #[tokio::test] async fn malformed_and_oversized_frames_are_rejected_before_body_allocation() { - let (mut writer, mut reader) = duplex(16); + let (mut writer, reader) = duplex(16); writer.write_all(b"Q\0\0\0\x03").await.unwrap(); + let mut reader = Buffered::new_frontend(reader); assert!(read_frontend_message(&mut reader, None).await.is_err()); - let (mut writer, mut reader) = duplex(16); + let (mut writer, reader) = duplex(16); let oversized = (DEFAULT_MAX_FRAME_LEN as u32).to_be_bytes(); writer.write_all(b"Q").await.unwrap(); writer.write_all(&oversized).await.unwrap(); + let mut reader = Buffered::new_frontend(reader); assert!(read_frontend_message(&mut reader, None).await.is_err()); } @@ -214,13 +202,14 @@ mod tests { ]; for authentication in messages { - let (mut writer, mut reader) = duplex(128); + let (mut writer, reader) = duplex(128); writer .write_all(&encode_backend(BackendMessage::Authentication( authentication, ))) .await .unwrap(); + let mut reader = Buffered::new(reader); read_auth_message(&mut reader).await.unwrap(); } } diff --git a/packages/cipherstash-proxy/src/postgresql/startup.rs b/packages/cipherstash-proxy/src/postgresql/startup.rs index eabb6261f..02d7ef1c0 100644 --- a/packages/cipherstash-proxy/src/postgresql/startup.rs +++ b/packages/cipherstash-proxy/src/postgresql/startup.rs @@ -1,7 +1,10 @@ use std::time::Duration; -use bytes::BytesMut; -use pg_proto::pre_startup::{decode_pre_startup, EncryptionReply, PreStartupMessage}; +use pg_proto::{ + codec::Frontend, + pre_startup::{EncryptionReply, PreStartupMessage}, + transport::Buffered, +}; use tokio::{ io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}, time::timeout, @@ -50,12 +53,12 @@ pub async fn with_tls(stream: AsyncStream, config: &TandemConfig) -> Result( - mut stream: S, + stream: &mut Buffered, connection_timeout: Option, ) -> Result { match connection_timeout { Some(duration) => read_message_with_timeout(stream, duration).await, - None => read(&mut stream).await, + None => read(stream).await, } } @@ -66,10 +69,10 @@ pub async fn read_message( /// /// async fn read_message_with_timeout( - mut stream: S, + stream: &mut Buffered, duration: Duration, ) -> Result { - timeout(duration, read(&mut stream)) + timeout(duration, read(stream)) .await .map_err(|_| Error::ConnectionTimeout { duration })? } @@ -80,20 +83,13 @@ async fn read_message_with_timeout( /// /// /// -async fn read(client: &mut C) -> Result +async fn read(client: &mut Buffered) -> Result where C: AsyncRead + Unpin, { - let mut bytes = BytesMut::with_capacity(4); - loop { - if let Some(message) = decode_pre_startup(&mut bytes)? { - debug!(target: PROTOCOL, pre_startup = ?message); - return Ok(message); - } - if client.read_buf(&mut bytes).await? == 0 { - return Err(Error::ConnectionClosed); - } - } + let message = client.receive_pre_startup().await?; + debug!(target: PROTOCOL, pre_startup = ?message); + Ok(message) } /// @@ -155,11 +151,12 @@ mod tests { #[tokio::test] async fn ssl_and_cancellation_packets_use_pg_proto_pre_startup_decoding() { - let (mut writer, mut reader) = duplex(64); + let (mut writer, reader) = duplex(64); writer .write_all(&PreStartupMessage::SslRequest.to_packet().unwrap()) .await .unwrap(); + let mut reader = Buffered::new_frontend(reader); let ssl = read_message(&mut reader, None).await.unwrap(); assert!(matches!(ssl, PreStartupMessage::SslRequest)); From ca97b5b45efcdd37a1f4ef7782ac894e3c059861 Mon Sep 17 00:00:00 2001 From: James Sadler Date: Tue, 4 Aug 2026 23:38:52 +1000 Subject: [PATCH 05/16] refactor(proxy): pass typed PostgreSQL messages --- .../src/connect/channel_writer.rs | 32 ++- packages/cipherstash-proxy/src/error.rs | 3 +- .../src/postgresql/backend.rs | 118 ++++----- .../src/postgresql/context/mod.rs | 39 ++- .../src/postgresql/frontend.rs | 159 +++++------- .../src/postgresql/handler.rs | 245 +++++++++++++----- .../src/postgresql/messages/bind.rs | 82 +++++- .../src/postgresql/messages/close.rs | 117 --------- .../src/postgresql/messages/data_row.rs | 40 ++- .../src/postgresql/messages/describe.rs | 52 ---- .../src/postgresql/messages/execute.rs | 30 --- .../src/postgresql/messages/mod.rs | 9 +- .../src/postgresql/messages/name.rs | 50 ---- .../postgresql/messages/param_description.rs | 27 +- .../src/postgresql/messages/parse.rs | 46 +++- .../src/postgresql/messages/query.rs | 24 +- .../postgresql/messages/row_description.rs | 54 +++- .../src/postgresql/messages/target.rs | 46 ---- .../src/postgresql/protocol.rs | 44 ++-- 19 files changed, 614 insertions(+), 603 deletions(-) delete mode 100644 packages/cipherstash-proxy/src/postgresql/messages/close.rs delete mode 100644 packages/cipherstash-proxy/src/postgresql/messages/describe.rs delete mode 100644 packages/cipherstash-proxy/src/postgresql/messages/execute.rs delete mode 100644 packages/cipherstash-proxy/src/postgresql/messages/name.rs delete mode 100644 packages/cipherstash-proxy/src/postgresql/messages/target.rs 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/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/postgresql/backend.rs b/packages/cipherstash-proxy/src/postgresql/backend.rs index a029af61f..eb09f2495 100644 --- a/packages/cipherstash-proxy/src/postgresql/backend.rs +++ b/packages/cipherstash-proxy/src/postgresql/backend.rs @@ -20,7 +20,6 @@ use crate::prometheus::{ }; use crate::proxy::EncryptionService; use crate::EqlCiphertext; -use bytes::BytesMut; use metrics::{counter, histogram}; use pg_proto::{ codec::{Backend as BackendDirection, BackendMessage}, @@ -154,16 +153,25 @@ where /// error occurs that should terminate the connection. pub async fn rewrite(&mut self) -> Result<(), Error> { let read_start = Instant::now(); - let (mut bytes, protocol_message) = protocol::read_backend_message( + let protocol_message = protocol::read_backend_message( &mut self.server_reader, self.context.connection_timeout(), ) .await?; + let mut outbound_message = protocol_message.clone(); + + let session_item = self.server_reader.project_backend(protocol_message.clone()); + if session_item.is_none() { + // The demux has recorded the asynchronous event and its ordering; + // forwarding still uses the original typed message below. + let _ = self.server_reader.demux_mut().pop_async_event(); + } let read_duration = read_start.elapsed(); self.context.record_execute_server_timing(read_duration); - let sent: u64 = bytes.len() as u64; + let frame = protocol_message.to_frame()?; + let sent: u64 = (frame.body.len() + 5) as u64; counter!(SERVER_BYTES_RECEIVED_TOTAL).increment(sent); // Log slow database responses (configurable threshold, default 100ms) @@ -181,7 +189,7 @@ where client_id = self.context.client_id, msg = "Passthrough enabled" ); - self.write_with_flush(bytes).await?; + self.write_with_flush(outbound_message).await?; // The frontend starts a session and enqueues an execute for every // statement (start_session / set_execute), regardless of whether @@ -210,10 +218,10 @@ where debug!(target: CONTEXT, client_id = ?self.context.client_id, ?keyset_id); match protocol_message { - BackendMessage::DataRow(_) => { + 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? { + if self.data_row_handler(DataRow::from(row)).await? { return Ok(()); } } @@ -237,9 +245,7 @@ where self.context.finish_session(); } BackendMessage::ErrorResponse(ref response) => { - if let Some(b) = self.error_response_handler(response, &bytes) { - bytes = b - } + self.error_response_handler(response); match self.flush().await { Ok(_) => (), @@ -255,9 +261,12 @@ where // Describe with Target:Statement // Returns a ParameterDescription followed by RowDescription // The Describe is complete after the RowDescription - BackendMessage::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 @@ -265,9 +274,12 @@ where // Target::Portal returns a RowDescription // If no rows are returned, NoData is returned instead of a RowDescription // Complete the Describe - BackendMessage::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(); } @@ -298,7 +310,7 @@ where } } - self.write_with_flush(bytes).await?; + self.write_with_flush(outbound_message).await?; Ok(()) } @@ -339,15 +351,10 @@ where /// /// Always returns `Some(bytes)` containing the original error response /// to forward to the client unchanged. - fn error_response_handler( - &mut self, - response: &pg_proto::codec::DiagnosticResponse, - bytes: &BytesMut, - ) -> Option { + 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"); - Some(bytes.to_owned()) } /// @@ -370,7 +377,7 @@ 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 { @@ -381,20 +388,23 @@ where } } - 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> { - self.context.protocol_backend_forwarded(&bytes).await?; - 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); @@ -533,8 +543,7 @@ where row.rewrite(&data)?; - let bytes = BytesMut::try_from(row)?; - self.write(bytes).await?; + self.write(BackendMessage::from(row)).await?; } Ok(()) @@ -575,10 +584,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() { @@ -613,9 +620,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) } @@ -630,10 +637,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() { @@ -649,9 +654,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) } @@ -691,13 +696,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); @@ -737,7 +741,7 @@ where // Ensure any buffered data is cleared before sending error self.buffer.clear(); - let message = protocol::encode_backend_message(&error_response.into_backend_message())?; + let message = error_response.into_backend_message(); debug!( target: "PROTOCOL", @@ -746,7 +750,9 @@ where ?message, ); - self.context.protocol_backend_forwarded(&message).await?; + self.context + .protocol_backend_forwarded(message.clone()) + .await?; self.client_sender.send(message)?; self.context.protocol_backend_sent(); @@ -760,9 +766,9 @@ mod tests { use crate::config::{LogConfig, TandemConfig}; use crate::log; use crate::postgresql::context::KeysetIdentifier; - use crate::postgresql::messages::Name; use crate::proxy::{EncryptConfig, EncryptionService}; - use bytes::Bytes; + use bytes::Bytes as Name; + use bytes::{Bytes, BytesMut}; use eql_mapper::Schema; use pg_proto::codec::FrontendMessage; use std::io::Cursor; @@ -897,9 +903,7 @@ mod tests { 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::unnamed(), Some(session_id)); + backend.context.set_execute(Name::new(), Some(session_id)); backend .context .protocol_frontend_received(FrontendMessage::Execute( @@ -921,13 +925,11 @@ mod tests { .protocol_frontend_received(FrontendMessage::Sync) .await .unwrap(); - let ready = protocol::encode_backend_message(&BackendMessage::ReadyForQuery( - pg_proto::codec::TransactionStatus::Idle, - )) - .unwrap(); + let ready = + BackendMessage::ReadyForQuery(pg_proto::codec::TransactionStatus::Idle); backend .context - .protocol_backend_forwarded(&ready) + .protocol_backend_forwarded(ready) .await .unwrap(); backend.context.protocol_backend_sent(); diff --git a/packages/cipherstash-proxy/src/postgresql/context/mod.rs b/packages/cipherstash-proxy/src/postgresql/context/mod.rs index 38345cb2c..30a38b4ed 100644 --- a/packages/cipherstash-proxy/src/postgresql/context/mod.rs +++ b/packages/cipherstash-proxy/src/postgresql/context/mod.rs @@ -4,12 +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}, - protocol::decode_backend_frame, - Column, -}; +use super::{column_mapper::ColumnMapper, messages::Name, Column}; use crate::{ config::TandemConfig, error::{EncryptError, Error}, @@ -20,12 +15,11 @@ use crate::{ }, proxy::{EncryptConfig, EncryptionService, ReloadCommand, ReloadSender}, }; -use bytes::BytesMut; use cipherstash_client::IdentifiedBy; use eql_mapper::{Schema, TableResolver}; use metrics::{counter, histogram}; use pg_proto::{ - codec::FrontendMessage, + codec::{BackendMessage, Describe, DescribeTarget, FrontendMessage}, intermediary::Intermediary, pipeline::{BackendAction, BoundedPipeline, FrontendAction, FrontendHandling}, }; @@ -272,8 +266,10 @@ where /// 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, bytes: &BytesMut) -> Result<(), Error> { - let mut message = decode_backend_frame(bytes)?; + pub async fn protocol_backend_forwarded( + &self, + mut message: BackendMessage, + ) -> Result<(), Error> { loop { let notified = self.protocol_changed.notified(); tokio::pin!(notified); @@ -487,7 +483,7 @@ where ) .record(execute.duration()); - if execute.name.is_unnamed() { + if execute.name.is_empty() { self.close_portal(&execute.name); } } @@ -562,7 +558,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" ); } @@ -632,11 +628,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), } } @@ -1161,10 +1157,7 @@ mod tests { config::LogConfig, error::Error, log, - postgresql::{ - messages::{Name, Target}, - Column, - }, + postgresql::{messages::Name, Column}, proxy::{EncryptConfig, EncryptionService}, TandemConfig, }; @@ -1172,7 +1165,7 @@ mod tests { use cipherstash_client::IdentifiedBy; use eql_mapper::Schema; use pg_proto::codec::{ - BackendMessage, Bind, Execute, FrontendMessage, Parse, TransactionStatus, + BackendMessage, Bind, DescribeTarget, Execute, FrontendMessage, Parse, TransactionStatus, }; use pg_proto::pipeline::{BackendAction, FrontendAction, FrontendHandling, PipelineState}; use sqltk::parser::{dialect::PostgreSqlDialect, parser::Parser}; @@ -1312,7 +1305,7 @@ mod tests { let describe = Describe { name, - target: Target::Statement, + target: DescribeTarget::Statement, }; context.set_describe(describe); @@ -1381,7 +1374,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 / …). @@ -1448,10 +1441,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/frontend.rs b/packages/cipherstash-proxy/src/postgresql/frontend.rs index ed8e9ab65..99ced1778 100644 --- a/packages/cipherstash-proxy/src/postgresql/frontend.rs +++ b/packages/cipherstash-proxy/src/postgresql/frontend.rs @@ -2,8 +2,6 @@ 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::parser::SqlParser; @@ -20,8 +18,7 @@ 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::{Name, Target}; +use crate::postgresql::messages::Name; use crate::prometheus::{ CLIENTS_BYTES_RECEIVED_TOTAL, ENCRYPTED_VALUES_TOTAL, ENCRYPTION_DURATION_SECONDS, ENCRYPTION_ERROR_TOTAL, ENCRYPTION_REQUESTS_TOTAL, SERVER_BYTES_SENT_TOTAL, @@ -30,13 +27,14 @@ use crate::prometheus::{ }; 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, Frontend as FrontendDirection, FrontendMessage, TransactionStatus}, + codec::{ + Backend as BackendDirection, BackendMessage, Close, Describe, DescribeTarget, Execute, + Frontend as FrontendDirection, FrontendMessage, TransactionStatus, + }, transport::Buffered, }; use serde::Serialize; @@ -45,8 +43,8 @@ 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 tokio::io::{AsyncRead, AsyncWrite}; +use tracing::{debug, info, warn}; /// The PostgreSQL proxy frontend that handles client-to-server message processing. /// @@ -102,7 +100,7 @@ where /// Sender for outgoing messages to client client_sender: Sender, /// Writer for forwarding messages to server - server_writer: W, + server_writer: Buffered, /// Session context tracking statements, portals, and keyset IDs context: Context, } @@ -130,7 +128,7 @@ where Frontend { client_reader: Buffered::new_frontend(client_reader), client_sender, - server_writer, + server_writer: Buffered::new(server_writer), context, } } @@ -159,22 +157,24 @@ where /// 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 (mut bytes, protocol_message) = protocol::read_frontend_message( + let protocol_message = protocol::read_frontend_message( &mut self.client_reader, self.context.connection_timeout(), ) .await?; + let mut outbound_message = protocol_message.clone(); let recovering_from_extended_error = self.context.protocol_in_extended_error()?; self.context .protocol_frontend_received(protocol_message.clone()) .await?; - let sent: u64 = bytes.len() as u64; + let frame = protocol_message.to_frame()?; + let sent: u64 = (frame.body.len() + 5) as u64; counter!(CLIENTS_BYTES_RECEIVED_TOTAL).increment(sent); if self.context.mapping_disabled() { - self.write_to_server(bytes).await?; + self.write_to_server(outbound_message).await?; return Ok(()); } @@ -191,9 +191,9 @@ where } match protocol_message { - FrontendMessage::Query(_) => { - match self.query_handler(&bytes).await { - Ok(Some(mapped)) => bytes = mapped, + 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) => { @@ -208,15 +208,15 @@ where } } } - FrontendMessage::Describe(_) => { - self.describe_handler(&bytes).await?; + FrontendMessage::Describe(describe) => { + self.describe_handler(describe).await?; } - FrontendMessage::Execute(_) => { - self.execute_handler(&bytes).await?; + FrontendMessage::Execute(execute) => { + self.execute_handler(execute).await?; } - FrontendMessage::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) => { @@ -230,9 +230,9 @@ where } } } - FrontendMessage::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 { @@ -281,8 +281,8 @@ where return Ok(()); } } - FrontendMessage::Close(_) => { - self.close_handler(&bytes).await?; + FrontendMessage::Close(close) => { + self.close_handler(close).await?; } _ => { debug!(target: PROTOCOL, @@ -293,17 +293,19 @@ where } } - self.write_to_server(bytes).await?; + self.write_to_server(outbound_message).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; + pub async fn write_to_server(&mut self, message: FrontendMessage) -> Result<(), Error> { + debug!(target: PROTOCOL, msg = "Write to server", ?message); + let frame = message.to_frame()?; + let sent: u64 = (frame.body.len() + 5) as u64; counter!(SERVER_BYTES_SENT_TOTAL).increment(sent); let start = Instant::now(); - self.server_writer.write_all(&bytes).await?; + self.server_writer.push(frame)?; + self.server_writer.flush().await?; let duration = start.elapsed(); if let Some(session_id) = self.context.latest_session_id() { self.context.add_server_write_duration(session_id, duration); @@ -314,30 +316,26 @@ where pub async fn terminate(&mut self) -> Result<(), Error> { debug!(target: PROTOCOL, msg = "Terminate server connection"); - let bytes = protocol::encode_frontend_message(&FrontendMessage::Terminate)?; - self.write_to_server(bytes).await?; + self.write_to_server(FrontendMessage::Terminate).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()); @@ -376,7 +374,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(); @@ -387,8 +385,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![]; @@ -521,8 +517,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 @@ -533,7 +529,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, @@ -541,7 +537,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!( @@ -550,7 +546,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 { @@ -736,7 +732,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 @@ -746,7 +745,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); @@ -878,14 +876,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) } @@ -1030,7 +1028,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); @@ -1038,8 +1036,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); @@ -1075,15 +1071,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) } @@ -1191,9 +1187,7 @@ where /// Send an ReadyForQuery to the client and remove error state. /// async fn send_ready_for_query(&mut self) -> Result<(), Error> { - let message = protocol::encode_backend_message(&BackendMessage::ReadyForQuery( - TransactionStatus::Idle, - ))?; + let message = BackendMessage::ReadyForQuery(TransactionStatus::Idle); debug!(target: PROTOCOL, client_id = self.context.client_id, @@ -1201,34 +1195,13 @@ where ?message, ); - self.context.protocol_backend_forwarded(&message).await?; + self.context + .protocol_backend_forwarded(message.clone()) + .await?; self.client_sender.send(message)?; 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 @@ -1350,7 +1323,7 @@ where { async fn send_error_response(&mut self, err: Error) -> Result<(), Error> { let error_response = self.error_to_response(err); - let message = protocol::encode_backend_message(&error_response.into_backend_message())?; + let message = error_response.into_backend_message(); debug!(target: PROTOCOL, client_id = self.context.client_id, @@ -1358,7 +1331,9 @@ where ?message, ); - self.context.protocol_backend_forwarded(&message).await?; + self.context + .protocol_backend_forwarded(message.clone()) + .await?; self.client_sender.send(message)?; self.context.protocol_backend_sent(); Ok(()) diff --git a/packages/cipherstash-proxy/src/postgresql/handler.rs b/packages/cipherstash-proxy/src/postgresql/handler.rs index dc83cca92..69fd0ab5c 100644 --- a/packages/cipherstash-proxy/src/postgresql/handler.rs +++ b/packages/cipherstash-proxy/src/postgresql/handler.rs @@ -16,6 +16,7 @@ use bytes::{BufMut, Bytes, BytesMut}; use md5::{Digest, Md5}; use pg_proto::pre_startup::PreStartupMessage; use pg_proto::{ + auth::{AuthCompletion, AuthEvent, AuthOffer, SaslEvent}, codec::{ Authentication, Backend as BackendDirection, BackendMessage, Frontend as FrontendDirection, FrontendMessage, @@ -84,6 +85,169 @@ async fn receive_frontend_auth( } } +async fn receive_backend_conn( + mut conn: Conn, Phase>, +) -> Result<(Conn, Phase>, BackendMessage), Error> +where + S: AsyncRead + Unpin, +{ + let duration = Duration::from_secs(10); + let message = timeout(duration, conn.receive_backend_wire()) + .await + .map_err(|_| Error::ConnectionTimeout { duration })??; + Ok((conn, message)) +} + +async fn authenticate_upstream( + startup: Conn, pg_proto::pre_startup::Startup>, + context: &Context, + channel_binding: ChannelBinding, +) -> Result, Error> { + let mut auth = startup.authentication(); + let offer = loop { + let (current, message) = receive_backend_conn(auth).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.into_transport()), + 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).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).await + } + AuthOffer::Sasl { conn, mechanisms } => { + let mechanism = sasl_mechanism(&mechanisms)?; + if mechanism == SaslMechanism::ScramSha256Plus { + // pg-proto's SCRAM-PLUS entry point currently requires its own + // TLS transport trait. Keep only the cryptographic exchange as + // an adapter until custom channel-binding bytes are accepted. + let mut transport = conn.into_transport(); + scram_sha_256_plus_handler( + &mut transport, + mechanism, + context.database_password().as_bytes(), + channel_binding, + ) + .await?; + return Ok(transport); + } + + let mut scram = ScramSha256::new( + context.database_password().as_bytes(), + ChannelBinding::unsupported(), + ); + let (mut sasl, frame) = conn.scram_sha_256(scram.message())?; + sasl.push_frame(frame)?; + sasl.flush().await?; + + let (sasl, message) = receive_backend_conn(sasl).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).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()).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>, +) -> Result, Error> { + let (awaiting, message) = receive_backend_conn(awaiting).await?; + match awaiting.offer(message) { + Ok(AuthCompletion::Ok(conn)) => Ok(conn.into_transport()), + 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( + database: &mut Buffered, + client: &mut Buffered, +) -> Result<(), Error> { + loop { + let message = database.receive_backend().await?; + let ready = matches!(message, BackendMessage::ReadyForQuery(_)); + let _ = database.project_backend(message.clone()); + send_backend_message(client, message).await?; + if ready { + return Ok(()); + } + } +} + #[derive(Debug, Clone, Copy, PartialEq)] enum SaslMechanism { ScramSha256, @@ -159,13 +323,7 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R .await?; return Err(Error::CancelRequest); } - PreStartupOffer::Startup { conn, message } => { - let startup_message = PreStartupMessage::Startup(message.clone()); - database_stream - .write_all(&startup_message.to_packet()?) - .await?; - break (conn, message); - } + PreStartupOffer::Startup { conn, message } => break (conn, message), PreStartupOffer::Gss(conn) => { conn.into_transport(); return Err(ProtocolError::UnexpectedStartupMessage.into()); @@ -173,7 +331,13 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R } }; - let mut database_stream = Buffered::<_, BackendDirection>::new(database_stream); + 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 mut database_stream = + authenticate_upstream(database_startup, &context, database_channel_binding).await?; // Proxy -> Client Authentication // Uses MD5 @@ -246,64 +410,6 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R startup_ready.into_transport() }; - // 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).await?; - - match &auth { - Authentication::Ok => { - debug!(target: AUTHENTICATION, msg = "AuthenticationOk"); - } - Authentication::CleartextPassword => { - debug!(target: AUTHENTICATION, msg = "AuthenticationCleartextPassword"); - let password = context.database_password(); - send_frontend_message(&mut database_stream, password_message(password)?).await?; - } - Authentication::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); - send_frontend_message(&mut database_stream, password_message(hash)?).await?; - } - Authentication::Sasl { mechanisms } => { - debug!(target: AUTHENTICATION, msg = "Sasl"); - let mechanism = sasl_mechanism(mechanisms)?; - let password = context.database_password(); - let password = password.as_bytes(); - scram_sha_256_plus_handler( - &mut database_stream, - mechanism, - password, - database_channel_binding, - ) - .await?; - } - Authentication::KerberosV5 - | Authentication::Gss - | Authentication::GssContinue(_) - | Authentication::Sspi => { - debug!(target: AUTHENTICATION, msg = "UnsupportedAuthentication"); - return Err(ProtocolError::UnsupportedAuthentication { - method_code: authentication_method_code(&auth), - } - .into()); - } - Authentication::SaslContinue(_) | Authentication::SaslFinal(_) => { - debug!(target: AUTHENTICATION, msg = "UnexpectedStartupMessage", authentication_method = ?auth); - return Err(ProtocolError::UnexpectedStartupMessage.into()); - } - } - if context.require_tls() && !client_is_tls { let message = ErrorResponse::tls_required(); send_backend_message(&mut client_stream, message.into_backend_message()).await?; @@ -312,6 +418,8 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R return Err(ConfigError::TlsRequired.into()); } + drain_upstream_startup(&mut database_stream, &mut client_stream).await?; + let (client_reader, client_writer) = client_stream.into_inner().split(); let (server_reader, server_writer) = database_stream.into_inner().split(); @@ -369,10 +477,7 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R if let Err(ref err @ Error::ConnectionTimeout { .. }) = &result { let error_response = ErrorResponse::connection_timeout(err.to_string()); - if let Ok(bytes) = protocol::encode_backend_message(&error_response.into_backend_message()) - { - let _ = timeout_sender.send(bytes); - } + 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 diff --git a/packages/cipherstash-proxy/src/postgresql/messages/bind.rs b/packages/cipherstash-proxy/src/postgresql/messages/bind.rs index 94563c28c..3321b870a 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/bind.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/bind.rs @@ -9,9 +9,10 @@ use crate::postgresql::data::{ bind_param_from_sql, bind_param_json_value, json_value_selector_plaintext, }; use crate::postgresql::format_code::FormatCode; +#[cfg(test)] use crate::postgresql::protocol::{decode_frontend_frame, encode_frontend_message}; use crate::{EqlOutput, EqlQueryPayload}; -use bytes::{BufMut, Bytes, BytesMut}; +use bytes::{BufMut, BytesMut}; use cipherstash_client::encryption::Plaintext; use pg_proto::codec::{Bind as PgBind, FrontendMessage}; use postgres_types::Type; @@ -211,6 +212,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 @@ -337,6 +387,7 @@ impl Display for BindParam { } } +#[cfg(test)] impl TryFrom<&BytesMut> for Bind { type Error = Error; @@ -348,8 +399,8 @@ impl TryFrom<&BytesMut> for Bind { } .into()); }; - let portal = Name::from(String::from_utf8_lossy(&bind.portal).into_owned()); - let prepared_statement = Name::from(String::from_utf8_lossy(&bind.statement).into_owned()); + let portal = bind.portal; + let prepared_statement = bind.statement; let param_format_codes = bind .parameter_formats .iter() @@ -395,13 +446,14 @@ impl TryFrom<&BytesMut> for Bind { } } +#[cfg(test)] impl TryFrom for BytesMut { type Error = Error; fn try_from(bind: Bind) -> Result { encode_frontend_message(&FrontendMessage::Bind(PgBind { - portal: Bytes::copy_from_slice(bind.portal.as_str().as_bytes()), - statement: Bytes::copy_from_slice(bind.prepared_statement.as_str().as_bytes()), + portal: bind.portal, + statement: bind.prepared_statement, parameter_formats: bind.param_format_codes.into_iter().map(i16::from).collect(), parameters: bind .param_values @@ -417,6 +469,26 @@ impl TryFrom for BytesMut { } } +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(), + }) + } +} + #[cfg(test)] mod tests { use super::BindParam; 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 b3ee2c55c..000000000 --- a/packages/cipherstash-proxy/src/postgresql/messages/close.rs +++ /dev/null @@ -1,117 +0,0 @@ -use crate::error::{Error, ProtocolError}; -use crate::postgresql::protocol::{decode_frontend_frame, encode_frontend_message}; - -use bytes::{Bytes, BytesMut}; -use pg_proto::codec::{Close as PgClose, DescribeTarget, FrontendMessage}; -use std::convert::TryFrom; - -use super::target::Target; -use super::Name; - -/// Proxy state extracted from a typed frontend `Close` message. -#[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 FrontendMessage::Close(close) = decode_frontend_frame(bytes)? else { - return Err(ProtocolError::UnexpectedMessageCode { - expected: 'C', - received: bytes.first().copied().unwrap_or_default() as char, - } - .into()); - }; - let target = match close.target { - DescribeTarget::Statement => Target::Statement, - DescribeTarget::Portal => Target::Portal, - }; - let name = Name::from(String::from_utf8_lossy(&close.name).into_owned()); - - Ok(Close { target, name }) - } -} - -impl TryFrom for BytesMut { - type Error = Error; - - fn try_from(close: Close) -> Result { - let target = match close.target { - Target::Statement => DescribeTarget::Statement, - Target::Portal => DescribeTarget::Portal, - }; - encode_frontend_message(&FrontendMessage::Close(PgClose { - target, - name: Bytes::copy_from_slice(close.name.as_str().as_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/data_row.rs b/packages/cipherstash-proxy/src/postgresql/messages/data_row.rs index 843a657af..f6d914051 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/data_row.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/data_row.rs @@ -1,11 +1,13 @@ use crate::EqlCiphertext; +#[cfg(test)] +use crate::{ + error::ProtocolError, + postgresql::protocol::{decode_backend_frame, encode_backend_message}, +}; use crate::{ - error::{EncryptError, Error, ProtocolError}, + error::{EncryptError, Error}, log::DECRYPT, - postgresql::{ - protocol::{decode_backend_frame, encode_backend_message}, - Column, - }, + postgresql::Column, }; use bytes::BytesMut; use pg_proto::codec::{BackendMessage, DataRow as PgDataRow}; @@ -98,6 +100,7 @@ impl DataColumn { } } +#[cfg(test)] impl TryFrom<&BytesMut> for DataRow { type Error = Error; @@ -121,6 +124,21 @@ impl TryFrom<&BytesMut> for DataRow { } } +impl From for DataRow { + fn from(row: PgDataRow) -> Self { + Self { + columns: row + .columns + .into_iter() + .map(|bytes| DataColumn { + bytes: bytes.map(BytesMut::from), + }) + .collect(), + } + } +} + +#[cfg(test)] impl TryFrom for BytesMut { type Error = Error; @@ -135,6 +153,18 @@ impl TryFrom for BytesMut { } } +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(), + }) + } +} + impl DataColumn { /// Parse this column's bytes into an [`EqlCiphertext`]. /// 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 20af83058..000000000 --- a/packages/cipherstash-proxy/src/postgresql/messages/describe.rs +++ /dev/null @@ -1,52 +0,0 @@ -use crate::error::{Error, ProtocolError}; -use crate::postgresql::protocol::{decode_frontend_frame, encode_frontend_message}; - -use bytes::{Bytes, BytesMut}; -use pg_proto::codec::{Describe as PgDescribe, DescribeTarget, FrontendMessage}; -use std::convert::TryFrom; - -use super::target::Target; -use super::Name; - -/// Proxy state extracted from a typed frontend `Describe` message. -#[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 FrontendMessage::Describe(description) = decode_frontend_frame(bytes)? else { - return Err(ProtocolError::UnexpectedMessageCode { - expected: 'D', - received: bytes.first().copied().unwrap_or_default() as char, - } - .into()); - }; - let target = match description.target { - DescribeTarget::Statement => Target::Statement, - DescribeTarget::Portal => Target::Portal, - }; - let name = Name::from(String::from_utf8_lossy(&description.name).into_owned()); - - Ok(Describe { target, name }) - } -} - -impl TryFrom for BytesMut { - type Error = Error; - - fn try_from(describe: Describe) -> Result { - let target = match describe.target { - Target::Statement => DescribeTarget::Statement, - Target::Portal => DescribeTarget::Portal, - }; - encode_frontend_message(&FrontendMessage::Describe(PgDescribe { - target, - name: Bytes::copy_from_slice(describe.name.as_str().as_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 dc6d2aacd..000000000 --- a/packages/cipherstash-proxy/src/postgresql/messages/execute.rs +++ /dev/null @@ -1,30 +0,0 @@ -use super::Name; -use crate::error::{Error, ProtocolError}; -use crate::postgresql::protocol::decode_frontend_frame; -use bytes::BytesMut; -use pg_proto::codec::FrontendMessage; -use std::convert::TryFrom; - -#[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 FrontendMessage::Execute(execute) = decode_frontend_frame(bytes)? else { - return Err(ProtocolError::UnexpectedMessageCode { - expected: 'E', - received: bytes.first().copied().unwrap_or_default() as char, - } - .into()); - }; - let portal = Name::from(String::from_utf8_lossy(&execute.portal).into_owned()); - let max_rows = execute.max_rows; - - 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 index df7d98661..f172afdc1 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/mod.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/mod.rs @@ -1,20 +1,13 @@ use bytes::BytesMut; 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 row_description; -pub mod target; - -pub use name::Name; -pub use target::Target; +pub type Name = bytes::Bytes; pub const NULL: i32 = -1; 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/param_description.rs b/packages/cipherstash-proxy/src/postgresql/messages/param_description.rs index 50841a226..f2e36d188 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/param_description.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/param_description.rs @@ -1,8 +1,10 @@ +use crate::log::MAPPER; +#[cfg(test)] use crate::{ error::{Error, ProtocolError}, - log::MAPPER, postgresql::protocol::{decode_backend_frame, encode_backend_message}, }; +#[cfg(test)] use bytes::BytesMut; use pg_proto::codec::BackendMessage; use postgres_types::Type; @@ -65,6 +67,7 @@ impl ParamDescription { } } +#[cfg(test)] impl TryFrom<&BytesMut> for ParamDescription { type Error = Error; @@ -84,6 +87,16 @@ impl TryFrom<&BytesMut> for ParamDescription { } } +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; @@ -97,6 +110,18 @@ impl TryFrom for BytesMut { } } +impl From for BackendMessage { + fn from(parameter_description: ParamDescription) -> Self { + Self::ParameterDescription( + parameter_description + .types + .into_iter() + .map(|oid| oid as u32) + .collect(), + ) + } +} + #[cfg(test)] mod tests { diff --git a/packages/cipherstash-proxy/src/postgresql/messages/parse.rs b/packages/cipherstash-proxy/src/postgresql/messages/parse.rs index aef0d6b4d..421bbc857 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/parse.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/parse.rs @@ -1,12 +1,13 @@ use super::{Name, UNSPECIFIED_TYPE_OID}; +use crate::postgresql::context::statement::OutputParam; +#[cfg(test)] use crate::{ error::{Error, ProtocolError}, - postgresql::{ - context::statement::OutputParam, - protocol::{decode_frontend_frame, encode_frontend_message}, - }, + postgresql::protocol::{decode_frontend_frame, encode_frontend_message}, }; -use bytes::{Bytes, 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; @@ -87,6 +88,22 @@ 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; @@ -98,7 +115,7 @@ impl TryFrom<&BytesMut> for Parse { } .into()); }; - let name = Name::from(String::from_utf8_lossy(&parse.statement).into_owned()); + let name = parse.statement; let statement = String::from_utf8_lossy(&parse.query).into_owned(); let param_types = parse .parameter_types @@ -115,12 +132,13 @@ impl TryFrom<&BytesMut> for Parse { } } +#[cfg(test)] impl TryFrom for BytesMut { type Error = Error; fn try_from(parse: Parse) -> Result { encode_frontend_message(&FrontendMessage::Parse(PgParse { - statement: Bytes::copy_from_slice(parse.name.as_str().as_bytes()), + statement: parse.name, query: Bytes::from(parse.statement), parameter_types: parse .param_types @@ -131,6 +149,20 @@ impl TryFrom for BytesMut { } } +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(), + }) + } +} + #[cfg(test)] mod tests { use crate::{ diff --git a/packages/cipherstash-proxy/src/postgresql/messages/query.rs b/packages/cipherstash-proxy/src/postgresql/messages/query.rs index 919de8af8..99571bc4b 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/query.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/query.rs @@ -1,8 +1,13 @@ +#[cfg(test)] use crate::error::{Error, ProtocolError}; +#[cfg(test)] use crate::postgresql::protocol::{decode_frontend_frame, encode_frontend_message}; -use bytes::{Bytes, BytesMut}; +use bytes::Bytes; +#[cfg(test)] +use bytes::BytesMut; use pg_proto::codec::FrontendMessage; +#[cfg(test)] use std::convert::TryFrom; #[derive(Debug, Clone)] @@ -30,6 +35,16 @@ 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; @@ -49,6 +64,7 @@ impl TryFrom<&BytesMut> for Query { } } +#[cfg(test)] impl TryFrom for BytesMut { type Error = Error; @@ -56,3 +72,9 @@ impl TryFrom for BytesMut { encode_frontend_message(&FrontendMessage::Query(Bytes::from(query.statement))) } } + +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/messages/row_description.rs index 06145de98..5e4e256fc 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/row_description.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/row_description.rs @@ -1,15 +1,16 @@ -use bytes::{Bytes, BytesMut}; +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::{decode_backend_frame, encode_backend_message}, - }, + postgresql::protocol::{decode_backend_frame, encode_backend_message}, }; #[derive(Debug)] @@ -57,6 +58,7 @@ impl RowDescriptionField { } } +#[cfg(test)] impl TryFrom<&BytesMut> for RowDescription { type Error = Error; @@ -88,6 +90,28 @@ impl TryFrom<&BytesMut> for RowDescription { } } +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; @@ -110,6 +134,26 @@ impl TryFrom for BytesMut { } } +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(), + }) + } +} + #[cfg(test)] mod tests { 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/protocol.rs b/packages/cipherstash-proxy/src/postgresql/protocol.rs index 472402281..4c017dc3d 100644 --- a/packages/cipherstash-proxy/src/postgresql/protocol.rs +++ b/packages/cipherstash-proxy/src/postgresql/protocol.rs @@ -1,14 +1,16 @@ use crate::error::{Error, ProtocolError}; +#[cfg(test)] use bytes::BytesMut; -use pg_proto::codec::{ - Authentication, Backend, BackendMessage, Frontend, FrontendMessage, PgCodec, -}; +#[cfg(test)] +use pg_proto::codec::PgCodec; +use pg_proto::codec::{Authentication, Backend, BackendMessage, Frontend, FrontendMessage}; use pg_proto::transport::Buffered; use std::time::Duration; use tokio::{io::AsyncRead, time::timeout}; -use tokio_util::codec::Decoder; -use tokio_util::codec::Encoder; +#[cfg(test)] +use tokio_util::codec::{Decoder, Encoder}; +#[cfg(test)] pub fn decode_frontend_frame(bytes: &BytesMut) -> Result { let mut bytes = bytes.clone(); PgCodec::::default() @@ -18,6 +20,7 @@ pub fn decode_frontend_frame(bytes: &BytesMut) -> Result }) } +#[cfg(test)] pub fn decode_backend_frame(bytes: &BytesMut) -> Result { let mut bytes = bytes.clone(); PgCodec::::default() @@ -27,12 +30,14 @@ pub fn decode_backend_frame(bytes: &BytesMut) -> Result { }) } +#[cfg(test)] pub fn encode_frontend_message(message: &FrontendMessage) -> Result { let mut bytes = BytesMut::new(); PgCodec::::default().encode(message.to_frame()?, &mut bytes)?; Ok(bytes) } +#[cfg(test)] pub fn encode_backend_message(message: &BackendMessage) -> Result { let mut bytes = BytesMut::new(); PgCodec::::default().encode(message.to_frame()?, &mut bytes)?; @@ -50,7 +55,7 @@ pub async fn read_auth_message( stream: &mut Buffered, ) -> Result { let connection_timeout = Duration::from_millis(1000 * 10); - let (_bytes, message) = read_backend_message_with_timeout(stream, connection_timeout).await?; + let message = read_backend_message_with_timeout(stream, connection_timeout).await?; match message { BackendMessage::Authentication(authentication) => Ok(authentication), _ => Err(ProtocolError::UnexpectedAuthenticationResponse { @@ -70,7 +75,7 @@ pub async fn read_auth_message( pub async fn read_frontend_message( stream: &mut Buffered, connection_timeout: Option, -) -> Result<(BytesMut, FrontendMessage), Error> { +) -> Result { match connection_timeout { Some(duration) => read_frontend_message_with_timeout(stream, duration).await, None => read_frontend(stream).await, @@ -80,7 +85,7 @@ pub async fn read_frontend_message( pub async fn read_backend_message( stream: &mut Buffered, connection_timeout: Option, -) -> Result<(BytesMut, BackendMessage), Error> { +) -> Result { match connection_timeout { Some(duration) => read_backend_message_with_timeout(stream, duration).await, None => read_backend(stream).await, @@ -96,7 +101,7 @@ pub async fn read_backend_message( async fn read_frontend_message_with_timeout( stream: &mut Buffered, duration: Duration, -) -> Result<(BytesMut, FrontendMessage), Error> { +) -> Result { timeout(duration, read_frontend(stream)) .await .map_err(|_| Error::ConnectionTimeout { duration })? @@ -105,7 +110,7 @@ async fn read_frontend_message_with_timeout( async fn read_backend_message_with_timeout( stream: &mut Buffered, duration: Duration, -) -> Result<(BytesMut, BackendMessage), Error> { +) -> Result { timeout(duration, read_backend(stream)) .await .map_err(|_| Error::ConnectionTimeout { duration })? @@ -120,18 +125,14 @@ async fn read_backend_message_with_timeout( /// async fn read_frontend( stream: &mut Buffered, -) -> Result<(BytesMut, FrontendMessage), Error> { - let message = stream.receive_wire().await?; - let bytes = encode_frontend_message(&message)?; - Ok((bytes, message)) +) -> Result { + Ok(stream.receive_wire().await?) } async fn read_backend( stream: &mut Buffered, -) -> Result<(BytesMut, BackendMessage), Error> { - let message = stream.receive_backend().await?; - let bytes = encode_backend_message(&message)?; - Ok((bytes, message)) +) -> Result { + Ok(stream.receive_backend().await?) } #[cfg(test)] @@ -160,8 +161,11 @@ mod tests { }); let mut reader = Buffered::new_frontend(reader); - let (bytes, _) = read_frontend_message(&mut reader, None).await.unwrap(); - assert_eq!(&bytes[..], b"Q\0\0\0\x0dselect 1\0"); + let message = read_frontend_message(&mut reader, None).await.unwrap(); + assert_eq!( + message, + FrontendMessage::Query(bytes::Bytes::from_static(b"select 1")) + ); task.await.unwrap(); } From 34bdc788fd2d8a328288f5a7baebcb6dd0e66229 Mon Sep 17 00:00:00 2001 From: James Sadler Date: Tue, 4 Aug 2026 23:42:44 +1000 Subject: [PATCH 06/16] refactor(proxy): remove legacy protocol plumbing --- Cargo.lock | 34 --- PG_PROTO_FOLLOWUPS.md | 37 +++ packages/cipherstash-proxy/Cargo.toml | 1 - .../src/postgresql/backend.rs | 21 +- .../src/postgresql/frontend.rs | 23 +- .../src/postgresql/handler.rs | 45 ++-- .../src/postgresql/messages/bind.rs | 2 +- .../src/postgresql/messages/data_row.rs | 2 +- .../src/postgresql/messages/error_response.rs | 2 +- .../postgresql/messages/param_description.rs | 2 +- .../src/postgresql/messages/parse.rs | 2 +- .../src/postgresql/messages/query.rs | 2 +- .../postgresql/messages/row_description.rs | 2 +- .../cipherstash-proxy/src/postgresql/mod.rs | 3 +- .../src/postgresql/protocol.rs | 220 ------------------ .../src/postgresql/startup.rs | 189 ++------------- .../src/postgresql/test_codec.rs | 31 +++ 17 files changed, 160 insertions(+), 458 deletions(-) create mode 100644 PG_PROTO_FOLLOWUPS.md delete mode 100644 packages/cipherstash-proxy/src/postgresql/protocol.rs create mode 100644 packages/cipherstash-proxy/src/postgresql/test_codec.rs diff --git a/Cargo.lock b/Cargo.lock index d2fae3c95..253906731 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -842,7 +842,6 @@ dependencies = [ "moka", "oid-registry", "pg-proto", - "pg_escape", "postgres-protocol", "postgres-types", "rand 0.9.2", @@ -3059,48 +3058,15 @@ dependencies = [ "syn 3.0.3", ] -[[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" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1fd6780a80ae0c52cc120a26a1a42c1ae51b247a253e4e06113d23d2c2edd078" dependencies = [ - "phf_macros", "phf_shared", ] -[[package]] -name = "phf_generator" -version = "0.11.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3c80231409c20246a13fddb31776fb942c38553c51e871f8cbd687a4cfb5843d" -dependencies = [ - "phf_shared", - "rand 0.8.6", -] - -[[package]] -name = "phf_macros" -version = "0.11.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f84ac04429c13a7ff43785d75ad27569f2951ce0ffd30a3321230db2fc727216" -dependencies = [ - "phf_generator", - "phf_shared", - "proc-macro2", - "quote", - "syn 2.0.117", -] - [[package]] name = "phf_shared" version = "0.11.3" diff --git a/PG_PROTO_FOLLOWUPS.md b/PG_PROTO_FOLLOWUPS.md new file mode 100644 index 000000000..49e8db9cb --- /dev/null +++ b/PG_PROTO_FOLLOWUPS.md @@ -0,0 +1,37 @@ +# pg-proto follow-ups + +The proxy migration delegates framing, typed messages, startup/authentication +state, demultiplexing, and bounded pipeline scheduling to `pg-proto` 0.2.1. +Three remaining adapters 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. + +## Accept application-provided SCRAM channel binding + +The SCRAM-SHA-256-PLUS typestate obtains channel binding through pg-proto's TLS +transport trait. CipherStash uses its own `AsyncStream` TLS abstraction, which +already exposes the RFC 5929 binding bytes but cannot supply them through that +trait after type erasure. + +Allowing callers to provide validated channel-binding bytes (or a small adapter +trait independent of the transport type) would remove the proxy's final manual +SCRAM-PLUS exchange. The SCRAM cryptographic engine itself should remain +application-owned. + +## 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. diff --git a/packages/cipherstash-proxy/Cargo.toml b/packages/cipherstash-proxy/Cargo.toml index c2bb7ffcf..5c59da691 100644 --- a/packages/cipherstash-proxy/Cargo.toml +++ b/packages/cipherstash-proxy/Cargo.toml @@ -30,7 +30,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.2.1" postgres-protocol = "0.6.7" postgres-types = { version = "0.2.8", features = ["with-serde_json-1"] } diff --git a/packages/cipherstash-proxy/src/postgresql/backend.rs b/packages/cipherstash-proxy/src/postgresql/backend.rs index eb09f2495..b4ba2d89a 100644 --- a/packages/cipherstash-proxy/src/postgresql/backend.rs +++ b/packages/cipherstash-proxy/src/postgresql/backend.rs @@ -12,7 +12,6 @@ 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::prometheus::{ CLIENTS_BYTES_SENT_TOTAL, DECRYPTED_VALUES_TOTAL, DECRYPTION_DURATION_SECONDS, DECRYPTION_ERROR_TOTAL, DECRYPTION_REQUESTS_TOTAL, ROWS_ENCRYPTED_TOTAL, @@ -29,6 +28,19 @@ use std::time::Instant; use tokio::io::AsyncRead; use tracing::{debug, error, info, warn}; +async fn receive_backend( + reader: &mut Buffered, + connection_timeout: Option, +) -> Result { + match connection_timeout { + Some(duration) => tokio::time::timeout(duration, reader.receive_backend()) + .await + .map_err(|_| Error::ConnectionTimeout { duration })? + .map_err(Into::into), + None => reader.receive_backend().await.map_err(Into::into), + } +} + /// The PostgreSQL proxy backend that handles server-to-client message processing. /// /// The Backend intercepts messages from PostgreSQL servers, identifies encrypted data @@ -153,11 +165,8 @@ where /// error occurs that should terminate the connection. pub async fn rewrite(&mut self) -> Result<(), Error> { let read_start = Instant::now(); - let protocol_message = protocol::read_backend_message( - &mut self.server_reader, - self.context.connection_timeout(), - ) - .await?; + let protocol_message = + receive_backend(&mut self.server_reader, self.context.connection_timeout()).await?; let mut outbound_message = protocol_message.clone(); let session_item = self.server_reader.project_backend(protocol_message.clone()); diff --git a/packages/cipherstash-proxy/src/postgresql/frontend.rs b/packages/cipherstash-proxy/src/postgresql/frontend.rs index 99ced1778..0ac932fa6 100644 --- a/packages/cipherstash-proxy/src/postgresql/frontend.rs +++ b/packages/cipherstash-proxy/src/postgresql/frontend.rs @@ -5,7 +5,6 @@ use super::messages::bind::Bind; use super::messages::parse::Parse; use super::messages::query::Query; use super::parser::SqlParser; -use super::protocol::{self}; use crate::connect::Sender; use crate::error::{EncryptError, Error, MappingError}; use crate::log::{MAPPER, PROTOCOL}; @@ -42,10 +41,23 @@ use sqltk::parser::ast::{self, Value}; use sqltk::NodeKey; use std::collections::HashMap; use std::sync::Arc; -use std::time::Instant; +use std::time::{Duration, Instant}; use tokio::io::{AsyncRead, AsyncWrite}; use tracing::{debug, info, warn}; +async fn receive_frontend( + reader: &mut Buffered, + connection_timeout: Option, +) -> Result { + match connection_timeout { + Some(duration) => tokio::time::timeout(duration, reader.receive_wire()) + .await + .map_err(|_| Error::ConnectionTimeout { duration })? + .map_err(Into::into), + None => reader.receive_wire().await.map_err(Into::into), + } +} + /// The PostgreSQL proxy frontend that handles client-to-server message processing. /// /// The Frontend intercepts messages from PostgreSQL clients, analyzes SQL statements for @@ -157,11 +169,8 @@ where /// 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 protocol_message = protocol::read_frontend_message( - &mut self.client_reader, - self.context.connection_timeout(), - ) - .await?; + let protocol_message = + receive_frontend(&mut self.client_reader, self.context.connection_timeout()).await?; let mut outbound_message = protocol_message.clone(); let recovering_from_extended_error = self.context.protocol_in_extended_error()?; diff --git a/packages/cipherstash-proxy/src/postgresql/handler.rs b/packages/cipherstash-proxy/src/postgresql/handler.rs index 69fd0ab5c..95bf9f468 100644 --- a/packages/cipherstash-proxy/src/postgresql/handler.rs +++ b/packages/cipherstash-proxy/src/postgresql/handler.rs @@ -4,7 +4,7 @@ use crate::connect::ChannelWriter; use crate::error::ConfigError; use crate::log::AUTHENTICATION; use crate::postgresql::messages::error_response::ErrorResponse; -use crate::postgresql::{protocol, startup}; +use crate::postgresql::startup; use crate::proxy::ZeroKms; use crate::{ connect::AsyncStream, @@ -98,6 +98,23 @@ where Ok((conn, message)) } +async fn receive_auth_message( + stream: &mut Buffered, +) -> Result { + let duration = Duration::from_secs(10); + let message = timeout(duration, stream.receive_backend()) + .await + .map_err(|_| Error::ConnectionTimeout { duration })??; + let BackendMessage::Authentication(authentication) = message else { + return Err(ProtocolError::UnexpectedAuthenticationResponse { + expected: "Authentication".into(), + received: -1, + } + .into()); + }; + Ok(authentication) +} + async fn authenticate_upstream( startup: Conn, pg_proto::pre_startup::Startup>, context: &Context, @@ -291,8 +308,17 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R match pre_startup.offer_pre_startup(startup_message) { PreStartupOffer::Ssl(decision) => { - let mut stream = decision.into_transport().into_inner(); - startup::send_ssl_response(&mut stream, context.use_tls()).await?; + 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() { stream = match stream { AsyncStream::Tcp(tcp_stream) => { @@ -532,13 +558,6 @@ fn authentication_method_code(authentication: &Authentication) -> i32 { } } -fn password_message(password: String) -> Result { - let password = std::ffi::CString::new(password)?; - Ok(FrontendMessage::PasswordResponse(Bytes::copy_from_slice( - password.as_bytes_with_nul(), - ))) -} - pub fn md5_hash(username: &[u8], password: &[u8], salt: &[u8; 4]) -> String { let mut md5 = Md5::new(); md5.update(password); @@ -581,7 +600,7 @@ async fn scram_sha_256_plus_handler( initial.extend_from_slice(&bytes); send_frontend_message(stream, FrontendMessage::PasswordResponse(initial.freeze())).await?; - let auth = protocol::read_auth_message(stream).await?; + let auth = receive_auth_message(stream).await?; let Authentication::SaslContinue(bytes) = auth else { return Err(ProtocolError::UnexpectedAuthenticationResponse { @@ -598,7 +617,7 @@ async fn scram_sha_256_plus_handler( ) .await?; - let auth = protocol::read_auth_message(stream).await?; + let auth = receive_auth_message(stream).await?; let Authentication::SaslFinal(bytes) = auth else { return Err(ProtocolError::UnexpectedAuthenticationResponse { expected: "SaslFinal".into(), @@ -608,7 +627,7 @@ async fn scram_sha_256_plus_handler( }; scram.finish(&bytes)?; - let auth = protocol::read_auth_message(stream).await?; + let auth = receive_auth_message(stream).await?; if matches!(auth, Authentication::Ok) { debug!(target: AUTHENTICATION, msg = "SASL authentication successful"); diff --git a/packages/cipherstash-proxy/src/postgresql/messages/bind.rs b/packages/cipherstash-proxy/src/postgresql/messages/bind.rs index 3321b870a..b83700ce8 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/bind.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/bind.rs @@ -10,7 +10,7 @@ use crate::postgresql::data::{ }; use crate::postgresql::format_code::FormatCode; #[cfg(test)] -use crate::postgresql::protocol::{decode_frontend_frame, encode_frontend_message}; +use crate::postgresql::test_codec::{decode_frontend_frame, encode_frontend_message}; use crate::{EqlOutput, EqlQueryPayload}; use bytes::{BufMut, BytesMut}; use cipherstash_client::encryption::Plaintext; diff --git a/packages/cipherstash-proxy/src/postgresql/messages/data_row.rs b/packages/cipherstash-proxy/src/postgresql/messages/data_row.rs index f6d914051..2dbd1ca72 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/data_row.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/data_row.rs @@ -2,7 +2,7 @@ use crate::EqlCiphertext; #[cfg(test)] use crate::{ error::ProtocolError, - postgresql::protocol::{decode_backend_frame, encode_backend_message}, + postgresql::test_codec::{decode_backend_frame, encode_backend_message}, }; use crate::{ error::{EncryptError, Error}, diff --git a/packages/cipherstash-proxy/src/postgresql/messages/error_response.rs b/packages/cipherstash-proxy/src/postgresql/messages/error_response.rs index c54bd9764..4eacc23dd 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/error_response.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/error_response.rs @@ -454,7 +454,7 @@ impl From for ErrorResponseCode { mod tests { use super::ErrorResponseCode; use crate::postgresql::messages::error_response::ErrorResponse; - use crate::postgresql::protocol::{decode_backend_frame, encode_backend_message}; + use crate::postgresql::test_codec::{decode_backend_frame, encode_backend_message}; use bytes::BytesMut; use pg_proto::codec::BackendMessage; diff --git a/packages/cipherstash-proxy/src/postgresql/messages/param_description.rs b/packages/cipherstash-proxy/src/postgresql/messages/param_description.rs index f2e36d188..f2b270598 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/param_description.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/param_description.rs @@ -2,7 +2,7 @@ use crate::log::MAPPER; #[cfg(test)] use crate::{ error::{Error, ProtocolError}, - postgresql::protocol::{decode_backend_frame, encode_backend_message}, + postgresql::test_codec::{decode_backend_frame, encode_backend_message}, }; #[cfg(test)] use bytes::BytesMut; diff --git a/packages/cipherstash-proxy/src/postgresql/messages/parse.rs b/packages/cipherstash-proxy/src/postgresql/messages/parse.rs index 421bbc857..77ca64160 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/parse.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/parse.rs @@ -3,7 +3,7 @@ use crate::postgresql::context::statement::OutputParam; #[cfg(test)] use crate::{ error::{Error, ProtocolError}, - postgresql::protocol::{decode_frontend_frame, encode_frontend_message}, + postgresql::test_codec::{decode_frontend_frame, encode_frontend_message}, }; use bytes::Bytes; #[cfg(test)] diff --git a/packages/cipherstash-proxy/src/postgresql/messages/query.rs b/packages/cipherstash-proxy/src/postgresql/messages/query.rs index 99571bc4b..2476380fc 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/query.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/query.rs @@ -1,7 +1,7 @@ #[cfg(test)] use crate::error::{Error, ProtocolError}; #[cfg(test)] -use crate::postgresql::protocol::{decode_frontend_frame, encode_frontend_message}; +use crate::postgresql::test_codec::{decode_frontend_frame, encode_frontend_message}; use bytes::Bytes; #[cfg(test)] diff --git a/packages/cipherstash-proxy/src/postgresql/messages/row_description.rs b/packages/cipherstash-proxy/src/postgresql/messages/row_description.rs index 5e4e256fc..e02eb1818 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/row_description.rs +++ b/packages/cipherstash-proxy/src/postgresql/messages/row_description.rs @@ -10,7 +10,7 @@ use crate::postgresql::format_code::FormatCode; #[cfg(test)] use crate::{ error::{Error, ProtocolError}, - postgresql::protocol::{decode_backend_frame, encode_backend_message}, + postgresql::test_codec::{decode_backend_frame, encode_backend_message}, }; #[derive(Debug)] diff --git a/packages/cipherstash-proxy/src/postgresql/mod.rs b/packages/cipherstash-proxy/src/postgresql/mod.rs index d09a7a2ee..7e1502d4d 100644 --- a/packages/cipherstash-proxy/src/postgresql/mod.rs +++ b/packages/cipherstash-proxy/src/postgresql/mod.rs @@ -9,8 +9,9 @@ mod handler; mod message_buffer; mod messages; mod parser; -mod protocol; mod startup; +#[cfg(test)] +mod test_codec; pub use context::column::Column; pub use context::Context; diff --git a/packages/cipherstash-proxy/src/postgresql/protocol.rs b/packages/cipherstash-proxy/src/postgresql/protocol.rs deleted file mode 100644 index 4c017dc3d..000000000 --- a/packages/cipherstash-proxy/src/postgresql/protocol.rs +++ /dev/null @@ -1,220 +0,0 @@ -use crate::error::{Error, ProtocolError}; -#[cfg(test)] -use bytes::BytesMut; -#[cfg(test)] -use pg_proto::codec::PgCodec; -use pg_proto::codec::{Authentication, Backend, BackendMessage, Frontend, FrontendMessage}; -use pg_proto::transport::Buffered; -use std::time::Duration; -use tokio::{io::AsyncRead, time::timeout}; -#[cfg(test)] -use tokio_util::codec::{Decoder, Encoder}; - -#[cfg(test)] -pub fn decode_frontend_frame(bytes: &BytesMut) -> Result { - let mut bytes = bytes.clone(); - PgCodec::::default() - .decode(&mut bytes)? - .ok_or_else(|| { - std::io::Error::new(std::io::ErrorKind::UnexpectedEof, "partial frontend frame").into() - }) -} - -#[cfg(test)] -pub fn decode_backend_frame(bytes: &BytesMut) -> Result { - let mut bytes = bytes.clone(); - PgCodec::::default() - .decode(&mut bytes)? - .ok_or_else(|| { - std::io::Error::new(std::io::ErrorKind::UnexpectedEof, "partial backend frame").into() - }) -} - -#[cfg(test)] -pub fn encode_frontend_message(message: &FrontendMessage) -> Result { - let mut bytes = BytesMut::new(); - PgCodec::::default().encode(message.to_frame()?, &mut bytes)?; - Ok(bytes) -} - -#[cfg(test)] -pub fn encode_backend_message(message: &BackendMessage) -> Result { - let mut bytes = BytesMut::new(); - PgCodec::::default().encode(message.to_frame()?, &mut bytes)?; - Ok(bytes) -} - -/// -/// 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( - stream: &mut Buffered, -) -> Result { - let connection_timeout = Duration::from_millis(1000 * 10); - let message = read_backend_message_with_timeout(stream, connection_timeout).await?; - match message { - BackendMessage::Authentication(authentication) => Ok(authentication), - _ => Err(ProtocolError::UnexpectedAuthenticationResponse { - expected: "Authentication".into(), - received: -1, - } - .into()), - } -} - -/// -/// Reads a Postgres message from client with an optional timeout -/// -/// Timeout values are in config -/// -/// -pub async fn read_frontend_message( - stream: &mut Buffered, - connection_timeout: Option, -) -> Result { - match connection_timeout { - Some(duration) => read_frontend_message_with_timeout(stream, duration).await, - None => read_frontend(stream).await, - } -} - -pub async fn read_backend_message( - stream: &mut Buffered, - connection_timeout: Option, -) -> Result { - match connection_timeout { - Some(duration) => read_backend_message_with_timeout(stream, duration).await, - None => read_backend(stream).await, - } -} - -/// -/// Reads a Postgres message from client with a timeout -/// -/// Timeout values are in config -/// -/// -async fn read_frontend_message_with_timeout( - stream: &mut Buffered, - duration: Duration, -) -> Result { - timeout(duration, read_frontend(stream)) - .await - .map_err(|_| Error::ConnectionTimeout { duration })? -} - -async fn read_backend_message_with_timeout( - stream: &mut Buffered, - duration: Duration, -) -> Result { - timeout(duration, read_backend(stream)) - .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_frontend( - stream: &mut Buffered, -) -> Result { - Ok(stream.receive_wire().await?) -} - -async fn read_backend( - stream: &mut Buffered, -) -> Result { - Ok(stream.receive_backend().await?) -} - -#[cfg(test)] -mod tests { - use super::*; - use pg_proto::codec::{ - Authentication as PgAuthentication, BackendMessage, DEFAULT_MAX_FRAME_LEN, - }; - use tokio::io::{duplex, AsyncWriteExt}; - use tokio_util::codec::Encoder; - - fn encode_backend(message: BackendMessage) -> BytesMut { - let mut bytes = BytesMut::new(); - PgCodec::::default() - .encode(message.to_frame().unwrap(), &mut bytes) - .unwrap(); - bytes - } - - #[tokio::test] - async fn frontend_frame_can_arrive_in_partial_writes() { - let (mut writer, reader) = duplex(64); - let task = tokio::spawn(async move { - writer.write_all(b"Q\0\0").await.unwrap(); - writer.write_all(b"\0\x0dselect 1\0").await.unwrap(); - }); - let mut reader = Buffered::new_frontend(reader); - - let message = read_frontend_message(&mut reader, None).await.unwrap(); - assert_eq!( - message, - FrontendMessage::Query(bytes::Bytes::from_static(b"select 1")) - ); - task.await.unwrap(); - } - - #[tokio::test] - async fn unknown_frontend_tag_is_rejected() { - let (mut writer, reader) = duplex(16); - writer.write_all(b"?\0\0\0\x04").await.unwrap(); - let mut reader = Buffered::new_frontend(reader); - - let error = read_frontend_message(&mut reader, None).await.unwrap_err(); - assert!(error.to_string().contains("unknown frontend message tag")); - } - - #[tokio::test] - async fn malformed_and_oversized_frames_are_rejected_before_body_allocation() { - let (mut writer, reader) = duplex(16); - writer.write_all(b"Q\0\0\0\x03").await.unwrap(); - let mut reader = Buffered::new_frontend(reader); - assert!(read_frontend_message(&mut reader, None).await.is_err()); - - let (mut writer, reader) = duplex(16); - let oversized = (DEFAULT_MAX_FRAME_LEN as u32).to_be_bytes(); - writer.write_all(b"Q").await.unwrap(); - writer.write_all(&oversized).await.unwrap(); - let mut reader = Buffered::new_frontend(reader); - assert!(read_frontend_message(&mut reader, None).await.is_err()); - } - - #[tokio::test] - async fn authentication_modes_are_validated_by_the_backend_codec() { - let messages = [ - PgAuthentication::Ok, - PgAuthentication::CleartextPassword, - PgAuthentication::Md5Password { salt: *b"salt" }, - PgAuthentication::Sasl { - mechanisms: vec![bytes::Bytes::from_static(b"SCRAM-SHA-256")], - }, - ]; - - for authentication in messages { - let (mut writer, reader) = duplex(128); - writer - .write_all(&encode_backend(BackendMessage::Authentication( - authentication, - ))) - .await - .unwrap(); - let mut reader = Buffered::new(reader); - read_auth_message(&mut reader).await.unwrap(); - } - } -} diff --git a/packages/cipherstash-proxy/src/postgresql/startup.rs b/packages/cipherstash-proxy/src/postgresql/startup.rs index 02d7ef1c0..0464989be 100644 --- a/packages/cipherstash-proxy/src/postgresql/startup.rs +++ b/packages/cipherstash-proxy/src/postgresql/startup.rs @@ -1,187 +1,38 @@ -use std::time::Duration; - -use pg_proto::{ - codec::Frontend, - pre_startup::{EncryptionReply, PreStartupMessage}, - transport::Buffered, -}; -use tokio::{ - io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}, - time::timeout, -}; -use tracing::{debug, error, warn}; +use pg_proto::{codec::Backend, pre_startup::Negotiation, transport::Buffered, Conn}; +use tracing::warn; use crate::{ connect::AsyncStream, error::{Error, ProtocolError}, - log::PROTOCOL, tls, TandemConfig, }; +/// Applies CipherStash's upstream TLS policy while pg-proto owns the wire +/// negotiation and its legal pre-startup transitions. 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) + let mut request = Conn::new(Buffered::<_, Backend>::new(stream)).request_ssl(); + request.flush().await?; + match request.receive_ssl_reply().await? { + Negotiation::Accepted(conn) => { + let AsyncStream::Tcp(stream) = conn.into_transport().into_inner() else { + return Err(ProtocolError::UnexpectedStartupMessage.into()); + }; + Ok(AsyncStream::Tls(Box::new( + tls::client(stream, config).await?, + ))) } - } -} - -/// -/// Reads a Postgres startup message from client with an optional timeout -/// -/// Timeout values are in config -/// -/// -pub async fn read_message( - stream: &mut Buffered, - connection_timeout: Option, -) -> Result { - match connection_timeout { - Some(duration) => read_message_with_timeout(stream, duration).await, - None => read(stream).await, - } -} - -/// -/// Reads a Postgres message from client with a timeout -/// -/// Timeout values are in config -/// -/// -async fn read_message_with_timeout( - stream: &mut Buffered, - duration: Duration, -) -> Result { - timeout(duration, read(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 Buffered) -> Result -where - C: AsyncRead + Unpin, -{ - let message = client.receive_pre_startup().await?; - debug!(target: PROTOCOL, pre_startup = ?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 { - stream - .write_all(&PreStartupMessage::SslRequest.to_packet()?) - .await?; - - // Server supports TLS - let response = match EncryptionReply::try_from(stream.read_u8().await?) { - Ok(EncryptionReply::Accepted) => true, - Ok(EncryptionReply::Rejected) => false, - Ok(EncryptionReply::LegacyError) => { - return Err(ProtocolError::UnexpectedStartupMessage.into()); + Negotiation::Rejected(conn) => { + warn!(msg = "Connecting to database without Transport Layer Security (TLS)"); + Ok(conn.into_transport().into_inner()) } - Err(err) => { - let code = err.0; - error!(msg = "Unexpected startup message", code = ?(code as char)); - return Err(ProtocolError::UnexpectedStartupMessage.into()); + Negotiation::LegacyError(conn) => { + conn.into_transport(); + 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 { - EncryptionReply::Accepted - } else { - EncryptionReply::Rejected - }; - - debug!(target: PROTOCOL, msg = "SSLResponse to Client", SSLResponse = ?response); - - stream.write_all(&[response.as_byte()]).await?; - - Ok(()) -} - -#[cfg(test)] -mod tests { - use super::*; - use tokio::io::{duplex, AsyncReadExt, AsyncWriteExt}; - - #[tokio::test] - async fn ssl_and_cancellation_packets_use_pg_proto_pre_startup_decoding() { - let (mut writer, reader) = duplex(64); - writer - .write_all(&PreStartupMessage::SslRequest.to_packet().unwrap()) - .await - .unwrap(); - let mut reader = Buffered::new_frontend(reader); - let ssl = read_message(&mut reader, None).await.unwrap(); - assert!(matches!(ssl, PreStartupMessage::SslRequest)); - - let cancel = PreStartupMessage::CancelRequest { - process_id: 42, - secret_key: bytes::Bytes::from_static(b"key!"), - } - .to_packet() - .unwrap(); - writer.write_all(&cancel).await.unwrap(); - let decoded = read_message(&mut reader, None).await.unwrap(); - assert!(matches!(decoded, PreStartupMessage::CancelRequest { .. })); - assert_eq!(decoded.to_packet().unwrap(), cancel); - } - - #[tokio::test] - async fn ssl_reply_rejects_unknown_bytes() { - let (mut writer, mut reader) = duplex(8); - writer.write_all(b"?").await.unwrap(); - assert!(send_ssl_request(&mut reader).await.is_err()); - - let (mut client, mut server) = duplex(8); - send_ssl_response(&mut server, true).await.unwrap(); - let mut response = [0]; - client.read_exact(&mut response).await.unwrap(); - assert_eq!(response, [EncryptionReply::Accepted.as_byte()]); } } 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) +} From 9014adcbdc916e8bf57f609c7be9f2590706ce44 Mon Sep 17 00:00:00 2001 From: James Sadler Date: Wed, 5 Aug 2026 11:39:03 +1000 Subject: [PATCH 07/16] fix(proxy): order locally handled pipeline responses --- .../src/postgresql/backend.rs | 14 ++- .../src/postgresql/context/mod.rs | 107 +++++++++++++++++- .../src/postgresql/frontend.rs | 101 ++++++++++------- 3 files changed, 174 insertions(+), 48 deletions(-) diff --git a/packages/cipherstash-proxy/src/postgresql/backend.rs b/packages/cipherstash-proxy/src/postgresql/backend.rs index b4ba2d89a..846e00e76 100644 --- a/packages/cipherstash-proxy/src/postgresql/backend.rs +++ b/packages/cipherstash-proxy/src/postgresql/backend.rs @@ -915,12 +915,13 @@ mod tests { backend.context.set_execute(Name::new(), Some(session_id)); backend .context - .protocol_frontend_received(FrontendMessage::Execute( - pg_proto::codec::Execute { + .protocol_frontend_received( + FrontendMessage::Execute(pg_proto::codec::Execute { portal: Bytes::new(), max_rows: 0, - }, - )) + }), + pg_proto::pipeline::FrontendHandling::Forward, + ) .await .unwrap(); @@ -931,7 +932,10 @@ mod tests { if label == "ErrorResponse" { backend .context - .protocol_frontend_received(FrontendMessage::Sync) + .protocol_frontend_received( + FrontendMessage::Sync, + pg_proto::pipeline::FrontendHandling::Forward, + ) .await .unwrap(); let ready = diff --git a/packages/cipherstash-proxy/src/postgresql/context/mod.rs b/packages/cipherstash-proxy/src/postgresql/context/mod.rs index 30a38b4ed..406985849 100644 --- a/packages/cipherstash-proxy/src/postgresql/context/mod.rs +++ b/packages/cipherstash-proxy/src/postgresql/context/mod.rs @@ -21,7 +21,7 @@ use metrics::{counter, histogram}; use pg_proto::{ codec::{BackendMessage, Describe, DescribeTarget, FrontendMessage}, intermediary::Intermediary, - pipeline::{BackendAction, BoundedPipeline, FrontendAction, FrontendHandling}, + pipeline::{BackendAction, BoundedPipeline, FrontendAction, FrontendHandling, OperationId}, }; use serde_json::json; use sqltk::parser::ast::{Expr, Ident, ObjectName, ObjectNamePart, Set, Value, ValueWithSpan}; @@ -233,7 +233,8 @@ where pub async fn protocol_frontend_received( &self, mut message: FrontendMessage, - ) -> Result<(), Error> { + handling: FrontendHandling, + ) -> Result { loop { let notified = self.protocol_changed.notified(); tokio::pin!(notified); @@ -243,11 +244,13 @@ where protocol .sides .pipeline_mut() - .frontend_action(message, FrontendHandling::Forward) + .frontend_action(message, handling) .map_err(protocol_transition_error)? }; match action { - FrontendAction::Forward { .. } | FrontendAction::Discard { .. } => return Ok(()), + FrontendAction::Forward { id, .. } | FrontendAction::Discard { id } => { + return Ok(id); + } FrontendAction::Backpressure(returned) => { message = returned; notified.await; @@ -292,6 +295,34 @@ where } } + /// 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().map_err(protocol_lock_error)?; + protocol + .sides + .pipeline_mut() + .try_emit_local(id, message) + .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(); } @@ -1165,7 +1196,8 @@ mod tests { use cipherstash_client::IdentifiedBy; use eql_mapper::Schema; use pg_proto::codec::{ - BackendMessage, Bind, DescribeTarget, Execute, FrontendMessage, Parse, TransactionStatus, + BackendMessage, Bind, DescribeTarget, DiagnosticResponse, Execute, FrontendMessage, Parse, + TransactionStatus, }; use pg_proto::pipeline::{BackendAction, FrontendAction, FrontendHandling, PipelineState}; use sqltk::parser::{dialect::PostgreSqlDialect, parser::Parser}; @@ -1173,6 +1205,7 @@ mod tests { use tokio::sync::mpsc; use uuid::Uuid; + #[derive(Clone)] struct TestService {} #[async_trait::async_trait] @@ -1268,6 +1301,70 @@ mod tests { 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().unwrap().sides.pipeline().state(), + PipelineState::Ready + ); + } + fn statement() -> Statement { Statement { param_columns: vec![], diff --git a/packages/cipherstash-proxy/src/postgresql/frontend.rs b/packages/cipherstash-proxy/src/postgresql/frontend.rs index 0ac932fa6..389893cba 100644 --- a/packages/cipherstash-proxy/src/postgresql/frontend.rs +++ b/packages/cipherstash-proxy/src/postgresql/frontend.rs @@ -34,6 +34,7 @@ use pg_proto::{ Backend as BackendDirection, BackendMessage, Close, Describe, DescribeTarget, Execute, Frontend as FrontendDirection, FrontendMessage, TransactionStatus, }, + pipeline::{FrontendHandling, OperationId}, transport::Buffered, }; use serde::Serialize; @@ -174,15 +175,14 @@ where let mut outbound_message = protocol_message.clone(); let recovering_from_extended_error = self.context.protocol_in_extended_error()?; - self.context - .protocol_frontend_received(protocol_message.clone()) - .await?; - let frame = protocol_message.to_frame()?; let sent: u64 = (frame.body.len() + 5) as u64; counter!(CLIENTS_BYTES_RECEIVED_TOTAL).increment(sent); if self.context.mapping_disabled() { + self.context + .protocol_frontend_received(protocol_message.clone(), FrontendHandling::Forward) + .await?; self.write_to_server(outbound_message).await?; return Ok(()); } @@ -195,10 +195,14 @@ where message = ?protocol_message, ); if !matches!(protocol_message, FrontendMessage::Sync) { + self.context + .protocol_frontend_received(protocol_message, FrontendHandling::Local) + .await?; return Ok(()); } } + let tracking_message = protocol_message.clone(); match protocol_message { FrontendMessage::Query(query) => { match self.query_handler(Query::from(query)).await { @@ -211,8 +215,12 @@ where msg = "Query Handler Error", error = ?err.to_string(), ); - self.send_error_response(err).await?; - self.send_ready_for_query().await?; + 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?; return Ok(()); } } @@ -234,7 +242,11 @@ where msg = "Parse Handler Error", error = ?err.to_string(), ); - self.send_error_response(err).await?; + let id = self + .context + .protocol_frontend_received(tracking_message, FrontendHandling::Local) + .await?; + self.send_error_response(id, err).await?; return Ok(()); } } @@ -244,33 +256,39 @@ where 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).await?; - return Ok(()); - } - Error::Encrypt(EncryptError::UnknownKeysetIdentifier { .. }) => { - warn!(target: PROTOCOL, - client_id = self.context.client_id, - msg = "EncryptError::UnknownKeysetIdentifier", - ); - self.send_error_response(err).await?; - return Ok(()); - } - _ => { - warn!(target: PROTOCOL, - client_id = self.context.client_id, - msg = "Bind Error", - err = err.to_string() - ); - self.send_error_response(err).await?; - 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?; + return Ok(()); + } + Error::Encrypt(EncryptError::UnknownKeysetIdentifier { .. }) => { + warn!(target: PROTOCOL, + client_id = self.context.client_id, + msg = "EncryptError::UnknownKeysetIdentifier", + ); + self.send_error_response(id, err).await?; + return Ok(()); + } + _ => { + warn!(target: PROTOCOL, + client_id = self.context.client_id, + msg = "Bind Error", + err = err.to_string() + ); + self.send_error_response(id, err).await?; + return Ok(()); + } } - }, + } } } FrontendMessage::Sync => { @@ -286,7 +304,11 @@ where client_id = self.context.client_id, msg = "Ready for Query", ); - self.send_ready_for_query().await?; + let id = self + .context + .protocol_frontend_received(tracking_message, FrontendHandling::Local) + .await?; + self.send_ready_for_query(id).await?; return Ok(()); } } @@ -302,6 +324,9 @@ where } } + self.context + .protocol_frontend_received(tracking_message, FrontendHandling::Forward) + .await?; self.write_to_server(outbound_message).await?; Ok(()) } @@ -1195,7 +1220,7 @@ where /// /// Send an ReadyForQuery to the client and remove error state. /// - async fn send_ready_for_query(&mut self) -> Result<(), Error> { + async fn send_ready_for_query(&mut self, id: OperationId) -> Result<(), Error> { let message = BackendMessage::ReadyForQuery(TransactionStatus::Idle); debug!(target: PROTOCOL, @@ -1205,7 +1230,7 @@ where ); self.context - .protocol_backend_forwarded(message.clone()) + .protocol_backend_local(id, message.clone()) .await?; self.client_sender.send(message)?; self.context.protocol_backend_sent(); @@ -1330,7 +1355,7 @@ where W: AsyncWrite + Unpin, S: EncryptionService, { - async fn send_error_response(&mut self, err: Error) -> Result<(), Error> { + async fn send_error_response(&mut self, id: OperationId, err: Error) -> Result<(), Error> { let error_response = self.error_to_response(err); let message = error_response.into_backend_message(); @@ -1341,7 +1366,7 @@ where ); self.context - .protocol_backend_forwarded(message.clone()) + .protocol_backend_local(id, message.clone()) .await?; self.client_sender.send(message)?; self.context.protocol_backend_sent(); From 07f8e859f33a7f5a70afb2e125fb12fb3eb3ba70 Mon Sep 17 00:00:00 2001 From: James Sadler Date: Wed, 5 Aug 2026 22:42:28 +1000 Subject: [PATCH 08/16] refactor(proxy): use pg-proto network transports --- Cargo.lock | 31 +-- PG_PROTO_FOLLOWUPS.md | 14 +- packages/cipherstash-proxy/Cargo.toml | 7 +- .../src/connect/async_stream.rs | 148 -------------- packages/cipherstash-proxy/src/connect/mod.rs | 123 ++++-------- packages/cipherstash-proxy/src/main.rs | 4 +- .../src/postgresql/handler.rs | 185 ++++-------------- .../src/postgresql/startup.rs | 30 ++- packages/cipherstash-proxy/src/tls/mod.rs | 46 ++--- 9 files changed, 119 insertions(+), 469 deletions(-) delete mode 100644 packages/cipherstash-proxy/src/connect/async_stream.rs diff --git a/Cargo.lock b/Cargo.lock index 253906731..c3932320f 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -822,7 +822,6 @@ version = "2.2.4" dependencies = [ "arc-swap", "async-trait", - "aws-lc-rs", "bigdecimal", "blake3", "bytes", @@ -840,7 +839,6 @@ dependencies = [ "metrics", "metrics-exporter-prometheus", "moka", - "oid-registry", "pg-proto", "postgres-protocol", "postgres-types", @@ -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 0.17.0", ] [[package]] @@ -3027,8 +3022,7 @@ checksum = "e3148f5046208a5d56bcfc03053e3ca6334e51da8dfb19b6cdc8b306fae3283e" [[package]] name = "pg-proto" version = "0.2.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6e850c5837b91e4dfd30bd4fb4923e4173f175b48d6d421bdfd3f5316b42a198" +source = "git+https://github.com/freshtonic/pg-proto?rev=d9f201097f70f6b98b89f72d71728812463c860e#d9f201097f70f6b98b89f72d71728812463c860e" dependencies = [ "base64", "bytes", @@ -3038,19 +3032,19 @@ dependencies = [ "rand 0.10.2", "rustls", "sha2 0.11.0", + "socket2 0.6.5", "stringprep", "subtle", "tokio", "tokio-rustls", "tokio-util", - "x509-parser 0.18.1", + "x509-parser", ] [[package]] name = "pg-proto-fsm" version = "0.2.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4460a8f63f7626bd5c0a446bd5f7d0d3b72bcd41d67efea1eeee9e6ca69396ec" +source = "git+https://github.com/freshtonic/pg-proto?rev=d9f201097f70f6b98b89f72d71728812463c860e#d9f201097f70f6b98b89f72d71728812463c860e" dependencies = [ "proc-macro2", "quote", @@ -6167,23 +6161,6 @@ dependencies = [ "tls_codec", ] -[[package]] -name = "x509-parser" -version = "0.17.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4569f339c0c402346d4a75a9e39cf8dad310e287eef1ff56d4c68e5067f53460" -dependencies = [ - "asn1-rs", - "data-encoding", - "der-parser", - "lazy_static", - "nom 7.1.3", - "oid-registry", - "rusticata-macros", - "thiserror 2.0.18", - "time", -] - [[package]] name = "x509-parser" version = "0.18.1" diff --git a/PG_PROTO_FOLLOWUPS.md b/PG_PROTO_FOLLOWUPS.md index 49e8db9cb..018dc01b6 100644 --- a/PG_PROTO_FOLLOWUPS.md +++ b/PG_PROTO_FOLLOWUPS.md @@ -2,7 +2,7 @@ The proxy migration delegates framing, typed messages, startup/authentication state, demultiplexing, and bounded pipeline scheduling to `pg-proto` 0.2.1. -Three remaining adapters would be better eliminated in `pg-proto` itself. +Two remaining adapters would be better eliminated in `pg-proto` itself. ## Preserve buffered transport state across a split @@ -15,18 +15,6 @@ 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. -## Accept application-provided SCRAM channel binding - -The SCRAM-SHA-256-PLUS typestate obtains channel binding through pg-proto's TLS -transport trait. CipherStash uses its own `AsyncStream` TLS abstraction, which -already exposes the RFC 5929 binding bytes but cannot supply them through that -trait after type erasure. - -Allowing callers to provide validated channel-binding bytes (or a small adapter -trait independent of the transport type) would remove the proxy's final manual -SCRAM-PLUS exchange. The SCRAM cryptographic engine itself should remain -application-owned. - ## Transfer a demultiplexer between transport owners The proxy routes backend messages through pg-proto's demultiplexer, but startup diff --git a/packages/cipherstash-proxy/Cargo.toml b/packages/cipherstash-proxy/Cargo.toml index 5c59da691..9ceb43608 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-proto = "0.2.1" +pg-proto = { git = "https://github.com/freshtonic/pg-proto", rev = "d9f201097f70f6b98b89f72d71728812463c860e" } 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/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/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/handler.rs b/packages/cipherstash-proxy/src/postgresql/handler.rs index 95bf9f468..7e8c10e57 100644 --- a/packages/cipherstash-proxy/src/postgresql/handler.rs +++ b/packages/cipherstash-proxy/src/postgresql/handler.rs @@ -1,26 +1,25 @@ use super::backend::Backend; use super::frontend::Frontend; -use crate::connect::ChannelWriter; +use crate::connect::{self, ChannelWriter}; use crate::error::ConfigError; use crate::log::AUTHENTICATION; use crate::postgresql::messages::error_response::ErrorResponse; use crate::postgresql::startup; use crate::proxy::ZeroKms; use crate::{ - connect::AsyncStream, error::{Error, ProtocolError}, postgresql::context::Context, tls, }; -use bytes::{BufMut, Bytes, BytesMut}; +use bytes::Bytes; use md5::{Digest, Md5}; use pg_proto::pre_startup::PreStartupMessage; use pg_proto::{ auth::{AuthCompletion, AuthEvent, AuthOffer, SaslEvent}, codec::{ - Authentication, Backend as BackendDirection, BackendMessage, Frontend as FrontendDirection, - FrontendMessage, + Backend as BackendDirection, BackendMessage, Frontend as FrontendDirection, FrontendMessage, }, + net::NetworkStream, pre_startup::{PreStartup, PreStartupOffer}, server_auth::{ServerPassword, ServerProtocolOffer}, startup::ProtocolVersion, @@ -32,12 +31,14 @@ use rand::Rng; use std::time::Duration; 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; async fn receive_pre_startup( mut conn: Conn, PreStartup>, @@ -98,28 +99,10 @@ where Ok((conn, message)) } -async fn receive_auth_message( - stream: &mut Buffered, -) -> Result { - let duration = Duration::from_secs(10); - let message = timeout(duration, stream.receive_backend()) - .await - .map_err(|_| Error::ConnectionTimeout { duration })??; - let BackendMessage::Authentication(authentication) = message else { - return Err(ProtocolError::UnexpectedAuthenticationResponse { - expected: "Authentication".into(), - received: -1, - } - .into()); - }; - Ok(authentication) -} - async fn authenticate_upstream( - startup: Conn, pg_proto::pre_startup::Startup>, + startup: Conn, pg_proto::pre_startup::Startup>, context: &Context, - channel_binding: ChannelBinding, -) -> Result, Error> { +) -> Result, Error> { let mut auth = startup.authentication(); let offer = loop { let (current, message) = receive_backend_conn(auth).await?; @@ -165,26 +148,19 @@ async fn authenticate_upstream( } AuthOffer::Sasl { conn, mechanisms } => { let mechanism = sasl_mechanism(&mechanisms)?; - if mechanism == SaslMechanism::ScramSha256Plus { - // pg-proto's SCRAM-PLUS entry point currently requires its own - // TLS transport trait. Keep only the cryptographic exchange as - // an adapter until custom channel-binding bytes are accepted. - let mut transport = conn.into_transport(); - scram_sha_256_plus_handler( - &mut transport, - mechanism, - context.database_password().as_bytes(), - channel_binding, - ) - .await?; - return Ok(transport); - } - let mut scram = ScramSha256::new( context.database_password().as_bytes(), - ChannelBinding::unsupported(), + match mechanism { + SaslMechanism::ScramSha256 => ChannelBinding::unsupported(), + SaslMechanism::ScramSha256Plus => { + ChannelBinding::tls_server_end_point(conn.tls_server_end_point().to_vec()) + } + }, ); - let (mut sasl, frame) = conn.scram_sha_256(scram.message())?; + 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?; @@ -234,8 +210,8 @@ async fn authenticate_upstream( } async fn complete_upstream_auth( - awaiting: Conn, pg_proto::auth::AwaitingAuthOk>, -) -> Result, Error> { + awaiting: Conn, pg_proto::auth::AwaitingAuthOk>, +) -> Result, Error> { let (awaiting, message) = receive_backend_conn(awaiting).await?; match awaiting.offer(message) { Ok(AuthCompletion::Ok(conn)) => Ok(conn.into_transport()), @@ -251,8 +227,8 @@ async fn complete_upstream_auth( } async fn drain_upstream_startup( - database: &mut Buffered, - client: &mut Buffered, + database: &mut Buffered, + client: &mut Buffered, ) -> Result<(), Error> { loop { let message = database.receive_backend().await?; @@ -274,7 +250,7 @@ enum SaslMechanism { /// /// 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: AsyncStream, context: Context) -> Result<(), Error> { +pub async fn handler(client_stream: PgStream, context: Context) -> Result<(), Error> { let mut client_is_tls = client_stream.is_tls(); let mut client = Conn::new(Buffered::<_, FrontendDirection>::new_frontend( client_stream, @@ -282,9 +258,8 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R 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 stream = connect::connect(&context.database_socket_address()).await?; let mut database_stream = startup::with_tls(stream, context.config()).await?; - let database_channel_binding = database_stream.channel_binding(); info!( msg = "Client connected", database = context.database_socket_address(), @@ -320,17 +295,18 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R stream }; if let Some(ref tls) = context.tls_config() { - stream = match stream { - AsyncStream::Tcp(tcp_stream) => { - // The Client is connecting to our Server - let tls_stream = tls::server(tcp_stream, tls).await?; - client_is_tls = true; - AsyncStream::Tls(Box::new(tls_stream)) - } - AsyncStream::Tls(_) => { - unreachable!(); - } - }; + 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)); } @@ -362,8 +338,7 @@ pub async fn handler(client_stream: AsyncStream, context: Context) -> R .startup(&startup_message)?; database_startup.push_startup_packet(&startup_packet); database_startup.flush().await?; - let mut database_stream = - authenticate_upstream(database_startup, &context, database_channel_binding).await?; + let mut database_stream = authenticate_upstream(database_startup, &context).await?; // Proxy -> Client Authentication // Uses MD5 @@ -543,21 +518,6 @@ fn sasl_mechanism(mechanisms: &[Bytes]) -> Result { } } -fn authentication_method_code(authentication: &Authentication) -> i32 { - match authentication { - Authentication::Ok => 0, - Authentication::KerberosV5 => 2, - Authentication::CleartextPassword => 3, - Authentication::Md5Password { .. } => 5, - Authentication::Gss => 7, - Authentication::GssContinue(_) => 8, - Authentication::Sspi => 9, - Authentication::Sasl { .. } => 10, - Authentication::SaslContinue(_) => 11, - Authentication::SaslFinal(_) => 12, - } -} - pub fn md5_hash(username: &[u8], password: &[u8], salt: &[u8; 4]) -> String { let mut md5 = Md5::new(); md5.update(password); @@ -575,68 +535,6 @@ fn generate_md5_password_salt() -> [u8; 4] { bytes } -async fn scram_sha_256_plus_handler( - stream: &mut Buffered, - mechanism: SaslMechanism, - password: &[u8], - channel_binding: ChannelBinding, -) -> Result<(), Error> { - let mut scram = ScramSha256::new(password, channel_binding); - let bytes = scram.message().to_vec(); - - let mechanism = match mechanism { - SaslMechanism::ScramSha256 => SCRAM_SHA_256, - SaslMechanism::ScramSha256Plus => SCRAM_SHA_256_PLUS, - }; - let mut initial = BytesMut::new(); - initial.extend_from_slice(mechanism); - initial.put_u8(0); - initial.put_i32(bytes.len().try_into().map_err(|_| { - std::io::Error::new( - std::io::ErrorKind::InvalidInput, - "SASL response is too large", - ) - })?); - initial.extend_from_slice(&bytes); - send_frontend_message(stream, FrontendMessage::PasswordResponse(initial.freeze())).await?; - - let auth = receive_auth_message(stream).await?; - - let Authentication::SaslContinue(bytes) = auth else { - return Err(ProtocolError::UnexpectedAuthenticationResponse { - expected: "SaslContinue".into(), - received: authentication_method_code(&auth), - } - .into()); - }; - scram.update(&bytes)?; - - send_frontend_message( - stream, - FrontendMessage::PasswordResponse(Bytes::copy_from_slice(scram.message())), - ) - .await?; - - let auth = receive_auth_message(stream).await?; - let Authentication::SaslFinal(bytes) = auth else { - return Err(ProtocolError::UnexpectedAuthenticationResponse { - expected: "SaslFinal".into(), - received: authentication_method_code(&auth), - } - .into()); - }; - scram.finish(&bytes)?; - - let auth = receive_auth_message(stream).await?; - - if matches!(auth, Authentication::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( @@ -655,12 +553,3 @@ async fn send_backend_message( stream.flush().await?; Ok(()) } - -async fn send_frontend_message( - stream: &mut Buffered, - message: FrontendMessage, -) -> Result<(), Error> { - stream.push(message.to_frame()?)?; - stream.flush().await?; - Ok(()) -} diff --git a/packages/cipherstash-proxy/src/postgresql/startup.rs b/packages/cipherstash-proxy/src/postgresql/startup.rs index 0464989be..ba424128b 100644 --- a/packages/cipherstash-proxy/src/postgresql/startup.rs +++ b/packages/cipherstash-proxy/src/postgresql/startup.rs @@ -1,15 +1,21 @@ -use pg_proto::{codec::Backend, pre_startup::Negotiation, transport::Buffered, Conn}; +use pg_proto::{ + codec::Backend, net::NetworkStream, pre_startup::Negotiation, transport::Buffered, Conn, +}; +use std::sync::Arc; +use tokio::net::TcpStream; use tracing::warn; use crate::{ - connect::AsyncStream, error::{Error, ProtocolError}, tls, TandemConfig, }; /// Applies CipherStash's upstream TLS policy while pg-proto owns the wire /// negotiation and its legal pre-startup transitions. -pub async fn with_tls(stream: AsyncStream, config: &TandemConfig) -> Result { +pub async fn with_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); @@ -19,12 +25,18 @@ pub async fn with_tls(stream: AsyncStream, config: &TandemConfig) -> Result { - let AsyncStream::Tcp(stream) = conn.into_transport().into_inner() else { - return Err(ProtocolError::UnexpectedStartupMessage.into()); - }; - Ok(AsyncStream::Tls(Box::new( - tls::client(stream, config).await?, - ))) + 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)"); 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)) } /// From 1fa5bb8819bfabef749ce7ce38a28b85af1c61af Mon Sep 17 00:00:00 2001 From: James Sadler Date: Wed, 5 Aug 2026 22:45:06 +1000 Subject: [PATCH 09/16] build(proxy): pin verified pg-proto transport --- Cargo.lock | 6 +++--- packages/cipherstash-proxy/Cargo.toml | 2 +- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index c3932320f..398d147a5 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3022,7 +3022,7 @@ checksum = "e3148f5046208a5d56bcfc03053e3ca6334e51da8dfb19b6cdc8b306fae3283e" [[package]] name = "pg-proto" version = "0.2.1" -source = "git+https://github.com/freshtonic/pg-proto?rev=d9f201097f70f6b98b89f72d71728812463c860e#d9f201097f70f6b98b89f72d71728812463c860e" +source = "git+https://github.com/freshtonic/pg-proto?rev=dc30c3c86a364262b9b95e036cfac2f25b39e688#dc30c3c86a364262b9b95e036cfac2f25b39e688" dependencies = [ "base64", "bytes", @@ -3044,7 +3044,7 @@ dependencies = [ [[package]] name = "pg-proto-fsm" version = "0.2.1" -source = "git+https://github.com/freshtonic/pg-proto?rev=d9f201097f70f6b98b89f72d71728812463c860e#d9f201097f70f6b98b89f72d71728812463c860e" +source = "git+https://github.com/freshtonic/pg-proto?rev=dc30c3c86a364262b9b95e036cfac2f25b39e688#dc30c3c86a364262b9b95e036cfac2f25b39e688" dependencies = [ "proc-macro2", "quote", @@ -3900,7 +3900,7 @@ dependencies = [ "security-framework", "security-framework-sys", "webpki-root-certs 1.0.5", - "windows-sys 0.61.2", + "windows-sys 0.59.0", ] [[package]] diff --git a/packages/cipherstash-proxy/Cargo.toml b/packages/cipherstash-proxy/Cargo.toml index 9ceb43608..dd5904f38 100644 --- a/packages/cipherstash-proxy/Cargo.toml +++ b/packages/cipherstash-proxy/Cargo.toml @@ -28,7 +28,7 @@ md-5 = "0.10.6" metrics = "0.24.3" metrics-exporter-prometheus = "0.17" moka = { version = "0.12", features = ["future"] } -pg-proto = { git = "https://github.com/freshtonic/pg-proto", rev = "d9f201097f70f6b98b89f72d71728812463c860e" } +pg-proto = { git = "https://github.com/freshtonic/pg-proto", rev = "dc30c3c86a364262b9b95e036cfac2f25b39e688" } postgres-protocol = "0.6.7" postgres-types = { version = "0.2.8", features = ["with-serde_json-1"] } rand = "0.9" From 04b5ce4423659aa0ce5449a619dc08538facb34c Mon Sep 17 00:00:00 2001 From: James Sadler Date: Wed, 5 Aug 2026 23:11:18 +1000 Subject: [PATCH 10/16] build(proxy): use pg-proto 0.2.2 --- Cargo.lock | 12 +++++++----- packages/cipherstash-proxy/Cargo.toml | 2 +- 2 files changed, 8 insertions(+), 6 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 398d147a5..e6e41b52c 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3021,8 +3021,9 @@ checksum = "e3148f5046208a5d56bcfc03053e3ca6334e51da8dfb19b6cdc8b306fae3283e" [[package]] name = "pg-proto" -version = "0.2.1" -source = "git+https://github.com/freshtonic/pg-proto?rev=dc30c3c86a364262b9b95e036cfac2f25b39e688#dc30c3c86a364262b9b95e036cfac2f25b39e688" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0d0c8bbd7d0f0458aa4d9ece0e6ffe28e2c0e306fdc51c5f267400980ec589d5" dependencies = [ "base64", "bytes", @@ -3043,8 +3044,9 @@ dependencies = [ [[package]] name = "pg-proto-fsm" -version = "0.2.1" -source = "git+https://github.com/freshtonic/pg-proto?rev=dc30c3c86a364262b9b95e036cfac2f25b39e688#dc30c3c86a364262b9b95e036cfac2f25b39e688" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e381acb6b853a785a8a92ce609ae0b4f77a6a08d561c814d8a4b2109997a5b5b" dependencies = [ "proc-macro2", "quote", @@ -3900,7 +3902,7 @@ dependencies = [ "security-framework", "security-framework-sys", "webpki-root-certs 1.0.5", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] diff --git a/packages/cipherstash-proxy/Cargo.toml b/packages/cipherstash-proxy/Cargo.toml index dd5904f38..c9ddb5667 100644 --- a/packages/cipherstash-proxy/Cargo.toml +++ b/packages/cipherstash-proxy/Cargo.toml @@ -28,7 +28,7 @@ md-5 = "0.10.6" metrics = "0.24.3" metrics-exporter-prometheus = "0.17" moka = { version = "0.12", features = ["future"] } -pg-proto = { git = "https://github.com/freshtonic/pg-proto", rev = "dc30c3c86a364262b9b95e036cfac2f25b39e688" } +pg-proto = "0.2.2" postgres-protocol = "0.6.7" postgres-types = { version = "0.2.8", features = ["with-serde_json-1"] } rand = "0.9" From 3c4e9f2deb1ff60acebb199175e8d21c50a7b92a Mon Sep 17 00:00:00 2001 From: James Sadler Date: Thu, 6 Aug 2026 12:42:52 +1000 Subject: [PATCH 11/16] refactor(proxy): adopt pg-proto typed middleware --- Cargo.lock | 8 +- packages/cipherstash-proxy/Cargo.toml | 2 +- .../src/postgresql/handler.rs | 174 ++++++++++++------ .../src/postgresql/startup.rs | 23 ++- 4 files changed, 142 insertions(+), 65 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index e6e41b52c..ff24bad3e 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3021,9 +3021,9 @@ checksum = "e3148f5046208a5d56bcfc03053e3ca6334e51da8dfb19b6cdc8b306fae3283e" [[package]] name = "pg-proto" -version = "0.2.2" +version = "0.2.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0d0c8bbd7d0f0458aa4d9ece0e6ffe28e2c0e306fdc51c5f267400980ec589d5" +checksum = "b6264dfbe018e8c8b34752fa64174df87a5f1bc8669ac5a3bbcce038c7671e58" dependencies = [ "base64", "bytes", @@ -3044,9 +3044,9 @@ dependencies = [ [[package]] name = "pg-proto-fsm" -version = "0.2.2" +version = "0.2.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e381acb6b853a785a8a92ce609ae0b4f77a6a08d561c814d8a4b2109997a5b5b" +checksum = "af49ce51ac27499ddbd2741609d99514d7bf5850d3a10b7fdf0a0709cb1178f1" dependencies = [ "proc-macro2", "quote", diff --git a/packages/cipherstash-proxy/Cargo.toml b/packages/cipherstash-proxy/Cargo.toml index c9ddb5667..57d202ea0 100644 --- a/packages/cipherstash-proxy/Cargo.toml +++ b/packages/cipherstash-proxy/Cargo.toml @@ -28,7 +28,7 @@ md-5 = "0.10.6" metrics = "0.24.3" metrics-exporter-prometheus = "0.17" moka = { version = "0.12", features = ["future"] } -pg-proto = "0.2.2" +pg-proto = "0.2.3" postgres-protocol = "0.6.7" postgres-types = { version = "0.2.8", features = ["with-serde_json-1"] } rand = "0.9" diff --git a/packages/cipherstash-proxy/src/postgresql/handler.rs b/packages/cipherstash-proxy/src/postgresql/handler.rs index 7e8c10e57..7fbdd5f5e 100644 --- a/packages/cipherstash-proxy/src/postgresql/handler.rs +++ b/packages/cipherstash-proxy/src/postgresql/handler.rs @@ -15,10 +15,11 @@ use bytes::Bytes; use md5::{Digest, Md5}; use pg_proto::pre_startup::PreStartupMessage; use pg_proto::{ - auth::{AuthCompletion, AuthEvent, AuthOffer, SaslEvent}, + auth::{AuthCompletion, AuthEvent, AuthOffer, AwaitingStartupReady, SaslEvent}, codec::{ Backend as BackendDirection, BackendMessage, Frontend as FrontendDirection, FrontendMessage, }, + middleware::{Identity, Middleware, ServerRole, TypedPhase, TypedReceiveError}, net::NetworkStream, pre_startup::{PreStartup, PreStartupOffer}, server_auth::{ServerPassword, ServerProtocolOffer}, @@ -39,9 +40,33 @@ 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< ( @@ -51,20 +76,22 @@ async fn receive_pre_startup( (Conn, PreStartup>, Error), > { let received = match connection_timeout { - Some(duration) => match timeout(duration, conn.receive_pre_startup_wire()).await { + 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_wire().await, + None => conn.receive_pre_startup_typed(middleware).await, }; match received { - Ok(message) => Ok((conn, message)), - Err(error) => Err((conn, error.into())), + 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< ( @@ -74,38 +101,43 @@ async fn receive_frontend_auth( (Conn, ServerPassword>, Error), > { let received = match connection_timeout { - Some(duration) => match timeout(duration, conn.receive_frontend_wire()).await { + 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_wire().await, + None => conn.receive_frontend_typed(middleware).await, }; match received { - Ok(message) => Ok((conn, message)), - Err(error) => Err((conn, error.into())), + 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: TypedPhase, + >::Message: Into, { let duration = Duration::from_secs(10); - let message = timeout(duration, conn.receive_backend_wire()) + let message = timeout(duration, conn.receive_backend_typed(middleware)) .await - .map_err(|_| Error::ConnectionTimeout { duration })??; - Ok((conn, message)) + .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, -) -> Result, Error> { + middleware: &mut ProtocolMiddleware, +) -> Result, AwaitingStartupReady>, Error> { let mut auth = startup.authentication(); let offer = loop { - let (current, message) = receive_backend_conn(auth).await?; + 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, @@ -128,12 +160,12 @@ async fn authenticate_upstream( }; match offer { - AuthOffer::Ok(conn) => Ok(conn.into_transport()), + 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).await + complete_upstream_auth(awaiting, middleware).await } AuthOffer::Md5 { conn, salt } => { let hash = md5_hash( @@ -144,7 +176,7 @@ async fn authenticate_upstream( let (mut awaiting, frame) = conn.password(hash.as_bytes())?; awaiting.push_frame(frame)?; awaiting.flush().await?; - complete_upstream_auth(awaiting).await + complete_upstream_auth(awaiting, middleware).await } AuthOffer::Sasl { conn, mechanisms } => { let mechanism = sasl_mechanism(&mechanisms)?; @@ -164,7 +196,7 @@ async fn authenticate_upstream( sasl.push_frame(frame)?; sasl.flush().await?; - let (sasl, message) = receive_backend_conn(sasl).await?; + let (sasl, message) = receive_backend_conn(sasl, middleware).await?; let BackendMessage::Authentication(authentication) = message else { sasl.into_transport(); return Err(ProtocolError::UnexpectedStartupMessage.into()); @@ -184,7 +216,7 @@ async fn authenticate_upstream( sasl.push_frame(frame)?; sasl.flush().await?; - let (sasl, message) = receive_backend_conn(sasl).await?; + let (sasl, message) = receive_backend_conn(sasl, middleware).await?; let BackendMessage::Authentication(authentication) = message else { sasl.into_transport(); return Err(ProtocolError::UnexpectedStartupMessage.into()); @@ -200,7 +232,7 @@ async fn authenticate_upstream( return Err(ProtocolError::AuthenticationFailed.into()); }; scram.finish(&server_final)?; - complete_upstream_auth(final_state.verified()).await + complete_upstream_auth(final_state.verified(), middleware).await } AuthOffer::Gss(conn) | AuthOffer::Sspi(conn) | AuthOffer::KerberosV5(conn) => { conn.into_transport(); @@ -211,10 +243,11 @@ async fn authenticate_upstream( async fn complete_upstream_auth( awaiting: Conn, pg_proto::auth::AwaitingAuthOk>, -) -> Result, Error> { - let (awaiting, message) = receive_backend_conn(awaiting).await?; + 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.into_transport()), + Ok(AuthCompletion::Ok(conn)) => Ok(conn), Ok(AuthCompletion::Error { conn, .. }) => { conn.into_transport(); Err(ProtocolError::AuthenticationFailed.into()) @@ -227,16 +260,23 @@ async fn complete_upstream_auth( } async fn drain_upstream_startup( - database: &mut Buffered, + mut database: Conn, AwaitingStartupReady>, client: &mut Buffered, -) -> Result<(), Error> { + middleware: &mut ProtocolMiddleware, +) -> Result, Error> { loop { - let message = database.receive_backend().await?; - let ready = matches!(message, BackendMessage::ReadyForQuery(_)); - let _ = database.project_backend(message.clone()); + 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 ready { - return Ok(()); + if let Some(item) = session_item { + match database.offer_ready(item) { + Ok(ready) => return Ok(ready.into_transport()), + Err((conn, _)) => database = conn, + } } } } @@ -251,6 +291,8 @@ enum SaslMechanism { /// 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, @@ -267,19 +309,24 @@ pub async fn handler(client_stream: PgStream, context: Context) -> Resu ); let (client_startup, startup_message) = loop { - let (pre_startup, startup_message) = - match receive_pre_startup(client, 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); - } - }; + 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) => { @@ -338,7 +385,8 @@ pub async fn handler(client_stream: PgStream, context: Context) -> Resu .startup(&startup_message)?; database_startup.push_startup_packet(&startup_packet); database_startup.flush().await?; - let mut database_stream = authenticate_upstream(database_startup, &context).await?; + let database_startup = + authenticate_upstream(database_startup, &context, &mut upstream_middleware).await?; // Proxy -> Client Authentication // Uses MD5 @@ -371,19 +419,24 @@ pub async fn handler(client_stream: PgStream, context: Context) -> Resu password_state.flush().await?; let connection_timeout = context.connection_timeout(); - let (password_state, message) = - match receive_frontend_auth(password_state, 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 (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 @@ -419,7 +472,12 @@ pub async fn handler(client_stream: PgStream, context: Context) -> Resu return Err(ConfigError::TlsRequired.into()); } - drain_upstream_startup(&mut database_stream, &mut client_stream).await?; + 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(); diff --git a/packages/cipherstash-proxy/src/postgresql/startup.rs b/packages/cipherstash-proxy/src/postgresql/startup.rs index ba424128b..c3b473620 100644 --- a/packages/cipherstash-proxy/src/postgresql/startup.rs +++ b/packages/cipherstash-proxy/src/postgresql/startup.rs @@ -1,5 +1,10 @@ use pg_proto::{ - codec::Backend, net::NetworkStream, pre_startup::Negotiation, transport::Buffered, Conn, + codec::Backend, + middleware::{Identity, Middleware, TypedReceiveError}, + net::NetworkStream, + pre_startup::Negotiation, + transport::Buffered, + Conn, }; use std::sync::Arc; use tokio::net::TcpStream; @@ -23,7 +28,21 @@ pub async fn with_tls( let mut request = Conn::new(Buffered::<_, Backend>::new(stream)).request_ssl(); request.flush().await?; - match request.receive_ssl_reply().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() From 84c427f07793350cadc69a11dbf4b3adc443b765 Mon Sep 17 00:00:00 2001 From: James Sadler Date: Thu, 6 Aug 2026 12:43:38 +1000 Subject: [PATCH 12/16] docs: note async typed middleware follow-up --- PG_PROTO_FOLLOWUPS.md | 20 ++++++++++++++++++-- 1 file changed, 18 insertions(+), 2 deletions(-) diff --git a/PG_PROTO_FOLLOWUPS.md b/PG_PROTO_FOLLOWUPS.md index 018dc01b6..8e1130ccb 100644 --- a/PG_PROTO_FOLLOWUPS.md +++ b/PG_PROTO_FOLLOWUPS.md @@ -1,8 +1,24 @@ # pg-proto follow-ups The proxy migration delegates framing, typed messages, startup/authentication -state, demultiplexing, and bounded pipeline scheduling to `pg-proto` 0.2.1. -Two remaining adapters would be better eliminated in `pg-proto` itself. +state, typed startup middleware, demultiplexing, and bounded pipeline scheduling +to `pg-proto` 0.2.3. The remaining adapters below would be better eliminated in +`pg-proto` itself. + +## Support asynchronous typed middleware for pipelined runtime sessions + +CipherStash query mapping, encryption, and batched decryption are asynchronous. +`TypedMiddleware::intercept_typed` is synchronous, and its phase index is derived +from a compile-time `Conn` typestate. After startup, the proxy deliberately runs +client-to-server and server-to-client processing concurrently through +`BoundedPipeline`, whose exact projected phases are runtime-selected. + +The proxy therefore uses typed middleware throughout pre-startup, TLS, +authentication, and startup completion, but retains its asynchronous runtime +rewrite handlers behind the bounded pipeline's legality checks. An asynchronous +typed middleware interface integrated with `Pipeline` admissions and responses +would let those handlers receive and return phase-legal generated message types +without serializing the two traffic directions. ## Preserve buffered transport state across a split From 4833dd959d1b47266fc1be4d2c5b961a00f67b63 Mon Sep 17 00:00:00 2001 From: James Sadler Date: Thu, 6 Aug 2026 22:04:55 +1000 Subject: [PATCH 13/16] refactor(proxy): adopt pg-proto async middleware --- Cargo.lock | 8 +- PG_PROTO_FOLLOWUPS.md | 29 +++--- packages/cipherstash-proxy/Cargo.toml | 2 +- .../src/postgresql/backend.rs | 68 ++++++++++++-- .../src/postgresql/frontend.rs | 92 +++++++++++++++---- 5 files changed, 155 insertions(+), 44 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index ff24bad3e..2857f3584 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3021,9 +3021,9 @@ checksum = "e3148f5046208a5d56bcfc03053e3ca6334e51da8dfb19b6cdc8b306fae3283e" [[package]] name = "pg-proto" -version = "0.2.3" +version = "0.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b6264dfbe018e8c8b34752fa64174df87a5f1bc8669ac5a3bbcce038c7671e58" +checksum = "ab8c5d419269cf3cb686939a183172a323230b2a438e3609c35361e3bb026d4f" dependencies = [ "base64", "bytes", @@ -3044,9 +3044,9 @@ dependencies = [ [[package]] name = "pg-proto-fsm" -version = "0.2.3" +version = "0.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "af49ce51ac27499ddbd2741609d99514d7bf5850d3a10b7fdf0a0709cb1178f1" +checksum = "0a70723cdea3619b778ba0a1e3d8747b086d0029513804802babf6b31b27123c" dependencies = [ "proc-macro2", "quote", diff --git a/PG_PROTO_FOLLOWUPS.md b/PG_PROTO_FOLLOWUPS.md index 8e1130ccb..b52974875 100644 --- a/PG_PROTO_FOLLOWUPS.md +++ b/PG_PROTO_FOLLOWUPS.md @@ -1,24 +1,29 @@ # pg-proto follow-ups The proxy migration delegates framing, typed messages, startup/authentication -state, typed startup middleware, demultiplexing, and bounded pipeline scheduling -to `pg-proto` 0.2.3. The remaining adapters below would be better eliminated in -`pg-proto` itself. +state, typed startup middleware, async runtime middleware, demultiplexing, and +bounded pipeline scheduling to `pg-proto` 0.3.0. The remaining adapters below +would be better eliminated in `pg-proto` itself. -## Support asynchronous typed middleware for pipelined runtime sessions +## Integrate typed middleware with pipelined runtime sessions CipherStash query mapping, encryption, and batched decryption are asynchronous. -`TypedMiddleware::intercept_typed` is synchronous, and its phase index is derived -from a compile-time `Conn` typestate. After startup, the proxy deliberately runs -client-to-server and server-to-client processing concurrently through -`BoundedPipeline`, whose exact projected phases are runtime-selected. +Version 0.3.0 allows those operations to run inside async middleware. However, +`TypedMiddleware` derives its phase index from a single compile-time `Conn` +typestate. After startup, the proxy deliberately runs client-to-server and +server-to-client processing concurrently through `BoundedPipeline`, whose exact +projected phases are runtime-selected and may include multiple outstanding +operations. The proxy therefore uses typed middleware throughout pre-startup, TLS, -authentication, and startup completion, but retains its asynchronous runtime -rewrite handlers behind the bounded pipeline's legality checks. An asynchronous -typed middleware interface integrated with `Pipeline` admissions and responses +authentication, and startup completion. Runtime rewriting uses async +direction-specific `MessageMiddleware`, followed by the bounded pipeline's +legality checks. Typed middleware hooks on `Pipeline` admissions and responses would let those handlers receive and return phase-legal generated message types -without serializing the two traffic directions. +without serializing or de-pipelining the two traffic directions. Those hooks +also need explicit outcomes for locally handled frontend operations and +suppressed/buffered backend messages, which cannot be represented by a +same-message-in/same-message-out middleware result. ## Preserve buffered transport state across a split diff --git a/packages/cipherstash-proxy/Cargo.toml b/packages/cipherstash-proxy/Cargo.toml index 57d202ea0..1bede91a1 100644 --- a/packages/cipherstash-proxy/Cargo.toml +++ b/packages/cipherstash-proxy/Cargo.toml @@ -28,7 +28,7 @@ md-5 = "0.10.6" metrics = "0.24.3" metrics-exporter-prometheus = "0.17" moka = { version = "0.12", features = ["future"] } -pg-proto = "0.2.3" +pg-proto = "0.3.0" postgres-protocol = "0.6.7" postgres-types = { version = "0.2.8", features = ["with-serde_json-1"] } rand = "0.9" diff --git a/packages/cipherstash-proxy/src/postgresql/backend.rs b/packages/cipherstash-proxy/src/postgresql/backend.rs index 846e00e76..ad5214379 100644 --- a/packages/cipherstash-proxy/src/postgresql/backend.rs +++ b/packages/cipherstash-proxy/src/postgresql/backend.rs @@ -22,6 +22,7 @@ use crate::EqlCiphertext; use metrics::{counter, histogram}; use pg_proto::{ codec::{Backend as BackendDirection, BackendMessage}, + middleware::{MessageMiddleware, Middleware}, transport::Buffered, }; use std::time::Instant; @@ -100,6 +101,37 @@ where buffer: MessageBuffer, } +#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)] +enum BackendDisposition { + #[default] + Emit, + Suppress, +} + +struct BackendInterceptor<'a, R, S> +where + R: AsyncRead + Unpin, + S: EncryptionService, +{ + backend: &'a mut Backend, +} + +impl MessageMiddleware for BackendInterceptor<'_, R, S> +where + R: AsyncRead + Unpin, + S: EncryptionService, +{ + type Error = Error; + + async fn intercept( + &mut self, + disposition: &mut BackendDisposition, + message: BackendMessage, + ) -> Result { + self.backend.intercept_backend(disposition, message).await + } +} + impl Backend where R: AsyncRead + Unpin, @@ -167,8 +199,6 @@ where let read_start = Instant::now(); let protocol_message = receive_backend(&mut self.server_reader, self.context.connection_timeout()).await?; - let mut outbound_message = protocol_message.clone(); - let session_item = self.server_reader.project_backend(protocol_message.clone()); if session_item.is_none() { // The demux has recorded the asynchronous event and its ordering; @@ -193,13 +223,34 @@ where ); } + let (outbound_message, disposition) = { + let mut middleware = Middleware::new( + BackendDisposition::Emit, + BackendInterceptor { backend: self }, + ); + let outbound_message = middleware.intercept(protocol_message).await?; + (outbound_message, *middleware.state()) + }; + + if disposition == BackendDisposition::Emit { + self.write_with_flush(outbound_message).await?; + } + + Ok(()) + } + + 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(outbound_message).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 @@ -220,7 +271,7 @@ where _ => {} } - return Ok(()); + return Ok(outbound_message); } let keyset_id = self.context.keyset_identifier(); @@ -231,7 +282,8 @@ where // Encrypted DataRows are added to the buffer and we return early // Otherwise, continue and write if self.data_row_handler(DataRow::from(row)).await? { - return Ok(()); + *disposition = BackendDisposition::Suppress; + return Ok(outbound_message); } } @@ -319,9 +371,7 @@ where } } - self.write_with_flush(outbound_message).await?; - - Ok(()) + Ok(outbound_message) } /// Handles PostgreSQL ErrorResponse messages from the server. diff --git a/packages/cipherstash-proxy/src/postgresql/frontend.rs b/packages/cipherstash-proxy/src/postgresql/frontend.rs index 389893cba..6268a7602 100644 --- a/packages/cipherstash-proxy/src/postgresql/frontend.rs +++ b/packages/cipherstash-proxy/src/postgresql/frontend.rs @@ -34,6 +34,7 @@ use pg_proto::{ Backend as BackendDirection, BackendMessage, Close, Describe, DescribeTarget, Execute, Frontend as FrontendDirection, FrontendMessage, TransactionStatus, }, + middleware::{MessageMiddleware, Middleware}, pipeline::{FrontendHandling, OperationId}, transport::Buffered, }; @@ -118,6 +119,40 @@ where context: Context, } +#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)] +enum FrontendDisposition { + #[default] + Forward, + Local, +} + +struct FrontendInterceptor<'a, R, W, S> +where + R: AsyncRead + Unpin, + W: AsyncWrite + Unpin, + S: EncryptionService, +{ + frontend: &'a mut Frontend, +} + +impl MessageMiddleware + for FrontendInterceptor<'_, R, W, S> +where + R: AsyncRead + Unpin, + W: AsyncWrite + Unpin, + S: EncryptionService, +{ + type Error = Error; + + async fn intercept( + &mut self, + disposition: &mut FrontendDisposition, + message: FrontendMessage, + ) -> Result { + self.frontend.intercept_frontend(disposition, message).await + } +} + impl Frontend where R: AsyncRead + Unpin, @@ -172,21 +207,43 @@ where pub async fn rewrite(&mut self) -> Result<(), Error> { let protocol_message = receive_frontend(&mut self.client_reader, self.context.connection_timeout()).await?; - let mut outbound_message = protocol_message.clone(); - - let recovering_from_extended_error = self.context.protocol_in_extended_error()?; let frame = protocol_message.to_frame()?; let sent: u64 = (frame.body.len() + 5) as u64; counter!(CLIENTS_BYTES_RECEIVED_TOTAL).increment(sent); - if self.context.mapping_disabled() { + let tracking_message = protocol_message.clone(); + let (outbound_message, disposition) = { + let mut middleware = Middleware::new( + FrontendDisposition::Forward, + FrontendInterceptor { frontend: self }, + ); + let outbound_message = middleware.intercept(protocol_message).await?; + (outbound_message, *middleware.state()) + }; + + if disposition == FrontendDisposition::Forward { self.context - .protocol_frontend_received(protocol_message.clone(), FrontendHandling::Forward) + .protocol_frontend_received(tracking_message, FrontendHandling::Forward) .await?; self.write_to_server(outbound_message).await?; - return Ok(()); } + Ok(()) + } + + 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() { + return Ok(outbound_message); + } + + let recovering_from_extended_error = self.context.protocol_in_extended_error()?; + // 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 recovering_from_extended_error { @@ -198,7 +255,8 @@ where self.context .protocol_frontend_received(protocol_message, FrontendHandling::Local) .await?; - return Ok(()); + *disposition = FrontendDisposition::Local; + return Ok(outbound_message); } } @@ -221,7 +279,8 @@ where .await?; self.send_error_response(id, err).await?; self.send_ready_for_query(id).await?; - return Ok(()); + *disposition = FrontendDisposition::Local; + return Ok(outbound_message); } } } @@ -247,7 +306,8 @@ where .protocol_frontend_received(tracking_message, FrontendHandling::Local) .await?; self.send_error_response(id, err).await?; - return Ok(()); + *disposition = FrontendDisposition::Local; + return Ok(outbound_message); } } } @@ -268,7 +328,6 @@ where msg = "EncryptError::InvalidParameter", ); self.send_error_response(id, err).await?; - return Ok(()); } Error::Encrypt(EncryptError::UnknownKeysetIdentifier { .. }) => { warn!(target: PROTOCOL, @@ -276,7 +335,6 @@ where msg = "EncryptError::UnknownKeysetIdentifier", ); self.send_error_response(id, err).await?; - return Ok(()); } _ => { warn!(target: PROTOCOL, @@ -285,9 +343,10 @@ where err = err.to_string() ); self.send_error_response(id, err).await?; - return Ok(()); } } + *disposition = FrontendDisposition::Local; + return Ok(outbound_message); } } } @@ -309,7 +368,8 @@ where .protocol_frontend_received(tracking_message, FrontendHandling::Local) .await?; self.send_ready_for_query(id).await?; - return Ok(()); + *disposition = FrontendDisposition::Local; + return Ok(outbound_message); } } FrontendMessage::Close(close) => { @@ -324,11 +384,7 @@ where } } - self.context - .protocol_frontend_received(tracking_message, FrontendHandling::Forward) - .await?; - self.write_to_server(outbound_message).await?; - Ok(()) + Ok(outbound_message) } pub async fn write_to_server(&mut self, message: FrontendMessage) -> Result<(), Error> { From ce1d80eda4709e518b21f2b37663c1a06a1424b8 Mon Sep 17 00:00:00 2001 From: James Sadler Date: Sat, 8 Aug 2026 11:59:52 +1000 Subject: [PATCH 14/16] refactor(proxy): use pg-proto typed pipeline dispatch --- Cargo.lock | 8 +- PG_PROTO_FOLLOWUPS.md | 25 +----- packages/cipherstash-proxy/Cargo.toml | 2 +- .../src/postgresql/context/mod.rs | 77 ++++++++++++------- .../src/postgresql/frontend.rs | 2 +- .../src/postgresql/handler.rs | 6 +- 6 files changed, 60 insertions(+), 60 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 2857f3584..fc82ab6e9 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3021,9 +3021,9 @@ checksum = "e3148f5046208a5d56bcfc03053e3ca6334e51da8dfb19b6cdc8b306fae3283e" [[package]] name = "pg-proto" -version = "0.3.0" +version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ab8c5d419269cf3cb686939a183172a323230b2a438e3609c35361e3bb026d4f" +checksum = "5275df13ca49c9ef28ad5d63006b6d93a38d6f1ab9f9655475b305be0a8c4e58" dependencies = [ "base64", "bytes", @@ -3044,9 +3044,9 @@ dependencies = [ [[package]] name = "pg-proto-fsm" -version = "0.3.0" +version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0a70723cdea3619b778ba0a1e3d8747b086d0029513804802babf6b31b27123c" +checksum = "3574ed461ac3cbc508cb2c0d44e8280c346a46856e3760bb544c0bc7ec4bcc64" dependencies = [ "proc-macro2", "quote", diff --git a/PG_PROTO_FOLLOWUPS.md b/PG_PROTO_FOLLOWUPS.md index b52974875..cee88a508 100644 --- a/PG_PROTO_FOLLOWUPS.md +++ b/PG_PROTO_FOLLOWUPS.md @@ -2,28 +2,9 @@ The proxy migration delegates framing, typed messages, startup/authentication state, typed startup middleware, async runtime middleware, demultiplexing, and -bounded pipeline scheduling to `pg-proto` 0.3.0. The remaining adapters below -would be better eliminated in `pg-proto` itself. - -## Integrate typed middleware with pipelined runtime sessions - -CipherStash query mapping, encryption, and batched decryption are asynchronous. -Version 0.3.0 allows those operations to run inside async middleware. However, -`TypedMiddleware` derives its phase index from a single compile-time `Conn` -typestate. After startup, the proxy deliberately runs client-to-server and -server-to-client processing concurrently through `BoundedPipeline`, whose exact -projected phases are runtime-selected and may include multiple outstanding -operations. - -The proxy therefore uses typed middleware throughout pre-startup, TLS, -authentication, and startup completion. Runtime rewriting uses async -direction-specific `MessageMiddleware`, followed by the bounded pipeline's -legality checks. Typed middleware hooks on `Pipeline` admissions and responses -would let those handlers receive and return phase-legal generated message types -without serializing or de-pipelining the two traffic directions. Those hooks -also need explicit outcomes for locally handled frontend operations and -suppressed/buffered backend messages, which cannot be represented by a -same-message-in/same-message-out middleware result. +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 diff --git a/packages/cipherstash-proxy/Cargo.toml b/packages/cipherstash-proxy/Cargo.toml index 1bede91a1..e7a0e925b 100644 --- a/packages/cipherstash-proxy/Cargo.toml +++ b/packages/cipherstash-proxy/Cargo.toml @@ -28,7 +28,7 @@ md-5 = "0.10.6" metrics = "0.24.3" metrics-exporter-prometheus = "0.17" moka = { version = "0.12", features = ["future"] } -pg-proto = "0.3.0" +pg-proto = "0.5.0" postgres-protocol = "0.6.7" postgres-types = { version = "0.2.8", features = ["with-serde_json-1"] } rand = "0.9" diff --git a/packages/cipherstash-proxy/src/postgresql/context/mod.rs b/packages/cipherstash-proxy/src/postgresql/context/mod.rs index 406985849..6037312c8 100644 --- a/packages/cipherstash-proxy/src/postgresql/context/mod.rs +++ b/packages/cipherstash-proxy/src/postgresql/context/mod.rs @@ -21,7 +21,11 @@ use metrics::{counter, histogram}; use pg_proto::{ codec::{BackendMessage, Describe, DescribeTarget, FrontendMessage}, intermediary::Intermediary, - pipeline::{BackendAction, BoundedPipeline, FrontendAction, FrontendHandling, OperationId}, + 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}; @@ -30,11 +34,11 @@ use std::{ collections::{HashMap, VecDeque}, sync::{ atomic::{AtomicU64, Ordering}, - Arc, LazyLock, Mutex, RwLock, + Arc, LazyLock, RwLock, }, time::{Duration, Instant}, }; -use tokio::sync::{oneshot, Notify}; +use tokio::sync::{oneshot, Mutex, Notify}; use tracing::{debug, error, warn}; use uuid::Uuid; @@ -43,10 +47,6 @@ type ExecuteQueue = Queue; type SessionMetricsQueue = Queue; type PortalQueue = Queue>; -fn protocol_lock_error(_: std::sync::PoisonError) -> Error { - std::io::Error::other("PostgreSQL protocol state lock poisoned").into() -} - fn protocol_transition_error(error: impl std::fmt::Debug) -> Error { std::io::Error::new(std::io::ErrorKind::InvalidData, format!("{error:?}")).into() } @@ -239,19 +239,27 @@ where let notified = self.protocol_changed.notified(); tokio::pin!(notified); notified.as_mut().enable(); - let action = { - let mut protocol = self.protocol.lock().map_err(protocol_lock_error)?; - protocol + let admission = { + let mut protocol = self.protocol.lock().await; + let mut middleware = Middleware::new((), Identity); + match protocol .sides .pipeline_mut() - .frontend_action(message, handling) - .map_err(protocol_transition_error)? + .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 action { - FrontendAction::Forward { id, .. } | FrontendAction::Discard { id } => { + match admission { + Ok(FrontendAction::Forward { id, .. } | FrontendAction::Discard { id }) => { return Ok(id); } - FrontendAction::Backpressure(returned) => { + Err(returned) => { message = returned; notified.await; } @@ -259,12 +267,12 @@ where } } - pub fn protocol_in_extended_error(&self) -> Result { - let protocol = self.protocol.lock().map_err(protocol_lock_error)?; - Ok(matches!( + 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. @@ -278,11 +286,13 @@ where tokio::pin!(notified); notified.as_mut().enable(); let action = { - let mut protocol = self.protocol.lock().map_err(protocol_lock_error)?; + let mut protocol = self.protocol.lock().await; + let mut middleware = Middleware::new((), Identity); protocol .sides .pipeline_mut() - .accept_backend(message) + .accept_backend_typed(&mut middleware, message) + .await .map_err(protocol_transition_error)? }; match action { @@ -306,11 +316,13 @@ where tokio::pin!(notified); notified.as_mut().enable(); let action = { - let mut protocol = self.protocol.lock().map_err(protocol_lock_error)?; + let mut protocol = self.protocol.lock().await; + let mut middleware = Middleware::new((), Identity); protocol .sides .pipeline_mut() - .try_emit_local(id, message) + .try_emit_local_typed(&mut middleware, id, message) + .await .map_err(protocol_transition_error)? }; match action { @@ -1199,6 +1211,7 @@ mod tests { 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; @@ -1248,9 +1261,10 @@ mod tests { ) } - #[test] - fn pg_proto_ledger_tracks_pipelined_extended_messages_in_processing_order() { + #[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(), @@ -1275,9 +1289,13 @@ mod tests { protocol .sides .pipeline_mut() - .frontend_action(message, FrontendHandling::Forward) + .accept_frontend_typed(&mut middleware, message, FrontendHandling::Forward,) + .await .unwrap(), - FrontendAction::Forward { .. } + pg_proto::pipeline::FrontendAdmission::Immediate(FrontendAction::Forward { .. }) + | pg_proto::pipeline::FrontendAdmission::Waiting( + FrontendAction::Forward { .. } + ) )); } @@ -1291,7 +1309,8 @@ mod tests { protocol .sides .pipeline_mut() - .accept_backend(response) + .accept_backend_typed(&mut middleware, response) + .await .unwrap(), BackendAction::Emit(_) )); @@ -1360,7 +1379,7 @@ mod tests { .await .unwrap(); assert_eq!( - context.protocol.lock().unwrap().sides.pipeline().state(), + context.protocol.lock().await.sides.pipeline().state(), PipelineState::Ready ); } diff --git a/packages/cipherstash-proxy/src/postgresql/frontend.rs b/packages/cipherstash-proxy/src/postgresql/frontend.rs index 6268a7602..28a994e19 100644 --- a/packages/cipherstash-proxy/src/postgresql/frontend.rs +++ b/packages/cipherstash-proxy/src/postgresql/frontend.rs @@ -242,7 +242,7 @@ where return Ok(outbound_message); } - let recovering_from_extended_error = self.context.protocol_in_extended_error()?; + 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 diff --git a/packages/cipherstash-proxy/src/postgresql/handler.rs b/packages/cipherstash-proxy/src/postgresql/handler.rs index 7fbdd5f5e..603b2ac6e 100644 --- a/packages/cipherstash-proxy/src/postgresql/handler.rs +++ b/packages/cipherstash-proxy/src/postgresql/handler.rs @@ -19,7 +19,7 @@ use pg_proto::{ codec::{ Backend as BackendDirection, BackendMessage, Frontend as FrontendDirection, FrontendMessage, }, - middleware::{Identity, Middleware, ServerRole, TypedPhase, TypedReceiveError}, + middleware::{Identity, Inbound, Middleware, PhaseAssociation, ServerRole, TypedReceiveError}, net::NetworkStream, pre_startup::{PreStartup, PreStartupOffer}, server_auth::{ServerPassword, ServerProtocolOffer}, @@ -119,8 +119,8 @@ async fn receive_backend_conn( ) -> Result<(Conn, Phase>, BackendMessage), Error> where S: AsyncRead + Unpin, - Phase: TypedPhase, - >::Message: Into, + Phase: PhaseAssociation, + >::Message: Into, { let duration = Duration::from_secs(10); let message = timeout(duration, conn.receive_backend_typed(middleware)) From 27bc379c8cc46c35c14227346e96f1bbc5a479de Mon Sep 17 00:00:00 2001 From: James Sadler Date: Sat, 8 Aug 2026 12:33:39 +1000 Subject: [PATCH 15/16] refactor postgres proxy around pg-proto middleware --- .../src/postgresql/context/mod.rs | 4 +- .../src/postgresql/data/from_sql.rs | 6 +- .../error_response.rs => diagnostics.rs} | 3 +- .../src/postgresql/{handler.rs => driver.rs} | 207 +++++++++++++++--- .../src/postgresql/error_handler.rs | 4 +- .../src/postgresql/message_buffer.rs | 42 ---- .../postgresql/{ => middleware}/backend.rs | 206 +++-------------- .../postgresql/{ => middleware}/frontend.rs | 182 ++------------- .../src/postgresql/middleware/mod.rs | 5 + .../cipherstash-proxy/src/postgresql/mod.rs | 12 +- .../postgresql/{messages => rewrite}/bind.rs | 3 +- .../{messages => rewrite}/data_row.rs | 3 +- .../postgresql/{messages => rewrite}/mod.rs | 8 +- .../param_description.rs | 1 + .../postgresql/{messages => rewrite}/parse.rs | 3 +- .../postgresql/{messages => rewrite}/query.rs | 1 + .../{messages => rewrite}/row_description.rs | 3 +- .../src/postgresql/startup.rs | 69 ------ 18 files changed, 272 insertions(+), 490 deletions(-) rename packages/cipherstash-proxy/src/postgresql/{messages/error_response.rs => diagnostics.rs} (99%) rename packages/cipherstash-proxy/src/postgresql/{handler.rs => driver.rs} (74%) delete mode 100644 packages/cipherstash-proxy/src/postgresql/message_buffer.rs rename packages/cipherstash-proxy/src/postgresql/{ => middleware}/backend.rs (83%) rename packages/cipherstash-proxy/src/postgresql/{ => middleware}/frontend.rs (89%) create mode 100644 packages/cipherstash-proxy/src/postgresql/middleware/mod.rs rename packages/cipherstash-proxy/src/postgresql/{messages => rewrite}/bind.rs (99%) rename packages/cipherstash-proxy/src/postgresql/{messages => rewrite}/data_row.rs (99%) rename packages/cipherstash-proxy/src/postgresql/{messages => rewrite}/mod.rs (60%) rename packages/cipherstash-proxy/src/postgresql/{messages => rewrite}/param_description.rs (99%) rename packages/cipherstash-proxy/src/postgresql/{messages => rewrite}/parse.rs (99%) rename packages/cipherstash-proxy/src/postgresql/{messages => rewrite}/query.rs (97%) rename packages/cipherstash-proxy/src/postgresql/{messages => rewrite}/row_description.rs (98%) delete mode 100644 packages/cipherstash-proxy/src/postgresql/startup.rs diff --git a/packages/cipherstash-proxy/src/postgresql/context/mod.rs b/packages/cipherstash-proxy/src/postgresql/context/mod.rs index 6037312c8..7382f2c90 100644 --- a/packages/cipherstash-proxy/src/postgresql/context/mod.rs +++ b/packages/cipherstash-proxy/src/postgresql/context/mod.rs @@ -4,7 +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::Name, Column}; +use super::{column_mapper::ColumnMapper, rewrite::Name, Column}; use crate::{ config::TandemConfig, error::{EncryptError, Error}, @@ -1200,7 +1200,7 @@ mod tests { config::LogConfig, error::Error, log, - postgresql::{messages::Name, Column}, + postgresql::{rewrite::Name, Column}, proxy::{EncryptConfig, EncryptionService}, TandemConfig, }; 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 99% rename from packages/cipherstash-proxy/src/postgresql/messages/error_response.rs rename to packages/cipherstash-proxy/src/postgresql/diagnostics.rs index 4eacc23dd..29558a4be 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/error_response.rs +++ b/packages/cipherstash-proxy/src/postgresql/diagnostics.rs @@ -1,3 +1,4 @@ +//! CipherStash diagnostic response factories. use bytes::Bytes; use core::fmt; use pg_proto::codec::{BackendMessage, DiagnosticField, DiagnosticResponse}; @@ -453,7 +454,7 @@ 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; diff --git a/packages/cipherstash-proxy/src/postgresql/handler.rs b/packages/cipherstash-proxy/src/postgresql/driver.rs similarity index 74% rename from packages/cipherstash-proxy/src/postgresql/handler.rs rename to packages/cipherstash-proxy/src/postgresql/driver.rs index 603b2ac6e..7476dd0eb 100644 --- a/packages/cipherstash-proxy/src/postgresql/handler.rs +++ b/packages/cipherstash-proxy/src/postgresql/driver.rs @@ -1,18 +1,20 @@ -use super::backend::Backend; -use super::frontend::Frontend; +use super::middleware::{Backend, BackendDisposition, Frontend, FrontendDisposition}; use crate::connect::{self, ChannelWriter}; use crate::error::ConfigError; use crate::log::AUTHENTICATION; -use crate::postgresql::messages::error_response::ErrorResponse; -use crate::postgresql::startup; +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, + 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}, @@ -21,7 +23,7 @@ use pg_proto::{ }, middleware::{Identity, Inbound, Middleware, PhaseAssociation, ServerRole, TypedReceiveError}, net::NetworkStream, - pre_startup::{PreStartup, PreStartupOffer}, + pre_startup::{Negotiation, PreStartup, PreStartupOffer}, server_auth::{ServerPassword, ServerProtocolOffer}, startup::ProtocolVersion, transport::Buffered, @@ -29,7 +31,10 @@ use pg_proto::{ }; use postgres_protocol::authentication::sasl::{ChannelBinding, ScramSha256}; use rand::Rng; -use std::time::Duration; +use std::{ + sync::Arc, + time::{Duration, Instant}, +}; use tokio::{ io::{AsyncRead, AsyncWrite, AsyncWriteExt}, net::TcpStream, @@ -301,7 +306,7 @@ pub async fn handler(client_stream: PgStream, context: Context) -> Resu // Connect to the database server, using TLS if configured let stream = connect::connect(&context.database_socket_address()).await?; - let mut database_stream = startup::with_tls(stream, context.config()).await?; + let mut database_stream = connect_upstream_tls(stream, context.config()).await?; info!( msg = "Client connected", database = context.database_socket_address(), @@ -484,13 +489,17 @@ pub async fn handler(client_stream: PgStream, context: Context) -> Resu 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 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()), ); - let mut backend = Backend::new(channel_writer.sender(), server_reader, context.clone()); if context.is_passthrough() { if context.use_structured_logging() { @@ -506,26 +515,76 @@ pub async fn handler(client_stream: PgStream, context: Context) -> Resu 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 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? + 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?; } - 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?; + 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>(()) @@ -564,6 +623,104 @@ pub async fn handler(client_stream: PgStream, context: Context) -> Resu 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), diff --git a/packages/cipherstash-proxy/src/postgresql/error_handler.rs b/packages/cipherstash-proxy/src/postgresql/error_handler.rs index bde5f1e0d..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. @@ -62,7 +62,7 @@ pub trait PostgreSqlErrorHandler { #[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; 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/backend.rs b/packages/cipherstash-proxy/src/postgresql/middleware/backend.rs similarity index 83% rename from packages/cipherstash-proxy/src/postgresql/backend.rs rename to packages/cipherstash-proxy/src/postgresql/middleware/backend.rs index ad5214379..66ac91330 100644 --- a/packages/cipherstash-proxy/src/postgresql/backend.rs +++ b/packages/cipherstash-proxy/src/postgresql/middleware/backend.rs @@ -1,47 +1,28 @@ -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::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::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 metrics::{counter, histogram}; -use pg_proto::{ - codec::{Backend as BackendDirection, BackendMessage}, - middleware::{MessageMiddleware, Middleware}, - transport::Buffered, -}; +use pg_proto::{codec::BackendMessage, middleware::MessageMiddleware}; use std::time::Instant; -use tokio::io::AsyncRead; use tracing::{debug, error, info, warn}; -async fn receive_backend( - reader: &mut Buffered, - connection_timeout: Option, -) -> Result { - match connection_timeout { - Some(duration) => tokio::time::timeout(duration, reader.receive_backend()) - .await - .map_err(|_| Error::ConnectionTimeout { duration })? - .map_err(Into::into), - None => reader.receive_backend().await.map_err(Into::into), - } -} - /// The PostgreSQL proxy backend that handles server-to-client message processing. /// /// The Backend intercepts messages from PostgreSQL servers, identifies encrypted data @@ -86,41 +67,23 @@ async fn receive_backend( /// - `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: Buffered, /// Session context with portal and statement metadata context: Context, /// Buffer for batching DataRow messages before decryption - buffer: MessageBuffer, + buffer: Vec, } #[derive(Clone, Copy, Debug, Default, Eq, PartialEq)] -enum BackendDisposition { +pub(crate) enum BackendDisposition { #[default] Emit, Suppress, } -struct BackendInterceptor<'a, R, S> -where - R: AsyncRead + Unpin, - S: EncryptionService, -{ - backend: &'a mut Backend, -} - -impl MessageMiddleware for BackendInterceptor<'_, R, S> -where - R: AsyncRead + Unpin, - S: EncryptionService, -{ +impl MessageMiddleware for Backend { type Error = Error; async fn intercept( @@ -128,15 +91,13 @@ where disposition: &mut BackendDisposition, message: BackendMessage, ) -> Result { - self.backend.intercept_backend(disposition, message).await + self.intercept_backend(disposition, message).await } } -impl Backend -where - R: AsyncRead + Unpin, - S: EncryptionService, -{ +impl Backend { + const RESPONSE_BUFFER_SIZE: usize = 4096; + /// Creates a new Backend instance. /// /// # Arguments @@ -145,100 +106,15 @@ 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: Buffered::new(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 protocol_message = - receive_backend(&mut self.server_reader, self.context.connection_timeout()).await?; - let session_item = self.server_reader.project_backend(protocol_message.clone()); - if session_item.is_none() { - // The demux has recorded the asynchronous event and its ordering; - // forwarding still uses the original typed message below. - let _ = self.server_reader.demux_mut().pop_async_event(); - } - - let read_duration = read_start.elapsed(); - self.context.record_execute_server_timing(read_duration); - - let frame = protocol_message.to_frame()?; - let sent: u64 = (frame.body.len() + 5) 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 = ?protocol_message, - ); - } - - let (outbound_message, disposition) = { - let mut middleware = Middleware::new( - BackendDisposition::Emit, - BackendInterceptor { backend: self }, - ); - let outbound_message = middleware.intercept(protocol_message).await?; - (outbound_message, *middleware.state()) - }; - - if disposition == BackendDisposition::Emit { - self.write_with_flush(outbound_message).await?; - } - - Ok(()) - } - async fn intercept_backend( &mut self, disposition: &mut BackendDisposition, @@ -425,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?; } @@ -530,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() { @@ -776,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 } @@ -790,11 +662,7 @@ where } } -impl Backend -where - R: AsyncRead + Unpin, - S: EncryptionService, -{ +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 @@ -825,12 +693,12 @@ mod tests { use crate::config::{LogConfig, TandemConfig}; use crate::log; use crate::postgresql::context::KeysetIdentifier; + 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 pg_proto::codec::FrontendMessage; - use std::io::Cursor; use std::sync::Arc; use tokio::sync::mpsc; @@ -946,18 +814,11 @@ 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. @@ -975,9 +836,14 @@ mod tests { .await .unwrap(); - // Backend: process the terminating message via the passthrough - // path, which must drain the queues. - backend.rewrite().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(); + } if label == "ErrorResponse" { backend diff --git a/packages/cipherstash-proxy/src/postgresql/frontend.rs b/packages/cipherstash-proxy/src/postgresql/middleware/frontend.rs similarity index 89% rename from packages/cipherstash-proxy/src/postgresql/frontend.rs rename to packages/cipherstash-proxy/src/postgresql/middleware/frontend.rs index 28a994e19..196593265 100644 --- a/packages/cipherstash-proxy/src/postgresql/frontend.rs +++ b/packages/cipherstash-proxy/src/postgresql/middleware/frontend.rs @@ -1,10 +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::parse::Parse; -use super::messages::query::Query; -use super::parser::SqlParser; +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}; @@ -17,12 +17,12 @@ use crate::postgresql::context::Portal; use crate::postgresql::data::{ json_value_selector_plaintext, literal_from_sql, literal_json_value, }; -use crate::postgresql::messages::Name; +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}; @@ -31,35 +31,20 @@ use eql_mapper::{self, EqlMapperError, EqlTermVariant, JsonSelectorSource, TypeC use metrics::{counter, histogram}; use pg_proto::{ codec::{ - Backend as BackendDirection, BackendMessage, Close, Describe, DescribeTarget, Execute, - Frontend as FrontendDirection, FrontendMessage, TransactionStatus, + BackendMessage, Close, Describe, DescribeTarget, Execute, FrontendMessage, + TransactionStatus, }, - middleware::{MessageMiddleware, Middleware}, + middleware::MessageMiddleware, pipeline::{FrontendHandling, OperationId}, - transport::Buffered, }; use serde::Serialize; use sqltk::parser::ast::{self, Value}; use sqltk::NodeKey; use std::collections::HashMap; use std::sync::Arc; -use std::time::{Duration, Instant}; -use tokio::io::{AsyncRead, AsyncWrite}; +use std::time::Instant; use tracing::{debug, info, warn}; -async fn receive_frontend( - reader: &mut Buffered, - connection_timeout: Option, -) -> Result { - match connection_timeout { - Some(duration) => tokio::time::timeout(duration, reader.receive_wire()) - .await - .map_err(|_| Error::ConnectionTimeout { duration })? - .map_err(Into::into), - None => reader.receive_wire().await.map_err(Into::into), - } -} - /// The PostgreSQL proxy frontend that handles client-to-server message processing. /// /// The Frontend intercepts messages from PostgreSQL clients, analyzes SQL statements for @@ -103,45 +88,21 @@ async fn receive_frontend( /// 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: Buffered, +pub struct Frontend { /// Sender for outgoing messages to client client_sender: Sender, - /// Writer for forwarding messages to server - server_writer: Buffered, /// Session context tracking statements, portals, and keyset IDs context: Context, } #[derive(Clone, Copy, Debug, Default, Eq, PartialEq)] -enum FrontendDisposition { +pub(crate) enum FrontendDisposition { #[default] Forward, Local, } -struct FrontendInterceptor<'a, R, W, S> -where - R: AsyncRead + Unpin, - W: AsyncWrite + Unpin, - S: EncryptionService, -{ - frontend: &'a mut Frontend, -} - -impl MessageMiddleware - for FrontendInterceptor<'_, R, W, S> -where - R: AsyncRead + Unpin, - W: AsyncWrite + Unpin, - S: EncryptionService, -{ +impl MessageMiddleware for Frontend { type Error = Error; async fn intercept( @@ -149,16 +110,11 @@ where disposition: &mut FrontendDisposition, message: FrontendMessage, ) -> Result { - self.frontend.intercept_frontend(disposition, message).await + self.intercept_frontend(disposition, message).await } } -impl Frontend -where - R: AsyncRead + Unpin, - W: AsyncWrite + Unpin, - S: EncryptionService, -{ +impl Frontend { /// Creates a new Frontend instance. /// /// # Arguments @@ -167,70 +123,13 @@ 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: Buffered::new_frontend(client_reader), client_sender, - server_writer: Buffered::new(server_writer), context, } } - /// 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 protocol_message = - receive_frontend(&mut self.client_reader, self.context.connection_timeout()).await?; - let frame = protocol_message.to_frame()?; - let sent: u64 = (frame.body.len() + 5) as u64; - counter!(CLIENTS_BYTES_RECEIVED_TOTAL).increment(sent); - - let tracking_message = protocol_message.clone(); - let (outbound_message, disposition) = { - let mut middleware = Middleware::new( - FrontendDisposition::Forward, - FrontendInterceptor { frontend: self }, - ); - let outbound_message = middleware.intercept(protocol_message).await?; - (outbound_message, *middleware.state()) - }; - - if disposition == FrontendDisposition::Forward { - self.context - .protocol_frontend_received(tracking_message, FrontendHandling::Forward) - .await?; - self.write_to_server(outbound_message).await?; - } - - Ok(()) - } - async fn intercept_frontend( &mut self, disposition: &mut FrontendDisposition, @@ -387,29 +286,6 @@ where Ok(outbound_message) } - pub async fn write_to_server(&mut self, message: FrontendMessage) -> Result<(), Error> { - debug!(target: PROTOCOL, msg = "Write to server", ?message); - let frame = message.to_frame()?; - let sent: u64 = (frame.body.len() + 5) as u64; - counter!(SERVER_BYTES_SENT_TOTAL).increment(sent); - - let start = Instant::now(); - self.server_writer.push(frame)?; - self.server_writer.flush().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(()) - } - - pub async fn terminate(&mut self) -> Result<(), Error> { - debug!(target: PROTOCOL, msg = "Terminate server connection"); - self.write_to_server(FrontendMessage::Terminate).await?; - Ok(()) - } - 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); @@ -1390,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 } @@ -1405,12 +1276,7 @@ where } } -impl Frontend -where - R: AsyncRead + Unpin, - W: AsyncWrite + Unpin, - S: EncryptionService, -{ +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 = error_response.into_backend_message(); 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 7e1502d4d..12e21feb9 100644 --- a/packages/cipherstash-proxy/src/postgresql/mod.rs +++ b/packages/cipherstash-proxy/src/postgresql/mod.rs @@ -1,19 +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 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 use driver::handler; diff --git a/packages/cipherstash-proxy/src/postgresql/messages/bind.rs b/packages/cipherstash-proxy/src/postgresql/rewrite/bind.rs similarity index 99% rename from packages/cipherstash-proxy/src/postgresql/messages/bind.rs rename to packages/cipherstash-proxy/src/postgresql/rewrite/bind.rs index b83700ce8..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; @@ -495,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; diff --git a/packages/cipherstash-proxy/src/postgresql/messages/data_row.rs b/packages/cipherstash-proxy/src/postgresql/rewrite/data_row.rs similarity index 99% rename from packages/cipherstash-proxy/src/postgresql/messages/data_row.rs rename to packages/cipherstash-proxy/src/postgresql/rewrite/data_row.rs index 2dbd1ca72..54239e757 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/data_row.rs +++ b/packages/cipherstash-proxy/src/postgresql/rewrite/data_row.rs @@ -1,3 +1,4 @@ +//! CipherStash DataRow rewriting. use crate::EqlCiphertext; #[cfg(test)] use crate::{ @@ -265,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/messages/mod.rs b/packages/cipherstash-proxy/src/postgresql/rewrite/mod.rs similarity index 60% rename from packages/cipherstash-proxy/src/postgresql/messages/mod.rs rename to packages/cipherstash-proxy/src/postgresql/rewrite/mod.rs index f172afdc1..3637cdbee 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/mod.rs +++ b/packages/cipherstash-proxy/src/postgresql/rewrite/mod.rs @@ -2,25 +2,19 @@ use bytes::BytesMut; pub mod bind; pub mod data_row; -pub mod error_response; pub mod param_description; pub mod parse; pub mod query; pub mod row_description; -pub type Name = bytes::Bytes; +pub type Name = bytes::Bytes; 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; -/// Returns whether a text value may contain a JSON object. pub fn maybe_json(bytes: &BytesMut) -> bool { bytes.first() == Some(&b'{') } -/// Returns whether a binary value may contain a JSONB object. 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 99% rename from packages/cipherstash-proxy/src/postgresql/messages/param_description.rs rename to packages/cipherstash-proxy/src/postgresql/rewrite/param_description.rs index f2b270598..10e6f6df5 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/param_description.rs +++ b/packages/cipherstash-proxy/src/postgresql/rewrite/param_description.rs @@ -1,3 +1,4 @@ +//! CipherStash ParameterDescription rewriting. use crate::log::MAPPER; #[cfg(test)] use crate::{ diff --git a/packages/cipherstash-proxy/src/postgresql/messages/parse.rs b/packages/cipherstash-proxy/src/postgresql/rewrite/parse.rs similarity index 99% rename from packages/cipherstash-proxy/src/postgresql/messages/parse.rs rename to packages/cipherstash-proxy/src/postgresql/rewrite/parse.rs index 77ca64160..bb5cc451a 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/parse.rs +++ b/packages/cipherstash-proxy/src/postgresql/rewrite/parse.rs @@ -1,3 +1,4 @@ +//! CipherStash Parse rewriting. use super::{Name, UNSPECIFIED_TYPE_OID}; use crate::postgresql::context::statement::OutputParam; #[cfg(test)] @@ -170,7 +171,7 @@ mod tests { log, postgresql::{ context::statement::{OutputParam, OutputParamSource}, - messages::parse::Parse, + rewrite::parse::Parse, Column, }, Identifier, diff --git a/packages/cipherstash-proxy/src/postgresql/messages/query.rs b/packages/cipherstash-proxy/src/postgresql/rewrite/query.rs similarity index 97% rename from packages/cipherstash-proxy/src/postgresql/messages/query.rs rename to packages/cipherstash-proxy/src/postgresql/rewrite/query.rs index 2476380fc..ea311d005 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/query.rs +++ b/packages/cipherstash-proxy/src/postgresql/rewrite/query.rs @@ -1,3 +1,4 @@ +//! CipherStash simple Query rewriting. #[cfg(test)] use crate::error::{Error, ProtocolError}; #[cfg(test)] diff --git a/packages/cipherstash-proxy/src/postgresql/messages/row_description.rs b/packages/cipherstash-proxy/src/postgresql/rewrite/row_description.rs similarity index 98% rename from packages/cipherstash-proxy/src/postgresql/messages/row_description.rs rename to packages/cipherstash-proxy/src/postgresql/rewrite/row_description.rs index e02eb1818..c60db234d 100644 --- a/packages/cipherstash-proxy/src/postgresql/messages/row_description.rs +++ b/packages/cipherstash-proxy/src/postgresql/rewrite/row_description.rs @@ -1,3 +1,4 @@ +//! CipherStash RowDescription rewriting. use bytes::Bytes; #[cfg(test)] use bytes::BytesMut; @@ -157,7 +158,7 @@ impl From for BackendMessage { #[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 c3b473620..000000000 --- a/packages/cipherstash-proxy/src/postgresql/startup.rs +++ /dev/null @@ -1,69 +0,0 @@ -use pg_proto::{ - codec::Backend, - middleware::{Identity, Middleware, TypedReceiveError}, - net::NetworkStream, - pre_startup::Negotiation, - transport::Buffered, - Conn, -}; -use std::sync::Arc; -use tokio::net::TcpStream; -use tracing::warn; - -use crate::{ - error::{Error, ProtocolError}, - tls, TandemConfig, -}; - -/// Applies CipherStash's upstream TLS policy while pg-proto owns the wire -/// negotiation and its legal pre-startup transitions. -pub async fn with_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::<_, Backend>::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()) - } - } -} From 5498a23f0aa6bdbdf0ab42348f771b17ac8c1ecf Mon Sep 17 00:00:00 2001 From: James Sadler Date: Sat, 8 Aug 2026 12:34:04 +1000 Subject: [PATCH 16/16] docs: flag reusable pg-proto proxy driver --- PG_PROTO_FOLLOWUPS.md | 13 +++++++++++++ 1 file changed, 13 insertions(+) diff --git a/PG_PROTO_FOLLOWUPS.md b/PG_PROTO_FOLLOWUPS.md index cee88a508..d41fab3af 100644 --- a/PG_PROTO_FOLLOWUPS.md +++ b/PG_PROTO_FOLLOWUPS.md @@ -25,3 +25,16 @@ 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.