From 45f20f521096845b313f10774d4e2a121b652862 Mon Sep 17 00:00:00 2001 From: bernie-g Date: Fri, 25 Sep 2026 14:58:06 -0400 Subject: [PATCH 01/16] feat(pam): clickhouse native protocol Serve ClickHouse's native TCP protocol alongside the HTTP interface, so clickhouse-client, clickhouse-driver and clickhouse-go can be used with a PAM account, and so a server with HTTP disabled can be used at all. A session hands out one local port and routes each connection by its first byte, so the driver decides the protocol rather than the user. The native handler swaps the client's credentials for the account's during the handshake, then parses every client packet: the statement arrives in its own length-prefixed Query packet, which is matched against the command blocking policy and written to the session recording. Data blocks are decoded only far enough to find where they end and are forwarded byte for byte, because re-encoding them would mean reproducing a serialization we do not own. A packet the loop cannot read ends the session rather than being relayed uninspected. For a server with no HTTP interface, the gateway answers HTTP itself by running the statement over the native protocol. ClickHouse formats the values through formatRow, so every column type keeps working without the gateway decoding one. That path exists for Web Access, which is HTTP-only; a third-party HTTP client is turned away with an explanation. An account now carries an HTTP port, a native port, or both, with at least one required. The connection test probes each port that is set and names the one that failed, and a gateway that does not report native support refuses to save an account with a native port rather than failing at session time. --- go.mod | 12 +- go.sum | 23 +- packages/api/model.go | 1 + packages/gateway-v2/capabilities.go | 5 + packages/gateway-v2/discovery_handler.go | 20 +- packages/gateway-v2/gateway.go | 11 +- .../gateway-v2/test_connection_handler.go | 93 ++- .../test_connection_handler_test.go | 75 ++ packages/pam/handlers/clickhouse/bridge.go | 617 +++++++++++++++ .../pam/handlers/clickhouse/bridge_test.go | 470 ++++++++++++ .../pam/handlers/clickhouse/clients_test.go | 46 ++ .../pam/handlers/clickhouse/contract_test.go | 61 ++ .../handlers/clickhouse/edge_cases_test.go | 724 ++++++++++++++++++ packages/pam/handlers/clickhouse/native.go | 683 +++++++++++++++++ .../clickhouse/native_integration_test.go | 321 ++++++++ .../pam/handlers/clickhouse/native_outcome.go | 106 +++ .../clickhouse/native_outcome_test.go | 83 ++ .../handlers/clickhouse/native_unit_test.go | 358 +++++++++ packages/pam/handlers/clickhouse/proxy.go | 103 ++- .../pam/handlers/clickhouse/proxy_test.go | 25 + packages/pam/handlers/clickhouse/sniff.go | 53 ++ .../pam/handlers/clickhouse/sniff_test.go | 121 +++ packages/pam/handlers/clickhouse/tap_test.go | 51 ++ .../handlers/clickhouse/testhelpers_test.go | 97 +++ packages/pam/local/access.go | 8 +- packages/pam/pam-proxy.go | 20 +- packages/pam/session/credentials.go | 2 + 27 files changed, 4143 insertions(+), 46 deletions(-) create mode 100644 packages/gateway-v2/test_connection_handler_test.go create mode 100644 packages/pam/handlers/clickhouse/bridge.go create mode 100644 packages/pam/handlers/clickhouse/bridge_test.go create mode 100644 packages/pam/handlers/clickhouse/clients_test.go create mode 100644 packages/pam/handlers/clickhouse/contract_test.go create mode 100644 packages/pam/handlers/clickhouse/edge_cases_test.go create mode 100644 packages/pam/handlers/clickhouse/native.go create mode 100644 packages/pam/handlers/clickhouse/native_integration_test.go create mode 100644 packages/pam/handlers/clickhouse/native_outcome.go create mode 100644 packages/pam/handlers/clickhouse/native_outcome_test.go create mode 100644 packages/pam/handlers/clickhouse/native_unit_test.go create mode 100644 packages/pam/handlers/clickhouse/sniff.go create mode 100644 packages/pam/handlers/clickhouse/sniff_test.go create mode 100644 packages/pam/handlers/clickhouse/tap_test.go create mode 100644 packages/pam/handlers/clickhouse/testhelpers_test.go diff --git a/go.mod b/go.mod index 25291073f..2df48cc23 100644 --- a/go.mod +++ b/go.mod @@ -5,6 +5,7 @@ go 1.25.14 require ( cloud.google.com/go/iam v1.1.11 github.com/Azure/go-ntlmssp v0.1.1 + github.com/ClickHouse/ch-go v0.74.0 github.com/Masterminds/sprig/v3 v3.3.0 github.com/awnumar/memguard v0.23.0 github.com/aws/aws-sdk-go-v2 v1.27.2 @@ -114,7 +115,7 @@ require ( github.com/dgraph-io/ristretto v0.1.1 // indirect github.com/dlclark/regexp2 v1.10.0 // indirect github.com/dsnet/compress v0.0.2-0.20230904184137-39efe44ab707 // indirect - github.com/dustin/go-humanize v1.0.0 // indirect + github.com/dustin/go-humanize v1.0.1 // indirect github.com/dvsekhvalnov/jose2go v1.7.0 // indirect github.com/emicklei/go-restful/v3 v3.11.0 // indirect github.com/emirpasic/gods v1.18.1 // indirect @@ -122,6 +123,8 @@ require ( github.com/fxamacker/cbor/v2 v2.7.0 // indirect github.com/geoffgarside/ber v1.1.0 // indirect github.com/go-asn1-ber/asn1-ber v1.5.8 // indirect + github.com/go-faster/city v1.0.1 // indirect + github.com/go-faster/errors v0.7.1 // indirect github.com/go-logr/logr v1.4.3 // indirect github.com/go-logr/stdr v1.2.2 // indirect github.com/go-openapi/errors v0.20.2 // indirect @@ -166,7 +169,7 @@ require ( github.com/jcmturner/rpc/v2 v2.0.3 // indirect github.com/josharian/intern v1.0.0 // indirect github.com/json-iterator/go v1.1.12 // indirect - github.com/klauspost/compress v1.18.7 // indirect + github.com/klauspost/compress v1.19.1 // indirect github.com/klauspost/pgzip v1.2.6 // indirect github.com/lucasb-eyer/go-colorful v1.4.0 // indirect github.com/mailru/easyjson v0.7.7 // indirect @@ -194,11 +197,12 @@ require ( github.com/onsi/ginkgo/v2 v2.22.2 // indirect github.com/onsi/gomega v1.36.2 // indirect github.com/oracle/oci-go-sdk/v65 v65.95.2 // indirect - github.com/pierrec/lz4/v4 v4.1.22 // indirect + github.com/pierrec/lz4/v4 v4.1.27 // indirect github.com/pingcap/errors v0.11.5-0.20250318082626-8f80e5cb09ec // indirect github.com/pingcap/log v1.1.1-0.20241212030209-7e3ff8601a2a // indirect github.com/pingcap/tidb/pkg/parser v0.0.0-20250421232622-526b2c79173d // indirect github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 // indirect + github.com/segmentio/asm v1.2.1 // indirect github.com/shopspring/decimal v1.4.0 // indirect github.com/shurcooL/graphql v0.0.0-20240915155400-7ee5256398cf // indirect github.com/smartystreets/goconvey v1.6.4 // indirect @@ -225,7 +229,7 @@ require ( go.opentelemetry.io/otel/trace v1.44.0 // indirect go.uber.org/atomic v1.11.0 // indirect go.uber.org/multierr v1.11.0 // indirect - go.uber.org/zap v1.27.0 // indirect + go.uber.org/zap v1.28.0 // indirect go.yaml.in/yaml/v3 v3.0.5 // indirect go4.org v0.0.0-20230225012048-214862532bf5 // indirect golang.org/x/net v0.58.0 // indirect diff --git a/go.sum b/go.sum index f6e59ee65..3cadc4f69 100644 --- a/go.sum +++ b/go.sum @@ -45,6 +45,8 @@ github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03 github.com/BurntSushi/xgb v0.0.0-20160522181843-27f122750802/go.mod h1:IVnqGOEym/WlBOVXweHU+Q+/VP0lqqI8lqeDx9IjBqo= github.com/ChrisTrenkamp/goxpath v0.0.0-20210404020558-97928f7e12b6 h1:w0E0fgc1YafGEh5cROhlROMWXiNoZqApk2PDN0M1+Ns= github.com/ChrisTrenkamp/goxpath v0.0.0-20210404020558-97928f7e12b6/go.mod h1:nuWgzSkT5PnyOd+272uUmV0dnAnAn42Mk7PiQC5VzN4= +github.com/ClickHouse/ch-go v0.74.0 h1:uYs2m4wIt0ZHSM1E72rg0maCfzhR2V3xWb/vZEgpeWE= +github.com/ClickHouse/ch-go v0.74.0/go.mod h1:sZ/r+8ttZMjyrP9PuFbgoVbth1ywIu2LIQNA2vgko6M= github.com/Infisical/go-keyring v1.0.2 h1:dWOkI/pB/7RocfSJgGXbXxLDcVYsdslgjEPmVhb+nl8= github.com/Infisical/go-keyring v1.0.2/go.mod h1:LWOnn/sw9FxDW/0VY+jHFAfOFEe03xmwBVSfJnBowto= github.com/Masterminds/goutils v1.1.1 h1:5nUrii3FMTL5diU80unEVvNevw1nH4+ZV4DSLVJLSYI= @@ -168,8 +170,9 @@ github.com/dlclark/regexp2 v1.10.0/go.mod h1:DHkYz0B9wPfa6wondMfaivmHpzrQ3v9q8cn github.com/dsnet/compress v0.0.2-0.20230904184137-39efe44ab707 h1:2tV76y6Q9BB+NEBasnqvs7e49aEBFI8ejC89PSnWH+4= github.com/dsnet/compress v0.0.2-0.20230904184137-39efe44ab707/go.mod h1:qssHWj60/X5sZFNxpG4HBPDHVqxNm4DfnCKgrbZOT+s= github.com/dsnet/golib v0.0.0-20171103203638-1ea166775780/go.mod h1:Lj+Z9rebOhdfkVLjJ8T6VcRQv3SXugXy999NBtR9aFY= -github.com/dustin/go-humanize v1.0.0 h1:VSnTsYCnlFHaM2/igO1h6X3HA71jcobQuxemgkq4zYo= github.com/dustin/go-humanize v1.0.0/go.mod h1:HtrtbFcZ19U5GC7JDqmcUSB87Iq5E25KnS6fMYU6eOk= +github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= +github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= github.com/dvsekhvalnov/jose2go v1.7.0 h1:bnQc8+GMnidJZA8zc6lLEAb4xNrIqHwO+9TzqvtQZPo= github.com/dvsekhvalnov/jose2go v1.7.0/go.mod h1:QsHjhyTlD/lAVqn/NSbVZmSCGeDehTB/mPZadG+mhXU= github.com/emicklei/go-restful/v3 v3.11.0 h1:rAQeMHw1c7zTmncogyy8VvRZwtkmkZ4FxERmMY4rD+g= @@ -204,6 +207,10 @@ github.com/gitleaks/go-gitdiff v0.9.1 h1:ni6z6/3i9ODT685OLCTf+s/ERlWUNWQF4x1pvoN github.com/gitleaks/go-gitdiff v0.9.1/go.mod h1:pKz0X4YzCKZs30BL+weqBIG7mx0jl4tF1uXV9ZyNvrA= github.com/go-asn1-ber/asn1-ber v1.5.8 h1:H9AZkK22UOmfX8J84ubyaZxKJZ3FMHVwn8swoMML7iQ= github.com/go-asn1-ber/asn1-ber v1.5.8/go.mod h1:hEBeB/ic+5LoWskz+yKT7vGhhPYkProFKoKdwZRWMe0= +github.com/go-faster/city v1.0.1 h1:4WAxSZ3V2Ws4QRDrscLEDcibJY8uf41H6AhXDrNDcGw= +github.com/go-faster/city v1.0.1/go.mod h1:jKcUJId49qdW3L1qKHH/3wPeUstCVpVSXTM6vO3VcTw= +github.com/go-faster/errors v0.7.1 h1:MkJTnDoEdi9pDabt1dpWf7AA8/BaSYZqibYyhZ20AYg= +github.com/go-faster/errors v0.7.1/go.mod h1:5ySTjWFiphBs07IKuiL69nxdfd5+fzh1u7FPGZP2quo= github.com/go-gl/glfw v0.0.0-20190409004039-e6da0acd62b1/go.mod h1:vR7hzQXu2zJy9AVAgeJqvqgH9Q5CA+iKCZ2gyEVpxRU= github.com/go-gl/glfw/v3.3/glfw v0.0.0-20191125211704-12ad95a8df72/go.mod h1:tQ2UAYgL5IevRw8kRxooKSPJfGvJ9fJQFa0TUsXzTg8= github.com/go-ldap/ldap/v3 v3.4.13 h1:+x1nG9h+MZN7h/lUi5Q3UZ0fJ1GyDQYbPvbuH38baDQ= @@ -399,8 +406,8 @@ github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+o github.com/klauspost/compress v1.4.1/go.mod h1:RyIbtBH6LamlWaDj8nUwkbUhJ87Yi3uG0guNDohfE1A= github.com/klauspost/compress v1.12.3/go.mod h1:8dP1Hq4DHOhN9w426knH3Rhby4rFm6D8eO+e+Dq5Gzg= github.com/klauspost/compress v1.13.6/go.mod h1:/3/Vjq9QcHkK5uEr5lBEmyoZ1iFhe47etQ6QUkpK6sk= -github.com/klauspost/compress v1.18.7 h1:aUyZsS4kH3QTKurYhAOwAHxllVPnOthb3vPfnF1Ehjw= -github.com/klauspost/compress v1.18.7/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= +github.com/klauspost/compress v1.19.1 h1:VsB4HPswih7mmZ8WleSFQ75c/Ui1M4trX5oAsJnhSlk= +github.com/klauspost/compress v1.19.1/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= github.com/klauspost/cpuid v1.2.0/go.mod h1:Pj4uuM528wm8OyEC2QMXAi2YiTZ96dNQPGgoMS4s3ek= github.com/klauspost/pgzip v1.2.6 h1:8RXeL5crjEUFnR2/Sn6GJNWtSQ3Dk8pq4CL3jvdDyjU= github.com/klauspost/pgzip v1.2.6/go.mod h1:Ch1tH69qFZu15pkjo5kYi6mth2Zzwzt50oCQKQE9RUs= @@ -499,8 +506,8 @@ github.com/oracle/oci-go-sdk/v65 v65.95.2/go.mod h1:u6XRPsw9tPziBh76K7GrrRXPa8P8 github.com/pelletier/go-toml v1.2.0/go.mod h1:5z9KED0ma1S8pY6P1sdut58dfprrGBbd/94hg7ilaic= github.com/pelletier/go-toml/v2 v2.4.3 h1:GTRvJQutkOSftxIFD5xw9aepkYNuPWmVJpffdDPYVpY= github.com/pelletier/go-toml/v2 v2.4.3/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY= -github.com/pierrec/lz4/v4 v4.1.22 h1:cKFw6uJDK+/gfw5BcDL0JL5aBsAFdsIT18eRtLj7VIU= -github.com/pierrec/lz4/v4 v4.1.22/go.mod h1:gZWDp/Ze/IJXGXf23ltt2EXimqmTUXEy0GFuRQyBid4= +github.com/pierrec/lz4/v4 v4.1.27 h1:+PhzhWDrjRj89TH2sw43nE3+4+W8lSxIuQadEHZyjUk= +github.com/pierrec/lz4/v4 v4.1.27/go.mod h1:EoQMVJgeeEOMsCqCzqFm2O0cJvljX2nGZjcRIPL34O4= github.com/pingcap/errors v0.11.0/go.mod h1:Oi8TUi2kEtXXLMJk9l1cGmz20kV3TaQ0usTwv5KuLY8= github.com/pingcap/errors v0.11.5-0.20250318082626-8f80e5cb09ec h1:3EiGmeJWoNixU+EwllIn26x6s4njiWRXewdx2zlYa84= github.com/pingcap/errors v0.11.5-0.20250318082626-8f80e5cb09ec/go.mod h1:X2r9ueLEUZgtx2cIogM0v4Zj5uvvzhuuiu7Pn8HzMPg= @@ -538,6 +545,8 @@ github.com/russross/blackfriday v1.5.2/go.mod h1:JO/DiYxRf+HjHt06OyowR9PTA263kcR github.com/russross/blackfriday/v2 v2.0.1/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM= github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM= github.com/rwcarlsen/goexif v0.0.0-20190401172101-9e8deecbddbd/go.mod h1:hPqNNc0+uJM6H+SuU8sEs5K5IQeKccPqeSjfgcKGgPk= +github.com/segmentio/asm v1.2.1 h1:DTNbBqs57ioxAD4PrArqftgypG4/qNpXoJx8TVXxPR0= +github.com/segmentio/asm v1.2.1/go.mod h1:BqMnlJP91P8d+4ibuonYZw9mfnzI9HfxselHZr5aAcs= github.com/shopspring/decimal v1.4.0 h1:bxl37RwXBklmTi0C79JfXCEBD1cqqHt0bbgBAGFp81k= github.com/shopspring/decimal v1.4.0/go.mod h1:gawqmDU56v4yIKSwfBSFip1HdCCXN8/+DMd9qYNcwME= github.com/shurcooL/githubv4 v0.0.0-20260209031235-2402fdf4a9ed h1:KT7hI8vYXgU0s2qaMkrfq9tCA1w/iEPgfredVP+4Tzw= @@ -667,8 +676,8 @@ go.uber.org/multierr v1.7.0/go.mod h1:7EAYxJLBy9rStEaz58O2t4Uvip6FSURkq8/ppBp95a go.uber.org/multierr v1.11.0 h1:blXXJkSxSSfBVBlC76pxqeO+LN3aDfLQo+309xJstO0= go.uber.org/multierr v1.11.0/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN80Y= go.uber.org/zap v1.19.0/go.mod h1:xg/QME4nWcxGxrpdeYfq7UvYrLh66cuVKdrbD1XF/NI= -go.uber.org/zap v1.27.0 h1:aJMhYGrd5QSmlpLMr2MftRKl7t8J8PTZPA732ud/XR8= -go.uber.org/zap v1.27.0/go.mod h1:GB2qFLM7cTU87MWRP2mPIjqfIDnGu+VIO4V/SdhGo2E= +go.uber.org/zap v1.28.0 h1:IZzaP1Fv73/T/pBMLk4VutPl36uNC+OSUh3JLG3FIjo= +go.uber.org/zap v1.28.0/go.mod h1:rDLpOi171uODNm/mxFcuYWxDsqWSAVkFdX4XojSKg/Q= go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw= go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg= go4.org v0.0.0-20230225012048-214862532bf5 h1:nifaUDeh+rPaBCMPMQHZmvJf+QdpLFnuQPwx+LxVmtc= diff --git a/packages/api/model.go b/packages/api/model.go index e0b1c5fc4..c2787bbe0 100644 --- a/packages/api/model.go +++ b/packages/api/model.go @@ -1005,6 +1005,7 @@ type ChunkMetadataRequest struct { type PAMSessionCredentials struct { Host string `json:"host"` Port int `json:"port"` + NativePort int `json:"nativePort,omitempty"` Database string `json:"database"` ConnectionString string `json:"connectionString,omitempty"` // MongoDB: full URI (mongodb[+srv]://...) SSLEnabled bool `json:"sslEnabled"` diff --git a/packages/gateway-v2/capabilities.go b/packages/gateway-v2/capabilities.go index c1a798bba..51f9097c6 100644 --- a/packages/gateway-v2/capabilities.go +++ b/packages/gateway-v2/capabilities.go @@ -5,3 +5,8 @@ package gatewayv2 const CapabilitySessionLogMaskingBuiltInDetection = "sessionLogMaskingBuiltInDetection" const CapabilitySupportedAccountTypes = "supported_account_types" + +// Reported separately from the account type, because a gateway can support ClickHouse accounts and still +// predate the native protocol. Without it the platform cannot tell the difference, and an account with a +// native port would save against an old gateway and then fail every native client at session time. +const CapabilityClickhouseNativeProtocol = "clickhouseNativeProtocol" diff --git a/packages/gateway-v2/discovery_handler.go b/packages/gateway-v2/discovery_handler.go index bf04f16b0..4c56800aa 100644 --- a/packages/gateway-v2/discovery_handler.go +++ b/packages/gateway-v2/discovery_handler.go @@ -26,6 +26,23 @@ const ( type rpcTarget struct { host string port int + // Every port the signed certificate authorises. Empty means the certificate named only port, which is + // what a platform too old to send the list produces. + ports []int +} + +// allows reports whether the certificate authorises this port. A certificate that named only one port keeps +// the old behaviour, where that port is the only one a handler may reach. +func (t rpcTarget) allows(port int) bool { + if port == t.port { + return true + } + for _, allowed := range t.ports { + if port == allowed { + return true + } + } + return false } type rpcTargetContextKey struct{} @@ -63,7 +80,8 @@ func serveRPCOverTLS( opCtx, cancel := context.WithTimeout(ctx, requestDeadline) defer cancel() - opCtx = context.WithValue(opCtx, rpcTargetContextKey{}, rpcTarget{forwardConfig.TargetHost, forwardConfig.TargetPort}) + opCtx = context.WithValue(opCtx, rpcTargetContextKey{}, + rpcTarget{forwardConfig.TargetHost, forwardConfig.TargetPort, forwardConfig.TargetPorts}) req = req.WithContext(opCtx) rw := newBufferedResponseWriter() diff --git a/packages/gateway-v2/gateway.go b/packages/gateway-v2/gateway.go index 0612572a3..12f3b83b9 100644 --- a/packages/gateway-v2/gateway.go +++ b/packages/gateway-v2/gateway.go @@ -80,14 +80,19 @@ type ForwardConfig struct { VerifyTLS bool // Whether to verify TLS certificates TargetHost string TargetPort int - ActorType ActorType - PAMConfig pam.GatewayPAMConfig + // Additional ports the certificate authorises, empty when it names only TargetPort. + TargetPorts []int + ActorType ActorType + PAMConfig pam.GatewayPAMConfig } // RoutingInfo represents the routing information embedded in client certificates type RoutingInfo struct { TargetHost string `json:"targetHost"` TargetPort int `json:"targetPort"` + // Every port this certificate authorises, for the account types that reach one host on more than one. + // Absent from a certificate minted by an older platform, which means TargetPort is the only one. + TargetPorts []int `json:"targetPorts,omitempty"` } type PAMInfo struct { @@ -464,6 +469,7 @@ func (g *Gateway) registerHeartBeat(ctx context.Context, errCh chan error) { capabilities[CapabilityPkcs11] = true } capabilities[CapabilitySupportedAccountTypes] = pam.GetSupportedResourceTypes() + capabilities[CapabilityClickhouseNativeProtocol] = true req := api.GatewayHeartbeatRequest{Capabilities: capabilities} if err := api.CallGatewayHeartBeatV2(g.httpClient, req); err != nil { log.Warn().Msgf("Heartbeat failed: %v", err) @@ -1451,6 +1457,7 @@ func (g *Gateway) parseDetailsFromCertificate(tlsConn *tls.Conn, config *Forward config.TargetHost = routingInfo.TargetHost config.TargetPort = routingInfo.TargetPort + config.TargetPorts = routingInfo.TargetPorts } // Extract actor type from client certificate custom extension if ext.Id.String() == GATEWAY_ACTOR_OID { diff --git a/packages/gateway-v2/test_connection_handler.go b/packages/gateway-v2/test_connection_handler.go index 3395d401a..688e0a39d 100644 --- a/packages/gateway-v2/test_connection_handler.go +++ b/packages/gateway-v2/test_connection_handler.go @@ -124,6 +124,8 @@ type clickhouseTestParams struct { Username string `json:"username"` Password string `json:"password"` Database string `json:"database"` + HttpPort int `json:"httpPort"` + NativePort int `json:"nativePort"` SslEnabled bool `json:"sslEnabled"` SslRejectUnauthorized *bool `json:"sslRejectUnauthorized"` SslCertificate string `json:"sslCertificate"` @@ -714,17 +716,78 @@ func handleTestConnection(w http.ResponseWriter, r *http.Request) { return connectFailure(err) } } - if err := dialTarget(ctx, target.host, target.port); err != nil { - return connectFailure(err) + config := clickhousehandler.ClickHouseProxyConfig{ + Username: params.Username, + Password: params.Password, + Database: params.Database, + EnableTLS: params.SslEnabled, + TLSConfig: tlsConfig, + } + + // A server can have either interface turned off, so only the ones the account names are tested. + // Both are checked up front, so an account that would only ever fail for one kind of client is + // caught here rather than at the first session. + // One ClickHouse account can expose two ports, so the body names which to probe. The signed + // certificate still decides which are allowed, so the body cannot point the gateway at a port + // the platform did not authorise. + for _, port := range []int{params.HttpPort, params.NativePort} { + if port > 0 && !target.allows(port) { + return fmt.Errorf("port %d is not authorised for this connection test", port) + } + } + + httpPort := params.HttpPort + if httpPort <= 0 && params.NativePort <= 0 { + // An API old enough not to send the ports still means the cert-bound one. + httpPort = target.port + } + + // Each probe gets its own slice of the budget. Sharing one deadline across up to four network + // round trips means a slow first probe swallows the second one's specific error message. + probes := 0 + if httpPort > 0 { + probes++ + } + if params.NativePort > 0 { + probes++ + } + probeCtx := func() (context.Context, context.CancelFunc) { + if deadline, ok := ctx.Deadline(); ok && probes > 1 { + return context.WithTimeout(ctx, time.Until(deadline)/time.Duration(probes)) + } + return context.WithCancel(ctx) } - return authFailure(clickhousehandler.TestConnection(ctx, clickhousehandler.ClickHouseProxyConfig{ - TargetAddr: net.JoinHostPort(target.host, strconv.Itoa(target.port)), - Username: params.Username, - Password: params.Password, - Database: params.Database, - EnableTLS: params.SslEnabled, - TLSConfig: tlsConfig, - })) + + if httpPort > 0 { + httpCtx, cancel := probeCtx() + err := func() error { + defer cancel() + if err := dialTarget(httpCtx, target.host, httpPort); err != nil { + return connectFailure(err) + } + config.TargetAddr = net.JoinHostPort(target.host, strconv.Itoa(httpPort)) + if err := clickhousehandler.TestConnection(httpCtx, config); err != nil { + return authFailure(err) + } + return nil + }() + if err != nil { + return err + } + } + + if params.NativePort > 0 { + nativeCtx, cancel := probeCtx() + defer cancel() + if err := dialTarget(nativeCtx, target.host, params.NativePort); err != nil { + return connectFailure(nativePortError(params.NativePort, err, httpPort > 0)) + } + config.NativeAddr = net.JoinHostPort(target.host, strconv.Itoa(params.NativePort)) + if err := clickhousehandler.TestNativeConnection(nativeCtx, config); err != nil { + return authFailure(nativePortError(params.NativePort, err, httpPort > 0)) + } + } + return nil } case testConnModeSSH: var params sshTestParams @@ -771,3 +834,13 @@ func redactProbeSecrets(msg string, secrets ...string) string { } return urlUserinfoPattern.ReplaceAllString(msg, "${1}******@") } + +// A server can be reachable over HTTP and not over the native protocol, so the failure has to say which port +// it was and that the account can be saved without one. +func nativePortError(port int, err error, httpWorks bool) error { + if !httpWorks { + return fmt.Errorf("ClickHouse's native port %d did not answer: %w", port, err) + } + return fmt.Errorf("the HTTP interface works, but ClickHouse's native port %d did not: %w. "+ + "Clear the native port to use this account over HTTP only, which clickhouse-client cannot do", port, err) +} diff --git a/packages/gateway-v2/test_connection_handler_test.go b/packages/gateway-v2/test_connection_handler_test.go new file mode 100644 index 000000000..e2e2150a9 --- /dev/null +++ b/packages/gateway-v2/test_connection_handler_test.go @@ -0,0 +1,75 @@ +package gatewayv2 + +import ( + "encoding/json" + "testing" + + "github.com/stretchr/testify/require" +) + +// The backend builds this request in TypeScript, so nothing checks the field names match at compile time. +// These payloads were captured from buildGatewayConnectionTest itself. +func TestClickhouseTestParamsContract(t *testing.T) { + cases := []struct { + name string + payload string + wantHTTPPort int + wantNativePort int + wantSSL bool + }{ + { + name: "both interfaces", + payload: `{"mode":"clickhouse","username":"default","password":"pw","database":"analytics","httpPort":8123,"nativePort":9000,"sslEnabled":false,"sslRejectUnauthorized":true}`, + wantHTTPPort: 8123, + wantNativePort: 9000, + }, + { + name: "native only", + payload: `{"mode":"clickhouse","username":"default","password":"pw","database":"analytics","nativePort":9440,"sslEnabled":true,"sslRejectUnauthorized":true}`, + wantHTTPPort: 0, + wantNativePort: 9440, + wantSSL: true, + }, + { + name: "http only", + payload: `{"mode":"clickhouse","username":"default","password":"pw","database":"analytics","httpPort":8443,"sslEnabled":true,"sslRejectUnauthorized":false}`, + wantHTTPPort: 8443, + wantNativePort: 0, + wantSSL: true, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + var params clickhouseTestParams + require.NoError(t, json.Unmarshal([]byte(tc.payload), ¶ms)) + + require.Equal(t, tc.wantHTTPPort, params.HttpPort) + require.Equal(t, tc.wantNativePort, params.NativePort) + require.Equal(t, "default", params.Username) + require.Equal(t, "analytics", params.Database) + require.Equal(t, "pw", params.Password) + // A rename here would decode as false and quietly disable TLS for the probe. + require.Equal(t, tc.wantSSL, params.SslEnabled) + require.NotNil(t, params.SslRejectUnauthorized) + }) + } +} + +// The ports to probe come from the request body, so the signed certificate is what stops a caller pointing +// the gateway at a port the platform never authorised. +func TestRPCTargetAllows(t *testing.T) { + t.Run("a certificate naming one port authorises only that port", func(t *testing.T) { + target := rpcTarget{host: "db.internal", port: 8123} + require.True(t, target.allows(8123)) + require.False(t, target.allows(9000)) + require.False(t, target.allows(22)) + }) + + t.Run("a certificate naming several authorises each of them", func(t *testing.T) { + target := rpcTarget{host: "db.internal", port: 8123, ports: []int{8123, 9000}} + require.True(t, target.allows(8123)) + require.True(t, target.allows(9000)) + require.False(t, target.allows(22), "a port outside the certificate stays refused") + }) +} diff --git a/packages/pam/handlers/clickhouse/bridge.go b/packages/pam/handlers/clickhouse/bridge.go new file mode 100644 index 000000000..2130cb5c8 --- /dev/null +++ b/packages/pam/handlers/clickhouse/bridge.go @@ -0,0 +1,617 @@ +package clickhouse + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net/http" + "strconv" + "strings" + "time" + + "github.com/ClickHouse/ch-go" + "github.com/ClickHouse/ch-go/proto" + "github.com/rs/zerolog" +) + +// Serves ClickHouse's HTTP interface over the native protocol, for a server that has HTTP disabled. Web Access +// reaches the gateway over HTTP and @clickhouse/client cannot speak native, so without this a native-only +// account would work from the CLI and nowhere else. +// +// ClickHouse serialises the values itself through formatRow, so every column type keeps working without this +// having to decode one. +// +// Only the two envelopes Web Access asks for are produced, and that is the intended scope rather than a gap. +// A native-only account is reached by native clients; a third-party HTTP client such as JDBC, which +// negotiates a binary format, is expected to fail against it rather than be translated for. + +const ( + formatJSON = "JSON" + formatJSONCompact = "JSONCompact" + + // What formatRow is asked for, so each row arrives as the exact fragment the envelope needs. + rowFormatJSON = "JSONEachRow" + rowFormatJSONCompact = "JSONCompactEachRow" + + bridgeReadTimeout = 5 * time.Minute + maxBridgeRows = 100_000 + // A row count alone does not bound memory: one row can be hundreds of megabytes, and the envelope is + // assembled in memory. The gateway is shared, so one session must not be able to exhaust it. + maxBridgeResultBytes = 64 << 20 + + // The formatted column has to be named, because its default name is the whole call expression. + formattedRowAlias = "__infisical_row" +) + +type bridgeColumn struct { + Name string `json:"name"` + Type string `json:"type"` +} + +type bridgeStatistics struct { + Elapsed float64 `json:"elapsed"` + RowsRead uint64 `json:"rows_read"` + BytesRead uint64 `json:"bytes_read"` +} + +type bridgeEnvelope struct { + Meta []bridgeColumn `json:"meta"` + Data json.RawMessage `json:"data"` + Rows int `json:"rows"` + Statistics bridgeStatistics `json:"statistics"` +} + +func (p *ClickHouseProxy) dialNative(ctx context.Context) (*ch.Client, error) { + options := ch.Options{ + Address: p.config.NativeAddr, + Database: p.config.Database, + User: p.config.Username, + Password: p.config.Password, + ClientName: "Infisical PAM", + DialTimeout: nativeDialTimeout, + ReadTimeout: bridgeReadTimeout, + Compression: ch.CompressionDisabled, + ProtocolVersion: maxNativeRevision, + HandshakeTimeout: nativeHandshakeTimeout, + } + if p.config.EnableTLS { + options.TLS = p.config.TLSConfig + } + return ch.Dial(ctx, options) +} + +// serveBridge answers one HTTP request by running its statement over the native protocol. +func (p *ClickHouseProxy) serveBridge(w http.ResponseWriter, r *http.Request, state *requestState, l zerolog.Logger) { + // The health endpoint carries no statement, so it would otherwise be refused as an empty one. + if r.URL.Path == pingPath { + w.Header().Set("Content-Type", "text/plain; charset=UTF-8") + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte("Ok.\n")) + return + } + + // The statement was already decoded during inspection, compression and all, so it is used rather than + // the body, which the bridge cannot hand to ClickHouse the way the reverse proxy can. + if state.truncated { + message := fmt.Sprintf( + "this account reaches ClickHouse over the native protocol, so a statement larger than %d MB is "+ + "not supported here", maxInspectBytes>>20) + p.logStatement(state.statement, "ERROR: "+message) + writeClickHouseError(w, http.StatusBadRequest, codeNotImplemented, message) + return + } + + body, format := splitFormatClause(state.sql) + if format == "" { + // A client can ask for the format as a setting instead of a clause, which is what the SQL editor does. + format = r.URL.Query().Get("default_format") + } + if body == "" { + writeClickHouseError(w, http.StatusBadRequest, codeNotImplemented, "No statement was sent.") + return + } + + client, err := p.dialNative(r.Context()) + if err != nil { + l.Error().Err(err).Msg("Failed to reach ClickHouse over the native protocol") + p.logStatement(state.statement, fmt.Sprintf("ERROR: %s", err)) + writeClickHouseError(w, http.StatusBadGateway, codeNetworkError, + fmt.Sprintf("The gateway could not reach ClickHouse: %v", err)) + return + } + defer client.Close() + + envelope, err := p.runBridgeQuery(r.Context(), client, body, format, bridgeParameters(r), bridgeSettings(r)) + if err != nil { + code, status, message := classifyNativeError(err) + p.logStatement(state.statement, fmt.Sprintf("ERROR: %s", message)) + writeClickHouseError(w, status, code, message) + return + } + + envelope.Statistics.Elapsed = time.Since(state.started).Seconds() + + if format == "" { + p.logStatement(state.statement, summarizeBridge(envelope, state.started)) + w.Header().Set("Content-Type", "text/plain; charset=UTF-8") + w.WriteHeader(http.StatusOK) + return + } + + encoded, err := json.Marshal(envelope) + if err != nil { + writeClickHouseError(w, http.StatusInternalServerError, codeNotImplemented, err.Error()) + return + } + + p.logStatement(state.statement, summarizeBridge(envelope, state.started)) + + summary, _ := json.Marshal(map[string]string{ + "read_rows": strconv.FormatUint(envelope.Statistics.RowsRead, 10), + "read_bytes": strconv.FormatUint(envelope.Statistics.BytesRead, 10), + "result_rows": strconv.Itoa(envelope.Rows), + }) + + w.Header().Set("Content-Type", "application/json; charset=UTF-8") + w.Header().Set("X-ClickHouse-Summary", string(summary)) + w.Header().Set("Content-Length", strconv.Itoa(len(encoded))) + w.WriteHeader(http.StatusOK) + _, _ = w.Write(encoded) +} + +func bridgeParameters(r *http.Request) []proto.Parameter { + var parameters []proto.Parameter + for name, values := range r.URL.Query() { + if !strings.HasPrefix(name, "param_") || len(values) == 0 { + continue + } + parameters = append(parameters, proto.Parameter{ + Key: strings.TrimPrefix(name, "param_"), + Value: quoteFieldDump(values[0]), + }) + } + return parameters +} + +// The native protocol carries a query parameter as a custom setting, whose value ClickHouse reads as a Field +// dump rather than as the plain string the HTTP interface takes. An unquoted value is rejected outright. +func quoteFieldDump(value string) string { + var quoted strings.Builder + quoted.Grow(len(value) + 2) + quoted.WriteByte('\'') + // Byte-wise, because ranging over the string would turn every invalid byte into U+FFFD and silently + // change the value. The two escapes are ASCII, and UTF-8 is self-synchronising. + for i := 0; i < len(value); i++ { + if c := value[i]; c == '\\' || c == '\'' { + quoted.WriteByte('\\') + } + quoted.WriteByte(value[i]) + } + quoted.WriteByte('\'') + return quoted.String() +} + +// splitFormatClause peels off the trailing FORMAT clause @clickhouse/client appends, which decides the +// envelope rather than anything the server should see. The scan runs over the original bytes: uppercasing +// first would shift offsets, because some runes shrink when folded. +func splitFormatClause(sql string) (string, string) { + trimmed := trimTrailingSemicolons(sql) + + idx := lastIndexFold(trimmed, "FORMAT") + if idx <= 0 { + return trimmed, "" + } + // A word boundary is needed on both sides, or a trailing identifier such as `format_events` is read as + // the clause and the operand before it is thrown away. + if !isSQLSpace(trimmed[idx-1]) { + return trimmed, "" + } + after := trimmed[idx+len("FORMAT"):] + if after == "" || !isSQLSpace(after[0]) { + return trimmed, "" + } + + name := strings.TrimSpace(after) + if !isFormatName(name) { + return trimmed, "" + } + + return trimTrailingSemicolons(trimmed[:idx]), name +} + +// A format name is a bare identifier. Anything else means the word FORMAT was part of the statement. +func isFormatName(name string) bool { + if name == "" { + return false + } + for i := 0; i < len(name); i++ { + c := name[i] + if c >= 'a' && c <= 'z' || c >= 'A' && c <= 'Z' || c >= '0' && c <= '9' || c == '_' { + continue + } + return false + } + return true +} + +func isSQLSpace(c byte) bool { + return c == ' ' || c == '\t' || c == '\r' || c == '\n' +} + +// trimTrailingSemicolons removes any run of trailing semicolons and the whitespace around them, so +// `SELECT 1 ; ;` does not end up inside the subquery wrapper. +func trimTrailingSemicolons(sql string) string { + trimmed := strings.TrimSpace(sql) + for strings.HasSuffix(trimmed, ";") { + trimmed = strings.TrimSpace(strings.TrimSuffix(trimmed, ";")) + } + return trimmed +} + +// lastIndexFold is strings.LastIndex with ASCII case folding, returning an offset into s itself. +func lastIndexFold(s string, substr string) int { + for i := len(s) - len(substr); i >= 0; i-- { + if strings.EqualFold(s[i:i+len(substr)], substr) { + return i + } + } + return -1 +} + +// checkSpliceable rejects a statement that would not stay inside the parentheses it is wrapped in. Quotes +// and comments are tracked so that a semicolon or bracket inside a string literal is left alone. +func checkSpliceable(body string) error { + depth := 0 + for i := 0; i < len(body); i++ { + switch c := body[i]; c { + case '\'', '"', '`': + end := skipQuoted(body, i, c) + if end < 0 { + return refuseBridge(codeNotImplemented, "this statement has an unterminated string, so the gateway could not read it") + } + i = end + case '-': + if i+1 < len(body) && body[i+1] == '-' { + if idx := strings.IndexByte(body[i:], '\n'); idx != -1 { + i += idx + } else { + i = len(body) + } + } + case '/': + if i+1 < len(body) && body[i+1] == '*' { + idx := strings.Index(body[i+2:], "*/") + if idx == -1 { + return refuseBridge(codeNotImplemented, "this statement has an unterminated comment, so the gateway could not read it") + } + i += 2 + idx + 1 + } + case '(': + depth++ + case ')': + depth-- + if depth < 0 { + return refuseBridge(codeNotImplemented, + "this statement closes more parentheses than it opens, which the gateway cannot run over "+ + "ClickHouse's native protocol") + } + case ';': + return refuseBridge(codeNotImplemented, + "this account reaches ClickHouse over the native protocol, which runs one statement at a "+ + "time, so a semicolon inside a statement is not supported") + } + } + if depth != 0 { + return refuseBridge(codeNotImplemented, + "this statement leaves %d parenthesis open, so the gateway could not run it", depth) + } + return nil +} + +// skipQuoted returns the index of the closing quote, honouring doubled and backslash escapes. +func skipQuoted(body string, start int, quote byte) int { + for i := start + 1; i < len(body); i++ { + switch body[i] { + case '\\': + i++ + case quote: + if i+1 < len(body) && body[i+1] == quote { + i++ + continue + } + return i + } + } + return -1 +} + +// humanList renders an allowlist the way the error messages read, so the message and the list cannot drift. +func humanList(items []string) string { + switch len(items) { + case 0: + return "" + case 1: + return items[0] + default: + return strings.Join(items[:len(items)-1], ", ") + " or " + items[len(items)-1] + } +} + +func rowFormatFor(format string) (string, error) { + switch strings.ToUpper(format) { + case strings.ToUpper(formatJSON): + return rowFormatJSON, nil + case strings.ToUpper(formatJSONCompact): + return rowFormatJSONCompact, nil + default: + return "", refuseBridge(codeNotImplemented, + "this account has no HTTP port, so it is reached over ClickHouse's native protocol and only %s and "+ + "%s can be returned. This client asked for %s. Use a native client such as clickhouse-client, "+ + "or give the account an HTTP port", + formatJSON, formatJSONCompact, format) + } +} + +// Settings a client sends to bound a statement. They are forwarded so a browser session costs the server no +// more over the native protocol than it does over HTTP; anything else a client asks for is dropped. +var forwardedSettings = map[string]bool{ + "max_execution_time": true, + "max_result_rows": true, + "max_result_bytes": true, + "result_overflow_mode": true, + "max_rows_to_read": true, + "readonly": true, +} + +func bridgeSettings(r *http.Request) []ch.Setting { + var settings []ch.Setting + for name, values := range r.URL.Query() { + if !forwardedSettings[name] || len(values) == 0 { + continue + } + settings = append(settings, ch.Setting{Key: name, Value: values[0], Important: true}) + } + return settings +} + +func (p *ClickHouseProxy) runBridgeQuery( + ctx context.Context, + client *ch.Client, + body string, + format string, + parameters []proto.Parameter, + settings []ch.Setting, +) (*bridgeEnvelope, error) { + envelope := &bridgeEnvelope{Meta: []bridgeColumn{}, Data: json.RawMessage("[]")} + + // A statement with no FORMAT clause is one nothing reads the rows of, so it only has to run. + if format == "" { + var discard proto.Results + return envelope, client.Do(ctx, ch.Query{ + Body: body, + Parameters: parameters, + Settings: settings, + Result: discard.Auto(), + // ch-go refuses a second data block unless a handler is present, so a statement that returns + // more than one block would fail even though nothing here reads the rows. + OnResult: func(context.Context, proto.Block) error { return nil }, + OnProgress: func(_ context.Context, pr proto.Progress) error { + envelope.Statistics.RowsRead += pr.Rows + envelope.Statistics.BytesRead += pr.Bytes + return nil + }, + }) + } + + rowFormat, err := rowFormatFor(format) + if err != nil { + return nil, err + } + + if !isWrappable(body) { + return nil, refuseBridge(codeNotImplemented, + "this account reaches ClickHouse over the native protocol, where the gateway can only return rows "+ + "for a %s. Run this statement from the CLI instead", humanList(wrappableStatements)) + } + + meta, err := describeStatement(ctx, client, body, parameters, settings) + if err != nil { + return nil, err + } + envelope.Meta = meta + + rows, err := selectFormattedRows(ctx, client, body, rowFormat, parameters, settings, envelope) + if err != nil { + return nil, err + } + + envelope.Rows = len(rows) + envelope.Data = json.RawMessage("[" + strings.Join(rows, ",") + "]") + return envelope, nil +} + +// Only these can sit inside a subquery, which is what both halves of the bridge rely on. +var wrappableStatements = []string{"SELECT", "WITH", "EXPLAIN"} + +func isWrappable(body string) bool { + rest := strings.TrimLeft(stripLeadingNoise(body), "(") + for _, prefix := range wrappableStatements { + if len(rest) < len(prefix) || !strings.EqualFold(rest[:len(prefix)], prefix) { + continue + } + // A prefix match is not a keyword match: SELECTFOO is an identifier, not a SELECT. + if len(rest) == len(prefix) || !isIdentifierByte(rest[len(prefix)]) { + return true + } + } + return false +} + +func isIdentifierByte(c byte) bool { + return c >= 'a' && c <= 'z' || c >= 'A' && c <= 'Z' || c >= '0' && c <= '9' || c == '_' +} + +// stripLeadingNoise drops a byte-order mark, whitespace and leading comments, which editors add freely and +// which would otherwise make a perfectly ordinary SELECT look like something the bridge cannot serve. +func stripLeadingNoise(body string) string { + rest := strings.TrimPrefix(body, "\ufeff") + for { + rest = strings.TrimLeft(rest, " \t\r\n") + switch { + case strings.HasPrefix(rest, "--"): + if idx := strings.IndexByte(rest, '\n'); idx != -1 { + rest = rest[idx+1:] + continue + } + return "" + case strings.HasPrefix(rest, "/*"): + idx := strings.Index(rest[2:], "*/") + if idx == -1 { + return "" + } + rest = rest[2+idx+2:] + continue + default: + return rest + } + } +} + +// DESCRIBE resolves the statement's header without running it, so the column types are the server's own. +func describeStatement( + ctx context.Context, + client *ch.Client, + body string, + parameters []proto.Parameter, + settings []ch.Setting, +) ([]bridgeColumn, error) { + // DESCRIBE returns more columns than are needed here and the set has grown between versions, so they are + // inferred and picked by name rather than bound positionally. + var described proto.Results + columns := []bridgeColumn{} + + err := client.Do(ctx, ch.Query{ + Body: "DESCRIBE (\n" + body + "\n)", + Parameters: parameters, + Settings: settings, + Result: described.Auto(), + OnResult: func(_ context.Context, block proto.Block) error { + names, err := stringColumn(described, "name") + if err != nil { + return err + } + types, err := stringColumn(described, "type") + if err != nil { + return err + } + for i := 0; i < block.Rows; i++ { + columns = append(columns, bridgeColumn{Name: names.Row(i), Type: types.Row(i)}) + } + return nil + }, + }) + if err != nil { + return nil, err + } + return columns, nil +} + +func stringColumn(results proto.Results, name string) (*proto.ColStr, error) { + for _, column := range results { + if column.Name != name { + continue + } + if typed, ok := column.Data.(*proto.ColStr); ok { + return typed, nil + } + return nil, fmt.Errorf("DESCRIBE returned %q as %s rather than String", name, column.Data.Type()) + } + return nil, fmt.Errorf("DESCRIBE returned no %q column", name) +} + +func selectFormattedRows( + ctx context.Context, + client *ch.Client, + body string, + rowFormat string, + parameters []proto.Parameter, + settings []ch.Setting, + envelope *bridgeEnvelope, +) ([]string, error) { + var formatted proto.ColStr + rows := make([]string, 0, 64) + resultBytes := 0 + + err := client.Do(ctx, ch.Query{ + Body: "SELECT formatRowNoNewline('" + rowFormat + "', *) AS " + formattedRowAlias + + " FROM (\n" + body + "\n)", + Parameters: parameters, + Settings: settings, + Result: proto.Results{{Name: formattedRowAlias, Data: &formatted}}, + OnResult: func(_ context.Context, block proto.Block) error { + for i := 0; i < block.Rows; i++ { + if len(rows) >= maxBridgeRows { + return refuseBridge(codeTooManyRows, + "this statement returned more than %d rows, which is more than a browser session "+ + "returns over the native protocol. Add a LIMIT, or use the CLI", maxBridgeRows) + } + row := formatted.Row(i) + resultBytes += len(row) + 1 + if resultBytes > maxBridgeResultBytes { + return refuseBridge(codeTooManyRows, + "this statement returned more than %d MB, which is more than a browser session "+ + "returns over the native protocol. Narrow the result, or use the CLI", + maxBridgeResultBytes>>20) + } + rows = append(rows, row) + } + return nil + }, + OnProgress: func(_ context.Context, pr proto.Progress) error { + envelope.Statistics.RowsRead += pr.Rows + envelope.Statistics.BytesRead += pr.Bytes + return nil + }, + }) + if err != nil { + return nil, err + } + return rows, nil +} + +// bridgeRefusal is something the gateway decided about the request itself, so it carries its own ClickHouse +// code and is reported as a bad request rather than as a network error wrapped in ch-go's decoding context. +type bridgeRefusal struct { + code int + message string +} + +func (e *bridgeRefusal) Error() string { return e.message } + +func refuseBridge(code int, format string, args ...any) error { + return &bridgeRefusal{code: code, message: fmt.Sprintf(format, args...)} +} + +// classifyNativeError turns a ch-go error back into the code, message and status a ClickHouse client expects. +func classifyNativeError(err error) (int, int, string) { + var refusal *bridgeRefusal + if errors.As(err, &refusal) { + return refusal.code, http.StatusBadRequest, refusal.message + } + if exception, ok := ch.AsException(err); ok { + return int(exception.Code), http.StatusBadRequest, exception.Message + } + return codeNetworkError, http.StatusBadGateway, err.Error() +} + +func summarizeBridge(envelope *bridgeEnvelope, started time.Time) string { + parts := []string{"200 OK"} + if envelope.Rows > 0 { + parts = append(parts, fmt.Sprintf("%d row(s) returned", envelope.Rows)) + } + if envelope.Statistics.RowsRead > 0 { + parts = append(parts, fmt.Sprintf("%d row(s) read", envelope.Statistics.RowsRead)) + } + return strings.Join(append(parts, fmt.Sprintf("%dms", time.Since(started).Milliseconds())), ", ") +} diff --git a/packages/pam/handlers/clickhouse/bridge_test.go b/packages/pam/handlers/clickhouse/bridge_test.go new file mode 100644 index 000000000..6128ef158 --- /dev/null +++ b/packages/pam/handlers/clickhouse/bridge_test.go @@ -0,0 +1,470 @@ +package clickhouse + +import ( + "context" + "encoding/json" + "fmt" + "io" + "net" + "net/http" + "os" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestSplitFormatClause(t *testing.T) { + cases := []struct { + name string + sql string + wantBody string + wantFormat string + }{ + { + name: "the clause the node client appends", + sql: "SELECT 1 \nFORMAT JSON", + wantBody: "SELECT 1", + wantFormat: "JSON", + }, + {name: "trailing semicolon", sql: "SELECT 1 FORMAT JSONCompact;", wantBody: "SELECT 1", wantFormat: "JSONCompact"}, + {name: "no clause", sql: "SELECT 1", wantBody: "SELECT 1", wantFormat: ""}, + { + // "format" inside the statement is not the clause, and peeling it off would corrupt the query. + name: "a call to formatDateTime is not a clause", + sql: "SELECT formatDateTime(now(), '%F')", + wantBody: "SELECT formatDateTime(now(), '%F')", + }, + {name: "a column named format", sql: "SELECT format FROM t", wantBody: "SELECT format FROM t"}, + {name: "format with no name", sql: "SELECT 1 FORMAT", wantBody: "SELECT 1 FORMAT"}, + // A trailing identifier that merely starts with "format" is not a clause, and splitting it would + // throw away the operand in front of it. + {name: "a table whose name starts with format", sql: "SELECT * FROM format_events", wantBody: "SELECT * FROM format_events"}, + {name: "an alias that starts with format", sql: "SELECT 1 AS format_id", wantBody: "SELECT 1 AS format_id"}, + {name: "ordering by a column called formatted", sql: "SELECT x FROM t ORDER BY formatted", wantBody: "SELECT x FROM t ORDER BY formatted"}, + {name: "lowercase clause", sql: "select 1 format json", wantBody: "select 1", wantFormat: "json"}, + {name: "mixed case clause", sql: "SELECT 1 FoRmAt JSONCompact", wantBody: "SELECT 1", wantFormat: "JSONCompact"}, + {name: "format alone is not a clause", sql: "FORMAT JSON", wantBody: "FORMAT JSON"}, + {name: "format inside a string literal", sql: "SELECT 'FORMAT JSON'", wantBody: "SELECT 'FORMAT JSON'"}, + {name: "several trailing semicolons", sql: "SELECT 1 ; ;", wantBody: "SELECT 1"}, + // A rune that shrinks when uppercased would shift byte offsets if the scan ran over a folded copy. + {name: "a value whose uppercase form is shorter", sql: "SELECT 'ı' AS x FORMAT JSON", wantBody: "SELECT 'ı' AS x", wantFormat: "JSON"}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + body, format := splitFormatClause(tc.sql) + require.Equal(t, tc.wantBody, body) + require.Equal(t, tc.wantFormat, format) + }) + } +} + +func TestIsWrappable(t *testing.T) { + for _, body := range []string{"SELECT 1", " select 1", "WITH x AS (SELECT 1) SELECT * FROM x", "EXPLAIN SELECT 1"} { + require.True(t, isWrappable(body), body) + } + for _, body := range []string{"SHOW TABLES", "DESCRIBE TABLE t", "INSERT INTO t VALUES (1)", "CREATE TABLE t (a UInt8) ENGINE = Memory"} { + require.False(t, isWrappable(body), body) + } + + // Editors put comments and byte-order marks in front of perfectly ordinary statements. + for _, body := range []string{"-- a note\nSELECT 1", "/* a note */ SELECT 1", "\ufeffSELECT 1", "/* a */ -- b\n select 1"} { + require.True(t, isWrappable(body), body) + } + // A prefix match is not a keyword match. + require.False(t, isWrappable("SELECTFOO 1")) + require.False(t, isWrappable("WITHOUT_ROWS()")) +} + +func TestCheckSpliceable(t *testing.T) { + ok := []string{ + "SELECT 1", + "SELECT (1 + 2) AS x", + "SELECT 'a;b) --' AS s", + "SELECT \"col;)\" FROM t", + "SELECT 1 -- a trailing ; comment", + "SELECT 1 /* a ) comment */", + } + for _, body := range ok { + require.NoError(t, checkSpliceable(body), body) + } + + // Each of these would otherwise run something other than what was recorded and policy-checked. + bad := []string{ + "SELECT 1) ; DROP TABLE users; --", + "SELECT 1) UNION ALL (SELECT 2", + "SELECT 1; SELECT 2", + "SELECT (1", + "SELECT 'unterminated", + } + for _, body := range bad { + require.Error(t, checkSpliceable(body), body) + } +} + +// The bridge has to answer what ClickHouse's own HTTP interface answers, so both are asked the same thing +// and the envelopes compared. +func TestBridgeMatchesHTTPInterface(t *testing.T) { + if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { + t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") + } + + statements := []string{ + "SELECT 1 AS n, 'x' AS s", + "SELECT map('a', 1::UInt64) AS m, tuple('p', 2) AS t", + "SELECT number, toString(number) AS s FROM numbers(5)", + "SELECT id, note FROM pam_write_test ORDER BY id LIMIT 3", + "SELECT * FROM exotic ORDER BY id", + "SELECT count() AS c FROM users", + "SELECT NULL::Nullable(String) AS nothing", + "WITH 2 AS x SELECT x * 3 AS y", + } + + for _, format := range []string{formatJSON, formatJSONCompact} { + for _, statement := range statements { + t.Run(format+": "+statement, func(t *testing.T) { + viaBridge := queryBridge(t, statement, format) + viaHTTP := queryRealHTTP(t, statement, format) + + require.Equal(t, viaHTTP.Meta, viaBridge.Meta, "column metadata should match ClickHouse") + require.Equal(t, viaHTTP.Rows, viaBridge.Rows, "row count should match ClickHouse") + require.JSONEq(t, string(viaHTTP.Data), string(viaBridge.Data), "rows should match ClickHouse") + }) + } + } +} + +func TestBridgeReportsClickHouseErrors(t *testing.T) { + if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { + t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") + } + + addr := startBridgeProxy(t, nil, &recordingLogger{}) + status, body := postStatement(t, addr, "SELECT * FROM table_that_is_not_there FORMAT JSON") + + require.Equal(t, http.StatusBadRequest, status) + require.Contains(t, body, "table_that_is_not_there") + require.Contains(t, body, "DB::Exception") +} + +func TestBridgeAppliesCommandBlocking(t *testing.T) { + if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { + t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") + } + + recorder := &recordingLogger{} + addr := startBridgeProxy(t, []string{`(?i)\bdrop\b`}, recorder) + + status, body := postStatement(t, addr, "DROP TABLE pam_write_test FORMAT JSON") + require.Equal(t, http.StatusForbidden, status) + require.Contains(t, body, "blocked by the command blocking policy") + require.True(t, recorder.contains("DROP TABLE pam_write_test")) +} + +func TestBridgeRecordsStatements(t *testing.T) { + if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { + t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") + } + + recorder := &recordingLogger{} + addr := startBridgeProxy(t, nil, recorder) + + status, _ := postStatement(t, addr, "SELECT 42 AS answer FORMAT JSON") + require.Equal(t, http.StatusOK, status) + require.True(t, recorder.contains("SELECT 42 AS answer"), recorder.dump()) + require.Contains(t, recorder.dump(), "1 row(s) returned") +} + +// A native-only account has no HTTP upstream, so TargetAddr is deliberately empty. +func startBridgeProxy(t *testing.T, blocked []string, logger *recordingLogger) string { + t.Helper() + + patterns := compileForTest(t, blocked) + proxy := NewClickHouseProxy(ClickHouseProxyConfig{ + NativeAddr: envOr("PAM_CLICKHOUSE_NATIVE", "127.0.0.1:9000"), + Username: envOr("PAM_CLICKHOUSE_USER", "default"), + Password: envOr("PAM_CLICKHOUSE_PASSWORD", "clickhouse"), + Database: envOr("PAM_CLICKHOUSE_DB", "analytics"), + SessionID: "bridge-test", + SessionLogger: logger, + BlockedCommands: patterns, + }) + + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + t.Cleanup(func() { listener.Close() }) + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + go func() { + for { + conn, err := listener.Accept() + if err != nil { + return + } + go func() { _ = proxy.HandleConnection(ctx, conn) }() + } + }() + + return listener.Addr().String() +} + +func postStatement(t *testing.T, addr string, sql string) (int, string) { + t.Helper() + + req, err := http.NewRequest(http.MethodPost, "http://"+addr+"/", strings.NewReader(sql)) + require.NoError(t, err) + + resp, err := (&http.Client{Timeout: 60 * time.Second}).Do(req) + require.NoError(t, err) + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + return resp.StatusCode, string(body) +} + +func queryBridge(t *testing.T, statement string, format string) bridgeEnvelope { + t.Helper() + + addr := startBridgeProxy(t, nil, &recordingLogger{}) + status, body := postStatement(t, addr, statement+" \nFORMAT "+format) + require.Equal(t, http.StatusOK, status, body) + + var envelope bridgeEnvelope + require.NoError(t, json.Unmarshal([]byte(body), &envelope), body) + return envelope +} + +func queryRealHTTP(t *testing.T, statement string, format string) bridgeEnvelope { + t.Helper() + + target := fmt.Sprintf("http://%s/?database=%s", + envOr("PAM_CLICKHOUSE_HTTP", "127.0.0.1:8123"), envOr("PAM_CLICKHOUSE_DB", "analytics")) + + req, err := http.NewRequest(http.MethodPost, target, strings.NewReader(statement+" \nFORMAT "+format)) + require.NoError(t, err) + req.Header.Set("X-ClickHouse-User", envOr("PAM_CLICKHOUSE_USER", "default")) + req.Header.Set("X-ClickHouse-Key", envOr("PAM_CLICKHOUSE_PASSWORD", "clickhouse")) + + resp, err := (&http.Client{Timeout: 60 * time.Second}).Do(req) + require.NoError(t, err) + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode, string(body)) + + var envelope bridgeEnvelope + require.NoError(t, json.Unmarshal(body, &envelope), string(body)) + return envelope +} + +// The bridge exists for a server with HTTP genuinely turned off, so it is also exercised against one. +func TestBridgeAgainstHTTPDisabledServer(t *testing.T) { + if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { + t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") + } + + native := envOr("PAM_CLICKHOUSE_NOHTTP_NATIVE", "127.0.0.1:19001") + probe, err := net.DialTimeout("tcp", native, 2*time.Second) + if err != nil { + t.Skipf("no HTTP-disabled ClickHouse on %s: %v", native, err) + } + probe.Close() + + bridged := NewClickHouseProxy(ClickHouseProxyConfig{ + NativeAddr: native, + Username: "default", + Password: envOr("PAM_CLICKHOUSE_NOHTTP_PASSWORD", "clickhouse"), + Database: "default", + SessionID: "bridge-nohttp-test", + SessionLogger: &recordingLogger{}, + }) + + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + defer listener.Close() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go func() { + for { + conn, acceptErr := listener.Accept() + if acceptErr != nil { + return + } + go func() { _ = bridged.HandleConnection(ctx, conn) }() + } + }() + + status, body := postStatement(t, listener.Addr().String(), + "SELECT 1 AS n, map('k', 'v') AS m \nFORMAT JSON") + require.Equal(t, http.StatusOK, status, body) + + var envelope bridgeEnvelope + require.NoError(t, json.Unmarshal([]byte(body), &envelope), body) + require.Equal(t, 1, envelope.Rows) + require.Equal(t, []bridgeColumn{{Name: "n", Type: "UInt8"}, {Name: "m", Type: "Map(String, String)"}}, envelope.Meta) + require.JSONEq(t, `[{"n":1,"m":{"k":"v"}}]`, string(envelope.Data)) +} + +// The SQL editor asks for its format with the default_format setting rather than a FORMAT clause, which is +// a different code path and was returning nothing at all. +func TestBridgeHonoursDefaultFormatSetting(t *testing.T) { + if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { + t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") + } + + addr := startBridgeProxy(t, nil, &recordingLogger{}) + + for _, format := range []string{formatJSON, formatJSONCompact} { + t.Run(format, func(t *testing.T) { + status, body := postGET(t, addr, + "/?default_format="+format+"&max_execution_time=30&query="+ + urlEscape("SELECT id, tags FROM exotic ORDER BY id")) + require.Equal(t, http.StatusOK, status, body) + + var envelope bridgeEnvelope + require.NoError(t, json.Unmarshal([]byte(body), &envelope), body) + require.Equal(t, 1, envelope.Rows) + require.Equal(t, []bridgeColumn{ + {Name: "id", Type: "UInt64"}, + {Name: "tags", Type: "Map(String, UInt64)"}, + }, envelope.Meta) + }) + } +} + +// A statement whose last line is a comment would otherwise swallow the wrapper's closing parenthesis. +func TestBridgeHandlesATrailingComment(t *testing.T) { + if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { + t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") + } + + addr := startBridgeProxy(t, nil, &recordingLogger{}) + status, body := postStatement(t, addr, "SELECT 1 AS n\n-- a trailing note\nFORMAT JSON") + require.Equal(t, http.StatusOK, status, body) + require.Contains(t, body, `"n":1`) +} + +// A result bigger than one block used to fail because ch-go refuses a second block with no handler. +func TestBridgeHandlesAMultiBlockResult(t *testing.T) { + if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { + t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") + } + + addr := startBridgeProxy(t, nil, &recordingLogger{}) + + // Comfortably more than one block (max_block_size defaults to ~65k) and under the row cap. + const rows = 80000 + + t.Run("with a format", func(t *testing.T) { + status, body := postStatement(t, addr, fmt.Sprintf("SELECT number FROM numbers(%d) \nFORMAT JSONCompact", rows)) + require.Equal(t, http.StatusOK, status, body[:min(len(body), 400)]) + + var envelope bridgeEnvelope + require.NoError(t, json.Unmarshal([]byte(body), &envelope)) + require.Equal(t, rows, envelope.Rows) + }) + + t.Run("without a format", func(t *testing.T) { + status, body := postStatement(t, addr, fmt.Sprintf("SELECT number FROM numbers(%d)", rows)) + require.Equal(t, http.StatusOK, status, body) + }) +} + +func TestBridgeRefusesAResultBeyondTheRowCap(t *testing.T) { + if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { + t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") + } + + addr := startBridgeProxy(t, nil, &recordingLogger{}) + status, body := postStatement(t, addr, + fmt.Sprintf("SELECT number FROM numbers(%d) \nFORMAT JSONCompact", maxBridgeRows+1000)) + + require.Equal(t, http.StatusBadRequest, status) + require.Contains(t, body, "more than 100000 rows") + // The refusal is the gateway's own, so it must not be dressed up as a network error. + require.Contains(t, body, "TOO_MANY_ROWS") + require.NotContains(t, body, "decode block", "ch-go's internal wrapping should not reach the client") +} + +// /ping is a health check, and on a native-only account it used to be refused as an empty statement. +func TestBridgeAnswersPing(t *testing.T) { + if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { + t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") + } + + addr := startBridgeProxy(t, nil, &recordingLogger{}) + status, body := postGET(t, addr, "/ping") + require.Equal(t, http.StatusOK, status) + require.Contains(t, body, "Ok.") +} + +// A parameter is data. Values that are awkward to quote must survive unchanged. +func TestBridgeParameterRoundTrip(t *testing.T) { + if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { + t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") + } + + addr := startBridgeProxy(t, nil, &recordingLogger{}) + + for _, value := range []string{"plain", "it's", `back\slash`, "a,b", " spaced ", "ünïcødé", "0", ""} { + t.Run(fmt.Sprintf("%q", value), func(t *testing.T) { + status, body := postGET(t, addr, + "/?param_v="+urlEscape(value)+"&query="+urlEscape("SELECT {v:String} AS got FORMAT JSON")) + require.Equal(t, http.StatusOK, status, body) + + var envelope bridgeEnvelope + require.NoError(t, json.Unmarshal([]byte(body), &envelope), body) + + var rows []struct { + Got string `json:"got"` + } + require.NoError(t, json.Unmarshal(envelope.Data, &rows)) + require.Len(t, rows, 1) + require.Equal(t, value, rows[0].Got) + }) + } +} + +// @clickhouse/client.insert() sends `INSERT INTO t FORMAT JSONEachRow` with the rows in the body, which goes +// down the no-format path as one native query carrying inline data. It must complete rather than sit until +// the server's receive timeout, and the rows have to actually land. +func TestBridgeInsertWithInlineData(t *testing.T) { + if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { + t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") + } + + addr := startBridgeProxy(t, nil, &recordingLogger{}) + marker := fmt.Sprintf("bridge-insert-%d", time.Now().UnixNano()) + + type result struct { + status int + body string + } + done := make(chan result, 1) + go func() { + status, body, err := postStatementE(addr, + fmt.Sprintf("INSERT INTO pam_write_test (id, note) FORMAT JSONEachRow\n{\"id\":9100,\"note\":%q}", marker)) + if err != nil { + done <- result{status: -1, body: err.Error()} + return + } + done <- result{status: status, body: body} + }() + + select { + case got := <-done: + require.Equal(t, http.StatusOK, got.status, got.body) + case <-time.After(45 * time.Second): + t.Fatal("the bridge hung on an INSERT carrying inline data") + } + + // An empty 200 that quietly inserted nothing would be worse than hanging. + require.Equal(t, "1", queryDirect(t, fmt.Sprintf("SELECT count() FROM pam_write_test WHERE note = '%s'", marker))) +} diff --git a/packages/pam/handlers/clickhouse/clients_test.go b/packages/pam/handlers/clickhouse/clients_test.go new file mode 100644 index 000000000..ee74c3e60 --- /dev/null +++ b/packages/pam/handlers/clickhouse/clients_test.go @@ -0,0 +1,46 @@ +package clickhouse + +import ( + "context" + "os" + "os/exec" + "strings" + "testing" + "time" +) + +// clickhouse-client is one native implementation. Python's clickhouse-driver is an independent one, so it +// catches assumptions that happen to match ClickHouse's own client. +func TestPythonClickHouseDriver(t *testing.T) { + if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { + t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") + } + + port := startProxy(t, baseConfig(&recordingLogger{})) + + script := ` +import clickhouse_driver, sys +c = clickhouse_driver.Client(host=sys.argv[1], port=int(sys.argv[2]), user='wrong', password='wrong', database='ignored') +print("currentUser:", c.execute("SELECT currentUser()")[0][0]) +print("count:", c.execute("SELECT count() FROM users")[0][0]) +print("exotic:", c.execute("SELECT map('a', 1::UInt64), tuple('p', 2)")[0]) +print("multi:", c.execute("SELECT 1")[0][0], c.execute("SELECT 2")[0][0]) +` + ctx, cancel := context.WithTimeout(context.Background(), 180*time.Second) + defer cancel() + + cmd := exec.CommandContext(ctx, "docker", "run", "--rm", "-i", "python:3.12-slim", "bash", "-lc", + "pip install --quiet clickhouse-driver >/dev/null 2>&1 && python -c \""+strings.ReplaceAll(script, `"`, `\"`)+"\" "+ + envOr("PAM_CLICKHOUSE_CLIENT_HOST", "host.docker.internal")+" "+port) + + out, err := cmd.CombinedOutput() + t.Logf("%s", out) + if err != nil { + t.Fatalf("clickhouse-driver failed: %v", err) + } + for _, want := range []string{"currentUser: default", "count: 250", "multi: 1 2"} { + if !strings.Contains(string(out), want) { + t.Fatalf("expected %q in output", want) + } + } +} diff --git a/packages/pam/handlers/clickhouse/contract_test.go b/packages/pam/handlers/clickhouse/contract_test.go new file mode 100644 index 000000000..f49c4b6de --- /dev/null +++ b/packages/pam/handlers/clickhouse/contract_test.go @@ -0,0 +1,61 @@ +package clickhouse + +import ( + "encoding/json" + "testing" + + "github.com/Infisical/infisical-merge/packages/api" + "github.com/stretchr/testify/require" +) + +// The API and the gateway agree on these shapes only by convention, and a renamed field would not fail to +// compile in either repo. The payloads below were captured from the backend's own builders. +func TestSessionCredentialsContract(t *testing.T) { + cases := []struct { + name string + payload string + wantPort int + wantNativePort int + wantSSL bool + wantRejectUnauthorized bool + }{ + { + name: "both interfaces", + payload: `{"host":"ch.example.com","port":8123,"nativePort":9000,"database":"analytics","sslEnabled":false,"sslRejectUnauthorized":true,"username":"default","password":"pw"}`, + wantPort: 8123, + wantNativePort: 9000, + wantRejectUnauthorized: true, + }, + { + name: "native only, so no HTTP port is sent at all", + payload: `{"host":"ch.example.com","nativePort":9440,"database":"analytics","sslEnabled":true,"sslRejectUnauthorized":true,"username":"default","password":"pw"}`, + wantPort: 0, + wantNativePort: 9440, + wantSSL: true, + wantRejectUnauthorized: true, + }, + { + name: "http only, so no native port is sent at all", + payload: `{"host":"ch.example.com","port":8443,"database":"analytics","sslEnabled":true,"sslRejectUnauthorized":false,"username":"default","password":"pw"}`, + wantPort: 8443, + wantNativePort: 0, + wantSSL: true, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + var credentials api.PAMSessionCredentials + require.NoError(t, json.Unmarshal([]byte(tc.payload), &credentials)) + + require.Equal(t, tc.wantPort, credentials.Port) + require.Equal(t, tc.wantNativePort, credentials.NativePort) + require.Equal(t, "ch.example.com", credentials.Host) + require.Equal(t, "default", credentials.Username) + require.Equal(t, "pw", credentials.Password) + // A rename here would decode as false, silently dropping TLS and sending the password in clear. + require.Equal(t, tc.wantSSL, credentials.SSLEnabled) + require.Equal(t, tc.wantRejectUnauthorized, credentials.SSLRejectUnauthorized) + }) + } +} diff --git a/packages/pam/handlers/clickhouse/edge_cases_test.go b/packages/pam/handlers/clickhouse/edge_cases_test.go new file mode 100644 index 000000000..733ebe3a9 --- /dev/null +++ b/packages/pam/handlers/clickhouse/edge_cases_test.go @@ -0,0 +1,724 @@ +package clickhouse + +import ( + "context" + "crypto/tls" + "crypto/x509" + "encoding/json" + "fmt" + "io" + "net" + "net/http" + "os" + "regexp" + "strconv" + "strings" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func itOnly(t *testing.T) { + t.Helper() + if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { + t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") + } +} + +// startProxy serves one local port for whatever config it is given, which is what a session does. +func startProxy(t *testing.T, config ClickHouseProxyConfig) string { + t.Helper() + + proxy := NewClickHouseProxy(config) + + listener, err := net.Listen("tcp", "0.0.0.0:0") + require.NoError(t, err) + t.Cleanup(func() { listener.Close() }) + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + go func() { + for { + conn, acceptErr := listener.Accept() + if acceptErr != nil { + return + } + go func() { _ = proxy.HandleConnection(ctx, conn) }() + } + }() + + return fmt.Sprintf("%d", listener.Addr().(*net.TCPAddr).Port) +} + +func baseConfig(logger *recordingLogger, blocked ...string) ClickHouseProxyConfig { + patterns := make([]*regexp.Regexp, 0, len(blocked)) + for _, p := range blocked { + patterns = append(patterns, regexp.MustCompile(p)) + } + return ClickHouseProxyConfig{ + TargetAddr: envOr("PAM_CLICKHOUSE_HTTP", "127.0.0.1:8123"), + NativeAddr: envOr("PAM_CLICKHOUSE_NATIVE", "127.0.0.1:9000"), + Username: envOr("PAM_CLICKHOUSE_USER", "default"), + Password: envOr("PAM_CLICKHOUSE_PASSWORD", "clickhouse"), + Database: envOr("PAM_CLICKHOUSE_DB", "analytics"), + SessionID: "edge-test", + SessionLogger: logger, + BlockedCommands: patterns, + } +} + +// ---------- compression ---------- + +// clickhouse-client compresses data blocks by default, which is a different decode path from an +// uncompressed session. Both have to keep inspecting every statement. +func TestNativeCompressionBothWays(t *testing.T) { + itOnly(t) + + for _, compression := range []string{"1", "0"} { + t.Run("compression="+compression, func(t *testing.T) { + recorder := &recordingLogger{} + port := startProxy(t, baseConfig(recorder, `(?i)\bdrop\b`)) + + out, err := runClient(t, port, "SELECT 1;\nSELECT 'second';\nDROP TABLE pam_write_test;", + "--compression", compression) + t.Logf("%s", out) + + require.Error(t, err, "the blocked statement should fail the client") + require.Contains(t, out, "second", "earlier statements should still run") + require.Contains(t, out, "blocked by the command blocking policy") + require.True(t, recorder.contains("SELECT 'second'"), recorder.dump()) + }) + } +} + +// An INSERT pushes binary rows from the client, so it exercises block decoding rather than framing alone. +func TestNativeInsertWithCompressionBothWays(t *testing.T) { + itOnly(t) + + for _, compression := range []string{"1", "0"} { + t.Run("compression="+compression, func(t *testing.T) { + port := startProxy(t, baseConfig(&recordingLogger{})) + marker := fmt.Sprintf("edge-%s-%d", compression, time.Now().UnixNano()) + + out, err := runClient(t, port, + fmt.Sprintf("INSERT INTO pam_write_test (id, note) VALUES (7001, '%s');\nSELECT note FROM pam_write_test WHERE note = '%s';", marker, marker), + "--compression", compression) + require.NoError(t, err, out) + require.Contains(t, out, marker) + }) + } +} + +// ---------- the documented fail-closed path ---------- + +// A column type ch-go cannot infer can only appear in the client direction on an INSERT. That has to fail +// closed with a message naming the type, never relay uninspected. +func TestNativeInsertIntoUnreadableColumnFailsClosed(t *testing.T) { + itOnly(t) + + recorder := &recordingLogger{} + port := startProxy(t, baseConfig(recorder)) + + out, _ := runClient(t, port, + "INSERT INTO exotic (id, tags, pair, arr, lc, dec, ts, en) VALUES (99, {'x':1}, ('p',2), ['a'], 'low', 1.0, '2026-01-01 00:00:00.000', 'a');") + t.Logf("%s", out) + + require.Contains(t, out, "could not read the data block", + "an unreadable INSERT should be refused with an explanation") + require.Contains(t, out, "Map(String, UInt64)", "the message should name the offending type") + require.Contains(t, out, "HTTP interface", "the message should point at the way that works") + + // Refusing has to mean the rows never reach ClickHouse. A refusal that still commits the INSERT is + // worse than no check at all. + require.Equal(t, 0, countExotic(t, 99), "the refused rows should not have been written") +} + +// countExotic reads straight from ClickHouse rather than through the proxy. +func countExotic(t *testing.T, id int) int { + t.Helper() + + target := fmt.Sprintf("http://%s/?database=%s", + envOr("PAM_CLICKHOUSE_HTTP", "127.0.0.1:8123"), envOr("PAM_CLICKHOUSE_DB", "analytics")) + + req, err := http.NewRequest(http.MethodPost, target, + strings.NewReader(fmt.Sprintf("SELECT count() FROM exotic WHERE id = %d", id))) + require.NoError(t, err) + req.Header.Set("X-ClickHouse-User", envOr("PAM_CLICKHOUSE_USER", "default")) + req.Header.Set("X-ClickHouse-Key", envOr("PAM_CLICKHOUSE_PASSWORD", "clickhouse")) + + resp, err := (&http.Client{Timeout: 30 * time.Second}).Do(req) + require.NoError(t, err) + defer resp.Body.Close() + + raw, err := io.ReadAll(resp.Body) + require.NoError(t, err) + + count, err := strconv.Atoi(strings.TrimSpace(string(raw))) + require.NoError(t, err, string(raw)) + return count +} + +// A refusal ends the session. Anything else leaves the parser reading a stream it has lost its place in, +// where a later packet can flush the refused bytes upstream and run the statement that was just refused. +func TestBlockedStatementEndsTheSession(t *testing.T) { + itOnly(t) + + recorder := &recordingLogger{} + port := startProxy(t, baseConfig(recorder, `(?i)\bdrop\b`)) + + before := countWriteTest(t) + + out, err := runClient(t, port, + "SELECT 1;\nDROP TABLE pam_write_test;\nINSERT INTO pam_write_test (id, note) VALUES (7777, 'after-block');") + t.Logf("%s", out) + require.Error(t, err) + require.Contains(t, out, "blocked by the command blocking policy") + + // Neither the blocked DROP nor the statement behind it may have run. + require.Equal(t, before, countWriteTest(t), "nothing after a refusal should reach ClickHouse") + require.NotContains(t, out, "after-block") +} + +func countWriteTest(t *testing.T) int { + t.Helper() + + target := fmt.Sprintf("http://%s/?database=%s", + envOr("PAM_CLICKHOUSE_HTTP", "127.0.0.1:8123"), envOr("PAM_CLICKHOUSE_DB", "analytics")) + + req, err := http.NewRequest(http.MethodPost, target, strings.NewReader("SELECT count() FROM pam_write_test")) + require.NoError(t, err) + req.Header.Set("X-ClickHouse-User", envOr("PAM_CLICKHOUSE_USER", "default")) + req.Header.Set("X-ClickHouse-Key", envOr("PAM_CLICKHOUSE_PASSWORD", "clickhouse")) + + resp, err := (&http.Client{Timeout: 30 * time.Second}).Do(req) + require.NoError(t, err) + defer resp.Body.Close() + + raw, err := io.ReadAll(resp.Body) + require.NoError(t, err) + + count, err := strconv.Atoi(strings.TrimSpace(string(raw))) + require.NoError(t, err, string(raw)) + return count +} + +// ---------- TLS ---------- + +func tlsConfigFor(t *testing.T, insecure bool) *tls.Config { + t.Helper() + config := &tls.Config{ServerName: "localhost", InsecureSkipVerify: insecure} + if insecure { + return config + } + pem, err := os.ReadFile(os.Getenv("PAM_CLICKHOUSE_TLS_CERT")) + require.NoError(t, err) + pool := x509.NewCertPool() + require.True(t, pool.AppendCertsFromPEM(pem)) + config.RootCAs = pool + return config +} + +func tlsConfigSkipOrConfig(t *testing.T) (ClickHouseProxyConfig, bool) { + t.Helper() + native := os.Getenv("PAM_CLICKHOUSE_TLS_NATIVE") + httpAddr := os.Getenv("PAM_CLICKHOUSE_TLS_HTTP") + if native == "" || httpAddr == "" { + return ClickHouseProxyConfig{}, false + } + return ClickHouseProxyConfig{ + TargetAddr: httpAddr, + NativeAddr: native, + Username: "default", + Password: "clickhouse", + Database: "analytics", + EnableTLS: true, + TLSConfig: tlsConfigFor(t, true), + SessionID: "edge-tls-test", + SessionLogger: &recordingLogger{}, + }, true +} + +// One SSL toggle covers both interfaces, because ClickHouse shares the certificate between them. +func TestTLSUpstream(t *testing.T) { + itOnly(t) + + config, ok := tlsConfigSkipOrConfig(t) + if !ok { + t.Skip("set PAM_CLICKHOUSE_TLS_NATIVE and PAM_CLICKHOUSE_TLS_HTTP to run") + } + + t.Run("native client over TLS to the server", func(t *testing.T) { + port := startProxy(t, config) + out, err := runClient(t, port, "SELECT note FROM t ORDER BY id;") + require.NoError(t, err, out) + require.Contains(t, out, "tls-one") + }) + + t.Run("http client over TLS to the server", func(t *testing.T) { + port := startProxy(t, config) + status, body := postStatement(t, "127.0.0.1:"+port, "SELECT count() AS c FROM t") + require.Equal(t, http.StatusOK, status, body) + require.Contains(t, body, "2") + }) + + t.Run("connection tests reach both interfaces over TLS", func(t *testing.T) { + require.NoError(t, TestConnection(context.Background(), config)) + require.NoError(t, TestNativeConnection(context.Background(), config)) + }) + + t.Run("bridging over TLS for a server with no HTTP", func(t *testing.T) { + bridged := config + bridged.TargetAddr = "" + port := startProxy(t, bridged) + + status, body := postStatement(t, "127.0.0.1:"+port, "SELECT id, note FROM t ORDER BY id \nFORMAT JSON") + require.Equal(t, http.StatusOK, status, body) + + var envelope bridgeEnvelope + require.NoError(t, json.Unmarshal([]byte(body), &envelope), body) + require.Equal(t, 2, envelope.Rows) + require.JSONEq(t, `[{"id":1,"note":"tls-one"},{"id":2,"note":"tls-two"}]`, string(envelope.Data)) + }) + + t.Run("a pinned CA verifies rather than skipping", func(t *testing.T) { + if os.Getenv("PAM_CLICKHOUSE_TLS_CERT") == "" { + t.Skip("set PAM_CLICKHOUSE_TLS_CERT to run") + } + verified := config + verified.TLSConfig = tlsConfigFor(t, false) + require.NoError(t, TestNativeConnection(context.Background(), verified)) + }) + + t.Run("an untrusted certificate is refused when verification is on", func(t *testing.T) { + strict := config + strict.TLSConfig = &tls.Config{ServerName: "localhost"} + err := TestNativeConnection(context.Background(), strict) + require.Error(t, err, "a self-signed certificate should not verify") + require.Contains(t, strings.ToLower(err.Error()), "certificate") + }) +} + +// ---------- account shapes ---------- + +func TestAccountWithoutNativePortRefusesNativeClients(t *testing.T) { + itOnly(t) + + config := baseConfig(&recordingLogger{}) + config.NativeAddr = "" + port := startProxy(t, config) + + out, err := runClient(t, port, "SELECT 1;") + t.Logf("%s", out) + require.Error(t, err) + require.Contains(t, out, "native port") + + // The HTTP interface has to keep working on the same account. + status, body := postStatement(t, "127.0.0.1:"+port, "SELECT 1") + require.Equal(t, http.StatusOK, status, body) +} + +func TestAccountWithoutHTTPPortStillServesBothClients(t *testing.T) { + itOnly(t) + + config := baseConfig(&recordingLogger{}) + config.TargetAddr = "" + port := startProxy(t, config) + + out, err := runClient(t, port, "SELECT 'native-on-bridged-account';") + require.NoError(t, err, out) + require.Contains(t, out, "native-on-bridged-account") + + status, body := postStatement(t, "127.0.0.1:"+port, "SELECT 1 AS n \nFORMAT JSON") + require.Equal(t, http.StatusOK, status, body) + require.Contains(t, body, `"n":1`) +} + +func TestAccountWithNeitherPortFailsClearly(t *testing.T) { + itOnly(t) + + config := baseConfig(&recordingLogger{}) + config.TargetAddr = "" + config.NativeAddr = "" + port := startProxy(t, config) + + // The session layer rejects this config before a handler ever runs, so the handler's own guard simply + // refuses the connection rather than answering it. + _, _, err := postStatementE("127.0.0.1:"+port, "SELECT 1") + require.Error(t, err, "a session with neither port must not serve anything") +} + +// ---------- bridge edge cases ---------- + +func TestBridgeEdgeCases(t *testing.T) { + itOnly(t) + + config := baseConfig(&recordingLogger{}) + config.TargetAddr = "" + + cases := []struct { + name string + sql string + wantStatus int + wantBody string + }{ + { + name: "a statement shape that cannot be wrapped says so", + sql: "SHOW TABLES \nFORMAT JSON", + wantStatus: http.StatusBadRequest, + wantBody: "SELECT, WITH or EXPLAIN", + }, + { + name: "a format the bridge does not produce says so", + sql: "SELECT 1 \nFORMAT TabSeparated", + wantStatus: http.StatusBadRequest, + wantBody: "JSON and JSONCompact", + }, + { + name: "a statement with no format runs and returns nothing to parse", + sql: "CREATE TABLE IF NOT EXISTS bridge_ddl (a UInt8) ENGINE = Memory", + wantStatus: http.StatusOK, + }, + { + name: "an empty statement is refused", + sql: "", + wantStatus: http.StatusBadRequest, + wantBody: "No statement was sent", + }, + { + name: "a syntax error comes back as ClickHouse wrote it", + sql: "SELECT FROM WHERE \nFORMAT JSON", + wantStatus: http.StatusBadRequest, + wantBody: "Syntax error", + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + port := startProxy(t, config) + status, body := postStatement(t, "127.0.0.1:"+port, tc.sql) + require.Equal(t, tc.wantStatus, status, body) + if tc.wantBody != "" { + require.Contains(t, body, tc.wantBody) + } + }) + } +} + +// A parameterized statement has to run and be recorded with its values, the same as over HTTP. +func TestBridgePassesQueryParameters(t *testing.T) { + itOnly(t) + + recorder := &recordingLogger{} + config := baseConfig(recorder) + config.TargetAddr = "" + port := startProxy(t, config) + + status, body := postGET(t, "127.0.0.1:"+port, + "/?param_wanted=7&query="+urlEscape("SELECT {wanted:UInt8} AS got FORMAT JSON")) + require.Equal(t, http.StatusOK, status, body) + require.Contains(t, body, `"got":7`) + require.Contains(t, recorder.dump(), "wanted=7", "the parameter belongs in the recording") +} + +// A gzipped body has to work on both paths. The reverse proxy hands it to ClickHouse untouched, while the +// bridge has to run the decoded statement itself. +func TestCompressedRequestBodies(t *testing.T) { + itOnly(t) + + for _, native := range []bool{false, true} { + name := "http interface" + if native { + name = "bridged to native" + } + t.Run(name, func(t *testing.T) { + config := baseConfig(&recordingLogger{}) + if native { + config.TargetAddr = "" + } + port := startProxy(t, config) + + status, body := postGzipped(t, "127.0.0.1:"+port, "SELECT 5 AS five \nFORMAT JSON") + require.Equal(t, http.StatusOK, status, body) + + // ClickHouse pretty-prints its JSON and the bridge writes it compact, so the rows are compared + // rather than the bytes. + var envelope bridgeEnvelope + require.NoError(t, json.Unmarshal([]byte(body), &envelope), body) + require.JSONEq(t, `[{"five":5}]`, string(envelope.Data)) + }) + } +} + +// ---------- sniffing ---------- + +func TestSnifferEdgeCases(t *testing.T) { + itOnly(t) + + port := startProxy(t, baseConfig(&recordingLogger{})) + + t.Run("a client that connects and says nothing is eventually dropped", func(t *testing.T) { + conn, err := net.Dial("tcp", "127.0.0.1:"+port) + require.NoError(t, err) + defer conn.Close() + + // The sniff deadline is what guarantees this; without it the handler would hold the session open. + require.NoError(t, conn.SetReadDeadline(time.Now().Add(sniffTimeout+10*time.Second))) + buf := make([]byte, 64) + _, err = conn.Read(buf) + require.Error(t, err, "the gateway should close a connection that never says anything") + require.NotErrorIs(t, err, os.ErrDeadlineExceeded, "the gateway held the connection open past the sniff timeout") + }) + + t.Run("garbage is not mistaken for either protocol", func(t *testing.T) { + conn, err := net.Dial("tcp", "127.0.0.1:"+port) + require.NoError(t, err) + defer conn.Close() + _, err = conn.Write([]byte{0xFF, 0xFE, 0xFD, 0xFC}) + require.NoError(t, err) + // Bytes that are not a request line leave net/http waiting for headers, so the bound here is its + // ReadHeaderTimeout rather than the sniff timeout. + require.NoError(t, conn.SetReadDeadline(time.Now().Add(45*time.Second))) + + buf := make([]byte, 256) + n, err := conn.Read(buf) + require.NotErrorIs(t, err, os.ErrDeadlineExceeded, "garbage must not leave the handler hanging") + if err == nil || n > 0 { + require.Contains(t, string(buf[:n]), "400", "garbage should be answered as a bad HTTP request") + } else { + require.ErrorIs(t, err, io.EOF) + } + }) + + t.Run("the /ping path answers", func(t *testing.T) { + resp, err := (&http.Client{Timeout: 10 * time.Second}).Get("http://127.0.0.1:" + port + "/ping") + require.NoError(t, err) + defer resp.Body.Close() + require.Equal(t, http.StatusOK, resp.StatusCode) + }) + + t.Run("a path outside the query endpoint is refused", func(t *testing.T) { + resp, err := (&http.Client{Timeout: 10 * time.Second}).Get("http://127.0.0.1:" + port + "/play") + require.NoError(t, err) + defer resp.Body.Close() + require.Equal(t, http.StatusNotFound, resp.StatusCode) + }) +} + +// ---------- concurrency ---------- + +// A session hands out one port that many clients share, so the handler has to hold up under parallel use +// of both protocols at once. +func TestConcurrentMixedProtocolSessions(t *testing.T) { + itOnly(t) + + recorder := &recordingLogger{} + port := startProxy(t, baseConfig(recorder)) + + var wg sync.WaitGroup + errs := make(chan error, 16) + + for i := range 6 { + wg.Add(1) + go func(n int) { + defer wg.Done() + marker := fmt.Sprintf("concurrent-native-%d", n) + out, err := runClient(t, port, fmt.Sprintf("SELECT '%s';", marker)) + if err != nil { + errs <- fmt.Errorf("native %d: %v\n%s", n, err, out) + return + } + if !strings.Contains(out, marker) { + errs <- fmt.Errorf("native %d: missing marker in %s", n, out) + } + }(i) + } + + for i := range 6 { + wg.Add(1) + go func(n int) { + defer wg.Done() + status, body, err := postStatementE("127.0.0.1:"+port, fmt.Sprintf("SELECT %d AS n", n)) + if err != nil { + errs <- fmt.Errorf("http %d: %v", n, err) + return + } + if status != http.StatusOK { + errs <- fmt.Errorf("http %d: status %d: %s", n, status, body) + } + }(i) + } + + wg.Wait() + close(errs) + for err := range errs { + t.Error(err) + } +} + +// ---------- volume ---------- + +func TestNativeLargeResultSet(t *testing.T) { + itOnly(t) + + port := startProxy(t, baseConfig(&recordingLogger{})) + // The rows have to actually cross the proxy, or none of the multi-block relay is exercised. + out, err := runClient(t, port, "SELECT number FROM numbers(300000);") + require.NoError(t, err, out) + + lines := strings.Count(strings.TrimSpace(out), "\n") + 1 + require.Equal(t, 300000, lines, "every row should reach the client") + require.Contains(t, out, "299999", "the last row should survive the relay") +} + +func TestNativeWideRowsStreamThrough(t *testing.T) { + itOnly(t) + + port := startProxy(t, baseConfig(&recordingLogger{})) + + direct := queryDirect(t, "SELECT sum(length(payload)) FROM wide_blobs") + require.NotEqual(t, "0", direct, "wide_blobs must be seeded for this to test anything") + + out, err := runClient(t, port, "SELECT sum(length(payload)) FROM wide_blobs;") + require.NoError(t, err, out) + require.Equal(t, direct, strings.TrimSpace(out), "the proxied total must match the server's") +} + +// queryDirect bypasses the proxy so a test can state what the answer should be. +func queryDirect(t *testing.T, sql string) string { + t.Helper() + + target := fmt.Sprintf("http://%s/?database=%s", + envOr("PAM_CLICKHOUSE_HTTP", "127.0.0.1:8123"), envOr("PAM_CLICKHOUSE_DB", "analytics")) + req, err := http.NewRequest(http.MethodPost, target, strings.NewReader(sql)) + require.NoError(t, err) + req.Header.Set("X-ClickHouse-User", envOr("PAM_CLICKHOUSE_USER", "default")) + req.Header.Set("X-ClickHouse-Key", envOr("PAM_CLICKHOUSE_PASSWORD", "clickhouse")) + + resp, err := (&http.Client{Timeout: 60 * time.Second}).Do(req) + require.NoError(t, err) + defer resp.Body.Close() + + raw, err := io.ReadAll(resp.Body) + require.NoError(t, err) + return strings.TrimSpace(string(raw)) +} + +// ---------- upstream failures ---------- + +func TestUpstreamUnreachable(t *testing.T) { + itOnly(t) + + t.Run("native client gets a native exception", func(t *testing.T) { + config := baseConfig(&recordingLogger{}) + config.NativeAddr = "127.0.0.1:1" + port := startProxy(t, config) + + out, err := runClient(t, port, "SELECT 1;") + t.Logf("%s", out) + require.Error(t, err) + require.Contains(t, out, "could not reach ClickHouse") + }) + + t.Run("bridge reports a bad gateway", func(t *testing.T) { + config := baseConfig(&recordingLogger{}) + config.TargetAddr = "" + config.NativeAddr = "127.0.0.1:1" + port := startProxy(t, config) + + status, body := postStatement(t, "127.0.0.1:"+port, "SELECT 1 \nFORMAT JSON") + require.Equal(t, http.StatusBadGateway, status) + require.Contains(t, body, "could not reach ClickHouse") + }) +} + +func TestWrongAccountCredentialsSurfaceCleanly(t *testing.T) { + itOnly(t) + + config := baseConfig(&recordingLogger{}) + config.Password = "definitely-not-the-password" + port := startProxy(t, config) + + out, err := runClient(t, port, "SELECT 1;") + t.Logf("%s", out) + require.Error(t, err) + require.Contains(t, out, "refused the account") +} + +// ---------- recording ---------- + +func TestNativeRecordingCapturesOutcomes(t *testing.T) { + itOnly(t) + + recorder := &recordingLogger{} + port := startProxy(t, baseConfig(recorder)) + + out, err := runClient(t, port, "SELECT 1;\nSELECT * FROM nope_not_here;") + t.Logf("%s", out) + require.Error(t, err) + + waitFor(t, func() bool { return strings.Contains(recorder.dump(), "ERROR:") }) + dump := recorder.dump() + require.Contains(t, dump, "SELECT 1") + require.Contains(t, dump, "nope_not_here") + require.Equal(t, 1, strings.Count(dump, "=> OK"), "exactly the one successful statement keeps its outcome") + require.Equal(t, 1, strings.Count(dump, "ERROR:"), "the failed statement is recorded once") +} + +// The revision the client speaks is newer than ch-go's, so the handshake has to pin it and say so. +func TestNativeRevisionPinning(t *testing.T) { + itOnly(t) + + port := startProxy(t, baseConfig(&recordingLogger{})) + + // The server is newer than ch-go, so this only passes if the pinned revision is honoured end to end. + out, err := runClient(t, port, "SELECT version();") + require.NoError(t, err, out) + require.Equal(t, queryDirect(t, "SELECT version()"), strings.TrimSpace(out)) +} + +func TestQuoteFieldDump(t *testing.T) { + cases := []struct{ in, want string }{ + {"7", `'7'`}, + {"plain", `'plain'`}, + {"it's", `'it\'s'`}, + {`back\slash`, `'back\\slash'`}, + {`'; DROP TABLE users; --`, `'\'; DROP TABLE users; --'`}, + {"", `''`}, + } + for _, tc := range cases { + require.Equal(t, tc.want, quoteFieldDump(tc.in), tc.in) + } +} + +// A parameter is data, so a value full of quotes has to come back as that value rather than changing +// the statement around it. +func TestBridgeParameterCannotEscapeItsQuotes(t *testing.T) { + itOnly(t) + + config := baseConfig(&recordingLogger{}) + config.TargetAddr = "" + port := startProxy(t, config) + + hostile := `'; DROP TABLE pam_write_test; --` + status, body := postGET(t, "127.0.0.1:"+port, + "/?param_v="+urlEscape(hostile)+"&query="+urlEscape("SELECT {v:String} AS got FORMAT JSON")) + + require.Equal(t, http.StatusOK, status, body) + + var envelope bridgeEnvelope + require.NoError(t, json.Unmarshal([]byte(body), &envelope), body) + + var rows []struct { + Got string `json:"got"` + } + require.NoError(t, json.Unmarshal(envelope.Data, &rows)) + require.Len(t, rows, 1) + require.Equal(t, hostile, rows[0].Got, "the value should survive intact, not be executed") + + // The table the payload tried to drop is still there. + okStatus, okBody := postStatement(t, "127.0.0.1:"+port, "SELECT count() AS c FROM pam_write_test \nFORMAT JSON") + require.Equal(t, http.StatusOK, okStatus, okBody) +} diff --git a/packages/pam/handlers/clickhouse/native.go b/packages/pam/handlers/clickhouse/native.go new file mode 100644 index 000000000..407651adc --- /dev/null +++ b/packages/pam/handlers/clickhouse/native.go @@ -0,0 +1,683 @@ +package clickhouse + +import ( + "bufio" + "context" + "crypto/tls" + "errors" + "fmt" + "io" + "net" + "os" + "strings" + "sync" + "sync/atomic" + "time" + + "github.com/ClickHouse/ch-go/proto" + "github.com/rs/zerolog" +) + +// Native brokers ClickHouse's TCP protocol on 9000/9440, which the HTTP interface cannot serve: clickhouse-client, +// clickhouse-driver and clickhouse-go all speak it exclusively. A statement arrives length-prefixed in its own +// Query packet, so unlike HTTP there is no inspection window to overrun and an INSERT's rows never reach the policy. +// +// Only the client direction is parsed. The server direction is relayed byte for byte, because nothing in it is +// inspected and decoding it would make a SELECT fail on any column type ch-go cannot infer. + +const ( + // ch-go decodes up to its own revision, and a newer server sends Hello fields it cannot read. Both sides are + // pinned to what we can parse, which downgrades the client too. + maxNativeRevision = proto.Version + + nativeDialTimeout = 30 * time.Second + nativeWriteTimeout = 60 * time.Second + // A port that accepts TCP and then says nothing is the symptom of the HTTP port entered as the native one, + // so the handshake gives up well before the dial would. + nativeHandshakeTimeout = 10 * time.Second + // Long enough that a slow client is not dropped mid-statement, short enough to reap an abandoned session. + nativeIdleTimeout = 12 * time.Hour +) + +type nativeProxy struct { + *ClickHouseProxy +} + +func newNativeProxy(owner *ClickHouseProxy) *nativeProxy { + return &nativeProxy{ClickHouseProxy: owner} +} + +// tap records every byte a decoder consumes so the exact wire bytes can be replayed upstream. Reads are handed +// out one at a time: proto.Reader buffers 128 KB, which would swallow packets we have not parsed yet. +type tap struct { + src *bufio.Reader + buf []byte +} + +func newTap(r io.Reader) *tap { + return &tap{src: bufio.NewReaderSize(r, 64<<10)} +} + +func (t *tap) Read(p []byte) (int, error) { + if len(p) == 0 { + return 0, nil + } + b, err := t.src.ReadByte() + if err != nil { + return 0, err + } + p[0] = b + t.buf = append(t.buf, b) + return 1, nil +} + +// take returns the bytes consumed since the last call and starts a new run. +func (t *tap) take() []byte { + b := t.buf + t.buf = nil + return b +} + +// rest is the reader the tap is draining, buffered bytes included. Relaying from the socket directly would +// silently drop whatever the buffer already holds. +func (t *tap) rest() io.Reader { + return t.src +} + +func (t *tap) discard() { + t.buf = nil +} + +// Writing an exception ends the session. The client's stream is mid-packet once a refusal happens, so +// carrying on would parse garbage, and anything still in the tap could be flushed upstream by a later +// packet: a refused statement would reach ClickHouse after being refused. +var errSessionRefused = errors.New("the session was refused") + +type nativeSession struct { + proxy *nativeProxy + log zerolog.Logger + client net.Conn + upstream net.Conn + rev int + + // One tap over the upstream, shared by the handshake and the server loop. A second tap would start its + // own buffer and lose whatever the first had already read off the socket. + upstreamTap *tap + upstreamReader *proto.Reader + + outcomes *outcomeRecorder + + // Both directions write to the client, so a refusal must not land inside a packet the server loop is + // still writing. Once refused, the server loop stops writing entirely. + writeMu sync.Mutex + refused atomic.Bool + + // Set by the Query packet and read by the data-block decoder + compressed atomic.Bool +} + +// writeToClient serialises the two directions and bounds a client that has stopped reading. +func (s *nativeSession) writeToClient(payload []byte) error { + if len(payload) == 0 { + return nil + } + s.writeMu.Lock() + defer s.writeMu.Unlock() + + _ = s.client.SetWriteDeadline(time.Now().Add(nativeWriteTimeout)) + defer func() { _ = s.client.SetWriteDeadline(time.Time{}) }() + + _, err := s.client.Write(payload) + return err +} + +func (s *nativeSession) writeToUpstream(payload []byte) error { + if len(payload) == 0 { + return nil + } + _ = s.upstream.SetWriteDeadline(time.Now().Add(nativeWriteTimeout)) + defer func() { _ = s.upstream.SetWriteDeadline(time.Time{}) }() + + _, err := s.upstream.Write(payload) + return err +} + +func (p *nativeProxy) HandleConnection(ctx context.Context, clientConn net.Conn, l zerolog.Logger) error { + defer clientConn.Close() + // A malformed packet must not take the gateway, and every other session with it, down. + defer func() { + if r := recover(); r != nil { + l.Error().Interface("panic", r).Msg("Recovered from a panic in the ClickHouse native handler") + } + }() + + upstream, err := p.dialUpstream(ctx) + if err != nil { + l.Error().Err(err).Msg("Failed to reach ClickHouse over the native protocol") + return writeNativeError(clientConn, maxNativeRevision, codeNetworkError, + fmt.Sprintf("The gateway could not reach ClickHouse: %v", err)) + } + defer upstream.Close() + + stop := context.AfterFunc(ctx, func() { + clientConn.Close() + upstream.Close() + }) + defer stop() + + s := &nativeSession{proxy: p, log: l, client: clientConn, upstream: upstream} + s.outcomes = newOutcomeRecorder(p.ClickHouseProxy) + s.upstreamTap = newTap(upstream) + s.upstreamReader = proto.NewReader(s.upstreamTap) + + clientTap := newTap(clientConn) + clientReader := proto.NewReader(clientTap) + + // A port that accepts TCP and then says nothing is the symptom of the HTTP port entered as the native + // one, so the handshake gives up rather than holding both sockets until the session expires. + deadline := time.Now().Add(nativeHandshakeTimeout) + _ = clientConn.SetDeadline(deadline) + _ = upstream.SetDeadline(deadline) + + if err := s.handshake(clientTap, clientReader); err != nil { + l.Debug().Err(err).Msg("ClickHouse native handshake ended") + return nil + } + + _ = clientConn.SetDeadline(time.Time{}) + _ = upstream.SetDeadline(time.Time{}) + + // Nothing in the server direction gates a statement, so a packet it cannot read costs only the outcome of + // that statement: it flushes what it read and streams the rest untouched. + serverDone := make(chan struct{}) + go func() { + defer close(serverDone) + defer func() { + if r := recover(); r != nil { + l.Error().Interface("panic", r).Msg("Recovered from a panic reading the ClickHouse server direction") + } + }() + s.serverLoop() + }() + + if err := s.clientLoop(clientTap, clientReader); err != nil { + l.Debug().Err(err).Msg("ClickHouse native session ended") + } + + // Closing the upstream unblocks the server loop. Its outcomes have to land before the recorder drains, + // or a statement that finished is written to the recording as interrupted. + upstream.Close() + select { + case <-serverDone: + case <-time.After(nativeWriteTimeout): + l.Warn().Msg("The ClickHouse server direction did not stop in time") + } + s.outcomes.finish() + return nil +} + +// refuse reports the refusal to the client and ends the session. +func (s *nativeSession) refuse(t *tap, code int, message string) error { + if t != nil { + t.discard() + } + // Set before writing, so the server loop cannot interleave a packet with the exception. + s.refused.Store(true) + + if err := s.writeToClient(nativeErrorPacket(s.rev, code, message)); err != nil { + return err + } + return errSessionRefused +} + +func (p *nativeProxy) dialUpstream(ctx context.Context) (net.Conn, error) { + dialer := &net.Dialer{Timeout: nativeDialTimeout} + if !p.config.EnableTLS { + return dialer.DialContext(ctx, "tcp", p.config.NativeAddr) + } + return (&tls.Dialer{NetDialer: dialer, Config: p.config.TLSConfig}).DialContext(ctx, "tcp", p.config.NativeAddr) +} + +// handshake swaps the client's credentials for the account's, so nothing the client holds works outside a +// recorded session, and pins the protocol revision both ways. +func (s *nativeSession) handshake(t *tap, r *proto.Reader) error { + code, err := r.UVarInt() + if err != nil { + return fmt.Errorf("read client packet code: %w", err) + } + if proto.ClientCode(code) != proto.ClientCodeHello { + return fmt.Errorf("expected Hello, got client packet %d", code) + } + + var hello proto.ClientHello + if err := hello.Decode(r); err != nil { + return fmt.Errorf("decode client hello: %w", err) + } + t.discard() + + // The revision drives feature gating on both sides, and it is unvalidated client input: a huge uvarint + // decodes to a negative int, which would be re-encoded upstream as an enormous revision. + if hello.ProtocolVersion <= 0 { + return s.refuse(t, codeNotImplemented, + fmt.Sprintf("This session could not read the protocol revision %d the client asked for.", + hello.ProtocolVersion)) + } + s.rev = min(hello.ProtocolVersion, maxNativeRevision) + + var b proto.Buffer + proto.ClientHello{ + Name: hello.Name, + Major: hello.Major, + Minor: hello.Minor, + ProtocolVersion: s.rev, + Database: s.proxy.config.Database, + User: s.proxy.config.Username, + Password: s.proxy.config.Password, + }.Encode(&b) + if err := s.writeToUpstream(b.Buf); err != nil { + return fmt.Errorf("write upstream hello: %w", err) + } + + serverReader := s.upstreamReader + serverCode, err := serverReader.UVarInt() + if err != nil { + return fmt.Errorf("read server packet code: %w", err) + } + + if proto.ServerCode(serverCode) == proto.ServerCodeException { + var e proto.Exception + if err := e.DecodeAware(serverReader, s.rev); err != nil { + return fmt.Errorf("decode upstream exception: %w", err) + } + s.log.Warn().Str("upstreamError", e.Message).Msg("ClickHouse refused the account credentials") + return s.refuse(t, codeAccessDenied, + fmt.Sprintf("ClickHouse refused the account this session uses: %s", e.Message)) + } + if proto.ServerCode(serverCode) != proto.ServerCodeHello { + return fmt.Errorf("expected server Hello, got packet %d", serverCode) + } + + var serverHello proto.ServerHello + if err := serverHello.DecodeAware(serverReader, s.rev); err != nil { + return fmt.Errorf("decode server hello: %w", err) + } + // The client keys its own encoding off the revision it is told, so it has to see the pinned one. + serverHello.Revision = s.rev + + s.upstreamTap.discard() + + b.Reset() + serverHello.EncodeAware(&b, s.rev) + if err := s.writeToClient(b.Buf); err != nil { + return fmt.Errorf("write client hello response: %w", err) + } + + // At rev >= 54458 the client follows the handshake with its quota key, written as a bare string rather than + // a coded packet. Ours is empty: the account's quota is not the client's to choose. + if proto.FeatureAddendum.In(s.rev) { + if _, err := r.Str(); err != nil { + return fmt.Errorf("read client addendum: %w", err) + } + t.discard() + + b.Reset() + b.PutString("") + if err := s.writeToUpstream(b.Buf); err != nil { + return fmt.Errorf("write upstream addendum: %w", err) + } + } + + s.log.Info(). + Str("clientName", hello.Name). + Int("revision", s.rev). + Msg("ClickHouse native session established") + return nil +} + +// clientLoop parses every packet the client sends. A statement that is never parsed is a statement the policy +// never sees, so an unreadable stream ends the session rather than being relayed blind. +func (s *nativeSession) clientLoop(t *tap, r *proto.Reader) error { + for { + _ = s.client.SetReadDeadline(time.Now().Add(nativeIdleTimeout)) + + code, err := r.UVarInt() + if err != nil { + return fmt.Errorf("client hung up: %w", err) + } + + switch proto.ClientCode(code) { + case proto.ClientCodePing, proto.ClientCodeCancel: + if err := s.forward(t.take()); err != nil { + return err + } + + case proto.ClientTablesStatusRequest: + // This one carries a table list the loop does not decode. Forwarding just the code would leave + // the stream one packet out of step and every later statement unreadable. + return s.refuse(t, codeNotImplemented, + "This session does not support ClickHouse's tables-status request.") + + case proto.ClientCodeQuery: + if err := s.handleQuery(t, r); err != nil { + return err + } + + case proto.ClientCodeData: + if err := s.handleData(t, r); err != nil { + return err + } + + default: + s.log.Warn().Uint64("packetCode", code).Msg("Refused an unreadable ClickHouse client packet") + return s.refuse(t, codeNotImplemented, + fmt.Sprintf("This session could not read ClickHouse client packet %d, so the command blocking "+ + "policy could not be applied to it.", code)) + } + } +} + +func (s *nativeSession) handleQuery(t *tap, r *proto.Reader) error { + var q proto.Query + if err := q.DecodeAware(r, s.rev); err != nil { + s.log.Warn().Err(err).Msg("Could not read a ClickHouse query packet") + return s.refuse(t, codeNotImplemented, + "This session could not read the query packet, so the command blocking policy could not be "+ + "applied to it.") + } + t.discard() + + // EncodeAware always writes StageComplete, so a client asking for a partial stage would have its query + // silently upgraded to a full execution. Refusing is the honest answer. + if q.Stage != proto.StageComplete { + return s.refuse(t, codeNotImplemented, + fmt.Sprintf("This session only runs statements to completion, and this client asked for stage %d.", + int(q.Stage))) + } + + s.compressed.Store(q.Compression == proto.CompressionEnabled) + + statement := q.Body + nativeParameterSuffix(q.Parameters) + + if blocked := s.proxy.blockedBy(statement); blocked != nil { + s.proxy.logStatement(statement, fmt.Sprintf("BLOCKED: %s", blocked.String())) + s.log.Info().Str("pattern", blocked.String()).Msg("Blocked a statement by policy") + return s.refuse(t, codeAccessDenied, + "This statement is blocked by the command blocking policy on this account.") + } + + s.outcomes.begin(statement) + + // The identities a client could otherwise choose for itself, matching the set the HTTP interface strips. + // An inter-server secret in particular must never be something a session gets to pick. + // InitialUser and InitialAddress are deliberately left alone: forcing the query kind to Initial already + // makes ClickHouse authorise as the account, and it asserts on an empty initial address. + q.Info.QuotaKey = "" + q.Info.Query = proto.ClientQueryInitial + q.Secret = "" + + var b proto.Buffer + q.EncodeAware(&b, s.rev) + return s.forward(b.Buf) +} + +// handleData decodes a block only far enough to find where it ends, then replays the client's own bytes. The +// decoded values are discarded: re-encoding them would have to reproduce a serialization we do not own. +func (s *nativeSession) handleData(t *tap, r *proto.Reader) error { + table, err := r.Str() + if err != nil { + s.log.Warn().Err(err).Msg("Could not read a ClickHouse data packet") + return s.refuse(t, codeNotImplemented, + "This session could not read the data packet that followed this statement.") + } + + compressed := s.compressed.Load() + if compressed { + r.EnableCompression() + } + var ( + block proto.Block + discard proto.Results + ) + decodeErr := block.DecodeBlock(r, s.rev, discard.Auto()) + if compressed { + r.DisableCompression() + } + + if decodeErr != nil { + s.log.Warn().Err(decodeErr).Str("table", table).Msg("Could not read a ClickHouse data block") + return s.refuse(t, codeNotImplemented, + fmt.Sprintf("This session could not read the data block sent with this statement, so it was not "+ + "forwarded: %v. Sending this data over ClickHouse's HTTP interface avoids the limitation.", decodeErr)) + } + + return s.forward(t.take()) +} + +// serverLoop reads the server direction for the sake of the recording only. Every packet is replayed to the +// client byte for byte, and the first one it cannot read ends the parsing rather than the session. +func (s *nativeSession) serverLoop() { + t := s.upstreamTap + r := s.upstreamReader + + relayRest := func(reason string) { + s.outcomes.degrade(reason) + if s.refused.Load() { + return + } + if err := s.writeToClient(t.take()); err != nil { + return + } + // Reading the socket directly here would skip whatever the tap has already buffered off it, which + // is most of the in-flight response. + _, _ = io.Copy(newRefusalAwareWriter(s), t.rest()) + } + + for { + code, err := r.UVarInt() + if err != nil { + return + } + + switch proto.ServerCode(code) { + case proto.ServerCodePong: + + case proto.ServerCodeEndOfStream: + s.outcomes.complete("OK") + + case proto.ServerCodeException: + var e proto.Exception + if err := e.DecodeAware(r, s.rev); err != nil { + relayRest(err.Error()) + return + } + message := firstLine(e.Message, e.Name) + if len(message) > maxLoggedErrorBytes { + message = message[:maxLoggedErrorBytes] + "... [truncated]" + } + s.outcomes.complete(fmt.Sprintf("ERROR: Code %d: %s", e.Code, message)) + + case proto.ServerCodeProgress: + var p proto.Progress + if err := p.DecodeAware(r, s.rev); err != nil { + relayRest(err.Error()) + return + } + s.outcomes.progress(p.Rows, p.Bytes) + + case proto.ServerCodeProfile: + var p proto.Profile + if err := p.DecodeAware(r, s.rev); err != nil { + relayRest(err.Error()) + return + } + + case proto.ServerCodeTableColumns: + if _, err := r.Str(); err != nil { + relayRest(err.Error()) + return + } + if _, err := r.Str(); err != nil { + relayRest(err.Error()) + return + } + + case proto.ServerCodeData, proto.ServerCodeTotals, proto.ServerCodeExtremes, proto.ServerCodeLog, + proto.ServerProfileEvents: + if err := s.skipServerBlock(r, proto.ServerCode(code)); err != nil { + relayRest(err.Error()) + return + } + + default: + relayRest(fmt.Sprintf("unreadable server packet %d", code)) + return + } + + if s.refused.Load() { + return + } + if err := s.writeToClient(t.take()); err != nil { + return + } + } +} + +// refusalAwareWriter stops relaying once the client loop has refused the session, so a raw relay cannot +// append bytes after the exception the client was just sent. +type refusalAwareWriter struct{ s *nativeSession } + +func newRefusalAwareWriter(s *nativeSession) io.Writer { return refusalAwareWriter{s: s} } + +func (w refusalAwareWriter) Write(p []byte) (int, error) { + if w.s.refused.Load() { + return 0, errSessionRefused + } + if err := w.s.writeToClient(p); err != nil { + return 0, err + } + return len(p), nil +} + +// Log and profile-event blocks are never compressed, whatever the query asked for. +func (s *nativeSession) skipServerBlock(r *proto.Reader, code proto.ServerCode) error { + if _, err := r.Str(); err != nil { + return fmt.Errorf("read block table name: %w", err) + } + + compressed := s.compressed.Load() && + code != proto.ServerCodeLog && code != proto.ServerProfileEvents + if compressed { + r.EnableCompression() + defer r.DisableCompression() + } + + var ( + block proto.Block + discard proto.Results + ) + return block.DecodeBlock(r, s.rev, discard.Auto()) +} + +func (s *nativeSession) forward(payload []byte) error { + if err := s.writeToUpstream(payload); err != nil { + return fmt.Errorf("forward to ClickHouse: %w", err) + } + return nil +} + +// nativeParameterSuffix mirrors the HTTP handler, so a parameterized statement reads the same in a recording +// whichever interface ran it. +func nativeParameterSuffix(parameters []proto.Parameter) string { + if len(parameters) == 0 { + return "" + } + pairs := make([]string, 0, len(parameters)) + for _, p := range parameters { + pairs = append(pairs, p.Key+"="+p.Value) + } + return "\n-- parameters: " + strings.Join(pairs, " ") +} + +// writeNativeError reports a gateway refusal the way ClickHouse reports its own, so a driver surfaces it as a +// server exception rather than a broken connection. +func writeNativeError(w io.Writer, revision int, code int, message string) error { + if _, err := w.Write(nativeErrorPacket(revision, code, message)); err != nil { + return fmt.Errorf("write native exception: %w", err) + } + return nil +} + +// nativeErrorPacket builds the exception and the end-of-stream that closes it out. +func nativeErrorPacket(revision int, code int, message string) []byte { + if revision <= 0 { + revision = maxNativeRevision + } + + var b proto.Buffer + proto.ServerCodeException.Encode(&b) + exception := proto.Exception{ + Code: proto.Error(code), + Name: "DB::Exception", + Message: message, + } + exception.EncodeAware(&b, revision) + proto.ServerCodeEndOfStream.Encode(&b) + return b.Buf +} + +// TestNativeConnection proves the account can log in over the native port. ClickHouse validates credentials +// during the handshake, so a successful Hello exchange is a real auth check rather than a reachability probe. +func TestNativeConnection(ctx context.Context, config ClickHouseProxyConfig) error { + dialCtx, cancel := context.WithTimeout(ctx, nativeDialTimeout) + defer cancel() + + conn, err := (&nativeProxy{ClickHouseProxy: &ClickHouseProxy{config: config}}).dialUpstream(dialCtx) + if err != nil { + return err + } + defer conn.Close() + + _ = conn.SetDeadline(time.Now().Add(nativeHandshakeTimeout)) + + var b proto.Buffer + proto.ClientHello{ + Name: "Infisical PAM", + Major: 1, + Minor: 0, + ProtocolVersion: maxNativeRevision, + Database: config.Database, + User: config.Username, + Password: config.Password, + }.Encode(&b) + if _, err := conn.Write(b.Buf); err != nil { + return fmt.Errorf("write hello: %w", err) + } + + r := proto.NewReader(newTap(conn)) + code, err := r.UVarInt() + if err != nil { + if errors.Is(err, os.ErrDeadlineExceeded) { + return fmt.Errorf("the port accepted the connection but did not answer ClickHouse's native "+ + "handshake within %s, which is what the HTTP port does when it is entered as the native one", + nativeHandshakeTimeout) + } + return fmt.Errorf("read hello response: %w", err) + } + + switch proto.ServerCode(code) { + case proto.ServerCodeHello: + var hello proto.ServerHello + if err := hello.DecodeAware(r, maxNativeRevision); err != nil { + return fmt.Errorf("decode hello response: %w", err) + } + return nil + case proto.ServerCodeException: + var e proto.Exception + if err := e.DecodeAware(r, maxNativeRevision); err != nil { + return fmt.Errorf("decode exception: %w", err) + } + return fmt.Errorf("clickhouse rejected the connection: %s", e.Message) + default: + return fmt.Errorf("unexpected server packet %d during the native handshake", code) + } +} diff --git a/packages/pam/handlers/clickhouse/native_integration_test.go b/packages/pam/handlers/clickhouse/native_integration_test.go new file mode 100644 index 000000000..83f6c9b75 --- /dev/null +++ b/packages/pam/handlers/clickhouse/native_integration_test.go @@ -0,0 +1,321 @@ +package clickhouse + +import ( + "context" + "fmt" + "net" + "os" + "regexp" + "strings" + "testing" + "time" +) + +// Exercises the native handler against a real ClickHouse over the real clickhouse-client, which is the only way +// to cover revision pinning, the addendum and block framing. Opt in with PAM_CLICKHOUSE_NATIVE_IT=1. +// +// docker run -d --name pam-clickhouse-target -p 8123:8123 -p 9000:9000 \ +// -e CLICKHOUSE_PASSWORD=clickhouse -e CLICKHOUSE_DB=analytics clickhouse/clickhouse-server:24.8 +func TestNativeIntegration(t *testing.T) { + if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { + t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") + } + + type testCase struct { + name string + sql string + blocked []string + wantOutput []string + wantFailure string + } + + cases := []testCase{ + { + name: "multiple statements on one connection", + sql: "SELECT 1;\nSELECT 2;\nSELECT 3;", + wantOutput: []string{"1", "2", "3"}, + }, + { + name: "credentials come from the account, not the client", + sql: "SELECT currentUser();", + wantOutput: []string{"default"}, + }, + { + name: "column types ch-go cannot infer still stream back", + sql: "SELECT map('a', 1::UInt64) AS m, tuple('p', 2) AS t;", + wantOutput: []string{"{'a':1}", "('p',2)"}, + }, + { + // A fixed marker would be satisfied by a previous run's row even if the insert path regressed. + name: "insert pushes a client data block", + sql: insertMarkerSQL(), + wantOutput: []string{insertMarker}, + }, + { + name: "a blocked statement is refused as a native exception", + sql: "DROP TABLE pam_write_test;", + blocked: []string{`(?i)\bdrop\b`}, + wantFailure: "blocked by the command blocking policy", + }, + { + name: "blocking still applies after an earlier statement on the same connection", + sql: "SELECT 1;\nDROP TABLE pam_write_test;", + blocked: []string{`(?i)\bdrop\b`}, + wantFailure: "blocked by the command blocking policy", + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + recorder := &recordingLogger{} + addr := startNativeProxy(t, tc.blocked, recorder) + + out, err := runClient(t, addr, tc.sql) + t.Logf("client output:\n%s", out) + + if tc.wantFailure != "" { + if err == nil { + t.Fatalf("expected the client to fail, got success:\n%s", out) + } + if !strings.Contains(out, tc.wantFailure) { + t.Fatalf("expected %q in the client output, got:\n%s", tc.wantFailure, out) + } + return + } + + if err != nil { + t.Fatalf("clickhouse-client failed: %v\n%s", err, out) + } + for _, want := range tc.wantOutput { + if !strings.Contains(out, want) { + t.Fatalf("expected %q in the client output, got:\n%s", want, out) + } + } + }) + } +} + +// TestNativeRecordsEveryStatement proves the packet loop keeps inspecting after the first statement, which a +// handler that degrades into a raw relay would silently stop doing. +func TestNativeRecordsEveryStatement(t *testing.T) { + if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { + t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") + } + + recorder := &recordingLogger{} + addr := startNativeProxy(t, nil, recorder) + + if out, err := runClient(t, addr, "SELECT 1;\nSELECT 2;\nSELECT 3;"); err != nil { + t.Fatalf("clickhouse-client failed: %v\n%s", err, out) + } + + for _, want := range []string{"SELECT 1", "SELECT 2", "SELECT 3"} { + if !recorder.contains(want) { + t.Fatalf("expected %q in the session recording, got:\n%s", want, recorder.dump()) + } + } + + if !strings.Contains(recorder.dump(), "=> OK") { + t.Fatalf("expected the outcome of each statement to be recorded, got:\n%s", recorder.dump()) + } +} + +// TestNativeRecordsFailedStatement proves a statement ClickHouse rejects is recorded with its error rather +// than as a success. +func TestNativeRecordsFailedStatement(t *testing.T) { + if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { + t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") + } + + recorder := &recordingLogger{} + addr := startNativeProxy(t, nil, recorder) + + if _, err := runClient(t, addr, "SELECT * FROM does_not_exist;"); err == nil { + t.Fatal("expected the statement to fail") + } + + waitFor(t, func() bool { return strings.Contains(recorder.dump(), "ERROR:") }) + + if !strings.Contains(recorder.dump(), "does_not_exist") { + t.Fatalf("expected the failed statement in the recording, got:\n%s", recorder.dump()) + } +} + +// A column type ch-go cannot infer costs the outcome of that statement, never the statement or the session. +func TestNativeDegradesOnUnreadableResultBlock(t *testing.T) { + if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { + t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") + } + + recorder := &recordingLogger{} + addr := startNativeProxy(t, []string{`(?i)\bdrop\b`}, recorder) + + out, err := runClient(t, addr, + "SELECT map('a', 1::UInt64) AS m;\nSELECT 'after-the-map';\nDROP TABLE pam_write_test;") + t.Logf("client output:\n%s", out) + + if err == nil { + t.Fatalf("expected the blocked statement to fail the client:\n%s", out) + } + if !strings.Contains(out, "{'a':1}") { + t.Fatalf("expected the unreadable column type to still reach the client:\n%s", out) + } + if !strings.Contains(out, "after-the-map") { + t.Fatalf("expected the session to survive the unreadable block:\n%s", out) + } + // The security control has to keep working after the recorder degrades. + if !strings.Contains(out, "blocked by the command blocking policy") { + t.Fatalf("expected blocking to still apply after degrading:\n%s", out) + } + if !recorder.contains("after-the-map") { + t.Fatalf("expected statements to still be recorded after degrading, got:\n%s", recorder.dump()) + } + if !strings.Contains(recorder.dump(), "outcome could not be read") { + t.Fatalf("expected the recording to say outcomes stopped, got:\n%s", recorder.dump()) + } +} + +func waitFor(t *testing.T, condition func() bool) { + t.Helper() + for range 100 { + if condition() { + return + } + time.Sleep(20 * time.Millisecond) + } + t.Fatal("timed out waiting for the session recording") +} + +var insertMarker = fmt.Sprintf("native-it-%d", time.Now().UnixNano()) + +func insertMarkerSQL() string { + return fmt.Sprintf( + "INSERT INTO pam_write_test (id, note) VALUES (4242, '%s');\nSELECT note FROM pam_write_test WHERE note = '%s';", + insertMarker, insertMarker) +} + +func compileForTest(t *testing.T, blocked []string) []*regexp.Regexp { + t.Helper() + patterns := make([]*regexp.Regexp, 0, len(blocked)) + for _, p := range blocked { + patterns = append(patterns, regexp.MustCompile(p)) + } + return patterns +} + +func startNativeProxy(t *testing.T, blocked []string, logger *recordingLogger) string { + t.Helper() + + patterns := compileForTest(t, blocked) + + proxy := NewClickHouseProxy(ClickHouseProxyConfig{ + TargetAddr: envOr("PAM_CLICKHOUSE_HTTP", "127.0.0.1:8123"), + NativeAddr: envOr("PAM_CLICKHOUSE_NATIVE", "127.0.0.1:9000"), + Username: envOr("PAM_CLICKHOUSE_USER", "default"), + Password: envOr("PAM_CLICKHOUSE_PASSWORD", "clickhouse"), + Database: envOr("PAM_CLICKHOUSE_DB", "analytics"), + SessionID: "native-integration-test", + SessionLogger: logger, + BlockedCommands: patterns, + }) + + // The client runs in a container, so the listener has to be reachable from outside the loopback. + listener, err := net.Listen("tcp", "0.0.0.0:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + t.Cleanup(func() { listener.Close() }) + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + go func() { + for { + conn, err := listener.Accept() + if err != nil { + return + } + go func() { _ = proxy.HandleConnection(ctx, conn) }() + } + }() + + return fmt.Sprintf("%d", listener.Addr().(*net.TCPAddr).Port) +} + +func envOr(name string, fallback string) string { + if v := os.Getenv(name); v != "" { + return v + } + return fallback +} + +func TestNativeConnectionTest(t *testing.T) { + if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { + t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") + } + + native := envOr("PAM_CLICKHOUSE_NATIVE", "127.0.0.1:9000") + + cases := []struct { + name string + addr string + username string + password string + wantErr string + }{ + {name: "valid account", addr: native, username: "default", password: "clickhouse"}, + { + name: "wrong password is an auth failure, not a timeout", + addr: native, username: "default", password: "wrong", + wantErr: "clickhouse rejected the connection", + }, + { + name: "unknown user is reported as ClickHouse reported it", + addr: native, username: "nobody", password: "x", + wantErr: "clickhouse rejected the connection", + }, + { + name: "a port with nothing on it fails to dial", + addr: "127.0.0.1:1", username: "default", password: "clickhouse", + wantErr: "connect", + }, + { + // The HTTP port answers, so this proves the check is a real handshake rather than a dial. + name: "pointing the native check at the HTTP port fails", + addr: envOr("PAM_CLICKHOUSE_HTTP", "127.0.0.1:8123"), username: "default", password: "clickhouse", + wantErr: "", + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + err := TestNativeConnection(context.Background(), ClickHouseProxyConfig{ + NativeAddr: tc.addr, + Username: tc.username, + Password: tc.password, + Database: envOr("PAM_CLICKHOUSE_DB", "analytics"), + }) + + if tc.name == "pointing the native check at the HTTP port fails" { + if err == nil { + t.Fatal("expected the HTTP port to fail a native handshake") + } + t.Logf("got: %v", err) + return + } + + if tc.wantErr == "" { + if err != nil { + t.Fatalf("expected success, got %v", err) + } + return + } + if err == nil { + t.Fatalf("expected an error containing %q, got success", tc.wantErr) + } + if !strings.Contains(strings.ToLower(err.Error()), tc.wantErr) { + t.Fatalf("expected %q in %v", tc.wantErr, err) + } + }) + } +} diff --git a/packages/pam/handlers/clickhouse/native_outcome.go b/packages/pam/handlers/clickhouse/native_outcome.go new file mode 100644 index 000000000..d9fa2ce48 --- /dev/null +++ b/packages/pam/handlers/clickhouse/native_outcome.go @@ -0,0 +1,106 @@ +package clickhouse + +import ( + "fmt" + "strings" + "sync" + "time" +) + +// outcomeRecorder pairs a statement with how it ended. The two directions of a native session are read by +// separate goroutines, and ClickHouse answers statements in order, so the queue is what joins them back up. +// +// Reading the server direction is best effort: a block it cannot decode costs the outcome, never the statement. +// Once that happens the recorder degrades and every later statement is written as soon as it is sent. +type outcomeRecorder struct { + proxy *ClickHouseProxy + + mu sync.Mutex + pending []pendingStatement + degraded bool +} + +type pendingStatement struct { + statement string + started time.Time + rows uint64 + bytes uint64 +} + +func newOutcomeRecorder(proxy *ClickHouseProxy) *outcomeRecorder { + return &outcomeRecorder{proxy: proxy} +} + +func (r *outcomeRecorder) begin(statement string) { + r.mu.Lock() + if r.degraded { + r.mu.Unlock() + r.proxy.logStatement(statement, "SENT") + return + } + r.pending = append(r.pending, pendingStatement{statement: statement, started: time.Now()}) + r.mu.Unlock() +} + +// progress folds ClickHouse's running counters into the statement in flight. +func (r *outcomeRecorder) progress(rows uint64, bytes uint64) { + r.mu.Lock() + defer r.mu.Unlock() + if len(r.pending) == 0 { + return + } + r.pending[0].rows += rows + r.pending[0].bytes += bytes +} + +func (r *outcomeRecorder) complete(outcome string) { + r.mu.Lock() + if len(r.pending) == 0 { + r.mu.Unlock() + return + } + next := r.pending[0] + r.pending = r.pending[1:] + r.mu.Unlock() + + r.proxy.logStatement(next.statement, next.describe(outcome)) +} + +func (p pendingStatement) describe(outcome string) string { + parts := []string{outcome} + if p.rows > 0 { + parts = append(parts, fmt.Sprintf("%d row(s) read", p.rows)) + } + return strings.Join(append(parts, fmt.Sprintf("%dms", time.Since(p.started).Milliseconds())), ", ") +} + +// degrade stops pairing outcomes for the rest of the session and says so in the recording, so a log that +// carries outcomes for some statements and not others is never read as if the rest simply did nothing. +func (r *outcomeRecorder) degrade(reason string) { + r.mu.Lock() + if r.degraded { + r.mu.Unlock() + return + } + r.degraded = true + drained := r.pending + r.pending = nil + r.mu.Unlock() + + note := fmt.Sprintf("SENT: the outcome could not be read (%s), so the rest of this session records "+ + "statements without one", reason) + for _, statement := range drained { + r.proxy.logStatement(statement.statement, note) + } +} + +func (r *outcomeRecorder) finish() { + r.mu.Lock() + drained := r.pending + r.pending = nil + r.mu.Unlock() + + for _, statement := range drained { + r.proxy.logStatement(statement.statement, statement.describe("INTERRUPTED: the session ended first")) + } +} diff --git a/packages/pam/handlers/clickhouse/native_outcome_test.go b/packages/pam/handlers/clickhouse/native_outcome_test.go new file mode 100644 index 000000000..4117db18d --- /dev/null +++ b/packages/pam/handlers/clickhouse/native_outcome_test.go @@ -0,0 +1,83 @@ +package clickhouse + +import ( + "strings" + "testing" + + "github.com/ClickHouse/ch-go/proto" + + "github.com/stretchr/testify/require" +) + +// The recorder decides what an auditor reads, and it joins two goroutines, so its queue behaviour is worth +// pinning down without needing a database. +func newTestRecorder() (*outcomeRecorder, *recordingLogger) { + logger := &recordingLogger{} + return newOutcomeRecorder(&ClickHouseProxy{config: ClickHouseProxyConfig{SessionLogger: logger}}), logger +} + +func TestOutcomeRecorderPairsStatementsInOrder(t *testing.T) { + recorder, logger := newTestRecorder() + + recorder.begin("SELECT 1") + recorder.begin("SELECT 2") + recorder.progress(7, 70) + recorder.complete("OK") + recorder.complete("OK") + + dump := logger.dump() + require.Contains(t, dump, "SELECT 1 => OK, 7 row(s) read") + require.Contains(t, dump, "SELECT 2 => OK") + // The progress landed before either completed, so it belongs to the first statement only. + require.Equal(t, 1, strings.Count(dump, "row(s) read")) +} + +func TestOutcomeRecorderMarksWhatTheSessionCutShort(t *testing.T) { + recorder, logger := newTestRecorder() + + recorder.begin("SELECT 1") + recorder.begin("SELECT 2") + recorder.complete("OK") + recorder.finish() + + dump := logger.dump() + require.Contains(t, dump, "SELECT 1 => OK") + require.Contains(t, dump, "SELECT 2 => INTERRUPTED: the session ended first") +} + +func TestOutcomeRecorderDegradesLoudly(t *testing.T) { + recorder, logger := newTestRecorder() + + recorder.begin("SELECT before") + recorder.degrade("a column type it could not read") + + // Anything still queued has to say why it has no outcome, rather than look like it did nothing. + require.Contains(t, logger.dump(), "SELECT before => SENT: the outcome could not be read") + + recorder.begin("SELECT after") + require.Contains(t, logger.dump(), "SELECT after => SENT") + + // Degrading twice must not double-log, and completing afterwards must not resurrect pairing. + recorder.degrade("again") + recorder.complete("OK") + require.Equal(t, 1, strings.Count(logger.dump(), "outcome could not be read")) + require.NotContains(t, logger.dump(), "=> OK") +} + +func TestOutcomeRecorderIgnoresAnUnmatchedCompletion(t *testing.T) { + recorder, logger := newTestRecorder() + + // ClickHouse sends packets that are not tied to a statement we queued; they must not panic or invent one. + recorder.complete("OK") + recorder.progress(5, 5) + recorder.finish() + + require.Empty(t, strings.TrimSpace(logger.dump())) +} + +func TestNativeParameterSuffix(t *testing.T) { + require.Empty(t, nativeParameterSuffix(nil)) + require.Equal(t, "\n-- parameters: a=1", nativeParameterSuffix([]proto.Parameter{{Key: "a", Value: "1"}})) + require.Equal(t, "\n-- parameters: a=1 b=two", + nativeParameterSuffix([]proto.Parameter{{Key: "a", Value: "1"}, {Key: "b", Value: "two"}})) +} diff --git a/packages/pam/handlers/clickhouse/native_unit_test.go b/packages/pam/handlers/clickhouse/native_unit_test.go new file mode 100644 index 000000000..797d7dee1 --- /dev/null +++ b/packages/pam/handlers/clickhouse/native_unit_test.go @@ -0,0 +1,358 @@ +package clickhouse + +import ( + "context" + "net" + "regexp" + "sync" + "testing" + "time" + + "github.com/ClickHouse/ch-go/proto" + "github.com/stretchr/testify/require" +) + +// fakeClickHouse stands in for a server so the security-critical parts of the handshake and packet loop can +// be tested without docker: what the gateway sends upstream is recorded, and nothing needs a real database. +type fakeClickHouse struct { + listener net.Listener + + mu sync.Mutex + hello proto.ClientHello + quotaKey string + queries []proto.Query + // Bytes seen after the handshake, which is what proves nothing was relayed once a refusal happened. + bytesAfterHandshake int + done chan struct{} +} + +func startFakeClickHouse(t *testing.T) *fakeClickHouse { + t.Helper() + + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + + f := &fakeClickHouse{listener: listener, done: make(chan struct{})} + t.Cleanup(func() { listener.Close() }) + + go func() { + defer close(f.done) + conn, acceptErr := listener.Accept() + if acceptErr != nil { + return + } + defer conn.Close() + f.serve(conn) + }() + + return f +} + +func (f *fakeClickHouse) addr() string { return f.listener.Addr().String() } + +func (f *fakeClickHouse) serve(conn net.Conn) { + r := proto.NewReader(newTap(conn)) + + code, err := r.UVarInt() + if err != nil || proto.ClientCode(code) != proto.ClientCodeHello { + return + } + + var hello proto.ClientHello + if err := hello.Decode(r); err != nil { + return + } + + f.mu.Lock() + f.hello = hello + f.mu.Unlock() + + rev := min(hello.ProtocolVersion, proto.Version) + + var b proto.Buffer + serverHello := proto.ServerHello{Name: "FakeClickHouse", Major: 24, Minor: 8, Revision: rev} + serverHello.EncodeAware(&b, rev) + if _, err := conn.Write(b.Buf); err != nil { + return + } + + if proto.FeatureAddendum.In(rev) { + quotaKey, err := r.Str() + if err != nil { + return + } + f.mu.Lock() + f.quotaKey = quotaKey + f.mu.Unlock() + } + + for { + packet, err := r.UVarInt() + if err != nil { + return + } + f.mu.Lock() + f.bytesAfterHandshake++ + f.mu.Unlock() + + if proto.ClientCode(packet) != proto.ClientCodeQuery { + continue + } + var q proto.Query + if err := q.DecodeAware(r, rev); err != nil { + return + } + f.mu.Lock() + f.queries = append(f.queries, q) + f.mu.Unlock() + } +} + +func (f *fakeClickHouse) snapshot() (proto.ClientHello, string, []proto.Query, int) { + f.mu.Lock() + defer f.mu.Unlock() + return f.hello, f.quotaKey, append([]proto.Query(nil), f.queries...), f.bytesAfterHandshake +} + +// dialProxy runs one session against a proxy configured to reach the fake server. +func dialProxy(t *testing.T, config ClickHouseProxyConfig) net.Conn { + t.Helper() + + proxy := NewClickHouseProxy(config) + client, server := net.Pipe() + + ctx, cancel := context.WithCancel(context.Background()) + handled := make(chan struct{}) + + // Cleanup runs last-registered-first, so this waits only after the conn is closed and ctx cancelled. + t.Cleanup(func() { + select { + case <-handled: + case <-time.After(5 * time.Second): + t.Error("the handler did not finish") + } + }) + t.Cleanup(cancel) + t.Cleanup(func() { client.Close() }) + + go func() { + defer close(handled) + _ = proxy.HandleConnection(ctx, server) + }() + + require.NoError(t, client.SetDeadline(time.Now().Add(10*time.Second))) + return client +} + +func clientHandshake(t *testing.T, conn net.Conn, user, password string) *proto.Reader { + t.Helper() + + var b proto.Buffer + proto.ClientHello{ + Name: "unit-test client", + Major: 24, + Minor: 8, + ProtocolVersion: proto.Version, + Database: "whatever-the-client-wants", + User: user, + Password: password, + }.Encode(&b) + _, err := conn.Write(b.Buf) + require.NoError(t, err) + + r := proto.NewReader(newTap(conn)) + code, err := r.UVarInt() + require.NoError(t, err) + require.Equal(t, proto.ServerCodeHello, proto.ServerCode(code)) + + var serverHello proto.ServerHello + require.NoError(t, serverHello.DecodeAware(r, proto.Version)) + + if proto.FeatureAddendum.In(proto.Version) { + b.Reset() + b.PutString("the-client-quota-key") + _, err = conn.Write(b.Buf) + require.NoError(t, err) + } + return r +} + +// The whole point of the proxy: what the client presents is dropped and the account's own identity is used. +func TestNativeHandshakeInjectsAccountCredentials(t *testing.T) { + upstream := startFakeClickHouse(t) + + conn := dialProxy(t, ClickHouseProxyConfig{ + NativeAddr: upstream.addr(), + Username: "the-account", + Password: "the-account-password", + Database: "the-account-database", + SessionID: "unit", + }) + clientHandshake(t, conn, "attacker", "attacker-password") + + require.Eventually(t, func() bool { + hello, _, _, _ := upstream.snapshot() + return hello.User != "" + }, 5*time.Second, 20*time.Millisecond) + + hello, quotaKey, _, _ := upstream.snapshot() + require.Equal(t, "the-account", hello.User) + require.Equal(t, "the-account-password", hello.Password) + require.Equal(t, "the-account-database", hello.Database) + require.NotEqual(t, "attacker", hello.User) + require.Empty(t, quotaKey, "the client's quota key is not the account's to choose") +} + +// A packet the loop cannot read must end the session, and nothing may reach the server after it. +func TestNativeUnreadablePacketFailsClosed(t *testing.T) { + upstream := startFakeClickHouse(t) + + conn := dialProxy(t, ClickHouseProxyConfig{ + NativeAddr: upstream.addr(), + Username: "account", + SessionID: "unit", + }) + reader := clientHandshake(t, conn, "someone", "") + + var b proto.Buffer + b.PutUVarInt(99) + _, err := conn.Write(b.Buf) + require.NoError(t, err) + + code, message := decodeException(t, reader) + require.Equal(t, codeNotImplemented, code) + require.Contains(t, message, "could not read ClickHouse client packet 99") + + _, _, queries, seen := upstream.snapshot() + require.Empty(t, queries) + require.Zero(t, seen, "no packet may reach ClickHouse after a refusal") +} + +// A blocked statement must be refused before it is forwarded, and must end the session. +func TestNativeBlockedStatementNeverReachesUpstream(t *testing.T) { + upstream := startFakeClickHouse(t) + recorder := &recordingLogger{} + + conn := dialProxy(t, ClickHouseProxyConfig{ + NativeAddr: upstream.addr(), + Username: "account", + SessionID: "unit", + SessionLogger: recorder, + BlockedCommands: []*regexp.Regexp{regexp.MustCompile(`(?i)\bdrop\b`)}, + }) + reader := clientHandshake(t, conn, "someone", "") + + writeQuery(t, conn, proto.Query{Body: "DROP TABLE important"}) + + code, message := decodeException(t, reader) + require.Equal(t, codeAccessDenied, code) + require.Contains(t, message, "blocked by the command blocking policy") + + _, _, queries, _ := upstream.snapshot() + require.Empty(t, queries, "a blocked statement must not reach ClickHouse") + require.True(t, recorder.contains("DROP TABLE important")) + require.Contains(t, recorder.dump(), "BLOCKED") +} + +// The client's quota key rides on the Query packet as well as the addendum, and both are the account's. +func TestNativeStripsTheClientQuotaKeyFromTheQuery(t *testing.T) { + upstream := startFakeClickHouse(t) + + conn := dialProxy(t, ClickHouseProxyConfig{ + NativeAddr: upstream.addr(), + Username: "account", + SessionID: "unit", + }) + clientHandshake(t, conn, "someone", "") + + writeQuery(t, conn, proto.Query{ + Body: "SELECT 1", + Info: proto.ClientInfo{QuotaKey: "quota-the-client-picked"}, + }) + + require.Eventually(t, func() bool { + _, _, queries, _ := upstream.snapshot() + return len(queries) == 1 + }, 5*time.Second, 20*time.Millisecond) + + _, _, queries, _ := upstream.snapshot() + require.Equal(t, "SELECT 1", queries[0].Body) + require.Empty(t, queries[0].Info.QuotaKey) +} + +// The revision is pinned to what ch-go can parse, and the client has to be told the pinned one so it +// encodes to match. +func TestNativeHandshakePinsTheRevision(t *testing.T) { + upstream := startFakeClickHouse(t) + + conn := dialProxy(t, ClickHouseProxyConfig{ + NativeAddr: upstream.addr(), + Username: "account", + SessionID: "unit", + }) + + var b proto.Buffer + proto.ClientHello{ + Name: "a client newer than ch-go", + Major: 99, + Minor: 9, + ProtocolVersion: proto.Version + 500, + Database: "db", + User: "someone", + }.Encode(&b) + _, err := conn.Write(b.Buf) + require.NoError(t, err) + + r := proto.NewReader(newTap(conn)) + code, err := r.UVarInt() + require.NoError(t, err) + require.Equal(t, proto.ServerCodeHello, proto.ServerCode(code)) + + var serverHello proto.ServerHello + require.NoError(t, serverHello.DecodeAware(r, proto.Version)) + require.Equal(t, maxNativeRevision, serverHello.Revision, + "the client must be told the pinned revision, not the server's") + + require.Eventually(t, func() bool { + hello, _, _, _ := upstream.snapshot() + return hello.ProtocolVersion != 0 + }, 5*time.Second, 20*time.Millisecond) + + hello, _, _, _ := upstream.snapshot() + require.Equal(t, maxNativeRevision, hello.ProtocolVersion) +} + +func writeQuery(t *testing.T, conn net.Conn, q proto.Query) { + t.Helper() + + q.Info.ProtocolVersion = proto.Version + q.Info.Major, q.Info.Minor = 24, 8 + q.Info.Interface = proto.InterfaceTCP + q.Info.Query = proto.ClientQueryInitial + q.Info.InitialAddress = "127.0.0.1:0" + q.Stage = proto.StageComplete + + var b proto.Buffer + q.EncodeAware(&b, proto.Version) + _, err := conn.Write(b.Buf) + require.NoError(t, err) +} + +func decodeException(t *testing.T, r *proto.Reader) (int, string) { + t.Helper() + + code, err := r.UVarInt() + require.NoError(t, err) + require.Equal(t, proto.ServerCodeException, proto.ServerCode(code)) + + var e proto.Exception + require.NoError(t, e.DecodeAware(r, proto.Version)) + return int(e.Code), e.Message +} + +// ch-go trailing the server is the reason the handshake pins a revision at all. If ch-go ever catches up, +// the pinning becomes a no-op and this is the reminder to re-check it. +func TestChGoRevisionIsStillBehindTheServers(t *testing.T) { + require.LessOrEqual(t, maxNativeRevision, 54469, + "ch-go has caught up with ClickHouse; revision pinning needs revisiting") +} diff --git a/packages/pam/handlers/clickhouse/proxy.go b/packages/pam/handlers/clickhouse/proxy.go index 60cf4e85a..b0f3b7039 100644 --- a/packages/pam/handlers/clickhouse/proxy.go +++ b/packages/pam/handlers/clickhouse/proxy.go @@ -27,11 +27,13 @@ import ( "github.com/rs/zerolog/log" ) -// Brokered over ClickHouse's HTTP interface rather than the native protocol on 9000, where a statement is -// text and can be blocked and recorded without parsing a version-gated binary format. The client's own -// credentials are dropped and the account's injected, so nothing it holds works outside a recorded session. +// Brokers both of ClickHouse's interfaces: HTTP on TargetAddr and the native TCP protocol on NativeAddr. A +// session listens on one local port and routes by the first byte the client sends, so the driver decides the +// protocol rather than the user. The client's own credentials are dropped and the account's injected on either +// path, so nothing it holds works outside a recorded session. type ClickHouseProxyConfig struct { TargetAddr string + NativeAddr string Username string Password string Database string @@ -59,12 +61,14 @@ const ( const ( codeNotImplemented = 48 codeNetworkError = 210 + codeTooManyRows = 396 codeAccessDenied = 497 ) var errorNames = map[int]string{ codeNotImplemented: "NOT_IMPLEMENTED", codeNetworkError: "NETWORK_ERROR", + codeTooManyRows: "TOO_MANY_ROWS", codeAccessDenied: "ACCESS_DENIED", } @@ -82,7 +86,9 @@ var strippedAuthParams = []string{"user", "password"} var strippedExecutionParams = []string{"role", "quota_key"} -var allowedPaths = map[string]bool{"/": true, "/ping": true} +const pingPath = "/ping" + +var allowedPaths = map[string]bool{"/": true, pingPath: true} type ClickHouseProxy struct { config ClickHouseProxyConfig @@ -93,6 +99,10 @@ type stateKey struct{} type requestState struct { statement string + // The statement without the recorded parameter suffix, and whether it was cut short by the inspection + // window. The bridge runs this rather than re-reading a body that may be compressed. + sql string + truncated bool started time.Time } @@ -134,12 +144,39 @@ func (p *ClickHouseProxy) HandleConnection(ctx context.Context, clientConn net.C l := log.With().Str("sessionId", p.config.SessionID).Str("resourceType", "clickhouse").Logger() + if p.config.TargetAddr == "" && p.config.NativeAddr == "" { + l.Error().Msg("Refused a ClickHouse session with neither a HTTP nor a native port") + return nil + } + + conn, isNative, err := sniffProtocol(clientConn) + if err != nil { + if errors.Is(err, io.EOF) { + l.Debug().Msg("Client closed before sending anything") + } else { + l.Warn().Err(err).Msg("Could not read the first byte of a ClickHouse connection") + } + return nil + } + + if isNative { + if p.config.NativeAddr == "" { + l.Info().Msg("Refused a native connection on an account with no native port") + return writeNativeError(conn, maxNativeRevision, codeNotImplemented, + "This account does not have ClickHouse's native port configured, so only the HTTP interface is "+ + "available in this session.") + } + return newNativeProxy(p).HandleConnection(ctx, conn, l.With().Str("protocol", "native").Logger()) + } + + l = l.With().Str("protocol", "http").Logger() + server := &http.Server{ Handler: p.handler(l), ReadHeaderTimeout: 30 * time.Second, } - listener := newSingleConnListener(clientConn) + listener := newSingleConnListener(conn) done := make(chan struct{}) defer close(done) @@ -167,11 +204,12 @@ func (p *ClickHouseProxy) handler(l zerolog.Logger) http.Handler { return } - statement, body, err := p.inspect(r) + inspected, body, err := p.inspect(r) if err != nil { writeClickHouseError(w, http.StatusBadRequest, codeNotImplemented, err.Error()) return } + statement := inspected.statement if blocked := p.blockedBy(statement); blocked != nil { p.logStatement(statement, fmt.Sprintf("BLOCKED: %s", blocked.String())) @@ -182,7 +220,19 @@ func (p *ClickHouseProxy) handler(l zerolog.Logger) http.Handler { } r.Body = body - state := &requestState{statement: statement, started: time.Now()} + state := &requestState{ + statement: statement, + sql: inspected.sql, + truncated: inspected.truncated, + started: time.Now(), + } + + // A server with HTTP disabled still has to serve Web Access, which only speaks HTTP. + if p.config.TargetAddr == "" { + p.serveBridge(w, r, state, l) + return + } + p.reverse.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), stateKey{}, state))) }) } @@ -192,17 +242,26 @@ type bodyReadCloser struct { io.Closer } +type inspectedRequest struct { + statement string + sql string + truncated bool +} + // Returns the statement and a body that still replays in full. ClickHouse concatenates `query` and the body. -func (p *ClickHouseProxy) inspect(r *http.Request) (string, io.ReadCloser, error) { +func (p *ClickHouseProxy) inspect(r *http.Request) (inspectedRequest, io.ReadCloser, error) { queryParam := strings.TrimSpace(r.URL.Query().Get("query")) if r.Body == nil || r.ContentLength == 0 { - return queryParam + parameterSuffix(r.URL.Query()), http.NoBody, nil + return inspectedRequest{ + statement: queryParam + parameterSuffix(r.URL.Query()), + sql: queryParam, + }, http.NoBody, nil } // ClickHouse's own block compression is opaque to anything but a ClickHouse client if r.URL.Query().Get("decompress") == "1" { - return "", nil, fmt.Errorf( + return inspectedRequest{}, nil, fmt.Errorf( "this session cannot read a ClickHouse-compressed request body, so decompress=1 is not supported here. " + "Send the statement uncompressed or with Content-Encoding: gzip") } @@ -211,7 +270,7 @@ func (p *ClickHouseProxy) inspect(r *http.Request) (string, io.ReadCloser, error switch encoding { case "", "identity", "gzip", "deflate": default: - return "", nil, fmt.Errorf( + return inspectedRequest{}, nil, fmt.Errorf( "this session cannot read a %q-encoded request body, so the command blocking policy could not be applied to it. "+ "Use gzip, deflate, or no compression", encoding) } @@ -220,7 +279,7 @@ func (p *ClickHouseProxy) inspect(r *http.Request) (string, io.ReadCloser, error head := make([]byte, maxInspectBytes+1) n, err := io.ReadFull(r.Body, head) if err != nil && err != io.EOF && err != io.ErrUnexpectedEOF { - return "", nil, fmt.Errorf("the gateway could not read the request body: %v", err) + return inspectedRequest{}, nil, fmt.Errorf("the gateway could not read the request body: %v", err) } head = head[:n] @@ -228,18 +287,24 @@ func (p *ClickHouseProxy) inspect(r *http.Request) (string, io.ReadCloser, error decoded, decodedOverflow, decodeErr := decodeHead(head, encoding) if decodeErr != nil { - return "", nil, fmt.Errorf( + return inspectedRequest{}, nil, fmt.Errorf( "the gateway could not decompress the request body to apply the command blocking policy: %v", decodeErr) } - if (len(head) > maxInspectBytes || decodedOverflow) && len(p.config.BlockedCommands) > 0 { - return "", nil, fmt.Errorf( + truncated := len(head) > maxInspectBytes || decodedOverflow + if truncated && len(p.config.BlockedCommands) > 0 { + return inspectedRequest{}, nil, fmt.Errorf( "this account blocks commands, so a request body larger than %d MB is refused: the gateway has to read "+ "the whole statement to apply the policy. Send the data in smaller batches", maxInspectBytes>>20) } - return joinStatement(queryParam, string(decoded)) + parameterSuffix(r.URL.Query()), forwarded, nil + sql := joinStatement(queryParam, string(decoded)) + return inspectedRequest{ + statement: sql + parameterSuffix(r.URL.Query()), + sql: sql, + truncated: truncated, + }, forwarded, nil } func parameterSuffix(query url.Values) string { @@ -321,7 +386,11 @@ func (p *ClickHouseProxy) rewrite(pr *httputil.ProxyRequest) { for _, param := range append(append([]string{}, strippedAuthParams...), strippedExecutionParams...) { query.Del(param) } - if p.config.Database != "" { + // ClickHouse's health endpoint refuses any query string, so a database parameter turns it into a 404. + if req.URL.Path == pingPath { + query = nil + req.URL.ForceQuery = false + } else if p.config.Database != "" { query.Set("database", p.config.Database) } req.URL.RawQuery = query.Encode() diff --git a/packages/pam/handlers/clickhouse/proxy_test.go b/packages/pam/handlers/clickhouse/proxy_test.go index e1d6f945a..45f7ab965 100644 --- a/packages/pam/handlers/clickhouse/proxy_test.go +++ b/packages/pam/handlers/clickhouse/proxy_test.go @@ -9,6 +9,7 @@ import ( "net/url" "regexp" "strings" + "sync" "testing" "github.com/Infisical/infisical-merge/packages/pam/session" @@ -17,13 +18,37 @@ import ( ) type recordingLogger struct { + mu sync.Mutex entries []session.SessionLogEntry } func (r *recordingLogger) LogEntry(entry session.SessionLogEntry) error { + r.mu.Lock() + defer r.mu.Unlock() r.entries = append(r.entries, entry) return nil } + +func (r *recordingLogger) contains(want string) bool { + r.mu.Lock() + defer r.mu.Unlock() + for _, entry := range r.entries { + if strings.Contains(entry.Input, want) { + return true + } + } + return false +} + +func (r *recordingLogger) dump() string { + r.mu.Lock() + defer r.mu.Unlock() + var out strings.Builder + for _, entry := range r.entries { + out.WriteString(entry.Input + " => " + entry.Output + "\n") + } + return out.String() +} func (r *recordingLogger) LogSessionEvent(session.SessionEvent) error { return nil } func (r *recordingLogger) LogHttpEvent(session.HttpEvent) error { return nil } func (r *recordingLogger) Close() error { return nil } diff --git a/packages/pam/handlers/clickhouse/sniff.go b/packages/pam/handlers/clickhouse/sniff.go new file mode 100644 index 000000000..ddff3fb61 --- /dev/null +++ b/packages/pam/handlers/clickhouse/sniff.go @@ -0,0 +1,53 @@ +package clickhouse + +import ( + "bufio" + "net" + "time" +) + +// ClickHouse serves two interfaces on two ports, and which one a client speaks depends on the driver rather than +// on anything the user chose: clickhouse-client is native-only, the JDBC driver is HTTP-only. A session hands out +// one local port and reads the first byte to tell them apart, so nobody has to pick a protocol. +// +// A native session opens with the Hello packet code, a uvarint 0. Every HTTP request opens with the ASCII letter +// of its method, so the two can never be confused. +const nativeHelloByte = 0x00 + +// A client that connects and then says nothing would otherwise park the session handler forever: the peek +// happens before any HTTP server exists, so ReadHeaderTimeout does not cover it. +const sniffTimeout = 30 * time.Second + +// peekConn replays the sniffed byte to whichever handler takes the connection. +type peekConn struct { + net.Conn + reader *bufio.Reader +} + +func (c *peekConn) Read(p []byte) (int, error) { + return c.reader.Read(p) +} + +// CloseWrite is promoted explicitly: embedding the net.Conn interface hides it, and net/http uses it to +// half-close rather than resetting a connection it is finished with. +func (c *peekConn) CloseWrite() error { + if cw, ok := c.Conn.(interface{ CloseWrite() error }); ok { + return cw.CloseWrite() + } + return nil +} + +// sniffProtocol reports whether the client opened a native session, and returns a connection that still replays +// the byte it read to decide. +func sniffProtocol(conn net.Conn) (net.Conn, bool, error) { + reader := bufio.NewReaderSize(conn, 64<<10) + + _ = conn.SetReadDeadline(time.Now().Add(sniffTimeout)) + first, err := reader.Peek(1) + _ = conn.SetReadDeadline(time.Time{}) + if err != nil { + return conn, false, err + } + + return &peekConn{Conn: conn, reader: reader}, first[0] == nativeHelloByte, nil +} diff --git a/packages/pam/handlers/clickhouse/sniff_test.go b/packages/pam/handlers/clickhouse/sniff_test.go new file mode 100644 index 000000000..4816e454f --- /dev/null +++ b/packages/pam/handlers/clickhouse/sniff_test.go @@ -0,0 +1,121 @@ +package clickhouse + +import ( + "bufio" + "context" + "fmt" + "net" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/ClickHouse/ch-go/proto" + "github.com/stretchr/testify/require" +) + +func TestSniffProtocol(t *testing.T) { + cases := []struct { + name string + first []byte + wantNative bool + }{ + {name: "native hello", first: []byte{0x00}, wantNative: true}, + {name: "http get", first: []byte("GET / HTTP/1.1\r\n"), wantNative: false}, + {name: "http post", first: []byte("POST / HTTP/1.1\r\n"), wantNative: false}, + {name: "http head", first: []byte("HEAD / HTTP/1.1\r\n"), wantNative: false}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + client, server := net.Pipe() + defer client.Close() + defer server.Close() + + go func() { _, _ = client.Write(tc.first) }() + + conn, isNative, err := sniffProtocol(server) + require.NoError(t, err) + require.Equal(t, tc.wantNative, isNative) + + // The byte used to decide has to still be readable by the handler that takes the connection. + buf := make([]byte, len(tc.first)) + _, err = bufio.NewReader(conn).Read(buf[:1]) + require.NoError(t, err) + require.Equal(t, tc.first[0], buf[0]) + }) + } +} + +// The HTTP interface has to keep working now that a connection is routed by its first byte. +func TestHandleConnectionRoutesHTTP(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = w.Write([]byte("ok")) + })) + defer upstream.Close() + + proxy := NewClickHouseProxy(ClickHouseProxyConfig{ + TargetAddr: strings.TrimPrefix(upstream.URL, "http://"), + Username: "account", + SessionID: "sniff-test", + }) + + client, server := net.Pipe() + defer client.Close() + + go func() { _ = proxy.HandleConnection(context.Background(), server) }() + + require.NoError(t, client.SetDeadline(time.Now().Add(10*time.Second))) + _, err := client.Write([]byte("POST / HTTP/1.1\r\nHost: x\r\nContent-Length: 8\r\n\r\nSELECT 1")) + require.NoError(t, err) + + resp, err := http.ReadResponse(bufio.NewReader(client), nil) + require.NoError(t, err) + defer resp.Body.Close() + require.Equal(t, http.StatusOK, resp.StatusCode) +} + +// A native client reaching an account with no native port has to get a native exception, not a dead socket. +func TestHandleConnectionRefusesNativeWithoutPort(t *testing.T) { + proxy := NewClickHouseProxy(ClickHouseProxyConfig{ + TargetAddr: "127.0.0.1:1", + Username: "account", + SessionID: "sniff-test", + }) + + client, server := net.Pipe() + defer client.Close() + + go func() { _ = proxy.HandleConnection(context.Background(), server) }() + + require.NoError(t, client.SetDeadline(time.Now().Add(10*time.Second))) + + var b proto.Buffer + proto.ClientHello{ + Name: "test", + ProtocolVersion: proto.Version, + Database: "default", + User: "someone", + }.Encode(&b) + _, err := client.Write(b.Buf) + require.NoError(t, err) + + code, message := readNativeException(t, client) + require.Equal(t, codeNotImplemented, code) + require.Contains(t, message, "native port") +} + +func readNativeException(t *testing.T, conn net.Conn) (int, string) { + t.Helper() + + r := proto.NewReader(newTap(conn)) + + code, err := r.UVarInt() + require.NoError(t, err) + require.Equal(t, proto.ServerCodeException, proto.ServerCode(code), fmt.Sprintf("unexpected packet %d", code)) + + var e proto.Exception + require.NoError(t, e.DecodeAware(r, proto.Version)) + return int(e.Code), e.Message +} diff --git a/packages/pam/handlers/clickhouse/tap_test.go b/packages/pam/handlers/clickhouse/tap_test.go new file mode 100644 index 000000000..fdd232a5b --- /dev/null +++ b/packages/pam/handlers/clickhouse/tap_test.go @@ -0,0 +1,51 @@ +package clickhouse + +import ( + "bytes" + "io" + "net" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +// The degrade path stops decoding and relays the rest of the stream. Anything the tap's buffered reader +// already pulled off the socket has to go with it, or the client gets a truncated packet. +func TestTapRelaysWhatItHasAlreadyBuffered(t *testing.T) { + upstreamRead, upstreamWrite := net.Pipe() + defer upstreamRead.Close() + defer upstreamWrite.Close() + + payload := bytes.Repeat([]byte("abcdefghij"), 20) // 200 bytes in one burst + go func() { + _, _ = upstreamWrite.Write(payload) + upstreamWrite.Close() + }() + + tp := newTap(upstreamRead) + + // Consume 10 bytes through the tap, as a decoder would before it fails. + consumed := make([]byte, 10) + _, err := io.ReadFull(tp, consumed) + require.NoError(t, err) + + var client bytes.Buffer + _, err = client.Write(tp.take()) + require.NoError(t, err) + + done := make(chan struct{}) + go func() { + defer close(done) + _, _ = io.Copy(&client, tp.rest()) + }() + select { + case <-done: + case <-time.After(5 * time.Second): + t.Fatal("relay did not finish") + } + + require.Equal(t, len(payload), client.Len(), + "every byte the upstream sent must reach the client, including what the tap buffered") + require.Equal(t, payload, client.Bytes()) +} diff --git a/packages/pam/handlers/clickhouse/testhelpers_test.go b/packages/pam/handlers/clickhouse/testhelpers_test.go new file mode 100644 index 000000000..ed4e295c4 --- /dev/null +++ b/packages/pam/handlers/clickhouse/testhelpers_test.go @@ -0,0 +1,97 @@ +package clickhouse + +import ( + "bytes" + "compress/gzip" + "context" + "io" + "net/http" + "net/url" + "os/exec" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +// runClient drives the real clickhouse-client against the session's port, which is the only way to cover +// revision pinning, the addendum and block framing as a real driver produces them. +func runClient(t *testing.T, port string, sql string, extra ...string) (string, error) { + t.Helper() + + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + + args := []string{ + "run", "--rm", "-i", "clickhouse/clickhouse-server:24.8", "clickhouse-client", + "--host", envOr("PAM_CLICKHOUSE_CLIENT_HOST", "host.docker.internal"), + "--port", port, + // Deliberately wrong: the gateway replaces them with the account's. + "--user", "not-the-account", "--password", "not-the-password", + "--multiquery", + } + args = append(args, extra...) + + cmd := exec.CommandContext(ctx, "docker", args...) + cmd.Stdin = strings.NewReader(sql) + + out, err := cmd.CombinedOutput() + return string(out), err +} + +// postStatementE is the non-asserting form. require.* calls t.FailNow, which is illegal off the test +// goroutine, so anything running in parallel has to report failures over a channel instead. +func postStatementE(addr string, sql string) (int, string, error) { + req, err := http.NewRequest(http.MethodPost, "http://"+addr+"/", strings.NewReader(sql)) + if err != nil { + return 0, "", err + } + resp, err := (&http.Client{Timeout: 60 * time.Second}).Do(req) + if err != nil { + return 0, "", err + } + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + if err != nil { + return resp.StatusCode, "", err + } + return resp.StatusCode, string(body), nil +} + +func postGET(t *testing.T, addr string, path string) (int, string) { + t.Helper() + + resp, err := (&http.Client{Timeout: 60 * time.Second}).Get("http://" + addr + path) + require.NoError(t, err) + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + return resp.StatusCode, string(body) +} + +func postGzipped(t *testing.T, addr string, sql string) (int, string) { + t.Helper() + + var buf bytes.Buffer + writer := gzip.NewWriter(&buf) + _, err := writer.Write([]byte(sql)) + require.NoError(t, err) + require.NoError(t, writer.Close()) + + req, err := http.NewRequest(http.MethodPost, "http://"+addr+"/", &buf) + require.NoError(t, err) + req.Header.Set("Content-Encoding", "gzip") + + resp, err := (&http.Client{Timeout: 60 * time.Second}).Do(req) + require.NoError(t, err) + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + return resp.StatusCode, string(body) +} + +func urlEscape(v string) string { return url.QueryEscape(v) } diff --git a/packages/pam/local/access.go b/packages/pam/local/access.go index f795c7b6c..18e2fe2e8 100644 --- a/packages/pam/local/access.go +++ b/packages/pam/local/access.go @@ -397,9 +397,9 @@ var accountDisplays = map[string]AccountConnectionDisplay{ TypeLabel: "ClickHouse", DefaultPort: 8123, Note: []string{ - "This is ClickHouse's HTTP interface, so a client that only speaks the", - "native protocol, clickhouse-client included, cannot use this port.", - "JDBC, clickhouse-connect and curl all can.", + "This port serves whichever of ClickHouse's two interfaces the account", + "has configured, detected per connection. JDBC, clickhouse-connect and", + "curl use the HTTP one; clickhouse-client needs a native port set.", }, ConnectionString: func(username, database string, port int) string { return fmt.Sprintf("jdbc:clickhouse://127.0.0.1:%d/%s", port, database) @@ -407,6 +407,8 @@ var accountDisplays = map[string]AccountConnectionDisplay{ UsageExamples: func(username, database string, port int) []string { return []string{ fmt.Sprintf("curl 'http://127.0.0.1:%d/?database=%s&query=SELECT+1'", port, url.QueryEscape(database)), + fmt.Sprintf("clickhouse-client --host 127.0.0.1 --port %d --database %s # needs a native port", + port, database), } }, }, diff --git a/packages/pam/pam-proxy.go b/packages/pam/pam-proxy.go index ace544671..a1f76af3f 100644 --- a/packages/pam/pam-proxy.go +++ b/packages/pam/pam-proxy.go @@ -10,6 +10,7 @@ import ( "net/url" "os" "regexp" + "strconv" "time" "github.com/Infisical/infisical-merge/packages/api" @@ -593,8 +594,22 @@ func HandlePAMProxy(ctx context.Context, conn *tls.Conn, pamConfig *GatewayPAMCo blockedCommands = compilePolicyPatterns(rulePatterns(credentials.PolicyRules.CommandBlocking), pamConfig.SessionId, "command-blocking") } + // Either interface can be absent: an empty address is what tells the handler that one is not served. + nativeAddr := "" + if credentials.NativePort > 0 { + nativeAddr = net.JoinHostPort(credentials.Host, strconv.Itoa(credentials.NativePort)) + } + httpAddr := "" + if credentials.Port > 0 { + httpAddr = net.JoinHostPort(credentials.Host, strconv.Itoa(credentials.Port)) + } + if httpAddr == "" && nativeAddr == "" { + return fmt.Errorf("clickhouse account has neither a HTTP port nor a native port configured") + } + proxy := clickhouse.NewClickHouseProxy(clickhouse.ClickHouseProxyConfig{ - TargetAddr: fmt.Sprintf("%s:%d", credentials.Host, credentials.Port), + TargetAddr: httpAddr, + NativeAddr: nativeAddr, Username: credentials.Username, Password: credentials.Password, Database: credentials.Database, @@ -606,7 +621,8 @@ func HandlePAMProxy(ctx context.Context, conn *tls.Conn, pamConfig *GatewayPAMCo }) log.Info(). Str("sessionId", pamConfig.SessionId). - Str("target", fmt.Sprintf("%s:%d", credentials.Host, credentials.Port)). + Str("target", httpAddr). + Str("nativeTarget", nativeAddr). Bool("sslEnabled", credentials.SSLEnabled). Msg("Starting ClickHouse PAM proxy") return proxy.HandleConnection(ctx, handlerConn) diff --git a/packages/pam/session/credentials.go b/packages/pam/session/credentials.go index 49544e51f..bd872ad74 100644 --- a/packages/pam/session/credentials.go +++ b/packages/pam/session/credentials.go @@ -31,6 +31,7 @@ type PAMCredentials struct { Certificate string Host string Port int + NativePort int SSLEnabled bool SSLRejectUnauthorized bool SSLCertificate string @@ -195,6 +196,7 @@ func (cm *CredentialsManager) GetPAMSessionCredentials(sessionId string, expiryT Certificate: response.Credentials.Certificate, Host: response.Credentials.Host, Port: response.Credentials.Port, + NativePort: response.Credentials.NativePort, SSLEnabled: response.Credentials.SSLEnabled, SSLRejectUnauthorized: response.Credentials.SSLRejectUnauthorized, SSLCertificate: response.Credentials.SSLCertificate, From 7f838fa0d234f4a6d407e8d9324377a3a92fa125 Mon Sep 17 00:00:00 2001 From: bernie-g Date: Fri, 25 Sep 2026 15:13:10 -0400 Subject: [PATCH 02/16] chore(clickhouse): trim comments --- packages/gateway-v2/capabilities.go | 4 +- packages/gateway-v2/discovery_handler.go | 6 +- packages/gateway-v2/gateway.go | 8 +- .../gateway-v2/test_connection_handler.go | 15 +-- .../test_connection_handler_test.go | 4 +- packages/pam/handlers/clickhouse/bridge.go | 77 ++++++--------- .../pam/handlers/clickhouse/bridge_test.go | 17 +--- .../pam/handlers/clickhouse/clients_test.go | 2 - .../pam/handlers/clickhouse/contract_test.go | 3 +- .../handlers/clickhouse/edge_cases_test.go | 53 ++-------- packages/pam/handlers/clickhouse/native.go | 99 ++++++------------- .../clickhouse/native_integration_test.go | 10 -- .../pam/handlers/clickhouse/native_outcome.go | 11 +-- .../clickhouse/native_outcome_test.go | 2 - .../handlers/clickhouse/native_unit_test.go | 13 +-- packages/pam/handlers/clickhouse/proxy.go | 11 +-- packages/pam/handlers/clickhouse/sniff.go | 18 +--- .../pam/handlers/clickhouse/sniff_test.go | 2 - packages/pam/handlers/clickhouse/tap_test.go | 2 - .../handlers/clickhouse/testhelpers_test.go | 6 +- packages/pam/pam-proxy.go | 2 +- 21 files changed, 94 insertions(+), 271 deletions(-) diff --git a/packages/gateway-v2/capabilities.go b/packages/gateway-v2/capabilities.go index 51f9097c6..87bc9134f 100644 --- a/packages/gateway-v2/capabilities.go +++ b/packages/gateway-v2/capabilities.go @@ -6,7 +6,5 @@ const CapabilitySessionLogMaskingBuiltInDetection = "sessionLogMaskingBuiltInDet const CapabilitySupportedAccountTypes = "supported_account_types" -// Reported separately from the account type, because a gateway can support ClickHouse accounts and still -// predate the native protocol. Without it the platform cannot tell the difference, and an account with a -// native port would save against an old gateway and then fail every native client at session time. +// Separate from the account type: a gateway can support ClickHouse accounts and still predate native. const CapabilityClickhouseNativeProtocol = "clickhouseNativeProtocol" diff --git a/packages/gateway-v2/discovery_handler.go b/packages/gateway-v2/discovery_handler.go index 4c56800aa..5c174ebc7 100644 --- a/packages/gateway-v2/discovery_handler.go +++ b/packages/gateway-v2/discovery_handler.go @@ -26,13 +26,11 @@ const ( type rpcTarget struct { host string port int - // Every port the signed certificate authorises. Empty means the certificate named only port, which is - // what a platform too old to send the list produces. + // Empty when the certificate named only port, which is what an older platform produces. ports []int } -// allows reports whether the certificate authorises this port. A certificate that named only one port keeps -// the old behaviour, where that port is the only one a handler may reach. +// A certificate naming only one port keeps the old behaviour: that port is the only one reachable. func (t rpcTarget) allows(port int) bool { if port == t.port { return true diff --git a/packages/gateway-v2/gateway.go b/packages/gateway-v2/gateway.go index 12f3b83b9..27495e84b 100644 --- a/packages/gateway-v2/gateway.go +++ b/packages/gateway-v2/gateway.go @@ -80,17 +80,15 @@ type ForwardConfig struct { VerifyTLS bool // Whether to verify TLS certificates TargetHost string TargetPort int - // Additional ports the certificate authorises, empty when it names only TargetPort. - TargetPorts []int - ActorType ActorType - PAMConfig pam.GatewayPAMConfig + TargetPorts []int + ActorType ActorType + PAMConfig pam.GatewayPAMConfig } // RoutingInfo represents the routing information embedded in client certificates type RoutingInfo struct { TargetHost string `json:"targetHost"` TargetPort int `json:"targetPort"` - // Every port this certificate authorises, for the account types that reach one host on more than one. // Absent from a certificate minted by an older platform, which means TargetPort is the only one. TargetPorts []int `json:"targetPorts,omitempty"` } diff --git a/packages/gateway-v2/test_connection_handler.go b/packages/gateway-v2/test_connection_handler.go index 688e0a39d..17e567b2e 100644 --- a/packages/gateway-v2/test_connection_handler.go +++ b/packages/gateway-v2/test_connection_handler.go @@ -724,12 +724,7 @@ func handleTestConnection(w http.ResponseWriter, r *http.Request) { TLSConfig: tlsConfig, } - // A server can have either interface turned off, so only the ones the account names are tested. - // Both are checked up front, so an account that would only ever fail for one kind of client is - // caught here rather than at the first session. - // One ClickHouse account can expose two ports, so the body names which to probe. The signed - // certificate still decides which are allowed, so the body cannot point the gateway at a port - // the platform did not authorise. + // The body names which ports to probe; the signed certificate still decides which are allowed. for _, port := range []int{params.HttpPort, params.NativePort} { if port > 0 && !target.allows(port) { return fmt.Errorf("port %d is not authorised for this connection test", port) @@ -738,12 +733,11 @@ func handleTestConnection(w http.ResponseWriter, r *http.Request) { httpPort := params.HttpPort if httpPort <= 0 && params.NativePort <= 0 { - // An API old enough not to send the ports still means the cert-bound one. + // An API too old to send the ports still means the cert-bound one. httpPort = target.port } - // Each probe gets its own slice of the budget. Sharing one deadline across up to four network - // round trips means a slow first probe swallows the second one's specific error message. + // One shared deadline would let a slow first probe swallow the second one's specific error. probes := 0 if httpPort > 0 { probes++ @@ -835,8 +829,7 @@ func redactProbeSecrets(msg string, secrets ...string) string { return urlUserinfoPattern.ReplaceAllString(msg, "${1}******@") } -// A server can be reachable over HTTP and not over the native protocol, so the failure has to say which port -// it was and that the account can be saved without one. +// The failure has to name the port, and say the account can be saved without one. func nativePortError(port int, err error, httpWorks bool) error { if !httpWorks { return fmt.Errorf("ClickHouse's native port %d did not answer: %w", port, err) diff --git a/packages/gateway-v2/test_connection_handler_test.go b/packages/gateway-v2/test_connection_handler_test.go index e2e2150a9..a2ff5f7eb 100644 --- a/packages/gateway-v2/test_connection_handler_test.go +++ b/packages/gateway-v2/test_connection_handler_test.go @@ -8,7 +8,6 @@ import ( ) // The backend builds this request in TypeScript, so nothing checks the field names match at compile time. -// These payloads were captured from buildGatewayConnectionTest itself. func TestClickhouseTestParamsContract(t *testing.T) { cases := []struct { name string @@ -56,8 +55,7 @@ func TestClickhouseTestParamsContract(t *testing.T) { } } -// The ports to probe come from the request body, so the signed certificate is what stops a caller pointing -// the gateway at a port the platform never authorised. +// The ports to probe come from the request body, so the signed certificate is what stops a caller pointing... func TestRPCTargetAllows(t *testing.T) { t.Run("a certificate naming one port authorises only that port", func(t *testing.T) { target := rpcTarget{host: "db.internal", port: 8123} diff --git a/packages/pam/handlers/clickhouse/bridge.go b/packages/pam/handlers/clickhouse/bridge.go index 2130cb5c8..466b14fcd 100644 --- a/packages/pam/handlers/clickhouse/bridge.go +++ b/packages/pam/handlers/clickhouse/bridge.go @@ -15,32 +15,25 @@ import ( "github.com/rs/zerolog" ) -// Serves ClickHouse's HTTP interface over the native protocol, for a server that has HTTP disabled. Web Access -// reaches the gateway over HTTP and @clickhouse/client cannot speak native, so without this a native-only -// account would work from the CLI and nowhere else. +// Serves ClickHouse's HTTP interface over the native protocol, for a server with HTTP disabled. ClickHouse +// serialises the values through formatRow, so every column type works without this decoding one. // -// ClickHouse serialises the values itself through formatRow, so every column type keeps working without this -// having to decode one. -// -// Only the two envelopes Web Access asks for are produced, and that is the intended scope rather than a gap. -// A native-only account is reached by native clients; a third-party HTTP client such as JDBC, which -// negotiates a binary format, is expected to fail against it rather than be translated for. +// Only the two envelopes Web Access asks for are produced. That is the intended scope: a third-party HTTP +// client such as JDBC is expected to fail against a native-only account rather than be translated for. const ( formatJSON = "JSON" formatJSONCompact = "JSONCompact" - // What formatRow is asked for, so each row arrives as the exact fragment the envelope needs. rowFormatJSON = "JSONEachRow" rowFormatJSONCompact = "JSONCompactEachRow" bridgeReadTimeout = 5 * time.Minute maxBridgeRows = 100_000 - // A row count alone does not bound memory: one row can be hundreds of megabytes, and the envelope is - // assembled in memory. The gateway is shared, so one session must not be able to exhaust it. + // A row count alone does not bound memory, and the gateway is shared across sessions. maxBridgeResultBytes = 64 << 20 - // The formatted column has to be named, because its default name is the whole call expression. + // Needs a name: the default is the whole call expression. formattedRowAlias = "__infisical_row" ) @@ -81,9 +74,8 @@ func (p *ClickHouseProxy) dialNative(ctx context.Context) (*ch.Client, error) { return ch.Dial(ctx, options) } -// serveBridge answers one HTTP request by running its statement over the native protocol. func (p *ClickHouseProxy) serveBridge(w http.ResponseWriter, r *http.Request, state *requestState, l zerolog.Logger) { - // The health endpoint carries no statement, so it would otherwise be refused as an empty one. + // Carries no statement, so it would otherwise be refused as an empty one. if r.URL.Path == pingPath { w.Header().Set("Content-Type", "text/plain; charset=UTF-8") w.WriteHeader(http.StatusOK) @@ -91,8 +83,7 @@ func (p *ClickHouseProxy) serveBridge(w http.ResponseWriter, r *http.Request, st return } - // The statement was already decoded during inspection, compression and all, so it is used rather than - // the body, which the bridge cannot hand to ClickHouse the way the reverse proxy can. + // Already decoded during inspection, compression and all; the raw body cannot be handed to ClickHouse. if state.truncated { message := fmt.Sprintf( "this account reaches ClickHouse over the native protocol, so a statement larger than %d MB is "+ @@ -104,7 +95,7 @@ func (p *ClickHouseProxy) serveBridge(w http.ResponseWriter, r *http.Request, st body, format := splitFormatClause(state.sql) if format == "" { - // A client can ask for the format as a setting instead of a clause, which is what the SQL editor does. + // The SQL editor asks for the format as a setting rather than a clause. format = r.URL.Query().Get("default_format") } if body == "" { @@ -174,14 +165,12 @@ func bridgeParameters(r *http.Request) []proto.Parameter { return parameters } -// The native protocol carries a query parameter as a custom setting, whose value ClickHouse reads as a Field -// dump rather than as the plain string the HTTP interface takes. An unquoted value is rejected outright. +// Native carries a parameter as a custom setting, read as a Field dump rather than the plain string HTTP takes. func quoteFieldDump(value string) string { var quoted strings.Builder quoted.Grow(len(value) + 2) quoted.WriteByte('\'') - // Byte-wise, because ranging over the string would turn every invalid byte into U+FFFD and silently - // change the value. The two escapes are ASCII, and UTF-8 is self-synchronising. + // Byte-wise: ranging would turn every invalid byte into U+FFFD and silently change the value. for i := 0; i < len(value); i++ { if c := value[i]; c == '\\' || c == '\'' { quoted.WriteByte('\\') @@ -192,8 +181,7 @@ func quoteFieldDump(value string) string { return quoted.String() } -// splitFormatClause peels off the trailing FORMAT clause @clickhouse/client appends, which decides the -// envelope rather than anything the server should see. The scan runs over the original bytes: uppercasing +// Peels off the trailing FORMAT clause, which decides the envelope. Scans the original bytes: uppercasing // first would shift offsets, because some runes shrink when folded. func splitFormatClause(sql string) (string, string) { trimmed := trimTrailingSemicolons(sql) @@ -202,8 +190,7 @@ func splitFormatClause(sql string) (string, string) { if idx <= 0 { return trimmed, "" } - // A word boundary is needed on both sides, or a trailing identifier such as `format_events` is read as - // the clause and the operand before it is thrown away. + // Boundary on both sides, or a trailing identifier such as `format_events` is read as the clause. if !isSQLSpace(trimmed[idx-1]) { return trimmed, "" } @@ -220,7 +207,7 @@ func splitFormatClause(sql string) (string, string) { return trimTrailingSemicolons(trimmed[:idx]), name } -// A format name is a bare identifier. Anything else means the word FORMAT was part of the statement. +// Anything but a bare identifier means FORMAT was part of the statement. func isFormatName(name string) bool { if name == "" { return false @@ -239,8 +226,7 @@ func isSQLSpace(c byte) bool { return c == ' ' || c == '\t' || c == '\r' || c == '\n' } -// trimTrailingSemicolons removes any run of trailing semicolons and the whitespace around them, so -// `SELECT 1 ; ;` does not end up inside the subquery wrapper. +// So `SELECT 1 ; ;` does not end up inside the subquery wrapper. func trimTrailingSemicolons(sql string) string { trimmed := strings.TrimSpace(sql) for strings.HasSuffix(trimmed, ";") { @@ -249,7 +235,6 @@ func trimTrailingSemicolons(sql string) string { return trimmed } -// lastIndexFold is strings.LastIndex with ASCII case folding, returning an offset into s itself. func lastIndexFold(s string, substr string) int { for i := len(s) - len(substr); i >= 0; i-- { if strings.EqualFold(s[i:i+len(substr)], substr) { @@ -259,8 +244,8 @@ func lastIndexFold(s string, substr string) int { return -1 } -// checkSpliceable rejects a statement that would not stay inside the parentheses it is wrapped in. Quotes -// and comments are tracked so that a semicolon or bracket inside a string literal is left alone. +// Rejects a statement that would not stay inside its wrapper. Quotes and comments are tracked so a +// semicolon or bracket inside a string literal is left alone. func checkSpliceable(body string) error { depth := 0 for i := 0; i < len(body); i++ { @@ -309,7 +294,6 @@ func checkSpliceable(body string) error { return nil } -// skipQuoted returns the index of the closing quote, honouring doubled and backslash escapes. func skipQuoted(body string, start int, quote byte) int { for i := start + 1; i < len(body); i++ { switch body[i] { @@ -326,7 +310,7 @@ func skipQuoted(body string, start int, quote byte) int { return -1 } -// humanList renders an allowlist the way the error messages read, so the message and the list cannot drift. +// Keeps the error message and the allowlist from drifting apart. func humanList(items []string) string { switch len(items) { case 0: @@ -353,8 +337,7 @@ func rowFormatFor(format string) (string, error) { } } -// Settings a client sends to bound a statement. They are forwarded so a browser session costs the server no -// more over the native protocol than it does over HTTP; anything else a client asks for is dropped. +// Forwarded so a browser session costs the server no more over native than over HTTP; anything else is dropped. var forwardedSettings = map[string]bool{ "max_execution_time": true, "max_result_rows": true, @@ -385,7 +368,7 @@ func (p *ClickHouseProxy) runBridgeQuery( ) (*bridgeEnvelope, error) { envelope := &bridgeEnvelope{Meta: []bridgeColumn{}, Data: json.RawMessage("[]")} - // A statement with no FORMAT clause is one nothing reads the rows of, so it only has to run. + // Nothing reads the rows, so it only has to run. if format == "" { var discard proto.Results return envelope, client.Do(ctx, ch.Query{ @@ -393,8 +376,7 @@ func (p *ClickHouseProxy) runBridgeQuery( Parameters: parameters, Settings: settings, Result: discard.Auto(), - // ch-go refuses a second data block unless a handler is present, so a statement that returns - // more than one block would fail even though nothing here reads the rows. + // ch-go refuses a second data block unless a handler is present. OnResult: func(context.Context, proto.Block) error { return nil }, OnProgress: func(_ context.Context, pr proto.Progress) error { envelope.Statistics.RowsRead += pr.Rows @@ -431,7 +413,7 @@ func (p *ClickHouseProxy) runBridgeQuery( return envelope, nil } -// Only these can sit inside a subquery, which is what both halves of the bridge rely on. +// Only these can sit inside a subquery, which both halves of the bridge rely on. var wrappableStatements = []string{"SELECT", "WITH", "EXPLAIN"} func isWrappable(body string) bool { @@ -440,7 +422,7 @@ func isWrappable(body string) bool { if len(rest) < len(prefix) || !strings.EqualFold(rest[:len(prefix)], prefix) { continue } - // A prefix match is not a keyword match: SELECTFOO is an identifier, not a SELECT. + // SELECTFOO is an identifier, not a SELECT. if len(rest) == len(prefix) || !isIdentifierByte(rest[len(prefix)]) { return true } @@ -452,8 +434,7 @@ func isIdentifierByte(c byte) bool { return c >= 'a' && c <= 'z' || c >= 'A' && c <= 'Z' || c >= '0' && c <= '9' || c == '_' } -// stripLeadingNoise drops a byte-order mark, whitespace and leading comments, which editors add freely and -// which would otherwise make a perfectly ordinary SELECT look like something the bridge cannot serve. +// Editors add leading comments freely, which would otherwise make an ordinary SELECT look unservable. func stripLeadingNoise(body string) string { rest := strings.TrimPrefix(body, "\ufeff") for { @@ -478,7 +459,8 @@ func stripLeadingNoise(body string) string { } } -// DESCRIBE resolves the statement's header without running it, so the column types are the server's own. +// DESCRIBE resolves the header, so the column types are the server's own. It is not free: schema inference +// for a table function such as url() or s3() does fetch. func describeStatement( ctx context.Context, client *ch.Client, @@ -486,8 +468,7 @@ func describeStatement( parameters []proto.Parameter, settings []ch.Setting, ) ([]bridgeColumn, error) { - // DESCRIBE returns more columns than are needed here and the set has grown between versions, so they are - // inferred and picked by name rather than bound positionally. + // The column set has grown between versions, so they are picked by name rather than position. var described proto.Results columns := []bridgeColumn{} @@ -580,8 +561,7 @@ func selectFormattedRows( return rows, nil } -// bridgeRefusal is something the gateway decided about the request itself, so it carries its own ClickHouse -// code and is reported as a bad request rather than as a network error wrapped in ch-go's decoding context. +// A refusal the gateway made itself, so it reports its own code rather than ch-go's decoding context. type bridgeRefusal struct { code int message string @@ -593,7 +573,6 @@ func refuseBridge(code int, format string, args ...any) error { return &bridgeRefusal{code: code, message: fmt.Sprintf(format, args...)} } -// classifyNativeError turns a ch-go error back into the code, message and status a ClickHouse client expects. func classifyNativeError(err error) (int, int, string) { var refusal *bridgeRefusal if errors.As(err, &refusal) { diff --git a/packages/pam/handlers/clickhouse/bridge_test.go b/packages/pam/handlers/clickhouse/bridge_test.go index 6128ef158..49ff790b2 100644 --- a/packages/pam/handlers/clickhouse/bridge_test.go +++ b/packages/pam/handlers/clickhouse/bridge_test.go @@ -38,8 +38,7 @@ func TestSplitFormatClause(t *testing.T) { }, {name: "a column named format", sql: "SELECT format FROM t", wantBody: "SELECT format FROM t"}, {name: "format with no name", sql: "SELECT 1 FORMAT", wantBody: "SELECT 1 FORMAT"}, - // A trailing identifier that merely starts with "format" is not a clause, and splitting it would - // throw away the operand in front of it. + // A trailing identifier that merely starts with "format" is not a clause, and splitting it would throw... {name: "a table whose name starts with format", sql: "SELECT * FROM format_events", wantBody: "SELECT * FROM format_events"}, {name: "an alias that starts with format", sql: "SELECT 1 AS format_id", wantBody: "SELECT 1 AS format_id"}, {name: "ordering by a column called formatted", sql: "SELECT x FROM t ORDER BY formatted", wantBody: "SELECT x FROM t ORDER BY formatted"}, @@ -104,8 +103,6 @@ func TestCheckSpliceable(t *testing.T) { } } -// The bridge has to answer what ClickHouse's own HTTP interface answers, so both are asked the same thing -// and the envelopes compared. func TestBridgeMatchesHTTPInterface(t *testing.T) { if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") @@ -177,7 +174,6 @@ func TestBridgeRecordsStatements(t *testing.T) { require.Contains(t, recorder.dump(), "1 row(s) returned") } -// A native-only account has no HTTP upstream, so TargetAddr is deliberately empty. func startBridgeProxy(t *testing.T, blocked []string, logger *recordingLogger) string { t.Helper() @@ -263,7 +259,6 @@ func queryRealHTTP(t *testing.T, statement string, format string) bridgeEnvelope return envelope } -// The bridge exists for a server with HTTP genuinely turned off, so it is also exercised against one. func TestBridgeAgainstHTTPDisabledServer(t *testing.T) { if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") @@ -312,8 +307,6 @@ func TestBridgeAgainstHTTPDisabledServer(t *testing.T) { require.JSONEq(t, `[{"n":1,"m":{"k":"v"}}]`, string(envelope.Data)) } -// The SQL editor asks for its format with the default_format setting rather than a FORMAT clause, which is -// a different code path and was returning nothing at all. func TestBridgeHonoursDefaultFormatSetting(t *testing.T) { if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") @@ -339,7 +332,6 @@ func TestBridgeHonoursDefaultFormatSetting(t *testing.T) { } } -// A statement whose last line is a comment would otherwise swallow the wrapper's closing parenthesis. func TestBridgeHandlesATrailingComment(t *testing.T) { if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") @@ -351,7 +343,6 @@ func TestBridgeHandlesATrailingComment(t *testing.T) { require.Contains(t, body, `"n":1`) } -// A result bigger than one block used to fail because ch-go refuses a second block with no handler. func TestBridgeHandlesAMultiBlockResult(t *testing.T) { if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") @@ -393,7 +384,6 @@ func TestBridgeRefusesAResultBeyondTheRowCap(t *testing.T) { require.NotContains(t, body, "decode block", "ch-go's internal wrapping should not reach the client") } -// /ping is a health check, and on a native-only account it used to be refused as an empty statement. func TestBridgeAnswersPing(t *testing.T) { if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") @@ -405,7 +395,6 @@ func TestBridgeAnswersPing(t *testing.T) { require.Contains(t, body, "Ok.") } -// A parameter is data. Values that are awkward to quote must survive unchanged. func TestBridgeParameterRoundTrip(t *testing.T) { if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") @@ -432,9 +421,7 @@ func TestBridgeParameterRoundTrip(t *testing.T) { } } -// @clickhouse/client.insert() sends `INSERT INTO t FORMAT JSONEachRow` with the rows in the body, which goes -// down the no-format path as one native query carrying inline data. It must complete rather than sit until -// the server's receive timeout, and the rows have to actually land. +// @clickhouse/client.insert() sends `INSERT INTO t FORMAT JSONEachRow` with the rows in the body, which... func TestBridgeInsertWithInlineData(t *testing.T) { if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") diff --git a/packages/pam/handlers/clickhouse/clients_test.go b/packages/pam/handlers/clickhouse/clients_test.go index ee74c3e60..c01b520b1 100644 --- a/packages/pam/handlers/clickhouse/clients_test.go +++ b/packages/pam/handlers/clickhouse/clients_test.go @@ -9,8 +9,6 @@ import ( "time" ) -// clickhouse-client is one native implementation. Python's clickhouse-driver is an independent one, so it -// catches assumptions that happen to match ClickHouse's own client. func TestPythonClickHouseDriver(t *testing.T) { if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") diff --git a/packages/pam/handlers/clickhouse/contract_test.go b/packages/pam/handlers/clickhouse/contract_test.go index f49c4b6de..c303d774c 100644 --- a/packages/pam/handlers/clickhouse/contract_test.go +++ b/packages/pam/handlers/clickhouse/contract_test.go @@ -8,8 +8,7 @@ import ( "github.com/stretchr/testify/require" ) -// The API and the gateway agree on these shapes only by convention, and a renamed field would not fail to -// compile in either repo. The payloads below were captured from the backend's own builders. +// The API and the gateway agree on these shapes only by convention, and a renamed field would not fail to... func TestSessionCredentialsContract(t *testing.T) { cases := []struct { name string diff --git a/packages/pam/handlers/clickhouse/edge_cases_test.go b/packages/pam/handlers/clickhouse/edge_cases_test.go index 733ebe3a9..30a485fbf 100644 --- a/packages/pam/handlers/clickhouse/edge_cases_test.go +++ b/packages/pam/handlers/clickhouse/edge_cases_test.go @@ -27,7 +27,6 @@ func itOnly(t *testing.T) { } } -// startProxy serves one local port for whatever config it is given, which is what a session does. func startProxy(t *testing.T, config ClickHouseProxyConfig) string { t.Helper() @@ -70,10 +69,6 @@ func baseConfig(logger *recordingLogger, blocked ...string) ClickHouseProxyConfi } } -// ---------- compression ---------- - -// clickhouse-client compresses data blocks by default, which is a different decode path from an -// uncompressed session. Both have to keep inspecting every statement. func TestNativeCompressionBothWays(t *testing.T) { itOnly(t) @@ -94,7 +89,6 @@ func TestNativeCompressionBothWays(t *testing.T) { } } -// An INSERT pushes binary rows from the client, so it exercises block decoding rather than framing alone. func TestNativeInsertWithCompressionBothWays(t *testing.T) { itOnly(t) @@ -112,10 +106,7 @@ func TestNativeInsertWithCompressionBothWays(t *testing.T) { } } -// ---------- the documented fail-closed path ---------- - -// A column type ch-go cannot infer can only appear in the client direction on an INSERT. That has to fail -// closed with a message naming the type, never relay uninspected. +// A column type ch-go cannot infer can only appear in the client direction on an INSERT. func TestNativeInsertIntoUnreadableColumnFailsClosed(t *testing.T) { itOnly(t) @@ -131,12 +122,10 @@ func TestNativeInsertIntoUnreadableColumnFailsClosed(t *testing.T) { require.Contains(t, out, "Map(String, UInt64)", "the message should name the offending type") require.Contains(t, out, "HTTP interface", "the message should point at the way that works") - // Refusing has to mean the rows never reach ClickHouse. A refusal that still commits the INSERT is - // worse than no check at all. + // Refusing has to mean the rows never reach ClickHouse. require.Equal(t, 0, countExotic(t, 99), "the refused rows should not have been written") } -// countExotic reads straight from ClickHouse rather than through the proxy. func countExotic(t *testing.T, id int) int { t.Helper() @@ -161,8 +150,6 @@ func countExotic(t *testing.T, id int) int { return count } -// A refusal ends the session. Anything else leaves the parser reading a stream it has lost its place in, -// where a later packet can flush the refused bytes upstream and run the statement that was just refused. func TestBlockedStatementEndsTheSession(t *testing.T) { itOnly(t) @@ -205,8 +192,6 @@ func countWriteTest(t *testing.T) int { return count } -// ---------- TLS ---------- - func tlsConfigFor(t *testing.T, insecure bool) *tls.Config { t.Helper() config := &tls.Config{ServerName: "localhost", InsecureSkipVerify: insecure} @@ -241,7 +226,6 @@ func tlsConfigSkipOrConfig(t *testing.T) (ClickHouseProxyConfig, bool) { }, true } -// One SSL toggle covers both interfaces, because ClickHouse shares the certificate between them. func TestTLSUpstream(t *testing.T) { itOnly(t) @@ -301,8 +285,6 @@ func TestTLSUpstream(t *testing.T) { }) } -// ---------- account shapes ---------- - func TestAccountWithoutNativePortRefusesNativeClients(t *testing.T) { itOnly(t) @@ -344,14 +326,11 @@ func TestAccountWithNeitherPortFailsClearly(t *testing.T) { config.NativeAddr = "" port := startProxy(t, config) - // The session layer rejects this config before a handler ever runs, so the handler's own guard simply - // refuses the connection rather than answering it. + // The session layer rejects this config before a handler ever runs, so the handler's own guard simply... _, _, err := postStatementE("127.0.0.1:"+port, "SELECT 1") require.Error(t, err, "a session with neither port must not serve anything") } -// ---------- bridge edge cases ---------- - func TestBridgeEdgeCases(t *testing.T) { itOnly(t) @@ -407,7 +386,6 @@ func TestBridgeEdgeCases(t *testing.T) { } } -// A parameterized statement has to run and be recorded with its values, the same as over HTTP. func TestBridgePassesQueryParameters(t *testing.T) { itOnly(t) @@ -423,8 +401,6 @@ func TestBridgePassesQueryParameters(t *testing.T) { require.Contains(t, recorder.dump(), "wanted=7", "the parameter belongs in the recording") } -// A gzipped body has to work on both paths. The reverse proxy hands it to ClickHouse untouched, while the -// bridge has to run the decoded statement itself. func TestCompressedRequestBodies(t *testing.T) { itOnly(t) @@ -443,8 +419,7 @@ func TestCompressedRequestBodies(t *testing.T) { status, body := postGzipped(t, "127.0.0.1:"+port, "SELECT 5 AS five \nFORMAT JSON") require.Equal(t, http.StatusOK, status, body) - // ClickHouse pretty-prints its JSON and the bridge writes it compact, so the rows are compared - // rather than the bytes. + // ClickHouse pretty-prints its JSON and the bridge writes it compact, so the rows are compared rather than... var envelope bridgeEnvelope require.NoError(t, json.Unmarshal([]byte(body), &envelope), body) require.JSONEq(t, `[{"five":5}]`, string(envelope.Data)) @@ -452,8 +427,6 @@ func TestCompressedRequestBodies(t *testing.T) { } } -// ---------- sniffing ---------- - func TestSnifferEdgeCases(t *testing.T) { itOnly(t) @@ -478,8 +451,7 @@ func TestSnifferEdgeCases(t *testing.T) { defer conn.Close() _, err = conn.Write([]byte{0xFF, 0xFE, 0xFD, 0xFC}) require.NoError(t, err) - // Bytes that are not a request line leave net/http waiting for headers, so the bound here is its - // ReadHeaderTimeout rather than the sniff timeout. + // Bytes that are not a request line leave net/http waiting for headers, so the bound here is its... require.NoError(t, conn.SetReadDeadline(time.Now().Add(45*time.Second))) buf := make([]byte, 256) @@ -507,10 +479,6 @@ func TestSnifferEdgeCases(t *testing.T) { }) } -// ---------- concurrency ---------- - -// A session hands out one port that many clients share, so the handler has to hold up under parallel use -// of both protocols at once. func TestConcurrentMixedProtocolSessions(t *testing.T) { itOnly(t) @@ -558,8 +526,6 @@ func TestConcurrentMixedProtocolSessions(t *testing.T) { } } -// ---------- volume ---------- - func TestNativeLargeResultSet(t *testing.T) { itOnly(t) @@ -586,7 +552,6 @@ func TestNativeWideRowsStreamThrough(t *testing.T) { require.Equal(t, direct, strings.TrimSpace(out), "the proxied total must match the server's") } -// queryDirect bypasses the proxy so a test can state what the answer should be. func queryDirect(t *testing.T, sql string) string { t.Helper() @@ -606,8 +571,6 @@ func queryDirect(t *testing.T, sql string) string { return strings.TrimSpace(string(raw)) } -// ---------- upstream failures ---------- - func TestUpstreamUnreachable(t *testing.T) { itOnly(t) @@ -647,8 +610,6 @@ func TestWrongAccountCredentialsSurfaceCleanly(t *testing.T) { require.Contains(t, out, "refused the account") } -// ---------- recording ---------- - func TestNativeRecordingCapturesOutcomes(t *testing.T) { itOnly(t) @@ -667,7 +628,6 @@ func TestNativeRecordingCapturesOutcomes(t *testing.T) { require.Equal(t, 1, strings.Count(dump, "ERROR:"), "the failed statement is recorded once") } -// The revision the client speaks is newer than ch-go's, so the handshake has to pin it and say so. func TestNativeRevisionPinning(t *testing.T) { itOnly(t) @@ -693,8 +653,7 @@ func TestQuoteFieldDump(t *testing.T) { } } -// A parameter is data, so a value full of quotes has to come back as that value rather than changing -// the statement around it. +// A parameter is data, so a value full of quotes has to come back as that value rather than changing the... func TestBridgeParameterCannotEscapeItsQuotes(t *testing.T) { itOnly(t) diff --git a/packages/pam/handlers/clickhouse/native.go b/packages/pam/handlers/clickhouse/native.go index 407651adc..5d97e61ff 100644 --- a/packages/pam/handlers/clickhouse/native.go +++ b/packages/pam/handlers/clickhouse/native.go @@ -18,25 +18,17 @@ import ( "github.com/rs/zerolog" ) -// Native brokers ClickHouse's TCP protocol on 9000/9440, which the HTTP interface cannot serve: clickhouse-client, -// clickhouse-driver and clickhouse-go all speak it exclusively. A statement arrives length-prefixed in its own -// Query packet, so unlike HTTP there is no inspection window to overrun and an INSERT's rows never reach the policy. -// -// Only the client direction is parsed. The server direction is relayed byte for byte, because nothing in it is -// inspected and decoding it would make a SELECT fail on any column type ch-go cannot infer. +// Only the client direction gates anything. The server direction is read solely to pair an outcome with a +// statement, and gives that up rather than fail a SELECT on a column type ch-go cannot infer. const ( - // ch-go decodes up to its own revision, and a newer server sends Hello fields it cannot read. Both sides are - // pinned to what we can parse, which downgrades the client too. + // A newer server sends Hello fields ch-go cannot read, so both sides are pinned to what we can parse. maxNativeRevision = proto.Version - nativeDialTimeout = 30 * time.Second - nativeWriteTimeout = 60 * time.Second - // A port that accepts TCP and then says nothing is the symptom of the HTTP port entered as the native one, - // so the handshake gives up well before the dial would. + nativeDialTimeout = 30 * time.Second + nativeWriteTimeout = 60 * time.Second nativeHandshakeTimeout = 10 * time.Second - // Long enough that a slow client is not dropped mid-statement, short enough to reap an abandoned session. - nativeIdleTimeout = 12 * time.Hour + nativeIdleTimeout = 12 * time.Hour ) type nativeProxy struct { @@ -47,8 +39,8 @@ func newNativeProxy(owner *ClickHouseProxy) *nativeProxy { return &nativeProxy{ClickHouseProxy: owner} } -// tap records every byte a decoder consumes so the exact wire bytes can be replayed upstream. Reads are handed -// out one at a time: proto.Reader buffers 128 KB, which would swallow packets we have not parsed yet. +// tap records every byte a decoder consumes so the exact wire bytes can be replayed upstream. One byte at a +// time, because proto.Reader buffers 128 KB and would swallow packets we have not parsed yet. type tap struct { src *bufio.Reader buf []byte @@ -71,15 +63,13 @@ func (t *tap) Read(p []byte) (int, error) { return 1, nil } -// take returns the bytes consumed since the last call and starts a new run. func (t *tap) take() []byte { b := t.buf t.buf = nil return b } -// rest is the reader the tap is draining, buffered bytes included. Relaying from the socket directly would -// silently drop whatever the buffer already holds. +// Relaying from the socket instead would drop whatever the tap's buffer already holds. func (t *tap) rest() io.Reader { return t.src } @@ -88,9 +78,8 @@ func (t *tap) discard() { t.buf = nil } -// Writing an exception ends the session. The client's stream is mid-packet once a refusal happens, so -// carrying on would parse garbage, and anything still in the tap could be flushed upstream by a later -// packet: a refused statement would reach ClickHouse after being refused. +// A refusal ends the session: the stream is mid-packet, so carrying on would let a later packet flush the +// refused bytes upstream. var errSessionRefused = errors.New("the session was refused") type nativeSession struct { @@ -100,23 +89,19 @@ type nativeSession struct { upstream net.Conn rev int - // One tap over the upstream, shared by the handshake and the server loop. A second tap would start its - // own buffer and lose whatever the first had already read off the socket. + // One tap for both the handshake and the server loop; a second would lose what the first buffered. upstreamTap *tap upstreamReader *proto.Reader outcomes *outcomeRecorder - // Both directions write to the client, so a refusal must not land inside a packet the server loop is - // still writing. Once refused, the server loop stops writing entirely. + // Both directions write to the client, so a refusal must not land inside a packet mid-write. writeMu sync.Mutex refused atomic.Bool - // Set by the Query packet and read by the data-block decoder compressed atomic.Bool } -// writeToClient serialises the two directions and bounds a client that has stopped reading. func (s *nativeSession) writeToClient(payload []byte) error { if len(payload) == 0 { return nil @@ -144,7 +129,6 @@ func (s *nativeSession) writeToUpstream(payload []byte) error { func (p *nativeProxy) HandleConnection(ctx context.Context, clientConn net.Conn, l zerolog.Logger) error { defer clientConn.Close() - // A malformed packet must not take the gateway, and every other session with it, down. defer func() { if r := recover(); r != nil { l.Error().Interface("panic", r).Msg("Recovered from a panic in the ClickHouse native handler") @@ -173,8 +157,7 @@ func (p *nativeProxy) HandleConnection(ctx context.Context, clientConn net.Conn, clientTap := newTap(clientConn) clientReader := proto.NewReader(clientTap) - // A port that accepts TCP and then says nothing is the symptom of the HTTP port entered as the native - // one, so the handshake gives up rather than holding both sockets until the session expires. + // A port that accepts TCP then says nothing is the HTTP port entered as the native one. deadline := time.Now().Add(nativeHandshakeTimeout) _ = clientConn.SetDeadline(deadline) _ = upstream.SetDeadline(deadline) @@ -187,8 +170,6 @@ func (p *nativeProxy) HandleConnection(ctx context.Context, clientConn net.Conn, _ = clientConn.SetDeadline(time.Time{}) _ = upstream.SetDeadline(time.Time{}) - // Nothing in the server direction gates a statement, so a packet it cannot read costs only the outcome of - // that statement: it flushes what it read and streams the rest untouched. serverDone := make(chan struct{}) go func() { defer close(serverDone) @@ -204,8 +185,7 @@ func (p *nativeProxy) HandleConnection(ctx context.Context, clientConn net.Conn, l.Debug().Err(err).Msg("ClickHouse native session ended") } - // Closing the upstream unblocks the server loop. Its outcomes have to land before the recorder drains, - // or a statement that finished is written to the recording as interrupted. + // Outcomes must land before the recorder drains, or a finished statement is recorded as interrupted. upstream.Close() select { case <-serverDone: @@ -216,12 +196,10 @@ func (p *nativeProxy) HandleConnection(ctx context.Context, clientConn net.Conn, return nil } -// refuse reports the refusal to the client and ends the session. func (s *nativeSession) refuse(t *tap, code int, message string) error { if t != nil { t.discard() } - // Set before writing, so the server loop cannot interleave a packet with the exception. s.refused.Store(true) if err := s.writeToClient(nativeErrorPacket(s.rev, code, message)); err != nil { @@ -238,8 +216,7 @@ func (p *nativeProxy) dialUpstream(ctx context.Context) (net.Conn, error) { return (&tls.Dialer{NetDialer: dialer, Config: p.config.TLSConfig}).DialContext(ctx, "tcp", p.config.NativeAddr) } -// handshake swaps the client's credentials for the account's, so nothing the client holds works outside a -// recorded session, and pins the protocol revision both ways. +// handshake swaps the client's credentials for the account's and pins the revision both ways. func (s *nativeSession) handshake(t *tap, r *proto.Reader) error { code, err := r.UVarInt() if err != nil { @@ -255,8 +232,7 @@ func (s *nativeSession) handshake(t *tap, r *proto.Reader) error { } t.discard() - // The revision drives feature gating on both sides, and it is unvalidated client input: a huge uvarint - // decodes to a negative int, which would be re-encoded upstream as an enormous revision. + // Unvalidated client input: a huge uvarint decodes negative and would be re-encoded as an enormous revision. if hello.ProtocolVersion <= 0 { return s.refuse(t, codeNotImplemented, fmt.Sprintf("This session could not read the protocol revision %d the client asked for.", @@ -301,7 +277,6 @@ func (s *nativeSession) handshake(t *tap, r *proto.Reader) error { if err := serverHello.DecodeAware(serverReader, s.rev); err != nil { return fmt.Errorf("decode server hello: %w", err) } - // The client keys its own encoding off the revision it is told, so it has to see the pinned one. serverHello.Revision = s.rev s.upstreamTap.discard() @@ -312,8 +287,7 @@ func (s *nativeSession) handshake(t *tap, r *proto.Reader) error { return fmt.Errorf("write client hello response: %w", err) } - // At rev >= 54458 the client follows the handshake with its quota key, written as a bare string rather than - // a coded packet. Ours is empty: the account's quota is not the client's to choose. + // At rev >= 54458 the quota key follows the handshake as a bare string. Ours is empty: not the client's to pick. if proto.FeatureAddendum.In(s.rev) { if _, err := r.Str(); err != nil { return fmt.Errorf("read client addendum: %w", err) @@ -334,8 +308,7 @@ func (s *nativeSession) handshake(t *tap, r *proto.Reader) error { return nil } -// clientLoop parses every packet the client sends. A statement that is never parsed is a statement the policy -// never sees, so an unreadable stream ends the session rather than being relayed blind. +// A statement that is never parsed is one the policy never sees, so an unreadable stream ends the session. func (s *nativeSession) clientLoop(t *tap, r *proto.Reader) error { for { _ = s.client.SetReadDeadline(time.Now().Add(nativeIdleTimeout)) @@ -352,8 +325,7 @@ func (s *nativeSession) clientLoop(t *tap, r *proto.Reader) error { } case proto.ClientTablesStatusRequest: - // This one carries a table list the loop does not decode. Forwarding just the code would leave - // the stream one packet out of step and every later statement unreadable. + // Carries a table list the loop does not decode; forwarding just the code would desync the stream. return s.refuse(t, codeNotImplemented, "This session does not support ClickHouse's tables-status request.") @@ -386,8 +358,7 @@ func (s *nativeSession) handleQuery(t *tap, r *proto.Reader) error { } t.discard() - // EncodeAware always writes StageComplete, so a client asking for a partial stage would have its query - // silently upgraded to a full execution. Refusing is the honest answer. + // EncodeAware always writes StageComplete, so a partial stage would be silently upgraded to a full run. if q.Stage != proto.StageComplete { return s.refuse(t, codeNotImplemented, fmt.Sprintf("This session only runs statements to completion, and this client asked for stage %d.", @@ -407,10 +378,8 @@ func (s *nativeSession) handleQuery(t *tap, r *proto.Reader) error { s.outcomes.begin(statement) - // The identities a client could otherwise choose for itself, matching the set the HTTP interface strips. - // An inter-server secret in particular must never be something a session gets to pick. - // InitialUser and InitialAddress are deliberately left alone: forcing the query kind to Initial already - // makes ClickHouse authorise as the account, and it asserts on an empty initial address. + // The identities a client could otherwise pick for itself. InitialAddress is left alone: ClickHouse + // asserts on an empty one, and forcing the kind to Initial already authorises as the account. q.Info.QuotaKey = "" q.Info.Query = proto.ClientQueryInitial q.Secret = "" @@ -420,8 +389,8 @@ func (s *nativeSession) handleQuery(t *tap, r *proto.Reader) error { return s.forward(b.Buf) } -// handleData decodes a block only far enough to find where it ends, then replays the client's own bytes. The -// decoded values are discarded: re-encoding them would have to reproduce a serialization we do not own. +// Decodes a block only far enough to find its end, then replays the client's bytes: re-encoding would mean +// reproducing a serialization we do not own. func (s *nativeSession) handleData(t *tap, r *proto.Reader) error { table, err := r.Str() if err != nil { @@ -453,8 +422,7 @@ func (s *nativeSession) handleData(t *tap, r *proto.Reader) error { return s.forward(t.take()) } -// serverLoop reads the server direction for the sake of the recording only. Every packet is replayed to the -// client byte for byte, and the first one it cannot read ends the parsing rather than the session. +// Read for the recording only: the first packet it cannot read ends the parsing, not the session. func (s *nativeSession) serverLoop() { t := s.upstreamTap r := s.upstreamReader @@ -467,8 +435,6 @@ func (s *nativeSession) serverLoop() { if err := s.writeToClient(t.take()); err != nil { return } - // Reading the socket directly here would skip whatever the tap has already buffered off it, which - // is most of the in-flight response. _, _ = io.Copy(newRefusalAwareWriter(s), t.rest()) } @@ -542,8 +508,7 @@ func (s *nativeSession) serverLoop() { } } -// refusalAwareWriter stops relaying once the client loop has refused the session, so a raw relay cannot -// append bytes after the exception the client was just sent. +// Stops the raw relay appending bytes after the exception the client was just sent. type refusalAwareWriter struct{ s *nativeSession } func newRefusalAwareWriter(s *nativeSession) io.Writer { return refusalAwareWriter{s: s} } @@ -585,8 +550,7 @@ func (s *nativeSession) forward(payload []byte) error { return nil } -// nativeParameterSuffix mirrors the HTTP handler, so a parameterized statement reads the same in a recording -// whichever interface ran it. +// Mirrors the HTTP handler, so a parameterized statement reads the same in either recording. func nativeParameterSuffix(parameters []proto.Parameter) string { if len(parameters) == 0 { return "" @@ -598,8 +562,7 @@ func nativeParameterSuffix(parameters []proto.Parameter) string { return "\n-- parameters: " + strings.Join(pairs, " ") } -// writeNativeError reports a gateway refusal the way ClickHouse reports its own, so a driver surfaces it as a -// server exception rather than a broken connection. +// Reports a gateway refusal as ClickHouse would, so a driver surfaces it rather than a broken connection. func writeNativeError(w io.Writer, revision int, code int, message string) error { if _, err := w.Write(nativeErrorPacket(revision, code, message)); err != nil { return fmt.Errorf("write native exception: %w", err) @@ -607,7 +570,6 @@ func writeNativeError(w io.Writer, revision int, code int, message string) error return nil } -// nativeErrorPacket builds the exception and the end-of-stream that closes it out. func nativeErrorPacket(revision int, code int, message string) []byte { if revision <= 0 { revision = maxNativeRevision @@ -625,8 +587,7 @@ func nativeErrorPacket(revision int, code int, message string) []byte { return b.Buf } -// TestNativeConnection proves the account can log in over the native port. ClickHouse validates credentials -// during the handshake, so a successful Hello exchange is a real auth check rather than a reachability probe. +// ClickHouse validates credentials during the handshake, so a Hello exchange is a real auth check. func TestNativeConnection(ctx context.Context, config ClickHouseProxyConfig) error { dialCtx, cancel := context.WithTimeout(ctx, nativeDialTimeout) defer cancel() diff --git a/packages/pam/handlers/clickhouse/native_integration_test.go b/packages/pam/handlers/clickhouse/native_integration_test.go index 83f6c9b75..dd636d101 100644 --- a/packages/pam/handlers/clickhouse/native_integration_test.go +++ b/packages/pam/handlers/clickhouse/native_integration_test.go @@ -11,11 +11,6 @@ import ( "time" ) -// Exercises the native handler against a real ClickHouse over the real clickhouse-client, which is the only way -// to cover revision pinning, the addendum and block framing. Opt in with PAM_CLICKHOUSE_NATIVE_IT=1. -// -// docker run -d --name pam-clickhouse-target -p 8123:8123 -p 9000:9000 \ -// -e CLICKHOUSE_PASSWORD=clickhouse -e CLICKHOUSE_DB=analytics clickhouse/clickhouse-server:24.8 func TestNativeIntegration(t *testing.T) { if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") @@ -95,8 +90,6 @@ func TestNativeIntegration(t *testing.T) { } } -// TestNativeRecordsEveryStatement proves the packet loop keeps inspecting after the first statement, which a -// handler that degrades into a raw relay would silently stop doing. func TestNativeRecordsEveryStatement(t *testing.T) { if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") @@ -120,8 +113,6 @@ func TestNativeRecordsEveryStatement(t *testing.T) { } } -// TestNativeRecordsFailedStatement proves a statement ClickHouse rejects is recorded with its error rather -// than as a success. func TestNativeRecordsFailedStatement(t *testing.T) { if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") @@ -141,7 +132,6 @@ func TestNativeRecordsFailedStatement(t *testing.T) { } } -// A column type ch-go cannot infer costs the outcome of that statement, never the statement or the session. func TestNativeDegradesOnUnreadableResultBlock(t *testing.T) { if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") diff --git a/packages/pam/handlers/clickhouse/native_outcome.go b/packages/pam/handlers/clickhouse/native_outcome.go index d9fa2ce48..de81bd531 100644 --- a/packages/pam/handlers/clickhouse/native_outcome.go +++ b/packages/pam/handlers/clickhouse/native_outcome.go @@ -7,11 +7,8 @@ import ( "time" ) -// outcomeRecorder pairs a statement with how it ended. The two directions of a native session are read by -// separate goroutines, and ClickHouse answers statements in order, so the queue is what joins them back up. -// -// Reading the server direction is best effort: a block it cannot decode costs the outcome, never the statement. -// Once that happens the recorder degrades and every later statement is written as soon as it is sent. +// Pairs a statement with how it ended. The two directions are separate goroutines and ClickHouse answers in +// order, so the queue is what joins them back up. Best effort: a block it cannot decode costs only the outcome. type outcomeRecorder struct { proxy *ClickHouseProxy @@ -42,7 +39,6 @@ func (r *outcomeRecorder) begin(statement string) { r.mu.Unlock() } -// progress folds ClickHouse's running counters into the statement in flight. func (r *outcomeRecorder) progress(rows uint64, bytes uint64) { r.mu.Lock() defer r.mu.Unlock() @@ -74,8 +70,7 @@ func (p pendingStatement) describe(outcome string) string { return strings.Join(append(parts, fmt.Sprintf("%dms", time.Since(p.started).Milliseconds())), ", ") } -// degrade stops pairing outcomes for the rest of the session and says so in the recording, so a log that -// carries outcomes for some statements and not others is never read as if the rest simply did nothing. +// Says so in the recording, so a log with outcomes for only some statements is not read as the rest doing nothing. func (r *outcomeRecorder) degrade(reason string) { r.mu.Lock() if r.degraded { diff --git a/packages/pam/handlers/clickhouse/native_outcome_test.go b/packages/pam/handlers/clickhouse/native_outcome_test.go index 4117db18d..0d54f6d68 100644 --- a/packages/pam/handlers/clickhouse/native_outcome_test.go +++ b/packages/pam/handlers/clickhouse/native_outcome_test.go @@ -9,8 +9,6 @@ import ( "github.com/stretchr/testify/require" ) -// The recorder decides what an auditor reads, and it joins two goroutines, so its queue behaviour is worth -// pinning down without needing a database. func newTestRecorder() (*outcomeRecorder, *recordingLogger) { logger := &recordingLogger{} return newOutcomeRecorder(&ClickHouseProxy{config: ClickHouseProxyConfig{SessionLogger: logger}}), logger diff --git a/packages/pam/handlers/clickhouse/native_unit_test.go b/packages/pam/handlers/clickhouse/native_unit_test.go index 797d7dee1..9b592effc 100644 --- a/packages/pam/handlers/clickhouse/native_unit_test.go +++ b/packages/pam/handlers/clickhouse/native_unit_test.go @@ -12,8 +12,7 @@ import ( "github.com/stretchr/testify/require" ) -// fakeClickHouse stands in for a server so the security-critical parts of the handshake and packet loop can -// be tested without docker: what the gateway sends upstream is recorded, and nothing needs a real database. +// fakeClickHouse stands in for a server so the security-critical parts of the handshake and packet loop... type fakeClickHouse struct { listener net.Listener @@ -114,7 +113,6 @@ func (f *fakeClickHouse) snapshot() (proto.ClientHello, string, []proto.Query, i return f.hello, f.quotaKey, append([]proto.Query(nil), f.queries...), f.bytesAfterHandshake } -// dialProxy runs one session against a proxy configured to reach the fake server. func dialProxy(t *testing.T, config ClickHouseProxyConfig) net.Conn { t.Helper() @@ -177,7 +175,6 @@ func clientHandshake(t *testing.T, conn net.Conn, user, password string) *proto. return r } -// The whole point of the proxy: what the client presents is dropped and the account's own identity is used. func TestNativeHandshakeInjectsAccountCredentials(t *testing.T) { upstream := startFakeClickHouse(t) @@ -203,7 +200,6 @@ func TestNativeHandshakeInjectsAccountCredentials(t *testing.T) { require.Empty(t, quotaKey, "the client's quota key is not the account's to choose") } -// A packet the loop cannot read must end the session, and nothing may reach the server after it. func TestNativeUnreadablePacketFailsClosed(t *testing.T) { upstream := startFakeClickHouse(t) @@ -228,7 +224,6 @@ func TestNativeUnreadablePacketFailsClosed(t *testing.T) { require.Zero(t, seen, "no packet may reach ClickHouse after a refusal") } -// A blocked statement must be refused before it is forwarded, and must end the session. func TestNativeBlockedStatementNeverReachesUpstream(t *testing.T) { upstream := startFakeClickHouse(t) recorder := &recordingLogger{} @@ -254,7 +249,6 @@ func TestNativeBlockedStatementNeverReachesUpstream(t *testing.T) { require.Contains(t, recorder.dump(), "BLOCKED") } -// The client's quota key rides on the Query packet as well as the addendum, and both are the account's. func TestNativeStripsTheClientQuotaKeyFromTheQuery(t *testing.T) { upstream := startFakeClickHouse(t) @@ -280,8 +274,6 @@ func TestNativeStripsTheClientQuotaKeyFromTheQuery(t *testing.T) { require.Empty(t, queries[0].Info.QuotaKey) } -// The revision is pinned to what ch-go can parse, and the client has to be told the pinned one so it -// encodes to match. func TestNativeHandshakePinsTheRevision(t *testing.T) { upstream := startFakeClickHouse(t) @@ -350,8 +342,7 @@ func decodeException(t *testing.T, r *proto.Reader) (int, string) { return int(e.Code), e.Message } -// ch-go trailing the server is the reason the handshake pins a revision at all. If ch-go ever catches up, -// the pinning becomes a no-op and this is the reminder to re-check it. +// ch-go trailing the server is the reason the handshake pins a revision at all. func TestChGoRevisionIsStillBehindTheServers(t *testing.T) { require.LessOrEqual(t, maxNativeRevision, 54469, "ch-go has caught up with ClickHouse; revision pinning needs revisiting") diff --git a/packages/pam/handlers/clickhouse/proxy.go b/packages/pam/handlers/clickhouse/proxy.go index b0f3b7039..2cc213fa9 100644 --- a/packages/pam/handlers/clickhouse/proxy.go +++ b/packages/pam/handlers/clickhouse/proxy.go @@ -27,10 +27,8 @@ import ( "github.com/rs/zerolog/log" ) -// Brokers both of ClickHouse's interfaces: HTTP on TargetAddr and the native TCP protocol on NativeAddr. A -// session listens on one local port and routes by the first byte the client sends, so the driver decides the -// protocol rather than the user. The client's own credentials are dropped and the account's injected on either -// path, so nothing it holds works outside a recorded session. +// Brokers both of ClickHouse's interfaces on one local port, routed by the first byte the client sends. The +// client's own credentials are dropped and the account's injected on either path. type ClickHouseProxyConfig struct { TargetAddr string NativeAddr string @@ -99,8 +97,7 @@ type stateKey struct{} type requestState struct { statement string - // The statement without the recorded parameter suffix, and whether it was cut short by the inspection - // window. The bridge runs this rather than re-reading a body that may be compressed. + // The bridge runs this rather than re-reading a body that may be compressed. sql string truncated bool started time.Time @@ -386,7 +383,7 @@ func (p *ClickHouseProxy) rewrite(pr *httputil.ProxyRequest) { for _, param := range append(append([]string{}, strippedAuthParams...), strippedExecutionParams...) { query.Del(param) } - // ClickHouse's health endpoint refuses any query string, so a database parameter turns it into a 404. + // ClickHouse's health endpoint refuses any query string, so a parameter turns it into a 404. if req.URL.Path == pingPath { query = nil req.URL.ForceQuery = false diff --git a/packages/pam/handlers/clickhouse/sniff.go b/packages/pam/handlers/clickhouse/sniff.go index ddff3fb61..827d7de9c 100644 --- a/packages/pam/handlers/clickhouse/sniff.go +++ b/packages/pam/handlers/clickhouse/sniff.go @@ -6,19 +6,13 @@ import ( "time" ) -// ClickHouse serves two interfaces on two ports, and which one a client speaks depends on the driver rather than -// on anything the user chose: clickhouse-client is native-only, the JDBC driver is HTTP-only. A session hands out -// one local port and reads the first byte to tell them apart, so nobody has to pick a protocol. -// -// A native session opens with the Hello packet code, a uvarint 0. Every HTTP request opens with the ASCII letter -// of its method, so the two can never be confused. +// Which protocol a client speaks is decided by its driver, not the user, so one port serves both. A native +// session opens with the Hello code, a uvarint 0; every HTTP request opens with an ASCII method letter. const nativeHelloByte = 0x00 -// A client that connects and then says nothing would otherwise park the session handler forever: the peek -// happens before any HTTP server exists, so ReadHeaderTimeout does not cover it. +// The peek happens before any HTTP server exists, so ReadHeaderTimeout does not cover it. const sniffTimeout = 30 * time.Second -// peekConn replays the sniffed byte to whichever handler takes the connection. type peekConn struct { net.Conn reader *bufio.Reader @@ -28,8 +22,7 @@ func (c *peekConn) Read(p []byte) (int, error) { return c.reader.Read(p) } -// CloseWrite is promoted explicitly: embedding the net.Conn interface hides it, and net/http uses it to -// half-close rather than resetting a connection it is finished with. +// Embedding the net.Conn interface hides this, and net/http uses it to half-close rather than reset. func (c *peekConn) CloseWrite() error { if cw, ok := c.Conn.(interface{ CloseWrite() error }); ok { return cw.CloseWrite() @@ -37,8 +30,7 @@ func (c *peekConn) CloseWrite() error { return nil } -// sniffProtocol reports whether the client opened a native session, and returns a connection that still replays -// the byte it read to decide. +// Returns a connection that still replays the byte it read to decide. func sniffProtocol(conn net.Conn) (net.Conn, bool, error) { reader := bufio.NewReaderSize(conn, 64<<10) diff --git a/packages/pam/handlers/clickhouse/sniff_test.go b/packages/pam/handlers/clickhouse/sniff_test.go index 4816e454f..0598c76fc 100644 --- a/packages/pam/handlers/clickhouse/sniff_test.go +++ b/packages/pam/handlers/clickhouse/sniff_test.go @@ -48,7 +48,6 @@ func TestSniffProtocol(t *testing.T) { } } -// The HTTP interface has to keep working now that a connection is routed by its first byte. func TestHandleConnectionRoutesHTTP(t *testing.T) { upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { _, _ = w.Write([]byte("ok")) @@ -76,7 +75,6 @@ func TestHandleConnectionRoutesHTTP(t *testing.T) { require.Equal(t, http.StatusOK, resp.StatusCode) } -// A native client reaching an account with no native port has to get a native exception, not a dead socket. func TestHandleConnectionRefusesNativeWithoutPort(t *testing.T) { proxy := NewClickHouseProxy(ClickHouseProxyConfig{ TargetAddr: "127.0.0.1:1", diff --git a/packages/pam/handlers/clickhouse/tap_test.go b/packages/pam/handlers/clickhouse/tap_test.go index fdd232a5b..7c8d688af 100644 --- a/packages/pam/handlers/clickhouse/tap_test.go +++ b/packages/pam/handlers/clickhouse/tap_test.go @@ -10,8 +10,6 @@ import ( "github.com/stretchr/testify/require" ) -// The degrade path stops decoding and relays the rest of the stream. Anything the tap's buffered reader -// already pulled off the socket has to go with it, or the client gets a truncated packet. func TestTapRelaysWhatItHasAlreadyBuffered(t *testing.T) { upstreamRead, upstreamWrite := net.Pipe() defer upstreamRead.Close() diff --git a/packages/pam/handlers/clickhouse/testhelpers_test.go b/packages/pam/handlers/clickhouse/testhelpers_test.go index ed4e295c4..c659fb4e7 100644 --- a/packages/pam/handlers/clickhouse/testhelpers_test.go +++ b/packages/pam/handlers/clickhouse/testhelpers_test.go @@ -15,8 +15,6 @@ import ( "github.com/stretchr/testify/require" ) -// runClient drives the real clickhouse-client against the session's port, which is the only way to cover -// revision pinning, the addendum and block framing as a real driver produces them. func runClient(t *testing.T, port string, sql string, extra ...string) (string, error) { t.Helper() @@ -27,7 +25,7 @@ func runClient(t *testing.T, port string, sql string, extra ...string) (string, "run", "--rm", "-i", "clickhouse/clickhouse-server:24.8", "clickhouse-client", "--host", envOr("PAM_CLICKHOUSE_CLIENT_HOST", "host.docker.internal"), "--port", port, - // Deliberately wrong: the gateway replaces them with the account's. + // Deliberately wrong: "--user", "not-the-account", "--password", "not-the-password", "--multiquery", } @@ -40,8 +38,6 @@ func runClient(t *testing.T, port string, sql string, extra ...string) (string, return string(out), err } -// postStatementE is the non-asserting form. require.* calls t.FailNow, which is illegal off the test -// goroutine, so anything running in parallel has to report failures over a channel instead. func postStatementE(addr string, sql string) (int, string, error) { req, err := http.NewRequest(http.MethodPost, "http://"+addr+"/", strings.NewReader(sql)) if err != nil { diff --git a/packages/pam/pam-proxy.go b/packages/pam/pam-proxy.go index a1f76af3f..d2156b6ae 100644 --- a/packages/pam/pam-proxy.go +++ b/packages/pam/pam-proxy.go @@ -594,7 +594,7 @@ func HandlePAMProxy(ctx context.Context, conn *tls.Conn, pamConfig *GatewayPAMCo blockedCommands = compilePolicyPatterns(rulePatterns(credentials.PolicyRules.CommandBlocking), pamConfig.SessionId, "command-blocking") } - // Either interface can be absent: an empty address is what tells the handler that one is not served. + // An empty address is what tells the handler an interface is not served. nativeAddr := "" if credentials.NativePort > 0 { nativeAddr = net.JoinHostPort(credentials.Host, strconv.Itoa(credentials.NativePort)) From f978cb9027b5ae9b28dfadafdf79e379f8379aec Mon Sep 17 00:00:00 2001 From: bernie-g Date: Fri, 25 Sep 2026 15:38:51 -0400 Subject: [PATCH 03/16] refactor(clickhouse): drop the HTTP-to-native bridge --- packages/pam/handlers/clickhouse/bridge.go | 596 ------------------ .../pam/handlers/clickhouse/bridge_test.go | 457 -------------- .../handlers/clickhouse/edge_cases_test.go | 178 +----- packages/pam/handlers/clickhouse/proxy.go | 55 +- .../handlers/clickhouse/testhelpers_test.go | 8 + 5 files changed, 23 insertions(+), 1271 deletions(-) delete mode 100644 packages/pam/handlers/clickhouse/bridge.go delete mode 100644 packages/pam/handlers/clickhouse/bridge_test.go diff --git a/packages/pam/handlers/clickhouse/bridge.go b/packages/pam/handlers/clickhouse/bridge.go deleted file mode 100644 index 466b14fcd..000000000 --- a/packages/pam/handlers/clickhouse/bridge.go +++ /dev/null @@ -1,596 +0,0 @@ -package clickhouse - -import ( - "context" - "encoding/json" - "errors" - "fmt" - "net/http" - "strconv" - "strings" - "time" - - "github.com/ClickHouse/ch-go" - "github.com/ClickHouse/ch-go/proto" - "github.com/rs/zerolog" -) - -// Serves ClickHouse's HTTP interface over the native protocol, for a server with HTTP disabled. ClickHouse -// serialises the values through formatRow, so every column type works without this decoding one. -// -// Only the two envelopes Web Access asks for are produced. That is the intended scope: a third-party HTTP -// client such as JDBC is expected to fail against a native-only account rather than be translated for. - -const ( - formatJSON = "JSON" - formatJSONCompact = "JSONCompact" - - rowFormatJSON = "JSONEachRow" - rowFormatJSONCompact = "JSONCompactEachRow" - - bridgeReadTimeout = 5 * time.Minute - maxBridgeRows = 100_000 - // A row count alone does not bound memory, and the gateway is shared across sessions. - maxBridgeResultBytes = 64 << 20 - - // Needs a name: the default is the whole call expression. - formattedRowAlias = "__infisical_row" -) - -type bridgeColumn struct { - Name string `json:"name"` - Type string `json:"type"` -} - -type bridgeStatistics struct { - Elapsed float64 `json:"elapsed"` - RowsRead uint64 `json:"rows_read"` - BytesRead uint64 `json:"bytes_read"` -} - -type bridgeEnvelope struct { - Meta []bridgeColumn `json:"meta"` - Data json.RawMessage `json:"data"` - Rows int `json:"rows"` - Statistics bridgeStatistics `json:"statistics"` -} - -func (p *ClickHouseProxy) dialNative(ctx context.Context) (*ch.Client, error) { - options := ch.Options{ - Address: p.config.NativeAddr, - Database: p.config.Database, - User: p.config.Username, - Password: p.config.Password, - ClientName: "Infisical PAM", - DialTimeout: nativeDialTimeout, - ReadTimeout: bridgeReadTimeout, - Compression: ch.CompressionDisabled, - ProtocolVersion: maxNativeRevision, - HandshakeTimeout: nativeHandshakeTimeout, - } - if p.config.EnableTLS { - options.TLS = p.config.TLSConfig - } - return ch.Dial(ctx, options) -} - -func (p *ClickHouseProxy) serveBridge(w http.ResponseWriter, r *http.Request, state *requestState, l zerolog.Logger) { - // Carries no statement, so it would otherwise be refused as an empty one. - if r.URL.Path == pingPath { - w.Header().Set("Content-Type", "text/plain; charset=UTF-8") - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte("Ok.\n")) - return - } - - // Already decoded during inspection, compression and all; the raw body cannot be handed to ClickHouse. - if state.truncated { - message := fmt.Sprintf( - "this account reaches ClickHouse over the native protocol, so a statement larger than %d MB is "+ - "not supported here", maxInspectBytes>>20) - p.logStatement(state.statement, "ERROR: "+message) - writeClickHouseError(w, http.StatusBadRequest, codeNotImplemented, message) - return - } - - body, format := splitFormatClause(state.sql) - if format == "" { - // The SQL editor asks for the format as a setting rather than a clause. - format = r.URL.Query().Get("default_format") - } - if body == "" { - writeClickHouseError(w, http.StatusBadRequest, codeNotImplemented, "No statement was sent.") - return - } - - client, err := p.dialNative(r.Context()) - if err != nil { - l.Error().Err(err).Msg("Failed to reach ClickHouse over the native protocol") - p.logStatement(state.statement, fmt.Sprintf("ERROR: %s", err)) - writeClickHouseError(w, http.StatusBadGateway, codeNetworkError, - fmt.Sprintf("The gateway could not reach ClickHouse: %v", err)) - return - } - defer client.Close() - - envelope, err := p.runBridgeQuery(r.Context(), client, body, format, bridgeParameters(r), bridgeSettings(r)) - if err != nil { - code, status, message := classifyNativeError(err) - p.logStatement(state.statement, fmt.Sprintf("ERROR: %s", message)) - writeClickHouseError(w, status, code, message) - return - } - - envelope.Statistics.Elapsed = time.Since(state.started).Seconds() - - if format == "" { - p.logStatement(state.statement, summarizeBridge(envelope, state.started)) - w.Header().Set("Content-Type", "text/plain; charset=UTF-8") - w.WriteHeader(http.StatusOK) - return - } - - encoded, err := json.Marshal(envelope) - if err != nil { - writeClickHouseError(w, http.StatusInternalServerError, codeNotImplemented, err.Error()) - return - } - - p.logStatement(state.statement, summarizeBridge(envelope, state.started)) - - summary, _ := json.Marshal(map[string]string{ - "read_rows": strconv.FormatUint(envelope.Statistics.RowsRead, 10), - "read_bytes": strconv.FormatUint(envelope.Statistics.BytesRead, 10), - "result_rows": strconv.Itoa(envelope.Rows), - }) - - w.Header().Set("Content-Type", "application/json; charset=UTF-8") - w.Header().Set("X-ClickHouse-Summary", string(summary)) - w.Header().Set("Content-Length", strconv.Itoa(len(encoded))) - w.WriteHeader(http.StatusOK) - _, _ = w.Write(encoded) -} - -func bridgeParameters(r *http.Request) []proto.Parameter { - var parameters []proto.Parameter - for name, values := range r.URL.Query() { - if !strings.HasPrefix(name, "param_") || len(values) == 0 { - continue - } - parameters = append(parameters, proto.Parameter{ - Key: strings.TrimPrefix(name, "param_"), - Value: quoteFieldDump(values[0]), - }) - } - return parameters -} - -// Native carries a parameter as a custom setting, read as a Field dump rather than the plain string HTTP takes. -func quoteFieldDump(value string) string { - var quoted strings.Builder - quoted.Grow(len(value) + 2) - quoted.WriteByte('\'') - // Byte-wise: ranging would turn every invalid byte into U+FFFD and silently change the value. - for i := 0; i < len(value); i++ { - if c := value[i]; c == '\\' || c == '\'' { - quoted.WriteByte('\\') - } - quoted.WriteByte(value[i]) - } - quoted.WriteByte('\'') - return quoted.String() -} - -// Peels off the trailing FORMAT clause, which decides the envelope. Scans the original bytes: uppercasing -// first would shift offsets, because some runes shrink when folded. -func splitFormatClause(sql string) (string, string) { - trimmed := trimTrailingSemicolons(sql) - - idx := lastIndexFold(trimmed, "FORMAT") - if idx <= 0 { - return trimmed, "" - } - // Boundary on both sides, or a trailing identifier such as `format_events` is read as the clause. - if !isSQLSpace(trimmed[idx-1]) { - return trimmed, "" - } - after := trimmed[idx+len("FORMAT"):] - if after == "" || !isSQLSpace(after[0]) { - return trimmed, "" - } - - name := strings.TrimSpace(after) - if !isFormatName(name) { - return trimmed, "" - } - - return trimTrailingSemicolons(trimmed[:idx]), name -} - -// Anything but a bare identifier means FORMAT was part of the statement. -func isFormatName(name string) bool { - if name == "" { - return false - } - for i := 0; i < len(name); i++ { - c := name[i] - if c >= 'a' && c <= 'z' || c >= 'A' && c <= 'Z' || c >= '0' && c <= '9' || c == '_' { - continue - } - return false - } - return true -} - -func isSQLSpace(c byte) bool { - return c == ' ' || c == '\t' || c == '\r' || c == '\n' -} - -// So `SELECT 1 ; ;` does not end up inside the subquery wrapper. -func trimTrailingSemicolons(sql string) string { - trimmed := strings.TrimSpace(sql) - for strings.HasSuffix(trimmed, ";") { - trimmed = strings.TrimSpace(strings.TrimSuffix(trimmed, ";")) - } - return trimmed -} - -func lastIndexFold(s string, substr string) int { - for i := len(s) - len(substr); i >= 0; i-- { - if strings.EqualFold(s[i:i+len(substr)], substr) { - return i - } - } - return -1 -} - -// Rejects a statement that would not stay inside its wrapper. Quotes and comments are tracked so a -// semicolon or bracket inside a string literal is left alone. -func checkSpliceable(body string) error { - depth := 0 - for i := 0; i < len(body); i++ { - switch c := body[i]; c { - case '\'', '"', '`': - end := skipQuoted(body, i, c) - if end < 0 { - return refuseBridge(codeNotImplemented, "this statement has an unterminated string, so the gateway could not read it") - } - i = end - case '-': - if i+1 < len(body) && body[i+1] == '-' { - if idx := strings.IndexByte(body[i:], '\n'); idx != -1 { - i += idx - } else { - i = len(body) - } - } - case '/': - if i+1 < len(body) && body[i+1] == '*' { - idx := strings.Index(body[i+2:], "*/") - if idx == -1 { - return refuseBridge(codeNotImplemented, "this statement has an unterminated comment, so the gateway could not read it") - } - i += 2 + idx + 1 - } - case '(': - depth++ - case ')': - depth-- - if depth < 0 { - return refuseBridge(codeNotImplemented, - "this statement closes more parentheses than it opens, which the gateway cannot run over "+ - "ClickHouse's native protocol") - } - case ';': - return refuseBridge(codeNotImplemented, - "this account reaches ClickHouse over the native protocol, which runs one statement at a "+ - "time, so a semicolon inside a statement is not supported") - } - } - if depth != 0 { - return refuseBridge(codeNotImplemented, - "this statement leaves %d parenthesis open, so the gateway could not run it", depth) - } - return nil -} - -func skipQuoted(body string, start int, quote byte) int { - for i := start + 1; i < len(body); i++ { - switch body[i] { - case '\\': - i++ - case quote: - if i+1 < len(body) && body[i+1] == quote { - i++ - continue - } - return i - } - } - return -1 -} - -// Keeps the error message and the allowlist from drifting apart. -func humanList(items []string) string { - switch len(items) { - case 0: - return "" - case 1: - return items[0] - default: - return strings.Join(items[:len(items)-1], ", ") + " or " + items[len(items)-1] - } -} - -func rowFormatFor(format string) (string, error) { - switch strings.ToUpper(format) { - case strings.ToUpper(formatJSON): - return rowFormatJSON, nil - case strings.ToUpper(formatJSONCompact): - return rowFormatJSONCompact, nil - default: - return "", refuseBridge(codeNotImplemented, - "this account has no HTTP port, so it is reached over ClickHouse's native protocol and only %s and "+ - "%s can be returned. This client asked for %s. Use a native client such as clickhouse-client, "+ - "or give the account an HTTP port", - formatJSON, formatJSONCompact, format) - } -} - -// Forwarded so a browser session costs the server no more over native than over HTTP; anything else is dropped. -var forwardedSettings = map[string]bool{ - "max_execution_time": true, - "max_result_rows": true, - "max_result_bytes": true, - "result_overflow_mode": true, - "max_rows_to_read": true, - "readonly": true, -} - -func bridgeSettings(r *http.Request) []ch.Setting { - var settings []ch.Setting - for name, values := range r.URL.Query() { - if !forwardedSettings[name] || len(values) == 0 { - continue - } - settings = append(settings, ch.Setting{Key: name, Value: values[0], Important: true}) - } - return settings -} - -func (p *ClickHouseProxy) runBridgeQuery( - ctx context.Context, - client *ch.Client, - body string, - format string, - parameters []proto.Parameter, - settings []ch.Setting, -) (*bridgeEnvelope, error) { - envelope := &bridgeEnvelope{Meta: []bridgeColumn{}, Data: json.RawMessage("[]")} - - // Nothing reads the rows, so it only has to run. - if format == "" { - var discard proto.Results - return envelope, client.Do(ctx, ch.Query{ - Body: body, - Parameters: parameters, - Settings: settings, - Result: discard.Auto(), - // ch-go refuses a second data block unless a handler is present. - OnResult: func(context.Context, proto.Block) error { return nil }, - OnProgress: func(_ context.Context, pr proto.Progress) error { - envelope.Statistics.RowsRead += pr.Rows - envelope.Statistics.BytesRead += pr.Bytes - return nil - }, - }) - } - - rowFormat, err := rowFormatFor(format) - if err != nil { - return nil, err - } - - if !isWrappable(body) { - return nil, refuseBridge(codeNotImplemented, - "this account reaches ClickHouse over the native protocol, where the gateway can only return rows "+ - "for a %s. Run this statement from the CLI instead", humanList(wrappableStatements)) - } - - meta, err := describeStatement(ctx, client, body, parameters, settings) - if err != nil { - return nil, err - } - envelope.Meta = meta - - rows, err := selectFormattedRows(ctx, client, body, rowFormat, parameters, settings, envelope) - if err != nil { - return nil, err - } - - envelope.Rows = len(rows) - envelope.Data = json.RawMessage("[" + strings.Join(rows, ",") + "]") - return envelope, nil -} - -// Only these can sit inside a subquery, which both halves of the bridge rely on. -var wrappableStatements = []string{"SELECT", "WITH", "EXPLAIN"} - -func isWrappable(body string) bool { - rest := strings.TrimLeft(stripLeadingNoise(body), "(") - for _, prefix := range wrappableStatements { - if len(rest) < len(prefix) || !strings.EqualFold(rest[:len(prefix)], prefix) { - continue - } - // SELECTFOO is an identifier, not a SELECT. - if len(rest) == len(prefix) || !isIdentifierByte(rest[len(prefix)]) { - return true - } - } - return false -} - -func isIdentifierByte(c byte) bool { - return c >= 'a' && c <= 'z' || c >= 'A' && c <= 'Z' || c >= '0' && c <= '9' || c == '_' -} - -// Editors add leading comments freely, which would otherwise make an ordinary SELECT look unservable. -func stripLeadingNoise(body string) string { - rest := strings.TrimPrefix(body, "\ufeff") - for { - rest = strings.TrimLeft(rest, " \t\r\n") - switch { - case strings.HasPrefix(rest, "--"): - if idx := strings.IndexByte(rest, '\n'); idx != -1 { - rest = rest[idx+1:] - continue - } - return "" - case strings.HasPrefix(rest, "/*"): - idx := strings.Index(rest[2:], "*/") - if idx == -1 { - return "" - } - rest = rest[2+idx+2:] - continue - default: - return rest - } - } -} - -// DESCRIBE resolves the header, so the column types are the server's own. It is not free: schema inference -// for a table function such as url() or s3() does fetch. -func describeStatement( - ctx context.Context, - client *ch.Client, - body string, - parameters []proto.Parameter, - settings []ch.Setting, -) ([]bridgeColumn, error) { - // The column set has grown between versions, so they are picked by name rather than position. - var described proto.Results - columns := []bridgeColumn{} - - err := client.Do(ctx, ch.Query{ - Body: "DESCRIBE (\n" + body + "\n)", - Parameters: parameters, - Settings: settings, - Result: described.Auto(), - OnResult: func(_ context.Context, block proto.Block) error { - names, err := stringColumn(described, "name") - if err != nil { - return err - } - types, err := stringColumn(described, "type") - if err != nil { - return err - } - for i := 0; i < block.Rows; i++ { - columns = append(columns, bridgeColumn{Name: names.Row(i), Type: types.Row(i)}) - } - return nil - }, - }) - if err != nil { - return nil, err - } - return columns, nil -} - -func stringColumn(results proto.Results, name string) (*proto.ColStr, error) { - for _, column := range results { - if column.Name != name { - continue - } - if typed, ok := column.Data.(*proto.ColStr); ok { - return typed, nil - } - return nil, fmt.Errorf("DESCRIBE returned %q as %s rather than String", name, column.Data.Type()) - } - return nil, fmt.Errorf("DESCRIBE returned no %q column", name) -} - -func selectFormattedRows( - ctx context.Context, - client *ch.Client, - body string, - rowFormat string, - parameters []proto.Parameter, - settings []ch.Setting, - envelope *bridgeEnvelope, -) ([]string, error) { - var formatted proto.ColStr - rows := make([]string, 0, 64) - resultBytes := 0 - - err := client.Do(ctx, ch.Query{ - Body: "SELECT formatRowNoNewline('" + rowFormat + "', *) AS " + formattedRowAlias + - " FROM (\n" + body + "\n)", - Parameters: parameters, - Settings: settings, - Result: proto.Results{{Name: formattedRowAlias, Data: &formatted}}, - OnResult: func(_ context.Context, block proto.Block) error { - for i := 0; i < block.Rows; i++ { - if len(rows) >= maxBridgeRows { - return refuseBridge(codeTooManyRows, - "this statement returned more than %d rows, which is more than a browser session "+ - "returns over the native protocol. Add a LIMIT, or use the CLI", maxBridgeRows) - } - row := formatted.Row(i) - resultBytes += len(row) + 1 - if resultBytes > maxBridgeResultBytes { - return refuseBridge(codeTooManyRows, - "this statement returned more than %d MB, which is more than a browser session "+ - "returns over the native protocol. Narrow the result, or use the CLI", - maxBridgeResultBytes>>20) - } - rows = append(rows, row) - } - return nil - }, - OnProgress: func(_ context.Context, pr proto.Progress) error { - envelope.Statistics.RowsRead += pr.Rows - envelope.Statistics.BytesRead += pr.Bytes - return nil - }, - }) - if err != nil { - return nil, err - } - return rows, nil -} - -// A refusal the gateway made itself, so it reports its own code rather than ch-go's decoding context. -type bridgeRefusal struct { - code int - message string -} - -func (e *bridgeRefusal) Error() string { return e.message } - -func refuseBridge(code int, format string, args ...any) error { - return &bridgeRefusal{code: code, message: fmt.Sprintf(format, args...)} -} - -func classifyNativeError(err error) (int, int, string) { - var refusal *bridgeRefusal - if errors.As(err, &refusal) { - return refusal.code, http.StatusBadRequest, refusal.message - } - if exception, ok := ch.AsException(err); ok { - return int(exception.Code), http.StatusBadRequest, exception.Message - } - return codeNetworkError, http.StatusBadGateway, err.Error() -} - -func summarizeBridge(envelope *bridgeEnvelope, started time.Time) string { - parts := []string{"200 OK"} - if envelope.Rows > 0 { - parts = append(parts, fmt.Sprintf("%d row(s) returned", envelope.Rows)) - } - if envelope.Statistics.RowsRead > 0 { - parts = append(parts, fmt.Sprintf("%d row(s) read", envelope.Statistics.RowsRead)) - } - return strings.Join(append(parts, fmt.Sprintf("%dms", time.Since(started).Milliseconds())), ", ") -} diff --git a/packages/pam/handlers/clickhouse/bridge_test.go b/packages/pam/handlers/clickhouse/bridge_test.go deleted file mode 100644 index 49ff790b2..000000000 --- a/packages/pam/handlers/clickhouse/bridge_test.go +++ /dev/null @@ -1,457 +0,0 @@ -package clickhouse - -import ( - "context" - "encoding/json" - "fmt" - "io" - "net" - "net/http" - "os" - "strings" - "testing" - "time" - - "github.com/stretchr/testify/require" -) - -func TestSplitFormatClause(t *testing.T) { - cases := []struct { - name string - sql string - wantBody string - wantFormat string - }{ - { - name: "the clause the node client appends", - sql: "SELECT 1 \nFORMAT JSON", - wantBody: "SELECT 1", - wantFormat: "JSON", - }, - {name: "trailing semicolon", sql: "SELECT 1 FORMAT JSONCompact;", wantBody: "SELECT 1", wantFormat: "JSONCompact"}, - {name: "no clause", sql: "SELECT 1", wantBody: "SELECT 1", wantFormat: ""}, - { - // "format" inside the statement is not the clause, and peeling it off would corrupt the query. - name: "a call to formatDateTime is not a clause", - sql: "SELECT formatDateTime(now(), '%F')", - wantBody: "SELECT formatDateTime(now(), '%F')", - }, - {name: "a column named format", sql: "SELECT format FROM t", wantBody: "SELECT format FROM t"}, - {name: "format with no name", sql: "SELECT 1 FORMAT", wantBody: "SELECT 1 FORMAT"}, - // A trailing identifier that merely starts with "format" is not a clause, and splitting it would throw... - {name: "a table whose name starts with format", sql: "SELECT * FROM format_events", wantBody: "SELECT * FROM format_events"}, - {name: "an alias that starts with format", sql: "SELECT 1 AS format_id", wantBody: "SELECT 1 AS format_id"}, - {name: "ordering by a column called formatted", sql: "SELECT x FROM t ORDER BY formatted", wantBody: "SELECT x FROM t ORDER BY formatted"}, - {name: "lowercase clause", sql: "select 1 format json", wantBody: "select 1", wantFormat: "json"}, - {name: "mixed case clause", sql: "SELECT 1 FoRmAt JSONCompact", wantBody: "SELECT 1", wantFormat: "JSONCompact"}, - {name: "format alone is not a clause", sql: "FORMAT JSON", wantBody: "FORMAT JSON"}, - {name: "format inside a string literal", sql: "SELECT 'FORMAT JSON'", wantBody: "SELECT 'FORMAT JSON'"}, - {name: "several trailing semicolons", sql: "SELECT 1 ; ;", wantBody: "SELECT 1"}, - // A rune that shrinks when uppercased would shift byte offsets if the scan ran over a folded copy. - {name: "a value whose uppercase form is shorter", sql: "SELECT 'ı' AS x FORMAT JSON", wantBody: "SELECT 'ı' AS x", wantFormat: "JSON"}, - } - - for _, tc := range cases { - t.Run(tc.name, func(t *testing.T) { - body, format := splitFormatClause(tc.sql) - require.Equal(t, tc.wantBody, body) - require.Equal(t, tc.wantFormat, format) - }) - } -} - -func TestIsWrappable(t *testing.T) { - for _, body := range []string{"SELECT 1", " select 1", "WITH x AS (SELECT 1) SELECT * FROM x", "EXPLAIN SELECT 1"} { - require.True(t, isWrappable(body), body) - } - for _, body := range []string{"SHOW TABLES", "DESCRIBE TABLE t", "INSERT INTO t VALUES (1)", "CREATE TABLE t (a UInt8) ENGINE = Memory"} { - require.False(t, isWrappable(body), body) - } - - // Editors put comments and byte-order marks in front of perfectly ordinary statements. - for _, body := range []string{"-- a note\nSELECT 1", "/* a note */ SELECT 1", "\ufeffSELECT 1", "/* a */ -- b\n select 1"} { - require.True(t, isWrappable(body), body) - } - // A prefix match is not a keyword match. - require.False(t, isWrappable("SELECTFOO 1")) - require.False(t, isWrappable("WITHOUT_ROWS()")) -} - -func TestCheckSpliceable(t *testing.T) { - ok := []string{ - "SELECT 1", - "SELECT (1 + 2) AS x", - "SELECT 'a;b) --' AS s", - "SELECT \"col;)\" FROM t", - "SELECT 1 -- a trailing ; comment", - "SELECT 1 /* a ) comment */", - } - for _, body := range ok { - require.NoError(t, checkSpliceable(body), body) - } - - // Each of these would otherwise run something other than what was recorded and policy-checked. - bad := []string{ - "SELECT 1) ; DROP TABLE users; --", - "SELECT 1) UNION ALL (SELECT 2", - "SELECT 1; SELECT 2", - "SELECT (1", - "SELECT 'unterminated", - } - for _, body := range bad { - require.Error(t, checkSpliceable(body), body) - } -} - -func TestBridgeMatchesHTTPInterface(t *testing.T) { - if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { - t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") - } - - statements := []string{ - "SELECT 1 AS n, 'x' AS s", - "SELECT map('a', 1::UInt64) AS m, tuple('p', 2) AS t", - "SELECT number, toString(number) AS s FROM numbers(5)", - "SELECT id, note FROM pam_write_test ORDER BY id LIMIT 3", - "SELECT * FROM exotic ORDER BY id", - "SELECT count() AS c FROM users", - "SELECT NULL::Nullable(String) AS nothing", - "WITH 2 AS x SELECT x * 3 AS y", - } - - for _, format := range []string{formatJSON, formatJSONCompact} { - for _, statement := range statements { - t.Run(format+": "+statement, func(t *testing.T) { - viaBridge := queryBridge(t, statement, format) - viaHTTP := queryRealHTTP(t, statement, format) - - require.Equal(t, viaHTTP.Meta, viaBridge.Meta, "column metadata should match ClickHouse") - require.Equal(t, viaHTTP.Rows, viaBridge.Rows, "row count should match ClickHouse") - require.JSONEq(t, string(viaHTTP.Data), string(viaBridge.Data), "rows should match ClickHouse") - }) - } - } -} - -func TestBridgeReportsClickHouseErrors(t *testing.T) { - if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { - t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") - } - - addr := startBridgeProxy(t, nil, &recordingLogger{}) - status, body := postStatement(t, addr, "SELECT * FROM table_that_is_not_there FORMAT JSON") - - require.Equal(t, http.StatusBadRequest, status) - require.Contains(t, body, "table_that_is_not_there") - require.Contains(t, body, "DB::Exception") -} - -func TestBridgeAppliesCommandBlocking(t *testing.T) { - if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { - t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") - } - - recorder := &recordingLogger{} - addr := startBridgeProxy(t, []string{`(?i)\bdrop\b`}, recorder) - - status, body := postStatement(t, addr, "DROP TABLE pam_write_test FORMAT JSON") - require.Equal(t, http.StatusForbidden, status) - require.Contains(t, body, "blocked by the command blocking policy") - require.True(t, recorder.contains("DROP TABLE pam_write_test")) -} - -func TestBridgeRecordsStatements(t *testing.T) { - if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { - t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") - } - - recorder := &recordingLogger{} - addr := startBridgeProxy(t, nil, recorder) - - status, _ := postStatement(t, addr, "SELECT 42 AS answer FORMAT JSON") - require.Equal(t, http.StatusOK, status) - require.True(t, recorder.contains("SELECT 42 AS answer"), recorder.dump()) - require.Contains(t, recorder.dump(), "1 row(s) returned") -} - -func startBridgeProxy(t *testing.T, blocked []string, logger *recordingLogger) string { - t.Helper() - - patterns := compileForTest(t, blocked) - proxy := NewClickHouseProxy(ClickHouseProxyConfig{ - NativeAddr: envOr("PAM_CLICKHOUSE_NATIVE", "127.0.0.1:9000"), - Username: envOr("PAM_CLICKHOUSE_USER", "default"), - Password: envOr("PAM_CLICKHOUSE_PASSWORD", "clickhouse"), - Database: envOr("PAM_CLICKHOUSE_DB", "analytics"), - SessionID: "bridge-test", - SessionLogger: logger, - BlockedCommands: patterns, - }) - - listener, err := net.Listen("tcp", "127.0.0.1:0") - require.NoError(t, err) - t.Cleanup(func() { listener.Close() }) - - ctx, cancel := context.WithCancel(context.Background()) - t.Cleanup(cancel) - - go func() { - for { - conn, err := listener.Accept() - if err != nil { - return - } - go func() { _ = proxy.HandleConnection(ctx, conn) }() - } - }() - - return listener.Addr().String() -} - -func postStatement(t *testing.T, addr string, sql string) (int, string) { - t.Helper() - - req, err := http.NewRequest(http.MethodPost, "http://"+addr+"/", strings.NewReader(sql)) - require.NoError(t, err) - - resp, err := (&http.Client{Timeout: 60 * time.Second}).Do(req) - require.NoError(t, err) - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - return resp.StatusCode, string(body) -} - -func queryBridge(t *testing.T, statement string, format string) bridgeEnvelope { - t.Helper() - - addr := startBridgeProxy(t, nil, &recordingLogger{}) - status, body := postStatement(t, addr, statement+" \nFORMAT "+format) - require.Equal(t, http.StatusOK, status, body) - - var envelope bridgeEnvelope - require.NoError(t, json.Unmarshal([]byte(body), &envelope), body) - return envelope -} - -func queryRealHTTP(t *testing.T, statement string, format string) bridgeEnvelope { - t.Helper() - - target := fmt.Sprintf("http://%s/?database=%s", - envOr("PAM_CLICKHOUSE_HTTP", "127.0.0.1:8123"), envOr("PAM_CLICKHOUSE_DB", "analytics")) - - req, err := http.NewRequest(http.MethodPost, target, strings.NewReader(statement+" \nFORMAT "+format)) - require.NoError(t, err) - req.Header.Set("X-ClickHouse-User", envOr("PAM_CLICKHOUSE_USER", "default")) - req.Header.Set("X-ClickHouse-Key", envOr("PAM_CLICKHOUSE_PASSWORD", "clickhouse")) - - resp, err := (&http.Client{Timeout: 60 * time.Second}).Do(req) - require.NoError(t, err) - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - require.Equal(t, http.StatusOK, resp.StatusCode, string(body)) - - var envelope bridgeEnvelope - require.NoError(t, json.Unmarshal(body, &envelope), string(body)) - return envelope -} - -func TestBridgeAgainstHTTPDisabledServer(t *testing.T) { - if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { - t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") - } - - native := envOr("PAM_CLICKHOUSE_NOHTTP_NATIVE", "127.0.0.1:19001") - probe, err := net.DialTimeout("tcp", native, 2*time.Second) - if err != nil { - t.Skipf("no HTTP-disabled ClickHouse on %s: %v", native, err) - } - probe.Close() - - bridged := NewClickHouseProxy(ClickHouseProxyConfig{ - NativeAddr: native, - Username: "default", - Password: envOr("PAM_CLICKHOUSE_NOHTTP_PASSWORD", "clickhouse"), - Database: "default", - SessionID: "bridge-nohttp-test", - SessionLogger: &recordingLogger{}, - }) - - listener, err := net.Listen("tcp", "127.0.0.1:0") - require.NoError(t, err) - defer listener.Close() - - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() - go func() { - for { - conn, acceptErr := listener.Accept() - if acceptErr != nil { - return - } - go func() { _ = bridged.HandleConnection(ctx, conn) }() - } - }() - - status, body := postStatement(t, listener.Addr().String(), - "SELECT 1 AS n, map('k', 'v') AS m \nFORMAT JSON") - require.Equal(t, http.StatusOK, status, body) - - var envelope bridgeEnvelope - require.NoError(t, json.Unmarshal([]byte(body), &envelope), body) - require.Equal(t, 1, envelope.Rows) - require.Equal(t, []bridgeColumn{{Name: "n", Type: "UInt8"}, {Name: "m", Type: "Map(String, String)"}}, envelope.Meta) - require.JSONEq(t, `[{"n":1,"m":{"k":"v"}}]`, string(envelope.Data)) -} - -func TestBridgeHonoursDefaultFormatSetting(t *testing.T) { - if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { - t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") - } - - addr := startBridgeProxy(t, nil, &recordingLogger{}) - - for _, format := range []string{formatJSON, formatJSONCompact} { - t.Run(format, func(t *testing.T) { - status, body := postGET(t, addr, - "/?default_format="+format+"&max_execution_time=30&query="+ - urlEscape("SELECT id, tags FROM exotic ORDER BY id")) - require.Equal(t, http.StatusOK, status, body) - - var envelope bridgeEnvelope - require.NoError(t, json.Unmarshal([]byte(body), &envelope), body) - require.Equal(t, 1, envelope.Rows) - require.Equal(t, []bridgeColumn{ - {Name: "id", Type: "UInt64"}, - {Name: "tags", Type: "Map(String, UInt64)"}, - }, envelope.Meta) - }) - } -} - -func TestBridgeHandlesATrailingComment(t *testing.T) { - if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { - t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") - } - - addr := startBridgeProxy(t, nil, &recordingLogger{}) - status, body := postStatement(t, addr, "SELECT 1 AS n\n-- a trailing note\nFORMAT JSON") - require.Equal(t, http.StatusOK, status, body) - require.Contains(t, body, `"n":1`) -} - -func TestBridgeHandlesAMultiBlockResult(t *testing.T) { - if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { - t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") - } - - addr := startBridgeProxy(t, nil, &recordingLogger{}) - - // Comfortably more than one block (max_block_size defaults to ~65k) and under the row cap. - const rows = 80000 - - t.Run("with a format", func(t *testing.T) { - status, body := postStatement(t, addr, fmt.Sprintf("SELECT number FROM numbers(%d) \nFORMAT JSONCompact", rows)) - require.Equal(t, http.StatusOK, status, body[:min(len(body), 400)]) - - var envelope bridgeEnvelope - require.NoError(t, json.Unmarshal([]byte(body), &envelope)) - require.Equal(t, rows, envelope.Rows) - }) - - t.Run("without a format", func(t *testing.T) { - status, body := postStatement(t, addr, fmt.Sprintf("SELECT number FROM numbers(%d)", rows)) - require.Equal(t, http.StatusOK, status, body) - }) -} - -func TestBridgeRefusesAResultBeyondTheRowCap(t *testing.T) { - if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { - t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") - } - - addr := startBridgeProxy(t, nil, &recordingLogger{}) - status, body := postStatement(t, addr, - fmt.Sprintf("SELECT number FROM numbers(%d) \nFORMAT JSONCompact", maxBridgeRows+1000)) - - require.Equal(t, http.StatusBadRequest, status) - require.Contains(t, body, "more than 100000 rows") - // The refusal is the gateway's own, so it must not be dressed up as a network error. - require.Contains(t, body, "TOO_MANY_ROWS") - require.NotContains(t, body, "decode block", "ch-go's internal wrapping should not reach the client") -} - -func TestBridgeAnswersPing(t *testing.T) { - if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { - t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") - } - - addr := startBridgeProxy(t, nil, &recordingLogger{}) - status, body := postGET(t, addr, "/ping") - require.Equal(t, http.StatusOK, status) - require.Contains(t, body, "Ok.") -} - -func TestBridgeParameterRoundTrip(t *testing.T) { - if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { - t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") - } - - addr := startBridgeProxy(t, nil, &recordingLogger{}) - - for _, value := range []string{"plain", "it's", `back\slash`, "a,b", " spaced ", "ünïcødé", "0", ""} { - t.Run(fmt.Sprintf("%q", value), func(t *testing.T) { - status, body := postGET(t, addr, - "/?param_v="+urlEscape(value)+"&query="+urlEscape("SELECT {v:String} AS got FORMAT JSON")) - require.Equal(t, http.StatusOK, status, body) - - var envelope bridgeEnvelope - require.NoError(t, json.Unmarshal([]byte(body), &envelope), body) - - var rows []struct { - Got string `json:"got"` - } - require.NoError(t, json.Unmarshal(envelope.Data, &rows)) - require.Len(t, rows, 1) - require.Equal(t, value, rows[0].Got) - }) - } -} - -// @clickhouse/client.insert() sends `INSERT INTO t FORMAT JSONEachRow` with the rows in the body, which... -func TestBridgeInsertWithInlineData(t *testing.T) { - if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { - t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") - } - - addr := startBridgeProxy(t, nil, &recordingLogger{}) - marker := fmt.Sprintf("bridge-insert-%d", time.Now().UnixNano()) - - type result struct { - status int - body string - } - done := make(chan result, 1) - go func() { - status, body, err := postStatementE(addr, - fmt.Sprintf("INSERT INTO pam_write_test (id, note) FORMAT JSONEachRow\n{\"id\":9100,\"note\":%q}", marker)) - if err != nil { - done <- result{status: -1, body: err.Error()} - return - } - done <- result{status: status, body: body} - }() - - select { - case got := <-done: - require.Equal(t, http.StatusOK, got.status, got.body) - case <-time.After(45 * time.Second): - t.Fatal("the bridge hung on an INSERT carrying inline data") - } - - // An empty 200 that quietly inserted nothing would be worse than hanging. - require.Equal(t, "1", queryDirect(t, fmt.Sprintf("SELECT count() FROM pam_write_test WHERE note = '%s'", marker))) -} diff --git a/packages/pam/handlers/clickhouse/edge_cases_test.go b/packages/pam/handlers/clickhouse/edge_cases_test.go index 30a485fbf..088d08ac9 100644 --- a/packages/pam/handlers/clickhouse/edge_cases_test.go +++ b/packages/pam/handlers/clickhouse/edge_cases_test.go @@ -4,7 +4,6 @@ import ( "context" "crypto/tls" "crypto/x509" - "encoding/json" "fmt" "io" "net" @@ -253,20 +252,6 @@ func TestTLSUpstream(t *testing.T) { require.NoError(t, TestNativeConnection(context.Background(), config)) }) - t.Run("bridging over TLS for a server with no HTTP", func(t *testing.T) { - bridged := config - bridged.TargetAddr = "" - port := startProxy(t, bridged) - - status, body := postStatement(t, "127.0.0.1:"+port, "SELECT id, note FROM t ORDER BY id \nFORMAT JSON") - require.Equal(t, http.StatusOK, status, body) - - var envelope bridgeEnvelope - require.NoError(t, json.Unmarshal([]byte(body), &envelope), body) - require.Equal(t, 2, envelope.Rows) - require.JSONEq(t, `[{"id":1,"note":"tls-one"},{"id":2,"note":"tls-two"}]`, string(envelope.Data)) - }) - t.Run("a pinned CA verifies rather than skipping", func(t *testing.T) { if os.Getenv("PAM_CLICKHOUSE_TLS_CERT") == "" { t.Skip("set PAM_CLICKHOUSE_TLS_CERT to run") @@ -302,22 +287,6 @@ func TestAccountWithoutNativePortRefusesNativeClients(t *testing.T) { require.Equal(t, http.StatusOK, status, body) } -func TestAccountWithoutHTTPPortStillServesBothClients(t *testing.T) { - itOnly(t) - - config := baseConfig(&recordingLogger{}) - config.TargetAddr = "" - port := startProxy(t, config) - - out, err := runClient(t, port, "SELECT 'native-on-bridged-account';") - require.NoError(t, err, out) - require.Contains(t, out, "native-on-bridged-account") - - status, body := postStatement(t, "127.0.0.1:"+port, "SELECT 1 AS n \nFORMAT JSON") - require.Equal(t, http.StatusOK, status, body) - require.Contains(t, body, `"n":1`) -} - func TestAccountWithNeitherPortFailsClearly(t *testing.T) { itOnly(t) @@ -331,100 +300,14 @@ func TestAccountWithNeitherPortFailsClearly(t *testing.T) { require.Error(t, err, "a session with neither port must not serve anything") } -func TestBridgeEdgeCases(t *testing.T) { - itOnly(t) - - config := baseConfig(&recordingLogger{}) - config.TargetAddr = "" - - cases := []struct { - name string - sql string - wantStatus int - wantBody string - }{ - { - name: "a statement shape that cannot be wrapped says so", - sql: "SHOW TABLES \nFORMAT JSON", - wantStatus: http.StatusBadRequest, - wantBody: "SELECT, WITH or EXPLAIN", - }, - { - name: "a format the bridge does not produce says so", - sql: "SELECT 1 \nFORMAT TabSeparated", - wantStatus: http.StatusBadRequest, - wantBody: "JSON and JSONCompact", - }, - { - name: "a statement with no format runs and returns nothing to parse", - sql: "CREATE TABLE IF NOT EXISTS bridge_ddl (a UInt8) ENGINE = Memory", - wantStatus: http.StatusOK, - }, - { - name: "an empty statement is refused", - sql: "", - wantStatus: http.StatusBadRequest, - wantBody: "No statement was sent", - }, - { - name: "a syntax error comes back as ClickHouse wrote it", - sql: "SELECT FROM WHERE \nFORMAT JSON", - wantStatus: http.StatusBadRequest, - wantBody: "Syntax error", - }, - } - - for _, tc := range cases { - t.Run(tc.name, func(t *testing.T) { - port := startProxy(t, config) - status, body := postStatement(t, "127.0.0.1:"+port, tc.sql) - require.Equal(t, tc.wantStatus, status, body) - if tc.wantBody != "" { - require.Contains(t, body, tc.wantBody) - } - }) - } -} - -func TestBridgePassesQueryParameters(t *testing.T) { - itOnly(t) - - recorder := &recordingLogger{} - config := baseConfig(recorder) - config.TargetAddr = "" - port := startProxy(t, config) - - status, body := postGET(t, "127.0.0.1:"+port, - "/?param_wanted=7&query="+urlEscape("SELECT {wanted:UInt8} AS got FORMAT JSON")) - require.Equal(t, http.StatusOK, status, body) - require.Contains(t, body, `"got":7`) - require.Contains(t, recorder.dump(), "wanted=7", "the parameter belongs in the recording") -} - func TestCompressedRequestBodies(t *testing.T) { itOnly(t) - for _, native := range []bool{false, true} { - name := "http interface" - if native { - name = "bridged to native" - } - t.Run(name, func(t *testing.T) { - config := baseConfig(&recordingLogger{}) - if native { - config.TargetAddr = "" - } - port := startProxy(t, config) - - status, body := postGzipped(t, "127.0.0.1:"+port, "SELECT 5 AS five \nFORMAT JSON") - require.Equal(t, http.StatusOK, status, body) + port := startProxy(t, baseConfig(&recordingLogger{})) - // ClickHouse pretty-prints its JSON and the bridge writes it compact, so the rows are compared rather than... - var envelope bridgeEnvelope - require.NoError(t, json.Unmarshal([]byte(body), &envelope), body) - require.JSONEq(t, `[{"five":5}]`, string(envelope.Data)) - }) - } + status, body := postGzipped(t, "127.0.0.1:"+port, "SELECT 5 AS five \nFORMAT JSON") + require.Equal(t, http.StatusOK, status, body) + require.Contains(t, body, "\"five\": 5") } func TestSnifferEdgeCases(t *testing.T) { @@ -585,16 +468,6 @@ func TestUpstreamUnreachable(t *testing.T) { require.Contains(t, out, "could not reach ClickHouse") }) - t.Run("bridge reports a bad gateway", func(t *testing.T) { - config := baseConfig(&recordingLogger{}) - config.TargetAddr = "" - config.NativeAddr = "127.0.0.1:1" - port := startProxy(t, config) - - status, body := postStatement(t, "127.0.0.1:"+port, "SELECT 1 \nFORMAT JSON") - require.Equal(t, http.StatusBadGateway, status) - require.Contains(t, body, "could not reach ClickHouse") - }) } func TestWrongAccountCredentialsSurfaceCleanly(t *testing.T) { @@ -638,46 +511,3 @@ func TestNativeRevisionPinning(t *testing.T) { require.NoError(t, err, out) require.Equal(t, queryDirect(t, "SELECT version()"), strings.TrimSpace(out)) } - -func TestQuoteFieldDump(t *testing.T) { - cases := []struct{ in, want string }{ - {"7", `'7'`}, - {"plain", `'plain'`}, - {"it's", `'it\'s'`}, - {`back\slash`, `'back\\slash'`}, - {`'; DROP TABLE users; --`, `'\'; DROP TABLE users; --'`}, - {"", `''`}, - } - for _, tc := range cases { - require.Equal(t, tc.want, quoteFieldDump(tc.in), tc.in) - } -} - -// A parameter is data, so a value full of quotes has to come back as that value rather than changing the... -func TestBridgeParameterCannotEscapeItsQuotes(t *testing.T) { - itOnly(t) - - config := baseConfig(&recordingLogger{}) - config.TargetAddr = "" - port := startProxy(t, config) - - hostile := `'; DROP TABLE pam_write_test; --` - status, body := postGET(t, "127.0.0.1:"+port, - "/?param_v="+urlEscape(hostile)+"&query="+urlEscape("SELECT {v:String} AS got FORMAT JSON")) - - require.Equal(t, http.StatusOK, status, body) - - var envelope bridgeEnvelope - require.NoError(t, json.Unmarshal([]byte(body), &envelope), body) - - var rows []struct { - Got string `json:"got"` - } - require.NoError(t, json.Unmarshal(envelope.Data, &rows)) - require.Len(t, rows, 1) - require.Equal(t, hostile, rows[0].Got, "the value should survive intact, not be executed") - - // The table the payload tried to drop is still there. - okStatus, okBody := postStatement(t, "127.0.0.1:"+port, "SELECT count() AS c FROM pam_write_test \nFORMAT JSON") - require.Equal(t, http.StatusOK, okStatus, okBody) -} diff --git a/packages/pam/handlers/clickhouse/proxy.go b/packages/pam/handlers/clickhouse/proxy.go index 2cc213fa9..f4634b15b 100644 --- a/packages/pam/handlers/clickhouse/proxy.go +++ b/packages/pam/handlers/clickhouse/proxy.go @@ -59,14 +59,12 @@ const ( const ( codeNotImplemented = 48 codeNetworkError = 210 - codeTooManyRows = 396 codeAccessDenied = 497 ) var errorNames = map[int]string{ codeNotImplemented: "NOT_IMPLEMENTED", codeNetworkError: "NETWORK_ERROR", - codeTooManyRows: "TOO_MANY_ROWS", codeAccessDenied: "ACCESS_DENIED", } @@ -97,9 +95,6 @@ type stateKey struct{} type requestState struct { statement string - // The bridge runs this rather than re-reading a body that may be compressed. - sql string - truncated bool started time.Time } @@ -201,12 +196,11 @@ func (p *ClickHouseProxy) handler(l zerolog.Logger) http.Handler { return } - inspected, body, err := p.inspect(r) + statement, body, err := p.inspect(r) if err != nil { writeClickHouseError(w, http.StatusBadRequest, codeNotImplemented, err.Error()) return } - statement := inspected.statement if blocked := p.blockedBy(statement); blocked != nil { p.logStatement(statement, fmt.Sprintf("BLOCKED: %s", blocked.String())) @@ -217,19 +211,7 @@ func (p *ClickHouseProxy) handler(l zerolog.Logger) http.Handler { } r.Body = body - state := &requestState{ - statement: statement, - sql: inspected.sql, - truncated: inspected.truncated, - started: time.Now(), - } - - // A server with HTTP disabled still has to serve Web Access, which only speaks HTTP. - if p.config.TargetAddr == "" { - p.serveBridge(w, r, state, l) - return - } - + state := &requestState{statement: statement, started: time.Now()} p.reverse.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), stateKey{}, state))) }) } @@ -239,26 +221,17 @@ type bodyReadCloser struct { io.Closer } -type inspectedRequest struct { - statement string - sql string - truncated bool -} - // Returns the statement and a body that still replays in full. ClickHouse concatenates `query` and the body. -func (p *ClickHouseProxy) inspect(r *http.Request) (inspectedRequest, io.ReadCloser, error) { +func (p *ClickHouseProxy) inspect(r *http.Request) (string, io.ReadCloser, error) { queryParam := strings.TrimSpace(r.URL.Query().Get("query")) if r.Body == nil || r.ContentLength == 0 { - return inspectedRequest{ - statement: queryParam + parameterSuffix(r.URL.Query()), - sql: queryParam, - }, http.NoBody, nil + return queryParam + parameterSuffix(r.URL.Query()), http.NoBody, nil } // ClickHouse's own block compression is opaque to anything but a ClickHouse client if r.URL.Query().Get("decompress") == "1" { - return inspectedRequest{}, nil, fmt.Errorf( + return "", nil, fmt.Errorf( "this session cannot read a ClickHouse-compressed request body, so decompress=1 is not supported here. " + "Send the statement uncompressed or with Content-Encoding: gzip") } @@ -267,7 +240,7 @@ func (p *ClickHouseProxy) inspect(r *http.Request) (inspectedRequest, io.ReadClo switch encoding { case "", "identity", "gzip", "deflate": default: - return inspectedRequest{}, nil, fmt.Errorf( + return "", nil, fmt.Errorf( "this session cannot read a %q-encoded request body, so the command blocking policy could not be applied to it. "+ "Use gzip, deflate, or no compression", encoding) } @@ -276,7 +249,7 @@ func (p *ClickHouseProxy) inspect(r *http.Request) (inspectedRequest, io.ReadClo head := make([]byte, maxInspectBytes+1) n, err := io.ReadFull(r.Body, head) if err != nil && err != io.EOF && err != io.ErrUnexpectedEOF { - return inspectedRequest{}, nil, fmt.Errorf("the gateway could not read the request body: %v", err) + return "", nil, fmt.Errorf("the gateway could not read the request body: %v", err) } head = head[:n] @@ -284,24 +257,18 @@ func (p *ClickHouseProxy) inspect(r *http.Request) (inspectedRequest, io.ReadClo decoded, decodedOverflow, decodeErr := decodeHead(head, encoding) if decodeErr != nil { - return inspectedRequest{}, nil, fmt.Errorf( + return "", nil, fmt.Errorf( "the gateway could not decompress the request body to apply the command blocking policy: %v", decodeErr) } - truncated := len(head) > maxInspectBytes || decodedOverflow - if truncated && len(p.config.BlockedCommands) > 0 { - return inspectedRequest{}, nil, fmt.Errorf( + if (len(head) > maxInspectBytes || decodedOverflow) && len(p.config.BlockedCommands) > 0 { + return "", nil, fmt.Errorf( "this account blocks commands, so a request body larger than %d MB is refused: the gateway has to read "+ "the whole statement to apply the policy. Send the data in smaller batches", maxInspectBytes>>20) } - sql := joinStatement(queryParam, string(decoded)) - return inspectedRequest{ - statement: sql + parameterSuffix(r.URL.Query()), - sql: sql, - truncated: truncated, - }, forwarded, nil + return joinStatement(queryParam, string(decoded)) + parameterSuffix(r.URL.Query()), forwarded, nil } func parameterSuffix(query url.Values) string { diff --git a/packages/pam/handlers/clickhouse/testhelpers_test.go b/packages/pam/handlers/clickhouse/testhelpers_test.go index c659fb4e7..459c90cdc 100644 --- a/packages/pam/handlers/clickhouse/testhelpers_test.go +++ b/packages/pam/handlers/clickhouse/testhelpers_test.go @@ -56,6 +56,14 @@ func postStatementE(addr string, sql string) (int, string, error) { return resp.StatusCode, string(body), nil } +func postStatement(t *testing.T, addr string, sql string) (int, string) { + t.Helper() + + status, body, err := postStatementE(addr, sql) + require.NoError(t, err) + return status, body +} + func postGET(t *testing.T, addr string, path string) (int, string) { t.Helper() From b43bb35226a6c87cc00bd9c236f349d5ece0ad5c Mon Sep 17 00:00:00 2001 From: bernie-g Date: Fri, 25 Sep 2026 16:10:16 -0400 Subject: [PATCH 04/16] fix(clickhouse): refuse http clearly on a native-only account --- .../pam/handlers/clickhouse/edge_cases_test.go | 18 ++++++++++++++++++ packages/pam/handlers/clickhouse/proxy.go | 10 ++++++++++ 2 files changed, 28 insertions(+) diff --git a/packages/pam/handlers/clickhouse/edge_cases_test.go b/packages/pam/handlers/clickhouse/edge_cases_test.go index 088d08ac9..fc4025cbd 100644 --- a/packages/pam/handlers/clickhouse/edge_cases_test.go +++ b/packages/pam/handlers/clickhouse/edge_cases_test.go @@ -511,3 +511,21 @@ func TestNativeRevisionPinning(t *testing.T) { require.NoError(t, err, out) require.Equal(t, queryDirect(t, "SELECT version()"), strings.TrimSpace(out)) } + +func TestAccountWithoutHTTPPortRefusesHTTPClients(t *testing.T) { + itOnly(t) + + config := baseConfig(&recordingLogger{}) + config.TargetAddr = "" + port := startProxy(t, config) + + status, body := postStatement(t, "127.0.0.1:"+port, "SELECT 1") + require.Equal(t, http.StatusBadGateway, status) + require.Contains(t, body, "does not have ClickHouse's HTTP port configured") + require.NotContains(t, body, "no Host in request URL", "the internal proxy error should not reach the client") + + // The native protocol still works on the same port. + out, err := runClient(t, port, "SELECT 'native-still-works';") + require.NoError(t, err, out) + require.Contains(t, out, "native-still-works") +} diff --git a/packages/pam/handlers/clickhouse/proxy.go b/packages/pam/handlers/clickhouse/proxy.go index f4634b15b..a6ae35960 100644 --- a/packages/pam/handlers/clickhouse/proxy.go +++ b/packages/pam/handlers/clickhouse/proxy.go @@ -210,6 +210,16 @@ func (p *ClickHouseProxy) handler(l zerolog.Logger) http.Handler { return } + // Without an HTTP upstream the reverse proxy would fail on an empty host, which reads as a network + // fault rather than an account that does not serve this protocol. + if p.config.TargetAddr == "" { + l.Info().Msg("Refused an HTTP connection on an account with no HTTP port") + writeClickHouseError(w, http.StatusBadGateway, codeNotImplemented, + "This account does not have ClickHouse's HTTP port configured, so only the native protocol is "+ + "available in this session. Connect with a native client such as clickhouse-client.") + return + } + r.Body = body state := &requestState{statement: statement, started: time.Now()} p.reverse.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), stateKey{}, state))) From 4c44bb329d33ac2c01c26763e70e7a1b942f8f19 Mon Sep 17 00:00:00 2001 From: bernie-g Date: Fri, 25 Sep 2026 17:27:24 -0400 Subject: [PATCH 05/16] fix(clickhouse): harden the native handshake and probe classification ch-go allocates a declared string length before reading it, so an unauthenticated client could name a terabyte and take the process down. Handshake fields are now read through a bounded reader. The session revision was pinned to min(client, ch-go) and never clamped to the upstream's. Against a server older than ch-go that put feature-gated bytes on the wire it never reads, desynchronising the stream. A native port that accepts the connection and then says nothing dropped the deadline error, so it classified as a rejected credential and stopped the heartbeat schedule. It now wraps, and the handshake read is bounded by the probe's own budget rather than outliving it. Also refuses a connection test that was given no port, and classifies an unauthorised port as a transport failure instead of leaving it unknown. --- .../gateway-v2/test_connection_handler.go | 17 +- packages/pam/handlers/clickhouse/native.go | 68 ++++- .../handlers/clickhouse/native_unit_test.go | 250 +++++++++++++++++- .../pam/handlers/clickhouse/proxy_test.go | 44 +++ 4 files changed, 366 insertions(+), 13 deletions(-) diff --git a/packages/gateway-v2/test_connection_handler.go b/packages/gateway-v2/test_connection_handler.go index 17e567b2e..17bd7a90c 100644 --- a/packages/gateway-v2/test_connection_handler.go +++ b/packages/gateway-v2/test_connection_handler.go @@ -727,7 +727,7 @@ func handleTestConnection(w http.ResponseWriter, r *http.Request) { // The body names which ports to probe; the signed certificate still decides which are allowed. for _, port := range []int{params.HttpPort, params.NativePort} { if port > 0 && !target.allows(port) { - return fmt.Errorf("port %d is not authorised for this connection test", port) + return connectFailure(fmt.Errorf("port %d is not authorised for this connection test", port)) } } @@ -745,11 +745,20 @@ func handleTestConnection(w http.ResponseWriter, r *http.Request) { if params.NativePort > 0 { probes++ } + if probes == 0 { + return connectFailure(errors.New("no ClickHouse port was supplied for this connection test")) + } + + remaining := probes probeCtx := func() (context.Context, context.CancelFunc) { - if deadline, ok := ctx.Deadline(); ok && probes > 1 { - return context.WithTimeout(ctx, time.Until(deadline)/time.Duration(probes)) + deadline, ok := ctx.Deadline() + if !ok || remaining <= 1 { + remaining-- + return context.WithCancel(ctx) } - return context.WithCancel(ctx) + slice := time.Until(deadline) / time.Duration(remaining) + remaining-- + return context.WithTimeout(ctx, slice) } if httpPort > 0 { diff --git a/packages/pam/handlers/clickhouse/native.go b/packages/pam/handlers/clickhouse/native.go index 5d97e61ff..dfe4164a9 100644 --- a/packages/pam/handlers/clickhouse/native.go +++ b/packages/pam/handlers/clickhouse/native.go @@ -78,6 +78,52 @@ func (t *tap) discard() { t.buf = nil } +// ch-go allocates a declared string length before it reads a single byte, so an unauthenticated client could +// name a terabyte and take the process down with it. Every handshake field is a short identifier. +const maxHandshakeStringLen = 64 << 10 + +func readBoundedStr(r *proto.Reader) (string, error) { + n, err := r.UVarInt() + if err != nil { + return "", err + } + if n > maxHandshakeStringLen { + return "", fmt.Errorf("handshake field of %d bytes exceeds the %d byte cap", n, maxHandshakeStringLen) + } + buf := make([]byte, n) + if _, err := io.ReadFull(r, buf); err != nil { + return "", err + } + return string(buf), nil +} + +func decodeBoundedClientHello(r *proto.Reader) (proto.ClientHello, error) { + var h proto.ClientHello + var err error + if h.Name, err = readBoundedStr(r); err != nil { + return h, fmt.Errorf("name: %w", err) + } + if h.Major, err = r.Int(); err != nil { + return h, fmt.Errorf("major: %w", err) + } + if h.Minor, err = r.Int(); err != nil { + return h, fmt.Errorf("minor: %w", err) + } + if h.ProtocolVersion, err = r.Int(); err != nil { + return h, fmt.Errorf("protocol version: %w", err) + } + if h.Database, err = readBoundedStr(r); err != nil { + return h, fmt.Errorf("database: %w", err) + } + if h.User, err = readBoundedStr(r); err != nil { + return h, fmt.Errorf("user: %w", err) + } + if h.Password, err = readBoundedStr(r); err != nil { + return h, fmt.Errorf("password: %w", err) + } + return h, nil +} + // A refusal ends the session: the stream is mid-packet, so carrying on would let a later packet flush the // refused bytes upstream. var errSessionRefused = errors.New("the session was refused") @@ -226,8 +272,8 @@ func (s *nativeSession) handshake(t *tap, r *proto.Reader) error { return fmt.Errorf("expected Hello, got client packet %d", code) } - var hello proto.ClientHello - if err := hello.Decode(r); err != nil { + hello, err := decodeBoundedClientHello(r) + if err != nil { return fmt.Errorf("decode client hello: %w", err) } t.discard() @@ -277,6 +323,11 @@ func (s *nativeSession) handshake(t *tap, r *proto.Reader) error { if err := serverHello.DecodeAware(serverReader, s.rev); err != nil { return fmt.Errorf("decode server hello: %w", err) } + // The upstream can be older than the revision pinned from the client, and anything above what it + // speaks puts feature-gated bytes on the wire it never reads, desynchronising the stream. + if serverHello.Revision > 0 && serverHello.Revision < s.rev { + s.rev = serverHello.Revision + } serverHello.Revision = s.rev s.upstreamTap.discard() @@ -289,7 +340,7 @@ func (s *nativeSession) handshake(t *tap, r *proto.Reader) error { // At rev >= 54458 the quota key follows the handshake as a bare string. Ours is empty: not the client's to pick. if proto.FeatureAddendum.In(s.rev) { - if _, err := r.Str(); err != nil { + if _, err := readBoundedStr(r); err != nil { return fmt.Errorf("read client addendum: %w", err) } t.discard() @@ -598,7 +649,12 @@ func TestNativeConnection(ctx context.Context, config ClickHouseProxyConfig) err } defer conn.Close() - _ = conn.SetDeadline(time.Now().Add(nativeHandshakeTimeout)) + // The probe's own budget wins when it is shorter, so a slow handshake cannot outlive the test. + deadline := time.Now().Add(nativeHandshakeTimeout) + if probeDeadline, ok := ctx.Deadline(); ok && probeDeadline.Before(deadline) { + deadline = probeDeadline + } + _ = conn.SetDeadline(deadline) var b proto.Buffer proto.ClientHello{ @@ -619,8 +675,8 @@ func TestNativeConnection(ctx context.Context, config ClickHouseProxyConfig) err if err != nil { if errors.Is(err, os.ErrDeadlineExceeded) { return fmt.Errorf("the port accepted the connection but did not answer ClickHouse's native "+ - "handshake within %s, which is what the HTTP port does when it is entered as the native one", - nativeHandshakeTimeout) + "handshake within %s, which is what the HTTP port does when it is entered as the native one: %w", + nativeHandshakeTimeout, err) } return fmt.Errorf("read hello response: %w", err) } diff --git a/packages/pam/handlers/clickhouse/native_unit_test.go b/packages/pam/handlers/clickhouse/native_unit_test.go index 9b592effc..595a0fc56 100644 --- a/packages/pam/handlers/clickhouse/native_unit_test.go +++ b/packages/pam/handlers/clickhouse/native_unit_test.go @@ -2,13 +2,19 @@ package clickhouse import ( "context" + "io" "net" + "net/http" + "net/http/httptest" + "os" "regexp" + "strings" "sync" "testing" "time" "github.com/ClickHouse/ch-go/proto" + "github.com/rs/zerolog" "github.com/stretchr/testify/require" ) @@ -23,15 +29,23 @@ type fakeClickHouse struct { // Bytes seen after the handshake, which is what proves nothing was relayed once a refusal happened. bytesAfterHandshake int done chan struct{} + + // Non-zero caps what this server claims to speak, standing in for a ClickHouse older than ch-go. + serverRevision int + // Non-empty answers the handshake with an exception instead of a hello. + refuseWith string } -func startFakeClickHouse(t *testing.T) *fakeClickHouse { +func startFakeClickHouse(t *testing.T, serverRevision ...int) *fakeClickHouse { t.Helper() listener, err := net.Listen("tcp", "127.0.0.1:0") require.NoError(t, err) f := &fakeClickHouse{listener: listener, done: make(chan struct{})} + if len(serverRevision) > 0 { + f.serverRevision = serverRevision[0] + } t.Cleanup(func() { listener.Close() }) go func() { @@ -67,8 +81,18 @@ func (f *fakeClickHouse) serve(conn net.Conn) { f.mu.Unlock() rev := min(hello.ProtocolVersion, proto.Version) + if f.serverRevision > 0 { + rev = min(rev, f.serverRevision) + } var b proto.Buffer + if f.refuseWith != "" { + exception := proto.Exception{Code: 516, Name: "AUTHENTICATION_FAILED", Message: f.refuseWith} + proto.ServerCodeException.Encode(&b) + exception.EncodeAware(&b, rev) + _, _ = conn.Write(b.Buf) + return + } serverHello := proto.ServerHello{Name: "FakeClickHouse", Major: 24, Minor: 8, Revision: rev} serverHello.EncodeAware(&b, rev) if _, err := conn.Write(b.Buf); err != nil { @@ -315,9 +339,13 @@ func TestNativeHandshakePinsTheRevision(t *testing.T) { } func writeQuery(t *testing.T, conn net.Conn, q proto.Query) { + writeQueryAt(t, conn, q, proto.Version) +} + +func writeQueryAt(t *testing.T, conn net.Conn, q proto.Query, rev int) { t.Helper() - q.Info.ProtocolVersion = proto.Version + q.Info.ProtocolVersion = rev q.Info.Major, q.Info.Minor = 24, 8 q.Info.Interface = proto.InterfaceTCP q.Info.Query = proto.ClientQueryInitial @@ -325,7 +353,7 @@ func writeQuery(t *testing.T, conn net.Conn, q proto.Query) { q.Stage = proto.StageComplete var b proto.Buffer - q.EncodeAware(&b, proto.Version) + q.EncodeAware(&b, rev) _, err := conn.Write(b.Buf) require.NoError(t, err) } @@ -347,3 +375,219 @@ func TestChGoRevisionIsStillBehindTheServers(t *testing.T) { require.LessOrEqual(t, maxNativeRevision, 54469, "ch-go has caught up with ClickHouse; revision pinning needs revisiting") } + +// A server older than ch-go reads no addendum, so pinning above what it speaks desynchronises the stream. +func TestNativeHandshakeClampsToAnOlderServer(t *testing.T) { + const oldRevision = 54455 // ClickHouse 22.3 LTS, below FeatureAddendum (54458) + require.False(t, proto.FeatureAddendum.In(oldRevision)) + + upstream := startFakeClickHouse(t, oldRevision) + + conn := dialProxy(t, ClickHouseProxyConfig{ + NativeAddr: upstream.addr(), + Username: "account", + SessionID: "unit", + }) + + var b proto.Buffer + proto.ClientHello{ + Name: "a modern client", + Major: 24, + Minor: 8, + ProtocolVersion: proto.Version, + Database: "db", + User: "someone", + }.Encode(&b) + _, err := conn.Write(b.Buf) + require.NoError(t, err) + + r := proto.NewReader(newTap(conn)) + code, err := r.UVarInt() + require.NoError(t, err) + require.Equal(t, proto.ServerCodeHello, proto.ServerCode(code)) + + var serverHello proto.ServerHello + require.NoError(t, serverHello.DecodeAware(r, oldRevision)) + require.Equal(t, oldRevision, serverHello.Revision, + "the client must be told the upstream's revision, not one it cannot parse") + + // The fake reads no quota key at this revision, so a query must land as the next packet it sees. + writeQueryAt(t, conn, proto.Query{Body: "SELECT 1"}, oldRevision) + + require.Eventually(t, func() bool { + _, _, queries, _ := upstream.snapshot() + return len(queries) == 1 + }, 5*time.Second, 20*time.Millisecond, "the upstream stream desynchronised") + + _, quotaKey, queries, _ := upstream.snapshot() + require.Equal(t, "SELECT 1", queries[0].Body) + require.Empty(t, quotaKey, "no addendum may be written to a server that does not read one") +} + +func TestNativeHandshakeRefusesAnOversizedField(t *testing.T) { + upstream := startFakeClickHouse(t) + + conn := dialProxy(t, ClickHouseProxyConfig{ + NativeAddr: upstream.addr(), + Username: "account", + SessionID: "unit", + }) + + var b proto.Buffer + proto.ClientCodeHello.Encode(&b) + b.PutUVarInt(uint64(maxHandshakeStringLen) + 1) + _, err := conn.Write(b.Buf) + require.NoError(t, err) + + require.NoError(t, conn.SetReadDeadline(time.Now().Add(5*time.Second))) + answer, _ := io.ReadAll(conn) + require.Empty(t, answer, "an oversized handshake field must not be answered") + + hello, _, _, bytesAfterHandshake := upstream.snapshot() + require.Empty(t, hello.Name, "nothing may reach the upstream") + require.Zero(t, bytesAfterHandshake) +} + +// The server direction must stop writing once a statement has been refused, or the client sees bytes +// trailing the exception the proxy just sent it. +func newRefusedSession(t *testing.T, client net.Conn, upstream net.Conn) *nativeSession { + t.Helper() + + proxy := NewClickHouseProxy(ClickHouseProxyConfig{SessionID: "unit", SessionLogger: &recordingLogger{}}) + s := &nativeSession{ + proxy: newNativeProxy(proxy), + log: zerolog.Nop(), + client: client, + upstream: upstream, + rev: maxNativeRevision, + outcomes: newOutcomeRecorder(proxy), + } + s.upstreamTap = newTap(upstream) + s.upstreamReader = proto.NewReader(s.upstreamTap) + s.refused.Store(true) + return s +} + +func TestRefusedSessionRelaysNoServerBytes(t *testing.T) { + // A packet the loop can parse leaves by the per-iteration guard; one it cannot leaves by the relay. + for name, packet := range map[string][]byte{ + "a parseable packet": func() []byte { + var b proto.Buffer + proto.ServerCodePong.Encode(&b) + return b.Buf + }(), + "a packet the loop cannot read": func() []byte { + var b proto.Buffer + b.PutUVarInt(250) + b.PutString(strings.Repeat("x", 512)) + return b.Buf + }(), + } { + t.Run(name, func(t *testing.T) { + client, clientPeer := net.Pipe() + upstream, upstreamPeer := net.Pipe() + t.Cleanup(func() { + client.Close() + clientPeer.Close() + upstream.Close() + upstreamPeer.Close() + }) + + s := newRefusedSession(t, client, upstream) + + go func() { + _, _ = upstreamPeer.Write(packet) + upstreamPeer.Close() + }() + + done := make(chan struct{}) + go func() { + defer close(done) + s.serverLoop() + }() + + require.NoError(t, clientPeer.SetReadDeadline(time.Now().Add(2*time.Second))) + buf := make([]byte, 64) + n, err := clientPeer.Read(buf) + require.Error(t, err, "a refused session must write nothing further to the client") + require.Zero(t, n) + + client.Close() + select { + case <-done: + case <-time.After(5 * time.Second): + t.Fatal("serverLoop did not return") + } + }) + } +} + +func TestRefusalAwareWriterRefusesAfterARefusal(t *testing.T) { + client, clientPeer := net.Pipe() + upstream, upstreamPeer := net.Pipe() + t.Cleanup(func() { + client.Close() + clientPeer.Close() + upstream.Close() + upstreamPeer.Close() + }) + + s := newRefusedSession(t, client, upstream) + + n, err := newRefusalAwareWriter(s).Write([]byte("server bytes")) + require.ErrorIs(t, err, errSessionRefused) + require.Zero(t, n) +} + +func TestNativeConnectionTestClassifiesFailures(t *testing.T) { + t.Run("an http port is named as a handshake timeout, not a rejected credential", func(t *testing.T) { + // An HTTP server accepts the connection and then waits for a request, which is exactly what a + // misconfigured native port looks like. + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {})) + t.Cleanup(server.Close) + + // The probe's own deadline bounds the handshake, so this does not wait out nativeHandshakeTimeout. + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + + err := TestNativeConnection(ctx, ClickHouseProxyConfig{ + NativeAddr: strings.TrimPrefix(server.URL, "http://"), + Username: "account", + Database: "analytics", + }) + require.Error(t, err) + require.Contains(t, err.Error(), "did not answer ClickHouse's native handshake") + // The heartbeat stops scheduling on a rejected credential, so a silent port has to stay a + // transport failure rather than being read as one. + require.ErrorIs(t, err, os.ErrDeadlineExceeded) + }) + + t.Run("a refused credential surfaces the server's own message", func(t *testing.T) { + upstream := startFakeClickHouse(t) + upstream.refuseWith = "Authentication failed: password is incorrect" + + err := TestNativeConnection(context.Background(), ClickHouseProxyConfig{ + NativeAddr: upstream.addr(), + Username: "account", + Database: "analytics", + }) + require.Error(t, err) + require.Contains(t, err.Error(), "password is incorrect") + require.NotErrorIs(t, err, os.ErrDeadlineExceeded) + }) + + t.Run("an unreachable port fails as a dial error", func(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + addr := listener.Addr().String() + require.NoError(t, listener.Close()) + + err = TestNativeConnection(context.Background(), ClickHouseProxyConfig{ + NativeAddr: addr, + Username: "account", + Database: "analytics", + }) + require.Error(t, err) + require.NotContains(t, err.Error(), "did not answer ClickHouse's native handshake") + }) +} diff --git a/packages/pam/handlers/clickhouse/proxy_test.go b/packages/pam/handlers/clickhouse/proxy_test.go index 45f7ab965..9560cd3a1 100644 --- a/packages/pam/handlers/clickhouse/proxy_test.go +++ b/packages/pam/handlers/clickhouse/proxy_test.go @@ -3,6 +3,7 @@ package clickhouse import ( "bytes" "compress/gzip" + "compress/zlib" "io" "net/http" "net/http/httptest" @@ -446,3 +447,46 @@ func TestReportsAnUnreachableTargetAsAClickHouseError(t *testing.T) { require.Len(t, logger.entries, 1) require.Contains(t, logger.entries[0].Output, "ERROR:") } + +func deflated(t *testing.T, payload string) []byte { + t.Helper() + var buffer bytes.Buffer + writer := zlib.NewWriter(&buffer) + _, err := writer.Write([]byte(payload)) + require.NoError(t, err) + require.NoError(t, writer.Close()) + return buffer.Bytes() +} + +// deflate is on the accepted-encoding list, so a statement hidden in one has to be inspected too. +func TestBlocksAStatementInsideADeflatedBody(t *testing.T) { + reached := false + handler, _, closeUpstream := newTestProxy(t, func(w http.ResponseWriter, r *http.Request) { + reached = true + }, `(?i)\btruncate\b`) + defer closeUpstream() + + req := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(deflated(t, "TRUNCATE TABLE events"))) + req.Header.Set("Content-Encoding", "deflate") + recorder := httptest.NewRecorder() + handler.ServeHTTP(recorder, req) + + require.False(t, reached, "a blocked statement must not reach the upstream") + require.Equal(t, http.StatusForbidden, recorder.Code) +} + +func TestRefusesADeflatedBodyItCannotDecode(t *testing.T) { + reached := false + handler, _, closeUpstream := newTestProxy(t, func(w http.ResponseWriter, r *http.Request) { + reached = true + }, `(?i)\btruncate\b`) + defer closeUpstream() + + req := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader([]byte{0x00, 0x01, 0x02, 0x03, 0x04})) + req.Header.Set("Content-Encoding", "deflate") + recorder := httptest.NewRecorder() + handler.ServeHTTP(recorder, req) + + require.False(t, reached, "a body the gateway could not read must not be forwarded uninspected") + require.Equal(t, http.StatusBadRequest, recorder.Code, recorder.Body.String()) +} From 21628773fe1997efbb14b19d372b9b43da287cc3 Mon Sep 17 00:00:00 2001 From: bernie-g Date: Fri, 25 Sep 2026 17:27:37 -0400 Subject: [PATCH 06/16] test(clickhouse): move live-server tests to the e2e suite The handler package carried 23 tests behind PAM_CLICKHOUSE_NATIVE_IT that needed an ambient ClickHouse with hand-seeded tables. No fixtures were committed and no CI job ran them, so they only ever ran on one machine. Real-client interop moves to e2e/pam, where testcontainers starts the server and the container's own clickhouse-client drives the session, matching the Postgres and Redis suites. The handler package now runs with no containers and no environment variables, and CI runs it under -race along with gateway-v2. e2e is a separate module and its go.mod had gone stale against the new ch-go dependency, which was failing three CI jobs. --- .github/workflows/run-cli-tests.yml | 2 + e2e/go.mod | 10 +- e2e/go.sum | 20 +- e2e/pam/clickhouse_test.go | 298 ++++++++++ .../pam/handlers/clickhouse/clients_test.go | 44 -- .../handlers/clickhouse/edge_cases_test.go | 531 ------------------ .../clickhouse/native_integration_test.go | 311 ---------- .../handlers/clickhouse/testhelpers_test.go | 101 ---- 8 files changed, 321 insertions(+), 996 deletions(-) create mode 100644 e2e/pam/clickhouse_test.go delete mode 100644 packages/pam/handlers/clickhouse/clients_test.go delete mode 100644 packages/pam/handlers/clickhouse/edge_cases_test.go delete mode 100644 packages/pam/handlers/clickhouse/native_integration_test.go delete mode 100644 packages/pam/handlers/clickhouse/testhelpers_test.go diff --git a/.github/workflows/run-cli-tests.yml b/.github/workflows/run-cli-tests.yml index 3b48b6a22..ed583def9 100644 --- a/.github/workflows/run-cli-tests.yml +++ b/.github/workflows/run-cli-tests.yml @@ -40,6 +40,8 @@ jobs: run: go get . - name: Race-check the Agent Vault packages run: go test -race -count=1 ./packages/agentvault/ ./packages/api/ + - name: Race-check the PAM gateway packages + run: go test -race -count=1 ./packages/gateway-v2/... ./packages/pam/handlers/clickhouse/ - name: Test with the Go CLI env: CLI_TESTS_UA_CLIENT_ID: ${{ secrets.CLI_TESTS_UA_CLIENT_ID }} diff --git a/e2e/go.mod b/e2e/go.mod index 075246b1d..706837996 100644 --- a/e2e/go.mod +++ b/e2e/go.mod @@ -36,6 +36,7 @@ require ( github.com/Azure/go-ansiterm v0.0.0-20250102033503-faa5f7b0171c // indirect github.com/Azure/go-ntlmssp v0.1.1 // indirect github.com/ChrisTrenkamp/goxpath v0.0.0-20210404020558-97928f7e12b6 // indirect + github.com/ClickHouse/ch-go v0.74.0 // indirect github.com/DefangLabs/secret-detector v0.0.0-20250403165618-22662109213e // indirect github.com/Masterminds/goutils v1.1.1 // indirect github.com/Masterminds/semver/v3 v3.4.0 // indirect @@ -125,6 +126,8 @@ require ( github.com/getkin/kin-openapi v0.133.0 // indirect github.com/gitleaks/go-gitdiff v0.9.1 // indirect github.com/go-asn1-ber/asn1-ber v1.5.8 // indirect + github.com/go-faster/city v1.0.1 // indirect + github.com/go-faster/errors v0.7.1 // indirect github.com/go-ldap/ldap/v3 v3.4.14 // indirect github.com/go-logr/logr v1.4.3 // indirect github.com/go-logr/stdr v1.2.2 // indirect @@ -194,7 +197,7 @@ require ( github.com/jonboulle/clockwork v0.5.0 // indirect github.com/josharian/intern v1.0.0 // indirect github.com/json-iterator/go v1.1.12 // indirect - github.com/klauspost/compress v1.18.7 // indirect + github.com/klauspost/compress v1.19.1 // indirect github.com/klauspost/pgzip v1.2.6 // indirect github.com/lucasb-eyer/go-colorful v1.4.0 // indirect github.com/lufia/plan9stats v0.0.0-20211012122336-39d0f177ccd0 // indirect @@ -261,7 +264,7 @@ require ( github.com/pelletier/go-toml v1.9.5 // indirect github.com/pelletier/go-toml/v2 v2.4.3 // indirect github.com/perimeterx/marshmallow v1.1.5 // indirect - github.com/pierrec/lz4/v4 v4.1.22 // indirect + github.com/pierrec/lz4/v4 v4.1.27 // indirect github.com/pingcap/errors v0.11.5-0.20250318082626-8f80e5cb09ec // indirect github.com/pingcap/log v1.1.1-0.20241212030209-7e3ff8601a2a // indirect github.com/pingcap/tidb/pkg/parser v0.0.0-20250421232622-526b2c79173d // indirect @@ -279,6 +282,7 @@ require ( github.com/rs/cors v1.11.0 // indirect github.com/santhosh-tekuri/jsonschema/v6 v6.0.1 // indirect github.com/secure-systems-lab/go-securesystemslib v0.6.0 // indirect + github.com/segmentio/asm v1.2.1 // indirect github.com/serialx/hashring v0.0.0-20200727003509-22c0c7ab6b1b // indirect github.com/shibumi/go-pathspec v1.3.0 // indirect github.com/shirou/gopsutil/v4 v4.25.6 // indirect @@ -341,7 +345,7 @@ require ( go.opentelemetry.io/proto/otlp v1.9.0 // indirect go.uber.org/atomic v1.11.0 // indirect go.uber.org/multierr v1.11.0 // indirect - go.uber.org/zap v1.27.0 // indirect + go.uber.org/zap v1.28.0 // indirect go.yaml.in/yaml/v3 v3.0.5 // indirect go4.org v0.0.0-20230225012048-214862532bf5 // indirect golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0 // indirect diff --git a/e2e/go.sum b/e2e/go.sum index 451da42ea..517b4d2b9 100644 --- a/e2e/go.sum +++ b/e2e/go.sum @@ -71,6 +71,8 @@ github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03 github.com/BurntSushi/xgb v0.0.0-20160522181843-27f122750802/go.mod h1:IVnqGOEym/WlBOVXweHU+Q+/VP0lqqI8lqeDx9IjBqo= github.com/ChrisTrenkamp/goxpath v0.0.0-20210404020558-97928f7e12b6 h1:w0E0fgc1YafGEh5cROhlROMWXiNoZqApk2PDN0M1+Ns= github.com/ChrisTrenkamp/goxpath v0.0.0-20210404020558-97928f7e12b6/go.mod h1:nuWgzSkT5PnyOd+272uUmV0dnAnAn42Mk7PiQC5VzN4= +github.com/ClickHouse/ch-go v0.74.0 h1:uYs2m4wIt0ZHSM1E72rg0maCfzhR2V3xWb/vZEgpeWE= +github.com/ClickHouse/ch-go v0.74.0/go.mod h1:sZ/r+8ttZMjyrP9PuFbgoVbth1ywIu2LIQNA2vgko6M= github.com/DefangLabs/secret-detector v0.0.0-20250403165618-22662109213e h1:rd4bOvKmDIx0WeTv9Qz+hghsgyjikFiPrseXHlKepO0= github.com/DefangLabs/secret-detector v0.0.0-20250403165618-22662109213e/go.mod h1:blbwPQh4DTlCZEfk1BLU4oMIhLda2U+A840Uag9DsZw= github.com/Infisical/go-keyring v1.0.2 h1:dWOkI/pB/7RocfSJgGXbXxLDcVYsdslgjEPmVhb+nl8= @@ -377,6 +379,10 @@ github.com/go-asn1-ber/asn1-ber v1.5.8 h1:H9AZkK22UOmfX8J84ubyaZxKJZ3FMHVwn8swoM github.com/go-asn1-ber/asn1-ber v1.5.8/go.mod h1:hEBeB/ic+5LoWskz+yKT7vGhhPYkProFKoKdwZRWMe0= github.com/go-faker/faker/v4 v4.7.0 h1:VboC02cXHl/NuQh5lM2W8b87yp4iFXIu59x4w0RZi4E= github.com/go-faker/faker/v4 v4.7.0/go.mod h1:u1dIRP5neLB6kTzgyVjdBOV5R1uP7BdxkcWk7tiKQXk= +github.com/go-faster/city v1.0.1 h1:4WAxSZ3V2Ws4QRDrscLEDcibJY8uf41H6AhXDrNDcGw= +github.com/go-faster/city v1.0.1/go.mod h1:jKcUJId49qdW3L1qKHH/3wPeUstCVpVSXTM6vO3VcTw= +github.com/go-faster/errors v0.7.1 h1:MkJTnDoEdi9pDabt1dpWf7AA8/BaSYZqibYyhZ20AYg= +github.com/go-faster/errors v0.7.1/go.mod h1:5ySTjWFiphBs07IKuiL69nxdfd5+fzh1u7FPGZP2quo= github.com/go-gl/glfw v0.0.0-20190409004039-e6da0acd62b1/go.mod h1:vR7hzQXu2zJy9AVAgeJqvqgH9Q5CA+iKCZ2gyEVpxRU= github.com/go-gl/glfw/v3.3/glfw v0.0.0-20191125211704-12ad95a8df72/go.mod h1:tQ2UAYgL5IevRw8kRxooKSPJfGvJ9fJQFa0TUsXzTg8= github.com/go-gl/glfw/v3.3/glfw v0.0.0-20200222043503-6f7a984d4dc4/go.mod h1:tQ2UAYgL5IevRw8kRxooKSPJfGvJ9fJQFa0TUsXzTg8= @@ -670,8 +676,8 @@ github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+o github.com/klauspost/compress v1.4.1/go.mod h1:RyIbtBH6LamlWaDj8nUwkbUhJ87Yi3uG0guNDohfE1A= github.com/klauspost/compress v1.12.3/go.mod h1:8dP1Hq4DHOhN9w426knH3Rhby4rFm6D8eO+e+Dq5Gzg= github.com/klauspost/compress v1.13.6/go.mod h1:/3/Vjq9QcHkK5uEr5lBEmyoZ1iFhe47etQ6QUkpK6sk= -github.com/klauspost/compress v1.18.7 h1:aUyZsS4kH3QTKurYhAOwAHxllVPnOthb3vPfnF1Ehjw= -github.com/klauspost/compress v1.18.7/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= +github.com/klauspost/compress v1.19.1 h1:VsB4HPswih7mmZ8WleSFQ75c/Ui1M4trX5oAsJnhSlk= +github.com/klauspost/compress v1.19.1/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= github.com/klauspost/cpuid v1.2.0/go.mod h1:Pj4uuM528wm8OyEC2QMXAi2YiTZ96dNQPGgoMS4s3ek= github.com/klauspost/pgzip v1.2.6 h1:8RXeL5crjEUFnR2/Sn6GJNWtSQ3Dk8pq4CL3jvdDyjU= github.com/klauspost/pgzip v1.2.6/go.mod h1:Ch1tH69qFZu15pkjo5kYi6mth2Zzwzt50oCQKQE9RUs= @@ -884,8 +890,8 @@ github.com/pelletier/go-toml/v2 v2.4.3 h1:GTRvJQutkOSftxIFD5xw9aepkYNuPWmVJpffdD github.com/pelletier/go-toml/v2 v2.4.3/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY= github.com/perimeterx/marshmallow v1.1.5 h1:a2LALqQ1BlHM8PZblsDdidgv1mWi1DgC2UmX50IvK2s= github.com/perimeterx/marshmallow v1.1.5/go.mod h1:dsXbUu8CRzfYP5a87xpp0xq9S3u0Vchtcl8we9tYaXw= -github.com/pierrec/lz4/v4 v4.1.22 h1:cKFw6uJDK+/gfw5BcDL0JL5aBsAFdsIT18eRtLj7VIU= -github.com/pierrec/lz4/v4 v4.1.22/go.mod h1:gZWDp/Ze/IJXGXf23ltt2EXimqmTUXEy0GFuRQyBid4= +github.com/pierrec/lz4/v4 v4.1.27 h1:+PhzhWDrjRj89TH2sw43nE3+4+W8lSxIuQadEHZyjUk= +github.com/pierrec/lz4/v4 v4.1.27/go.mod h1:EoQMVJgeeEOMsCqCzqFm2O0cJvljX2nGZjcRIPL34O4= github.com/pingcap/errors v0.11.0/go.mod h1:Oi8TUi2kEtXXLMJk9l1cGmz20kV3TaQ0usTwv5KuLY8= github.com/pingcap/errors v0.11.5-0.20250318082626-8f80e5cb09ec h1:3EiGmeJWoNixU+EwllIn26x6s4njiWRXewdx2zlYa84= github.com/pingcap/errors v0.11.5-0.20250318082626-8f80e5cb09ec/go.mod h1:X2r9ueLEUZgtx2cIogM0v4Zj5uvvzhuuiu7Pn8HzMPg= @@ -957,6 +963,8 @@ github.com/santhosh-tekuri/jsonschema/v6 v6.0.1/go.mod h1:JXeL+ps8p7/KNMjDQk3TCw github.com/sean-/seed v0.0.0-20170313163322-e2103e2c3529/go.mod h1:DxrIzT+xaE7yg65j358z/aeFdxmN0P9QXhEzd20vsDc= github.com/secure-systems-lab/go-securesystemslib v0.6.0 h1:T65atpAVCJQK14UA57LMdZGpHi4QYSH/9FZyNGqMYIA= github.com/secure-systems-lab/go-securesystemslib v0.6.0/go.mod h1:8Mtpo9JKks/qhPG4HGZ2LGMvrPbzuxwfz/f/zLfEWkk= +github.com/segmentio/asm v1.2.1 h1:DTNbBqs57ioxAD4PrArqftgypG4/qNpXoJx8TVXxPR0= +github.com/segmentio/asm v1.2.1/go.mod h1:BqMnlJP91P8d+4ibuonYZw9mfnzI9HfxselHZr5aAcs= github.com/sergi/go-diff v1.1.0 h1:we8PVUC3FE2uYfodKH/nBHMSetSfHDR6scGdBi+erh0= github.com/sergi/go-diff v1.1.0/go.mod h1:STckp+ISIX8hZLjrqAeVduY0gWCT9IjLuqbuNXdaHfM= github.com/serialx/hashring v0.0.0-20200727003509-22c0c7ab6b1b h1:h+3JX2VoWTFuyQEo87pStk/a99dzIO1mM9KxIyLPGTU= @@ -1186,8 +1194,8 @@ go.uber.org/multierr v1.11.0 h1:blXXJkSxSSfBVBlC76pxqeO+LN3aDfLQo+309xJstO0= go.uber.org/multierr v1.11.0/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN80Y= go.uber.org/zap v1.17.0/go.mod h1:MXVU+bhUf/A7Xi2HNOnopQOrmycQ5Ih87HtOu4q5SSo= go.uber.org/zap v1.19.0/go.mod h1:xg/QME4nWcxGxrpdeYfq7UvYrLh66cuVKdrbD1XF/NI= -go.uber.org/zap v1.27.0 h1:aJMhYGrd5QSmlpLMr2MftRKl7t8J8PTZPA732ud/XR8= -go.uber.org/zap v1.27.0/go.mod h1:GB2qFLM7cTU87MWRP2mPIjqfIDnGu+VIO4V/SdhGo2E= +go.uber.org/zap v1.28.0 h1:IZzaP1Fv73/T/pBMLk4VutPl36uNC+OSUh3JLG3FIjo= +go.uber.org/zap v1.28.0/go.mod h1:rDLpOi171uODNm/mxFcuYWxDsqWSAVkFdX4XojSKg/Q= go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw= go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg= go4.org v0.0.0-20230225012048-214862532bf5 h1:nifaUDeh+rPaBCMPMQHZmvJf+QdpLFnuQPwx+LxVmtc= diff --git a/e2e/pam/clickhouse_test.go b/e2e/pam/clickhouse_test.go new file mode 100644 index 000000000..d3e33c9be --- /dev/null +++ b/e2e/pam/clickhouse_test.go @@ -0,0 +1,298 @@ +package pam + +import ( + "context" + "fmt" + "io" + "log/slog" + "net/http" + "strings" + "testing" + "time" + + "github.com/docker/docker/api/types/container" + "github.com/infisical/cli/e2e-tests/packages/client" + helpers "github.com/infisical/cli/e2e-tests/util" + openapitypes "github.com/oapi-codegen/runtime/types" + "github.com/stretchr/testify/require" + "github.com/testcontainers/testcontainers-go" + "github.com/testcontainers/testcontainers-go/wait" +) + +const ( + clickhouseImage = "clickhouse/clickhouse-server:24.8" + clickhouseDatabase = "analytics" + clickhouseUser = "default" + clickhousePassword = "clickhouse" +) + +func startClickHouseContainer(t *testing.T, ctx context.Context) (testcontainers.Container, string, int, int) { + t.Helper() + + ctr, err := testcontainers.GenericContainer(ctx, testcontainers.GenericContainerRequest{ + ContainerRequest: testcontainers.ContainerRequest{ + Image: clickhouseImage, + ExposedPorts: []string{"8123/tcp", "9000/tcp"}, + Env: map[string]string{ + "CLICKHOUSE_DB": clickhouseDatabase, + "CLICKHOUSE_USER": clickhouseUser, + "CLICKHOUSE_PASSWORD": clickhousePassword, + }, + HostConfigModifier: func(hc *container.HostConfig) { + hc.ExtraHosts = append(hc.ExtraHosts, "host.docker.internal:host-gateway") + }, + WaitingFor: wait.ForAll( + wait.ForListeningPort("8123/tcp"), + wait.ForListeningPort("9000/tcp"), + ).WithStartupTimeout(180 * time.Second), + }, + Started: true, + }) + require.NoError(t, err) + t.Cleanup(func() { + if err := ctr.Terminate(ctx); err != nil { + t.Logf("Failed to terminate ClickHouse container: %v", err) + } + }) + + host, err := ctr.Host(ctx) + require.NoError(t, err) + httpPort, err := ctr.MappedPort(ctx, "8123") + require.NoError(t, err) + nativePort, err := ctr.MappedPort(ctx, "9000") + require.NoError(t, err) + + seedClickHouse(t, ctx, ctr) + return ctr, host, httpPort.Int(), nativePort.Int() +} + +// clickhouseClientIn runs the container's own clickhouse-client, which is what makes this an interop +// check rather than a test of our own encoder. +func clickhouseClientIn(t *testing.T, ctx context.Context, ctr testcontainers.Container, + host string, port int, sql string, extra ...string) (int, string) { + t.Helper() + + args := []string{ + "clickhouse-client", + "--host", host, + "--port", fmt.Sprintf("%d", port), + "--database", clickhouseDatabase, + // The proxy injects the account's credentials, so whatever the client sends is discarded. + "--user", "not-the-account", "--password", "not-the-password", + "--query", sql, + } + args = append(args, extra...) + + exitCode, reader, err := ctr.Exec(ctx, args) + require.NoError(t, err) + out, err := io.ReadAll(reader) + require.NoError(t, err) + return exitCode, string(out) +} + +func seedClickHouse(t *testing.T, ctx context.Context, ctr testcontainers.Container) { + t.Helper() + + statements := []string{ + "CREATE TABLE IF NOT EXISTS events (id UInt64, name String) ENGINE = MergeTree ORDER BY id", + "INSERT INTO events VALUES (1, 'alpha'), (2, 'beta'), (3, 'gamma')", + } + for _, sql := range statements { + exitCode, _, err := ctr.Exec(ctx, []string{ + "clickhouse-client", + "--user", clickhouseUser, "--password", clickhousePassword, + "--database", clickhouseDatabase, + "--query", sql, + }) + require.NoError(t, err) + require.Zero(t, exitCode, "seeding failed for: %s", sql) + } +} + +func createClickHousePamAccount(t *testing.T, ctx context.Context, infra *PAMTestInfra, + folderId, templateId openapitypes.UUID, name, host string, httpPort, nativePort *int) { + t.Helper() + + connectionDetails := map[string]interface{}{ + "host": host, + "database": clickhouseDatabase, + "sslEnabled": false, + "sslRejectUnauthorized": false, + } + if httpPort != nil { + connectionDetails["port"] = *httpPort + } + if nativePort != nil { + connectionDetails["nativePort"] = *nativePort + } + + CreatePamAccount(t, ctx, infra, "clickhouse", name, folderId, templateId, connectionDetails, + map[string]interface{}{"username": clickhouseUser, "password": clickhousePassword}) +} + +func startClickHouseProxy(t *testing.T, ctx context.Context, infra *PAMTestInfra, + folderName, accountName string) (int, *helpers.Command) { + t.Helper() + + freePort := helpers.GetFreePort() + pamCmd := helpers.Command{ + Test: t, + RunMethod: helpers.RunMethodSubprocess, + DisableTempHomeDir: true, + Args: []string{ + "pam", "access", fmt.Sprintf("%s/%s", folderName, accountName), + "--duration", "5m", + "--port", fmt.Sprintf("%d", freePort), + }, + Env: map[string]string{ + "HOME": infra.SharedHomeDir, + "INFISICAL_API_URL": infra.Infisical.ApiUrl(t), + }, + } + pamCmd.Start(ctx) + t.Cleanup(pamCmd.Stop) + + result := helpers.WaitFor(t, helpers.WaitForOptions{ + EnsureCmdRunning: &pamCmd, + Condition: func() helpers.ConditionResult { + if strings.Contains(pamCmd.Stdout(), "ClickHouse Proxy Session Started") { + return helpers.ConditionSuccess + } + return helpers.ConditionWait + }, + }) + if result != helpers.WaitSuccess { + infra.DumpOutput(&pamCmd) + } + require.Equal(t, helpers.WaitSuccess, result, "ClickHouse proxy should start successfully") + + return freePort, &pamCmd +} + +func queryOverHTTP(t *testing.T, ctx context.Context, proxyPort int, sql string) (int, string) { + t.Helper() + + url := fmt.Sprintf("http://127.0.0.1:%d/?database=%s", proxyPort, clickhouseDatabase) + req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, strings.NewReader(sql)) + require.NoError(t, err) + + resp, err := (&http.Client{Timeout: 60 * time.Second}).Do(req) + require.NoError(t, err) + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + return resp.StatusCode, string(body) +} + +// waitForProxyHTTP absorbs the gap between the banner and the listener accepting. +func waitForProxyHTTP(t *testing.T, ctx context.Context, pamCmd *helpers.Command, proxyPort int) { + t.Helper() + + result := helpers.WaitFor(t, helpers.WaitForOptions{ + EnsureCmdRunning: pamCmd, + Interval: 2 * time.Second, + Timeout: 60 * time.Second, + Condition: func() helpers.ConditionResult { + status, _ := queryOverHTTP(t, ctx, proxyPort, "SELECT 1") + if status == http.StatusOK { + return helpers.ConditionSuccess + } + return helpers.ConditionWait + }, + }) + require.Equal(t, helpers.WaitSuccess, result, "the proxy should answer HTTP") +} + +func TestPAM_ClickHouse(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + infra := SetupPAMInfra(t, ctx) + LoginUser(t, ctx, infra) + + folderName := "clickhouse-folder" + folderId := CreatePamFolder(t, ctx, infra, folderName) + templateId := CreatePamTemplate(t, ctx, infra, "clickhouse-template", + client.CreatePamAccountTemplateJSONBodyType("clickhouse")) + + ctr, chHost, chHTTPPort, chNativePort := startClickHouseContainer(t, ctx) + // The container reaches the proxy running on this host, so a loopback address will not do. + hostIP := getOutboundIP(t) + + t.Run("both interfaces on one session port", func(t *testing.T) { + accountName := "clickhouse-dual-account" + createClickHousePamAccount(t, ctx, infra, folderId, templateId, accountName, chHost, + &chHTTPPort, &chNativePort) + + proxyPort, pamCmd := startClickHouseProxy(t, ctx, infra, folderName, accountName) + waitForProxyHTTP(t, ctx, pamCmd, proxyPort) + + status, body := queryOverHTTP(t, ctx, proxyPort, "SELECT name FROM events ORDER BY id") + require.Equal(t, http.StatusOK, status, body) + require.Equal(t, "alpha\nbeta\ngamma", strings.TrimSpace(body)) + slog.Info("HTTP interface answered through the proxy") + + exitCode, out := clickhouseClientIn(t, ctx, ctr, hostIP, proxyPort, "SELECT count() FROM events") + require.Zero(t, exitCode, out) + require.Equal(t, "3", strings.TrimSpace(out)) + slog.Info("native interface answered the same session port") + }) + + t.Run("compression is negotiated end to end", func(t *testing.T) { + accountName := "clickhouse-compression-account" + createClickHousePamAccount(t, ctx, infra, folderId, templateId, accountName, chHost, + &chHTTPPort, &chNativePort) + + proxyPort, pamCmd := startClickHouseProxy(t, ctx, infra, folderName, accountName) + waitForProxyHTTP(t, ctx, pamCmd, proxyPort) + + for _, compression := range []string{"1", "0"} { + exitCode, out := clickhouseClientIn(t, ctx, ctr, hostIP, proxyPort, + "SELECT count() FROM events", "--compression", compression) + require.Zero(t, exitCode, out) + require.Equal(t, "3", strings.TrimSpace(out), "compression=%s", compression) + } + }) + + t.Run("a native-only account turns HTTP clients away", func(t *testing.T) { + accountName := "clickhouse-native-only-account" + createClickHousePamAccount(t, ctx, infra, folderId, templateId, accountName, chHost, + nil, &chNativePort) + + proxyPort, pamCmd := startClickHouseProxy(t, ctx, infra, folderName, accountName) + + result := helpers.WaitFor(t, helpers.WaitForOptions{ + EnsureCmdRunning: pamCmd, + Interval: 2 * time.Second, + Timeout: 60 * time.Second, + Condition: func() helpers.ConditionResult { + exitCode, out := clickhouseClientIn(t, ctx, ctr, hostIP, proxyPort, "SELECT count() FROM events") + if exitCode == 0 && strings.TrimSpace(out) == "3" { + return helpers.ConditionSuccess + } + return helpers.ConditionWait + }, + }) + require.Equal(t, helpers.WaitSuccess, result, "a native client should still work") + + status, body := queryOverHTTP(t, ctx, proxyPort, "SELECT 1") + require.NotEqual(t, http.StatusOK, status, body) + require.Contains(t, body, "HTTP port", + "an HTTP client must be told why, not left with a transport error") + }) + + t.Run("an HTTP-only account turns native clients away", func(t *testing.T) { + accountName := "clickhouse-http-only-account" + createClickHousePamAccount(t, ctx, infra, folderId, templateId, accountName, chHost, + &chHTTPPort, nil) + + proxyPort, pamCmd := startClickHouseProxy(t, ctx, infra, folderName, accountName) + waitForProxyHTTP(t, ctx, pamCmd, proxyPort) + + exitCode, out := clickhouseClientIn(t, ctx, ctr, hostIP, proxyPort, "SELECT count() FROM events") + require.NotZero(t, exitCode, out) + require.Contains(t, out, "native port", + "a native client must be told why, not left with a transport error") + }) +} diff --git a/packages/pam/handlers/clickhouse/clients_test.go b/packages/pam/handlers/clickhouse/clients_test.go deleted file mode 100644 index c01b520b1..000000000 --- a/packages/pam/handlers/clickhouse/clients_test.go +++ /dev/null @@ -1,44 +0,0 @@ -package clickhouse - -import ( - "context" - "os" - "os/exec" - "strings" - "testing" - "time" -) - -func TestPythonClickHouseDriver(t *testing.T) { - if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { - t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") - } - - port := startProxy(t, baseConfig(&recordingLogger{})) - - script := ` -import clickhouse_driver, sys -c = clickhouse_driver.Client(host=sys.argv[1], port=int(sys.argv[2]), user='wrong', password='wrong', database='ignored') -print("currentUser:", c.execute("SELECT currentUser()")[0][0]) -print("count:", c.execute("SELECT count() FROM users")[0][0]) -print("exotic:", c.execute("SELECT map('a', 1::UInt64), tuple('p', 2)")[0]) -print("multi:", c.execute("SELECT 1")[0][0], c.execute("SELECT 2")[0][0]) -` - ctx, cancel := context.WithTimeout(context.Background(), 180*time.Second) - defer cancel() - - cmd := exec.CommandContext(ctx, "docker", "run", "--rm", "-i", "python:3.12-slim", "bash", "-lc", - "pip install --quiet clickhouse-driver >/dev/null 2>&1 && python -c \""+strings.ReplaceAll(script, `"`, `\"`)+"\" "+ - envOr("PAM_CLICKHOUSE_CLIENT_HOST", "host.docker.internal")+" "+port) - - out, err := cmd.CombinedOutput() - t.Logf("%s", out) - if err != nil { - t.Fatalf("clickhouse-driver failed: %v", err) - } - for _, want := range []string{"currentUser: default", "count: 250", "multi: 1 2"} { - if !strings.Contains(string(out), want) { - t.Fatalf("expected %q in output", want) - } - } -} diff --git a/packages/pam/handlers/clickhouse/edge_cases_test.go b/packages/pam/handlers/clickhouse/edge_cases_test.go deleted file mode 100644 index fc4025cbd..000000000 --- a/packages/pam/handlers/clickhouse/edge_cases_test.go +++ /dev/null @@ -1,531 +0,0 @@ -package clickhouse - -import ( - "context" - "crypto/tls" - "crypto/x509" - "fmt" - "io" - "net" - "net/http" - "os" - "regexp" - "strconv" - "strings" - "sync" - "testing" - "time" - - "github.com/stretchr/testify/require" -) - -func itOnly(t *testing.T) { - t.Helper() - if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { - t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") - } -} - -func startProxy(t *testing.T, config ClickHouseProxyConfig) string { - t.Helper() - - proxy := NewClickHouseProxy(config) - - listener, err := net.Listen("tcp", "0.0.0.0:0") - require.NoError(t, err) - t.Cleanup(func() { listener.Close() }) - - ctx, cancel := context.WithCancel(context.Background()) - t.Cleanup(cancel) - - go func() { - for { - conn, acceptErr := listener.Accept() - if acceptErr != nil { - return - } - go func() { _ = proxy.HandleConnection(ctx, conn) }() - } - }() - - return fmt.Sprintf("%d", listener.Addr().(*net.TCPAddr).Port) -} - -func baseConfig(logger *recordingLogger, blocked ...string) ClickHouseProxyConfig { - patterns := make([]*regexp.Regexp, 0, len(blocked)) - for _, p := range blocked { - patterns = append(patterns, regexp.MustCompile(p)) - } - return ClickHouseProxyConfig{ - TargetAddr: envOr("PAM_CLICKHOUSE_HTTP", "127.0.0.1:8123"), - NativeAddr: envOr("PAM_CLICKHOUSE_NATIVE", "127.0.0.1:9000"), - Username: envOr("PAM_CLICKHOUSE_USER", "default"), - Password: envOr("PAM_CLICKHOUSE_PASSWORD", "clickhouse"), - Database: envOr("PAM_CLICKHOUSE_DB", "analytics"), - SessionID: "edge-test", - SessionLogger: logger, - BlockedCommands: patterns, - } -} - -func TestNativeCompressionBothWays(t *testing.T) { - itOnly(t) - - for _, compression := range []string{"1", "0"} { - t.Run("compression="+compression, func(t *testing.T) { - recorder := &recordingLogger{} - port := startProxy(t, baseConfig(recorder, `(?i)\bdrop\b`)) - - out, err := runClient(t, port, "SELECT 1;\nSELECT 'second';\nDROP TABLE pam_write_test;", - "--compression", compression) - t.Logf("%s", out) - - require.Error(t, err, "the blocked statement should fail the client") - require.Contains(t, out, "second", "earlier statements should still run") - require.Contains(t, out, "blocked by the command blocking policy") - require.True(t, recorder.contains("SELECT 'second'"), recorder.dump()) - }) - } -} - -func TestNativeInsertWithCompressionBothWays(t *testing.T) { - itOnly(t) - - for _, compression := range []string{"1", "0"} { - t.Run("compression="+compression, func(t *testing.T) { - port := startProxy(t, baseConfig(&recordingLogger{})) - marker := fmt.Sprintf("edge-%s-%d", compression, time.Now().UnixNano()) - - out, err := runClient(t, port, - fmt.Sprintf("INSERT INTO pam_write_test (id, note) VALUES (7001, '%s');\nSELECT note FROM pam_write_test WHERE note = '%s';", marker, marker), - "--compression", compression) - require.NoError(t, err, out) - require.Contains(t, out, marker) - }) - } -} - -// A column type ch-go cannot infer can only appear in the client direction on an INSERT. -func TestNativeInsertIntoUnreadableColumnFailsClosed(t *testing.T) { - itOnly(t) - - recorder := &recordingLogger{} - port := startProxy(t, baseConfig(recorder)) - - out, _ := runClient(t, port, - "INSERT INTO exotic (id, tags, pair, arr, lc, dec, ts, en) VALUES (99, {'x':1}, ('p',2), ['a'], 'low', 1.0, '2026-01-01 00:00:00.000', 'a');") - t.Logf("%s", out) - - require.Contains(t, out, "could not read the data block", - "an unreadable INSERT should be refused with an explanation") - require.Contains(t, out, "Map(String, UInt64)", "the message should name the offending type") - require.Contains(t, out, "HTTP interface", "the message should point at the way that works") - - // Refusing has to mean the rows never reach ClickHouse. - require.Equal(t, 0, countExotic(t, 99), "the refused rows should not have been written") -} - -func countExotic(t *testing.T, id int) int { - t.Helper() - - target := fmt.Sprintf("http://%s/?database=%s", - envOr("PAM_CLICKHOUSE_HTTP", "127.0.0.1:8123"), envOr("PAM_CLICKHOUSE_DB", "analytics")) - - req, err := http.NewRequest(http.MethodPost, target, - strings.NewReader(fmt.Sprintf("SELECT count() FROM exotic WHERE id = %d", id))) - require.NoError(t, err) - req.Header.Set("X-ClickHouse-User", envOr("PAM_CLICKHOUSE_USER", "default")) - req.Header.Set("X-ClickHouse-Key", envOr("PAM_CLICKHOUSE_PASSWORD", "clickhouse")) - - resp, err := (&http.Client{Timeout: 30 * time.Second}).Do(req) - require.NoError(t, err) - defer resp.Body.Close() - - raw, err := io.ReadAll(resp.Body) - require.NoError(t, err) - - count, err := strconv.Atoi(strings.TrimSpace(string(raw))) - require.NoError(t, err, string(raw)) - return count -} - -func TestBlockedStatementEndsTheSession(t *testing.T) { - itOnly(t) - - recorder := &recordingLogger{} - port := startProxy(t, baseConfig(recorder, `(?i)\bdrop\b`)) - - before := countWriteTest(t) - - out, err := runClient(t, port, - "SELECT 1;\nDROP TABLE pam_write_test;\nINSERT INTO pam_write_test (id, note) VALUES (7777, 'after-block');") - t.Logf("%s", out) - require.Error(t, err) - require.Contains(t, out, "blocked by the command blocking policy") - - // Neither the blocked DROP nor the statement behind it may have run. - require.Equal(t, before, countWriteTest(t), "nothing after a refusal should reach ClickHouse") - require.NotContains(t, out, "after-block") -} - -func countWriteTest(t *testing.T) int { - t.Helper() - - target := fmt.Sprintf("http://%s/?database=%s", - envOr("PAM_CLICKHOUSE_HTTP", "127.0.0.1:8123"), envOr("PAM_CLICKHOUSE_DB", "analytics")) - - req, err := http.NewRequest(http.MethodPost, target, strings.NewReader("SELECT count() FROM pam_write_test")) - require.NoError(t, err) - req.Header.Set("X-ClickHouse-User", envOr("PAM_CLICKHOUSE_USER", "default")) - req.Header.Set("X-ClickHouse-Key", envOr("PAM_CLICKHOUSE_PASSWORD", "clickhouse")) - - resp, err := (&http.Client{Timeout: 30 * time.Second}).Do(req) - require.NoError(t, err) - defer resp.Body.Close() - - raw, err := io.ReadAll(resp.Body) - require.NoError(t, err) - - count, err := strconv.Atoi(strings.TrimSpace(string(raw))) - require.NoError(t, err, string(raw)) - return count -} - -func tlsConfigFor(t *testing.T, insecure bool) *tls.Config { - t.Helper() - config := &tls.Config{ServerName: "localhost", InsecureSkipVerify: insecure} - if insecure { - return config - } - pem, err := os.ReadFile(os.Getenv("PAM_CLICKHOUSE_TLS_CERT")) - require.NoError(t, err) - pool := x509.NewCertPool() - require.True(t, pool.AppendCertsFromPEM(pem)) - config.RootCAs = pool - return config -} - -func tlsConfigSkipOrConfig(t *testing.T) (ClickHouseProxyConfig, bool) { - t.Helper() - native := os.Getenv("PAM_CLICKHOUSE_TLS_NATIVE") - httpAddr := os.Getenv("PAM_CLICKHOUSE_TLS_HTTP") - if native == "" || httpAddr == "" { - return ClickHouseProxyConfig{}, false - } - return ClickHouseProxyConfig{ - TargetAddr: httpAddr, - NativeAddr: native, - Username: "default", - Password: "clickhouse", - Database: "analytics", - EnableTLS: true, - TLSConfig: tlsConfigFor(t, true), - SessionID: "edge-tls-test", - SessionLogger: &recordingLogger{}, - }, true -} - -func TestTLSUpstream(t *testing.T) { - itOnly(t) - - config, ok := tlsConfigSkipOrConfig(t) - if !ok { - t.Skip("set PAM_CLICKHOUSE_TLS_NATIVE and PAM_CLICKHOUSE_TLS_HTTP to run") - } - - t.Run("native client over TLS to the server", func(t *testing.T) { - port := startProxy(t, config) - out, err := runClient(t, port, "SELECT note FROM t ORDER BY id;") - require.NoError(t, err, out) - require.Contains(t, out, "tls-one") - }) - - t.Run("http client over TLS to the server", func(t *testing.T) { - port := startProxy(t, config) - status, body := postStatement(t, "127.0.0.1:"+port, "SELECT count() AS c FROM t") - require.Equal(t, http.StatusOK, status, body) - require.Contains(t, body, "2") - }) - - t.Run("connection tests reach both interfaces over TLS", func(t *testing.T) { - require.NoError(t, TestConnection(context.Background(), config)) - require.NoError(t, TestNativeConnection(context.Background(), config)) - }) - - t.Run("a pinned CA verifies rather than skipping", func(t *testing.T) { - if os.Getenv("PAM_CLICKHOUSE_TLS_CERT") == "" { - t.Skip("set PAM_CLICKHOUSE_TLS_CERT to run") - } - verified := config - verified.TLSConfig = tlsConfigFor(t, false) - require.NoError(t, TestNativeConnection(context.Background(), verified)) - }) - - t.Run("an untrusted certificate is refused when verification is on", func(t *testing.T) { - strict := config - strict.TLSConfig = &tls.Config{ServerName: "localhost"} - err := TestNativeConnection(context.Background(), strict) - require.Error(t, err, "a self-signed certificate should not verify") - require.Contains(t, strings.ToLower(err.Error()), "certificate") - }) -} - -func TestAccountWithoutNativePortRefusesNativeClients(t *testing.T) { - itOnly(t) - - config := baseConfig(&recordingLogger{}) - config.NativeAddr = "" - port := startProxy(t, config) - - out, err := runClient(t, port, "SELECT 1;") - t.Logf("%s", out) - require.Error(t, err) - require.Contains(t, out, "native port") - - // The HTTP interface has to keep working on the same account. - status, body := postStatement(t, "127.0.0.1:"+port, "SELECT 1") - require.Equal(t, http.StatusOK, status, body) -} - -func TestAccountWithNeitherPortFailsClearly(t *testing.T) { - itOnly(t) - - config := baseConfig(&recordingLogger{}) - config.TargetAddr = "" - config.NativeAddr = "" - port := startProxy(t, config) - - // The session layer rejects this config before a handler ever runs, so the handler's own guard simply... - _, _, err := postStatementE("127.0.0.1:"+port, "SELECT 1") - require.Error(t, err, "a session with neither port must not serve anything") -} - -func TestCompressedRequestBodies(t *testing.T) { - itOnly(t) - - port := startProxy(t, baseConfig(&recordingLogger{})) - - status, body := postGzipped(t, "127.0.0.1:"+port, "SELECT 5 AS five \nFORMAT JSON") - require.Equal(t, http.StatusOK, status, body) - require.Contains(t, body, "\"five\": 5") -} - -func TestSnifferEdgeCases(t *testing.T) { - itOnly(t) - - port := startProxy(t, baseConfig(&recordingLogger{})) - - t.Run("a client that connects and says nothing is eventually dropped", func(t *testing.T) { - conn, err := net.Dial("tcp", "127.0.0.1:"+port) - require.NoError(t, err) - defer conn.Close() - - // The sniff deadline is what guarantees this; without it the handler would hold the session open. - require.NoError(t, conn.SetReadDeadline(time.Now().Add(sniffTimeout+10*time.Second))) - buf := make([]byte, 64) - _, err = conn.Read(buf) - require.Error(t, err, "the gateway should close a connection that never says anything") - require.NotErrorIs(t, err, os.ErrDeadlineExceeded, "the gateway held the connection open past the sniff timeout") - }) - - t.Run("garbage is not mistaken for either protocol", func(t *testing.T) { - conn, err := net.Dial("tcp", "127.0.0.1:"+port) - require.NoError(t, err) - defer conn.Close() - _, err = conn.Write([]byte{0xFF, 0xFE, 0xFD, 0xFC}) - require.NoError(t, err) - // Bytes that are not a request line leave net/http waiting for headers, so the bound here is its... - require.NoError(t, conn.SetReadDeadline(time.Now().Add(45*time.Second))) - - buf := make([]byte, 256) - n, err := conn.Read(buf) - require.NotErrorIs(t, err, os.ErrDeadlineExceeded, "garbage must not leave the handler hanging") - if err == nil || n > 0 { - require.Contains(t, string(buf[:n]), "400", "garbage should be answered as a bad HTTP request") - } else { - require.ErrorIs(t, err, io.EOF) - } - }) - - t.Run("the /ping path answers", func(t *testing.T) { - resp, err := (&http.Client{Timeout: 10 * time.Second}).Get("http://127.0.0.1:" + port + "/ping") - require.NoError(t, err) - defer resp.Body.Close() - require.Equal(t, http.StatusOK, resp.StatusCode) - }) - - t.Run("a path outside the query endpoint is refused", func(t *testing.T) { - resp, err := (&http.Client{Timeout: 10 * time.Second}).Get("http://127.0.0.1:" + port + "/play") - require.NoError(t, err) - defer resp.Body.Close() - require.Equal(t, http.StatusNotFound, resp.StatusCode) - }) -} - -func TestConcurrentMixedProtocolSessions(t *testing.T) { - itOnly(t) - - recorder := &recordingLogger{} - port := startProxy(t, baseConfig(recorder)) - - var wg sync.WaitGroup - errs := make(chan error, 16) - - for i := range 6 { - wg.Add(1) - go func(n int) { - defer wg.Done() - marker := fmt.Sprintf("concurrent-native-%d", n) - out, err := runClient(t, port, fmt.Sprintf("SELECT '%s';", marker)) - if err != nil { - errs <- fmt.Errorf("native %d: %v\n%s", n, err, out) - return - } - if !strings.Contains(out, marker) { - errs <- fmt.Errorf("native %d: missing marker in %s", n, out) - } - }(i) - } - - for i := range 6 { - wg.Add(1) - go func(n int) { - defer wg.Done() - status, body, err := postStatementE("127.0.0.1:"+port, fmt.Sprintf("SELECT %d AS n", n)) - if err != nil { - errs <- fmt.Errorf("http %d: %v", n, err) - return - } - if status != http.StatusOK { - errs <- fmt.Errorf("http %d: status %d: %s", n, status, body) - } - }(i) - } - - wg.Wait() - close(errs) - for err := range errs { - t.Error(err) - } -} - -func TestNativeLargeResultSet(t *testing.T) { - itOnly(t) - - port := startProxy(t, baseConfig(&recordingLogger{})) - // The rows have to actually cross the proxy, or none of the multi-block relay is exercised. - out, err := runClient(t, port, "SELECT number FROM numbers(300000);") - require.NoError(t, err, out) - - lines := strings.Count(strings.TrimSpace(out), "\n") + 1 - require.Equal(t, 300000, lines, "every row should reach the client") - require.Contains(t, out, "299999", "the last row should survive the relay") -} - -func TestNativeWideRowsStreamThrough(t *testing.T) { - itOnly(t) - - port := startProxy(t, baseConfig(&recordingLogger{})) - - direct := queryDirect(t, "SELECT sum(length(payload)) FROM wide_blobs") - require.NotEqual(t, "0", direct, "wide_blobs must be seeded for this to test anything") - - out, err := runClient(t, port, "SELECT sum(length(payload)) FROM wide_blobs;") - require.NoError(t, err, out) - require.Equal(t, direct, strings.TrimSpace(out), "the proxied total must match the server's") -} - -func queryDirect(t *testing.T, sql string) string { - t.Helper() - - target := fmt.Sprintf("http://%s/?database=%s", - envOr("PAM_CLICKHOUSE_HTTP", "127.0.0.1:8123"), envOr("PAM_CLICKHOUSE_DB", "analytics")) - req, err := http.NewRequest(http.MethodPost, target, strings.NewReader(sql)) - require.NoError(t, err) - req.Header.Set("X-ClickHouse-User", envOr("PAM_CLICKHOUSE_USER", "default")) - req.Header.Set("X-ClickHouse-Key", envOr("PAM_CLICKHOUSE_PASSWORD", "clickhouse")) - - resp, err := (&http.Client{Timeout: 60 * time.Second}).Do(req) - require.NoError(t, err) - defer resp.Body.Close() - - raw, err := io.ReadAll(resp.Body) - require.NoError(t, err) - return strings.TrimSpace(string(raw)) -} - -func TestUpstreamUnreachable(t *testing.T) { - itOnly(t) - - t.Run("native client gets a native exception", func(t *testing.T) { - config := baseConfig(&recordingLogger{}) - config.NativeAddr = "127.0.0.1:1" - port := startProxy(t, config) - - out, err := runClient(t, port, "SELECT 1;") - t.Logf("%s", out) - require.Error(t, err) - require.Contains(t, out, "could not reach ClickHouse") - }) - -} - -func TestWrongAccountCredentialsSurfaceCleanly(t *testing.T) { - itOnly(t) - - config := baseConfig(&recordingLogger{}) - config.Password = "definitely-not-the-password" - port := startProxy(t, config) - - out, err := runClient(t, port, "SELECT 1;") - t.Logf("%s", out) - require.Error(t, err) - require.Contains(t, out, "refused the account") -} - -func TestNativeRecordingCapturesOutcomes(t *testing.T) { - itOnly(t) - - recorder := &recordingLogger{} - port := startProxy(t, baseConfig(recorder)) - - out, err := runClient(t, port, "SELECT 1;\nSELECT * FROM nope_not_here;") - t.Logf("%s", out) - require.Error(t, err) - - waitFor(t, func() bool { return strings.Contains(recorder.dump(), "ERROR:") }) - dump := recorder.dump() - require.Contains(t, dump, "SELECT 1") - require.Contains(t, dump, "nope_not_here") - require.Equal(t, 1, strings.Count(dump, "=> OK"), "exactly the one successful statement keeps its outcome") - require.Equal(t, 1, strings.Count(dump, "ERROR:"), "the failed statement is recorded once") -} - -func TestNativeRevisionPinning(t *testing.T) { - itOnly(t) - - port := startProxy(t, baseConfig(&recordingLogger{})) - - // The server is newer than ch-go, so this only passes if the pinned revision is honoured end to end. - out, err := runClient(t, port, "SELECT version();") - require.NoError(t, err, out) - require.Equal(t, queryDirect(t, "SELECT version()"), strings.TrimSpace(out)) -} - -func TestAccountWithoutHTTPPortRefusesHTTPClients(t *testing.T) { - itOnly(t) - - config := baseConfig(&recordingLogger{}) - config.TargetAddr = "" - port := startProxy(t, config) - - status, body := postStatement(t, "127.0.0.1:"+port, "SELECT 1") - require.Equal(t, http.StatusBadGateway, status) - require.Contains(t, body, "does not have ClickHouse's HTTP port configured") - require.NotContains(t, body, "no Host in request URL", "the internal proxy error should not reach the client") - - // The native protocol still works on the same port. - out, err := runClient(t, port, "SELECT 'native-still-works';") - require.NoError(t, err, out) - require.Contains(t, out, "native-still-works") -} diff --git a/packages/pam/handlers/clickhouse/native_integration_test.go b/packages/pam/handlers/clickhouse/native_integration_test.go deleted file mode 100644 index dd636d101..000000000 --- a/packages/pam/handlers/clickhouse/native_integration_test.go +++ /dev/null @@ -1,311 +0,0 @@ -package clickhouse - -import ( - "context" - "fmt" - "net" - "os" - "regexp" - "strings" - "testing" - "time" -) - -func TestNativeIntegration(t *testing.T) { - if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { - t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") - } - - type testCase struct { - name string - sql string - blocked []string - wantOutput []string - wantFailure string - } - - cases := []testCase{ - { - name: "multiple statements on one connection", - sql: "SELECT 1;\nSELECT 2;\nSELECT 3;", - wantOutput: []string{"1", "2", "3"}, - }, - { - name: "credentials come from the account, not the client", - sql: "SELECT currentUser();", - wantOutput: []string{"default"}, - }, - { - name: "column types ch-go cannot infer still stream back", - sql: "SELECT map('a', 1::UInt64) AS m, tuple('p', 2) AS t;", - wantOutput: []string{"{'a':1}", "('p',2)"}, - }, - { - // A fixed marker would be satisfied by a previous run's row even if the insert path regressed. - name: "insert pushes a client data block", - sql: insertMarkerSQL(), - wantOutput: []string{insertMarker}, - }, - { - name: "a blocked statement is refused as a native exception", - sql: "DROP TABLE pam_write_test;", - blocked: []string{`(?i)\bdrop\b`}, - wantFailure: "blocked by the command blocking policy", - }, - { - name: "blocking still applies after an earlier statement on the same connection", - sql: "SELECT 1;\nDROP TABLE pam_write_test;", - blocked: []string{`(?i)\bdrop\b`}, - wantFailure: "blocked by the command blocking policy", - }, - } - - for _, tc := range cases { - t.Run(tc.name, func(t *testing.T) { - recorder := &recordingLogger{} - addr := startNativeProxy(t, tc.blocked, recorder) - - out, err := runClient(t, addr, tc.sql) - t.Logf("client output:\n%s", out) - - if tc.wantFailure != "" { - if err == nil { - t.Fatalf("expected the client to fail, got success:\n%s", out) - } - if !strings.Contains(out, tc.wantFailure) { - t.Fatalf("expected %q in the client output, got:\n%s", tc.wantFailure, out) - } - return - } - - if err != nil { - t.Fatalf("clickhouse-client failed: %v\n%s", err, out) - } - for _, want := range tc.wantOutput { - if !strings.Contains(out, want) { - t.Fatalf("expected %q in the client output, got:\n%s", want, out) - } - } - }) - } -} - -func TestNativeRecordsEveryStatement(t *testing.T) { - if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { - t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") - } - - recorder := &recordingLogger{} - addr := startNativeProxy(t, nil, recorder) - - if out, err := runClient(t, addr, "SELECT 1;\nSELECT 2;\nSELECT 3;"); err != nil { - t.Fatalf("clickhouse-client failed: %v\n%s", err, out) - } - - for _, want := range []string{"SELECT 1", "SELECT 2", "SELECT 3"} { - if !recorder.contains(want) { - t.Fatalf("expected %q in the session recording, got:\n%s", want, recorder.dump()) - } - } - - if !strings.Contains(recorder.dump(), "=> OK") { - t.Fatalf("expected the outcome of each statement to be recorded, got:\n%s", recorder.dump()) - } -} - -func TestNativeRecordsFailedStatement(t *testing.T) { - if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { - t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") - } - - recorder := &recordingLogger{} - addr := startNativeProxy(t, nil, recorder) - - if _, err := runClient(t, addr, "SELECT * FROM does_not_exist;"); err == nil { - t.Fatal("expected the statement to fail") - } - - waitFor(t, func() bool { return strings.Contains(recorder.dump(), "ERROR:") }) - - if !strings.Contains(recorder.dump(), "does_not_exist") { - t.Fatalf("expected the failed statement in the recording, got:\n%s", recorder.dump()) - } -} - -func TestNativeDegradesOnUnreadableResultBlock(t *testing.T) { - if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { - t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") - } - - recorder := &recordingLogger{} - addr := startNativeProxy(t, []string{`(?i)\bdrop\b`}, recorder) - - out, err := runClient(t, addr, - "SELECT map('a', 1::UInt64) AS m;\nSELECT 'after-the-map';\nDROP TABLE pam_write_test;") - t.Logf("client output:\n%s", out) - - if err == nil { - t.Fatalf("expected the blocked statement to fail the client:\n%s", out) - } - if !strings.Contains(out, "{'a':1}") { - t.Fatalf("expected the unreadable column type to still reach the client:\n%s", out) - } - if !strings.Contains(out, "after-the-map") { - t.Fatalf("expected the session to survive the unreadable block:\n%s", out) - } - // The security control has to keep working after the recorder degrades. - if !strings.Contains(out, "blocked by the command blocking policy") { - t.Fatalf("expected blocking to still apply after degrading:\n%s", out) - } - if !recorder.contains("after-the-map") { - t.Fatalf("expected statements to still be recorded after degrading, got:\n%s", recorder.dump()) - } - if !strings.Contains(recorder.dump(), "outcome could not be read") { - t.Fatalf("expected the recording to say outcomes stopped, got:\n%s", recorder.dump()) - } -} - -func waitFor(t *testing.T, condition func() bool) { - t.Helper() - for range 100 { - if condition() { - return - } - time.Sleep(20 * time.Millisecond) - } - t.Fatal("timed out waiting for the session recording") -} - -var insertMarker = fmt.Sprintf("native-it-%d", time.Now().UnixNano()) - -func insertMarkerSQL() string { - return fmt.Sprintf( - "INSERT INTO pam_write_test (id, note) VALUES (4242, '%s');\nSELECT note FROM pam_write_test WHERE note = '%s';", - insertMarker, insertMarker) -} - -func compileForTest(t *testing.T, blocked []string) []*regexp.Regexp { - t.Helper() - patterns := make([]*regexp.Regexp, 0, len(blocked)) - for _, p := range blocked { - patterns = append(patterns, regexp.MustCompile(p)) - } - return patterns -} - -func startNativeProxy(t *testing.T, blocked []string, logger *recordingLogger) string { - t.Helper() - - patterns := compileForTest(t, blocked) - - proxy := NewClickHouseProxy(ClickHouseProxyConfig{ - TargetAddr: envOr("PAM_CLICKHOUSE_HTTP", "127.0.0.1:8123"), - NativeAddr: envOr("PAM_CLICKHOUSE_NATIVE", "127.0.0.1:9000"), - Username: envOr("PAM_CLICKHOUSE_USER", "default"), - Password: envOr("PAM_CLICKHOUSE_PASSWORD", "clickhouse"), - Database: envOr("PAM_CLICKHOUSE_DB", "analytics"), - SessionID: "native-integration-test", - SessionLogger: logger, - BlockedCommands: patterns, - }) - - // The client runs in a container, so the listener has to be reachable from outside the loopback. - listener, err := net.Listen("tcp", "0.0.0.0:0") - if err != nil { - t.Fatalf("listen: %v", err) - } - t.Cleanup(func() { listener.Close() }) - - ctx, cancel := context.WithCancel(context.Background()) - t.Cleanup(cancel) - - go func() { - for { - conn, err := listener.Accept() - if err != nil { - return - } - go func() { _ = proxy.HandleConnection(ctx, conn) }() - } - }() - - return fmt.Sprintf("%d", listener.Addr().(*net.TCPAddr).Port) -} - -func envOr(name string, fallback string) string { - if v := os.Getenv(name); v != "" { - return v - } - return fallback -} - -func TestNativeConnectionTest(t *testing.T) { - if os.Getenv("PAM_CLICKHOUSE_NATIVE_IT") != "1" { - t.Skip("set PAM_CLICKHOUSE_NATIVE_IT=1 to run") - } - - native := envOr("PAM_CLICKHOUSE_NATIVE", "127.0.0.1:9000") - - cases := []struct { - name string - addr string - username string - password string - wantErr string - }{ - {name: "valid account", addr: native, username: "default", password: "clickhouse"}, - { - name: "wrong password is an auth failure, not a timeout", - addr: native, username: "default", password: "wrong", - wantErr: "clickhouse rejected the connection", - }, - { - name: "unknown user is reported as ClickHouse reported it", - addr: native, username: "nobody", password: "x", - wantErr: "clickhouse rejected the connection", - }, - { - name: "a port with nothing on it fails to dial", - addr: "127.0.0.1:1", username: "default", password: "clickhouse", - wantErr: "connect", - }, - { - // The HTTP port answers, so this proves the check is a real handshake rather than a dial. - name: "pointing the native check at the HTTP port fails", - addr: envOr("PAM_CLICKHOUSE_HTTP", "127.0.0.1:8123"), username: "default", password: "clickhouse", - wantErr: "", - }, - } - - for _, tc := range cases { - t.Run(tc.name, func(t *testing.T) { - err := TestNativeConnection(context.Background(), ClickHouseProxyConfig{ - NativeAddr: tc.addr, - Username: tc.username, - Password: tc.password, - Database: envOr("PAM_CLICKHOUSE_DB", "analytics"), - }) - - if tc.name == "pointing the native check at the HTTP port fails" { - if err == nil { - t.Fatal("expected the HTTP port to fail a native handshake") - } - t.Logf("got: %v", err) - return - } - - if tc.wantErr == "" { - if err != nil { - t.Fatalf("expected success, got %v", err) - } - return - } - if err == nil { - t.Fatalf("expected an error containing %q, got success", tc.wantErr) - } - if !strings.Contains(strings.ToLower(err.Error()), tc.wantErr) { - t.Fatalf("expected %q in %v", tc.wantErr, err) - } - }) - } -} diff --git a/packages/pam/handlers/clickhouse/testhelpers_test.go b/packages/pam/handlers/clickhouse/testhelpers_test.go deleted file mode 100644 index 459c90cdc..000000000 --- a/packages/pam/handlers/clickhouse/testhelpers_test.go +++ /dev/null @@ -1,101 +0,0 @@ -package clickhouse - -import ( - "bytes" - "compress/gzip" - "context" - "io" - "net/http" - "net/url" - "os/exec" - "strings" - "testing" - "time" - - "github.com/stretchr/testify/require" -) - -func runClient(t *testing.T, port string, sql string, extra ...string) (string, error) { - t.Helper() - - ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) - defer cancel() - - args := []string{ - "run", "--rm", "-i", "clickhouse/clickhouse-server:24.8", "clickhouse-client", - "--host", envOr("PAM_CLICKHOUSE_CLIENT_HOST", "host.docker.internal"), - "--port", port, - // Deliberately wrong: - "--user", "not-the-account", "--password", "not-the-password", - "--multiquery", - } - args = append(args, extra...) - - cmd := exec.CommandContext(ctx, "docker", args...) - cmd.Stdin = strings.NewReader(sql) - - out, err := cmd.CombinedOutput() - return string(out), err -} - -func postStatementE(addr string, sql string) (int, string, error) { - req, err := http.NewRequest(http.MethodPost, "http://"+addr+"/", strings.NewReader(sql)) - if err != nil { - return 0, "", err - } - resp, err := (&http.Client{Timeout: 60 * time.Second}).Do(req) - if err != nil { - return 0, "", err - } - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - if err != nil { - return resp.StatusCode, "", err - } - return resp.StatusCode, string(body), nil -} - -func postStatement(t *testing.T, addr string, sql string) (int, string) { - t.Helper() - - status, body, err := postStatementE(addr, sql) - require.NoError(t, err) - return status, body -} - -func postGET(t *testing.T, addr string, path string) (int, string) { - t.Helper() - - resp, err := (&http.Client{Timeout: 60 * time.Second}).Get("http://" + addr + path) - require.NoError(t, err) - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - return resp.StatusCode, string(body) -} - -func postGzipped(t *testing.T, addr string, sql string) (int, string) { - t.Helper() - - var buf bytes.Buffer - writer := gzip.NewWriter(&buf) - _, err := writer.Write([]byte(sql)) - require.NoError(t, err) - require.NoError(t, writer.Close()) - - req, err := http.NewRequest(http.MethodPost, "http://"+addr+"/", &buf) - require.NoError(t, err) - req.Header.Set("Content-Encoding", "gzip") - - resp, err := (&http.Client{Timeout: 60 * time.Second}).Do(req) - require.NoError(t, err) - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - return resp.StatusCode, string(body) -} - -func urlEscape(v string) string { return url.QueryEscape(v) } From 5adc23bc27e9362d3b481f4d78f181005503755f Mon Sep 17 00:00:00 2001 From: bernie-g Date: Fri, 25 Sep 2026 17:45:27 -0400 Subject: [PATCH 07/16] fix(clickhouse): check the executable SQL against the blocking policy The policy matched the recorded form of a statement, which carries a "-- parameters:" suffix, while ClickHouse received the bare SQL. An end-anchored rule stopped matching the moment a client attached a parameter and the blocked statement ran. Both forms are now checked, on the HTTP path as well as the native one. The e2e test drove clickhouse-client from inside the container, which can never reach a PAM proxy: those bind loopback only, and TestLocalProxiesBindLoopback enforces it. It now drives ch-go's client from the host instead. A handshake bounded by a shorter probe budget also reported the ten-second limit rather than the one it applied. --- e2e/go.mod | 4 +- e2e/go.sum | 4 + e2e/pam/clickhouse_test.go | 78 +++++++++++-------- packages/pam/handlers/clickhouse/native.go | 14 ++-- .../handlers/clickhouse/native_unit_test.go | 26 +++++++ packages/pam/handlers/clickhouse/proxy.go | 34 ++++---- .../pam/handlers/clickhouse/proxy_test.go | 17 ++++ 7 files changed, 121 insertions(+), 56 deletions(-) diff --git a/e2e/go.mod b/e2e/go.mod index 706837996..616bda718 100644 --- a/e2e/go.mod +++ b/e2e/go.mod @@ -3,6 +3,7 @@ module github.com/infisical/cli/e2e-tests go 1.25.14 require ( + github.com/ClickHouse/ch-go v0.74.0 github.com/Infisical/infisical-merge v0.0.0 github.com/compose-spec/compose-go/v2 v2.9.0 github.com/docker/compose/v2 v2.40.2 @@ -36,7 +37,6 @@ require ( github.com/Azure/go-ansiterm v0.0.0-20250102033503-faa5f7b0171c // indirect github.com/Azure/go-ntlmssp v0.1.1 // indirect github.com/ChrisTrenkamp/goxpath v0.0.0-20210404020558-97928f7e12b6 // indirect - github.com/ClickHouse/ch-go v0.74.0 // indirect github.com/DefangLabs/secret-detector v0.0.0-20250403165618-22662109213e // indirect github.com/Masterminds/goutils v1.1.1 // indirect github.com/Masterminds/semver/v3 v3.4.0 // indirect @@ -99,6 +99,7 @@ require ( github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect github.com/distribution/reference v0.6.0 // indirect github.com/dlclark/regexp2 v1.11.0 // indirect + github.com/dmarkham/enumer v1.6.3 // indirect github.com/docker/buildx v0.29.1 // indirect github.com/docker/cli v28.5.1+incompatible // indirect github.com/docker/cli-docs-tool v0.10.0 // indirect @@ -261,6 +262,7 @@ require ( github.com/opencontainers/go-digest v1.0.0 // indirect github.com/opencontainers/image-spec v1.1.1 // indirect github.com/oracle/oci-go-sdk/v65 v65.95.2 // indirect + github.com/pascaldekloe/name v1.0.1 // indirect github.com/pelletier/go-toml v1.9.5 // indirect github.com/pelletier/go-toml/v2 v2.4.3 // indirect github.com/perimeterx/marshmallow v1.1.5 // indirect diff --git a/e2e/go.sum b/e2e/go.sum index 517b4d2b9..023c43c88 100644 --- a/e2e/go.sum +++ b/e2e/go.sum @@ -287,6 +287,8 @@ github.com/distribution/reference v0.6.0 h1:0IXCQ5g4/QMHHkarYzh5l+u8T3t73zM5Qvfr github.com/distribution/reference v0.6.0/go.mod h1:BbU0aIcezP1/5jX/8MP0YiH4SdvB5Y4f/wlDRiLyi3E= github.com/dlclark/regexp2 v1.11.0 h1:G/nrcoOa7ZXlpoa/91N3X7mM3r8eIlMBBJZvsz/mxKI= github.com/dlclark/regexp2 v1.11.0/go.mod h1:DHkYz0B9wPfa6wondMfaivmHpzrQ3v9q8cnmRbL6yW8= +github.com/dmarkham/enumer v1.6.3 h1:B4aV4OsfzbrS5rvjILt4mMjiWBA//cKxJUMsvHZ8mEI= +github.com/dmarkham/enumer v1.6.3/go.mod h1:DyjXaqCglj4GhELF73oWiparNkYkXvmOBLza/o4kO74= github.com/docker/buildx v0.29.1 h1:58hxM5Z4mnNje3G5NKfULT9xCr8ooM8XFtlfUK9bKaA= github.com/docker/buildx v0.29.1/go.mod h1:J4EFv6oxlPiV1MjO0VyJx2u5tLM7ImDEl9zyB8d4wPI= github.com/docker/cli v28.5.1+incompatible h1:ESutzBALAD6qyCLqbQSEf1a/U8Ybms5agw59yGVc+yY= @@ -882,6 +884,8 @@ github.com/opentracing/opentracing-go v1.1.0/go.mod h1:UkNAQd3GIcIGf0SeVgPpRdFSt github.com/oracle/oci-go-sdk/v65 v65.95.2 h1:0HJ0AgpLydp/DtvYrF2d4str2BjXOVAeNbuW7E07g94= github.com/oracle/oci-go-sdk/v65 v65.95.2/go.mod h1:u6XRPsw9tPziBh76K7GrrRXPa8P8W3BQeqJ6ZZt9VLA= github.com/pascaldekloe/goe v0.0.0-20180627143212-57f6aae5913c/go.mod h1:lzWF7FIEvWOWxwDKqyGYQf6ZUaNfKdP144TG7ZOy1lc= +github.com/pascaldekloe/name v1.0.1 h1:9lnXOHeqeHHnWLbKfH6X98+4+ETVqFqxN09UXSjcMb0= +github.com/pascaldekloe/name v1.0.1/go.mod h1:Z//MfYJnH4jVpQ9wkclwu2I2MkHmXTlT9wR5UZScttM= github.com/pelletier/go-toml v1.2.0/go.mod h1:5z9KED0ma1S8pY6P1sdut58dfprrGBbd/94hg7ilaic= github.com/pelletier/go-toml v1.9.3/go.mod h1:u1nR/EPcESfeI/szUZKdtJ0xRNbUoANCkoOuaOx1Y+c= github.com/pelletier/go-toml v1.9.5 h1:4yBQzkHv+7BHq2PQUZF3Mx0IYxG7LsP222s7Agd3ve8= diff --git a/e2e/pam/clickhouse_test.go b/e2e/pam/clickhouse_test.go index d3e33c9be..7aa980909 100644 --- a/e2e/pam/clickhouse_test.go +++ b/e2e/pam/clickhouse_test.go @@ -10,6 +10,8 @@ import ( "testing" "time" + "github.com/ClickHouse/ch-go" + "github.com/ClickHouse/ch-go/proto" "github.com/docker/docker/api/types/container" "github.com/infisical/cli/e2e-tests/packages/client" helpers "github.com/infisical/cli/e2e-tests/util" @@ -66,28 +68,40 @@ func startClickHouseContainer(t *testing.T, ctx context.Context) (testcontainers return ctr, host, httpPort.Int(), nativePort.Int() } -// clickhouseClientIn runs the container's own clickhouse-client, which is what makes this an interop -// check rather than a test of our own encoder. -func clickhouseClientIn(t *testing.T, ctx context.Context, ctr testcontainers.Container, - host string, port int, sql string, extra ...string) (int, string) { +// queryOverNative drives the session with ch-go's client, which performs a real native handshake and +// query exchange. It runs on this host because a PAM proxy binds loopback only +// (TestLocalProxiesBindLoopback), so nothing inside a container can reach it. +func queryOverNative(t *testing.T, ctx context.Context, proxyPort int, sql string, compress bool) (string, error) { t.Helper() - args := []string{ - "clickhouse-client", - "--host", host, - "--port", fmt.Sprintf("%d", port), - "--database", clickhouseDatabase, + options := ch.Options{ + Address: fmt.Sprintf("127.0.0.1:%d", proxyPort), + Database: clickhouseDatabase, // The proxy injects the account's credentials, so whatever the client sends is discarded. - "--user", "not-the-account", "--password", "not-the-password", - "--query", sql, + User: "not-the-account", + Password: "not-the-password", + } + if compress { + options.Compression = ch.CompressionLZ4 } - args = append(args, extra...) - exitCode, reader, err := ctr.Exec(ctx, args) - require.NoError(t, err) - out, err := io.ReadAll(reader) - require.NoError(t, err) - return exitCode, string(out) + client, err := ch.Dial(ctx, options) + if err != nil { + return "", err + } + defer client.Close() + + var answer proto.ColStr + if err := client.Do(ctx, ch.Query{ + Body: sql, + Result: proto.Results{{Name: "answer", Data: &answer}}, + }); err != nil { + return "", err + } + if answer.Rows() == 0 { + return "", fmt.Errorf("no rows returned") + } + return answer.First(), nil } func seedClickHouse(t *testing.T, ctx context.Context, ctr testcontainers.Container) { @@ -216,9 +230,7 @@ func TestPAM_ClickHouse(t *testing.T) { templateId := CreatePamTemplate(t, ctx, infra, "clickhouse-template", client.CreatePamAccountTemplateJSONBodyType("clickhouse")) - ctr, chHost, chHTTPPort, chNativePort := startClickHouseContainer(t, ctx) - // The container reaches the proxy running on this host, so a loopback address will not do. - hostIP := getOutboundIP(t) + _, chHost, chHTTPPort, chNativePort := startClickHouseContainer(t, ctx) t.Run("both interfaces on one session port", func(t *testing.T) { accountName := "clickhouse-dual-account" @@ -233,9 +245,9 @@ func TestPAM_ClickHouse(t *testing.T) { require.Equal(t, "alpha\nbeta\ngamma", strings.TrimSpace(body)) slog.Info("HTTP interface answered through the proxy") - exitCode, out := clickhouseClientIn(t, ctx, ctr, hostIP, proxyPort, "SELECT count() FROM events") - require.Zero(t, exitCode, out) - require.Equal(t, "3", strings.TrimSpace(out)) + answer, err := queryOverNative(t, ctx, proxyPort, "SELECT toString(count()) AS answer FROM events", false) + require.NoError(t, err) + require.Equal(t, "3", answer) slog.Info("native interface answered the same session port") }) @@ -247,11 +259,11 @@ func TestPAM_ClickHouse(t *testing.T) { proxyPort, pamCmd := startClickHouseProxy(t, ctx, infra, folderName, accountName) waitForProxyHTTP(t, ctx, pamCmd, proxyPort) - for _, compression := range []string{"1", "0"} { - exitCode, out := clickhouseClientIn(t, ctx, ctr, hostIP, proxyPort, - "SELECT count() FROM events", "--compression", compression) - require.Zero(t, exitCode, out) - require.Equal(t, "3", strings.TrimSpace(out), "compression=%s", compression) + for _, compression := range []bool{true, false} { + answer, err := queryOverNative(t, ctx, proxyPort, + "SELECT toString(count()) AS answer FROM events", compression) + require.NoError(t, err, "compression=%v", compression) + require.Equal(t, "3", answer, "compression=%v", compression) } }) @@ -267,8 +279,8 @@ func TestPAM_ClickHouse(t *testing.T) { Interval: 2 * time.Second, Timeout: 60 * time.Second, Condition: func() helpers.ConditionResult { - exitCode, out := clickhouseClientIn(t, ctx, ctr, hostIP, proxyPort, "SELECT count() FROM events") - if exitCode == 0 && strings.TrimSpace(out) == "3" { + answer, err := queryOverNative(t, ctx, proxyPort, "SELECT toString(count()) AS answer FROM events", false) + if err == nil && answer == "3" { return helpers.ConditionSuccess } return helpers.ConditionWait @@ -290,9 +302,9 @@ func TestPAM_ClickHouse(t *testing.T) { proxyPort, pamCmd := startClickHouseProxy(t, ctx, infra, folderName, accountName) waitForProxyHTTP(t, ctx, pamCmd, proxyPort) - exitCode, out := clickhouseClientIn(t, ctx, ctr, hostIP, proxyPort, "SELECT count() FROM events") - require.NotZero(t, exitCode, out) - require.Contains(t, out, "native port", + _, err := queryOverNative(t, ctx, proxyPort, "SELECT toString(count()) AS answer FROM events", false) + require.Error(t, err) + require.Contains(t, err.Error(), "native port", "a native client must be told why, not left with a transport error") }) } diff --git a/packages/pam/handlers/clickhouse/native.go b/packages/pam/handlers/clickhouse/native.go index dfe4164a9..8c211c7e9 100644 --- a/packages/pam/handlers/clickhouse/native.go +++ b/packages/pam/handlers/clickhouse/native.go @@ -420,7 +420,7 @@ func (s *nativeSession) handleQuery(t *tap, r *proto.Reader) error { statement := q.Body + nativeParameterSuffix(q.Parameters) - if blocked := s.proxy.blockedBy(statement); blocked != nil { + if blocked := s.proxy.blockedBy(q.Body, statement); blocked != nil { s.proxy.logStatement(statement, fmt.Sprintf("BLOCKED: %s", blocked.String())) s.log.Info().Str("pattern", blocked.String()).Msg("Blocked a statement by policy") return s.refuse(t, codeAccessDenied, @@ -650,11 +650,13 @@ func TestNativeConnection(ctx context.Context, config ClickHouseProxyConfig) err defer conn.Close() // The probe's own budget wins when it is shorter, so a slow handshake cannot outlive the test. - deadline := time.Now().Add(nativeHandshakeTimeout) - if probeDeadline, ok := ctx.Deadline(); ok && probeDeadline.Before(deadline) { - deadline = probeDeadline + budget := nativeHandshakeTimeout + if probeDeadline, ok := ctx.Deadline(); ok { + if remaining := time.Until(probeDeadline); remaining < budget { + budget = remaining + } } - _ = conn.SetDeadline(deadline) + _ = conn.SetDeadline(time.Now().Add(budget)) var b proto.Buffer proto.ClientHello{ @@ -676,7 +678,7 @@ func TestNativeConnection(ctx context.Context, config ClickHouseProxyConfig) err if errors.Is(err, os.ErrDeadlineExceeded) { return fmt.Errorf("the port accepted the connection but did not answer ClickHouse's native "+ "handshake within %s, which is what the HTTP port does when it is entered as the native one: %w", - nativeHandshakeTimeout, err) + budget.Round(time.Second), err) } return fmt.Errorf("read hello response: %w", err) } diff --git a/packages/pam/handlers/clickhouse/native_unit_test.go b/packages/pam/handlers/clickhouse/native_unit_test.go index 595a0fc56..21550ccbc 100644 --- a/packages/pam/handlers/clickhouse/native_unit_test.go +++ b/packages/pam/handlers/clickhouse/native_unit_test.go @@ -557,6 +557,7 @@ func TestNativeConnectionTestClassifiesFailures(t *testing.T) { }) require.Error(t, err) require.Contains(t, err.Error(), "did not answer ClickHouse's native handshake") + require.Contains(t, err.Error(), "within 1s", "the message must name the budget that was applied") // The heartbeat stops scheduling on a rejected credential, so a silent port has to stay a // transport failure rather than being read as one. require.ErrorIs(t, err, os.ErrDeadlineExceeded) @@ -591,3 +592,28 @@ func TestNativeConnectionTestClassifiesFailures(t *testing.T) { require.NotContains(t, err.Error(), "did not answer ClickHouse's native handshake") }) } + +func TestNativeAnchoredRuleStillBlocksAStatementCarryingParameters(t *testing.T) { + upstream := startFakeClickHouse(t) + + conn := dialProxy(t, ClickHouseProxyConfig{ + NativeAddr: upstream.addr(), + Username: "account", + SessionID: "unit", + SessionLogger: &recordingLogger{}, + BlockedCommands: []*regexp.Regexp{regexp.MustCompile(`(?i)^DROP TABLE important$`)}, + }) + r := clientHandshake(t, conn, "someone", "whatever") + + writeQuery(t, conn, proto.Query{ + Body: "DROP TABLE important", + Parameters: []proto.Parameter{{Key: "who", Value: "someone"}}, + }) + + code, message := decodeException(t, r) + require.Equal(t, codeAccessDenied, code) + require.Contains(t, message, "blocked by the command blocking policy") + + _, _, queries, _ := upstream.snapshot() + require.Empty(t, queries, "the blocked statement must not reach the upstream") +} diff --git a/packages/pam/handlers/clickhouse/proxy.go b/packages/pam/handlers/clickhouse/proxy.go index a6ae35960..0f8b11b86 100644 --- a/packages/pam/handlers/clickhouse/proxy.go +++ b/packages/pam/handlers/clickhouse/proxy.go @@ -196,13 +196,13 @@ func (p *ClickHouseProxy) handler(l zerolog.Logger) http.Handler { return } - statement, body, err := p.inspect(r) + sql, statement, body, err := p.inspect(r) if err != nil { writeClickHouseError(w, http.StatusBadRequest, codeNotImplemented, err.Error()) return } - if blocked := p.blockedBy(statement); blocked != nil { + if blocked := p.blockedBy(sql, statement); blocked != nil { p.logStatement(statement, fmt.Sprintf("BLOCKED: %s", blocked.String())) l.Info().Str("pattern", blocked.String()).Msg("Blocked a statement by policy") writeClickHouseError(w, http.StatusForbidden, codeAccessDenied, @@ -232,16 +232,16 @@ type bodyReadCloser struct { } // Returns the statement and a body that still replays in full. ClickHouse concatenates `query` and the body. -func (p *ClickHouseProxy) inspect(r *http.Request) (string, io.ReadCloser, error) { +func (p *ClickHouseProxy) inspect(r *http.Request) (sql string, recorded string, body io.ReadCloser, err error) { queryParam := strings.TrimSpace(r.URL.Query().Get("query")) if r.Body == nil || r.ContentLength == 0 { - return queryParam + parameterSuffix(r.URL.Query()), http.NoBody, nil + return queryParam, queryParam + parameterSuffix(r.URL.Query()), http.NoBody, nil } // ClickHouse's own block compression is opaque to anything but a ClickHouse client if r.URL.Query().Get("decompress") == "1" { - return "", nil, fmt.Errorf( + return "", "", nil, fmt.Errorf( "this session cannot read a ClickHouse-compressed request body, so decompress=1 is not supported here. " + "Send the statement uncompressed or with Content-Encoding: gzip") } @@ -250,7 +250,7 @@ func (p *ClickHouseProxy) inspect(r *http.Request) (string, io.ReadCloser, error switch encoding { case "", "identity", "gzip", "deflate": default: - return "", nil, fmt.Errorf( + return "", "", nil, fmt.Errorf( "this session cannot read a %q-encoded request body, so the command blocking policy could not be applied to it. "+ "Use gzip, deflate, or no compression", encoding) } @@ -259,7 +259,7 @@ func (p *ClickHouseProxy) inspect(r *http.Request) (string, io.ReadCloser, error head := make([]byte, maxInspectBytes+1) n, err := io.ReadFull(r.Body, head) if err != nil && err != io.EOF && err != io.ErrUnexpectedEOF { - return "", nil, fmt.Errorf("the gateway could not read the request body: %v", err) + return "", "", nil, fmt.Errorf("the gateway could not read the request body: %v", err) } head = head[:n] @@ -267,18 +267,19 @@ func (p *ClickHouseProxy) inspect(r *http.Request) (string, io.ReadCloser, error decoded, decodedOverflow, decodeErr := decodeHead(head, encoding) if decodeErr != nil { - return "", nil, fmt.Errorf( + return "", "", nil, fmt.Errorf( "the gateway could not decompress the request body to apply the command blocking policy: %v", decodeErr) } if (len(head) > maxInspectBytes || decodedOverflow) && len(p.config.BlockedCommands) > 0 { - return "", nil, fmt.Errorf( + return "", "", nil, fmt.Errorf( "this account blocks commands, so a request body larger than %d MB is refused: the gateway has to read "+ "the whole statement to apply the policy. Send the data in smaller batches", maxInspectBytes>>20) } - return joinStatement(queryParam, string(decoded)) + parameterSuffix(r.URL.Query()), forwarded, nil + joined := joinStatement(queryParam, string(decoded)) + return joined, joined + parameterSuffix(r.URL.Query()), forwarded, nil } func parameterSuffix(query url.Values) string { @@ -468,13 +469,14 @@ func (p *ClickHouseProxy) handleUpstreamError(w http.ResponseWriter, r *http.Req fmt.Sprintf("The gateway could not reach ClickHouse: %v", err)) } -func (p *ClickHouseProxy) blockedBy(statement string) *regexp.Regexp { - if statement == "" { - return nil - } +// The recorded form carries a "-- parameters:" suffix, so an end-anchored rule stops matching the moment +// a client attaches one. Both the executable SQL and the recorded form are checked. +func (p *ClickHouseProxy) blockedBy(statements ...string) *regexp.Regexp { for _, pattern := range p.config.BlockedCommands { - if pattern.MatchString(statement) { - return pattern + for _, statement := range statements { + if statement != "" && pattern.MatchString(statement) { + return pattern + } } } return nil diff --git a/packages/pam/handlers/clickhouse/proxy_test.go b/packages/pam/handlers/clickhouse/proxy_test.go index 9560cd3a1..28cc0dbab 100644 --- a/packages/pam/handlers/clickhouse/proxy_test.go +++ b/packages/pam/handlers/clickhouse/proxy_test.go @@ -490,3 +490,20 @@ func TestRefusesADeflatedBodyItCannotDecode(t *testing.T) { require.False(t, reached, "a body the gateway could not read must not be forwarded uninspected") require.Equal(t, http.StatusBadRequest, recorder.Code, recorder.Body.String()) } + +// The recorded form carries a parameter suffix, so an end-anchored rule would stop matching as soon as +// a client attached a parameter and the blocked statement would run. +func TestAnAnchoredRuleStillBlocksAStatementCarryingParameters(t *testing.T) { + reached := false + handler, _, closeUpstream := newTestProxy(t, func(w http.ResponseWriter, r *http.Request) { + reached = true + }, `(?i)^DROP TABLE important$`) + defer closeUpstream() + + req := httptest.NewRequest(http.MethodPost, "/?param_who=someone", strings.NewReader("DROP TABLE important")) + recorder := httptest.NewRecorder() + handler.ServeHTTP(recorder, req) + + require.False(t, reached, "the blocked statement must not reach the upstream") + require.Equal(t, http.StatusForbidden, recorder.Code) +} From b0e19a63e9985270d9f4668ae0b312c1e2193a6e Mon Sep 17 00:00:00 2001 From: bernie-g Date: Fri, 25 Sep 2026 21:07:30 -0400 Subject: [PATCH 08/16] fix(clickhouse): bound every declared length in a native packet ch-go allocates a declared string length before it reads a byte and rejects only a length that goes negative, so a client naming a terabyte took the gateway down with it. Verified in a 512 MB container: 4 GB allocates fine because untouched pages never become resident, while 1 TB is SIGKILL, which no recover can catch, and it takes every other session on that gateway. The query packet is now decoded field for field with every string read through a cap, and the settings and parameters lists carry a count bound so neither can grow without limit. A round-trip test encodes with ch-go and decodes with ours across three revisions, since a field read in the wrong order would desynchronise the stream, which is the failure this exists to prevent. The handshake was already bounded; those helpers move alongside the rest. --- .../pam/handlers/clickhouse/bounded_decode.go | 311 ++++++++++++++++++ .../clickhouse/bounded_decode_test.go | 161 +++++++++ packages/pam/handlers/clickhouse/native.go | 52 +-- .../handlers/clickhouse/native_unit_test.go | 41 +++ 4 files changed, 516 insertions(+), 49 deletions(-) create mode 100644 packages/pam/handlers/clickhouse/bounded_decode.go create mode 100644 packages/pam/handlers/clickhouse/bounded_decode_test.go diff --git a/packages/pam/handlers/clickhouse/bounded_decode.go b/packages/pam/handlers/clickhouse/bounded_decode.go new file mode 100644 index 000000000..d32994c84 --- /dev/null +++ b/packages/pam/handlers/clickhouse/bounded_decode.go @@ -0,0 +1,311 @@ +package clickhouse + +import ( + "fmt" + "io" + + "github.com/ClickHouse/ch-go/proto" + "github.com/segmentio/asm/bswap" + "go.opentelemetry.io/otel/trace" +) + +// ch-go allocates a declared string length before it reads a single byte, and rejects only a length that +// goes negative. A client that names a terabyte therefore kills the process outright: the allocation is a +// fatal runtime error rather than a panic anything can recover. These decoders mirror ch-go's field for +// field and differ only in reading every string through a cap. +const ( + // Short identifiers: names, users, hostnames, the quota key. + maxHandshakeStringLen = 64 << 10 + // Query bodies and setting values. ClickHouse's own max_query_size defaults to 256 KB. + maxQueryStringLen = 16 << 20 + // A settings or parameters list terminates on an empty key, so it also needs a count bound. + maxQuerySettings = 4096 +) + +func readCappedStr(r *proto.Reader, limit int) (string, error) { + n, err := r.UVarInt() + if err != nil { + return "", err + } + if n > uint64(limit) { + return "", fmt.Errorf("declared string of %d bytes exceeds the %d byte cap", n, limit) + } + buf := make([]byte, n) + if _, err := io.ReadFull(r, buf); err != nil { + return "", err + } + return string(buf), nil +} + +func readBoundedStr(r *proto.Reader) (string, error) { + return readCappedStr(r, maxHandshakeStringLen) +} + +func decodeBoundedClientHello(r *proto.Reader) (proto.ClientHello, error) { + var h proto.ClientHello + var err error + if h.Name, err = readBoundedStr(r); err != nil { + return h, fmt.Errorf("name: %w", err) + } + if h.Major, err = r.Int(); err != nil { + return h, fmt.Errorf("major: %w", err) + } + if h.Minor, err = r.Int(); err != nil { + return h, fmt.Errorf("minor: %w", err) + } + if h.ProtocolVersion, err = r.Int(); err != nil { + return h, fmt.Errorf("protocol version: %w", err) + } + if h.Database, err = readBoundedStr(r); err != nil { + return h, fmt.Errorf("database: %w", err) + } + if h.User, err = readBoundedStr(r); err != nil { + return h, fmt.Errorf("user: %w", err) + } + if h.Password, err = readBoundedStr(r); err != nil { + return h, fmt.Errorf("password: %w", err) + } + return h, nil +} + +// Mirrors proto.Setting.Decode. An empty key terminates the list and leaves the rest unread. +func decodeBoundedSetting(r *proto.Reader) (proto.Setting, error) { + var s proto.Setting + + key, err := readBoundedStr(r) + if err != nil { + return s, fmt.Errorf("key: %w", err) + } + if key == "" { + return s, nil + } + + flags, err := r.UVarInt() + if err != nil { + return s, fmt.Errorf("flags: %w", err) + } + value, err := readCappedStr(r, maxQueryStringLen) + if err != nil { + return s, fmt.Errorf("value (%s): %w", key, err) + } + + s.Key = key + s.Value = value + s.Important = flags&0x01 != 0 + s.Custom = flags&0x02 != 0 + s.Obsolete = flags&0x04 != 0 + return s, nil +} + +// Mirrors proto.ClientInfo.DecodeAware. +func decodeBoundedClientInfo(r *proto.Reader, version int) (proto.ClientInfo, error) { + var c proto.ClientInfo + + kind, err := r.UInt8() + if err != nil { + return c, fmt.Errorf("query kind: %w", err) + } + c.Query = proto.ClientQueryKind(kind) + if !c.Query.IsAClientQueryKind() { + return c, fmt.Errorf("unknown query kind %d", kind) + } + + if c.InitialUser, err = readBoundedStr(r); err != nil { + return c, fmt.Errorf("initial user: %w", err) + } + if c.InitialQueryID, err = readBoundedStr(r); err != nil { + return c, fmt.Errorf("initial query id: %w", err) + } + if c.InitialAddress, err = readBoundedStr(r); err != nil { + return c, fmt.Errorf("initial address: %w", err) + } + + if proto.FeatureQueryStartTime.In(version) { + if c.InitialTime, err = r.Int64(); err != nil { + return c, fmt.Errorf("query start time: %w", err) + } + } + + iface, err := r.UInt8() + if err != nil { + return c, fmt.Errorf("interface: %w", err) + } + c.Interface = proto.Interface(iface) + if !c.Interface.IsAInterface() { + return c, fmt.Errorf("unknown interface %d", iface) + } + if c.Interface != proto.InterfaceTCP { + return c, fmt.Errorf("only tcp interface is supported") + } + + if c.OSUser, err = readBoundedStr(r); err != nil { + return c, fmt.Errorf("os user: %w", err) + } + if c.ClientHostname, err = readBoundedStr(r); err != nil { + return c, fmt.Errorf("client hostname: %w", err) + } + if c.ClientName, err = readBoundedStr(r); err != nil { + return c, fmt.Errorf("client name: %w", err) + } + if c.Major, err = r.Int(); err != nil { + return c, fmt.Errorf("major version: %w", err) + } + if c.Minor, err = r.Int(); err != nil { + return c, fmt.Errorf("minor version: %w", err) + } + if c.ProtocolVersion, err = r.Int(); err != nil { + return c, fmt.Errorf("protocol version: %w", err) + } + + if proto.FeatureQuotaKeyInClientInfo.In(version) { + if c.QuotaKey, err = readBoundedStr(r); err != nil { + return c, fmt.Errorf("quota key: %w", err) + } + } + if proto.FeatureDistributedDepth.In(version) { + if c.DistributedDepth, err = r.Int(); err != nil { + return c, fmt.Errorf("distributed depth: %w", err) + } + } + if proto.FeatureVersionPatch.In(version) && c.Interface == proto.InterfaceTCP { + if c.Patch, err = r.Int(); err != nil { + return c, fmt.Errorf("patch version: %w", err) + } + } + + if proto.FeatureOpenTelemetry.In(version) { + hasTrace, err := r.Bool() + if err != nil { + return c, fmt.Errorf("open telemetry start: %w", err) + } + if hasTrace { + var cfg trace.SpanContextConfig + raw, err := r.ReadRaw(len(cfg.TraceID)) + if err != nil { + return c, fmt.Errorf("trace id: %w", err) + } + bswap.Swap64(raw) + copy(cfg.TraceID[:], raw) + + raw, err = r.ReadRaw(len(cfg.SpanID)) + if err != nil { + return c, fmt.Errorf("span id: %w", err) + } + bswap.Swap64(raw) + copy(cfg.SpanID[:], raw) + + state, err := readBoundedStr(r) + if err != nil { + return c, fmt.Errorf("trace state: %w", err) + } + parsed, err := trace.ParseTraceState(state) + if err != nil { + return c, fmt.Errorf("parse trace state: %w", err) + } + cfg.TraceState = parsed + + flags, err := r.Byte() + if err != nil { + return c, fmt.Errorf("trace flag: %w", err) + } + cfg.TraceFlags = trace.TraceFlags(flags) + c.Span = trace.NewSpanContext(cfg) + } + } + + if proto.FeatureParallelReplicas.In(version) { + collaborate, err := r.Int() + if err != nil { + return c, fmt.Errorf("parallel replicas: %w", err) + } + c.CollaborateWithInitiator = collaborate == 1 + if c.CountParticipatingReplicas, err = r.Int(); err != nil { + return c, fmt.Errorf("count participating replicas: %w", err) + } + if c.NumberOfCurrentReplica, err = r.Int(); err != nil { + return c, fmt.Errorf("number of current replica: %w", err) + } + } + + return c, nil +} + +// Mirrors proto.Query.DecodeAware. +func decodeBoundedQuery(r *proto.Reader, version int) (proto.Query, error) { + var q proto.Query + var err error + + if q.ID, err = readBoundedStr(r); err != nil { + return q, fmt.Errorf("query id: %w", err) + } + + if proto.FeatureClientWriteInfo.In(version) { + if q.Info, err = decodeBoundedClientInfo(r, version); err != nil { + return q, fmt.Errorf("client info: %w", err) + } + } + + if !proto.FeatureSettingsSerializedAsStrings.In(version) { + return q, fmt.Errorf("unsupported version") + } + + for { + s, err := decodeBoundedSetting(r) + if err != nil { + return q, fmt.Errorf("setting: %w", err) + } + if s.Key == "" { + break + } + if len(q.Settings) >= maxQuerySettings { + return q, fmt.Errorf("more than %d settings", maxQuerySettings) + } + q.Settings = append(q.Settings, s) + } + + if proto.FeatureInterServerSecret.In(version) { + if q.Secret, err = readBoundedStr(r); err != nil { + return q, fmt.Errorf("inter-server secret: %w", err) + } + } + + stage, err := r.UVarInt() + if err != nil { + return q, fmt.Errorf("stage: %w", err) + } + q.Stage = proto.Stage(stage) + if !q.Stage.IsAStage() { + return q, fmt.Errorf("unknown stage %d", stage) + } + + compression, err := r.UVarInt() + if err != nil { + return q, fmt.Errorf("compression: %w", err) + } + q.Compression = proto.Compression(compression) + if !q.Compression.IsACompression() { + return q, fmt.Errorf("unknown compression %d", compression) + } + + if q.Body, err = readCappedStr(r, maxQueryStringLen); err != nil { + return q, fmt.Errorf("query body: %w", err) + } + + if proto.FeatureParameters.In(version) { + for { + s, err := decodeBoundedSetting(r) + if err != nil { + return q, fmt.Errorf("parameter: %w", err) + } + if s.Key == "" { + break + } + if len(q.Parameters) >= maxQuerySettings { + return q, fmt.Errorf("more than %d parameters", maxQuerySettings) + } + q.Parameters = append(q.Parameters, proto.Parameter{Key: s.Key, Value: s.Value}) + } + } + + return q, nil +} diff --git a/packages/pam/handlers/clickhouse/bounded_decode_test.go b/packages/pam/handlers/clickhouse/bounded_decode_test.go new file mode 100644 index 000000000..97ec6bf2a --- /dev/null +++ b/packages/pam/handlers/clickhouse/bounded_decode_test.go @@ -0,0 +1,161 @@ +package clickhouse + +import ( + "bytes" + "fmt" + "strings" + "testing" + + "github.com/ClickHouse/ch-go/proto" + "github.com/stretchr/testify/require" + "go.opentelemetry.io/otel/trace" +) + +func readerOver(b []byte) *proto.Reader { + return proto.NewReader(bytes.NewReader(b)) +} + +// A field read in the wrong order desynchronises the stream, which is the failure this decoder exists to +// prevent. Encoding with ch-go and decoding with ours is what pins the two together. +func TestBoundedQueryDecodeMatchesChGo(t *testing.T) { + span := trace.NewSpanContext(trace.SpanContextConfig{ + TraceID: trace.TraceID{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16}, + SpanID: trace.SpanID{9, 8, 7, 6, 5, 4, 3, 2}, + TraceFlags: trace.FlagsSampled, + }) + + cases := map[string]proto.Query{ + "a bare statement": { + Body: "SELECT 1", + Stage: proto.StageComplete, + }, + "settings and parameters": { + ID: "query-id", + Body: "SELECT {who:String}", + Secret: "interserver", + Stage: proto.StageComplete, + Compression: proto.CompressionEnabled, + Settings: []proto.Setting{ + {Key: "max_block_size", Value: "1024"}, + {Key: "readonly", Value: "1", Important: true}, + {Key: "obsolete_one", Value: "x", Obsolete: true}, + }, + Parameters: []proto.Parameter{ + {Key: "who", Value: "someone"}, + {Key: "other", Value: "value"}, + }, + }, + "a fully populated client info": { + Body: "SELECT 2", + Stage: proto.StageComplete, + Info: proto.ClientInfo{ + InitialUser: "initial", + InitialQueryID: "initial-query", + InitialTime: 1727212345, + OSUser: "os-user", + ClientHostname: "hostname", + ClientName: "client-name", + QuotaKey: "quota", + DistributedDepth: 2, + Patch: 7, + Span: span, + CollaborateWithInitiator: true, + CountParticipatingReplicas: 3, + NumberOfCurrentReplica: 1, + }, + }, + "a large but legal body": { + Body: "SELECT '" + strings.Repeat("x", 1<<20) + "'", + Stage: proto.StageComplete, + }, + } + + for name, query := range cases { + t.Run(name, func(t *testing.T) { + for _, rev := range []int{54455, 54458, maxNativeRevision} { + q := query + q.Info.Query = proto.ClientQueryInitial + q.Info.Interface = proto.InterfaceTCP + q.Info.InitialAddress = "127.0.0.1:0" + q.Info.Major, q.Info.Minor, q.Info.ProtocolVersion = 24, 8, rev + + var b proto.Buffer + q.EncodeAware(&b, rev) + + // The packet code ch-go writes first is consumed by the caller, so skip it here. + r := readerOver(b.Buf) + code, err := r.UVarInt() + require.NoError(t, err) + require.Equal(t, proto.ClientCodeQuery, proto.ClientCode(code)) + + var theirs proto.Query + require.NoError(t, theirs.DecodeAware(readerSkippingCode(t, b.Buf), rev)) + + ours, err := decodeBoundedQuery(readerSkippingCode(t, b.Buf), rev) + require.NoError(t, err, "revision %d", rev) + require.Equal(t, theirs, ours, "revision %d", rev) + } + }) + } +} + +func readerSkippingCode(t *testing.T, payload []byte) *proto.Reader { + t.Helper() + r := readerOver(payload) + _, err := r.UVarInt() + require.NoError(t, err) + return r +} + +func TestBoundedQueryRefusesAnOversizedField(t *testing.T) { + for name, build := range map[string]func(*proto.Buffer){ + "query id": func(b *proto.Buffer) { + b.PutUVarInt(uint64(maxHandshakeStringLen) + 1) + }, + "query body": func(b *proto.Buffer) { + var q proto.Query + q.ID = "id" + q.Info.Query = proto.ClientQueryInitial + q.Info.Interface = proto.InterfaceTCP + q.Info.InitialAddress = "127.0.0.1:0" + q.Info.Major, q.Info.Minor, q.Info.ProtocolVersion = 24, 8, maxNativeRevision + q.Stage = proto.StageComplete + q.Body = "SELECT 1" + + var full proto.Buffer + q.EncodeAware(&full, maxNativeRevision) + + // Re-encode everything up to the body, then declare an absurd length in its place. + trimmed := full.Buf[:bytes.LastIndex(full.Buf, []byte("SELECT 1"))-1] + b.Buf = append(b.Buf, trimmed[1:]...) // drop the packet code + b.PutUVarInt(1 << 40) + }, + } { + t.Run(name, func(t *testing.T) { + var b proto.Buffer + build(&b) + + _, err := decodeBoundedQuery(readerOver(b.Buf), maxNativeRevision) + require.Error(t, err) + require.Contains(t, err.Error(), "exceeds the", + "an absurd length must be refused before it is allocated") + }) + } +} + +func TestBoundedQueryRefusesTooManySettings(t *testing.T) { + var b proto.Buffer + q := proto.Query{ID: "id", Body: "SELECT 1", Stage: proto.StageComplete} + q.Info.Query = proto.ClientQueryInitial + q.Info.Interface = proto.InterfaceTCP + q.Info.InitialAddress = "127.0.0.1:0" + q.Info.Major, q.Info.Minor, q.Info.ProtocolVersion = 24, 8, maxNativeRevision + for i := 0; i <= maxQuerySettings; i++ { + q.Settings = append(q.Settings, proto.Setting{Key: fmt.Sprintf("s%d", i), Value: "1"}) + } + q.EncodeAware(&b, maxNativeRevision) + + _, err := decodeBoundedQuery(readerSkippingCode(t, b.Buf), maxNativeRevision) + require.Error(t, err) + require.Contains(t, err.Error(), "more than") +} diff --git a/packages/pam/handlers/clickhouse/native.go b/packages/pam/handlers/clickhouse/native.go index 8c211c7e9..55ffde9f1 100644 --- a/packages/pam/handlers/clickhouse/native.go +++ b/packages/pam/handlers/clickhouse/native.go @@ -78,52 +78,6 @@ func (t *tap) discard() { t.buf = nil } -// ch-go allocates a declared string length before it reads a single byte, so an unauthenticated client could -// name a terabyte and take the process down with it. Every handshake field is a short identifier. -const maxHandshakeStringLen = 64 << 10 - -func readBoundedStr(r *proto.Reader) (string, error) { - n, err := r.UVarInt() - if err != nil { - return "", err - } - if n > maxHandshakeStringLen { - return "", fmt.Errorf("handshake field of %d bytes exceeds the %d byte cap", n, maxHandshakeStringLen) - } - buf := make([]byte, n) - if _, err := io.ReadFull(r, buf); err != nil { - return "", err - } - return string(buf), nil -} - -func decodeBoundedClientHello(r *proto.Reader) (proto.ClientHello, error) { - var h proto.ClientHello - var err error - if h.Name, err = readBoundedStr(r); err != nil { - return h, fmt.Errorf("name: %w", err) - } - if h.Major, err = r.Int(); err != nil { - return h, fmt.Errorf("major: %w", err) - } - if h.Minor, err = r.Int(); err != nil { - return h, fmt.Errorf("minor: %w", err) - } - if h.ProtocolVersion, err = r.Int(); err != nil { - return h, fmt.Errorf("protocol version: %w", err) - } - if h.Database, err = readBoundedStr(r); err != nil { - return h, fmt.Errorf("database: %w", err) - } - if h.User, err = readBoundedStr(r); err != nil { - return h, fmt.Errorf("user: %w", err) - } - if h.Password, err = readBoundedStr(r); err != nil { - return h, fmt.Errorf("password: %w", err) - } - return h, nil -} - // A refusal ends the session: the stream is mid-packet, so carrying on would let a later packet flush the // refused bytes upstream. var errSessionRefused = errors.New("the session was refused") @@ -400,8 +354,8 @@ func (s *nativeSession) clientLoop(t *tap, r *proto.Reader) error { } func (s *nativeSession) handleQuery(t *tap, r *proto.Reader) error { - var q proto.Query - if err := q.DecodeAware(r, s.rev); err != nil { + q, err := decodeBoundedQuery(r, s.rev) + if err != nil { s.log.Warn().Err(err).Msg("Could not read a ClickHouse query packet") return s.refuse(t, codeNotImplemented, "This session could not read the query packet, so the command blocking policy could not be "+ @@ -443,7 +397,7 @@ func (s *nativeSession) handleQuery(t *tap, r *proto.Reader) error { // Decodes a block only far enough to find its end, then replays the client's bytes: re-encoding would mean // reproducing a serialization we do not own. func (s *nativeSession) handleData(t *tap, r *proto.Reader) error { - table, err := r.Str() + table, err := readBoundedStr(r) if err != nil { s.log.Warn().Err(err).Msg("Could not read a ClickHouse data packet") return s.refuse(t, codeNotImplemented, diff --git a/packages/pam/handlers/clickhouse/native_unit_test.go b/packages/pam/handlers/clickhouse/native_unit_test.go index 21550ccbc..9d8e24002 100644 --- a/packages/pam/handlers/clickhouse/native_unit_test.go +++ b/packages/pam/handlers/clickhouse/native_unit_test.go @@ -1,6 +1,7 @@ package clickhouse import ( + "bytes" "context" "io" "net" @@ -617,3 +618,43 @@ func TestNativeAnchoredRuleStillBlocksAStatementCarryingParameters(t *testing.T) _, _, queries, _ := upstream.snapshot() require.Empty(t, queries, "the blocked statement must not reach the upstream") } + +// An absurd declared length is a fatal allocation inside ch-go, not a panic anything can recover, so the +// session has to refuse it before the decoder ever sees it. +func TestNativeRefusesAnOversizedQueryBody(t *testing.T) { + upstream := startFakeClickHouse(t) + + conn := dialProxy(t, ClickHouseProxyConfig{ + NativeAddr: upstream.addr(), + Username: "account", + SessionID: "unit", + SessionLogger: &recordingLogger{}, + }) + r := clientHandshake(t, conn, "someone", "whatever") + + var q proto.Query + q.ID = "id" + q.Info.Query = proto.ClientQueryInitial + q.Info.Interface = proto.InterfaceTCP + q.Info.InitialAddress = "127.0.0.1:0" + q.Info.Major, q.Info.Minor, q.Info.ProtocolVersion = 24, 8, proto.Version + q.Stage = proto.StageComplete + q.Body = "SELECT 1" + + var full proto.Buffer + q.EncodeAware(&full, proto.Version) + + // Everything up to the body, then a terabyte in place of its length. + var b proto.Buffer + b.Buf = append(b.Buf, full.Buf[:bytes.LastIndex(full.Buf, []byte("SELECT 1"))-1]...) + b.PutUVarInt(1 << 40) + _, err := conn.Write(b.Buf) + require.NoError(t, err) + + code, message := decodeException(t, r) + require.Equal(t, codeNotImplemented, code) + require.Contains(t, message, "could not read the query packet") + + _, _, queries, _ := upstream.snapshot() + require.Empty(t, queries, "nothing may be forwarded from a packet the gateway refused") +} From b0f0a5b7b79c4b1bfd48fa26e1e0f4d95d767168 Mon Sep 17 00:00:00 2001 From: bernie-g Date: Fri, 25 Sep 2026 22:38:30 -0400 Subject: [PATCH 09/16] fix(clickhouse): wait for the e2e database before seeding it The ClickHouse ports start listening before the entrypoint has created CLICKHOUSE_DB, so seeding raced initialisation and clickhouse-client exited 81, UNKNOWN_DATABASE. The wait now runs a query against the database itself, which is what proves it is usable. The HTTP helper also asserted on its request, and the wait for the proxy to bind called it in a retry loop, so a connection refused while the port was coming up would have failed the test instead of retrying. --- e2e/pam/clickhouse_test.go | 36 ++++++++++++++++++++++++++++-------- 1 file changed, 28 insertions(+), 8 deletions(-) diff --git a/e2e/pam/clickhouse_test.go b/e2e/pam/clickhouse_test.go index 7aa980909..003658191 100644 --- a/e2e/pam/clickhouse_test.go +++ b/e2e/pam/clickhouse_test.go @@ -43,9 +43,16 @@ func startClickHouseContainer(t *testing.T, ctx context.Context) (testcontainers HostConfigModifier: func(hc *container.HostConfig) { hc.ExtraHosts = append(hc.ExtraHosts, "host.docker.internal:host-gateway") }, + // The ports listen before the entrypoint has created CLICKHOUSE_DB, so waiting on them alone + // races initialisation and seeding fails with UNKNOWN_DATABASE. WaitingFor: wait.ForAll( wait.ForListeningPort("8123/tcp"), wait.ForListeningPort("9000/tcp"), + wait.ForExec([]string{ + "clickhouse-client", + "--user", clickhouseUser, "--password", clickhousePassword, + "--database", clickhouseDatabase, "--query", "SELECT 1", + }).WithExitCode(0), ).WithStartupTimeout(180 * time.Second), }, Started: true, @@ -183,20 +190,32 @@ func startClickHouseProxy(t *testing.T, ctx context.Context, infra *PAMTestInfra return freePort, &pamCmd } -func queryOverHTTP(t *testing.T, ctx context.Context, proxyPort int, sql string) (int, string) { - t.Helper() - +func tryQueryOverHTTP(ctx context.Context, proxyPort int, sql string) (int, string, error) { url := fmt.Sprintf("http://127.0.0.1:%d/?database=%s", proxyPort, clickhouseDatabase) req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, strings.NewReader(sql)) - require.NoError(t, err) + if err != nil { + return 0, "", err + } resp, err := (&http.Client{Timeout: 60 * time.Second}).Do(req) - require.NoError(t, err) + if err != nil { + return 0, "", err + } defer resp.Body.Close() body, err := io.ReadAll(resp.Body) + if err != nil { + return resp.StatusCode, "", err + } + return resp.StatusCode, string(body), nil +} + +func queryOverHTTP(t *testing.T, ctx context.Context, proxyPort int, sql string) (int, string) { + t.Helper() + + status, body, err := tryQueryOverHTTP(ctx, proxyPort, sql) require.NoError(t, err) - return resp.StatusCode, string(body) + return status, body } // waitForProxyHTTP absorbs the gap between the banner and the listener accepting. @@ -208,8 +227,9 @@ func waitForProxyHTTP(t *testing.T, ctx context.Context, pamCmd *helpers.Command Interval: 2 * time.Second, Timeout: 60 * time.Second, Condition: func() helpers.ConditionResult { - status, _ := queryOverHTTP(t, ctx, proxyPort, "SELECT 1") - if status == http.StatusOK { + // Must not assert: the proxy may not have bound the port yet and the wait has to retry. + status, _, err := tryQueryOverHTTP(ctx, proxyPort, "SELECT 1") + if err == nil && status == http.StatusOK { return helpers.ConditionSuccess } return helpers.ConditionWait From 650ec0224e8349324f76ef10122aad823d2398a8 Mon Sep 17 00:00:00 2001 From: bernie-g Date: Mon, 28 Sep 2026 10:21:21 -0400 Subject: [PATCH 10/16] test(clickhouse): drop the e2e test until the backend is on main The PAM e2e job checks out infisical/main with no ref, so a CLI test can only exercise backend behaviour that has already merged. Against main the account's nativePort is stripped as an unknown key and port is still required, so every ClickHouse account comes back HTTP-only and three of the four subtests fail on a backend that predates the feature. Every other PAM e2e test landed this way, in its own PR once the backend was in place: redis shipped in December and its test followed in April. This one follows the same route. The go.mod tidy stays, since that failure is real and independent. --- e2e/go.mod | 4 +- e2e/go.sum | 4 - e2e/pam/clickhouse_test.go | 330 ------------------------------------- 3 files changed, 1 insertion(+), 337 deletions(-) delete mode 100644 e2e/pam/clickhouse_test.go diff --git a/e2e/go.mod b/e2e/go.mod index 616bda718..706837996 100644 --- a/e2e/go.mod +++ b/e2e/go.mod @@ -3,7 +3,6 @@ module github.com/infisical/cli/e2e-tests go 1.25.14 require ( - github.com/ClickHouse/ch-go v0.74.0 github.com/Infisical/infisical-merge v0.0.0 github.com/compose-spec/compose-go/v2 v2.9.0 github.com/docker/compose/v2 v2.40.2 @@ -37,6 +36,7 @@ require ( github.com/Azure/go-ansiterm v0.0.0-20250102033503-faa5f7b0171c // indirect github.com/Azure/go-ntlmssp v0.1.1 // indirect github.com/ChrisTrenkamp/goxpath v0.0.0-20210404020558-97928f7e12b6 // indirect + github.com/ClickHouse/ch-go v0.74.0 // indirect github.com/DefangLabs/secret-detector v0.0.0-20250403165618-22662109213e // indirect github.com/Masterminds/goutils v1.1.1 // indirect github.com/Masterminds/semver/v3 v3.4.0 // indirect @@ -99,7 +99,6 @@ require ( github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect github.com/distribution/reference v0.6.0 // indirect github.com/dlclark/regexp2 v1.11.0 // indirect - github.com/dmarkham/enumer v1.6.3 // indirect github.com/docker/buildx v0.29.1 // indirect github.com/docker/cli v28.5.1+incompatible // indirect github.com/docker/cli-docs-tool v0.10.0 // indirect @@ -262,7 +261,6 @@ require ( github.com/opencontainers/go-digest v1.0.0 // indirect github.com/opencontainers/image-spec v1.1.1 // indirect github.com/oracle/oci-go-sdk/v65 v65.95.2 // indirect - github.com/pascaldekloe/name v1.0.1 // indirect github.com/pelletier/go-toml v1.9.5 // indirect github.com/pelletier/go-toml/v2 v2.4.3 // indirect github.com/perimeterx/marshmallow v1.1.5 // indirect diff --git a/e2e/go.sum b/e2e/go.sum index 023c43c88..517b4d2b9 100644 --- a/e2e/go.sum +++ b/e2e/go.sum @@ -287,8 +287,6 @@ github.com/distribution/reference v0.6.0 h1:0IXCQ5g4/QMHHkarYzh5l+u8T3t73zM5Qvfr github.com/distribution/reference v0.6.0/go.mod h1:BbU0aIcezP1/5jX/8MP0YiH4SdvB5Y4f/wlDRiLyi3E= github.com/dlclark/regexp2 v1.11.0 h1:G/nrcoOa7ZXlpoa/91N3X7mM3r8eIlMBBJZvsz/mxKI= github.com/dlclark/regexp2 v1.11.0/go.mod h1:DHkYz0B9wPfa6wondMfaivmHpzrQ3v9q8cnmRbL6yW8= -github.com/dmarkham/enumer v1.6.3 h1:B4aV4OsfzbrS5rvjILt4mMjiWBA//cKxJUMsvHZ8mEI= -github.com/dmarkham/enumer v1.6.3/go.mod h1:DyjXaqCglj4GhELF73oWiparNkYkXvmOBLza/o4kO74= github.com/docker/buildx v0.29.1 h1:58hxM5Z4mnNje3G5NKfULT9xCr8ooM8XFtlfUK9bKaA= github.com/docker/buildx v0.29.1/go.mod h1:J4EFv6oxlPiV1MjO0VyJx2u5tLM7ImDEl9zyB8d4wPI= github.com/docker/cli v28.5.1+incompatible h1:ESutzBALAD6qyCLqbQSEf1a/U8Ybms5agw59yGVc+yY= @@ -884,8 +882,6 @@ github.com/opentracing/opentracing-go v1.1.0/go.mod h1:UkNAQd3GIcIGf0SeVgPpRdFSt github.com/oracle/oci-go-sdk/v65 v65.95.2 h1:0HJ0AgpLydp/DtvYrF2d4str2BjXOVAeNbuW7E07g94= github.com/oracle/oci-go-sdk/v65 v65.95.2/go.mod h1:u6XRPsw9tPziBh76K7GrrRXPa8P8W3BQeqJ6ZZt9VLA= github.com/pascaldekloe/goe v0.0.0-20180627143212-57f6aae5913c/go.mod h1:lzWF7FIEvWOWxwDKqyGYQf6ZUaNfKdP144TG7ZOy1lc= -github.com/pascaldekloe/name v1.0.1 h1:9lnXOHeqeHHnWLbKfH6X98+4+ETVqFqxN09UXSjcMb0= -github.com/pascaldekloe/name v1.0.1/go.mod h1:Z//MfYJnH4jVpQ9wkclwu2I2MkHmXTlT9wR5UZScttM= github.com/pelletier/go-toml v1.2.0/go.mod h1:5z9KED0ma1S8pY6P1sdut58dfprrGBbd/94hg7ilaic= github.com/pelletier/go-toml v1.9.3/go.mod h1:u1nR/EPcESfeI/szUZKdtJ0xRNbUoANCkoOuaOx1Y+c= github.com/pelletier/go-toml v1.9.5 h1:4yBQzkHv+7BHq2PQUZF3Mx0IYxG7LsP222s7Agd3ve8= diff --git a/e2e/pam/clickhouse_test.go b/e2e/pam/clickhouse_test.go deleted file mode 100644 index 003658191..000000000 --- a/e2e/pam/clickhouse_test.go +++ /dev/null @@ -1,330 +0,0 @@ -package pam - -import ( - "context" - "fmt" - "io" - "log/slog" - "net/http" - "strings" - "testing" - "time" - - "github.com/ClickHouse/ch-go" - "github.com/ClickHouse/ch-go/proto" - "github.com/docker/docker/api/types/container" - "github.com/infisical/cli/e2e-tests/packages/client" - helpers "github.com/infisical/cli/e2e-tests/util" - openapitypes "github.com/oapi-codegen/runtime/types" - "github.com/stretchr/testify/require" - "github.com/testcontainers/testcontainers-go" - "github.com/testcontainers/testcontainers-go/wait" -) - -const ( - clickhouseImage = "clickhouse/clickhouse-server:24.8" - clickhouseDatabase = "analytics" - clickhouseUser = "default" - clickhousePassword = "clickhouse" -) - -func startClickHouseContainer(t *testing.T, ctx context.Context) (testcontainers.Container, string, int, int) { - t.Helper() - - ctr, err := testcontainers.GenericContainer(ctx, testcontainers.GenericContainerRequest{ - ContainerRequest: testcontainers.ContainerRequest{ - Image: clickhouseImage, - ExposedPorts: []string{"8123/tcp", "9000/tcp"}, - Env: map[string]string{ - "CLICKHOUSE_DB": clickhouseDatabase, - "CLICKHOUSE_USER": clickhouseUser, - "CLICKHOUSE_PASSWORD": clickhousePassword, - }, - HostConfigModifier: func(hc *container.HostConfig) { - hc.ExtraHosts = append(hc.ExtraHosts, "host.docker.internal:host-gateway") - }, - // The ports listen before the entrypoint has created CLICKHOUSE_DB, so waiting on them alone - // races initialisation and seeding fails with UNKNOWN_DATABASE. - WaitingFor: wait.ForAll( - wait.ForListeningPort("8123/tcp"), - wait.ForListeningPort("9000/tcp"), - wait.ForExec([]string{ - "clickhouse-client", - "--user", clickhouseUser, "--password", clickhousePassword, - "--database", clickhouseDatabase, "--query", "SELECT 1", - }).WithExitCode(0), - ).WithStartupTimeout(180 * time.Second), - }, - Started: true, - }) - require.NoError(t, err) - t.Cleanup(func() { - if err := ctr.Terminate(ctx); err != nil { - t.Logf("Failed to terminate ClickHouse container: %v", err) - } - }) - - host, err := ctr.Host(ctx) - require.NoError(t, err) - httpPort, err := ctr.MappedPort(ctx, "8123") - require.NoError(t, err) - nativePort, err := ctr.MappedPort(ctx, "9000") - require.NoError(t, err) - - seedClickHouse(t, ctx, ctr) - return ctr, host, httpPort.Int(), nativePort.Int() -} - -// queryOverNative drives the session with ch-go's client, which performs a real native handshake and -// query exchange. It runs on this host because a PAM proxy binds loopback only -// (TestLocalProxiesBindLoopback), so nothing inside a container can reach it. -func queryOverNative(t *testing.T, ctx context.Context, proxyPort int, sql string, compress bool) (string, error) { - t.Helper() - - options := ch.Options{ - Address: fmt.Sprintf("127.0.0.1:%d", proxyPort), - Database: clickhouseDatabase, - // The proxy injects the account's credentials, so whatever the client sends is discarded. - User: "not-the-account", - Password: "not-the-password", - } - if compress { - options.Compression = ch.CompressionLZ4 - } - - client, err := ch.Dial(ctx, options) - if err != nil { - return "", err - } - defer client.Close() - - var answer proto.ColStr - if err := client.Do(ctx, ch.Query{ - Body: sql, - Result: proto.Results{{Name: "answer", Data: &answer}}, - }); err != nil { - return "", err - } - if answer.Rows() == 0 { - return "", fmt.Errorf("no rows returned") - } - return answer.First(), nil -} - -func seedClickHouse(t *testing.T, ctx context.Context, ctr testcontainers.Container) { - t.Helper() - - statements := []string{ - "CREATE TABLE IF NOT EXISTS events (id UInt64, name String) ENGINE = MergeTree ORDER BY id", - "INSERT INTO events VALUES (1, 'alpha'), (2, 'beta'), (3, 'gamma')", - } - for _, sql := range statements { - exitCode, _, err := ctr.Exec(ctx, []string{ - "clickhouse-client", - "--user", clickhouseUser, "--password", clickhousePassword, - "--database", clickhouseDatabase, - "--query", sql, - }) - require.NoError(t, err) - require.Zero(t, exitCode, "seeding failed for: %s", sql) - } -} - -func createClickHousePamAccount(t *testing.T, ctx context.Context, infra *PAMTestInfra, - folderId, templateId openapitypes.UUID, name, host string, httpPort, nativePort *int) { - t.Helper() - - connectionDetails := map[string]interface{}{ - "host": host, - "database": clickhouseDatabase, - "sslEnabled": false, - "sslRejectUnauthorized": false, - } - if httpPort != nil { - connectionDetails["port"] = *httpPort - } - if nativePort != nil { - connectionDetails["nativePort"] = *nativePort - } - - CreatePamAccount(t, ctx, infra, "clickhouse", name, folderId, templateId, connectionDetails, - map[string]interface{}{"username": clickhouseUser, "password": clickhousePassword}) -} - -func startClickHouseProxy(t *testing.T, ctx context.Context, infra *PAMTestInfra, - folderName, accountName string) (int, *helpers.Command) { - t.Helper() - - freePort := helpers.GetFreePort() - pamCmd := helpers.Command{ - Test: t, - RunMethod: helpers.RunMethodSubprocess, - DisableTempHomeDir: true, - Args: []string{ - "pam", "access", fmt.Sprintf("%s/%s", folderName, accountName), - "--duration", "5m", - "--port", fmt.Sprintf("%d", freePort), - }, - Env: map[string]string{ - "HOME": infra.SharedHomeDir, - "INFISICAL_API_URL": infra.Infisical.ApiUrl(t), - }, - } - pamCmd.Start(ctx) - t.Cleanup(pamCmd.Stop) - - result := helpers.WaitFor(t, helpers.WaitForOptions{ - EnsureCmdRunning: &pamCmd, - Condition: func() helpers.ConditionResult { - if strings.Contains(pamCmd.Stdout(), "ClickHouse Proxy Session Started") { - return helpers.ConditionSuccess - } - return helpers.ConditionWait - }, - }) - if result != helpers.WaitSuccess { - infra.DumpOutput(&pamCmd) - } - require.Equal(t, helpers.WaitSuccess, result, "ClickHouse proxy should start successfully") - - return freePort, &pamCmd -} - -func tryQueryOverHTTP(ctx context.Context, proxyPort int, sql string) (int, string, error) { - url := fmt.Sprintf("http://127.0.0.1:%d/?database=%s", proxyPort, clickhouseDatabase) - req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, strings.NewReader(sql)) - if err != nil { - return 0, "", err - } - - resp, err := (&http.Client{Timeout: 60 * time.Second}).Do(req) - if err != nil { - return 0, "", err - } - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - if err != nil { - return resp.StatusCode, "", err - } - return resp.StatusCode, string(body), nil -} - -func queryOverHTTP(t *testing.T, ctx context.Context, proxyPort int, sql string) (int, string) { - t.Helper() - - status, body, err := tryQueryOverHTTP(ctx, proxyPort, sql) - require.NoError(t, err) - return status, body -} - -// waitForProxyHTTP absorbs the gap between the banner and the listener accepting. -func waitForProxyHTTP(t *testing.T, ctx context.Context, pamCmd *helpers.Command, proxyPort int) { - t.Helper() - - result := helpers.WaitFor(t, helpers.WaitForOptions{ - EnsureCmdRunning: pamCmd, - Interval: 2 * time.Second, - Timeout: 60 * time.Second, - Condition: func() helpers.ConditionResult { - // Must not assert: the proxy may not have bound the port yet and the wait has to retry. - status, _, err := tryQueryOverHTTP(ctx, proxyPort, "SELECT 1") - if err == nil && status == http.StatusOK { - return helpers.ConditionSuccess - } - return helpers.ConditionWait - }, - }) - require.Equal(t, helpers.WaitSuccess, result, "the proxy should answer HTTP") -} - -func TestPAM_ClickHouse(t *testing.T) { - ctx, cancel := context.WithCancel(context.Background()) - t.Cleanup(cancel) - - infra := SetupPAMInfra(t, ctx) - LoginUser(t, ctx, infra) - - folderName := "clickhouse-folder" - folderId := CreatePamFolder(t, ctx, infra, folderName) - templateId := CreatePamTemplate(t, ctx, infra, "clickhouse-template", - client.CreatePamAccountTemplateJSONBodyType("clickhouse")) - - _, chHost, chHTTPPort, chNativePort := startClickHouseContainer(t, ctx) - - t.Run("both interfaces on one session port", func(t *testing.T) { - accountName := "clickhouse-dual-account" - createClickHousePamAccount(t, ctx, infra, folderId, templateId, accountName, chHost, - &chHTTPPort, &chNativePort) - - proxyPort, pamCmd := startClickHouseProxy(t, ctx, infra, folderName, accountName) - waitForProxyHTTP(t, ctx, pamCmd, proxyPort) - - status, body := queryOverHTTP(t, ctx, proxyPort, "SELECT name FROM events ORDER BY id") - require.Equal(t, http.StatusOK, status, body) - require.Equal(t, "alpha\nbeta\ngamma", strings.TrimSpace(body)) - slog.Info("HTTP interface answered through the proxy") - - answer, err := queryOverNative(t, ctx, proxyPort, "SELECT toString(count()) AS answer FROM events", false) - require.NoError(t, err) - require.Equal(t, "3", answer) - slog.Info("native interface answered the same session port") - }) - - t.Run("compression is negotiated end to end", func(t *testing.T) { - accountName := "clickhouse-compression-account" - createClickHousePamAccount(t, ctx, infra, folderId, templateId, accountName, chHost, - &chHTTPPort, &chNativePort) - - proxyPort, pamCmd := startClickHouseProxy(t, ctx, infra, folderName, accountName) - waitForProxyHTTP(t, ctx, pamCmd, proxyPort) - - for _, compression := range []bool{true, false} { - answer, err := queryOverNative(t, ctx, proxyPort, - "SELECT toString(count()) AS answer FROM events", compression) - require.NoError(t, err, "compression=%v", compression) - require.Equal(t, "3", answer, "compression=%v", compression) - } - }) - - t.Run("a native-only account turns HTTP clients away", func(t *testing.T) { - accountName := "clickhouse-native-only-account" - createClickHousePamAccount(t, ctx, infra, folderId, templateId, accountName, chHost, - nil, &chNativePort) - - proxyPort, pamCmd := startClickHouseProxy(t, ctx, infra, folderName, accountName) - - result := helpers.WaitFor(t, helpers.WaitForOptions{ - EnsureCmdRunning: pamCmd, - Interval: 2 * time.Second, - Timeout: 60 * time.Second, - Condition: func() helpers.ConditionResult { - answer, err := queryOverNative(t, ctx, proxyPort, "SELECT toString(count()) AS answer FROM events", false) - if err == nil && answer == "3" { - return helpers.ConditionSuccess - } - return helpers.ConditionWait - }, - }) - require.Equal(t, helpers.WaitSuccess, result, "a native client should still work") - - status, body := queryOverHTTP(t, ctx, proxyPort, "SELECT 1") - require.NotEqual(t, http.StatusOK, status, body) - require.Contains(t, body, "HTTP port", - "an HTTP client must be told why, not left with a transport error") - }) - - t.Run("an HTTP-only account turns native clients away", func(t *testing.T) { - accountName := "clickhouse-http-only-account" - createClickHousePamAccount(t, ctx, infra, folderId, templateId, accountName, chHost, - &chHTTPPort, nil) - - proxyPort, pamCmd := startClickHouseProxy(t, ctx, infra, folderName, accountName) - waitForProxyHTTP(t, ctx, pamCmd, proxyPort) - - _, err := queryOverNative(t, ctx, proxyPort, "SELECT toString(count()) AS answer FROM events", false) - require.Error(t, err) - require.Contains(t, err.Error(), "native port", - "a native client must be told why, not left with a transport error") - }) -} From c3d89e19265640f69f082ccc2f27746b30e13e00 Mon Sep 17 00:00:00 2001 From: bernie-g Date: Mon, 28 Sep 2026 10:50:01 -0400 Subject: [PATCH 11/16] fix(clickhouse): bound the whole packet and close the session on both races Per-field caps only bound one field at a time, so thousands of individually legal settings still added up. The query packet now carries one budget across every field it decodes. The refusal check sat outside the lock that orders writes to the client, so a server packet cleared a moment before a refusal could still land after the exception. The check now happens under that lock. The server loop returning left the client blocked on a read until the idle deadline when the upstream went away, waiting for a result that could never arrive. It now ends the session instead. --- .../pam/handlers/clickhouse/bounded_decode.go | 66 +++++++++++++------ .../clickhouse/bounded_decode_test.go | 23 +++++++ packages/pam/handlers/clickhouse/native.go | 39 +++++++---- .../handlers/clickhouse/native_unit_test.go | 37 +++++++++++ 4 files changed, 133 insertions(+), 32 deletions(-) diff --git a/packages/pam/handlers/clickhouse/bounded_decode.go b/packages/pam/handlers/clickhouse/bounded_decode.go index d32994c84..8024bf02d 100644 --- a/packages/pam/handlers/clickhouse/bounded_decode.go +++ b/packages/pam/handlers/clickhouse/bounded_decode.go @@ -20,9 +20,27 @@ const ( maxQueryStringLen = 16 << 20 // A settings or parameters list terminates on an empty key, so it also needs a count bound. maxQuerySettings = 4096 + // Individually legal fields still add up, so the packet carries one budget across all of them. + maxQueryPacketBytes = 32 << 20 ) -func readCappedStr(r *proto.Reader, limit int) (string, error) { +// budget bounds a whole packet, where the per-field caps only bound one field at a time. +type budget struct{ remaining int } + +func newBudget() *budget { return &budget{remaining: maxQueryPacketBytes} } + +func (b *budget) take(n int) error { + if b == nil { + return nil + } + if n > b.remaining { + return fmt.Errorf("packet exceeds its %d byte budget", maxQueryPacketBytes) + } + b.remaining -= n + return nil +} + +func readCappedStr(r *proto.Reader, limit int, b *budget) (string, error) { n, err := r.UVarInt() if err != nil { return "", err @@ -30,6 +48,9 @@ func readCappedStr(r *proto.Reader, limit int) (string, error) { if n > uint64(limit) { return "", fmt.Errorf("declared string of %d bytes exceeds the %d byte cap", n, limit) } + if err := b.take(int(n)); err != nil { + return "", err + } buf := make([]byte, n) if _, err := io.ReadFull(r, buf); err != nil { return "", err @@ -38,7 +59,11 @@ func readCappedStr(r *proto.Reader, limit int) (string, error) { } func readBoundedStr(r *proto.Reader) (string, error) { - return readCappedStr(r, maxHandshakeStringLen) + return readCappedStr(r, maxHandshakeStringLen, nil) +} + +func readBudgetedStr(r *proto.Reader, b *budget) (string, error) { + return readCappedStr(r, maxHandshakeStringLen, b) } func decodeBoundedClientHello(r *proto.Reader) (proto.ClientHello, error) { @@ -69,10 +94,10 @@ func decodeBoundedClientHello(r *proto.Reader) (proto.ClientHello, error) { } // Mirrors proto.Setting.Decode. An empty key terminates the list and leaves the rest unread. -func decodeBoundedSetting(r *proto.Reader) (proto.Setting, error) { +func decodeBoundedSetting(r *proto.Reader, b *budget) (proto.Setting, error) { var s proto.Setting - key, err := readBoundedStr(r) + key, err := readBudgetedStr(r, b) if err != nil { return s, fmt.Errorf("key: %w", err) } @@ -84,7 +109,7 @@ func decodeBoundedSetting(r *proto.Reader) (proto.Setting, error) { if err != nil { return s, fmt.Errorf("flags: %w", err) } - value, err := readCappedStr(r, maxQueryStringLen) + value, err := readCappedStr(r, maxQueryStringLen, b) if err != nil { return s, fmt.Errorf("value (%s): %w", key, err) } @@ -98,7 +123,7 @@ func decodeBoundedSetting(r *proto.Reader) (proto.Setting, error) { } // Mirrors proto.ClientInfo.DecodeAware. -func decodeBoundedClientInfo(r *proto.Reader, version int) (proto.ClientInfo, error) { +func decodeBoundedClientInfo(r *proto.Reader, version int, b *budget) (proto.ClientInfo, error) { var c proto.ClientInfo kind, err := r.UInt8() @@ -110,13 +135,13 @@ func decodeBoundedClientInfo(r *proto.Reader, version int) (proto.ClientInfo, er return c, fmt.Errorf("unknown query kind %d", kind) } - if c.InitialUser, err = readBoundedStr(r); err != nil { + if c.InitialUser, err = readBudgetedStr(r, b); err != nil { return c, fmt.Errorf("initial user: %w", err) } - if c.InitialQueryID, err = readBoundedStr(r); err != nil { + if c.InitialQueryID, err = readBudgetedStr(r, b); err != nil { return c, fmt.Errorf("initial query id: %w", err) } - if c.InitialAddress, err = readBoundedStr(r); err != nil { + if c.InitialAddress, err = readBudgetedStr(r, b); err != nil { return c, fmt.Errorf("initial address: %w", err) } @@ -138,13 +163,13 @@ func decodeBoundedClientInfo(r *proto.Reader, version int) (proto.ClientInfo, er return c, fmt.Errorf("only tcp interface is supported") } - if c.OSUser, err = readBoundedStr(r); err != nil { + if c.OSUser, err = readBudgetedStr(r, b); err != nil { return c, fmt.Errorf("os user: %w", err) } - if c.ClientHostname, err = readBoundedStr(r); err != nil { + if c.ClientHostname, err = readBudgetedStr(r, b); err != nil { return c, fmt.Errorf("client hostname: %w", err) } - if c.ClientName, err = readBoundedStr(r); err != nil { + if c.ClientName, err = readBudgetedStr(r, b); err != nil { return c, fmt.Errorf("client name: %w", err) } if c.Major, err = r.Int(); err != nil { @@ -158,7 +183,7 @@ func decodeBoundedClientInfo(r *proto.Reader, version int) (proto.ClientInfo, er } if proto.FeatureQuotaKeyInClientInfo.In(version) { - if c.QuotaKey, err = readBoundedStr(r); err != nil { + if c.QuotaKey, err = readBudgetedStr(r, b); err != nil { return c, fmt.Errorf("quota key: %w", err) } } @@ -194,7 +219,7 @@ func decodeBoundedClientInfo(r *proto.Reader, version int) (proto.ClientInfo, er bswap.Swap64(raw) copy(cfg.SpanID[:], raw) - state, err := readBoundedStr(r) + state, err := readBudgetedStr(r, b) if err != nil { return c, fmt.Errorf("trace state: %w", err) } @@ -232,15 +257,16 @@ func decodeBoundedClientInfo(r *proto.Reader, version int) (proto.ClientInfo, er // Mirrors proto.Query.DecodeAware. func decodeBoundedQuery(r *proto.Reader, version int) (proto.Query, error) { + b := newBudget() var q proto.Query var err error - if q.ID, err = readBoundedStr(r); err != nil { + if q.ID, err = readBudgetedStr(r, b); err != nil { return q, fmt.Errorf("query id: %w", err) } if proto.FeatureClientWriteInfo.In(version) { - if q.Info, err = decodeBoundedClientInfo(r, version); err != nil { + if q.Info, err = decodeBoundedClientInfo(r, version, b); err != nil { return q, fmt.Errorf("client info: %w", err) } } @@ -250,7 +276,7 @@ func decodeBoundedQuery(r *proto.Reader, version int) (proto.Query, error) { } for { - s, err := decodeBoundedSetting(r) + s, err := decodeBoundedSetting(r, b) if err != nil { return q, fmt.Errorf("setting: %w", err) } @@ -264,7 +290,7 @@ func decodeBoundedQuery(r *proto.Reader, version int) (proto.Query, error) { } if proto.FeatureInterServerSecret.In(version) { - if q.Secret, err = readBoundedStr(r); err != nil { + if q.Secret, err = readBudgetedStr(r, b); err != nil { return q, fmt.Errorf("inter-server secret: %w", err) } } @@ -287,13 +313,13 @@ func decodeBoundedQuery(r *proto.Reader, version int) (proto.Query, error) { return q, fmt.Errorf("unknown compression %d", compression) } - if q.Body, err = readCappedStr(r, maxQueryStringLen); err != nil { + if q.Body, err = readCappedStr(r, maxQueryStringLen, b); err != nil { return q, fmt.Errorf("query body: %w", err) } if proto.FeatureParameters.In(version) { for { - s, err := decodeBoundedSetting(r) + s, err := decodeBoundedSetting(r, b) if err != nil { return q, fmt.Errorf("parameter: %w", err) } diff --git a/packages/pam/handlers/clickhouse/bounded_decode_test.go b/packages/pam/handlers/clickhouse/bounded_decode_test.go index 97ec6bf2a..257aa3d75 100644 --- a/packages/pam/handlers/clickhouse/bounded_decode_test.go +++ b/packages/pam/handlers/clickhouse/bounded_decode_test.go @@ -159,3 +159,26 @@ func TestBoundedQueryRefusesTooManySettings(t *testing.T) { require.Error(t, err) require.Contains(t, err.Error(), "more than") } + +// Each setting is individually legal, so only a whole-packet budget stops thousands of them adding up. +func TestBoundedQueryRefusesAPacketPastItsBudget(t *testing.T) { + q := proto.Query{ID: "id", Body: "SELECT 1", Stage: proto.StageComplete} + q.Info.Query = proto.ClientQueryInitial + q.Info.Interface = proto.InterfaceTCP + q.Info.InitialAddress = "127.0.0.1:0" + q.Info.Major, q.Info.Minor, q.Info.ProtocolVersion = 24, 8, maxNativeRevision + + // Well inside the per-setting cap and the count cap, but past the packet budget in aggregate. + value := strings.Repeat("x", 1<<20) + for i := 0; i < (maxQueryPacketBytes>>20)+2; i++ { + q.Settings = append(q.Settings, proto.Setting{Key: fmt.Sprintf("s%d", i), Value: value}) + } + require.Less(t, len(q.Settings), maxQuerySettings, "must not trip the count cap instead") + + var b proto.Buffer + q.EncodeAware(&b, maxNativeRevision) + + _, err := decodeBoundedQuery(readerSkippingCode(t, b.Buf), maxNativeRevision) + require.Error(t, err) + require.Contains(t, err.Error(), "budget") +} diff --git a/packages/pam/handlers/clickhouse/native.go b/packages/pam/handlers/clickhouse/native.go index 55ffde9f1..ee7a12c4e 100644 --- a/packages/pam/handlers/clickhouse/native.go +++ b/packages/pam/handlers/clickhouse/native.go @@ -108,7 +108,24 @@ func (s *nativeSession) writeToClient(payload []byte) error { } s.writeMu.Lock() defer s.writeMu.Unlock() + return s.writeClientLocked(payload) +} + +// A refusal can land between a caller's own check and its write, so the check belongs under the lock +// that orders the writes, or a packet cleared a moment earlier still trails the exception. +func (s *nativeSession) writeToClientUnlessRefused(payload []byte) error { + if len(payload) == 0 { + return nil + } + s.writeMu.Lock() + defer s.writeMu.Unlock() + if s.refused.Load() { + return errSessionRefused + } + return s.writeClientLocked(payload) +} +func (s *nativeSession) writeClientLocked(payload []byte) error { _ = s.client.SetWriteDeadline(time.Now().Add(nativeWriteTimeout)) defer func() { _ = s.client.SetWriteDeadline(time.Time{}) }() @@ -173,6 +190,13 @@ func (p *nativeProxy) HandleConnection(ctx context.Context, clientConn net.Conn, serverDone := make(chan struct{}) go func() { defer close(serverDone) + // Nothing can answer the client once the upstream is gone, so end the session rather than + // leave the client loop blocked on a read until the idle deadline. + defer func() { + if !s.refused.Load() { + _ = clientConn.Close() + } + }() defer func() { if r := recover(); r != nil { l.Error().Interface("panic", r).Msg("Recovered from a panic reading the ClickHouse server direction") @@ -434,10 +458,7 @@ func (s *nativeSession) serverLoop() { relayRest := func(reason string) { s.outcomes.degrade(reason) - if s.refused.Load() { - return - } - if err := s.writeToClient(t.take()); err != nil { + if err := s.writeToClientUnlessRefused(t.take()); err != nil { return } _, _ = io.Copy(newRefusalAwareWriter(s), t.rest()) @@ -504,10 +525,7 @@ func (s *nativeSession) serverLoop() { return } - if s.refused.Load() { - return - } - if err := s.writeToClient(t.take()); err != nil { + if err := s.writeToClientUnlessRefused(t.take()); err != nil { return } } @@ -519,10 +537,7 @@ type refusalAwareWriter struct{ s *nativeSession } func newRefusalAwareWriter(s *nativeSession) io.Writer { return refusalAwareWriter{s: s} } func (w refusalAwareWriter) Write(p []byte) (int, error) { - if w.s.refused.Load() { - return 0, errSessionRefused - } - if err := w.s.writeToClient(p); err != nil { + if err := w.s.writeToClientUnlessRefused(p); err != nil { return 0, err } return len(p), nil diff --git a/packages/pam/handlers/clickhouse/native_unit_test.go b/packages/pam/handlers/clickhouse/native_unit_test.go index 9d8e24002..3a9d0e8fb 100644 --- a/packages/pam/handlers/clickhouse/native_unit_test.go +++ b/packages/pam/handlers/clickhouse/native_unit_test.go @@ -35,6 +35,17 @@ type fakeClickHouse struct { serverRevision int // Non-empty answers the handshake with an exception instead of a hello. refuseWith string + // The accepted connection, so a test can drop the upstream mid-session. + conn net.Conn +} + +func (f *fakeClickHouse) disconnect() { + f.mu.Lock() + conn := f.conn + f.mu.Unlock() + if conn != nil { + _ = conn.Close() + } } func startFakeClickHouse(t *testing.T, serverRevision ...int) *fakeClickHouse { @@ -56,6 +67,9 @@ func startFakeClickHouse(t *testing.T, serverRevision ...int) *fakeClickHouse { return } defer conn.Close() + f.mu.Lock() + f.conn = conn + f.mu.Unlock() f.serve(conn) }() @@ -658,3 +672,26 @@ func TestNativeRefusesAnOversizedQueryBody(t *testing.T) { _, _, queries, _ := upstream.snapshot() require.Empty(t, queries, "nothing may be forwarded from a packet the gateway refused") } + +// Without this the client waits on a read until the idle deadline for a result that cannot arrive. +func TestUpstreamDisconnectEndsTheClientSession(t *testing.T) { + upstream := startFakeClickHouse(t) + + conn := dialProxy(t, ClickHouseProxyConfig{ + NativeAddr: upstream.addr(), + Username: "account", + SessionID: "unit", + SessionLogger: &recordingLogger{}, + }) + clientHandshake(t, conn, "someone", "whatever") + + upstream.disconnect() + + require.NoError(t, conn.SetReadDeadline(time.Now().Add(10*time.Second))) + buf := make([]byte, 16) + started := time.Now() + _, err := conn.Read(buf) + require.Error(t, err, "the client must not be left waiting once the upstream is gone") + // The deadline would also produce an error, so the point is that it ends well before one. + require.Less(t, time.Since(started), 3*time.Second, "the session should end promptly, not on a timeout") +} From 4e3c369db24764aeca2480c4734d3e1dd53fb079 Mon Sep 17 00:00:00 2001 From: bernie-g Date: Mon, 28 Sep 2026 12:20:05 -0400 Subject: [PATCH 12/16] fix(clickhouse): bound the data block and stop losing outcomes on a panic ch-go sizes a column from the declared row count before reading any of it, so a ~30 byte data block header committed 763 MB for Int64 and 3.2 GB for Int256 while the client sent nothing further. That allocation succeeds rather than panicking, so the handler's recover could not catch it. The block header is now scanned before the decoder sees it, without consuming it, so an absurd row or column count is refused. The scan is best-effort by design: anything it cannot parse falls through to ch-go, so a mistake in it can only miss an attack, never reject real traffic. It peeks only what has already arrived, since a fixed window would stall a session whose next block is smaller than it. Compressed blocks are left alone, being already capped by ch-go. Session teardown was straight-line after the client loop, so a panic there skipped it and every in-flight statement vanished from the session log without even an INTERRUPTED. It is deferred now. A near-exhausted probe budget rendered as "within 0s" and told the operator their port was misconfigured, which is a wrong diagnosis for a timeout the test itself caused. It now says the budget ran out, and the test that locked in the one duration that rendered correctly asserts the meaning instead. The upstream-disconnect test set its read deadline after triggering the teardown it was testing, so it failed about 1% of runs on the setup line rather than the assertion. 300 runs clean. --- .../pam/handlers/clickhouse/bounded_block.go | 126 ++++++++++++++++++ .../handlers/clickhouse/bounded_block_test.go | 103 ++++++++++++++ packages/pam/handlers/clickhouse/native.go | 43 ++++-- .../handlers/clickhouse/native_unit_test.go | 33 ++++- 4 files changed, 292 insertions(+), 13 deletions(-) create mode 100644 packages/pam/handlers/clickhouse/bounded_block.go create mode 100644 packages/pam/handlers/clickhouse/bounded_block_test.go diff --git a/packages/pam/handlers/clickhouse/bounded_block.go b/packages/pam/handlers/clickhouse/bounded_block.go new file mode 100644 index 000000000..a8dbffaf8 --- /dev/null +++ b/packages/pam/handlers/clickhouse/bounded_block.go @@ -0,0 +1,126 @@ +package clickhouse + +import ( + "bufio" + "encoding/binary" + "fmt" +) + +// ch-go sizes a column from the declared row count before it reads any of it: ColInt64.DecodeColumn does +// make([]int64, rows) and only then blocks in ReadFull. Its own ceiling is 100M rows, so a ~30 byte header +// commits 763 MB for Int64 and 3.2 GB for Int256 while the client sends nothing further. That allocation +// succeeds, so unlike an oversized string it is not a panic the handler can recover from. +const ( + maxBlockRows = 4 << 20 + maxBlockColumns = 4096 + // BlockInfo plus two varints; anything past this is the decoder's problem. + blockHeaderWindow = 64 +) + +// blockInfo field ids, from ch-go's proto/block.go. +const ( + blockInfoOverflows = 1 + blockInfoBucketNum = 2 + blockInfoEnd = 0 +) + +// checkBlockHeader inspects the block header without consuming it and reports a reason to refuse. The scan +// is best-effort on purpose: anything it cannot parse returns "", leaving the real decoder to judge, so a +// mistake here can only miss an attack and never reject legitimate traffic. +func checkBlockHeader(src *bufio.Reader, revision int, compressed bool) string { + // Compressed blocks arrive as LZ4 frames, which these raw bytes are not, and ch-go already caps a + // compressed frame at maxDataSize. + if compressed { + return "" + } + + // Peek only what has already arrived. A fixed size would block until that many bytes exist, which + // stalls a session whose next block is smaller than the window. + if _, err := src.Peek(1); err != nil { + return "" + } + window := src.Buffered() + if window > blockHeaderWindow { + window = blockHeaderWindow + } + head, err := src.Peek(window) + if err != nil && len(head) == 0 { + return "" + } + + p := &peeker{buf: head} + if featureBlockInfo(revision) && !p.skipBlockInfo() { + return "" + } + + columns, ok := p.uvarint() + if !ok { + return "" + } + rows, ok := p.uvarint() + if !ok { + return "" + } + + if columns > maxBlockColumns { + return fmt.Sprintf("the data block declares %d columns, more than the %d this session accepts", + columns, maxBlockColumns) + } + if rows > maxBlockRows { + return fmt.Sprintf("the data block declares %d rows, more than the %d this session accepts. "+ + "Send the data in smaller batches", rows, maxBlockRows) + } + return "" +} + +// FeatureBlockInfo is 51903 in ch-go; every revision this proxy speaks is above it, but keep the gate +// explicit so the scan stays aligned with DecodeBlock. +func featureBlockInfo(revision int) bool { return revision >= 51903 } + +type peeker struct { + buf []byte + pos int +} + +func (p *peeker) byteAt() (byte, bool) { + if p.pos >= len(p.buf) { + return 0, false + } + b := p.buf[p.pos] + p.pos++ + return b, true +} + +func (p *peeker) uvarint() (uint64, bool) { + v, n := binary.Uvarint(p.buf[p.pos:]) + if n <= 0 { + return 0, false + } + p.pos += n + return v, true +} + +// Mirrors BlockInfo.Decode: field-id/value pairs terminated by field 0. +func (p *peeker) skipBlockInfo() bool { + for { + field, ok := p.uvarint() + if !ok { + return false + } + switch field { + case blockInfoEnd: + return true + case blockInfoOverflows: + if _, ok := p.byteAt(); !ok { + return false + } + case blockInfoBucketNum: + if p.pos+4 > len(p.buf) { + return false + } + p.pos += 4 + default: + return false + } + } +} diff --git a/packages/pam/handlers/clickhouse/bounded_block_test.go b/packages/pam/handlers/clickhouse/bounded_block_test.go new file mode 100644 index 000000000..46c63149b --- /dev/null +++ b/packages/pam/handlers/clickhouse/bounded_block_test.go @@ -0,0 +1,103 @@ +package clickhouse + +import ( + "bufio" + "bytes" + "testing" + + "github.com/ClickHouse/ch-go/proto" + "github.com/stretchr/testify/require" +) + +// blockHeader builds the wire prefix of a Data block body: BlockInfo, then columns and rows. +func blockHeader(columns, rows uint64) []byte { + var b proto.Buffer + b.PutUVarInt(blockInfoOverflows) + b.PutBool(false) + b.PutUVarInt(blockInfoBucketNum) + b.PutInt32(-1) + b.PutUVarInt(blockInfoEnd) + b.PutUVarInt(columns) + b.PutUVarInt(rows) + return b.Buf +} + +func scan(payload []byte, compressed bool) string { + return checkBlockHeader(bufio.NewReaderSize(bytes.NewReader(payload), 64<<10), maxNativeRevision, compressed) +} + +func TestBlockHeaderScanRefusesAnOversizedBlock(t *testing.T) { + // ch-go allocates rows x width before reading, so this is committed memory, not a recoverable panic. + reason := scan(blockHeader(1, 100_000_000), false) + require.Contains(t, reason, "rows") + require.Contains(t, reason, "smaller batches") + + require.Contains(t, scan(blockHeader(1_000_000, 1), false), "columns") +} + +func TestBlockHeaderScanPassesALegitimateBlock(t *testing.T) { + for _, c := range []struct { + name string + columns, rows uint64 + }{ + {"an empty block", 0, 0}, + {"a single row", 1, 1}, + {"a default max_insert_block_size batch", 8, 1_048_545}, + {"exactly at the row cap", 1, maxBlockRows}, + {"exactly at the column cap", maxBlockColumns, 1}, + } { + t.Run(c.name, func(t *testing.T) { + require.Empty(t, scan(blockHeader(c.columns, c.rows), false)) + }) + } +} + +// The scan must never be the thing that rejects traffic: anything it cannot parse is left to the decoder. +func TestBlockHeaderScanFallsThroughWhenItCannotParse(t *testing.T) { + require.Empty(t, scan(nil, false), "empty input") + require.Empty(t, scan([]byte{0xFF}, false), "a truncated varint") + require.Empty(t, scan([]byte{0x09, 0x01}, false), "an unknown BlockInfo field") + require.Empty(t, scan(blockHeader(1, 100_000_000), true), "a compressed block is ch-go's to bound") +} + +// The bytes the scan inspects must still reach the decoder, or the packet would be truncated. +func TestBlockHeaderScanConsumesNothing(t *testing.T) { + payload := append(blockHeader(2, 7), []byte("trailing")...) + src := bufio.NewReaderSize(bytes.NewReader(payload), 64<<10) + + require.Empty(t, checkBlockHeader(src, maxNativeRevision, false)) + + got := make([]byte, len(payload)) + _, err := src.Read(got[:1]) + require.NoError(t, err) + require.Equal(t, payload[0], got[0], "the scan must not consume the header") + require.Equal(t, len(payload)-1, src.Buffered()+0, "everything after the first byte is still pending") +} + +// The whole point: a client must not be able to make the gateway size a column from a declared row count. +func TestNativeRefusesAnOversizedDataBlock(t *testing.T) { + upstream := startFakeClickHouse(t) + + conn := dialProxy(t, ClickHouseProxyConfig{ + NativeAddr: upstream.addr(), + Username: "account", + SessionID: "unit", + SessionLogger: &recordingLogger{}, + }) + r := clientHandshake(t, conn, "someone", "whatever") + + var b proto.Buffer + proto.ClientCodeData.Encode(&b) + b.PutString("") + b.Buf = append(b.Buf, blockHeader(1, 100_000_000)...) + _, err := conn.Write(b.Buf) + require.NoError(t, err) + + code, message := decodeException(t, r) + require.Equal(t, codeNotImplemented, code) + require.Contains(t, message, "rows") + + _, _, queries, bytesAfter := upstream.snapshot() + require.Empty(t, queries) + require.Zero(t, bytesAfter, "a refused block must not be relayed upstream") +} diff --git a/packages/pam/handlers/clickhouse/native.go b/packages/pam/handlers/clickhouse/native.go index ee7a12c4e..d470e7f09 100644 --- a/packages/pam/handlers/clickhouse/native.go +++ b/packages/pam/handlers/clickhouse/native.go @@ -78,6 +78,9 @@ func (t *tap) discard() { t.buf = nil } +// peeker exposes the buffered source so a header can be inspected without consuming it. +func (t *tap) peeker() *bufio.Reader { return t.src } + // A refusal ends the session: the stream is mid-packet, so carrying on would let a later packet flush the // refused bytes upstream. var errSessionRefused = errors.New("the session was refused") @@ -205,18 +208,22 @@ func (p *nativeProxy) HandleConnection(ctx context.Context, clientConn net.Conn, s.serverLoop() }() + // Deferred so a panic in the client loop still drains the recorder. Straight-line teardown would + // skip it and lose every in-flight statement from the session log. + defer func() { + // Outcomes must land before the recorder drains, or a finished statement is recorded as interrupted. + upstream.Close() + select { + case <-serverDone: + case <-time.After(nativeWriteTimeout): + l.Warn().Msg("The ClickHouse server direction did not stop in time") + } + s.outcomes.finish() + }() + if err := s.clientLoop(clientTap, clientReader); err != nil { l.Debug().Err(err).Msg("ClickHouse native session ended") } - - // Outcomes must land before the recorder drains, or a finished statement is recorded as interrupted. - upstream.Close() - select { - case <-serverDone: - case <-time.After(nativeWriteTimeout): - l.Warn().Msg("The ClickHouse server direction did not stop in time") - } - s.outcomes.finish() return nil } @@ -429,6 +436,12 @@ func (s *nativeSession) handleData(t *tap, r *proto.Reader) error { } compressed := s.compressed.Load() + + if reason := checkBlockHeader(t.peeker(), s.rev, compressed); reason != "" { + s.log.Warn().Str("table", table).Msg("Refused an oversized ClickHouse data block") + return s.refuse(t, codeNotImplemented, reason) + } + if compressed { r.EnableCompression() } @@ -620,11 +633,17 @@ func TestNativeConnection(ctx context.Context, config ClickHouseProxyConfig) err // The probe's own budget wins when it is shorter, so a slow handshake cannot outlive the test. budget := nativeHandshakeTimeout + budgetWasCapped := false if probeDeadline, ok := ctx.Deadline(); ok { if remaining := time.Until(probeDeadline); remaining < budget { - budget = remaining + budget, budgetWasCapped = remaining, true } } + // Too little left to tell a silent port from a probe that simply ran out of time. + if budget < time.Second { + return fmt.Errorf("the connection test ran out of time before ClickHouse's native port could be "+ + "checked; %s remained of the budget", budget.Round(time.Millisecond)) + } _ = conn.SetDeadline(time.Now().Add(budget)) var b proto.Buffer @@ -645,6 +664,10 @@ func TestNativeConnection(ctx context.Context, config ClickHouseProxyConfig) err code, err := r.UVarInt() if err != nil { if errors.Is(err, os.ErrDeadlineExceeded) { + if budgetWasCapped { + return fmt.Errorf("the port accepted the connection but did not answer ClickHouse's native "+ + "handshake in the %s left of the connection test's budget: %w", budget.Round(time.Millisecond), err) + } return fmt.Errorf("the port accepted the connection but did not answer ClickHouse's native "+ "handshake within %s, which is what the HTTP port does when it is entered as the native one: %w", budget.Round(time.Second), err) diff --git a/packages/pam/handlers/clickhouse/native_unit_test.go b/packages/pam/handlers/clickhouse/native_unit_test.go index 3a9d0e8fb..3cdcce582 100644 --- a/packages/pam/handlers/clickhouse/native_unit_test.go +++ b/packages/pam/handlers/clickhouse/native_unit_test.go @@ -562,7 +562,7 @@ func TestNativeConnectionTestClassifiesFailures(t *testing.T) { t.Cleanup(server.Close) // The probe's own deadline bounds the handshake, so this does not wait out nativeHandshakeTimeout. - ctx, cancel := context.WithTimeout(context.Background(), time.Second) + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() err := TestNativeConnection(ctx, ClickHouseProxyConfig{ @@ -572,7 +572,9 @@ func TestNativeConnectionTestClassifiesFailures(t *testing.T) { }) require.Error(t, err) require.Contains(t, err.Error(), "did not answer ClickHouse's native handshake") - require.Contains(t, err.Error(), "within 1s", "the message must name the budget that was applied") + // Capped by the probe, so it must name the remaining budget rather than blame the port. + require.Contains(t, err.Error(), "left of the connection test's budget") + require.NotContains(t, err.Error(), "entered as the native one") // The heartbeat stops scheduling on a rejected credential, so a silent port has to stay a // transport failure rather than being read as one. require.ErrorIs(t, err, os.ErrDeadlineExceeded) @@ -685,9 +687,12 @@ func TestUpstreamDisconnectEndsTheClientSession(t *testing.T) { }) clientHandshake(t, conn, "someone", "whatever") + // Set before the disconnect: the proxy may close this end first, which is the very teardown under + // test, and setting a deadline on a closed pipe errors. + require.NoError(t, conn.SetReadDeadline(time.Now().Add(10*time.Second))) + upstream.disconnect() - require.NoError(t, conn.SetReadDeadline(time.Now().Add(10*time.Second))) buf := make([]byte, 16) started := time.Now() _, err := conn.Read(buf) @@ -695,3 +700,25 @@ func TestUpstreamDisconnectEndsTheClientSession(t *testing.T) { // The deadline would also produce an error, so the point is that it ends well before one. require.Less(t, time.Since(started), 3*time.Second, "the session should end promptly, not on a timeout") } + +// A near-exhausted budget is not evidence that the port is misconfigured, so the message must not say so. +func TestNativeConnectionTestDoesNotBlameThePortWhenTheBudgetRanOut(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {})) + t.Cleanup(server.Close) + + for _, budget := range []time.Duration{20 * time.Millisecond, 100 * time.Millisecond, 400 * time.Millisecond} { + ctx, cancel := context.WithTimeout(context.Background(), budget) + err := TestNativeConnection(ctx, ClickHouseProxyConfig{ + NativeAddr: strings.TrimPrefix(server.URL, "http://"), + Username: "account", + Database: "analytics", + }) + cancel() + + require.Error(t, err, "budget=%s", budget) + require.NotContains(t, err.Error(), "within 0s", "budget=%s: a rounded-to-zero duration is nonsense", budget) + require.NotContains(t, err.Error(), "entered as the native one", + "budget=%s: an exhausted budget must not be reported as a misconfigured port", budget) + require.Contains(t, err.Error(), "ran out of time", "budget=%s", budget) + } +} From 851b6e9d83c75188c1510e101361de57feb92e49 Mon Sep 17 00:00:00 2001 From: bernie-g Date: Mon, 28 Sep 2026 12:32:49 -0400 Subject: [PATCH 13/16] chore(clickhouse): trim comments to one line --- .../pam/handlers/clickhouse/bounded_block.go | 17 +++--------- .../pam/handlers/clickhouse/bounded_decode.go | 5 +--- .../clickhouse/bounded_decode_test.go | 3 +-- packages/pam/handlers/clickhouse/native.go | 27 ++++++------------- .../pam/handlers/clickhouse/native_outcome.go | 2 -- .../handlers/clickhouse/native_unit_test.go | 11 +------- packages/pam/handlers/clickhouse/proxy.go | 5 +--- .../pam/handlers/clickhouse/proxy_test.go | 2 -- packages/pam/handlers/clickhouse/sniff.go | 3 +-- 9 files changed, 17 insertions(+), 58 deletions(-) diff --git a/packages/pam/handlers/clickhouse/bounded_block.go b/packages/pam/handlers/clickhouse/bounded_block.go index a8dbffaf8..518d67377 100644 --- a/packages/pam/handlers/clickhouse/bounded_block.go +++ b/packages/pam/handlers/clickhouse/bounded_block.go @@ -6,10 +6,7 @@ import ( "fmt" ) -// ch-go sizes a column from the declared row count before it reads any of it: ColInt64.DecodeColumn does -// make([]int64, rows) and only then blocks in ReadFull. Its own ceiling is 100M rows, so a ~30 byte header -// commits 763 MB for Int64 and 3.2 GB for Int256 while the client sends nothing further. That allocation -// succeeds, so unlike an oversized string it is not a panic the handler can recover from. +// ch-go allocates declared rows before reading, which recover cannot catch. const ( maxBlockRows = 4 << 20 maxBlockColumns = 4096 @@ -24,18 +21,14 @@ const ( blockInfoEnd = 0 ) -// checkBlockHeader inspects the block header without consuming it and reports a reason to refuse. The scan -// is best-effort on purpose: anything it cannot parse returns "", leaving the real decoder to judge, so a -// mistake here can only miss an attack and never reject legitimate traffic. +// Best-effort: anything unparseable falls through to ch-go rather than being refused. func checkBlockHeader(src *bufio.Reader, revision int, compressed bool) string { - // Compressed blocks arrive as LZ4 frames, which these raw bytes are not, and ch-go already caps a - // compressed frame at maxDataSize. + // Compressed frames are already capped by ch-go. if compressed { return "" } - // Peek only what has already arrived. A fixed size would block until that many bytes exist, which - // stalls a session whose next block is smaller than the window. + // Peek only what has arrived; a fixed window would stall on a small block. if _, err := src.Peek(1); err != nil { return "" } @@ -73,8 +66,6 @@ func checkBlockHeader(src *bufio.Reader, revision int, compressed bool) string { return "" } -// FeatureBlockInfo is 51903 in ch-go; every revision this proxy speaks is above it, but keep the gate -// explicit so the scan stays aligned with DecodeBlock. func featureBlockInfo(revision int) bool { return revision >= 51903 } type peeker struct { diff --git a/packages/pam/handlers/clickhouse/bounded_decode.go b/packages/pam/handlers/clickhouse/bounded_decode.go index 8024bf02d..a7ae1c62e 100644 --- a/packages/pam/handlers/clickhouse/bounded_decode.go +++ b/packages/pam/handlers/clickhouse/bounded_decode.go @@ -9,10 +9,7 @@ import ( "go.opentelemetry.io/otel/trace" ) -// ch-go allocates a declared string length before it reads a single byte, and rejects only a length that -// goes negative. A client that names a terabyte therefore kills the process outright: the allocation is a -// fatal runtime error rather than a panic anything can recover. These decoders mirror ch-go's field for -// field and differ only in reading every string through a cap. +// Mirrors ch-go's decoders, but every string is read through a cap. const ( // Short identifiers: names, users, hostnames, the quota key. maxHandshakeStringLen = 64 << 10 diff --git a/packages/pam/handlers/clickhouse/bounded_decode_test.go b/packages/pam/handlers/clickhouse/bounded_decode_test.go index 257aa3d75..c460c41fa 100644 --- a/packages/pam/handlers/clickhouse/bounded_decode_test.go +++ b/packages/pam/handlers/clickhouse/bounded_decode_test.go @@ -15,8 +15,7 @@ func readerOver(b []byte) *proto.Reader { return proto.NewReader(bytes.NewReader(b)) } -// A field read in the wrong order desynchronises the stream, which is the failure this decoder exists to -// prevent. Encoding with ch-go and decoding with ours is what pins the two together. +// Round-trips against ch-go so a misordered field fails. func TestBoundedQueryDecodeMatchesChGo(t *testing.T) { span := trace.NewSpanContext(trace.SpanContextConfig{ TraceID: trace.TraceID{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16}, diff --git a/packages/pam/handlers/clickhouse/native.go b/packages/pam/handlers/clickhouse/native.go index d470e7f09..c144466d4 100644 --- a/packages/pam/handlers/clickhouse/native.go +++ b/packages/pam/handlers/clickhouse/native.go @@ -18,9 +18,6 @@ import ( "github.com/rs/zerolog" ) -// Only the client direction gates anything. The server direction is read solely to pair an outcome with a -// statement, and gives that up rather than fail a SELECT on a column type ch-go cannot infer. - const ( // A newer server sends Hello fields ch-go cannot read, so both sides are pinned to what we can parse. maxNativeRevision = proto.Version @@ -39,8 +36,7 @@ func newNativeProxy(owner *ClickHouseProxy) *nativeProxy { return &nativeProxy{ClickHouseProxy: owner} } -// tap records every byte a decoder consumes so the exact wire bytes can be replayed upstream. One byte at a -// time, because proto.Reader buffers 128 KB and would swallow packets we have not parsed yet. +// One byte at a time, or proto.Reader buffers past the packet. type tap struct { src *bufio.Reader buf []byte @@ -81,8 +77,7 @@ func (t *tap) discard() { // peeker exposes the buffered source so a header can be inspected without consuming it. func (t *tap) peeker() *bufio.Reader { return t.src } -// A refusal ends the session: the stream is mid-packet, so carrying on would let a later packet flush the -// refused bytes upstream. +// A refusal ends the session: the stream is mid-packet. var errSessionRefused = errors.New("the session was refused") type nativeSession struct { @@ -114,8 +109,7 @@ func (s *nativeSession) writeToClient(payload []byte) error { return s.writeClientLocked(payload) } -// A refusal can land between a caller's own check and its write, so the check belongs under the lock -// that orders the writes, or a packet cleared a moment earlier still trails the exception. +// Checked under the lock, or a packet cleared just before a refusal trails it. func (s *nativeSession) writeToClientUnlessRefused(payload []byte) error { if len(payload) == 0 { return nil @@ -193,8 +187,7 @@ func (p *nativeProxy) HandleConnection(ctx context.Context, clientConn net.Conn, serverDone := make(chan struct{}) go func() { defer close(serverDone) - // Nothing can answer the client once the upstream is gone, so end the session rather than - // leave the client loop blocked on a read until the idle deadline. + // Upstream gone: end the session rather than block until the idle deadline. defer func() { if !s.refused.Load() { _ = clientConn.Close() @@ -208,8 +201,7 @@ func (p *nativeProxy) HandleConnection(ctx context.Context, clientConn net.Conn, s.serverLoop() }() - // Deferred so a panic in the client loop still drains the recorder. Straight-line teardown would - // skip it and lose every in-flight statement from the session log. + // Deferred so a panic still drains the recorder. defer func() { // Outcomes must land before the recorder drains, or a finished statement is recorded as interrupted. upstream.Close() @@ -308,8 +300,7 @@ func (s *nativeSession) handshake(t *tap, r *proto.Reader) error { if err := serverHello.DecodeAware(serverReader, s.rev); err != nil { return fmt.Errorf("decode server hello: %w", err) } - // The upstream can be older than the revision pinned from the client, and anything above what it - // speaks puts feature-gated bytes on the wire it never reads, desynchronising the stream. + // Clamp to the upstream's revision so it never gets fields it can't read. if serverHello.Revision > 0 && serverHello.Revision < s.rev { s.rev = serverHello.Revision } @@ -414,8 +405,7 @@ func (s *nativeSession) handleQuery(t *tap, r *proto.Reader) error { s.outcomes.begin(statement) - // The identities a client could otherwise pick for itself. InitialAddress is left alone: ClickHouse - // asserts on an empty one, and forcing the kind to Initial already authorises as the account. + // InitialAddress stays set: ClickHouse asserts on an empty one. q.Info.QuotaKey = "" q.Info.Query = proto.ClientQueryInitial q.Secret = "" @@ -425,8 +415,7 @@ func (s *nativeSession) handleQuery(t *tap, r *proto.Reader) error { return s.forward(b.Buf) } -// Decodes a block only far enough to find its end, then replays the client's bytes: re-encoding would mean -// reproducing a serialization we do not own. +// Replays the client's bytes rather than re-encoding a format we don't own. func (s *nativeSession) handleData(t *tap, r *proto.Reader) error { table, err := readBoundedStr(r) if err != nil { diff --git a/packages/pam/handlers/clickhouse/native_outcome.go b/packages/pam/handlers/clickhouse/native_outcome.go index de81bd531..ea3dd8365 100644 --- a/packages/pam/handlers/clickhouse/native_outcome.go +++ b/packages/pam/handlers/clickhouse/native_outcome.go @@ -7,8 +7,6 @@ import ( "time" ) -// Pairs a statement with how it ended. The two directions are separate goroutines and ClickHouse answers in -// order, so the queue is what joins them back up. Best effort: a block it cannot decode costs only the outcome. type outcomeRecorder struct { proxy *ClickHouseProxy diff --git a/packages/pam/handlers/clickhouse/native_unit_test.go b/packages/pam/handlers/clickhouse/native_unit_test.go index 3cdcce582..a30fc3cdc 100644 --- a/packages/pam/handlers/clickhouse/native_unit_test.go +++ b/packages/pam/handlers/clickhouse/native_unit_test.go @@ -463,8 +463,6 @@ func TestNativeHandshakeRefusesAnOversizedField(t *testing.T) { require.Zero(t, bytesAfterHandshake) } -// The server direction must stop writing once a statement has been refused, or the client sees bytes -// trailing the exception the proxy just sent it. func newRefusedSession(t *testing.T, client net.Conn, upstream net.Conn) *nativeSession { t.Helper() @@ -556,8 +554,6 @@ func TestRefusalAwareWriterRefusesAfterARefusal(t *testing.T) { func TestNativeConnectionTestClassifiesFailures(t *testing.T) { t.Run("an http port is named as a handshake timeout, not a rejected credential", func(t *testing.T) { - // An HTTP server accepts the connection and then waits for a request, which is exactly what a - // misconfigured native port looks like. server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {})) t.Cleanup(server.Close) @@ -575,8 +571,6 @@ func TestNativeConnectionTestClassifiesFailures(t *testing.T) { // Capped by the probe, so it must name the remaining budget rather than blame the port. require.Contains(t, err.Error(), "left of the connection test's budget") require.NotContains(t, err.Error(), "entered as the native one") - // The heartbeat stops scheduling on a rejected credential, so a silent port has to stay a - // transport failure rather than being read as one. require.ErrorIs(t, err, os.ErrDeadlineExceeded) }) @@ -635,8 +629,6 @@ func TestNativeAnchoredRuleStillBlocksAStatementCarryingParameters(t *testing.T) require.Empty(t, queries, "the blocked statement must not reach the upstream") } -// An absurd declared length is a fatal allocation inside ch-go, not a panic anything can recover, so the -// session has to refuse it before the decoder ever sees it. func TestNativeRefusesAnOversizedQueryBody(t *testing.T) { upstream := startFakeClickHouse(t) @@ -687,8 +679,7 @@ func TestUpstreamDisconnectEndsTheClientSession(t *testing.T) { }) clientHandshake(t, conn, "someone", "whatever") - // Set before the disconnect: the proxy may close this end first, which is the very teardown under - // test, and setting a deadline on a closed pipe errors. + // Set before disconnect: the proxy may close this end first. require.NoError(t, conn.SetReadDeadline(time.Now().Add(10*time.Second))) upstream.disconnect() diff --git a/packages/pam/handlers/clickhouse/proxy.go b/packages/pam/handlers/clickhouse/proxy.go index 0f8b11b86..41d8a62b3 100644 --- a/packages/pam/handlers/clickhouse/proxy.go +++ b/packages/pam/handlers/clickhouse/proxy.go @@ -210,8 +210,6 @@ func (p *ClickHouseProxy) handler(l zerolog.Logger) http.Handler { return } - // Without an HTTP upstream the reverse proxy would fail on an empty host, which reads as a network - // fault rather than an account that does not serve this protocol. if p.config.TargetAddr == "" { l.Info().Msg("Refused an HTTP connection on an account with no HTTP port") writeClickHouseError(w, http.StatusBadGateway, codeNotImplemented, @@ -469,8 +467,7 @@ func (p *ClickHouseProxy) handleUpstreamError(w http.ResponseWriter, r *http.Req fmt.Sprintf("The gateway could not reach ClickHouse: %v", err)) } -// The recorded form carries a "-- parameters:" suffix, so an end-anchored rule stops matching the moment -// a client attaches one. Both the executable SQL and the recorded form are checked. +// Also checks the bare SQL: the recorded suffix defeats end-anchored rules. func (p *ClickHouseProxy) blockedBy(statements ...string) *regexp.Regexp { for _, pattern := range p.config.BlockedCommands { for _, statement := range statements { diff --git a/packages/pam/handlers/clickhouse/proxy_test.go b/packages/pam/handlers/clickhouse/proxy_test.go index 28cc0dbab..4019c62db 100644 --- a/packages/pam/handlers/clickhouse/proxy_test.go +++ b/packages/pam/handlers/clickhouse/proxy_test.go @@ -491,8 +491,6 @@ func TestRefusesADeflatedBodyItCannotDecode(t *testing.T) { require.Equal(t, http.StatusBadRequest, recorder.Code, recorder.Body.String()) } -// The recorded form carries a parameter suffix, so an end-anchored rule would stop matching as soon as -// a client attached a parameter and the blocked statement would run. func TestAnAnchoredRuleStillBlocksAStatementCarryingParameters(t *testing.T) { reached := false handler, _, closeUpstream := newTestProxy(t, func(w http.ResponseWriter, r *http.Request) { diff --git a/packages/pam/handlers/clickhouse/sniff.go b/packages/pam/handlers/clickhouse/sniff.go index 827d7de9c..f1a790ed5 100644 --- a/packages/pam/handlers/clickhouse/sniff.go +++ b/packages/pam/handlers/clickhouse/sniff.go @@ -6,8 +6,7 @@ import ( "time" ) -// Which protocol a client speaks is decided by its driver, not the user, so one port serves both. A native -// session opens with the Hello code, a uvarint 0; every HTTP request opens with an ASCII method letter. +// Native opens with a uvarint 0; HTTP with an ASCII method letter. const nativeHelloByte = 0x00 // The peek happens before any HTTP server exists, so ReadHeaderTimeout does not cover it. From fd88e907683c656397ce8ea114e2c7759ee324af Mon Sep 17 00:00:00 2001 From: bernie-g Date: Mon, 28 Sep 2026 12:37:40 -0400 Subject: [PATCH 14/16] fix(clickhouse): enforce data block limits inside the decoder The header pre-scan could be skipped by a split header or a compressed block. The limits now run in the decoder's result hook, which sees the parsed row and column counts before any column is allocated. --- .../pam/handlers/clickhouse/bounded_block.go | 114 ++---------------- .../handlers/clickhouse/bounded_block_test.go | 106 ++++------------ packages/pam/handlers/clickhouse/native.go | 10 +- 3 files changed, 35 insertions(+), 195 deletions(-) diff --git a/packages/pam/handlers/clickhouse/bounded_block.go b/packages/pam/handlers/clickhouse/bounded_block.go index 518d67377..b04dbfc9b 100644 --- a/packages/pam/handlers/clickhouse/bounded_block.go +++ b/packages/pam/handlers/clickhouse/bounded_block.go @@ -1,117 +1,27 @@ package clickhouse import ( - "bufio" - "encoding/binary" "fmt" + + "github.com/ClickHouse/ch-go/proto" ) -// ch-go allocates declared rows before reading, which recover cannot catch. const ( maxBlockRows = 4 << 20 maxBlockColumns = 4096 - // BlockInfo plus two varints; anything past this is the decoder's problem. - blockHeaderWindow = 64 -) - -// blockInfo field ids, from ch-go's proto/block.go. -const ( - blockInfoOverflows = 1 - blockInfoBucketNum = 2 - blockInfoEnd = 0 ) -// Best-effort: anything unparseable falls through to ch-go rather than being refused. -func checkBlockHeader(src *bufio.Reader, revision int, compressed bool) string { - // Compressed frames are already capped by ch-go. - if compressed { - return "" - } - - // Peek only what has arrived; a fixed window would stall on a small block. - if _, err := src.Peek(1); err != nil { - return "" - } - window := src.Buffered() - if window > blockHeaderWindow { - window = blockHeaderWindow - } - head, err := src.Peek(window) - if err != nil && len(head) == 0 { - return "" - } - - p := &peeker{buf: head} - if featureBlockInfo(revision) && !p.skipBlockInfo() { - return "" - } - - columns, ok := p.uvarint() - if !ok { - return "" - } - rows, ok := p.uvarint() - if !ok { - return "" - } - - if columns > maxBlockColumns { - return fmt.Sprintf("the data block declares %d columns, more than the %d this session accepts", - columns, maxBlockColumns) - } - if rows > maxBlockRows { - return fmt.Sprintf("the data block declares %d rows, more than the %d this session accepts. "+ - "Send the data in smaller batches", rows, maxBlockRows) - } - return "" -} - -func featureBlockInfo(revision int) bool { return revision >= 51903 } +// Runs inside DecodeRawBlock after the header is read and before any column is allocated. +type boundedResult struct{ inner proto.Result } -type peeker struct { - buf []byte - pos int -} - -func (p *peeker) byteAt() (byte, bool) { - if p.pos >= len(p.buf) { - return 0, false - } - b := p.buf[p.pos] - p.pos++ - return b, true -} - -func (p *peeker) uvarint() (uint64, bool) { - v, n := binary.Uvarint(p.buf[p.pos:]) - if n <= 0 { - return 0, false +func (b boundedResult) DecodeResult(r *proto.Reader, version int, block proto.Block) error { + if block.Columns > maxBlockColumns { + return fmt.Errorf("the data block declares %d columns, more than the %d this session accepts", + block.Columns, maxBlockColumns) } - p.pos += n - return v, true -} - -// Mirrors BlockInfo.Decode: field-id/value pairs terminated by field 0. -func (p *peeker) skipBlockInfo() bool { - for { - field, ok := p.uvarint() - if !ok { - return false - } - switch field { - case blockInfoEnd: - return true - case blockInfoOverflows: - if _, ok := p.byteAt(); !ok { - return false - } - case blockInfoBucketNum: - if p.pos+4 > len(p.buf) { - return false - } - p.pos += 4 - default: - return false - } + if block.Rows > maxBlockRows { + return fmt.Errorf("the data block declares %d rows, more than the %d this session accepts; "+ + "send the data in smaller batches", block.Rows, maxBlockRows) } + return b.inner.DecodeResult(r, version, block) } diff --git a/packages/pam/handlers/clickhouse/bounded_block_test.go b/packages/pam/handlers/clickhouse/bounded_block_test.go index 46c63149b..7d268195f 100644 --- a/packages/pam/handlers/clickhouse/bounded_block_test.go +++ b/packages/pam/handlers/clickhouse/bounded_block_test.go @@ -1,103 +1,41 @@ package clickhouse import ( - "bufio" - "bytes" "testing" "github.com/ClickHouse/ch-go/proto" "github.com/stretchr/testify/require" ) -// blockHeader builds the wire prefix of a Data block body: BlockInfo, then columns and rows. -func blockHeader(columns, rows uint64) []byte { - var b proto.Buffer - b.PutUVarInt(blockInfoOverflows) - b.PutBool(false) - b.PutUVarInt(blockInfoBucketNum) - b.PutInt32(-1) - b.PutUVarInt(blockInfoEnd) - b.PutUVarInt(columns) - b.PutUVarInt(rows) - return b.Buf -} - -func scan(payload []byte, compressed bool) string { - return checkBlockHeader(bufio.NewReaderSize(bytes.NewReader(payload), 64<<10), maxNativeRevision, compressed) -} - -func TestBlockHeaderScanRefusesAnOversizedBlock(t *testing.T) { - // ch-go allocates rows x width before reading, so this is committed memory, not a recoverable panic. - reason := scan(blockHeader(1, 100_000_000), false) - require.Contains(t, reason, "rows") - require.Contains(t, reason, "smaller batches") +type recordingResult struct{ called bool } - require.Contains(t, scan(blockHeader(1_000_000, 1), false), "columns") +func (r *recordingResult) DecodeResult(*proto.Reader, int, proto.Block) error { + r.called = true + return nil } -func TestBlockHeaderScanPassesALegitimateBlock(t *testing.T) { +func TestBoundedResultEnforcesTheBlockLimits(t *testing.T) { for _, c := range []struct { - name string - columns, rows uint64 + name string + block proto.Block + wantErr string }{ - {"an empty block", 0, 0}, - {"a single row", 1, 1}, - {"a default max_insert_block_size batch", 8, 1_048_545}, - {"exactly at the row cap", 1, maxBlockRows}, - {"exactly at the column cap", maxBlockColumns, 1}, + {"within limits", proto.Block{Columns: 8, Rows: 1_048_545}, ""}, + {"at the row limit", proto.Block{Columns: 1, Rows: maxBlockRows}, ""}, + {"at the column limit", proto.Block{Columns: maxBlockColumns, Rows: 1}, ""}, + {"over the row limit", proto.Block{Columns: 1, Rows: maxBlockRows + 1}, "rows"}, + {"over the column limit", proto.Block{Columns: maxBlockColumns + 1, Rows: 1}, "columns"}, } { t.Run(c.name, func(t *testing.T) { - require.Empty(t, scan(blockHeader(c.columns, c.rows), false)) + inner := &recordingResult{} + err := boundedResult{inner}.DecodeResult(nil, maxNativeRevision, c.block) + if c.wantErr == "" { + require.NoError(t, err) + require.True(t, inner.called, "a block within limits must reach the decoder") + return + } + require.ErrorContains(t, err, c.wantErr) + require.False(t, inner.called, "an over-limit block must be refused before the decoder runs") }) } } - -// The scan must never be the thing that rejects traffic: anything it cannot parse is left to the decoder. -func TestBlockHeaderScanFallsThroughWhenItCannotParse(t *testing.T) { - require.Empty(t, scan(nil, false), "empty input") - require.Empty(t, scan([]byte{0xFF}, false), "a truncated varint") - require.Empty(t, scan([]byte{0x09, 0x01}, false), "an unknown BlockInfo field") - require.Empty(t, scan(blockHeader(1, 100_000_000), true), "a compressed block is ch-go's to bound") -} - -// The bytes the scan inspects must still reach the decoder, or the packet would be truncated. -func TestBlockHeaderScanConsumesNothing(t *testing.T) { - payload := append(blockHeader(2, 7), []byte("trailing")...) - src := bufio.NewReaderSize(bytes.NewReader(payload), 64<<10) - - require.Empty(t, checkBlockHeader(src, maxNativeRevision, false)) - - got := make([]byte, len(payload)) - _, err := src.Read(got[:1]) - require.NoError(t, err) - require.Equal(t, payload[0], got[0], "the scan must not consume the header") - require.Equal(t, len(payload)-1, src.Buffered()+0, "everything after the first byte is still pending") -} - -// The whole point: a client must not be able to make the gateway size a column from a declared row count. -func TestNativeRefusesAnOversizedDataBlock(t *testing.T) { - upstream := startFakeClickHouse(t) - - conn := dialProxy(t, ClickHouseProxyConfig{ - NativeAddr: upstream.addr(), - Username: "account", - SessionID: "unit", - SessionLogger: &recordingLogger{}, - }) - r := clientHandshake(t, conn, "someone", "whatever") - - var b proto.Buffer - proto.ClientCodeData.Encode(&b) - b.PutString("") - b.Buf = append(b.Buf, blockHeader(1, 100_000_000)...) - _, err := conn.Write(b.Buf) - require.NoError(t, err) - - code, message := decodeException(t, r) - require.Equal(t, codeNotImplemented, code) - require.Contains(t, message, "rows") - - _, _, queries, bytesAfter := upstream.snapshot() - require.Empty(t, queries) - require.Zero(t, bytesAfter, "a refused block must not be relayed upstream") -} diff --git a/packages/pam/handlers/clickhouse/native.go b/packages/pam/handlers/clickhouse/native.go index c144466d4..6eaf888f6 100644 --- a/packages/pam/handlers/clickhouse/native.go +++ b/packages/pam/handlers/clickhouse/native.go @@ -74,9 +74,6 @@ func (t *tap) discard() { t.buf = nil } -// peeker exposes the buffered source so a header can be inspected without consuming it. -func (t *tap) peeker() *bufio.Reader { return t.src } - // A refusal ends the session: the stream is mid-packet. var errSessionRefused = errors.New("the session was refused") @@ -426,11 +423,6 @@ func (s *nativeSession) handleData(t *tap, r *proto.Reader) error { compressed := s.compressed.Load() - if reason := checkBlockHeader(t.peeker(), s.rev, compressed); reason != "" { - s.log.Warn().Str("table", table).Msg("Refused an oversized ClickHouse data block") - return s.refuse(t, codeNotImplemented, reason) - } - if compressed { r.EnableCompression() } @@ -438,7 +430,7 @@ func (s *nativeSession) handleData(t *tap, r *proto.Reader) error { block proto.Block discard proto.Results ) - decodeErr := block.DecodeBlock(r, s.rev, discard.Auto()) + decodeErr := block.DecodeBlock(r, s.rev, boundedResult{discard.Auto()}) if compressed { r.DisableCompression() } From f33020f4d1b7fc5c4ecce6860451800466ed6ca8 Mon Sep 17 00:00:00 2001 From: bernie-g Date: Mon, 28 Sep 2026 12:46:05 -0400 Subject: [PATCH 15/16] test(clickhouse): restore the session-level oversized block check --- .../handlers/clickhouse/bounded_block_test.go | 40 +++++++++++++++++++ 1 file changed, 40 insertions(+) diff --git a/packages/pam/handlers/clickhouse/bounded_block_test.go b/packages/pam/handlers/clickhouse/bounded_block_test.go index 7d268195f..a58dd6b0e 100644 --- a/packages/pam/handlers/clickhouse/bounded_block_test.go +++ b/packages/pam/handlers/clickhouse/bounded_block_test.go @@ -39,3 +39,43 @@ func TestBoundedResultEnforcesTheBlockLimits(t *testing.T) { }) } } + +func blockHeader(columns, rows uint64) []byte { + var b proto.Buffer + b.PutUVarInt(1) + b.PutBool(false) + b.PutUVarInt(2) + b.PutInt32(-1) + b.PutUVarInt(0) + b.PutUVarInt(columns) + b.PutUVarInt(rows) + return b.Buf +} + +// Wires the limit to the packet path, which the unit test above cannot see. +func TestNativeRefusesAnOversizedDataBlock(t *testing.T) { + upstream := startFakeClickHouse(t) + + conn := dialProxy(t, ClickHouseProxyConfig{ + NativeAddr: upstream.addr(), + Username: "account", + SessionID: "unit", + SessionLogger: &recordingLogger{}, + }) + r := clientHandshake(t, conn, "someone", "whatever") + + var b proto.Buffer + proto.ClientCodeData.Encode(&b) + b.PutString("") + b.Buf = append(b.Buf, blockHeader(1, 100_000_000)...) + _, err := conn.Write(b.Buf) + require.NoError(t, err) + + code, message := decodeException(t, r) + require.Equal(t, codeNotImplemented, code) + require.Contains(t, message, "rows") + + _, _, queries, bytesAfter := upstream.snapshot() + require.Empty(t, queries) + require.Zero(t, bytesAfter, "a refused block must not be relayed upstream") +} From 599886d423d95a40b66d07fa5bfd5e2b0649deca Mon Sep 17 00:00:00 2001 From: bernie-g Date: Mon, 28 Sep 2026 14:23:42 -0400 Subject: [PATCH 16/16] chore(clickhouse): explain the bounded decoder and fix a truncated comment --- packages/pam/handlers/clickhouse/bounded_decode.go | 4 +++- packages/pam/handlers/clickhouse/contract_test.go | 2 +- 2 files changed, 4 insertions(+), 2 deletions(-) diff --git a/packages/pam/handlers/clickhouse/bounded_decode.go b/packages/pam/handlers/clickhouse/bounded_decode.go index a7ae1c62e..60010f847 100644 --- a/packages/pam/handlers/clickhouse/bounded_decode.go +++ b/packages/pam/handlers/clickhouse/bounded_decode.go @@ -1,3 +1,6 @@ +// ch-go allocates a client-declared length before reading it, so one packet could exhaust the shared gateway. +// These mirror its Hello and Query decoders with every length capped. Re-check them on ch-go upgrades. + package clickhouse import ( @@ -9,7 +12,6 @@ import ( "go.opentelemetry.io/otel/trace" ) -// Mirrors ch-go's decoders, but every string is read through a cap. const ( // Short identifiers: names, users, hostnames, the quota key. maxHandshakeStringLen = 64 << 10 diff --git a/packages/pam/handlers/clickhouse/contract_test.go b/packages/pam/handlers/clickhouse/contract_test.go index c303d774c..917f6826a 100644 --- a/packages/pam/handlers/clickhouse/contract_test.go +++ b/packages/pam/handlers/clickhouse/contract_test.go @@ -8,7 +8,7 @@ import ( "github.com/stretchr/testify/require" ) -// The API and the gateway agree on these shapes only by convention, and a renamed field would not fail to... +// The API and gateway share these field names by convention only; a rename would decode as zero, not fail. func TestSessionCredentialsContract(t *testing.T) { cases := []struct { name string