diff --git a/CHANGELOG-3.x.md b/CHANGELOG-3.x.md index c0914e0f..98d8ce20 100644 --- a/CHANGELOG-3.x.md +++ b/CHANGELOG-3.x.md @@ -1,3 +1,7 @@ +# v3.3.1 +- Schedule TLS cert refresh from credential expiration +- Don't denylist file handles for transient S3 errors on readbypass path + # v3.3.0 - Eliminate unnecessary data copies on the readbypass path - Consolidating dependencies diff --git a/amazon-efs-utils.spec b/amazon-efs-utils.spec index c5c3fe63..e61648a1 100644 --- a/amazon-efs-utils.spec +++ b/amazon-efs-utils.spec @@ -41,8 +41,8 @@ %{?!include_vendor_tarball:%define include_vendor_tarball true} Name : amazon-efs-utils -Version : 3.3.0 -Release : 3%{platform} +Version : 3.3.1 +Release : 1%{platform} Summary : This package provides utilities for simplifying the use of EFS file systems Group : Amazon/Tools @@ -221,6 +221,10 @@ fi %clean %changelog +* Sat Aug 23 2026 Yue Wang - 3.3.1 +- Schedule TLS cert refresh from credential expiration +- Don't denylist file handles for transient S3 errors on readbypass path + * Wed Aug 5 2026 Zachary Maguire 3.3.0 - Eliminate unnecessary data copies on the readbypass path - Consolidating dependencies diff --git a/build-deb.sh b/build-deb.sh index 3ad4f299..ea5c658d 100755 --- a/build-deb.sh +++ b/build-deb.sh @@ -11,8 +11,8 @@ set -ex BASE_DIR=$(pwd) BUILD_ROOT=${BASE_DIR}/build/debbuild -VERSION=3.3.0 -RELEASE=3 +VERSION=3.3.1 +RELEASE=1 ARCH=$(dpkg --print-architecture) DEB_SYSTEM_RELEASE_PATH=/etc/os-release export VERSION RELEASE ARCH diff --git a/config.ini b/config.ini index f420de01..18b27272 100644 --- a/config.ini +++ b/config.ini @@ -7,5 +7,5 @@ # [global] -version=3.3.0 -release=3 +version=3.3.1 +release=1 diff --git a/src/Cargo.lock b/src/Cargo.lock index 36cf0bcd..b8703665 100644 --- a/src/Cargo.lock +++ b/src/Cargo.lock @@ -64,11 +64,11 @@ dependencies = [ "dyn-clone", "futures", "libc", - "log 0.4.33", - "lru", + "log 0.4.34", + "lru 0.16.4", "moka", "onc-rpc", - "rand 0.8.7", + "rand 0.8.8", "regex-lite", "s2n-tls", "serde", @@ -170,7 +170,7 @@ checksum = "82f6aeea286b8eb4dd3431a1be1b59d290ace00f5bfd8e2a159bc2a05e2c1667" dependencies = [ "proc-macro2", "quote 1.0.47", - "syn 3.0.3", + "syn 3.0.4", ] [[package]] @@ -209,9 +209,9 @@ checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53" [[package]] name = "aws-config" -version = "1.10.1" +version = "1.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1b180a3c8b55960db3426d8964b8745e652466a1a49fe1a2eda828046d30b5e4" +checksum = "a767267da9e2c2e189b2f9df8b5657e850ecf5352644734ba130d4a57095cf1b" dependencies = [ "aws-credential-types", "aws-runtime", @@ -320,9 +320,9 @@ dependencies = [ [[package]] name = "aws-sdk-cloudwatch" -version = "1.123.0" +version = "1.127.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "72d33759772cc816ccd6d6c10a58f0c193405b6ed60e89ed3f5e3ea7f29f97d2" +checksum = "ac606233394937c6014a46fc776214f5eed0b6bde3bdd38c8fa1aa28e0f3ad09" dependencies = [ "arc-swap", "aws-credential-types", @@ -350,9 +350,9 @@ dependencies = [ [[package]] name = "aws-sdk-cloudwatchlogs" -version = "1.145.0" +version = "1.148.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "84f68887a6c2bcb78e0329dbbefef18fa792031110796e90789f370cb777b6d5" +checksum = "134d7e1265426faae1ebcc1d4ececb9f75a8049a676449c505b4c739a18e138e" dependencies = [ "arc-swap", "aws-credential-types", @@ -377,9 +377,9 @@ dependencies = [ [[package]] name = "aws-sdk-s3" -version = "1.141.0" +version = "1.144.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d9f9420d3a2467eed22ed3635ca653653162c386a0b0f65c78189f9bd3c1379e" +checksum = "30dc8bf6baaf7d46336a0ca2c69f223d9b90d7a801fb3e28f7ea17b00dc6b1de" dependencies = [ "arc-swap", "aws-credential-types", @@ -404,7 +404,7 @@ dependencies = [ "http 0.2.12", "http 1.5.0", "http-body 1.1.0", - "lru", + "lru 0.18.2", "percent-encoding", "regex-lite", "sha2 0.11.0", @@ -414,9 +414,9 @@ dependencies = [ [[package]] name = "aws-sdk-sso" -version = "1.105.0" +version = "1.108.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6ffd0fbe7873cb548a7aa60f9573c268fff94155397fd4f14dc9f1ecaaab8516" +checksum = "c15301b04372832947916607983b114b3374b9db0be058a00fb7513800de1f05" dependencies = [ "arc-swap", "aws-credential-types", @@ -440,9 +440,9 @@ dependencies = [ [[package]] name = "aws-sdk-ssooidc" -version = "1.107.0" +version = "1.110.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "175763eb222a46377df7aa257a3bca980ab3e96703fefc8f4d0b8da6ad2e254c" +checksum = "72cc2c205cb27108183cf1856333f7d584c2ba0f505421b4209ca5828f9ea899" dependencies = [ "arc-swap", "aws-credential-types", @@ -466,9 +466,9 @@ dependencies = [ [[package]] name = "aws-sdk-sts" -version = "1.110.0" +version = "1.113.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dd8b14781dfbff48984017d57167b6ea0b6471c6920ec52b44a2677c7feb3c13" +checksum = "68182ecb449f7537db0f4d5d25917789cf41e32074a9fe47b6a0b847fe1d2032" dependencies = [ "arc-swap", "aws-credential-types", @@ -616,9 +616,9 @@ dependencies = [ [[package]] name = "aws-smithy-http-client" -version = "1.3.0" +version = "1.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3c1c8a04cb31ba74d0115af5a890bb8c0d48fba64b52812fa13929a6ef0cc83c" +checksum = "ebfd138fac0337cee7516c352757ea73b9f2266e57d0bcb5bc70e9547e45aef1" dependencies = [ "aws-smithy-async", "aws-smithy-protocol-test", @@ -626,7 +626,7 @@ dependencies = [ "aws-smithy-types", "bytes", "h2 0.3.27", - "h2 0.4.15", + "h2 0.4.19", "http 0.2.12", "http 1.5.0", "http-body 0.4.6", @@ -716,9 +716,9 @@ dependencies = [ [[package]] name = "aws-smithy-runtime" -version = "1.13.1" +version = "1.14.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "483b858ff67522011c4786310c5cd8fd88d0be7ea3d5f1a48328446300c4269e" +checksum = "b82e438d30e02a825d363bd639a9efaed68a8089d86101054b0081e7e0d3e606" dependencies = [ "aws-smithy-async", "aws-smithy-http", @@ -743,9 +743,9 @@ dependencies = [ [[package]] name = "aws-smithy-runtime-api" -version = "1.14.0" +version = "1.15.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3b98f2e1fd67ec06618f9c291e5e495a468e60519e44c9c1979cd0521f3affdb" +checksum = "954c563ce84507722d2679f07a35d21b9c6466b3872d513020d0281fc8112ac9" dependencies = [ "aws-smithy-async", "aws-smithy-runtime-api-macros", @@ -887,7 +887,7 @@ dependencies = [ "cexpr", "clang-sys", "itertools 0.13.0", - "log 0.4.33", + "log 0.4.34", "prettyplease", "proc-macro2", "quote 1.0.47", @@ -991,9 +991,9 @@ dependencies = [ [[package]] name = "cc" -version = "1.4.2" +version = "1.4.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5d262e149917187838d5b42777c8253bcb64500067342904e7d429499a6f277e" +checksum = "0ad534f4357a5264cce5019c989cf66a4f0dc4e0d1b1d15f8aacec0ff7360273" dependencies = [ "find-msvc-tools", "jobserver", @@ -1203,9 +1203,9 @@ dependencies = [ [[package]] name = "crc32fast" -version = "1.5.0" +version = "1.5.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9481c1c90cbf2ac953f07c8d4a58aa3945c425b7185c9154d67a65e4230da511" +checksum = "8498c871161e1742aaa9d52551b2d6ebdd4c3d45a3be423e3728f33b955be550" dependencies = [ "cfg-if", ] @@ -1399,7 +1399,7 @@ checksum = "c6232dd377dcc64799954cbd3a9bb882e9cdc1308ccd87b1c098f1fb2eaf82a8" dependencies = [ "proc-macro2", "quote 1.0.47", - "syn 3.0.3", + "syn 3.0.4", ] [[package]] @@ -1436,7 +1436,7 @@ dependencies = [ [[package]] name = "efs-proxy" -version = "3.3.0" +version = "3.3.1" dependencies = [ "amzn-efs-client-core", "amzn-nfs-xdr-bindings", @@ -1461,14 +1461,14 @@ dependencies = [ "futures", "hex-literal", "libc", - "log 0.4.33", + "log 0.4.34", "log4rs", - "lru", + "lru 0.16.4", "mockall", "moka", "nix", "onc-rpc", - "rand 0.8.7", + "rand 0.8.8", "regex 1.13.1", "s2n-tls", "s2n-tls-tokio", @@ -1489,9 +1489,9 @@ dependencies = [ [[package]] name = "either" -version = "1.17.0" +version = "1.18.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9e5e8f6c15a24b9a3ee5efec809ccd006d3b30e8b3bb63c39af737c7f87daa1d" +checksum = "252afb9ae5eaa683babdc6a068b3f5726eb19e05070c731f9b2a23a7c3e8ed34" [[package]] name = "elliptic-curve" @@ -1531,7 +1531,7 @@ checksum = "4cd405aab171cb85d6735e5c8d9db038c17d3ca007a4d2c25f337935c3d90580" dependencies = [ "humantime", "is-terminal", - "log 0.4.33", + "log 0.4.34", "regex 1.13.1", "termcolor", ] @@ -1579,9 +1579,9 @@ dependencies = [ [[package]] name = "find-msvc-tools" -version = "0.1.10" +version = "0.1.11" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "26b73573e6edcd2af0cdf47bd6cb58f0b3839491263c314eaad1ccf24430e1de" +checksum = "d45db016d36b838f563236e9193d0ee6ce38f3f68b6c94e914b4929c96bbb890" [[package]] name = "flate2" @@ -1700,7 +1700,7 @@ checksum = "9fb9654ba8355388abeb8dcb4fc62f511300867002afc858860463bdd9fe0c44" dependencies = [ "proc-macro2", "quote 1.0.47", - "syn 3.0.3", + "syn 3.0.4", ] [[package]] @@ -1822,9 +1822,9 @@ dependencies = [ [[package]] name = "h2" -version = "0.4.15" +version = "0.4.19" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6cb093c84e8bd9b188d4c4a8cb6579fc016968d14c99882163cd3ff402a4f155" +checksum = "ef8e5e5a340588f4452631496976cf8636d4a7ecf600239fdc27615d2530bc16" dependencies = [ "atomic-waker", "bytes", @@ -1872,6 +1872,11 @@ name = "hashbrown" version = "0.17.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ed5909b6e89a2db4456e54cd5f673791d7eca6732202bbf2a9cc504fe2f9b84a" +dependencies = [ + "allocator-api2", + "equivalent", + "foldhash", +] [[package]] name = "heck" @@ -2046,7 +2051,7 @@ dependencies = [ "bytes", "futures-channel", "futures-core", - "h2 0.4.15", + "h2 0.4.19", "http 1.5.0", "http-body 1.1.0", "httparse", @@ -2066,7 +2071,7 @@ dependencies = [ "futures-util", "http 0.2.12", "hyper 0.14.32", - "log 0.4.33", + "log 0.4.34", "rustls 0.21.12", "tokio", "tokio-rustls 0.24.1", @@ -2121,7 +2126,7 @@ dependencies = [ "core-foundation-sys", "iana-time-zone-haiku", "js-sys", - "log 0.4.33", + "log 0.4.34", "wasm-bindgen", "windows-core 0.62.2", ] @@ -2137,9 +2142,9 @@ dependencies = [ [[package]] name = "icu_collections" -version = "2.2.0" +version = "2.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2984d1cd16c883d7935b9e07e44071dca8d917fd52ecc02c04d5fa0b5a3f191c" +checksum = "fa68d21081c4a05d5a901a1c62add574c77048b6a1c67be3b50ce0b60d4ca513" dependencies = [ "displaydoc", "potential_utf", @@ -2151,9 +2156,9 @@ dependencies = [ [[package]] name = "icu_locale_core" -version = "2.2.0" +version = "2.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "92219b62b3e2b4d88ac5119f8904c10f8f61bf7e95b640d25ba3075e6cac2c29" +checksum = "d56e28588da92eee5c3201a6eff33fabdd49b62269c8938d4ff050ce4d900deb" dependencies = [ "displaydoc", "litemap", @@ -2164,9 +2169,9 @@ dependencies = [ [[package]] name = "icu_normalizer" -version = "2.2.0" +version = "2.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c56e5ee99d6e3d33bd91c5d85458b6005a22140021cc324cea84dd0e72cff3b4" +checksum = "12f9cf5f235641ed274641dd81c3f28d870e276763d0797aeeab72317b1c646f" dependencies = [ "icu_collections", "icu_normalizer_data", @@ -2178,16 +2183,17 @@ dependencies = [ [[package]] name = "icu_normalizer_data" -version = "2.2.0" +version = "2.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "da3be0ae77ea334f4da67c12f149704f19f81d1adf7c51cf482943e84a2bad38" +checksum = "1563da1ed3e0b3bf3d74c9b85917ac9c56464d2f57242270c09c9e752f8021a0" [[package]] name = "icu_properties" -version = "2.2.0" +version = "2.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bee3b67d0ea5c2cca5003417989af8996f8604e34fb9ddf96208a033901e70de" +checksum = "7e7ca276ad3145661a65914e6daf131ca5120cd3dcee8f8f3214b8875184a148" dependencies = [ + "displaydoc", "icu_collections", "icu_locale_core", "icu_properties_data", @@ -2198,15 +2204,15 @@ dependencies = [ [[package]] name = "icu_properties_data" -version = "2.2.0" +version = "2.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8e2bbb201e0c04f7b4b3e14382af113e17ba4f63e2c9d2ee626b720cbce54a14" +checksum = "e590f038c1464a96894fd6d10127e90a8be4509f56ff7ecef851b15cee0b7caa" [[package]] name = "icu_provider" -version = "2.2.0" +version = "2.3.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "139c4cf31c8b5f33d7e199446eff9c1e02decfc2f0eec2c8d71f65befa45b421" +checksum = "d27bbb9d3abbefac45d55f647c9de1d44aafcd1186eb91879afef17c396c3e73" dependencies = [ "displaydoc", "icu_locale_core", @@ -2358,9 +2364,9 @@ checksum = "32a66949e030da00e8c7d4434b251670a91556f4144941d37452769c25d58a53" [[package]] name = "litemap" -version = "0.8.2" +version = "0.8.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "92daf443525c4cce67b150400bc2316076100ce0b3686209eb8cf3c31612e6f0" +checksum = "47d9d19d1d6efa0109d2f65ff4c85cddd50bd572e5a00127ab10987290bcefae" [[package]] name = "lock_api" @@ -2377,14 +2383,14 @@ version = "0.3.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e19e8d5c34a3e0e2223db8e060f9e8264aeeb5c5fc64a4ee9965c062211c024b" dependencies = [ - "log 0.4.33", + "log 0.4.34", ] [[package]] name = "log" -version = "0.4.33" +version = "0.4.34" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0ceec5bc11778974d1bcb055b18002eba7f4b3518b6a0081b3af5f21666da9ad" +checksum = "f9f8bd3e56ce4dfc153cf470fffbfa98c7620958b312ca5c3a4b8d5181fd13c6" dependencies = [ "serde_core", ] @@ -2408,7 +2414,7 @@ dependencies = [ "fnv", "humantime", "libc", - "log 0.4.33", + "log 0.4.34", "log-mdc", "mock_instant", "parking_lot", @@ -2433,6 +2439,15 @@ dependencies = [ "hashbrown 0.16.1", ] +[[package]] +name = "lru" +version = "0.18.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5d2f2f9b4ba7e6b24d95e7e899329d35be83bcded72c8540cdd5368932d1d90a" +dependencies = [ + "hashbrown 0.17.1", +] + [[package]] name = "matchers" version = "0.2.0" @@ -2799,9 +2814,9 @@ dependencies = [ [[package]] name = "pkg-config" -version = "0.3.33" +version = "0.3.34" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "19f132c84eca552bf34cab8ec81f1c1dcc229b811638f9d283dceabe58c5569e" +checksum = "f6b464fbc74e149a392436b17d523f769e057cb6877f6a5c4618bc6f11800548" [[package]] name = "portable-atomic" @@ -2811,9 +2826,9 @@ checksum = "05c8b63e8d9609db387f0324918f81d68fe27748f084ef092fb35954d0539a85" [[package]] name = "potential_utf" -version = "0.1.5" +version = "0.1.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0103b1cef7ec0cf76490e969665504990193874ea05c85ff9bab8b911d0a0564" +checksum = "d83eb9bc6d8e5cf568e7a1101d60ee05e81ed50ea106026f3d18deeb046d7661" dependencies = [ "zerovec", ] @@ -2988,9 +3003,9 @@ dependencies = [ [[package]] name = "rand" -version = "0.8.7" +version = "0.8.8" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "22f6172bdec972074665ed81ed53b71da00bfc44b65a753cfde883ec4c702a1a" +checksum = "e058c7de0b26af77780c769414d6257830bb240f3c38477dbc2c16e5f54d6d4c" dependencies = [ "libc", "rand_chacha 0.3.1", @@ -3257,7 +3272,7 @@ version = "0.21.12" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3f56a14d1f48b391359b22f731fd4bd7e43c97f3c50eee276f3aa09c94784d3e" dependencies = [ - "log 0.4.33", + "log 0.4.34", "ring", "rustls-webpki 0.101.7", "sct", @@ -3272,7 +3287,7 @@ dependencies = [ "aws-lc-rs", "once_cell", "rustls-pki-types", - "rustls-webpki 0.103.14", + "rustls-webpki 0.103.15", "subtle", "zeroize", ] @@ -3310,9 +3325,9 @@ dependencies = [ [[package]] name = "rustls-webpki" -version = "0.103.14" +version = "0.103.15" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0527518605e68109d875e248ea259b6758801cf165e4b2c2733ae3b51f12535a" +checksum = "f3c3cf1d8b1e7d4927e2d154c3fcb02979afb9939629c62cd9048d4f07b60ac2" dependencies = [ "aws-lc-rs", "ring", @@ -3481,7 +3496,7 @@ checksum = "e7a5d71263a5a7d47b41f6b3f06ba276f10cc18b0931f1799f710578e2309348" dependencies = [ "proc-macro2", "quote 1.0.47", - "syn 3.0.3", + "syn 3.0.4", ] [[package]] @@ -3530,7 +3545,7 @@ checksum = "699f4197115b8a7e7ff19c9a315a4bd6fffec26cc4626ef45ecaea389e081c6d" dependencies = [ "futures-executor", "futures-util", - "log 0.4.33", + "log 0.4.34", "once_cell", "parking_lot", "serial_test_derive", @@ -3752,9 +3767,9 @@ dependencies = [ [[package]] name = "syn" -version = "3.0.3" +version = "3.0.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "53e9bae58849f64dfa4f5d5ae372c8341f7305f82a3868709269343628b659a3" +checksum = "e6275cddf4610d1775e6d1fe9469b2e77d0f39fd98fb7450901b821e0c53649f" dependencies = [ "proc-macro2", "quote 1.0.47", @@ -3909,7 +3924,7 @@ checksum = "bc04cd3e1236dd4a98afca4569f2deb3f120e5422a4023be2cb683f8486292af" dependencies = [ "proc-macro2", "quote 1.0.47", - "syn 3.0.3", + "syn 3.0.4", ] [[package]] @@ -4002,9 +4017,9 @@ dependencies = [ [[package]] name = "tinystr" -version = "0.8.3" +version = "0.8.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c8323304221c2a851516f22236c5722a72eaa19749016521d6dff0824447d96d" +checksum = "b1e27c91459209c2986af3dcf603a5a74a4368754ce37414f59acc971167f643" dependencies = [ "displaydoc", "zerovec", @@ -4050,7 +4065,7 @@ checksum = "78773a2a397f451582ce068015985c33193cf6dea8b74d2a639fe457b2f07b0e" dependencies = [ "proc-macro2", "quote 1.0.47", - "syn 3.0.3", + "syn 3.0.4", ] [[package]] @@ -4147,7 +4162,7 @@ version = "0.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ee855f1f400bd0e5c02d150ae5de3840039a3f54b025156404e34c23c03f47c3" dependencies = [ - "log 0.4.33", + "log 0.4.34", "once_cell", "tracing-core", ] @@ -4287,9 +4302,9 @@ checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" [[package]] name = "uuid" -version = "1.24.0" +version = "1.25.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bf3923a6f5c4c6382e0b653c4117f48d631ea17f38ed86e2a828e6f7412f5239" +checksum = "f053576934f05a761a402421fbbe3d425d9366f75f978806a037b3ca481abecc" dependencies = [ "getrandom 0.4.3", "js-sys", @@ -4641,9 +4656,9 @@ checksum = "1ebf944e87a7c253233ad6766e082e3cd714b5d03812acc24c318f549614536e" [[package]] name = "writeable" -version = "0.6.3" +version = "0.6.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1ffae5123b2d3fc086436f8834ae3ab053a283cfac8fe0a0b8eaae044768a4c4" +checksum = "3ad82d2a33cdc9674dc7465672f271e096168fcdbe0f799d9e6db8c5892679dc" [[package]] name = "xmlparser" @@ -4729,9 +4744,9 @@ checksum = "e13c156562582aa81c60cb29407084cdb54c4164760106ab78e6c5b0858cf64e" [[package]] name = "zerotrie" -version = "0.2.4" +version = "0.2.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0f9152d31db0792fa83f70fb2f83148effb5c1f5b8c7686c3459e361d9bc20bf" +checksum = "4ea269c3bd32f0a32c321907a2ae912ba6f4649bb0fc764a15627e99a7095a3f" dependencies = [ "displaydoc", "yoke", @@ -4740,9 +4755,9 @@ dependencies = [ [[package]] name = "zerovec" -version = "0.11.6" +version = "0.11.8" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "90f911cbc359ab6af17377d242225f4d75119aec87ea711a880987b18cd7b239" +checksum = "bb0464e17806c1d976d5cba29399c7f08e516e279e2ba493f63123b5fca67dd8" dependencies = [ "yoke", "zerofrom", @@ -4751,13 +4766,13 @@ dependencies = [ [[package]] name = "zerovec-derive" -version = "0.11.3" +version = "0.11.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "625dc425cab0dca6dc3c3319506e6593dcb08a9f387ea3b284dbd52a92c40555" +checksum = "34df6fc39dbd26ddc9c10e6a2984476e13acce22e64e4487636ef494369225da" dependencies = [ "proc-macro2", "quote 1.0.47", - "syn 2.0.119", + "syn 3.0.4", ] [[package]] diff --git a/src/client-core/Cargo.toml b/src/client-core/Cargo.toml index 62f399d7..85d16128 100644 --- a/src/client-core/Cargo.toml +++ b/src/client-core/Cargo.toml @@ -56,6 +56,7 @@ uuid = { version = "1.4.1", features = ["v4", "fast-rng", "macro-diagnostics"] } xdr_codec = { package = "amzn-xdr-codec", path = "../nfs-xdr-bindings/rust-xdr/xdr-codec" } amzn-nfs-xdr-bindings = { path = "../nfs-xdr-bindings" } + [dev-dependencies] test-case = "*" tokio = { version = "1.29.0", features = ["test-util"] } diff --git a/src/client-core/src/aws/cw_publisher.rs b/src/client-core/src/aws/cw_publisher.rs index 6ceb0ac9..bbac5c66 100644 --- a/src/client-core/src/aws/cw_publisher.rs +++ b/src/client-core/src/aws/cw_publisher.rs @@ -1,6 +1,8 @@ //! ### CloudWatch metrics and logs emission //! +use crate::sync::atomic::{AtomicBool, Ordering}; +use crate::sync::Arc; use crate::{ aws::credentials::get_aws_config_loader, aws::utils::get_ec2_instance_id, config_parser::ProxyConfig, utils::is_running_on_lambda, @@ -11,8 +13,6 @@ use aws_sdk_cloudwatchlogs::{ types::InputLogEvent, }; use log::{info, warn}; -use std::sync::atomic::{AtomicBool, Ordering}; -use std::sync::Arc; use std::time::{SystemTime, UNIX_EPOCH}; /// Metrics names and namespaces @@ -794,7 +794,7 @@ mod tests { } // Tests for ensure_log_group_and_stream - use std::sync::Arc; + use crate::sync::Arc; use tokio::sync::Mutex; struct MockCWLogsHelper { diff --git a/src/client-core/src/aws/s3_client.rs b/src/client-core/src/aws/s3_client.rs index d5df9e88..4bf47280 100644 --- a/src/client-core/src/aws/s3_client.rs +++ b/src/client-core/src/aws/s3_client.rs @@ -1,12 +1,8 @@ #![allow(unused)] -use std::{ - sync::{ - atomic::{AtomicBool, Ordering}, - Arc, - }, - time::Duration, -}; +use crate::sync::atomic::{AtomicBool, Ordering}; +use crate::sync::Arc; +use std::time::Duration; use anyhow::{Error, Result}; use async_trait::async_trait; @@ -27,6 +23,7 @@ use crate::{ cw_publisher::{CloudWatchClient, CloudWatchPublisher, LogLevel, CW_NAMESPACE_S3FILES}, }, config_parser::{ProxyConfig, ReadBypassConfig}, + memory::memory_pool::{MemoryChunk, CHUNK_SIZE}, }; const DEFAULT_PERMISSION_VALIDATION_SECONDS: u64 = 60; @@ -45,7 +42,7 @@ pub struct PermissionCheckResult { pub should_enable: bool, } -#[derive(Debug, thiserror::Error)] +#[derive(Debug, PartialEq, Eq, thiserror::Error)] pub enum S3ClientError { #[error("S3 Bucket is currently inaccessible.")] NotEnabled, @@ -59,8 +56,36 @@ pub enum S3ClientError { #[error("S3 returned {actual} bytes, expected {expected}")] SizeMismatch { expected: u64, actual: u64 }, - #[error(transparent)] - NoAccess(#[from] Box), + // GetObject failures not picked off as `NoSuchKey` / `ETagMismatchError` + // are classified into the variants below at construction, while the + // typed evidence (SDK failure kind, raw HTTP status) still exists — see + // `from_wire_failure`. The underlying SDK error is logged with key/etag + // context at the construction site; the variants carry only the class. + // They are evidence, not policy: they say what happened; callers own + // the retry/denylist verdict. + /// The request was denied (HTTP 403). + #[error("S3 GetObject access denied")] + AccessDenied, + + /// The request timed out before completing. + #[error("S3 GetObject timed out")] + Timeout, + + /// The service pushed back (HTTP 429 / 5xx). + #[error("S3 GetObject throttled")] + Throttled, + + /// The request never got an HTTP answer (dispatch/connection failure), + /// or the connection died mid-body-stream. Mid-stream failures are + /// always this class: once body frames are flowing S3 has already sent + /// the 200 status line — auth and `If-Match` were evaluated before it, + /// and service pushback arrives as a status code, never mid-body. + #[error("S3 GetObject connection failure")] + Connection, + + /// No classification evidence available. + #[error("S3 GetObject failed")] + Unknown, #[error("Permission Validator is already running")] ValidatorAlreadyRunning, @@ -78,6 +103,44 @@ pub enum S3ClientError { SemaphoreClosed, } +impl S3ClientError { + /// Classify a GetObject SdkError from its structure — failure kind + /// first, then the parsed service-error code, then the raw HTTP status. + /// (`NoSuchKey` and the structured 412 are picked off by + /// `extract_s3_error` before this runs; the 404/412 arms cover stray + /// raw responses that carried no parsed service error.) The caller + /// owns logging the SdkError. + fn from_wire_failure(e: &aws_sdk_s3::error::SdkError) -> Self { + use aws_sdk_s3::error::{ProvideErrorMetadata, SdkError}; + match e { + SdkError::TimeoutError(_) => Self::Timeout, + SdkError::DispatchFailure(d) => { + if d.as_connector_error().is_some_and(|c| c.is_timeout()) { + Self::Timeout + } else { + Self::Connection + } + } + _ => { + // S3 reports a slow-client timeout as HTTP 400 with error + // code RequestTimeout (retryable per the S3 docs), so the + // parsed code must be consulted before the raw status — + // a bare 400 arm would misclassify genuine bad requests. + if e.as_service_error().and_then(|se| se.code()) == Some("RequestTimeout") { + return Self::Timeout; + } + match e.raw_response().map(|r| r.status().as_u16()) { + Some(403) => Self::AccessDenied, + Some(404) => Self::NoSuchKey, + Some(412) => Self::ETagMismatchError, + Some(429) | Some(500) | Some(502) | Some(503) | Some(504) => Self::Throttled, + _ => Self::Unknown, + } + } + } + } +} + pub struct S3Client { bucket: String, prefix: String, @@ -106,6 +169,15 @@ impl Clone for S3Client { } } +/// The object identity one GET runs against, within the client's +/// configured bucket. A set `version_id` pins the GET to that object +/// version; otherwise the request is conditional on `etag` (`If-Match`). +pub struct GetObjectTarget<'a> { + pub key: &'a str, + pub etag: &'a str, + pub version_id: &'a str, +} + impl S3Client { pub async fn new(bucket: &str, prefix: &str, proxy_config: &ProxyConfig) -> Result { let mut aws_config_loader = get_aws_config_loader(proxy_config).await; @@ -368,19 +440,18 @@ impl S3Client { ) -> S3ClientError { let error = if let Some(GetObjectError::NoSuchKey(_)) = e.as_service_error() { S3ClientError::NoSuchKey - } else if let Some(resp) = e.raw_response() { - if resp.status().as_u16() == 412 { - S3ClientError::ETagMismatchError - } else { - S3ClientError::NoAccess(e.into()) - } + } else if e + .raw_response() + .is_some_and(|resp| resp.status().as_u16() == 412) + { + S3ClientError::ETagMismatchError } else { - S3ClientError::NoAccess(e.into()) + S3ClientError::from_wire_failure(&e) }; let msg = format!( - "S3 GetObject failed: key='{}', etag='{}', {}, error='{:?}'", - object_name, etag, context, error + "S3 GetObject failed ({}): key='{}', etag='{}', {}, error='{:?}'", + error, object_name, etag, context, e ); if let Some(publisher) = &self.cw_publisher { publisher.emit_log(LogLevel::Error, &msg); @@ -458,7 +529,21 @@ impl S3Client { // within the response instead of reallocating. let mut body = response.body; while let Some(frame) = body.next().await { - let frame = frame.map_err(|e| S3ClientError::NoAccess(e.into()))?; + let frame = frame.map_err(|e| { + // Mid-stream failures are always Connection: after the + // 200 status line, terminal verdicts are impossible. + // Log with object context, like the wire path in + // extract_s3_error: a partial transfer is hard to + // root-cause without it. + let msg = format!( + "S3 GetObject failed mid-stream: key='{}', etag='{}', error='{:?}'", + object_name, etag, e + ); + if let Some(publisher) = &self.cw_publisher { + publisher.emit_log(LogLevel::Error, &msg); + } + S3ClientError::Connection + })?; result.extend_from_slice(&frame); } } @@ -470,6 +555,91 @@ impl S3Client { !version_id.is_empty() && !version_id.eq_ignore_ascii_case("null") } + /// One ranged GET for `[offset, offset + count)`, streamed directly into + /// the caller's pre-claimed pool `chunks` — no full-body + /// materialization. Each received frame slice is copied exactly once, at + /// its final chunk offset (a frame straddling a chunk boundary + /// split-copies across the two chunks), and dropped immediately so hyper + /// reuses its per-connection read buffer. + /// + /// Chunk claiming — and therefore pool backpressure policy — belongs to + /// the caller; `chunks` must cover `count` bytes. The call issues a + /// SINGLE request: request sizing and concurrency are the caller's. + /// + /// Returns the number of bytes filled (always a contiguous prefix of + /// `chunks`). Fewer than `count` means the body ended inside the range + /// (the object ends there) and is the caller's to interpret; a body + /// longer than `count` fails as [`S3ClientError::SizeMismatch`] — the + /// chunk claim is never overrun. + pub async fn get_object_into_chunks( + &self, + target: &GetObjectTarget<'_>, + offset: u64, + count: usize, + chunks: &mut [MemoryChunk], + ) -> Result { + if !self.is_enabled() { + return Err(S3ClientError::NotEnabled); + } + if count == 0 { + return Ok(0); + } + debug_assert!( + count <= chunks.len() * CHUNK_SIZE, + "chunk claim must cover the requested range" + ); + let mut req = self + .client + .get_object() + .bucket(&self.bucket) + .key(target.key) + .range(format!("bytes={}-{}", offset, offset + count as u64 - 1)); + req = if Self::is_version_id_set(target.version_id) { + req.version_id(target.version_id) + } else { + req.if_match(target.etag) + }; + let response = req + .send() + .await + .map_err(|e| self.extract_s3_error(e, target.key, target.etag, "GetObject"))?; + let mut body = response.body; + let mut filled = 0usize; + while let Some(frame) = body.next().await { + let frame = frame.map_err(|e| { + warn!( + "S3 streaming GET body failed mid-stream: key='{}' offset={} filled={} error='{:?}'", + target.key, offset, filled, e + ); + S3ClientError::Connection + })?; + if filled + frame.len() > count { + warn!( + "S3 streaming GET body exceeded the requested range: key='{}' offset={} count={} received>={}", + target.key, + offset, + count, + filled + frame.len() + ); + return Err(S3ClientError::SizeMismatch { + expected: count as u64, + actual: (filled + frame.len()) as u64, + }); + } + // Split-copy: each slice lands at its final chunk offset. + let mut src = &frame[..]; + while !src.is_empty() { + let chunk_index = filled / CHUNK_SIZE; + let chunk_offset = filled % CHUNK_SIZE; + let n = src.len().min(CHUNK_SIZE - chunk_offset); + chunks[chunk_index][chunk_offset..chunk_offset + n].copy_from_slice(&src[..n]); + src = &src[n..]; + filled += n; + } + } + Ok(filled) + } + pub fn is_bucket_name_valid(bucket_name: &str) -> bool { // Basic s3 bucket name checker. Lookahead not supported, so more complicated checks like // IP addresses and consecutive periods cannot be included @@ -854,6 +1024,232 @@ mod tests { assert!(matches!(result.unwrap_err(), S3ClientError::NotEnabled)); } + /// Claim `n` chunks from a throwaway pool sized exactly `n`. + fn test_chunks(n: usize) -> Vec { + let pool = crate::memory::memory_pool::MemoryPool::new( + crate::memory::memory_pool::MemoryPoolConfig { + initial_capacity: n, + min_capacity: n, + max_capacity: n, + ..Default::default() + }, + ); + pool.consume(n) + } + + #[tokio::test] + pub async fn test_get_object_into_chunks_fills_prefix_and_reports_len() { + let expected_content = b"streamed-content"; + let get_object_rule = mock!(aws_sdk_s3::Client::get_object) + .match_requests(|req| { + req.if_match() == Some("test-etag") + && req.range() == Some("bytes=7-38") + && req.bucket() == Some("test_bucket") + }) + .then_output(|| { + GetObjectOutput::builder() + .content_length(expected_content.len() as i64) + .body(ByteStream::from_static(expected_content)) + .build() + }); + let mock_client = + create_test_s3_client(mock_client!(aws_sdk_s3, [&get_object_rule]), None, true); + + let mut chunks = test_chunks(1); + // count (32) exceeds the body (16): the short body is legal and the + // return value reports the contiguous filled prefix, not `count`. + let filled = mock_client + .get_object_into_chunks( + &GetObjectTarget { + key: "test_object", + etag: "test-etag", + version_id: "", + }, + 7, + 32, + &mut chunks, + ) + .await + .unwrap(); + assert_eq!(filled, expected_content.len()); + assert_eq!(&chunks[0][..filled], expected_content); + } + + #[tokio::test] + pub async fn test_get_object_into_chunks_split_copies_across_chunk_boundary() { + // Body larger than one chunk: the fill must cross the chunk boundary + // with each byte at its final offset (no intermediate buffer). + let body: Vec = (0..CHUNK_SIZE + 512).map(|i| (i % 251) as u8).collect(); + let body_clone = body.clone(); + let get_object_rule = mock!(aws_sdk_s3::Client::get_object).then_output(move || { + GetObjectOutput::builder() + .content_length(body_clone.len() as i64) + .body(ByteStream::from(body_clone.clone())) + .build() + }); + let mock_client = + create_test_s3_client(mock_client!(aws_sdk_s3, [&get_object_rule]), None, true); + + let mut chunks = test_chunks(2); + let count = body.len(); + let filled = mock_client + .get_object_into_chunks( + &GetObjectTarget { + key: "test_object", + etag: "test-etag", + version_id: "", + }, + 0, + count, + &mut chunks, + ) + .await + .unwrap(); + assert_eq!(filled, count); + assert_eq!(&chunks[0][..], &body[..CHUNK_SIZE]); + assert_eq!(&chunks[1][..512], &body[CHUNK_SIZE..]); + } + + #[tokio::test] + pub async fn test_get_object_into_chunks_overflow_is_size_mismatch() { + // A body LONGER than `count` must fail without overrunning the claim. + let get_object_rule = mock!(aws_sdk_s3::Client::get_object).then_output(|| { + GetObjectOutput::builder() + .content_length(24) + .body(ByteStream::from_static(b"twenty-four bytes long!!")) + .build() + }); + let mock_client = + create_test_s3_client(mock_client!(aws_sdk_s3, [&get_object_rule]), None, true); + + let mut chunks = test_chunks(1); + let result = mock_client + .get_object_into_chunks( + &GetObjectTarget { + key: "test_object", + etag: "test-etag", + version_id: "", + }, + 0, + 16, + &mut chunks, + ) + .await; + assert!(matches!( + result.unwrap_err(), + S3ClientError::SizeMismatch { + expected: 16, + actual: 24 + } + )); + } + + #[tokio::test] + pub async fn test_get_object_into_chunks_version_id_pins_the_version() { + let expected_content = b"versioned"; + let get_object_rule = mock!(aws_sdk_s3::Client::get_object) + .match_requests(|req| { + req.version_id() == Some("v7") + && req.if_match().is_none() + && req.bucket() == Some("test_bucket") + }) + .then_output(|| { + GetObjectOutput::builder() + .content_length(expected_content.len() as i64) + .body(ByteStream::from_static(expected_content)) + .build() + }); + let mock_client = + create_test_s3_client(mock_client!(aws_sdk_s3, [&get_object_rule]), None, true); + + let mut chunks = test_chunks(1); + let filled = mock_client + .get_object_into_chunks( + &GetObjectTarget { + key: "test_object", + etag: "test-etag", + version_id: "v7", + }, + 0, + 32, + &mut chunks, + ) + .await + .unwrap(); + assert_eq!(&chunks[0][..filled], expected_content); + } + + #[tokio::test] + pub async fn test_get_object_into_chunks_zero_count() { + // No rules registered: a request would panic the mock. + let mock_client = create_test_s3_client(mock_client!(aws_sdk_s3, []), None, true); + let filled = mock_client + .get_object_into_chunks( + &GetObjectTarget { + key: "test_object", + etag: "test-etag", + version_id: "", + }, + 0, + 0, + &mut [], + ) + .await + .unwrap(); + assert_eq!(filled, 0); + } + + #[tokio::test] + pub async fn test_get_object_into_chunks_not_enabled() { + let mock_client = create_test_s3_client(mock_client!(aws_sdk_s3, []), None, false); + let mut chunks = test_chunks(1); + let result = mock_client + .get_object_into_chunks( + &GetObjectTarget { + key: "test_object", + etag: "test-etag", + version_id: "", + }, + 0, + 16, + &mut chunks, + ) + .await; + assert!(matches!(result.unwrap_err(), S3ClientError::NotEnabled)); + } + + #[tokio::test] + pub async fn test_get_object_into_chunks_etag_mismatch() { + // The same extractor serves both GET paths: a 412 must surface as + // ETagMismatchError here exactly as it does from get_object_if_match. + let get_object_rule = mock!(aws_sdk_s3::Client::get_object).then_http_response(|| { + HttpResponse::new( + StatusCode::try_from(412).unwrap(), + SdkBody::from("Precondition Failed"), + ) + }); + let mock_client = + create_test_s3_client(mock_client!(aws_sdk_s3, [&get_object_rule]), None, true); + + let mut chunks = test_chunks(1); + let result = mock_client + .get_object_into_chunks( + &GetObjectTarget { + key: "test_object", + etag: "wrong-etag", + version_id: "", + }, + 0, + 16, + &mut chunks, + ) + .await; + assert!(matches!( + result.unwrap_err(), + S3ClientError::ETagMismatchError + )); + } + #[tokio::test] pub async fn test_permission_validator_task() { // First call fails, second call succeeds @@ -1143,7 +1539,10 @@ mod tests { .await; assert!(result.is_err()); - assert!(matches!(result.unwrap_err(), S3ClientError::NoAccess(_))); + // The mock's error carries a dummy raw response (no real status), + // so the wire failure classifies as Unknown; what this test pins + // is that a mid-sequence chunk failure fails the whole request. + assert!(matches!(result.unwrap_err(), S3ClientError::Unknown)); } #[tokio::test] @@ -1263,18 +1662,18 @@ mod tests { // --- #2: Metric emission tests --- struct MockCloudWatchClient { - calls: Arc>>, - log_calls: Arc>>, + calls: Arc>>, + log_calls: Arc>>, } impl MockCloudWatchClient { fn new() -> ( Self, - Arc>>, - Arc>>, + Arc>>, + Arc>>, ) { - let calls = Arc::new(std::sync::Mutex::new(Vec::new())); - let log_calls = Arc::new(std::sync::Mutex::new(Vec::new())); + let calls = Arc::new(crate::sync::Mutex::new(Vec::new())); + let log_calls = Arc::new(crate::sync::Mutex::new(Vec::new())); ( Self { calls: calls.clone(), @@ -1617,4 +2016,55 @@ mod tests { tokio::time::resume(); } + + // ── Wire failure classification ── + + #[test_case(403, S3ClientError::AccessDenied; "403 access denied")] + #[test_case(404, S3ClientError::NoSuchKey; "404 no such key")] + #[test_case(412, S3ClientError::ETagMismatchError; "412 precondition")] + #[test_case(429, S3ClientError::Throttled; "429 throttled")] + #[test_case(500, S3ClientError::Throttled; "500 throttled")] + #[test_case(503, S3ClientError::Throttled; "503 throttled")] + #[test_case(418, S3ClientError::Unknown; "unmapped status is unknown")] + fn from_wire_failure_classifies_http_status(status: u16, expect: S3ClientError) { + // A real SdkError carrying the given HTTP status (the + // response-error shape has a raw response but no parsed service + // error — exactly what classification must handle). + let raw = HttpResponse::new( + StatusCode::try_from(status).unwrap(), + SdkBody::from("test body"), + ); + let sdk: SdkError = SdkError::response_error("test".to_string(), raw); + assert_eq!(S3ClientError::from_wire_failure(&sdk), expect); + } + + #[test] + fn from_wire_failure_classifies_sdk_timeout() { + let timeout: SdkError = + SdkError::timeout_error("request timed out".to_string()); + assert_eq!( + S3ClientError::from_wire_failure(&timeout), + S3ClientError::Timeout + ); + } + + #[test] + fn from_wire_failure_classifies_s3_request_timeout_code() { + // S3's slow-client timeout arrives as HTTP 400 with error code + // RequestTimeout; the parsed code must win over the (unmapped) + // raw status, or a retryable timeout lands in Unknown. + let meta = aws_sdk_s3::error::ErrorMetadata::builder() + .code("RequestTimeout") + .message("Your socket connection to the server was not read from or written to within the timeout period.") + .build(); + let raw = HttpResponse::new( + StatusCode::try_from(400).unwrap(), + SdkBody::from("test body"), + ); + let sdk = SdkError::service_error(GetObjectError::generic(meta), raw); + assert_eq!( + S3ClientError::from_wire_failure(&sdk), + S3ClientError::Timeout + ); + } } diff --git a/src/client-core/src/config_parser.rs b/src/client-core/src/config_parser.rs index 64975914..471fa7b2 100644 --- a/src/client-core/src/config_parser.rs +++ b/src/client-core/src/config_parser.rs @@ -1,6 +1,8 @@ use log::LevelFilter; use serde::{Deserialize, Serialize}; -use std::{error::Error, path::Path, str::FromStr}; +use std::error::Error; +use std::path::Path; +use std::str::FromStr; const DEFAULT_LOG_LEVEL: fn() -> String = || LevelFilter::Warn.to_string(); @@ -421,7 +423,8 @@ pub mod tests { use super::*; use crate::test_utils::TEST_CONFIG_PATH; use rand::random; - use std::{path::Path, string::String}; + use std::path::Path; + use std::string::String; #[test] fn test_read_config_from_file() { diff --git a/src/client-core/src/lib.rs b/src/client-core/src/lib.rs index 5eba8ddb..edb4612d 100644 --- a/src/client-core/src/lib.rs +++ b/src/client-core/src/lib.rs @@ -13,6 +13,7 @@ pub mod error; pub mod memory; pub mod proxy_identifier; pub mod read_ahead; +pub mod sync; pub mod util; pub mod utils; diff --git a/src/client-core/src/memory/memory_pool.rs b/src/client-core/src/memory/memory_pool.rs index 43f7211b..da6f9996 100644 --- a/src/client-core/src/memory/memory_pool.rs +++ b/src/client-core/src/memory/memory_pool.rs @@ -1,10 +1,10 @@ #![allow(unused)] +use crate::sync::atomic::{AtomicUsize, Ordering}; +use crate::sync::{Arc, Mutex, MutexGuard, Weak}; use log::debug; use std::mem::{ManuallyDrop, MaybeUninit}; use std::ops::{Deref, DerefMut}; -use std::sync::atomic::{AtomicUsize, Ordering}; -use std::sync::{Arc, Mutex, MutexGuard, Weak}; use std::time::{Duration, Instant}; const ONE_MB_IN_BYTES: usize = 1024 * 1024; @@ -470,8 +470,7 @@ impl MemoryPool { #[cfg(test)] mod tests { use super::*; - use std::sync::{Arc, Barrier}; - use std::thread; + use crate::sync::{thread, Arc, Barrier}; use std::time::Duration; #[test] diff --git a/src/client-core/src/read_ahead/cached_data.rs b/src/client-core/src/read_ahead/cached_data.rs index a5b3adf0..6d458944 100644 --- a/src/client-core/src/read_ahead/cached_data.rs +++ b/src/client-core/src/read_ahead/cached_data.rs @@ -5,12 +5,12 @@ //! while the data is initialized. This allows for concurrent reads of the same data, //! improving performance for read-heavy workloads. +use crate::sync::atomic::{AtomicU64, Ordering}; +use crate::sync::Arc; use atomic_enum::atomic_enum; use bytes::Bytes; use log::{error, warn}; use std::ops::Range; -use std::sync::atomic::{AtomicU64, Ordering}; -use std::sync::Arc; use std::time::{SystemTime, UNIX_EPOCH}; use tokio::sync::{Notify, OwnedRwLockReadGuard, RwLock, RwLockWriteGuard}; @@ -159,9 +159,9 @@ impl CachedData { { // If there's already data, don't overwrite it if !data_guard.is_empty() { - let error = ReadAheadCacheError { - message: "Cannot load data over existing cached data".to_string(), - }; + let error = ReadAheadCacheError::Other( + "Cannot load data over existing cached data".to_string(), + ); error!("{}", error); return Err(error); } @@ -191,9 +191,9 @@ impl CachedData { // Entry is already removed from cache index; memory chunks // will be returned to pool when outstanding Bytes references drop. warn!("Timeout acquiring write lock in clear()"); - return Err(ReadAheadCacheError { - message: "Timeout acquiring write lock in clear()".to_string(), - }); + return Err(ReadAheadCacheError::Other( + "Timeout acquiring write lock in clear()".to_string(), + )); } }; @@ -214,9 +214,8 @@ impl CachedData { last_read_end_position: u64, ) -> Result<(), ReadAheadCacheError> { if last_read_end_position == 0 { - let error = ReadAheadCacheError { - message: "Read position must be greater than 0".to_string(), - }; + let error = + ReadAheadCacheError::Other("Read position must be greater than 0".to_string()); error!("{}", error); return Err(error); } @@ -253,22 +252,18 @@ impl CachedData { match state { CacheEntryState::Loading => { - return Err(ReadAheadCacheError { - message: format!( - "Timeout waiting for cache entry to load (start: {} end: {}) within {}s", - range.start, range.end, DEFAULT_TIME_OUT_SECOND - ), - }); + return Err(ReadAheadCacheError::Other(format!( + "Timeout waiting for cache entry to load (start: {} end: {}) within {}s", + range.start, range.end, DEFAULT_TIME_OUT_SECOND + ))); } CacheEntryState::Failed => { - return Err(ReadAheadCacheError { - message: "Cache entry failed to load".to_string(), - }); + return Err(ReadAheadCacheError::Other( + "Cache entry failed to load".to_string(), + )); } CacheEntryState::Evicted => { - return Err(ReadAheadCacheError { - message: "Data evicted before read completed".to_string(), - }); + return Err(ReadAheadCacheError::DataEvicted); } CacheEntryState::Loaded => {} } @@ -286,23 +281,19 @@ impl CachedData { let guard = self.data.clone().read_owned().await; // Re-check state after acquiring lock - eviction sets state before clearing data if self.get_state() != CacheEntryState::Loaded { - return Err(ReadAheadCacheError { - message: "Data evicted before read completed".to_string(), - }); + return Err(ReadAheadCacheError::DataEvicted); } if guard.is_empty() { - return Err(ReadAheadCacheError { - message: "Cache entry has no data despite Loaded state".to_string(), - }); + return Err(ReadAheadCacheError::Other( + "Cache entry has no data despite Loaded state".to_string(), + )); } if start_chunk >= guard.len() { - return Err(ReadAheadCacheError { - message: format!( - "Start chunk index ({}) is beyond available data length ({})", - start_chunk, - guard.len() - ), - }); + return Err(ReadAheadCacheError::Other(format!( + "Start chunk index ({}) is beyond available data length ({})", + start_chunk, + guard.len() + ))); } let slice = ChunkSlice { guard, @@ -317,15 +308,13 @@ impl CachedData { let data_guard = self.acquire_read_lock().await?; // Re-check state after acquiring lock - eviction sets state before clearing data if self.get_state() != CacheEntryState::Loaded { - return Err(ReadAheadCacheError { - message: "Data evicted before read completed".to_string(), - }); + return Err(ReadAheadCacheError::DataEvicted); } if data_guard.is_empty() { // This shouldn't happen if state is Loaded, but handle defensively - return Err(ReadAheadCacheError { - message: "Cache entry has no data despite Loaded state".to_string(), - }); + return Err(ReadAheadCacheError::Other( + "Cache entry has no data despite Loaded state".to_string(), + )); } self.copy_data_from_chunks(&data_guard, offset, length) @@ -344,13 +333,11 @@ impl CachedData { let start_byte_in_chunk = offset_usize % memory_pool::CHUNK_SIZE; if start_chunk_idx >= data_guard.len() { - return Err(ReadAheadCacheError { - message: format!( - "Start chunk index ({}) is beyond available data length ({})", - start_chunk_idx, - data_guard.len() - ), - }); + return Err(ReadAheadCacheError::Other(format!( + "Start chunk index ({}) is beyond available data length ({})", + start_chunk_idx, + data_guard.len() + ))); } let mut result_data = Vec::with_capacity(length as usize); @@ -366,13 +353,11 @@ impl CachedData { }; if start_pos >= memory_pool::CHUNK_SIZE { - let error = ReadAheadCacheError { - message: format!( - "Invalid start position ({}) exceeds chunk size ({})", - start_pos, - memory_pool::CHUNK_SIZE - ), - }; + let error = ReadAheadCacheError::Other(format!( + "Invalid start position ({}) exceeds chunk size ({})", + start_pos, + memory_pool::CHUNK_SIZE + )); error!("{}", error); return Err(error); } @@ -383,10 +368,10 @@ impl CachedData { let bytes_available_in_chunk = memory_pool::CHUNK_SIZE - start_pos; let bytes_to_copy = bytes_remaining_usize.min(bytes_available_in_chunk); if start_pos + bytes_to_copy > memory_pool::CHUNK_SIZE { - let error = ReadAheadCacheError { - message: format!("Would read beyond chunk boundary: start_pos ({}) + bytes_to_copy ({}) > chunk_size ({})", - start_pos, bytes_to_copy, memory_pool::CHUNK_SIZE) - }; + let error = ReadAheadCacheError::Other(format!( + "Would read beyond chunk boundary: start_pos ({}) + bytes_to_copy ({}) > chunk_size ({})", + start_pos, bytes_to_copy, memory_pool::CHUNK_SIZE + )); error!("{}", error); return Err(error); } @@ -398,12 +383,10 @@ impl CachedData { // If we couldn't get all the requested data fail the request if bytes_copied < length { - let error = ReadAheadCacheError { - message: format!( - "Could not satisfy entire read request: bytes_copied ({}) < length ({})", - bytes_copied, length - ), - }; + let error = ReadAheadCacheError::Other(format!( + "Could not satisfy entire read request: bytes_copied ({}) < length ({})", + bytes_copied, length + )); warn!("{}", error); return Err(error); } @@ -415,12 +398,10 @@ impl CachedData { /// Returns the offset and length if valid, or an error if invalid fn validate_range(&self, range: Range) -> Result<(u64, u64), ReadAheadCacheError> { if range.start >= range.end { - let error = ReadAheadCacheError { - message: format!( - "Invalid range: start ({}) must be less than end ({})", - range.start, range.end - ), - }; + let error = ReadAheadCacheError::Other(format!( + "Invalid range: start ({}) must be less than end ({})", + range.start, range.end + )); error!("{}", error); return Err(error); } @@ -428,13 +409,11 @@ impl CachedData { let offset = range.start; let length = range.end - range.start; if length > usize::MAX as u64 { - let error = ReadAheadCacheError { - message: format!( - "Length ({}) too large for this system (max: {})", - length, - usize::MAX - ), - }; + let error = ReadAheadCacheError::Other(format!( + "Length ({}) too large for this system (max: {})", + length, + usize::MAX + )); error!("{}", error); return Err(error); } @@ -453,7 +432,7 @@ impl CachedData { mod tests { use super::*; use crate::memory::memory_pool::MemoryPoolConfig; - use std::sync::Arc; + use crate::sync::Arc; fn create_test_chunk(data: &[u8]) -> MemoryChunk { let memory_pool = create_test_memory_pool(); @@ -560,9 +539,8 @@ mod tests { let err = cached_data.get_data_range(0..5).await.unwrap_err(); assert!( - err.message.contains("Data evicted"), - "Error should trigger retry: {}", - err.message + matches!(err, ReadAheadCacheError::DataEvicted), + "Error should trigger retry: {err}" ); } diff --git a/src/client-core/src/read_ahead/error.rs b/src/client-core/src/read_ahead/error.rs index e21fdf67..84ef217d 100644 --- a/src/client-core/src/read_ahead/error.rs +++ b/src/client-core/src/read_ahead/error.rs @@ -1,7 +1,81 @@ use thiserror::Error as ThisError; +use crate::aws::s3_client::S3ClientError; + +#[derive(Debug, ThisError)] +pub enum ReadAheadCacheError { + /// The cache entry went away before the read could copy out of it. Retryable. + #[error("Data evicted before read completed")] + DataEvicted, + /// The memory pool has no room for the chunks this read needs. + #[error("Memory pool at capacity")] + MemoryPoolAtCapacity, + /// Failure that has no dedicated variant yet. + #[error("{0}")] + Other(String), +} + +/// Failure returned by the readahead cache's read path. `s3_error` and `cache_error` let the +/// caller pick a retry policy without parsing `message`. #[derive(Debug, ThisError)] -#[error("ReadAheadCache error: {message}")] -pub struct ReadAheadCacheError { +#[error("ReadAhead error: {message}")] +pub struct ReadAheadError { + /// Human-readable description of the failure, for callers that only log it. pub message: String, + pub s3_error: Option, + pub cache_error: Option, +} + +impl From for ReadAheadError { + fn from(error: ReadAheadCacheError) -> Self { + Self { + message: error.to_string(), + s3_error: None, + cache_error: Some(error), + } + } +} + +impl From for ReadAheadError { + fn from(error: S3ClientError) -> Self { + Self { + message: format!("S3 read failed: {}", error), + s3_error: Some(error), + cache_error: None, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_cache_error_conversion_keeps_variant() { + let error: ReadAheadError = ReadAheadCacheError::DataEvicted.into(); + assert!(matches!( + error.cache_error, + Some(ReadAheadCacheError::DataEvicted) + )); + assert!(error.s3_error.is_none()); + assert_eq!(error.message, "Data evicted before read completed"); + } + + #[test] + fn test_other_cache_error_conversion() { + let error: ReadAheadError = + ReadAheadCacheError::Other("Cache entry failed to load".to_string()).into(); + assert!(matches!( + error.cache_error, + Some(ReadAheadCacheError::Other(_)) + )); + assert_eq!(error.message, "Cache entry failed to load"); + } + + #[test] + fn test_s3_error_conversion_keeps_error_class() { + let error: ReadAheadError = S3ClientError::Throttled.into(); + assert!(matches!(error.s3_error, Some(S3ClientError::Throttled))); + assert!(error.cache_error.is_none()); + } } diff --git a/src/client-core/src/read_ahead/file_readahead_state.rs b/src/client-core/src/read_ahead/file_readahead_state.rs index 0da66710..cc4ee84b 100644 --- a/src/client-core/src/read_ahead/file_readahead_state.rs +++ b/src/client-core/src/read_ahead/file_readahead_state.rs @@ -6,25 +6,26 @@ //! mutex lock. //! +use crate::sync::atomic::{AtomicU64, Ordering}; +use crate::sync::{Arc, Mutex, Weak}; use bytes::{Bytes, BytesMut}; use log::{error, trace}; use lru::LruCache; use std::collections::BTreeMap; use std::ops::Range; -use std::sync::atomic::{AtomicU64, Ordering}; -use std::sync::{Arc, Mutex, Weak}; use tokio::sync::{RwLock, RwLockReadGuard, RwLockWriteGuard}; +use crate::aws::s3_client::S3ClientError; use crate::memory::memory_pool::{self, MemoryChunk, MemoryPool}; use crate::nfs::nfs4_1_xdr::awsfile_bypass_data_locator; use crate::read_ahead::cached_data::{ get_current_time_ms, CacheEntryState, CachedData, INVALID_U64, }; -use crate::read_ahead::error::ReadAheadCacheError; +use crate::read_ahead::error::{ReadAheadCacheError, ReadAheadError}; use crate::read_ahead::readahead_cache::FileReadAheadCache; use crate::util::read_bypass_request_context::ReadBypassRequestContext; use crate::util::s3_data_reader::S3DataReader; -use crate::{ctx_debug, ctx_error, ctx_warn}; +use crate::{ctx_debug, ctx_error, ctx_trace, ctx_warn}; /// Minimum time (ms) after loading before an entry can be evicted, protects // recently loaded items from being immediately evictedx @@ -173,21 +174,17 @@ impl FileReadAheadState { s3_data_locator: &awsfile_bypass_data_locator, ) -> Result<(), ReadAheadCacheError> { if s3_data_locator.s3_key != self.s3_key { - return Err(ReadAheadCacheError { - message: format!( - "S3 key mismatch: expected {:?}, got {:?}", - self.s3_key, s3_data_locator.s3_key - ), - }); + return Err(ReadAheadCacheError::Other(format!( + "S3 key mismatch: expected {:?}, got {:?}", + self.s3_key, s3_data_locator.s3_key + ))); } if s3_data_locator.etag != self.s3_etag { - return Err(ReadAheadCacheError { - message: format!( - "S3 etag mismatch: expected {:?}, got {:?}", - self.s3_etag, s3_data_locator.etag - ), - }); + return Err(ReadAheadCacheError::Other(format!( + "S3 etag mismatch: expected {:?}, got {:?}", + self.s3_etag, s3_data_locator.etag + ))); } Ok(()) @@ -203,7 +200,7 @@ impl FileReadAheadState { file_size: u64, s3_data_reader: Arc, suppress_readahead: bool, - ) -> Result<(Option, bool), ReadAheadCacheError> { + ) -> Result<(Option, bool), ReadAheadError> { // === STEP 1: Validate range === // Clamp to file size - NFS allows reads past EOF but we return data up to EOF only let end = (s3_data_locator.offset + s3_data_locator.count as u64).min(file_size); @@ -314,7 +311,7 @@ impl FileReadAheadState { s3_data_reader: Arc, range: Range, read_pattern: ReadPattern, - ) -> Result<(Option, bool), ReadAheadCacheError> { + ) -> Result<(Option, bool), ReadAheadError> { self.validate_request(s3_data_locator)?; // === STEP 1: Atomic planning @@ -351,9 +348,11 @@ impl FileReadAheadState { for prev_range in &missing_ranges[..i] { data_cache_write.remove(&prev_range.start); } - return Err(ReadAheadCacheError { - message: format!("Failed to insert missing range: {:?}", missing_range), - }); + return Err(ReadAheadCacheError::Other(format!( + "Failed to insert missing range: {:?}", + missing_range + )) + .into()); } } @@ -382,7 +381,7 @@ impl FileReadAheadState { self.remove_failed_entry(missing_range.start, CacheEntryState::Failed) .await; } - return Err(e); + return Err(e.into()); } }; @@ -611,12 +610,10 @@ impl FileReadAheadState { data_cache: &mut RwLockWriteGuard<'_, BTreeMap)>>, ) -> Result { if range.start >= range.end { - let error = ReadAheadCacheError { - message: format!( - "Invalid range: start ({}) must be less than end ({})", - range.start, range.end - ), - }; + let error = ReadAheadCacheError::Other(format!( + "Invalid range: start ({}) must be less than end ({})", + range.start, range.end + )); error!("{}", error); return Err(error); } @@ -643,12 +640,10 @@ impl FileReadAheadState { T: std::ops::Deref)>>, { if requested_range.start >= requested_range.end { - return Err(ReadAheadCacheError { - message: format!( - "Invalid range: start ({}) must be less than end ({})", - requested_range.start, requested_range.end - ), - }); + return Err(ReadAheadCacheError::Other(format!( + "Invalid range: start ({}) must be less than end ({})", + requested_range.start, requested_range.end + ))); } let mut ranges_to_fetch: Vec<(Range, u64, u64, Arc)> = Vec::new(); @@ -740,9 +735,10 @@ impl FileReadAheadState { expected: Range, ) -> Result { if ranges.is_empty() { - return Err(ReadAheadCacheError { - message: format!("No data for range {}..{}", expected.start, expected.end), - }); + return Err(ReadAheadCacheError::Other(format!( + "No data for range {}..{}", + expected.start, expected.end + ))); } // Fast path: a single range fully covering the request is served as a zero-copy view @@ -770,20 +766,16 @@ impl FileReadAheadState { continue; // Range entirely before current position } if range.start < pos && pos != expected.start { - return Err(ReadAheadCacheError { - message: format!( + return Err(ReadAheadCacheError::Other(format!( "Overlapping ranges detected: current position {} but range starts at {} (range end: {})", pos, range.start, range.end - ), - }); + ))); } if range.start > pos { - return Err(ReadAheadCacheError { - message: format!( - "Gap in data at offset {}, next range starts at {} (expected end: {})", - pos, range.start, expected.end - ), - }); + return Err(ReadAheadCacheError::Other(format!( + "Gap in data at offset {}, next range starts at {} (expected end: {})", + pos, range.start, expected.end + ))); } pos = range.end; if pos >= expected.end { @@ -791,9 +783,10 @@ impl FileReadAheadState { } } if pos < expected.end { - return Err(ReadAheadCacheError { - message: format!("Data ends at {} but expected {}", pos, expected.end), - }); + return Err(ReadAheadCacheError::Other(format!( + "Data ends at {} but expected {}", + pos, expected.end + ))); } // Extract just the expected range @@ -813,16 +806,14 @@ impl FileReadAheadState { // Final sanity check - result must be exactly the expected size if result.len() != expected_len { - return Err(ReadAheadCacheError { - message: format!( - "Result size mismatch for range {}..{}: got {} bytes but expected {} \ + return Err(ReadAheadCacheError::Other(format!( + "Result size mismatch for range {}..{}: got {} bytes but expected {} \ (possible overlap or gap in cached/fetched ranges)", - expected.start, - expected.end, - result.len(), - expected_len - ), - }); + expected.start, + expected.end, + result.len(), + expected_len + ))); } Ok(result.freeze()) @@ -905,7 +896,7 @@ impl FileReadAheadState { original_request_range: Range, s3_data_reader: &dyn S3DataReader, s3_data_locator: &awsfile_bypass_data_locator, - ) -> Result<(Vec<(Range, Bytes)>, bool), ReadAheadCacheError> { + ) -> Result<(Vec<(Range, Bytes)>, bool), ReadAheadError> { ctx_debug!( read_bypass_request_context, "Fetching {} missing ranges from S3", @@ -939,6 +930,8 @@ impl FileReadAheadState { } let mut required_failed = false; + // Keep the first S3 error that failed a required range. + let mut required_s3_error: Option = None; let mut any_required_cache_failed = false; let mut all_required_data: Vec<(Range, Bytes)> = Vec::new(); @@ -962,23 +955,39 @@ impl FileReadAheadState { ctx_warn!( read_bypass_request_context, "Caching failed ({}), will return data directly", - e.message + e ); any_required_cache_failed = true; } } } Ok(Err(s3_error)) => { - // S3 failure - use Failed so file gets denylisted + // S3 failure - drop the entry so waiters stop blocking on it. self.remove_failed_entry(missing_range.start, CacheEntryState::Failed) .await; if is_required { - ctx_error!( + // Throttling is reported once by the caller that decides to retry it. + if matches!(s3_error, S3ClientError::Throttled) { + ctx_trace!( + read_bypass_request_context, + "Required S3 read throttled: {:?}", + s3_error + ); + } else { + ctx_error!( + read_bypass_request_context, + "Required S3 read failed: {:?}", + s3_error + ); + } + required_failed = true; + required_s3_error.get_or_insert(s3_error); + } else if matches!(s3_error, S3ClientError::Throttled) { + ctx_trace!( read_bypass_request_context, - "Required S3 read failed: {:?}", + "Readahead S3 read throttled (non-critical): {:?}", s3_error ); - required_failed = true; } else { ctx_warn!( read_bypass_request_context, @@ -1010,8 +1019,9 @@ impl FileReadAheadState { } if required_failed { - return Err(ReadAheadCacheError { - message: "Required S3 read failed".to_string(), + return Err(match required_s3_error { + Some(s3_error) => s3_error.into(), + None => ReadAheadCacheError::Other("Required S3 read failed".to_string()).into(), }); } @@ -1056,8 +1066,8 @@ impl FileReadAheadState { .map(|(_, cached_data)| cached_data.clone()) }; - let cached_data = cached_data.ok_or_else(|| ReadAheadCacheError { - message: format!("Cached range not found for offset {}", range.start), + let cached_data = cached_data.ok_or_else(|| { + ReadAheadCacheError::Other(format!("Cached range not found for offset {}", range.start)) })?; let chunks = self.prepare_chunks_from_bytes(range, s3_data).await?; @@ -1073,13 +1083,11 @@ impl FileReadAheadState { let size = range.end - range.start; if size != s3_data.len() as u64 { - return Err(ReadAheadCacheError { - message: format!( - "Size mismatch: requested range size ({}) != bytes_to_write length ({})", - size, - s3_data.len() - ), - }); + return Err(ReadAheadCacheError::Other(format!( + "Size mismatch: requested range size ({}) != bytes_to_write length ({})", + size, + s3_data.len() + ))); } let chunk_size_u64 = memory_pool::CHUNK_SIZE as u64; @@ -1094,18 +1102,14 @@ impl FileReadAheadState { // If still over capacity, fail and let caller handle gracefully if self.memory_pool.would_exceed_capacity(num_chunks) { - return Err(ReadAheadCacheError { - message: "Memory pool at capacity".to_string(), - }); + return Err(ReadAheadCacheError::MemoryPoolAtCapacity); } let mut chunks = self.memory_pool.consume(num_chunks); // Race condition: another thread may have allocated between our check and consume if chunks.len() < num_chunks { - return Err(ReadAheadCacheError { - message: "Memory pool at capacity".to_string(), - }); + return Err(ReadAheadCacheError::MemoryPoolAtCapacity); } // Copy data into chunks @@ -1315,8 +1319,8 @@ mod tests { use super::*; use crate::config_parser::ProxyConfig; use crate::memory::memory_pool::{MemoryPoolConfig, CHUNK_SIZE}; + use crate::sync::Arc; use crate::util::read_bypass_context::ReadBypassContext; - use std::sync::Arc; // --- combine_ranges zero-copy fast path --- @@ -1474,10 +1478,10 @@ mod tests { } async fn create_test_read_bypass_context_with_size(size: usize) -> Arc { + use crate::sync::Arc; use aws_sdk_s3::operation::get_object::GetObjectOutput; use aws_sdk_s3::primitives::ByteStream; use aws_smithy_mocks::{mock, mock_client}; - use std::sync::Arc; let get_object_rule = mock!(aws_sdk_s3::Client::get_object) .match_requests(|_req| true) @@ -1625,7 +1629,9 @@ mod tests { assert_eq!(pattern, ReadPattern::SequentialRead); // Large backward read beyond window_size should be RandomRead - let window_size = state.window_size.load(std::sync::atomic::Ordering::SeqCst); + let window_size = state + .window_size + .load(crate::sync::atomic::Ordering::SeqCst); state.update_read_position(window_size + 1000); // Use offset 10 (not 0) to avoid StartOfFile pattern let pattern = recognize_read_pattern_with_lock(&state, 10..20).await; @@ -2696,7 +2702,7 @@ mod tests { .await; match result { - Err(e) => assert!(e.message.contains("capacity")), + Err(e) => assert!(matches!(e, ReadAheadCacheError::MemoryPoolAtCapacity)), Ok(_) => panic!("Expected error due to capacity"), } } diff --git a/src/client-core/src/read_ahead/readahead_cache.rs b/src/client-core/src/read_ahead/readahead_cache.rs index 0feab023..0bef78f6 100644 --- a/src/client-core/src/read_ahead/readahead_cache.rs +++ b/src/client-core/src/read_ahead/readahead_cache.rs @@ -8,12 +8,12 @@ //! This file contains the top-level FileReadAheadCache which manages the collection of file states //! and provides the main interface for the readahead system. +use crate::sync::atomic::{AtomicU64, Ordering}; +use crate::sync::{Arc, Mutex, Weak}; use bytes::Bytes; use dashmap::DashMap; use lru::LruCache; use std::num::NonZeroUsize; -use std::sync::atomic::{AtomicU64, Ordering}; -use std::sync::{Arc, Mutex, Weak}; use tokio::sync::RwLock; use crate::config_parser::ReadBypassConfig; @@ -21,7 +21,7 @@ use crate::ctx_debug; use crate::memory::memory_pool::{MemoryPool, MemoryPoolConfig}; use crate::nfs::nfs4_1_xdr::{awsfile_bypass_data_locator, nfs_fh4}; use crate::read_ahead::cached_data::get_current_time_ms; -use crate::read_ahead::error::ReadAheadCacheError; +use crate::read_ahead::error::{ReadAheadCacheError, ReadAheadError}; use crate::read_ahead::file_readahead_state::{FileReadAheadState, LruKey, LruValue}; use crate::util::read_bypass_request_context::ReadBypassRequestContext; use crate::util::s3_data_reader::S3DataReader; @@ -114,7 +114,7 @@ impl FileReadAheadCache { filehandle: nfs_fh4, file_size: u64, s3_data_locator: awsfile_bypass_data_locator, - ) -> Result, ReadAheadCacheError> { + ) -> Result, ReadAheadError> { // Skip cache for small files that fit within a single S3 chunk read if file_size <= self.small_file_caching_threshold && s3_data_locator.count as u64 == file_size @@ -171,7 +171,7 @@ impl FileReadAheadCache { &self, read_bypass_request_context: Arc, s3_data_locator: awsfile_bypass_data_locator, - ) -> Result, ReadAheadCacheError> { + ) -> Result, ReadAheadError> { let read_task = self .s3_data_reader .spawn_read_task( @@ -179,14 +179,10 @@ impl FileReadAheadCache { read_bypass_request_context.read_bypass_context.clone(), ) .await; - let data = read_task - .await - .map_err(|e| ReadAheadCacheError { - message: format!("Direct S3 read join error: {}", e), - })? - .map_err(|e| ReadAheadCacheError { - message: format!("Direct S3 read error: {:?}", e), - })?; + // The inner `?` keeps the S3 error class in `ReadAheadError::s3_error`. + let data = read_task.await.map_err(|e| { + ReadAheadCacheError::Other(format!("Direct S3 read join error: {}", e)) + })??; Ok(Some(data)) } @@ -343,11 +339,23 @@ impl FileReadAheadCache { let mut stale_removed = 0; let mut empty_files = Vec::new(); + // Snapshot the entries before awaiting. Holding the DashMap shard + // iterator across an .await is a lock-held-across-await hazard: the + // shard read lock blocks all writers to that shard for the duration + // of the awaited cleanup, and if a task on the same worker thread + // mutates the map at the suspension point it can deadlock (see the + // "Locking behaviour" notes on DashMap's mutating methods). + let entries: Vec<_> = self + .file_states + .iter() + .map(|entry| (entry.key().clone(), Arc::clone(entry.value()))) + .collect(); + // First pass: cleanup stale entries and identify empty files - for entry in self.file_states.iter() { - stale_removed += entry.value().cleanup_stale_entries(IDLE_TTL_MS).await; - if entry.value().is_empty() { - empty_files.push(entry.key().clone()); + for (key, state) in entries { + stale_removed += state.cleanup_stale_entries(IDLE_TTL_MS).await; + if state.is_empty() { + empty_files.push(key); } } @@ -414,13 +422,15 @@ impl FileReadAheadCache { #[cfg(test)] mod tests { use super::*; + use crate::aws::s3_client::S3ClientError; use crate::nfs::nfs4_1_xdr::nfs_fh4; + use crate::sync::atomic::Ordering; use crate::test_utils::{ create_test_read_bypass_context, create_test_s3_data_locator, CountingS3DataReader, + FailingS3DataReader, }; use crate::util::read_bypass_request_context::ReadBypassRequestContext; use crate::util::s3_data_reader::S3ReadBypassReader; - use std::sync::atomic::Ordering; use test_case::test_case; #[tokio::test] @@ -995,6 +1005,76 @@ mod tests { ); } + async fn read_with_failing_s3( + make_error: fn() -> S3ClientError, + file_size: u64, + count: u32, + ) -> ReadAheadError { + let cache = Arc::new(FileReadAheadCache::new( + 64 * 1024, + 64 * 1024, + Arc::new(FailingS3DataReader { make_error }), + &ReadBypassConfig::default(), + )); + cache.set_self_weak(Arc::downgrade(&cache)); + + let read_bypass_context = + Arc::new(crate::util::read_bypass_context::ReadBypassContext::default().await); + let ctx = Arc::new( + crate::util::read_bypass_request_context::ReadBypassRequestContext::new( + read_bypass_context, + 0, + ), + ); + + cache + .process_read_request( + ctx, + crate::nfs::nfs4_1_xdr::nfs_fh4(b"failing_fh".to_vec()), + file_size, + create_test_locator(0, count), + ) + .await + .expect_err("Read should fail when S3 fails") + } + + /// Small-file path (direct read, no cache state): the S3 error class must reach the caller. + #[tokio::test] + async fn test_direct_read_propagates_s3_error_class() { + // 512B <= default small_file_caching_threshold (1 MiB), so this takes the direct path. + let error = read_with_failing_s3(|| S3ClientError::Throttled, 512, 512).await; + assert!( + matches!(error.s3_error, Some(S3ClientError::Throttled)), + "Direct read should surface the S3 error class, got {error:?}" + ); + + let error = read_with_failing_s3(|| S3ClientError::AccessDenied, 512, 512).await; + assert!( + matches!(error.s3_error, Some(S3ClientError::AccessDenied)), + "Direct read should not rewrite the S3 error class, got {error:?}" + ); + } + + /// Cached path (the default production path): the S3 error class must reach the caller too, + /// otherwise a throttled GET is indistinguishable from a permanent cache failure and the + /// read bypass agent denylists the file handle for the full denylist TTL. + #[tokio::test] + async fn test_cached_read_propagates_s3_error_class() { + // 2 MiB > default small_file_caching_threshold (1 MiB), so this goes through the cache. + let file_size = 2 * 1024 * 1024u64; + let error = read_with_failing_s3(|| S3ClientError::Throttled, file_size, 8192).await; + assert!( + matches!(error.s3_error, Some(S3ClientError::Throttled)), + "Cached read should surface the S3 error class, got {error:?}" + ); + + let error = read_with_failing_s3(|| S3ClientError::AccessDenied, file_size, 8192).await; + assert!( + matches!(error.s3_error, Some(S3ClientError::AccessDenied)), + "Cached read should not rewrite the S3 error class, got {error:?}" + ); + } + /// Stress test for concurrent reads with cache eviction pressure /// to verify no race conditions between eviction and cache operations. #[test_case(false, false ; "single_file_random")] diff --git a/src/client-core/src/sync.rs b/src/client-core/src/sync.rs new file mode 100644 index 00000000..f64f2957 --- /dev/null +++ b/src/client-core/src/sync.rs @@ -0,0 +1,7 @@ +//! Re-exports of std synchronization and thread primitives. + +pub use std::sync::*; + +pub mod thread { + pub use std::thread::*; +} diff --git a/src/client-core/src/test_utils.rs b/src/client-core/src/test_utils.rs index 8bbe8b6f..25b4881d 100644 --- a/src/client-core/src/test_utils.rs +++ b/src/client-core/src/test_utils.rs @@ -21,7 +21,7 @@ pub fn get_test_config() -> ProxyConfig { /// Used by readahead cache and file readahead state tests. #[derive(Clone)] pub struct CountingS3DataReader { - pub call_count: std::sync::Arc, + pub call_count: crate::sync::Arc, } impl Default for CountingS3DataReader { @@ -33,12 +33,12 @@ impl Default for CountingS3DataReader { impl CountingS3DataReader { pub fn new() -> Self { Self { - call_count: std::sync::Arc::new(std::sync::atomic::AtomicU64::new(0)), + call_count: crate::sync::Arc::new(crate::sync::atomic::AtomicU64::new(0)), } } pub fn calls(&self) -> u64 { - self.call_count.load(std::sync::atomic::Ordering::SeqCst) + self.call_count.load(crate::sync::atomic::Ordering::SeqCst) } } @@ -47,10 +47,10 @@ impl S3DataReader for CountingS3DataReader { async fn spawn_read_task( &self, s3_data_locator: awsfile_bypass_data_locator, - _read_bypass_context: std::sync::Arc, + _read_bypass_context: crate::sync::Arc, ) -> tokio::task::JoinHandle> { self.call_count - .fetch_add(1, std::sync::atomic::Ordering::SeqCst); + .fetch_add(1, crate::sync::atomic::Ordering::SeqCst); let count = s3_data_locator.count as usize; let offset = s3_data_locator.offset; tokio::spawn(async move { @@ -62,6 +62,26 @@ impl S3DataReader for CountingS3DataReader { } } +/// Mock S3DataReader whose reads always fail with a fixed error class. Used to verify +/// the class survives the readahead cache's error plumbing instead of being flattened +/// into an opaque cache-error message. +#[derive(Clone)] +pub struct FailingS3DataReader { + pub make_error: fn() -> crate::aws::s3_client::S3ClientError, +} + +#[async_trait::async_trait] +impl S3DataReader for FailingS3DataReader { + async fn spawn_read_task( + &self, + _s3_data_locator: awsfile_bypass_data_locator, + _read_bypass_context: crate::sync::Arc, + ) -> tokio::task::JoinHandle> { + let make_error = self.make_error; + tokio::spawn(async move { Err(make_error()) }) + } +} + pub fn create_test_s3_data_locator(offset: u64, count: u32) -> awsfile_bypass_data_locator { awsfile_bypass_data_locator { bucket_name: b"test-bucket".to_vec(), @@ -77,7 +97,7 @@ pub fn create_test_s3_data_locator(offset: u64, count: u32) -> awsfile_bypass_da /// (not `test-util`) because `aws_smithy_mocks` is a dev-dependency; external /// consumers construct their own contexts via `ReadBypassContext::default()`. #[cfg(test)] -pub async fn create_test_read_bypass_context() -> std::sync::Arc { +pub async fn create_test_read_bypass_context() -> crate::sync::Arc { use aws_sdk_s3::operation::get_object::GetObjectOutput; use aws_sdk_s3::primitives::ByteStream; use aws_smithy_mocks::{mock, mock_client}; @@ -92,13 +112,13 @@ pub async fn create_test_read_bypass_context() -> std::sync::Arc, rpc_xid: u32) -> Self { - let thread_name = std::thread::current() + let thread_name = crate::sync::thread::current() .name() .unwrap_or("unknown") .to_string(); @@ -45,6 +45,13 @@ impl ReadBypassRequestContext { /// This prefix is automatically prepended by ctx_debug!, ctx_info!, ctx_warn!, ctx_error!, and ctx_trace! macros. pub fn log_prefix(&self) -> String { match self.task_id { + // shuttle-tokio's TaskId implements Debug but not Display. + // TODO: https://github.com/awslabs/shuttle/pull/307 adds the + // Display impl; once it merges and reaches the wrappers, drop + // this cfg split and keep the plain `{}` arm. + #[cfg(feature = "shuttle")] + Some(id) => format!("[{}-{:?}-{}]", self.thread_name, id, self.rpc_xid), + #[cfg(not(feature = "shuttle"))] Some(id) => format!("[{}-{}-{}]", self.thread_name, id, self.rpc_xid), None => format!("[{}-?-{}]", self.thread_name, self.rpc_xid), } diff --git a/src/client-core/src/util/s3_data_reader.rs b/src/client-core/src/util/s3_data_reader.rs index 213ba3c3..eaddba15 100644 --- a/src/client-core/src/util/s3_data_reader.rs +++ b/src/client-core/src/util/s3_data_reader.rs @@ -2,6 +2,7 @@ //! Auxiliary abstraction level between ReadBypassAgent and S3Client. //! +use crate::sync::Arc; use crate::{ aws::s3_client::S3ClientError, nfs::nfs4_1_xdr::awsfile_bypass_data_locator, util::read_bypass_context::ReadBypassContext, @@ -10,7 +11,6 @@ use async_trait::async_trait; use bytes::Bytes; use dyn_clone::{clone_trait_object, DynClone}; use log::warn; -use std::sync::Arc; use tokio::sync::Semaphore; use tokio::task::JoinHandle; @@ -90,8 +90,8 @@ impl S3DataReader for S3ReadBypassReader { #[cfg(test)] mod tests { use super::*; + use crate::sync::atomic::{AtomicUsize, Ordering}; use crate::util::read_bypass_context::ReadBypassContext; - use std::sync::atomic::{AtomicUsize, Ordering}; use tokio::sync::Notify; fn create_test_locator() -> awsfile_bypass_data_locator { diff --git a/src/client-core/src/utils.rs b/src/client-core/src/utils.rs index aea6f227..88162332 100644 --- a/src/client-core/src/utils.rs +++ b/src/client-core/src/utils.rs @@ -1,5 +1,5 @@ -use std::sync::atomic::{AtomicBool, Ordering}; -use std::sync::LazyLock; +use crate::sync::atomic::{AtomicBool, Ordering}; +use crate::sync::LazyLock; use std::time::Duration; use tokio::time::Instant; diff --git a/src/efs_utils_common/aws_credentials.py b/src/efs_utils_common/aws_credentials.py index 33b0eb1a..7fecf7ff 100644 --- a/src/efs_utils_common/aws_credentials.py +++ b/src/efs_utils_common/aws_credentials.py @@ -9,6 +9,7 @@ import logging import os +from datetime import datetime try: from configparser import NoOptionError, NoSectionError @@ -52,6 +53,25 @@ from efs_utils_common.error_reporting import fatal_error from efs_utils_common.metadata import get_dns_name_suffix, url_request_helper +# Format of the "Expiration" field in IMDS / ECS / STS credential responses. +CREDENTIALS_EXPIRATION_DATETIME_FORMAT = "%Y-%m-%dT%H:%M:%SZ" + + +def is_valid_credentials_expiration(expiration): + """Return True if the credential "Expiration" is a parseable ISO-8601 timestamp. + + Credential expirations are whole-second ISO-8601 UTC (e.g. "2019-10-25T21:17:24Z"). Callers + validate before persisting the value into the mount state file so the watchdog only reads + values it can schedule against. + """ + if not expiration: + return False + try: + datetime.strptime(expiration, CREDENTIALS_EXPIRATION_DATETIME_FORMAT) + return True + except (ValueError, TypeError): + return False + def get_aws_security_credentials( config, @@ -248,6 +268,7 @@ def get_aws_security_credentials_from_webidentity( "AccessKeyId": creds["AccessKeyId"], "SecretAccessKey": creds["SecretAccessKey"], "Token": creds["SessionToken"], + "Expiration": creds.get("Expiration"), }, "webidentity:" + ",".join([role_arn, token_file]) # Fail if credentials cannot be fetched from the given aws_creds_uri @@ -398,7 +419,8 @@ def botocore_credentials_helper(awsprofile): session.set_config_variable("profile", awsprofile) try: - frozen_credentials = session.get_credentials().get_frozen_credentials() + creds_object = session.get_credentials() + frozen_credentials = creds_object.get_frozen_credentials() except ProfileNotFound as e: fatal_error( "%s, please add the [profile %s] section in the aws config file following %s and %s." @@ -408,6 +430,14 @@ def botocore_credentials_helper(awsprofile): credentials["AccessKeyId"] = frozen_credentials.access_key credentials["SecretAccessKey"] = frozen_credentials.secret_key credentials["Token"] = frozen_credentials.token + + # Surface the expiration for temporary (assumed-role/session) profiles so the watchdog can + # refresh ahead of it. botocore exposes it only via the private _expiry_time; static profiles have none. + expiry_time = getattr(creds_object, "_expiry_time", None) + if isinstance(expiry_time, datetime): + credentials["Expiration"] = expiry_time.strftime( + CREDENTIALS_EXPIRATION_DATETIME_FORMAT + ) return credentials diff --git a/src/efs_utils_common/constants.py b/src/efs_utils_common/constants.py index 3a84fde6..bbaa8c10 100644 --- a/src/efs_utils_common/constants.py +++ b/src/efs_utils_common/constants.py @@ -11,7 +11,7 @@ import pwd import re -VERSION = "3.3.0" +VERSION = "3.3.1" AMAZON_LINUX_2_RELEASE_ID = "Amazon Linux release 2 (Karoo)" AMAZON_LINUX_2_PRETTY_NAME = "Amazon Linux 2" @@ -146,7 +146,7 @@ WATCHDOG_SERVICE = "amazon-efs-mount-watchdog" # MacOS instances use plist files. This files needs to be loaded on launchctl (init system of MacOS) -WATCHDOG_SERVICE_PLIST_PATH = "/Library/LaunchAgents/amazon-efs-mount-watchdog.plist" +WATCHDOG_SERVICE_PLIST_PATH = "/Library/LaunchDaemons/amazon-efs-mount-watchdog.plist" SYSTEM_RELEASE_PATH = "/etc/system-release" OS_RELEASE_PATH = "/etc/os-release" MACOS_BIG_SUR_RELEASE = "macOS-11" diff --git a/src/efs_utils_common/process_utils.py b/src/efs_utils_common/process_utils.py index 238198b6..df9f7959 100644 --- a/src/efs_utils_common/process_utils.py +++ b/src/efs_utils_common/process_utils.py @@ -99,7 +99,7 @@ def subprocess_call(cmd, error_message): process = subprocess.Popen( cmd.split(), stdout=subprocess.PIPE, stderr=subprocess.PIPE, close_fds=True ) - (output, err) = process.communicate() + output, err = process.communicate() rc = process.poll() if rc != 0: logging.error( diff --git a/src/efs_utils_common/proxy.py b/src/efs_utils_common/proxy.py index dd38edb3..8cb2c9dd 100644 --- a/src/efs_utils_common/proxy.py +++ b/src/efs_utils_common/proxy.py @@ -24,6 +24,7 @@ from efs_utils_common.aws_credentials import ( get_aws_profile, get_aws_security_credentials, + is_valid_credentials_expiration, ) from efs_utils_common.certificate_utils import create_certificate, get_private_key_path from efs_utils_common.config_utils import ( @@ -642,25 +643,35 @@ def start_watchdog(init_system): logging.debug("%s is already running", WATCHDOG_SERVICE) elif init_system == "launchd": - rc = subprocess.Popen( - ["sudo", "launchctl", "list", WATCHDOG_SERVICE], - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, + rc = subprocess.call( + ["launchctl", "list", WATCHDOG_SERVICE], + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, close_fds=True, ) - if rc != 0: + if rc == 0: + logging.debug("%s is already running", WATCHDOG_SERVICE) + else: if not os.path.exists(WATCHDOG_SERVICE_PLIST_PATH): fatal_error( - "Watchdog plist file missing. Copy the watchdog plist file in directory /Library/LaunchAgents" + "Watchdog plist file missing. Copy the watchdog plist file to /Library/LaunchDaemons/" ) - subprocess.Popen( - ["sudo", "launchctl", "load", WATCHDOG_SERVICE_PLIST_PATH], + return + + logging.debug("Loading watchdog from %s", WATCHDOG_SERVICE_PLIST_PATH) + rc = subprocess.call( + ["launchctl", "load", WATCHDOG_SERVICE_PLIST_PATH], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, close_fds=True, ) - else: - logging.debug("%s is already running", WATCHDOG_SERVICE) + if rc != 0: + fatal_error( + "Failed to load watchdog plist %s (exit code %d)" + % (WATCHDOG_SERVICE_PLIST_PATH, rc) + ) + return + logging.debug("Loaded watchdog from %s", WATCHDOG_SERVICE_PLIST_PATH) else: error_message = 'Could not start %s, unrecognized init system "%s"' % ( @@ -750,6 +761,27 @@ def bootstrap_proxy( credentials_source, ) + # Persist the credential expiration so the watchdog can schedule refresh + # ahead of it. + expiration = ( + security_credentials.get("Expiration") + if security_credentials + else None + ) + if not expiration: + logging.debug( + "No credential expiration available; the certificate will refresh on the " + "fixed tls_cert_renewal_interval_min interval" + ) + elif is_valid_credentials_expiration(expiration): + cert_details["certificateExpirationTime"] = expiration + else: + logging.warning( + 'Credential expiration "%s" is not valid ISO-8601; the watchdog will ' + "fall back to the tls_cert_renewal_interval_min fixed interval", + expiration, + ) + # Access points must be mounted over TLS if ap_id: cert_details["accessPoint"] = ap_id diff --git a/src/proxy/Cargo.toml b/src/proxy/Cargo.toml index bb5568cc..ed2d5bc4 100644 --- a/src/proxy/Cargo.toml +++ b/src/proxy/Cargo.toml @@ -3,7 +3,7 @@ name = "efs-proxy" edition = "2021" # The version of efs-proxy is tied to efs-utils. -version = "3.3.0" +version = "3.3.1" publish = false license = "MIT" diff --git a/src/proxy/src/controller.rs b/src/proxy/src/controller.rs index 3caea1a9..574efbf2 100644 --- a/src/proxy/src/controller.rs +++ b/src/proxy/src/controller.rs @@ -26,6 +26,11 @@ use tokio_util::sync::CancellationToken; pub const METRICS_EMISSION_PERIOD: Duration = Duration::from_secs(60); +// Max time to wait for the NFS client's first byte after we accept a reconnect. A connection that +// is accepted but never sends data would otherwise block peek() forever and wedge the mount; on +// timeout we drop it and return to accept(). +pub const PEEK_FIRST_BYTE_TIMEOUT: Duration = Duration::from_secs(2); + pub const AWSFILE_CHANNEL_INIT_MINOR_VERSION: u32 = 2; pub const DEFAULT_SCALE_UP_BACKOFF: Duration = Duration::from_secs(300); @@ -180,17 +185,36 @@ impl Controller { } }; - let peek_result = nfs_client.peek(&mut [0; 1]).await; - if let Ok(0) = peek_result { - // efs-utils performs a test in which it checks if a connection to the proxy port - // can be established. This connection is never used and is immediately closed. - // When this behavior is detected, this loops should be restarted so that another - // connection to the port can be established - debug!("Connection to nfs client was closed before any data was sent to the proxy. This is expected. Restarting controller"); - continue; - } else if let Err(e) = peek_result { - error!("Failed to check if data was sent by the NFS client. {}", e); - return Some(ShutdownReason::UnexpectedError); + let peek_result = + tokio::time::timeout(PEEK_FIRST_BYTE_TIMEOUT, nfs_client.peek(&mut [0; 1])).await; + match peek_result { + Ok(Ok(0)) => { + // efs-utils performs a test in which it checks if a connection to the proxy + // port can be established. This connection is never used and is immediately + // closed. When this behavior is detected, this loop should be restarted so + // that another connection to the port can be established. + debug!("Connection to nfs client was closed before any data was sent to the proxy. This is expected. Restarting controller"); + continue; + } + Ok(Err(e)) => { + error!("Failed to check if data was sent by the NFS client. {}", e); + return Some(ShutdownReason::UnexpectedError); + } + Err(_elapsed) => { + // The accepted connection sent no data within the timeout. A silent + // connection (e.g. a non-NFS local process that connected to the loopback + // listen port) would otherwise block here indefinitely and wedge the mount. + warn!( + "Accepted connection sent no data within {:?}; dropping it and re-accepting", + PEEK_FIRST_BYTE_TIMEOUT + ); + // Explicitly close the idle socket (sends FIN) before returning to accept(). + drop(nfs_client); + continue; + } + Ok(Ok(_)) => { + // The NFS client sent data; proceed to establish the upstream connection. + } } // Set init_deadline to be 1 second less than the `proxy_init_timeout_sec`, as we @@ -740,4 +764,72 @@ mod tests { // Very large counts must saturate to the cap, not overflow/panic. assert_eq!(reconnect_backoff(1000), DEFAULT_RECONNECT_BACKOFF_CAP); } + + // peek() reconnect-handoff tests: a connection that is accepted but never sends data is + // dropped after PEEK_FIRST_BYTE_TIMEOUT, and the NFS client's connection is then serviced. + use crate::test_utils::make_signaling_controller; + use tokio::io::AsyncWriteExt as _; + + fn spawn_run( + controller: Controller, + ) -> tokio::task::JoinHandle> { + tokio::spawn(controller.run( + CancellationToken::new(), + crate::awsfile_rpc::AwsFileRpcClient, + crate::aws::s3_client::S3ClientStandardBuilder, + )) + } + + // An idle connection accepted first is dropped after the peek timeout, and the NFS client + // waiting in the listen backlog is then serviced. + #[tokio::test(start_paused = true)] + async fn idle_connection_is_dropped_and_nfs_client_is_served() { + let (addr, controller, mut rx) = make_signaling_controller().await; + let _handle = spawn_run(controller); + + // Barrier: connect the idle socket and, before the peek timeout elapses, assert the + // controller has NOT established yet. This guarantees the idle socket (not the NFS + // client) is the one being peeked, so the test provably exercises the drop path. + let _idle = TcpStream::connect(addr).await.unwrap(); + tokio::time::sleep(PEEK_FIRST_BYTE_TIMEOUT / 2).await; + assert!( + rx.try_recv().is_err(), + "controller established before the idle socket's peek timed out" + ); + + // NFS client waits in the backlog; it must be served once the idle socket is dropped. + let mut nfs_client = TcpStream::connect(addr).await.unwrap(); + nfs_client.write_all(&[0x80, 0, 0, 0]).await.unwrap(); + + let reached = tokio::time::timeout(Duration::from_secs(30), rx.recv()).await; + assert!( + matches!(reached, Ok(Some(()))), + "controller did not recover: NFS client was not served after the idle socket was dropped" + ); + } + + // After the idle socket is dropped on timeout, a NFS client that connects afterwards is still + // served -- i.e. dropping the idle socket returns the controller to accept(). + #[tokio::test(start_paused = true)] + async fn nfs_client_connecting_after_idle_drop_is_served() { + let (addr, controller, mut rx) = make_signaling_controller().await; + let _handle = spawn_run(controller); + + let idle = TcpStream::connect(addr).await.unwrap(); + tokio::time::sleep(PEEK_FIRST_BYTE_TIMEOUT + Duration::from_secs(1)).await; + assert!( + rx.try_recv().is_err(), + "establish_connection reached with no data-sending client" + ); + drop(idle); + + let mut nfs_client = TcpStream::connect(addr).await.unwrap(); + nfs_client.write_all(&[0x80, 0, 0, 0]).await.unwrap(); + + let reached = tokio::time::timeout(Duration::from_secs(30), rx.recv()).await; + assert!( + matches!(reached, Ok(Some(()))), + "NFS client connecting after the idle socket was dropped was not served" + ); + } } diff --git a/src/proxy/src/read_bypass/read_bypass_agent.rs b/src/proxy/src/read_bypass/read_bypass_agent.rs index 486d0a49..4cce2db4 100644 --- a/src/proxy/src/read_bypass/read_bypass_agent.rs +++ b/src/proxy/src/read_bypass/read_bypass_agent.rs @@ -12,6 +12,7 @@ use std::{ use bytes::{BufMut, Bytes, BytesMut}; use futures::FutureExt; use log::{debug, error, info, trace, warn}; +use rand::Rng; use tokio::sync::mpsc; use xdr_codec::Pack; @@ -27,7 +28,10 @@ use crate::{ nfs_rpc_envelope::{NfsRpcEnvelope, NfsRpcInfo}, }, proxy_task::ConnectionMessage, - read_ahead::{error::ReadAheadCacheError, readahead_cache::FileReadAheadCache}, + read_ahead::{ + error::{ReadAheadCacheError, ReadAheadError}, + readahead_cache::FileReadAheadCache, + }, rpc::{rpc::RpcBatch, rpc_encoder::RpcEncoder, rpc_envelope::RpcMessageParams}, shutdown::ShutdownHandle, util::{ @@ -36,7 +40,13 @@ use crate::{ s3_data_reader::{S3DataReader, S3ReadBypassReader}, }, }; -use crate::{ctx_debug, ctx_error, ctx_trace, ctx_warn, util::fh_denylist::FileHandle}; +use crate::{ctx_debug, ctx_error, ctx_info, ctx_trace, ctx_warn, util::fh_denylist::FileHandle}; + +/// Bounds of the random delay applied before returning NFS4ERR_DELAY for a transient +/// read-bypass failure. The minimum matches the kernel's own 100ms retry interval; the maximum +/// keeps per-retry latency bounded while still spreading correlated retries. +const TRANSIENT_RETRY_DELAY_MIN: std::time::Duration = std::time::Duration::from_millis(100); +const TRANSIENT_RETRY_DELAY_MAX: std::time::Duration = std::time::Duration::from_secs(1); #[derive(Debug, thiserror::Error)] pub enum ReadBypassAgentError { @@ -60,6 +70,22 @@ pub enum ReadBypassAgentError { UnsupportedMessage, } +impl ReadBypassAgentError { + /// Whether the failure is expected to clear on its own, so the read gets an NFS4ERR_DELAY + /// and another read-bypass attempt instead of denylisting the file handle. + /// + /// Denylisting pushes the same reads onto the NFS server, which reads the same S3 object. + /// For a throttled GET that amplifies the throttling instead of shedding it. Timeouts and + /// connection failures are excluded: a host with broken S3 connectivity is better off on the + /// server for the denylist TTL than retrying reads that keep timing out. + pub fn is_transient(&self) -> bool { + matches!( + self, + Self::DataEvicted | Self::S3Error(S3ClientError::Throttled) + ) + } +} + #[async_trait::async_trait] pub trait S3Reader: Send + Sync { async fn read_data( @@ -140,8 +166,22 @@ impl S3Reader for CachedS3Reader { .await { Ok(data) => Ok(data), - Err(e) if e.message.contains("Data evicted") => Err(ReadBypassAgentError::DataEvicted), - Err(e) => Err(ReadBypassAgentError::CacheInternalError(e)), + // Keep the error class so the caller's denylist policy can see a throttled GET or an + // evicted cache entry. + Err(ReadAheadError { + s3_error: Some(e), .. + }) => Err(ReadBypassAgentError::S3Error(e)), + Err(ReadAheadError { + cache_error: Some(ReadAheadCacheError::DataEvicted), + .. + }) => Err(ReadBypassAgentError::DataEvicted), + Err(ReadAheadError { + cache_error: Some(e), + .. + }) => Err(ReadBypassAgentError::CacheInternalError(e)), + Err(e) => Err(ReadBypassAgentError::CacheInternalError( + ReadAheadCacheError::Other(e.message), + )), } } } @@ -381,18 +421,33 @@ impl ReadBypassAgent { software: consider disabling ReadBypass functionality." ); } - ReadBypassAgentError::DataEvicted => { - // Transient memory pressure - send delay, don't denylist - ctx_warn!( + _ if e.is_transient() => { + // The read is expected to succeed on retry, so send NFS4ERR_DELAY and + // leave the filehandle eligible for read bypass. + let delay = Self::transient_retry_delay(); + // INFO, not DEBUG: this is the only record that a read was retried + // instead of denylisted, and the default logging level is INFO. It is + // one line per transient failure, which is the same rate the read + // would have logged if it had been denylisted instead. + ctx_info!( read_bypass_request_context, - "Data evicted, sending NFS4ERR_DELAY without denylisting" + "Transient ReadBypass failure ({e}), sending NFS4ERR_DELAY after {delay:?} without denylisting" ); - let _ = Self::respond_failure_to_nfs_client( + // We are inside a per-request `tokio::spawn`, so sleeping here delays + // only this response, not the agent's message loop. + tokio::time::sleep(delay).await; + if let Err(e) = Self::respond_failure_to_nfs_client( read_bypass_request_context.clone(), message, nfs_client_sender, ) - .await; + .await + { + ctx_warn!( + read_bypass_request_context, + "Failed to send NFS4ERR_DELAY response for transient failure: {e}" + ); + } } _ => { ctx_warn!( @@ -400,10 +455,9 @@ impl ReadBypassAgent { "Error while processing ReadBypass compound: {e}" ); - // Denylist file handle, assuming that any failure during processing - // compound at this phase is caused by issues with S3 access and highly - // likely will repeat itself, so we want to deny list filehandle to - // avoid availability issues. + // The failure is not expected to clear on its own, so denylist the + // file handle and let the NFS server serve this file's reads for the + // denylist TTL. for (size, op) in compound_info.compound.resarray.iter_mut().enumerate() { if let nfs_resop4::OP_AWSFILE_READ_BYPASS( @@ -535,9 +589,9 @@ impl ReadBypassAgent { index ); return Err(ReadBypassAgentError::CacheInternalError( - ReadAheadCacheError { - message: "No data returned from read operation".to_string(), - }, + ReadAheadCacheError::Other( + "No data returned from read operation".to_string(), + ), )); } Err(ReadBypassAgentError::DataEvicted) => { @@ -616,6 +670,16 @@ impl ReadBypassAgent { .await; } + /// Random delay to wait before returning NFS4ERR_DELAY for a transient failure. + /// + /// This is the only backoff between NFS retry cycles: the S3 client's exponential backoff + /// applies only within a single GetObject, and the Linux NFS client retries NFS4ERR_DELAY on + /// a READ with a flat, unjittered 100ms, so clients that hit the same throttling event would + /// otherwise retry in lockstep. See `docs/read_bypass_design.md` for the kernel references. + fn transient_retry_delay() -> std::time::Duration { + rand::thread_rng().gen_range(TRANSIENT_RETRY_DELAY_MIN..=TRANSIENT_RETRY_DELAY_MAX) + } + async fn respond_failure_to_nfs_client( read_bypass_request_context: Arc, message: NfsRpcInfo, @@ -717,10 +781,11 @@ mod tests { use crate::rpc::rpc_envelope::{ EnvelopeHeader, RpcMessageParams, RpcMessageType, RpcReplyParams, }; - use crate::test_utils::get_test_config; + use crate::test_utils::{get_test_config, FailingS3Reader}; use crate::util::read_bypass_request_context; use bytes::{Bytes, BytesMut}; use mockall::predicate::le; + use std::time::Duration; use tokio::sync::mpsc; use tokio::task::JoinHandle; use tokio_util::sync::CancellationToken; @@ -1540,7 +1605,7 @@ mod tests { } } - #[tokio::test] + #[tokio::test(start_paused = true)] async fn test_data_evicted_does_not_denylist() { // DataEvicted should send NFS4ERR_DELAY but NOT denylist let mock_reader: Arc = Arc::new(DataEvictedMockReader::new(2)); @@ -1598,6 +1663,140 @@ mod tests { } } + /// Runs one read-bypass reply whose S3 read fails with `make_error`, and reports whether the + /// file handle ended up denylisted along with the status returned to the NFS client. + async fn denylist_and_status_for_s3_error( + make_error: fn() -> S3ClientError, + ) -> (bool, nfsstat4) { + let s3_reader: Arc = Arc::new(FailingS3Reader { make_error }); + let read_bypass_context = Arc::new(ReadBypassContext::default().await); + let read_bypass_request_context = + Arc::new(ReadBypassRequestContext::new(read_bypass_context, 0)); + let (nfs_client_sender, mut nfs_client_receiver) = mpsc::channel::(10); + + let compound_res = COMPOUND4res { + status: nfsstat4::NFS4_OK, + tag: utf8string(b"test".to_vec()), + resarray: vec![ + get_sample_op_sequence_res(), + get_sample_op_read_bypass_accepted_res(0, 1024, 2048), + get_sample_op_getattr_res(), + ], + }; + let message = create_nfs_rpc_envelope_batch_from_compound( + RpcMessageType::Reply, + compound_res.clone(), + ); + + ReadBypassAgent::process_message( + read_bypass_request_context.clone(), + message, + s3_reader, + nfs_client_sender, + ) + .await; + + let denylisted = if let nfs_resop4::OP_AWSFILE_READ_BYPASS( + AWSFILE_READ_BYPASS4res::NFS4ERR_AWSFILE_BYPASS(err_op), + ) = &compound_res.resarray[1] + { + read_bypass_request_context + .fh_denylist + .contains(&err_op.filehandle) + } else { + panic!("Expected read bypass file error"); + }; + + let response = nfs_client_receiver + .recv() + .await + .expect("Should receive a response"); + let ConnectionMessage::Response(batch) = response; + let nfs_envelope = NfsRpcEnvelope::try_from(batch.rpcs[0].clone()) + .expect("Failed to parse NfsRpcEnvelope"); + let status = if let RefNfsCompound::Compound4res(compound_info) = &nfs_envelope.body { + compound_info.compound.status + } else { + panic!("Expected Compound4res"); + }; + + (denylisted, status) + } + + /// A throttled S3 GET is the failure mode this policy exists for: the server would read + /// the very same S3 object and hit the same throttling, so denylisting the file handle + /// amplifies the event for the full denylist TTL instead of shedding it. + #[tokio::test(start_paused = true)] + async fn test_throttled_s3_does_not_denylist() { + let (denylisted, status) = + denylist_and_status_for_s3_error(|| S3ClientError::Throttled).await; + + assert!( + !denylisted, + "Should NOT denylist the file handle on a throttled S3 GET" + ); + assert!( + matches!(status, nfsstat4::NFS4ERR_DELAY), + "Expected NFS4ERR_DELAY, got {status:?}" + ); + } + + /// Counterpart to the test above: failures that will not clear on their own still denylist, + /// so reads for that file fall back to the NFS server instead of failing repeatedly. + #[tokio::test] + async fn test_non_transient_s3_errors_still_denylist() { + for make_error in [ + (|| S3ClientError::AccessDenied) as fn() -> S3ClientError, + || S3ClientError::NoSuchKey, + || S3ClientError::ETagMismatchError, + || S3ClientError::Timeout, + || S3ClientError::NotEnabled, + ] { + let (denylisted, status) = denylist_and_status_for_s3_error(make_error).await; + assert!( + denylisted, + "Should denylist on non-transient error {:?}", + make_error() + ); + assert!( + matches!(status, nfsstat4::NFS4ERR_DELAY), + "Expected NFS4ERR_DELAY, got {status:?}" + ); + } + } + + #[test] + fn test_is_transient_classification() { + assert!(ReadBypassAgentError::DataEvicted.is_transient()); + assert!(ReadBypassAgentError::S3Error(S3ClientError::Throttled).is_transient()); + + for error in [ + ReadBypassAgentError::S3Error(S3ClientError::AccessDenied), + ReadBypassAgentError::S3Error(S3ClientError::NoSuchKey), + ReadBypassAgentError::S3Error(S3ClientError::ETagMismatchError), + ReadBypassAgentError::S3Error(S3ClientError::Timeout), + ReadBypassAgentError::S3Error(S3ClientError::NotEnabled), + ReadBypassAgentError::S3Error(S3ClientError::SizeMismatch { + expected: 2, + actual: 1, + }), + ReadBypassAgentError::CacheInternalError(ReadAheadCacheError::Other( + "boom".to_string(), + )), + ReadBypassAgentError::DispatchingError, + ReadBypassAgentError::InvalidCompound, + ReadBypassAgentError::JoinFailure, + ReadBypassAgentError::NfsResponseEncodingError, + ReadBypassAgentError::OperationConversionFailure, + ReadBypassAgentError::UnsupportedMessage, + ] { + assert!( + !error.is_transient(), + "{error:?} should not be treated as transient" + ); + } + } + /// Mock S3Reader that panics to test panic recovery struct PanickingS3Reader; diff --git a/src/proxy/src/test_utils.rs b/src/proxy/src/test_utils.rs index 5453c45e..f75f73f1 100644 --- a/src/proxy/src/test_utils.rs +++ b/src/proxy/src/test_utils.rs @@ -6,6 +6,7 @@ use crate::{ aws::cw_publisher::{CloudWatchClient, LogLevel}, + aws::s3_client::S3ClientError, awsfile_prot::{ self, AwsFileChannelInitArgs, AwsFileChannelInitRes, BindClientResponse, BindResponse, ScaleUpConfig, @@ -17,21 +18,27 @@ use crate::{ error::RpcError, nfs::{ nfs4_1_xdr, + nfs4_1_xdr::{awsfile_bypass_data_locator, nfs_fh4}, nfs_compound::{NfsMetadata, RefNfsCompoundInfo}, }, proxy_identifier::ProxyIdentifier, + read_bypass::read_bypass_agent::{ReadBypassAgentError, S3Reader}, status_reporter::create_status_channel, tls::{create_config_builder, InsecureAcceptAllCertificatesHandler, TlsConfig}, + util::read_bypass_request_context::ReadBypassRequestContext, }; use anyhow::Result; -use bytes::BytesMut; +use bytes::{Bytes, BytesMut}; use rand::{Rng, RngCore}; use s2n_tls::config::Config; +use std::net::SocketAddr; use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::Arc; use std::time::Duration; use std::{io::Cursor, path::Path}; use tokio::net::{TcpListener, TcpStream}; +use tokio::sync::mpsc; +use tokio::time::Instant; // Proxy Configuration testing utils // @@ -317,3 +324,89 @@ pub async fn make_test_controller( consecutive_connect_failures: 0, } } + +// PartitionFinder that signals when establish_connection is reached (i.e. the controller got +// past peek()), then fails fast so the controller loops back to accept(). +pub struct SignalingPartitionFinder { + pub reached_establish: mpsc::UnboundedSender<()>, +} + +#[async_trait::async_trait] +impl PartitionFinder for SignalingPartitionFinder { + async fn create_connect_future( + &self, + ) -> futures::future::BoxFuture<'static, Result> { + unimplemented!("establish_connection is overridden; connect future is never used") + } + + async fn establish_connection( + &self, + _deadline: Instant, + _proxy_id: ProxyIdentifier, + ) -> Result< + ( + TcpStream, + Option, + Option, + ), + crate::error::ConnectError, + > { + let _ = self.reached_establish.send(()); + Err(crate::error::ConnectError::Timeout) + } +} + +// Build a Controller bound to an ephemeral localhost port whose PartitionFinder signals when +// establish_connection is reached. Returns the listen address, the controller (spawn its run +// loop in the test), and a receiver that fires on each establish_connection. +pub async fn make_signaling_controller() -> ( + SocketAddr, + Controller, + mpsc::UnboundedReceiver<()>, +) { + let (_status_requester, status_reporter) = create_status_channel(); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let (tx, rx) = mpsc::unbounded_channel(); + let mut proxy_config = ProxyConfig::default(); + proxy_config.nested_config.fs_id = "fs-test123".to_string(); + let controller = Controller:: { + listener, + partition_finder: Arc::new(SignalingPartitionFinder { + reached_establish: tx, + }), + proxy_id: ProxyIdentifier::new(), + scale_up_attempt_count: 0, + restart_count: 0, + scale_up_config: DEFAULT_SCALE_UP_CONFIG, + status_reporter, + proxy_config, + cw_publisher: None, + last_reachability_emitted: None, + last_reachability_emit_at: None, + reachability_emit_period: Duration::from_secs(1), + consecutive_connect_failures: 0, + }; + (addr, controller, rx) +} + +// Read-bypass testing utils +// + +/// Mock S3Reader whose reads always fail with a fixed S3 error class. +pub struct FailingS3Reader { + pub make_error: fn() -> S3ClientError, +} + +#[async_trait::async_trait] +impl S3Reader for FailingS3Reader { + async fn read_data( + &self, + _read_bypass_request_context: Arc, + _s3_data_locator: awsfile_bypass_data_locator, + _filehandle: nfs_fh4, + _file_size: u64, + ) -> Result, ReadBypassAgentError> { + Err(ReadBypassAgentError::S3Error((self.make_error)())) + } +} diff --git a/src/watchdog/__init__.py b/src/watchdog/__init__.py index ed10b5fe..62580735 100755 --- a/src/watchdog/__init__.py +++ b/src/watchdog/__init__.py @@ -56,7 +56,7 @@ AMAZON_LINUX_2_RELEASE_ID, AMAZON_LINUX_2_PRETTY_NAME, ] -VERSION = "3.3.0" +VERSION = "3.3.1" SERVICE = "elasticfilesystem" FS_PREFIX = "fs-" @@ -84,6 +84,8 @@ EFS_SERVICE_NAME = "elasticfilesystem" PRIVATE_KEY_FILE = "/etc/amazon/efs/privateKey.pem" DEFAULT_REFRESH_SELF_SIGNED_CERT_INTERVAL_MIN = 60 +# Refresh the certificate this many minutes before the credential (or the cert's own NOT_AFTER) expires. +CERT_EXPIRATION_SAFETY_MARGIN_MIN = 5 DEFAULT_STUNNEL_HEALTH_CHECK_INTERVAL_MIN = 5 DEFAULT_STUNNEL_HEALTH_CHECK_TIMEOUT_SEC = 30 NOT_BEFORE_MINS = 15 @@ -91,6 +93,7 @@ DATE_ONLY_FORMAT = "%Y%m%d" SIGV4_DATETIME_FORMAT = "%Y%m%dT%H%M%SZ" CERT_DATETIME_FORMAT = "%y%m%d%H%M%SZ" +CREDENTIALS_EXPIRATION_DATETIME_FORMAT = "%Y-%m-%dT%H:%M:%SZ" AWS_CREDENTIALS_FILES = { "credentials": os.path.expanduser( @@ -396,7 +399,8 @@ def botocore_credentials_helper(awsprofile): session.set_config_variable("profile", awsprofile) try: - frozen_credentials = session.get_credentials().get_frozen_credentials() + creds_object = session.get_credentials() + frozen_credentials = creds_object.get_frozen_credentials() except ProfileNotFound as e: logging.error( "%s, please add the [profile %s] section in the aws config file following %s and %s." @@ -407,6 +411,13 @@ def botocore_credentials_helper(awsprofile): credentials["AccessKeyId"] = frozen_credentials.access_key credentials["SecretAccessKey"] = frozen_credentials.secret_key credentials["Token"] = frozen_credentials.token + # Surface the expiration for temporary (assumed-role/session) profiles so refresh schedules + # ahead of it. botocore exposes it only via the private _expiry_time; static profiles have none. + expiry_time = getattr(creds_object, "_expiry_time", None) + if isinstance(expiry_time, datetime): + credentials["Expiration"] = expiry_time.strftime( + CREDENTIALS_EXPIRATION_DATETIME_FORMAT + ) return credentials @@ -520,6 +531,7 @@ def get_aws_security_credentials_from_webidentity(config, role_arn, token_file, "AccessKeyId": creds["AccessKeyId"], "SecretAccessKey": creds["SecretAccessKey"], "Token": creds["SessionToken"], + "Expiration": creds.get("Expiration"), } return None @@ -1664,20 +1676,69 @@ def read_config(config_file=CONFIG_FILE): return p +def parse_credentials_expiration(expiration): + """Parse a credential "Expiration" ISO-8601 string into a tz-aware UTC datetime, or None. + + Values are validated at the persist sites, so a failure here (absent or corrupted state) falls + back to the fixed interval silently. + """ + if not expiration: + return None + try: + return datetime.strptime( + expiration, CREDENTIALS_EXPIRATION_DATETIME_FORMAT + ).replace(tzinfo=timezone.utc) + except (ValueError, TypeError): + return None + + +def get_certificate_refresh_deadline( + certificate_creation_time, + renewal_interval_secs, + credentials_expiration, +): + """Return the UTC datetime at which the certificate should next be refreshed. + + The deadline is the earliest of: + - the fixed renewal interval, + - the credential's expiration, less the safety margin, and + - the certificate's own NOT_AFTER lifetime, less the safety margin. + """ + safety_margin = timedelta(minutes=CERT_EXPIRATION_SAFETY_MARGIN_MIN) + interval_deadline = certificate_creation_time + timedelta( + seconds=renewal_interval_secs + ) + cert_ttl_deadline = ( + certificate_creation_time + timedelta(hours=NOT_AFTER_HOURS) - safety_margin + ) + deadline = min(interval_deadline, cert_ttl_deadline) + if credentials_expiration is not None: + deadline = min(deadline, credentials_expiration - safety_margin) + return deadline + + def check_certificate( config, state, state_file_dir, state_file, service, base_path=STATE_FILE_DIR ): certificate_creation_time = datetime.strptime( state["certificateCreationTime"], CERT_DATETIME_FORMAT - ) + ).replace(tzinfo=timezone.utc) certificate_exists = os.path.isfile(state["certificate"]) certificate_renewal_interval_secs = ( get_certificate_renewal_interval_mins(config) * 60 ) + now = get_utc_now() + # creation instead of NOT_BEFORE datetime is used for refresh of cert because NOT_BEFORE derives from creation datetime - should_refresh_cert = ( - get_utc_now() - certificate_creation_time.replace(tzinfo=timezone.utc) - ).total_seconds() > certificate_renewal_interval_secs + credentials_expiration = parse_credentials_expiration( + state.get("certificateExpirationTime") + ) + refresh_deadline = get_certificate_refresh_deadline( + certificate_creation_time, + certificate_renewal_interval_secs, + credentials_expiration, + ) + should_refresh_cert = now > refresh_deadline if certificate_exists and not should_refresh_cert: return @@ -1699,21 +1760,38 @@ def check_certificate( logging.debug( "Refreshing self-signed certificate (at %s)" % state["certificate"] ) + interval_deadline = certificate_creation_time + timedelta( + seconds=certificate_renewal_interval_secs + ) + if refresh_deadline < interval_deadline: + logging.debug( + "Refreshing %s early: scheduled deadline %s precedes the fixed-interval deadline %s", + state["certificate"], + refresh_deadline.strftime(CREDENTIALS_EXPIRATION_DATETIME_FORMAT), + interval_deadline.strftime(CREDENTIALS_EXPIRATION_DATETIME_FORMAT), + ) credentials_source = state.get("awsCredentialsMethod") - updated_certificate_creation_time = recreate_certificate( - config, - state["mountStateDir"], - state["commonName"], - state["fsId"], - credentials_source, - ap_state, - state["region"], - service, - base_path=base_path, + updated_certificate_creation_time, updated_credentials_expiration = ( + recreate_certificate( + config, + state["mountStateDir"], + state["commonName"], + state["fsId"], + credentials_source, + ap_state, + state["region"], + service, + base_path=base_path, + ) ) if updated_certificate_creation_time: state["certificateCreationTime"] = updated_certificate_creation_time + # updated_credentials_expiration is pre-validated by create_ca_conf: valid ISO-8601 or None. + if updated_credentials_expiration: + state["certificateExpirationTime"] = updated_credentials_expiration + else: + state.pop("certificateExpirationTime", None) rewrite_state_file(state, state_file_dir, state_file) # send SIGHUP to force a reload of the configuration file to trigger the stunnel process to notice the new certificate @@ -1833,7 +1911,7 @@ def recreate_certificate( create_public_key(private_key, public_key) client_info = get_client_info(config) - config_body = create_ca_conf( + config_body, credentials_expiration = create_ca_conf( config, certificate_config, common_name, @@ -1850,7 +1928,7 @@ def recreate_certificate( if not config_body: logging.error("Cannot recreate self-signed certificate") - return None + return None, None create_certificate_signing_request( certificate_config, private_key, certificate_signing_request @@ -1870,7 +1948,7 @@ def recreate_certificate( ) ) subprocess_call(cmd, "Failed to create self-signed client-side certificate") - return current_time.strftime(CERT_DATETIME_FORMAT) + return current_time.strftime(CERT_DATETIME_FORMAT), credentials_expiration def get_private_key_path(): @@ -1964,7 +2042,13 @@ def create_ca_conf( ap_id=None, client_info=None, ): - """Populate ca/req configuration file with fresh configurations at every mount since SigV4 signature can change""" + """Populate ca/req configuration file with fresh configurations at every mount since SigV4 signature can change + + Returns a (full_config_body, credentials_expiration) tuple. credentials_expiration is the + credential's "Expiration" as a valid ISO-8601 string when the source provides a usable one, + else None (absent, or present but unparseable). It lets the caller schedule the next refresh + before the embedded credential expires. On failure returns (None, None). + """ public_key_path = os.path.join(directory, "publicKey.pem") security_credentials = ( get_aws_security_credentials(config, credentials_source, region) @@ -1977,7 +2061,39 @@ def create_ca_conf( "Failed to retrieve AWS security credentials using lookup method: %s", credentials_source, ) - return None + return None, None + + credentials_expiration = ( + security_credentials.get("Expiration") if security_credentials else None + ) + + # Validate the expiration once, here, so the caller can persist the returned value without + # re-checking it: it is normalized to a valid ISO-8601 string or None. + if credentials_expiration: + parsed_expiration = parse_credentials_expiration(credentials_expiration) + if parsed_expiration is None: + # Present but unparseable: mint the cert, but drop the value so the refresh falls back + # to the fixed interval (can't schedule against a timestamp we can't read). + logging.warning( + 'Credential expiration "%s" is not valid ISO-8601; certificate refresh will fall ' + "back to the fixed interval", + credentials_expiration, + ) + credentials_expiration = None + elif parsed_expiration <= date: + # A cert minted from already-expired credentials would be rejected on the next + # handshake; treat expired creds as a fetch failure rather than minting a doomed cert. + logging.error( + "Credential source returned already-expired credentials (expired %s); not " + "generating a certificate", + credentials_expiration, + ) + return None, None + elif security_credentials: + logging.debug( + "No expiration available for these credentials; certificate refresh falls back to " + "the fixed interval" + ) ca_extension_body = ca_extension_builder( ap_id, security_credentials, fs_id, client_info @@ -2001,7 +2117,7 @@ def create_ca_conf( "Failed to create AWS SigV4 signature section for OpenSSL config. Public Key path: %s", public_key_path, ) - return None + return None, None efs_client_info_body = efs_client_info_builder(client_info) if client_info else "" full_config_body = CA_CONFIG_BODY % ( directory, @@ -2015,7 +2131,7 @@ def create_ca_conf( with open(config_path, "w") as f: f.write(full_config_body) - return full_config_body + return full_config_body, credentials_expiration def ca_extension_builder(ap_id, security_credentials, fs_id, client_info): @@ -2084,7 +2200,7 @@ def subprocess_call(cmd, error_message): process = subprocess.Popen( cmd.split(), stdout=subprocess.PIPE, stderr=subprocess.PIPE, close_fds=True ) - (output, err) = process.communicate() + output, err = process.communicate() rc = process.poll() if rc != 0: logging.debug( diff --git a/test/global_test/test_watchdog_common_duplication_match.py b/test/global_test/test_watchdog_common_duplication_match.py new file mode 100644 index 00000000..fc342d25 --- /dev/null +++ b/test/global_test/test_watchdog_common_duplication_match.py @@ -0,0 +1,21 @@ +# +# Copyright 2017-2018 Amazon.com, Inc. and its affiliates. All Rights Reserved. +# +# Licensed under the MIT License. See the LICENSE accompanying this file +# for the specific language governing permissions and limitations under +# the License. +# +# The standalone watchdog cannot import from efs_utils_common (its install location has no access +# to the shared package), so the credential "Expiration" format is deliberately duplicated in both +# src/watchdog/__init__.py and src/efs_utils_common.aws_credentials. This test imports both copies +# and asserts they are equal, so the duplication cannot drift silently. + +import efs_utils_common.aws_credentials as aws_credentials +import watchdog + + +def test_credentials_expiration_datetime_format_match(): + assert ( + watchdog.CREDENTIALS_EXPIRATION_DATETIME_FORMAT + == aws_credentials.CREDENTIALS_EXPIRATION_DATETIME_FORMAT + ), "watchdog and efs_utils_common.aws_credentials disagree on the credential Expiration format" diff --git a/test/mount_common_test/test_bootstrap_proxy.py b/test/mount_common_test/test_bootstrap_proxy.py index 68430aa6..df57f8ba 100644 --- a/test/mount_common_test/test_bootstrap_proxy.py +++ b/test/mount_common_test/test_bootstrap_proxy.py @@ -221,6 +221,175 @@ def config_get_side_effect(section, field): assert os.path.exists(pk_path) +def test_bootstrap_proxy_persists_credential_expiration_for_iam_mount(mocker, tmpdir): + write_config_mock = setup_mocks_without_popen(mocker) + write_state_mock = mocker.patch( + "efs_utils_common.proxy.write_tunnel_state_file", return_value="~mocktempfile" + ) + mocker.patch( + "efs_utils_common.proxy.get_mount_specific_filename", return_value=DNS_NAME + ) + mocker.patch("efs_utils_common.proxy.get_target_region", return_value=REGION) + mocker.patch("efs_utils_common.proxy.is_ocsp_enabled", return_value=False) + mocker.patch("efs_utils_common.proxy.create_certificate", return_value="dummytime") + mocker.patch( + "efs_utils_common.proxy._efs_proxy_bin", return_value="/usr/bin/efs-proxy" + ) + pk_path = os.path.join(str(tmpdir), "privateKey.pem") + mocker.patch("efs_utils_common.proxy.get_private_key_path", return_value=pk_path) + + expiration = "2026-08-06T00:38:37Z" # whole-second ISO-8601 shape IMDS returns + mocker.patch( + "efs_utils_common.proxy.get_aws_security_credentials", + return_value=( + { + "AccessKeyId": "AKID", + "SecretAccessKey": "SECRET", + "Token": "TOKEN", + "Expiration": expiration, + }, + "metadata:", + ), + ) + + MOCK_CONFIG.get.side_effect = None + MOCK_CONFIG.get.return_value = "info" + MOCK_CONFIG.getboolean.return_value = True + + try: + with proxy.bootstrap_proxy( + MOCK_CONFIG, + INIT_SYSTEM, + DNS_NAME, + FS_ID, + MOUNT_POINT, + {"tls": None, "iam": None}, + str(tmpdir), + ): + pass + except OSError: + pass + + assert write_state_mock.called + cert_details = write_state_mock.call_args.kwargs["cert_details"] + assert cert_details["certificateExpirationTime"] == expiration + assert write_config_mock.called + + +def test_bootstrap_proxy_drops_malformed_credential_expiration(mocker, tmpdir, caplog): + import logging + + caplog.set_level(logging.WARNING) + setup_mocks_without_popen(mocker) + write_state_mock = mocker.patch( + "efs_utils_common.proxy.write_tunnel_state_file", return_value="~mocktempfile" + ) + mocker.patch( + "efs_utils_common.proxy.get_mount_specific_filename", return_value=DNS_NAME + ) + mocker.patch("efs_utils_common.proxy.get_target_region", return_value=REGION) + mocker.patch("efs_utils_common.proxy.is_ocsp_enabled", return_value=False) + mocker.patch("efs_utils_common.proxy.create_certificate", return_value="dummytime") + mocker.patch( + "efs_utils_common.proxy._efs_proxy_bin", return_value="/usr/bin/efs-proxy" + ) + pk_path = os.path.join(str(tmpdir), "privateKey.pem") + mocker.patch("efs_utils_common.proxy.get_private_key_path", return_value=pk_path) + + mocker.patch( + "efs_utils_common.proxy.get_aws_security_credentials", + return_value=( + { + "AccessKeyId": "AKID", + "SecretAccessKey": "SECRET", + "Token": "TOKEN", + "Expiration": "1786127616", # epoch, not ISO-8601 + }, + "metadata:", + ), + ) + + MOCK_CONFIG.get.side_effect = None + MOCK_CONFIG.get.return_value = "info" + MOCK_CONFIG.getboolean.return_value = True + + try: + with proxy.bootstrap_proxy( + MOCK_CONFIG, + INIT_SYSTEM, + DNS_NAME, + FS_ID, + MOUNT_POINT, + {"tls": None, "iam": None}, + str(tmpdir), + ): + pass + except OSError: + pass + + assert write_state_mock.called + cert_details = write_state_mock.call_args.kwargs["cert_details"] + assert "certificateExpirationTime" not in cert_details + assert "is not valid ISO-8601" in caplog.text + + +def test_bootstrap_proxy_no_expiration_logs_debug(mocker, tmpdir, caplog): + import logging + + caplog.set_level(logging.DEBUG) + setup_mocks_without_popen(mocker) + write_state_mock = mocker.patch( + "efs_utils_common.proxy.write_tunnel_state_file", return_value="~mocktempfile" + ) + mocker.patch( + "efs_utils_common.proxy.get_mount_specific_filename", return_value=DNS_NAME + ) + mocker.patch("efs_utils_common.proxy.get_target_region", return_value=REGION) + mocker.patch("efs_utils_common.proxy.is_ocsp_enabled", return_value=False) + mocker.patch("efs_utils_common.proxy.create_certificate", return_value="dummytime") + mocker.patch( + "efs_utils_common.proxy._efs_proxy_bin", return_value="/usr/bin/efs-proxy" + ) + pk_path = os.path.join(str(tmpdir), "privateKey.pem") + mocker.patch("efs_utils_common.proxy.get_private_key_path", return_value=pk_path) + + mocker.patch( + "efs_utils_common.proxy.get_aws_security_credentials", + return_value=( + { + "AccessKeyId": "AKID", + "SecretAccessKey": "SECRET", + "Token": "TOKEN", + # no Expiration key + }, + "metadata:", + ), + ) + + MOCK_CONFIG.get.side_effect = None + MOCK_CONFIG.get.return_value = "info" + MOCK_CONFIG.getboolean.return_value = True + + try: + with proxy.bootstrap_proxy( + MOCK_CONFIG, + INIT_SYSTEM, + DNS_NAME, + FS_ID, + MOUNT_POINT, + {"tls": None, "iam": None}, + str(tmpdir), + ): + pass + except OSError: + pass + + assert write_state_mock.called + cert_details = write_state_mock.call_args.kwargs["cert_details"] + assert "certificateExpirationTime" not in cert_details + assert "No credential expiration available" in caplog.text + + def test_bootstrap_proxy_cert_not_created_non_tls_mount(mocker, tmpdir): setup_mocks_without_popen(mocker) mocker.patch( diff --git a/test/mount_common_test/test_get_aws_security_credentials.py b/test/mount_common_test/test_get_aws_security_credentials.py index 7cf0a000..2fab9731 100644 --- a/test/mount_common_test/test_get_aws_security_credentials.py +++ b/test/mount_common_test/test_get_aws_security_credentials.py @@ -392,6 +392,81 @@ def test_get_aws_security_credentials_botocore_present_get_assumed_profile_crede utils.assert_called(botocore_get_assumed_profile_credentials_mock) +def _mock_botocore_session(mocker, expiry_time): + # Build a fake botocore session whose credentials object mimics RefreshableCredentials + # (carries _expiry_time) or plain Credentials (no _expiry_time) when expiry_time is None. + frozen = mocker.MagicMock() + frozen.access_key = ACCESS_KEY_ID_VAL + frozen.secret_key = SECRET_ACCESS_KEY_VAL + frozen.token = SESSION_TOKEN_VAL + + creds_object = mocker.MagicMock() + creds_object.get_frozen_credentials.return_value = frozen + if expiry_time is None: + # Plain Credentials have no _expiry_time attribute at all. + del creds_object._expiry_time + else: + creds_object._expiry_time = expiry_time + + session = mocker.MagicMock() + session.get_credentials.return_value = creds_object + + aws_credentials.BOTOCORE_PRESENT = True + mocker.patch("botocore.session.get_session", return_value=session) + + +def test_botocore_credentials_helper_surfaces_expiration(mocker): + # An assumed-role/session profile carries an expiry (tz-aware UTC datetime); it must be + # surfaced as an ISO-8601 Expiration string so the watchdog can refresh the cert ahead of it. + from datetime import datetime, timezone + + _mock_botocore_session(mocker, datetime(2026, 8, 6, 0, 38, 37, tzinfo=timezone.utc)) + + credentials = aws_credentials.botocore_credentials_helper("test-profile") + + assert credentials["AccessKeyId"] == ACCESS_KEY_ID_VAL + assert credentials["Expiration"] == "2026-08-06T00:38:37Z" + + +def test_botocore_credentials_helper_no_expiration_for_static_profile(mocker): + # A static profile has no expiry; no Expiration key is added and the caller falls back to + # the fixed renewal interval. + _mock_botocore_session(mocker, None) + + credentials = aws_credentials.botocore_credentials_helper("test-profile") + + assert credentials["AccessKeyId"] == ACCESS_KEY_ID_VAL + assert "Expiration" not in credentials + + +def test_botocore_credentials_helper_ignores_non_datetime_expiry(mocker): + # Defensive: _expiry_time is a private botocore attribute; if it is ever not a datetime, skip + # surfacing Expiration (fall back to the fixed interval) rather than crash on strftime. + _mock_botocore_session(mocker, "not-a-datetime") + + credentials = aws_credentials.botocore_credentials_helper("test-profile") + + assert credentials["AccessKeyId"] == ACCESS_KEY_ID_VAL + assert "Expiration" not in credentials + + +def test_is_valid_credentials_expiration(): + # Whole-second ISO-8601 (the format every credential source and STS use) is accepted; absent, + # non-ISO, fractional-second, and epoch values are rejected so callers do not persist a value + # the watchdog cannot parse. + assert ( + aws_credentials.is_valid_credentials_expiration("2026-08-06T00:38:37Z") is True + ) + assert ( + aws_credentials.is_valid_credentials_expiration("2026-08-06T00:38:37.123Z") + is False + ) + assert aws_credentials.is_valid_credentials_expiration(None) is False + assert aws_credentials.is_valid_credentials_expiration("") is False + assert aws_credentials.is_valid_credentials_expiration("not-a-timestamp") is False + assert aws_credentials.is_valid_credentials_expiration("1786127616") is False + + def test_get_aws_security_credentials_credentials_not_found_in_aws_creds_uri( mocker, capsys ): @@ -623,6 +698,45 @@ def test_get_aws_security_credentials_from_webidentity_real_builds_sts_url_and_p ) +def test_get_aws_security_credentials_from_webidentity_carries_expiration(mocker): + # The STS Credentials.Expiration must be surfaced so the watchdog can schedule the cert + # refresh ahead of it. When STS omits Expiration, the returned value is None. + config = get_fake_config_with_dns_suffix("amazonaws.com") + mocker.patch("builtins.open", mocker.mock_open(read_data="FAKE_WEB_IDENTITY_JWT")) + + response = _well_formed_webidentity_response() + expiration = "2019-10-25T21:17:24Z" + response["AssumeRoleWithWebIdentityResponse"]["AssumeRoleWithWebIdentityResult"][ + "Credentials" + ]["Expiration"] = expiration + mocker.patch( + "efs_utils_common.aws_credentials.url_request_helper", return_value=response + ) + + credentials, _ = aws_credentials.get_aws_security_credentials_from_webidentity( + config, + WEB_IDENTITY_ROLE_ARN, + WEB_IDENTITY_TOKEN_FILE, + "us-east-1", + is_fatal=False, + ) + assert credentials["Expiration"] == expiration + + # No Expiration in the STS response => Expiration key present but None. + mocker.patch( + "efs_utils_common.aws_credentials.url_request_helper", + return_value=_well_formed_webidentity_response(), + ) + credentials, _ = aws_credentials.get_aws_security_credentials_from_webidentity( + config, + WEB_IDENTITY_ROLE_ARN, + WEB_IDENTITY_TOKEN_FILE, + "us-east-1", + is_fatal=False, + ) + assert credentials["Expiration"] is None + + def test_get_aws_security_credentials_from_webidentity_real_is_fatal_failure( mocker, capsys ): diff --git a/test/mount_common_test/test_helper_function.py b/test/mount_common_test/test_helper_function.py index 590f596b..26e36098 100644 --- a/test/mount_common_test/test_helper_function.py +++ b/test/mount_common_test/test_helper_function.py @@ -445,6 +445,10 @@ def test_get_assumed_profile_credentials_via_botocore_botocore_present(mocker): get_credential_session_mock = MagicMock() boto_session_mock.get_credentials.return_value = get_credential_session_mock get_credential_session_mock.get_frozen_credentials.return_value = frozen_credentials + # Model a static (non-refreshable) profile: base botocore Credentials have no + # _expiry_time attribute, so no Expiration is surfaced. Without this delete, MagicMock + # would auto-create a truthy _expiry_time. + del get_credential_session_mock._expiry_time mocker.patch("botocore.session.get_session", return_value=boto_session_mock) diff --git a/test/mount_common_test/test_start_watchdog.py b/test/mount_common_test/test_start_watchdog.py index b3eca23e..6fe7479c 100644 --- a/test/mount_common_test/test_start_watchdog.py +++ b/test/mount_common_test/test_start_watchdog.py @@ -6,6 +6,7 @@ from unittest.mock import MagicMock import efs_utils_common.proxy as proxy +from efs_utils_common.constants import WATCHDOG_SERVICE, WATCHDOG_SERVICE_PLIST_PATH from .. import utils @@ -41,22 +42,58 @@ def test_systemd_system(mocker): assert "start" in popen_mock.call_args[0][0] -def test_launchd_system(mocker): - process_mock = MagicMock() - process_mock.communicate.return_value = ( - "stop", - "", - ) - process_mock.returncode = 0 - popen_mock = mocker.patch("subprocess.Popen", return_value=process_mock) +def test_launchd_canonical_label_loaded(mocker): + call_mock = mocker.patch("subprocess.call", side_effect=[0]) + + proxy.start_watchdog("launchd") + + utils.assert_called_once(call_mock) + assert call_mock.call_args_list[0][0][0] == [ + "launchctl", + "list", + WATCHDOG_SERVICE, + ] + + +def test_launchd_loads_canonical_plist(mocker): + call_mock = mocker.patch("subprocess.call", side_effect=[1, 0]) mocker.patch("os.path.exists", return_value=True) proxy.start_watchdog("launchd") - assert 2 == popen_mock.call_count - assert "sudo" in popen_mock.call_args[0][0] - assert "launchctl" in popen_mock.call_args[0][0] - assert "load" in popen_mock.call_args[0][0] + assert call_mock.call_args_list[-1][0][0] == [ + "launchctl", + "load", + WATCHDOG_SERVICE_PLIST_PATH, + ] + + +def test_launchd_missing_plists_is_fatal(mocker): + call_mock = mocker.patch("subprocess.call", side_effect=[1]) + mocker.patch("os.path.exists", return_value=False) + fatal_mock = mocker.patch("efs_utils_common.proxy.fatal_error") + + proxy.start_watchdog("launchd") + + utils.assert_called_once(call_mock) + utils.assert_called_once(fatal_mock) + assert "plist" in fatal_mock.call_args[0][0].lower() + + +def test_launchd_load_failure_is_fatal(mocker): + call_mock = mocker.patch("subprocess.call", side_effect=[1, 23]) + mocker.patch("os.path.exists", return_value=True) + fatal_mock = mocker.patch("efs_utils_common.proxy.fatal_error") + + proxy.start_watchdog("launchd") + + assert call_mock.call_args_list[-1][0][0] == [ + "launchctl", + "load", + WATCHDOG_SERVICE_PLIST_PATH, + ] + utils.assert_called_once(fatal_mock) + assert "exit code 23" in fatal_mock.call_args[0][0] def test_unknown_system(mocker): diff --git a/test/watchdog_test/test_helper_function.py b/test/watchdog_test/test_helper_function.py index 75fc77e3..d2725348 100644 --- a/test/watchdog_test/test_helper_function.py +++ b/test/watchdog_test/test_helper_function.py @@ -294,6 +294,7 @@ def test_get_assumed_profile_credentials_via_botocore_botocore_present(mocker): get_credential_session_mock = MagicMock() boto_session_mock.get_credentials.return_value = get_credential_session_mock get_credential_session_mock.get_frozen_credentials.return_value = frozen_credentials + get_credential_session_mock._expiry_time = None # static profile: no expiry mocker.patch("botocore.session.get_session", return_value=boto_session_mock) @@ -307,6 +308,56 @@ def test_get_assumed_profile_credentials_via_botocore_botocore_present(mocker): get_credential_session_mock.get_frozen_credentials.assert_called_once_with() +def test_get_assumed_profile_credentials_via_botocore_surfaces_expiration(mocker): + # A temporary (assumed-role/session) profile carries _expiry_time; the watchdog helper surfaces + # it as an ISO-8601 Expiration so the cert refresh can schedule ahead of it. + from datetime import datetime, timezone + + boto_session_mock = MagicMock() + ReadOnlyCredentials = namedtuple( + "ReadOnlyCredentials", ["access_key", "secret_key", "token"] + ) + frozen_credentials = ReadOnlyCredentials( + ACCESS_KEY_ID_VAL, SECRET_ACCESS_KEY_VAL, SESSION_TOKEN_VAL + ) + get_credential_session_mock = MagicMock() + get_credential_session_mock.get_frozen_credentials.return_value = frozen_credentials + get_credential_session_mock._expiry_time = datetime( + 2026, 8, 6, 0, 38, 37, tzinfo=timezone.utc + ) + boto_session_mock.get_credentials.return_value = get_credential_session_mock + mocker.patch("botocore.session.get_session", return_value=boto_session_mock) + + credentials = watchdog.botocore_credentials_helper("test_profile") + + assert credentials["AccessKeyId"] == ACCESS_KEY_ID_VAL + assert credentials["Expiration"] == "2026-08-06T00:38:37Z" + + +def test_get_assumed_profile_credentials_via_botocore_ignores_non_datetime_expiry( + mocker, +): + # Defensive: if botocore's private _expiry_time is ever not a datetime, skip surfacing + # Expiration rather than crash on strftime. + boto_session_mock = MagicMock() + ReadOnlyCredentials = namedtuple( + "ReadOnlyCredentials", ["access_key", "secret_key", "token"] + ) + frozen_credentials = ReadOnlyCredentials( + ACCESS_KEY_ID_VAL, SECRET_ACCESS_KEY_VAL, SESSION_TOKEN_VAL + ) + get_credential_session_mock = MagicMock() + get_credential_session_mock.get_frozen_credentials.return_value = frozen_credentials + get_credential_session_mock._expiry_time = "not-a-datetime" + boto_session_mock.get_credentials.return_value = get_credential_session_mock + mocker.patch("botocore.session.get_session", return_value=boto_session_mock) + + credentials = watchdog.botocore_credentials_helper("test_profile") + + assert credentials["AccessKeyId"] == ACCESS_KEY_ID_VAL + assert "Expiration" not in credentials + + def test_get_assumed_profile_credentials_via_botocore_botocore_present_profile_not_found( mocker, ): diff --git a/test/watchdog_test/test_refresh_self_signed_certificate.py b/test/watchdog_test/test_refresh_self_signed_certificate.py index b3675d0f..6ac135f6 100644 --- a/test/watchdog_test/test_refresh_self_signed_certificate.py +++ b/test/watchdog_test/test_refresh_self_signed_certificate.py @@ -194,7 +194,7 @@ def _create_ca_conf_helper( credentials = "dummy:lookup" if iam else None ap_id = AP_ID if ap else None client_info = CLIENT_INFO if client_info else None - full_config_body = watchdog.create_ca_conf( + full_config_body, _ = watchdog.create_ca_conf( config, tls_dict["certificate_path"], COMMON_NAME, @@ -655,7 +655,7 @@ def _test_recreate_certificate_with_valid_client_source_config( with open(os.path.join(tls_dict["mount_dir"], "config.conf")) as f: conf_body = f.read() - assert conf_body == watchdog.create_ca_conf( + recreated_conf_body, _ = watchdog.create_ca_conf( config, tmp_config_path, COMMON_NAME, @@ -669,6 +669,7 @@ def _test_recreate_certificate_with_valid_client_source_config( AP_ID, expected_client_info, ) + assert conf_body == recreated_conf_body assert os.path.exists(pk_path) assert os.path.exists(os.path.join(tls_dict["mount_dir"], "publicKey.pem")) assert os.path.exists(os.path.join(tls_dict["mount_dir"], "request.csr")) @@ -709,7 +710,7 @@ def _test_recreate_certificate_with_invalid_client_source_config( with open(os.path.join(tls_dict["mount_dir"], "config.conf")) as f: conf_body = f.read() - assert conf_body == watchdog.create_ca_conf( + recreated_conf_body, _ = watchdog.create_ca_conf( config, tmp_config_path, COMMON_NAME, @@ -723,6 +724,7 @@ def _test_recreate_certificate_with_invalid_client_source_config( AP_ID, expected_client_info, ) + assert conf_body == recreated_conf_body assert os.path.exists(pk_path) assert os.path.exists(os.path.join(tls_dict["mount_dir"], "publicKey.pem")) assert os.path.exists(os.path.join(tls_dict["mount_dir"], "request.csr")) @@ -954,3 +956,274 @@ def test_check_and_create_private_key_key_already_exists(mocker, tmpdir): state_file_dir = str(tmpdir) watchdog.check_and_create_private_key(state_file_dir) assert call_mock.call_count == 0 + + +# ---- Credential-expiration-driven certificate refresh ---- + +EXPIRATION_FORMAT = watchdog.CREDENTIALS_EXPIRATION_DATETIME_FORMAT +MARGIN = timedelta(minutes=watchdog.CERT_EXPIRATION_SAFETY_MARGIN_MIN) + + +def test_parse_credentials_expiration_valid(): + parsed = watchdog.parse_credentials_expiration(FIXED_DT.strftime(EXPIRATION_FORMAT)) + assert parsed == FIXED_DT + + +def test_parse_credentials_expiration_real_imds_value(): + # Exact whole-second ISO-8601 shape a real IMDS credential returns. + assert watchdog.parse_credentials_expiration("2026-08-06T00:38:37Z") == datetime( + 2026, 8, 6, 0, 38, 37, tzinfo=timezone.utc + ) + + +def test_parse_credentials_expiration_absent(): + assert watchdog.parse_credentials_expiration(None) is None + assert watchdog.parse_credentials_expiration("") is None + + +def test_parse_credentials_expiration_malformed(caplog): + # Values are validated before being persisted, so a malformed value reaching the parser + # (absent, or an externally corrupted state file) returns None silently and the caller falls + # back to the fixed interval; the warning lives at the persist sites, not here. + caplog.set_level(logging.WARNING) + assert watchdog.parse_credentials_expiration("not-a-timestamp") is None + assert caplog.text == "" + + +def test_refresh_deadline_no_expiration_uses_fixed_interval(): + deadline = watchdog.get_certificate_refresh_deadline(FIXED_DT, 60 * 60, None) + assert deadline == FIXED_DT + timedelta(minutes=60) + + +def test_refresh_deadline_far_expiration_uses_fixed_interval(): + far_expiration = FIXED_DT + timedelta(hours=6) + deadline = watchdog.get_certificate_refresh_deadline( + FIXED_DT, 60 * 60, far_expiration + ) + assert deadline == FIXED_DT + timedelta(minutes=60) + + +def test_refresh_deadline_near_expiration_refreshes_before_expiry(): + near_expiration = FIXED_DT + timedelta(minutes=15) + deadline = watchdog.get_certificate_refresh_deadline( + FIXED_DT, 60 * 60, near_expiration + ) + assert deadline == near_expiration - MARGIN + + +def test_refresh_deadline_already_expired_token_is_in_the_past(): + past_expiration = FIXED_DT - timedelta(minutes=13) + deadline = watchdog.get_certificate_refresh_deadline( + FIXED_DT, 60 * 60, past_expiration + ) + assert deadline == past_expiration - MARGIN + + +def test_refresh_deadline_capped_at_cert_ttl(): + # A renewal interval longer than the cert's own NOT_AFTER_HOURS lifetime would push the refresh + # past the cert's expiry; it is capped at NOT_AFTER - margin. + deadline = watchdog.get_certificate_refresh_deadline( + FIXED_DT, 4 * 60 * 60, None # 4h interval exceeds the 3h cert TTL + ) + assert deadline == FIXED_DT + timedelta(hours=watchdog.NOT_AFTER_HOURS) - MARGIN + + +def test_refresh_deadline_cert_ttl_cap_beats_far_expiration(): + far_expiration = FIXED_DT + timedelta(hours=6) + deadline = watchdog.get_certificate_refresh_deadline( + FIXED_DT, 4 * 60 * 60, far_expiration + ) + assert deadline == FIXED_DT + timedelta(hours=watchdog.NOT_AFTER_HOURS) - MARGIN + + +def test_check_certificate_refreshes_when_token_near_expiry(mocker, tmpdir, caplog): + # Cert created 20 min ago (well within the 60-min fixed interval, so fixed-interval mode would + # NOT refresh), but the embedded token expires in 3 min => inside the 5-min safety margin, + # so the refresh deadline is already past and the cert must be refreshed now. + caplog.set_level(logging.DEBUG) + mocker.patch("watchdog.get_utc_now", return_value=FIXED_DT) + config = _get_config() + pk_path = _get_mock_private_key_path(mocker, tmpdir) + created = (FIXED_DT - timedelta(minutes=20)).strftime(DT_PATTERN) + tls_dict = watchdog.tls_paths_dictionary(MOUNT_NAME, str(tmpdir)) + state = _create_certificate_and_state( + tls_dict, + str(tmpdir), + pk_path, + created, + security_credentials=CREDENTIALS, + credentials_source=CREDENTIALS_SOURCE, + ap_id=AP_ID, + ) + state["certificateExpirationTime"] = (FIXED_DT + timedelta(minutes=3)).strftime( + EXPIRATION_FORMAT + ) + + fresh_expiration = (FIXED_DT + timedelta(hours=1)).strftime(EXPIRATION_FORMAT) + fresh_credentials = dict(CREDENTIALS, Expiration=fresh_expiration) + mocker.patch( + "watchdog.get_aws_security_credentials", return_value=fresh_credentials + ) + + watchdog.check_certificate( + config, state, str(tmpdir), STATE_FILE, SERVICE, base_path=str(tmpdir) + ) + + with open(os.path.join(str(tmpdir), STATE_FILE), "r") as state_json: + state = json.load(state_json) + + assert datetime.strptime( + state["certificateCreationTime"], DT_PATTERN + ) > datetime.strptime(created, DT_PATTERN) + assert state["certificateExpirationTime"] == fresh_expiration + assert "early: scheduled deadline" in caplog.text + + +def test_check_certificate_does_not_refresh_when_token_far_from_expiry(mocker, tmpdir): + mocker.patch("watchdog.get_utc_now", return_value=FIXED_DT) + config = _get_config() + pk_path = _get_mock_private_key_path(mocker, tmpdir) + created = (FIXED_DT - timedelta(minutes=20)).strftime(DT_PATTERN) + tls_dict = watchdog.tls_paths_dictionary(MOUNT_NAME, str(tmpdir)) + state = _create_certificate_and_state( + tls_dict, str(tmpdir), pk_path, created, ap_id=AP_ID + ) + state["certificateExpirationTime"] = (FIXED_DT + timedelta(hours=6)).strftime( + EXPIRATION_FORMAT + ) + + recreate_mock = mocker.patch("watchdog.recreate_certificate") + + watchdog.check_certificate( + config, state, str(tmpdir), STATE_FILE, SERVICE, base_path=str(tmpdir) + ) + + recreate_mock.assert_not_called() + + +def test_check_certificate_clears_stale_expiration_when_none_returned( + mocker, tmpdir, caplog +): + caplog.set_level(logging.DEBUG) + mocker.patch("watchdog.get_utc_now", return_value=FIXED_DT) + config = _get_config() + pk_path = _get_mock_private_key_path(mocker, tmpdir) + created = (FIXED_DT - timedelta(minutes=90)).strftime( + DT_PATTERN + ) # past the interval + tls_dict = watchdog.tls_paths_dictionary(MOUNT_NAME, str(tmpdir)) + state = _create_certificate_and_state( + tls_dict, + str(tmpdir), + pk_path, + created, + security_credentials=CREDENTIALS, + credentials_source=CREDENTIALS_SOURCE, + ap_id=AP_ID, + ) + state["certificateExpirationTime"] = (FIXED_DT + timedelta(minutes=5)).strftime( + EXPIRATION_FORMAT + ) + + mocker.patch("watchdog.get_aws_security_credentials", return_value=CREDENTIALS) + + watchdog.check_certificate( + config, state, str(tmpdir), STATE_FILE, SERVICE, base_path=str(tmpdir) + ) + + with open(os.path.join(str(tmpdir), STATE_FILE), "r") as state_json: + state = json.load(state_json) + + assert "certificateExpirationTime" not in state + assert "No expiration available for these credentials" in caplog.text + + +def _setup_ca_conf_dir(tmpdir): + tls_dict = certificate_utils.tls_paths_dictionary(MOUNT_NAME, str(tmpdir)) + file_utils.create_required_directory({}, tls_dict["mount_dir"]) + with open(os.path.join(tls_dict["mount_dir"], "publicKey.pem"), "w") as f: + f.write(PUBLIC_KEY_BODY) + return tls_dict + + +def test_create_ca_conf_returns_credential_expiration(mocker, tmpdir): + mocker.patch("watchdog.get_utc_now", return_value=FIXED_DT) + expiration = (FIXED_DT + timedelta(hours=1)).strftime(EXPIRATION_FORMAT) + mocker.patch( + "watchdog.get_aws_security_credentials", + return_value=dict(CREDENTIALS, Expiration=expiration), + ) + tls_dict = _setup_ca_conf_dir(tmpdir) + config_body, credentials_expiration = watchdog.create_ca_conf( + _get_config(), + os.path.join(tls_dict["mount_dir"], "config.conf"), + COMMON_NAME, + tls_dict["mount_dir"], + os.path.join(tls_dict["mount_dir"], "privateKey.pem"), + FIXED_DT, + REGION, + FS_ID, + "dummy:lookup", + SERVICE, + AP_ID, + CLIENT_INFO, + ) + assert config_body + assert credentials_expiration == expiration + + +def test_create_ca_conf_rejects_already_expired_credentials(mocker, tmpdir, caplog): + caplog.set_level(logging.ERROR) + mocker.patch("watchdog.get_utc_now", return_value=FIXED_DT) + expired = (FIXED_DT - timedelta(minutes=1)).strftime(EXPIRATION_FORMAT) + mocker.patch( + "watchdog.get_aws_security_credentials", + return_value=dict(CREDENTIALS, Expiration=expired), + ) + tls_dict = _setup_ca_conf_dir(tmpdir) + config_body, credentials_expiration = watchdog.create_ca_conf( + _get_config(), + os.path.join(tls_dict["mount_dir"], "config.conf"), + COMMON_NAME, + tls_dict["mount_dir"], + os.path.join(tls_dict["mount_dir"], "privateKey.pem"), + FIXED_DT, + REGION, + FS_ID, + "dummy:lookup", + SERVICE, + AP_ID, + CLIENT_INFO, + ) + assert config_body is None + assert credentials_expiration is None + assert "already-expired credentials" in caplog.text + + +def test_create_ca_conf_drops_malformed_expiration(mocker, tmpdir, caplog): + # A present-but-unparseable expiration does not block minting; it is dropped (returned as None) + # with a warning so the caller schedules on the fixed interval instead. + caplog.set_level(logging.WARNING) + mocker.patch("watchdog.get_utc_now", return_value=FIXED_DT) + mocker.patch( + "watchdog.get_aws_security_credentials", + return_value=dict(CREDENTIALS, Expiration="1786127616"), # epoch, not ISO-8601 + ) + tls_dict = _setup_ca_conf_dir(tmpdir) + config_body, credentials_expiration = watchdog.create_ca_conf( + _get_config(), + os.path.join(tls_dict["mount_dir"], "config.conf"), + COMMON_NAME, + tls_dict["mount_dir"], + os.path.join(tls_dict["mount_dir"], "privateKey.pem"), + FIXED_DT, + REGION, + FS_ID, + "dummy:lookup", + SERVICE, + AP_ID, + CLIENT_INFO, + ) + assert config_body + assert credentials_expiration is None + assert "is not valid ISO-8601" in caplog.text