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/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..87bc9134f 100644 --- a/packages/gateway-v2/capabilities.go +++ b/packages/gateway-v2/capabilities.go @@ -5,3 +5,6 @@ package gatewayv2 const CapabilitySessionLogMaskingBuiltInDetection = "sessionLogMaskingBuiltInDetection" const CapabilitySupportedAccountTypes = "supported_account_types" + +// 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 bf04f16b0..5c174ebc7 100644 --- a/packages/gateway-v2/discovery_handler.go +++ b/packages/gateway-v2/discovery_handler.go @@ -26,6 +26,21 @@ const ( type rpcTarget struct { host string port int + // Empty when the certificate named only port, which is what an older platform produces. + ports []int +} + +// 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 + } + for _, allowed := range t.ports { + if port == allowed { + return true + } + } + return false } type rpcTargetContextKey struct{} @@ -63,7 +78,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..27495e84b 100644 --- a/packages/gateway-v2/gateway.go +++ b/packages/gateway-v2/gateway.go @@ -80,6 +80,7 @@ type ForwardConfig struct { VerifyTLS bool // Whether to verify TLS certificates TargetHost string TargetPort int + TargetPorts []int ActorType ActorType PAMConfig pam.GatewayPAMConfig } @@ -88,6 +89,8 @@ type ForwardConfig struct { type RoutingInfo struct { TargetHost string `json:"targetHost"` TargetPort int `json:"targetPort"` + // 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 +467,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 +1455,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..17bd7a90c 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,81 @@ 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, + } + + // 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 connectFailure(fmt.Errorf("port %d is not authorised for this connection test", port)) + } + } + + httpPort := params.HttpPort + if httpPort <= 0 && params.NativePort <= 0 { + // An API too old to send the ports still means the cert-bound one. + httpPort = target.port } - 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, - })) + + // One shared deadline would let a slow first probe swallow the second one's specific error. + probes := 0 + if httpPort > 0 { + probes++ + } + 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) { + deadline, ok := ctx.Deadline() + if !ok || remaining <= 1 { + remaining-- + return context.WithCancel(ctx) + } + slice := time.Until(deadline) / time.Duration(remaining) + remaining-- + return context.WithTimeout(ctx, slice) + } + + 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 +837,12 @@ func redactProbeSecrets(msg string, secrets ...string) string { } return urlUserinfoPattern.ReplaceAllString(msg, "${1}******@") } + +// 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) + } + 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..a2ff5f7eb --- /dev/null +++ b/packages/gateway-v2/test_connection_handler_test.go @@ -0,0 +1,73 @@ +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. +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... +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/bounded_block.go b/packages/pam/handlers/clickhouse/bounded_block.go new file mode 100644 index 000000000..b04dbfc9b --- /dev/null +++ b/packages/pam/handlers/clickhouse/bounded_block.go @@ -0,0 +1,27 @@ +package clickhouse + +import ( + "fmt" + + "github.com/ClickHouse/ch-go/proto" +) + +const ( + maxBlockRows = 4 << 20 + maxBlockColumns = 4096 +) + +// Runs inside DecodeRawBlock after the header is read and before any column is allocated. +type boundedResult struct{ inner proto.Result } + +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) + } + 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 new file mode 100644 index 000000000..a58dd6b0e --- /dev/null +++ b/packages/pam/handlers/clickhouse/bounded_block_test.go @@ -0,0 +1,81 @@ +package clickhouse + +import ( + "testing" + + "github.com/ClickHouse/ch-go/proto" + "github.com/stretchr/testify/require" +) + +type recordingResult struct{ called bool } + +func (r *recordingResult) DecodeResult(*proto.Reader, int, proto.Block) error { + r.called = true + return nil +} + +func TestBoundedResultEnforcesTheBlockLimits(t *testing.T) { + for _, c := range []struct { + name string + block proto.Block + wantErr string + }{ + {"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) { + 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") + }) + } +} + +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") +} diff --git a/packages/pam/handlers/clickhouse/bounded_decode.go b/packages/pam/handlers/clickhouse/bounded_decode.go new file mode 100644 index 000000000..60010f847 --- /dev/null +++ b/packages/pam/handlers/clickhouse/bounded_decode.go @@ -0,0 +1,336 @@ +// 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 ( + "fmt" + "io" + + "github.com/ClickHouse/ch-go/proto" + "github.com/segmentio/asm/bswap" + "go.opentelemetry.io/otel/trace" +) + +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 + // Individually legal fields still add up, so the packet carries one budget across all of them. + maxQueryPacketBytes = 32 << 20 +) + +// 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 + } + 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 + } + return string(buf), nil +} + +func readBoundedStr(r *proto.Reader) (string, error) { + 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) { + 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, b *budget) (proto.Setting, error) { + var s proto.Setting + + key, err := readBudgetedStr(r, b) + 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, b) + 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, b *budget) (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 = readBudgetedStr(r, b); err != nil { + return c, fmt.Errorf("initial user: %w", err) + } + if c.InitialQueryID, err = readBudgetedStr(r, b); err != nil { + return c, fmt.Errorf("initial query id: %w", err) + } + if c.InitialAddress, err = readBudgetedStr(r, b); 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 = readBudgetedStr(r, b); err != nil { + return c, fmt.Errorf("os user: %w", err) + } + if c.ClientHostname, err = readBudgetedStr(r, b); err != nil { + return c, fmt.Errorf("client hostname: %w", err) + } + 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 { + 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 = readBudgetedStr(r, b); 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 := readBudgetedStr(r, b) + 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) { + b := newBudget() + var q proto.Query + var err error + + 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, b); 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, b) + 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 = readBudgetedStr(r, b); 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, b); err != nil { + return q, fmt.Errorf("query body: %w", err) + } + + if proto.FeatureParameters.In(version) { + for { + s, err := decodeBoundedSetting(r, b) + 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..c460c41fa --- /dev/null +++ b/packages/pam/handlers/clickhouse/bounded_decode_test.go @@ -0,0 +1,183 @@ +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)) +} + +// 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}, + 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") +} + +// 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/contract_test.go b/packages/pam/handlers/clickhouse/contract_test.go new file mode 100644 index 000000000..917f6826a --- /dev/null +++ b/packages/pam/handlers/clickhouse/contract_test.go @@ -0,0 +1,60 @@ +package clickhouse + +import ( + "encoding/json" + "testing" + + "github.com/Infisical/infisical-merge/packages/api" + "github.com/stretchr/testify/require" +) + +// 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 + 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/native.go b/packages/pam/handlers/clickhouse/native.go new file mode 100644 index 000000000..6eaf888f6 --- /dev/null +++ b/packages/pam/handlers/clickhouse/native.go @@ -0,0 +1,675 @@ +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" +) + +const ( + // 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 + nativeHandshakeTimeout = 10 * time.Second + nativeIdleTimeout = 12 * time.Hour +) + +type nativeProxy struct { + *ClickHouseProxy +} + +func newNativeProxy(owner *ClickHouseProxy) *nativeProxy { + return &nativeProxy{ClickHouseProxy: owner} +} + +// One byte at a time, or proto.Reader buffers past the packet. +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 +} + +func (t *tap) take() []byte { + b := t.buf + t.buf = nil + return b +} + +// Relaying from the socket instead would drop whatever the tap's buffer already holds. +func (t *tap) rest() io.Reader { + return t.src +} + +func (t *tap) discard() { + t.buf = nil +} + +// A refusal ends the session: the stream is mid-packet. +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 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 mid-write. + writeMu sync.Mutex + refused atomic.Bool + + compressed atomic.Bool +} + +func (s *nativeSession) writeToClient(payload []byte) error { + if len(payload) == 0 { + return nil + } + s.writeMu.Lock() + defer s.writeMu.Unlock() + return s.writeClientLocked(payload) +} + +// 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 + } + 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{}) }() + + _, 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() + 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 then says nothing is the HTTP port entered as the native one. + 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{}) + + serverDone := make(chan struct{}) + go func() { + defer close(serverDone) + // Upstream gone: end the session rather than block 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") + } + }() + s.serverLoop() + }() + + // 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() + 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") + } + return nil +} + +func (s *nativeSession) refuse(t *tap, code int, message string) error { + if t != nil { + t.discard() + } + 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 and pins the 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) + } + + hello, err := decodeBoundedClientHello(r) + if err != nil { + return fmt.Errorf("decode client hello: %w", err) + } + t.discard() + + // 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.", + 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) + } + // 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 + } + 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 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 := readBoundedStr(r); 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 +} + +// 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)) + + 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: + // 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.") + + 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 { + 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 "+ + "applied to it.") + } + t.discard() + + // 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.", + int(q.Stage))) + } + + s.compressed.Store(q.Compression == proto.CompressionEnabled) + + statement := q.Body + nativeParameterSuffix(q.Parameters) + + 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, + "This statement is blocked by the command blocking policy on this account.") + } + + s.outcomes.begin(statement) + + // InitialAddress stays set: ClickHouse asserts on an empty one. + q.Info.QuotaKey = "" + q.Info.Query = proto.ClientQueryInitial + q.Secret = "" + + var b proto.Buffer + q.EncodeAware(&b, s.rev) + return s.forward(b.Buf) +} + +// 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 { + 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, boundedResult{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()) +} + +// 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 + + relayRest := func(reason string) { + s.outcomes.degrade(reason) + if err := s.writeToClientUnlessRefused(t.take()); err != nil { + return + } + _, _ = 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 err := s.writeToClientUnlessRefused(t.take()); err != nil { + return + } + } +} + +// 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} } + +func (w refusalAwareWriter) Write(p []byte) (int, error) { + if err := w.s.writeToClientUnlessRefused(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 +} + +// 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 "" + } + pairs := make([]string, 0, len(parameters)) + for _, p := range parameters { + pairs = append(pairs, p.Key+"="+p.Value) + } + return "\n-- parameters: " + strings.Join(pairs, " ") +} + +// 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) + } + return nil +} + +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 +} + +// 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() + + conn, err := (&nativeProxy{ClickHouseProxy: &ClickHouseProxy{config: config}}).dialUpstream(dialCtx) + if err != nil { + return err + } + defer conn.Close() + + // 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, 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 + 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) { + 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) + } + 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_outcome.go b/packages/pam/handlers/clickhouse/native_outcome.go new file mode 100644 index 000000000..ea3dd8365 --- /dev/null +++ b/packages/pam/handlers/clickhouse/native_outcome.go @@ -0,0 +1,99 @@ +package clickhouse + +import ( + "fmt" + "strings" + "sync" + "time" +) + +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() +} + +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())), ", ") +} + +// 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 { + 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..0d54f6d68 --- /dev/null +++ b/packages/pam/handlers/clickhouse/native_outcome_test.go @@ -0,0 +1,81 @@ +package clickhouse + +import ( + "strings" + "testing" + + "github.com/ClickHouse/ch-go/proto" + + "github.com/stretchr/testify/require" +) + +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..a30fc3cdc --- /dev/null +++ b/packages/pam/handlers/clickhouse/native_unit_test.go @@ -0,0 +1,715 @@ +package clickhouse + +import ( + "bytes" + "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" +) + +// fakeClickHouse stands in for a server so the security-critical parts of the handshake and packet loop... +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{} + + // 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 + // 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 { + 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() { + defer close(f.done) + conn, acceptErr := listener.Accept() + if acceptErr != nil { + return + } + defer conn.Close() + f.mu.Lock() + f.conn = conn + f.mu.Unlock() + 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) + 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 { + 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 +} + +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 +} + +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") +} + +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") +} + +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") +} + +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) +} + +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) { + writeQueryAt(t, conn, q, proto.Version) +} + +func writeQueryAt(t *testing.T, conn net.Conn, q proto.Query, rev int) { + t.Helper() + + q.Info.ProtocolVersion = rev + 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, rev) + _, 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. +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) +} + +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) { + 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(), 2*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") + // 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") + 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") + }) +} + +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") +} + +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") +} + +// 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") + + // Set before disconnect: the proxy may close this end first. + require.NoError(t, conn.SetReadDeadline(time.Now().Add(10*time.Second))) + + upstream.disconnect() + + 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") +} + +// 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) + } +} diff --git a/packages/pam/handlers/clickhouse/proxy.go b/packages/pam/handlers/clickhouse/proxy.go index 60cf4e85a..41d8a62b3 100644 --- a/packages/pam/handlers/clickhouse/proxy.go +++ b/packages/pam/handlers/clickhouse/proxy.go @@ -27,11 +27,11 @@ 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 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 Username string Password string Database string @@ -82,7 +82,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 @@ -134,12 +136,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,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, @@ -181,6 +210,14 @@ func (p *ClickHouseProxy) handler(l zerolog.Logger) http.Handler { return } + 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))) @@ -193,16 +230,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") } @@ -211,7 +248,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) } @@ -220,7 +257,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] @@ -228,18 +265,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 { @@ -321,7 +359,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 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() @@ -425,13 +467,13 @@ 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 - } +// 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 { - 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 e1d6f945a..4019c62db 100644 --- a/packages/pam/handlers/clickhouse/proxy_test.go +++ b/packages/pam/handlers/clickhouse/proxy_test.go @@ -3,12 +3,14 @@ package clickhouse import ( "bytes" "compress/gzip" + "compress/zlib" "io" "net/http" "net/http/httptest" "net/url" "regexp" "strings" + "sync" "testing" "github.com/Infisical/infisical-merge/packages/pam/session" @@ -17,13 +19,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 } @@ -421,3 +447,61 @@ 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()) +} + +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) +} diff --git a/packages/pam/handlers/clickhouse/sniff.go b/packages/pam/handlers/clickhouse/sniff.go new file mode 100644 index 000000000..f1a790ed5 --- /dev/null +++ b/packages/pam/handlers/clickhouse/sniff.go @@ -0,0 +1,44 @@ +package clickhouse + +import ( + "bufio" + "net" + "time" +) + +// 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. +const sniffTimeout = 30 * time.Second + +type peekConn struct { + net.Conn + reader *bufio.Reader +} + +func (c *peekConn) Read(p []byte) (int, error) { + return c.reader.Read(p) +} + +// 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() + } + return nil +} + +// 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..0598c76fc --- /dev/null +++ b/packages/pam/handlers/clickhouse/sniff_test.go @@ -0,0 +1,119 @@ +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]) + }) + } +} + +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) +} + +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..7c8d688af --- /dev/null +++ b/packages/pam/handlers/clickhouse/tap_test.go @@ -0,0 +1,49 @@ +package clickhouse + +import ( + "bytes" + "io" + "net" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +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/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..d2156b6ae 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") } + // 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)) + } + 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,