diff --git a/Cargo.lock b/Cargo.lock index 64c5c61875..7334be0337 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -97,56 +97,12 @@ dependencies = [ "libc", ] -[[package]] -name = "anstream" -version = "1.0.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "824a212faf96e9acacdbd09febd34438f8f711fb84e09a8916013cd7815ca28d" -dependencies = [ - "anstyle", - "anstyle-parse", - "anstyle-query", - "anstyle-wincon", - "colorchoice", - "is_terminal_polyfill", - "utf8parse", -] - [[package]] name = "anstyle" version = "1.0.14" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "940b3a0ca603d1eade50a4846a2afffd5ef57a9feac2c0e2ec2e14f9ead76000" -[[package]] -name = "anstyle-parse" -version = "1.0.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "52ce7f38b242319f7cabaa6813055467063ecdc9d355bbb4ce0c68908cd8130e" -dependencies = [ - "utf8parse", -] - -[[package]] -name = "anstyle-query" -version = "1.1.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "40c48f72fd53cd289104fc64099abca73db4166ad86ea0b4341abe65af83dadc" -dependencies = [ - "windows-sys 0.61.2", -] - -[[package]] -name = "anstyle-wincon" -version = "3.0.11" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "291e6a250ff86cd4a820112fb8898808a366d8f9f58ce16d1f538353ad55747d" -dependencies = [ - "anstyle", - "once_cell_polyfill", - "windows-sys 0.61.2", -] - [[package]] name = "anyhow" version = "1.0.103" @@ -511,6 +467,16 @@ dependencies = [ "syn", ] +[[package]] +name = "asyncband" +version = "0.6.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "94a214ba60d6231afd0e805e3c27c45a1626d9debaa5a5061c45a1ea1b2f1ed0" +dependencies = [ + "hashbrown 0.17.1", + "slab", +] + [[package]] name = "atoi" version = "2.0.0" @@ -1340,46 +1306,6 @@ dependencies = [ "inout", ] -[[package]] -name = "clap" -version = "4.6.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1ddb117e43bbf7dacf0a4190fef4d345b9bad68dfc649cb349e7d17d28428e51" -dependencies = [ - "clap_builder", - "clap_derive", -] - -[[package]] -name = "clap_builder" -version = "4.6.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "714a53001bf66416adb0e2ef5ac857140e7dc3a0c48fb28b2f10762fc4b5069f" -dependencies = [ - "anstream", - "anstyle", - "clap_lex", - "strsim", -] - -[[package]] -name = "clap_derive" -version = "4.6.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f2ce8604710f6733aa641a2b3731eaa1e8b3d9973d5e3565da11800813f997a9" -dependencies = [ - "heck", - "proc-macro2", - "quote", - "syn", -] - -[[package]] -name = "clap_lex" -version = "1.1.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c8d4a3bb8b1e0c1050499d1815f5ab16d04f0959b233085fb31653fbfc9d98f9" - [[package]] name = "cmake" version = "0.1.58" @@ -1395,12 +1321,6 @@ version = "0.5.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0c9ea0ac24bc397ab3c98583a3c9ba74fa56b09a4449bbe172b9b1ddb016027a" -[[package]] -name = "colorchoice" -version = "1.0.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1d07550c9036bf2ae0c684c4297d503f838287c83c53686d05370d0e139ae570" - [[package]] name = "colored" version = "3.1.1" @@ -1649,16 +1569,6 @@ dependencies = [ "memchr", ] -[[package]] -name = "ctor" -version = "0.6.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "424e0138278faeb2b401f174ad17e715c829512d74f3d1e81eb43365c2e0590e" -dependencies = [ - "ctor-proc-macro", - "dtor", -] - [[package]] name = "ctor" version = "1.0.8" @@ -1669,12 +1579,6 @@ dependencies = [ "linktime-proc-macro", ] -[[package]] -name = "ctor-proc-macro" -version = "0.0.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "52560adf09603e58c9a7ee1fe1dcb95a16927b17c127f0ac02d6e768a0e25bc1" - [[package]] name = "ctr" version = "0.9.2" @@ -1947,21 +1851,6 @@ version = "0.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1435fa1053d8b2fbbe9be7e97eca7f33d37b28409959813daefc1446a14247f1" -[[package]] -name = "dtor" -version = "0.1.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "404d02eeb088a82cfd873006cb713fe411306c7d182c344905e101fb1167d301" -dependencies = [ - "dtor-proc-macro", -] - -[[package]] -name = "dtor-proc-macro" -version = "0.0.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f678cf4a922c215c63e0de95eb1ff08a958a81d47e485cf9da1e27bf6305cfa5" - [[package]] name = "dunce" version = "1.0.5" @@ -2457,18 +2346,21 @@ checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70" [[package]] name = "hf-xet" -version = "1.5.2" +version = "1.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "430b33fa84f92796d4d263070b6c0d3ca219df7b9a0e1853ee431029b1612bcd" +checksum = "c237ef4fb0ce1962a5117f8bd8c74454b41629826a9df17d14a1840ca18f0754" dependencies = [ + "anyhow", "async-trait", "bytes", "http 1.5.0", "more-asserts", "serde", + "serde_json", "thiserror 2.0.18", "tokio", "tokio-util", + "tokio_with_wasm", "tracing", "uuid", "xet-client", @@ -2838,6 +2730,7 @@ dependencies = [ "iceberg_test_utils", "itertools 0.13.0", "mockito", + "rand 0.9.5", "reqwest 0.12.28", "serde", "serde_derive", @@ -2845,6 +2738,7 @@ dependencies = [ "tokio", "tracing", "typed-builder", + "typetag", "uuid", ] @@ -2938,7 +2832,9 @@ dependencies = [ "iceberg_test_utils", "opendal", "reqsign-aws-v4", + "reqsign-azure-storage", "reqsign-core", + "reqsign-google", "reqwest 0.12.28", "serde", "tempfile", @@ -3122,12 +3018,6 @@ version = "2.12.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d98f6fed1fde3f8c21bc40a1abb88dd75e67924f9cffc3ef95607bad8017f8e2" -[[package]] -name = "is_terminal_polyfill" -version = "1.70.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a6cb138bb79a146c1bd460005623e142ef0181e3d0219cb493e02f7d08a35695" - [[package]] name = "itertools" version = "0.13.0" @@ -3882,12 +3772,6 @@ version = "1.21.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" -[[package]] -name = "once_cell_polyfill" -version = "1.70.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe" - [[package]] name = "oneshot" version = "0.1.13" @@ -3902,11 +3786,11 @@ checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381" [[package]] name = "opendal" -version = "0.58.1" +version = "0.59.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4f20562cc7447fcc915fc5c23df305a412ea80a733c9f2fd9e2d267e2815be6d" +checksum = "f950151f9587a51a7bed70a15fa0cff464eae96e41ae7499f97067bdafdf43eb" dependencies = [ - "ctor 1.0.8", + "ctor", "opendal-core", "opendal-http-transport-reqwest", "opendal-layer-concurrent-limit", @@ -3923,11 +3807,12 @@ dependencies = [ [[package]] name = "opendal-core" -version = "0.58.1" +version = "0.59.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ec75551ff4cf3e57da98979f6a937aaa9ddb3915bf68cc17d03df733be6646ed" +checksum = "a43405d217dfdfb543f58847336d3af672897dd1939bb7dcf314b63cf364f1c9" dependencies = [ "anyhow", + "asyncband", "base64 0.23.0", "bytes", "futures", @@ -3935,7 +3820,6 @@ dependencies = [ "jiff", "log", "md-5 0.11.0", - "mea", "percent-encoding", "quick-xml", "reqsign-core", @@ -3949,9 +3833,9 @@ dependencies = [ [[package]] name = "opendal-http-transport-reqwest" -version = "0.58.1" +version = "0.59.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ad4d4f19c3ce01126a30611f8e544eaa217104a278c889ac17c9374fe4f9e4ef" +checksum = "401999057db611e592f883fcf2cbd6754ff37af587deaadd07b8c1398b2b6b06" dependencies = [ "bytes", "futures", @@ -3963,21 +3847,21 @@ dependencies = [ [[package]] name = "opendal-layer-concurrent-limit" -version = "0.58.1" +version = "0.59.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "249ac5b0aa5a7a6c3737342d10456067937f9c9a6f3f02544271f7908ab91081" +checksum = "fba1dd0742261925fc0eb910773ec39cbc1336c55d41e13b19f3af970ad5a126" dependencies = [ + "asyncband", "futures", "http 1.5.0", - "mea", "opendal-core", ] [[package]] name = "opendal-layer-logging" -version = "0.58.1" +version = "0.59.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5c75411ab00f77851ff086b686c1e9ca8175ac18c15afa2cb75b9036436cb06c" +checksum = "d0fd963f9d32dd276521479d7f1f3a265d669b03a75f2062a4570a26b1b17421" dependencies = [ "log", "opendal-core", @@ -3985,9 +3869,9 @@ dependencies = [ [[package]] name = "opendal-layer-retry" -version = "0.58.1" +version = "0.59.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "80b7738bd5f233ad8da39af9b9316b9b7a4eaddd91e8e32a1e19b7030688121d" +checksum = "06306202c97c54fb41bdbbdcb854c8798823f0022ae7aac7a40c8a1023aa83a6" dependencies = [ "backon", "log", @@ -3996,9 +3880,9 @@ dependencies = [ [[package]] name = "opendal-layer-timeout" -version = "0.58.1" +version = "0.59.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a704141924500f3803c05ed871b53305d2a2f11cb5ef20160c3ee688a1857f66" +checksum = "6bd334cbd0a0bc934146733e74a80a5faf8db8e014781a25f6d9d80d7b87c981" dependencies = [ "opendal-core", "tokio", @@ -4006,15 +3890,15 @@ dependencies = [ [[package]] name = "opendal-service-azdls" -version = "0.58.1" +version = "0.59.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2e3c406729935fe214ce574d68681a1ff7e0b322548f14094912bdbfe50e5c53" +checksum = "9064ed464286bffbea5955d470082a1e72b9b89f1e8f7a61b9e303c55ec08e4f" dependencies = [ + "asyncband", "base64 0.23.0", "bytes", "http 1.5.0", "log", - "mea", "opendal-core", "opendal-service-azure-common", "quick-xml", @@ -4027,9 +3911,9 @@ dependencies = [ [[package]] name = "opendal-service-azure-common" -version = "0.58.1" +version = "0.59.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7348c88edf15af435b7be930077746b569fac5e738c1bf6a363b675e7317c9df" +checksum = "6348e3c0d7ff77a9b05c5b7f3d744ed395c2b2511812d3b37a934218002248f8" dependencies = [ "http 1.5.0", "opendal-core", @@ -4037,9 +3921,9 @@ dependencies = [ [[package]] name = "opendal-service-fs" -version = "0.58.1" +version = "0.59.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "826c4e17a30643b888fe983897f9a4b23b07066e1d069727a923cc8fb419a702" +checksum = "fb9caf04d6d38713299dd4abac984b95ab16ee20e1ff09160f495c5a64644083" dependencies = [ "bytes", "log", @@ -4051,9 +3935,9 @@ dependencies = [ [[package]] name = "opendal-service-gcs" -version = "0.58.1" +version = "0.59.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "007f3fba63c21e516c956b891e96ff9892d8175662bfb781cdada9d3766a11e6" +checksum = "1ccbf8450652bfe7b3ae69b7decce090c95c120ad565048f9a5a77dac2917a19" dependencies = [ "async-trait", "bytes", @@ -4068,14 +3952,16 @@ dependencies = [ "serde", "serde_json", "tokio", + "uuid", ] [[package]] name = "opendal-service-hf" -version = "0.58.1" +version = "0.59.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b41fd41eb7ed03c5e66cefda61e8e117808ffd2908f2916737cb020a6beb02c7" +checksum = "5c17b59cf22bd2da9f751b8e66db595d5fe546b4c6b6ff0fb84c7bfd27455d7a" dependencies = [ + "asyncband", "bytes", "hf-xet", "http 1.5.0", @@ -4088,9 +3974,9 @@ dependencies = [ [[package]] name = "opendal-service-oss" -version = "0.58.1" +version = "0.59.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cd528ec2d49c5ca69e674ffed7b3e0686fb9cfcfea0596870de381467fda4f1b" +checksum = "284373c4a1143d8efaa7d010c1db05856d33a7cad475aaee8405dbcb0660cd96" dependencies = [ "bytes", "http 1.5.0", @@ -4105,9 +3991,9 @@ dependencies = [ [[package]] name = "opendal-service-s3" -version = "0.58.1" +version = "0.59.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "58e80cdf192d7eff05feed747894d64f81905ac4eaf132edf7ea270abdd2d663" +checksum = "388b1d39b62535c62803754ebef89808859558697366dbedd0299345887ba461" dependencies = [ "base64 0.23.0", "bytes", @@ -4262,11 +4148,11 @@ dependencies = [ [[package]] name = "pem" -version = "3.0.6" +version = "4.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1d30c53c26bc5b31a98cd02d20f25a7c8567146caf63ed593a9d87b2775291be" +checksum = "d354a98a3d1251555de99e8fdd8afda05573c31b82f59063a7b0a29b5527f120" dependencies = [ - "base64 0.22.1", + "base64 0.23.0", "serde_core", ] @@ -4922,9 +4808,9 @@ dependencies = [ [[package]] name = "reqsign-aws-core" -version = "3.1.0" +version = "3.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4d63b56638bb3cc7bd376a7cdce1ba3089777a08f47e4097888f2d784cc3f46c" +checksum = "bac4749b7dfa7bfaccd01eb03e9dc795ed37e3f20d6f0f38e2c67ee85ad6bc86" dependencies = [ "bytes", "form_urlencoded", @@ -4943,9 +4829,9 @@ dependencies = [ [[package]] name = "reqsign-aws-v4" -version = "3.2.0" +version = "3.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4a0c499f4ed12d04c3d4c78fe4cb01aee22c9dae22848c14db2c6313d9df9f43" +checksum = "ff250f0fd0b913fbd565e405acc553da0f13bde30bfb5403178c9d0313cdc15f" dependencies = [ "bytes", "http 1.5.0", @@ -4958,9 +4844,9 @@ dependencies = [ [[package]] name = "reqsign-azure-storage" -version = "3.1.2" +version = "3.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2824e7da3c2cc42ac3406c674eb57c89127fdcd97f3a73c608cfc680505ea134" +checksum = "c3f1fbc9add082f54e51e3bd9730f18d850a1f09d246baff5cd612c74ae4290e" dependencies = [ "anyhow", "base64 0.23.0", @@ -4979,9 +4865,9 @@ dependencies = [ [[package]] name = "reqsign-core" -version = "3.3.0" +version = "3.3.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f4ac1510872d9481205975d264deb39c109797e5068cc882ed9064270eaae5fa" +checksum = "ff052daffb0599681c50f85c59e7236438976efe991ab864edd9f3b235501a0f" dependencies = [ "anyhow", "base64 0.23.0", @@ -4992,6 +4878,7 @@ dependencies = [ "http 1.5.0", "jiff", "log", + "mea", "percent-encoding", "rsa", "serde", @@ -5003,9 +4890,9 @@ dependencies = [ [[package]] name = "reqsign-file-read-tokio" -version = "3.0.4" +version = "3.0.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "663d9d55abd0df0830ef0ae43708297cc1371cf4e8ca91f3ac813c309cca8c98" +checksum = "b3235df90a6bca681aa47dd86f2393d122a6d77042aa8a7c81e218cd45c5bfc0" dependencies = [ "anyhow", "reqsign-core", @@ -5633,7 +5520,6 @@ dependencies = [ "cfg-if 1.0.4", "cpufeatures 0.2.17", "digest 0.10.7", - "sha2-asm", ] [[package]] @@ -5647,15 +5533,6 @@ dependencies = [ "digest 0.11.3", ] -[[package]] -name = "sha2-asm" -version = "0.6.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b845214d6175804686b2bd482bcffe96651bb2d1200742b712003504a2dac1ab" -dependencies = [ - "cc", -] - [[package]] name = "sharded-slab" version = "0.1.7" @@ -6397,6 +6274,30 @@ dependencies = [ "tokio", ] +[[package]] +name = "tokio_with_wasm" +version = "0.8.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "34e40fbbbd95441133fe9483f522db15dbfd26dc636164ebd8f2dd28759a6aa6" +dependencies = [ + "js-sys", + "tokio", + "tokio_with_wasm_proc", + "wasm-bindgen", + "wasm-bindgen-futures", + "web-sys", +] + +[[package]] +name = "tokio_with_wasm_proc" +version = "0.8.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d01145a2c788d6aae4cd653afec1e8332534d7d783d01897cefcafe4428de992" +dependencies = [ + "quote", + "syn", +] + [[package]] name = "toml" version = "1.1.3+spec-1.1.0" @@ -6737,12 +6638,6 @@ version = "1.0.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" -[[package]] -name = "utf8parse" -version = "0.2.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" - [[package]] name = "uuid" version = "1.26.0" @@ -7377,20 +7272,18 @@ dependencies = [ [[package]] name = "xet-client" -version = "1.5.2" +version = "1.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3e1e496dcbe6a09017acdfaf48e1a646735e7ff5b2a49e2c7e081cca77a59bc8" +checksum = "c3b8da8cc70aa2e3c500c0400e012df82c656ab9fca47f9f939fffc5afd89aca" dependencies = [ "anyhow", "async-trait", "base64 0.22.1", "bytes", - "clap", "crc32fast", "futures", "http 1.5.0", "hyper", - "lazy_static", "more-asserts", "rand 0.10.2", "redb", @@ -7404,8 +7297,8 @@ dependencies = [ "thiserror 2.0.18", "tokio", "tokio-retry", + "tokio_with_wasm", "tracing", - "tracing-subscriber", "url", "urlencoding", "web-time", @@ -7415,24 +7308,21 @@ dependencies = [ [[package]] name = "xet-core-structures" -version = "1.5.2" +version = "1.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cb838aa8eb67d730af301584cf003caad407487606058292a6750711b603fbee" +checksum = "73503c223783dccc864abde22115e09d12f190448a0baf58ab2c54bc709e2f99" dependencies = [ "async-trait", "base64 0.22.1", "blake3", "bytemuck", "bytes", - "clap", "countio", - "csv", "futures", "futures-util", "getrandom 0.4.3", "heapify", "itertools 0.14.0", - "lazy_static", "lz4_flex 0.13.1", "more-asserts", "rand 0.10.2", @@ -7440,7 +7330,6 @@ dependencies = [ "safe-transmute", "serde", "static_assertions", - "tempfile", "thiserror 2.0.18", "tokio", "tokio-util", @@ -7452,32 +7341,31 @@ dependencies = [ [[package]] name = "xet-data" -version = "1.5.2" +version = "1.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "67fd409bef621411a9d9013798540bb8036cb2678f03ab39af89a5e88034ed8c" +checksum = "c89052ec5dec2187cad30b86af92cc24fd61c4a57a795f1ff7ff5f38d49184eb" dependencies = [ "anyhow", "async-trait", "bytes", "chrono", - "clap", "gearhash", "http 1.5.0", "itertools 0.14.0", - "lazy_static", "more-asserts", "rand 0.10.2", "serde", "serde_json", - "sha2 0.10.9", + "sha2 0.11.0", "tempfile", "thiserror 2.0.18", "tokio", "tokio-util", + "tokio_with_wasm", "tracing", "url", "uuid", - "walkdir", + "web-time", "xet-client", "xet-core-structures", "xet-runtime", @@ -7485,9 +7373,9 @@ dependencies = [ [[package]] name = "xet-runtime" -version = "1.5.2" +version = "1.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "15d8f121c33866f7648b737abe70d0e2dd9c0af4ffdd7219207531d0283aa63d" +checksum = "af5c60d5eed38ab4c576f4421bae835e7bd07631fb381705605529d2015c106b" dependencies = [ "anyhow", "async-trait", @@ -7495,13 +7383,12 @@ dependencies = [ "chrono", "colored", "const-str", - "ctor 0.6.3", + "ctor", "dirs", "futures", "git-version", "humantime", "konst", - "lazy_static", "libc", "more-asserts", "oneshot", @@ -7515,9 +7402,11 @@ dependencies = [ "thiserror 2.0.18", "tokio", "tokio-util", + "tokio_with_wasm", "tracing", "tracing-appender", "tracing-subscriber", + "web-time", "whoami 2.1.2", "winapi", ] diff --git a/Cargo.toml b/Cargo.toml index 661c56b094..426e322249 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -105,7 +105,7 @@ mockito = "1" motore-macros = "0.4.3" murmur3 = "0.5.2" once_cell = "1.20" -opendal = "0.58" +opendal = "0.59" ordered-float = "4" parquet = "59.2" pilota = "0.11.10" diff --git a/crates/catalog/rest/Cargo.toml b/crates/catalog/rest/Cargo.toml index 8dc9a86d7f..0f7a609786 100644 --- a/crates/catalog/rest/Cargo.toml +++ b/crates/catalog/rest/Cargo.toml @@ -35,6 +35,7 @@ chrono = { workspace = true } http = { workspace = true } iceberg = { workspace = true } itertools = { workspace = true } +rand = { workspace = true } reqwest = { workspace = true } serde = { workspace = true } serde_derive = { workspace = true } @@ -42,6 +43,7 @@ serde_json = { workspace = true } tokio = { workspace = true } tracing = { workspace = true } typed-builder = { workspace = true } +typetag = { workspace = true } uuid = { workspace = true, features = ["v4"] } [dev-dependencies] diff --git a/crates/catalog/rest/public-api.txt b/crates/catalog/rest/public-api.txt index 8afd581bcd..eb4638d809 100644 --- a/crates/catalog/rest/public-api.txt +++ b/crates/catalog/rest/public-api.txt @@ -181,6 +181,20 @@ impl serde_core::ser::Serialize for iceberg_catalog_rest::ListTablesResponse pub fn iceberg_catalog_rest::ListTablesResponse::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg_catalog_rest::ListTablesResponse pub fn iceberg_catalog_rest::ListTablesResponse::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> +pub struct iceberg_catalog_rest::LoadCredentialsResponse +pub iceberg_catalog_rest::LoadCredentialsResponse::storage_credentials: alloc::vec::Vec +impl core::clone::Clone for iceberg_catalog_rest::LoadCredentialsResponse +pub fn iceberg_catalog_rest::LoadCredentialsResponse::clone(&self) -> iceberg_catalog_rest::LoadCredentialsResponse +impl core::cmp::Eq for iceberg_catalog_rest::LoadCredentialsResponse +impl core::cmp::PartialEq for iceberg_catalog_rest::LoadCredentialsResponse +pub fn iceberg_catalog_rest::LoadCredentialsResponse::eq(&self, other: &iceberg_catalog_rest::LoadCredentialsResponse) -> bool +impl core::fmt::Debug for iceberg_catalog_rest::LoadCredentialsResponse +pub fn iceberg_catalog_rest::LoadCredentialsResponse::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result +impl core::marker::StructuralPartialEq for iceberg_catalog_rest::LoadCredentialsResponse +impl serde_core::ser::Serialize for iceberg_catalog_rest::LoadCredentialsResponse +pub fn iceberg_catalog_rest::LoadCredentialsResponse::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer +impl<'de> serde_core::de::Deserialize<'de> for iceberg_catalog_rest::LoadCredentialsResponse +pub fn iceberg_catalog_rest::LoadCredentialsResponse::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> pub struct iceberg_catalog_rest::LoadTableResult pub iceberg_catalog_rest::LoadTableResult::config: std::collections::hash::map::HashMap pub iceberg_catalog_rest::LoadTableResult::metadata: iceberg::spec::table_metadata::TableMetadata @@ -223,6 +237,7 @@ pub fn iceberg_catalog_rest::NoopAuthManager::fmt(&self, f: &mut core::fmt::Form impl iceberg_catalog_rest::AuthManager for iceberg_catalog_rest::NoopAuthManager pub fn iceberg_catalog_rest::NoopAuthManager::catalog_session<'life0, 'life1, 'life2, 'async_trait>(&'life0 self, _client: &'life1 iceberg_catalog_rest::HttpClient, _props: &'life2 std::collections::hash::map::HashMap) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait, 'life2: 'async_trait pub fn iceberg_catalog_rest::NoopAuthManager::init_session<'life0, 'life1, 'life2, 'async_trait>(&'life0 self, _client: &'life1 iceberg_catalog_rest::HttpClient, _props: &'life2 std::collections::hash::map::HashMap) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait, 'life2: 'async_trait +pub fn iceberg_catalog_rest::NoopAuthManager::table_session<'life0, 'life1, 'life2, 'life3, 'async_trait>(&'life0 self, _client: &'life1 iceberg_catalog_rest::HttpClient, _table: &'life2 iceberg::catalog::TableIdent, _props: &'life3 std::collections::hash::map::HashMap, parent: alloc::sync::Arc) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait, 'life2: 'async_trait, 'life3: 'async_trait pub struct iceberg_catalog_rest::OAuth2Manager impl iceberg_catalog_rest::OAuth2Manager pub fn iceberg_catalog_rest::OAuth2Manager::new(token_endpoint: impl core::convert::Into) -> Self @@ -235,6 +250,7 @@ pub fn iceberg_catalog_rest::OAuth2Manager::fmt(&self, f: &mut core::fmt::Format impl iceberg_catalog_rest::AuthManager for iceberg_catalog_rest::OAuth2Manager pub fn iceberg_catalog_rest::OAuth2Manager::catalog_session<'life0, 'life1, 'life2, 'async_trait>(&'life0 self, client: &'life1 iceberg_catalog_rest::HttpClient, props: &'life2 std::collections::hash::map::HashMap) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait, 'life2: 'async_trait pub fn iceberg_catalog_rest::OAuth2Manager::init_session<'life0, 'life1, 'life2, 'async_trait>(&'life0 self, client: &'life1 iceberg_catalog_rest::HttpClient, props: &'life2 std::collections::hash::map::HashMap) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait, 'life2: 'async_trait +pub fn iceberg_catalog_rest::OAuth2Manager::table_session<'life0, 'life1, 'life2, 'life3, 'async_trait>(&'life0 self, _client: &'life1 iceberg_catalog_rest::HttpClient, _table: &'life2 iceberg::catalog::TableIdent, props: &'life3 std::collections::hash::map::HashMap, parent: alloc::sync::Arc) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait, 'life2: 'async_trait, 'life3: 'async_trait pub struct iceberg_catalog_rest::RegisterTableRequest pub iceberg_catalog_rest::RegisterTableRequest::metadata_location: alloc::string::String pub iceberg_catalog_rest::RegisterTableRequest::name: alloc::string::String @@ -386,11 +402,14 @@ pub const iceberg_catalog_rest::REST_CATALOG_PROP_WAREHOUSE: &str pub trait iceberg_catalog_rest::AuthManager: core::fmt::Debug + core::marker::Send + core::marker::Sync pub fn iceberg_catalog_rest::AuthManager::catalog_session<'life0, 'life1, 'life2, 'async_trait>(&'life0 self, client: &'life1 iceberg_catalog_rest::HttpClient, props: &'life2 std::collections::hash::map::HashMap) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait, 'life2: 'async_trait pub fn iceberg_catalog_rest::AuthManager::init_session<'life0, 'life1, 'life2, 'async_trait>(&'life0 self, client: &'life1 iceberg_catalog_rest::HttpClient, props: &'life2 std::collections::hash::map::HashMap) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait, 'life2: 'async_trait +pub fn iceberg_catalog_rest::AuthManager::table_session<'life0, 'life1, 'life2, 'life3, 'async_trait>(&'life0 self, _client: &'life1 iceberg_catalog_rest::HttpClient, _table: &'life2 iceberg::catalog::TableIdent, _props: &'life3 std::collections::hash::map::HashMap, parent: alloc::sync::Arc) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait, 'life2: 'async_trait, 'life3: 'async_trait impl iceberg_catalog_rest::AuthManager for iceberg_catalog_rest::NoopAuthManager pub fn iceberg_catalog_rest::NoopAuthManager::catalog_session<'life0, 'life1, 'life2, 'async_trait>(&'life0 self, _client: &'life1 iceberg_catalog_rest::HttpClient, _props: &'life2 std::collections::hash::map::HashMap) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait, 'life2: 'async_trait pub fn iceberg_catalog_rest::NoopAuthManager::init_session<'life0, 'life1, 'life2, 'async_trait>(&'life0 self, _client: &'life1 iceberg_catalog_rest::HttpClient, _props: &'life2 std::collections::hash::map::HashMap) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait, 'life2: 'async_trait +pub fn iceberg_catalog_rest::NoopAuthManager::table_session<'life0, 'life1, 'life2, 'life3, 'async_trait>(&'life0 self, _client: &'life1 iceberg_catalog_rest::HttpClient, _table: &'life2 iceberg::catalog::TableIdent, _props: &'life3 std::collections::hash::map::HashMap, parent: alloc::sync::Arc) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait, 'life2: 'async_trait, 'life3: 'async_trait impl iceberg_catalog_rest::AuthManager for iceberg_catalog_rest::OAuth2Manager pub fn iceberg_catalog_rest::OAuth2Manager::catalog_session<'life0, 'life1, 'life2, 'async_trait>(&'life0 self, client: &'life1 iceberg_catalog_rest::HttpClient, props: &'life2 std::collections::hash::map::HashMap) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait, 'life2: 'async_trait pub fn iceberg_catalog_rest::OAuth2Manager::init_session<'life0, 'life1, 'life2, 'async_trait>(&'life0 self, client: &'life1 iceberg_catalog_rest::HttpClient, props: &'life2 std::collections::hash::map::HashMap) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait, 'life2: 'async_trait +pub fn iceberg_catalog_rest::OAuth2Manager::table_session<'life0, 'life1, 'life2, 'life3, 'async_trait>(&'life0 self, _client: &'life1 iceberg_catalog_rest::HttpClient, _table: &'life2 iceberg::catalog::TableIdent, props: &'life3 std::collections::hash::map::HashMap, parent: alloc::sync::Arc) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait, 'life2: 'async_trait, 'life3: 'async_trait pub trait iceberg_catalog_rest::AuthSession: core::fmt::Debug + core::marker::Send + core::marker::Sync pub fn iceberg_catalog_rest::AuthSession::authenticate<'life0, 'life1, 'async_trait>(&'life0 self, request: &'life1 mut iceberg_catalog_rest::HttpRequest) -> core::pin::Pin> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait diff --git a/crates/catalog/rest/src/auth/mod.rs b/crates/catalog/rest/src/auth/mod.rs index b15dcdb480..f6bf30ca2c 100644 --- a/crates/catalog/rest/src/auth/mod.rs +++ b/crates/catalog/rest/src/auth/mod.rs @@ -25,9 +25,10 @@ use std::fmt::Debug; use std::sync::Arc; use async_trait::async_trait; -use iceberg::Result; +use iceberg::{Error, ErrorKind, Result, TableIdent}; pub use oauth2::OAuth2Manager; +use crate::catalog::{REST_CATALOG_PROP_AUTH_TYPE, RestCatalogConfig}; use crate::client::HttpClient; use crate::request::HttpRequest; @@ -36,6 +37,32 @@ pub const AUTH_TYPE_NONE: &str = "none"; /// `rest.auth.type` value selecting OAuth2 token authentication. pub const AUTH_TYPE_OAUTH2: &str = "oauth2"; +/// Builds the auth manager selected by the `rest.auth.type` configuration, +/// like Java's `AuthManagers.loadAuthManager`. +pub(crate) fn load_auth_manager(config: &RestCatalogConfig) -> Result> { + let auth_type = config.auth_type(); + // Java parity (`AuthManagers`): make the inference visible so users + // configure the type explicitly. + if auth_type == AUTH_TYPE_OAUTH2 && !config.has_explicit_auth_type() { + tracing::warn!( + "Inferring {REST_CATALOG_PROP_AUTH_TYPE}={AUTH_TYPE_OAUTH2} from the configured \ + OAuth properties; set it explicitly to avoid this warning" + ); + } + match auth_type.as_str() { + AUTH_TYPE_NONE => Ok(Arc::new(NoopAuthManager)), + AUTH_TYPE_OAUTH2 => Ok(Arc::new(OAuth2Manager::from_config(config)?)), + other => Err(Error::new( + ErrorKind::DataInvalid, + format!( + "unknown '{REST_CATALOG_PROP_AUTH_TYPE}': {other}; use \ + `RestSessionCatalogBuilder::with_auth_manager` or \ + `RestCatalogBuilder::with_auth_manager` to inject a custom auth manager" + ), + )), + } +} + /// Creates the [`AuthSession`]s used to authenticate REST catalog requests. /// /// A manager is exclusively scoped to one catalog and must not be reused by @@ -46,9 +73,9 @@ pub const AUTH_TYPE_OAUTH2: &str = "oauth2"; /// Catalog initialization calls [`AuthManager::catalog_session`] exactly once; /// later sessions may rely on the state established by that call. /// -/// Both methods are handed the catalog's [`HttpClient`], which an -/// implementation may reuse for its own requests (e.g. a token exchange) so -/// that they share the catalog's connection pool and configuration. +/// Session-construction methods are handed the catalog's [`HttpClient`], which +/// an implementation may reuse for its own requests (e.g. a token exchange) +/// so that they share the catalog's connection pool and configuration. #[async_trait] pub trait AuthManager: Debug + Send + Sync { /// Session used for the initial `/v1/config` handshake, given the @@ -73,6 +100,22 @@ pub trait AuthManager: Debug + Send + Sync { client: &HttpClient, props: &HashMap, ) -> Result>; + + /// Returns a session for requests associated with `table`. + /// + /// `props` are the unmerged properties returned by the table endpoint. + /// The default preserves the catalog session; managers should return a + /// child session only when the table properties contain an authentication + /// override. + async fn table_session( + &self, + _client: &HttpClient, + _table: &TableIdent, + _props: &HashMap, + parent: Arc, + ) -> Result> { + Ok(parent) + } } /// Authenticates outgoing REST catalog requests. diff --git a/crates/catalog/rest/src/auth/oauth2.rs b/crates/catalog/rest/src/auth/oauth2.rs index aad91996a8..39a71b53d0 100644 --- a/crates/catalog/rest/src/auth/oauth2.rs +++ b/crates/catalog/rest/src/auth/oauth2.rs @@ -22,7 +22,7 @@ use std::sync::Arc; use async_trait::async_trait; use http::StatusCode; use iceberg::sensitive::SensitiveString; -use iceberg::{Error, ErrorKind, Result}; +use iceberg::{Error, ErrorKind, Result, TableIdent}; use reqwest::header::HeaderMap; use tokio::sync::Mutex; @@ -46,8 +46,9 @@ struct OAuth2Params { /// Iceberg REST catalogs. /// /// A configured `token` is used directly; otherwise `credential` is exchanged -/// for a token at the token endpoint and cached. The cached token is shared -/// across sessions so it survives the config handshake. +/// for a token at the token endpoint and cached. The cached token is shared by +/// the init and catalog sessions so it survives the config handshake; +/// table-specific tokens use isolated sessions. pub struct OAuth2Manager { token: Arc>>, init_params: OAuth2Params, @@ -144,13 +145,34 @@ impl AuthManager for OAuth2Manager { ) -> Result> { Ok(Arc::new(self.session_from(client, props).await?)) } + + /// Like Java, only a `token` in the table config overrides the parent + /// session; a table-level `credential` is ignored. Unlike Java, the token is + /// used as-is and never refreshed, and token-type exchange keys are ignored, + /// because this manager implements neither token refresh nor token exchange. + async fn table_session( + &self, + _client: &HttpClient, + _table: &TableIdent, + props: &HashMap, + parent: Arc, + ) -> Result> { + let Some(token) = props.get("token") else { + return Ok(parent); + }; + + Ok(Arc::new(OAuth2Session { + token: Arc::new(Mutex::new(Some(SensitiveString::from(token.clone())))), + token_source: TokenSource::StaticToken, + })) + } } impl OAuth2Manager { /// Builds a session from the manager's options with `props` merged onto /// them, so an injected manager keeps whatever a property doesn't - /// override. The manager's token cell is shared with every session it - /// builds, so a token cached during the handshake survives it. + /// override. The manager's token cell is shared by the init and catalog + /// sessions, so a token cached during the handshake survives it. async fn session_from( &self, client: &HttpClient, @@ -380,4 +402,56 @@ mod tests { "Bearer tok-static" ); } + + #[tokio::test] + async fn test_table_session_inherits_parent_unless_token_is_overridden() { + let manager = OAuth2Manager::new("http://localhost/unused").with_token("catalog-token"); + let client = test_client(); + let parent = manager + .catalog_session(&client, &HashMap::new()) + .await + .unwrap(); + let table = TableIdent::from_strs(["namespace", "table"]).unwrap(); + + let inherited = manager + .table_session(&client, &table, &HashMap::new(), Arc::clone(&parent)) + .await + .unwrap(); + assert!(Arc::ptr_eq(&parent, &inherited)); + + let overridden = manager + .table_session( + &client, + &table, + &HashMap::from([("token".to_string(), "table-token".to_string())]), + Arc::clone(&parent), + ) + .await + .unwrap(); + assert!(!Arc::ptr_eq(&parent, &overridden)); + + let mut parent_request = HttpRequest::new( + Client::new() + .get("https://rest.example.com/catalog") + .build() + .unwrap(), + ); + parent.authenticate(&mut parent_request).await.unwrap(); + assert_eq!( + parent_request.headers().get("authorization").unwrap(), + "Bearer catalog-token" + ); + + let mut table_request = HttpRequest::new( + Client::new() + .get("https://rest.example.com/table") + .build() + .unwrap(), + ); + overridden.authenticate(&mut table_request).await.unwrap(); + assert_eq!( + table_request.headers().get("authorization").unwrap(), + "Bearer table-token" + ); + } } diff --git a/crates/catalog/rest/src/catalog.rs b/crates/catalog/rest/src/catalog.rs index b7d3b3163c..b1a57ff156 100644 --- a/crates/catalog/rest/src/catalog.rs +++ b/crates/catalog/rest/src/catalog.rs @@ -39,10 +39,11 @@ use reqwest::{Client, Method, StatusCode, Url}; use tokio::sync::OnceCell; use typed_builder::TypedBuilder; -use crate::auth::{AUTH_TYPE_NONE, AUTH_TYPE_OAUTH2, AuthManager, NoopAuthManager, OAuth2Manager}; +use crate::auth::{AUTH_TYPE_NONE, AUTH_TYPE_OAUTH2, AuthManager, load_auth_manager}; use crate::client::{ HttpClient, deserialize_catalog_response, deserialize_unexpected_catalog_error, }; +use crate::credential::{RestVendedCredentialProviderFactory, build_vended_credential_provider}; use crate::endpoint::{Endpoint, V1_NAMESPACE_EXISTS, V1_TABLE_EXISTS}; use crate::request::HttpRequest; use crate::response::HttpResponse; @@ -59,6 +60,8 @@ pub const REST_CATALOG_PROP_WAREHOUSE: &str = "warehouse"; /// Disable header redaction in error logs and `Debug` output (defaults to /// false for security) pub const REST_CATALOG_PROP_DISABLE_HEADER_REDACTION: &str = "disable-header-redaction"; +/// Identifier for a server-side scan plan associated with credential requests. +pub(crate) const REST_CATALOG_PROP_SCAN_PLAN_ID: &str = "rest.scan.plan-id"; /// Authentication scheme: `none` or `oauth2`. When unset, `oauth2` is used /// if a `token`, `credential` or `oauth2-server-uri` is configured, `none` /// otherwise. @@ -289,10 +292,52 @@ impl RestCatalogConfig { /// Returns true if the `disable-header-redaction` property is set to "true". /// Defaults to false for security (headers are redacted by default). pub(crate) fn disable_header_redaction(&self) -> bool { + disable_header_redaction_from_props(&self.props).unwrap_or(false) + } + + /// The configured auth scheme: explicit `rest.auth.type` (matched + /// case-insensitively) when set; otherwise `oauth2` when a `token`, + /// `credential` or `oauth2-server-uri` is configured (preserving + /// pre-`rest.auth.type` setups), `none` when none is. + pub(crate) fn auth_type(&self) -> String { self.props - .get(REST_CATALOG_PROP_DISABLE_HEADER_REDACTION) - .map(|v| v.eq_ignore_ascii_case("true")) - .unwrap_or(false) + .get(REST_CATALOG_PROP_AUTH_TYPE) + // Matched case-insensitively, as the other flag properties are. + .map(|auth_type| auth_type.to_ascii_lowercase()) + .unwrap_or_else(|| { + if self.token().is_some() + || self.credential().is_some() + || self.explicit_oauth2_server_uri().is_some() + { + AUTH_TYPE_OAUTH2.to_string() + } else { + AUTH_TYPE_NONE.to_string() + } + }) + } + + /// Whether `rest.auth.type` is set explicitly rather than inferred. + pub(crate) fn has_explicit_auth_type(&self) -> bool { + self.props.contains_key(REST_CATALOG_PROP_AUTH_TYPE) + } + + /// The properties handed to the [`AuthManager`], with the catalog `uri` + /// and `warehouse` made explicit. + pub(crate) fn auth_props(&self) -> HashMap { + // `oauth2-server-uri` stays absent unless explicitly configured, so an + // injected manager keeps its own endpoint. The resolved `uri` and + // `warehouse` ARE passed: the builder moved them off the props, and + // the built-in manager recomputes its token endpoint from the URI. + let mut props = self.props.clone(); + props.insert(REST_CATALOG_PROP_URI.to_string(), self.uri.clone()); + if let Some(warehouse) = &self.warehouse { + // A fallback only: after the handshake the merged props hold + // the resolved warehouse, server override included. + props + .entry(REST_CATALOG_PROP_WAREHOUSE.to_string()) + .or_insert_with(|| warehouse.clone()); + } + props } /// Merge the `RestCatalogConfig` with the a [`CatalogConfig`] (fetched from the REST server). @@ -362,6 +407,13 @@ pub(crate) fn extra_headers_from_props(props: &HashMap) -> Resul Ok(headers) } +/// The `disable-header-redaction` property, when set. +pub(crate) fn disable_header_redaction_from_props(props: &HashMap) -> Option { + props + .get(REST_CATALOG_PROP_DISABLE_HEADER_REDACTION) + .map(|value| value.eq_ignore_ascii_case("true")) +} + /// The default OAuth2 token endpoint for a catalog `uri`. pub(crate) fn default_token_endpoint(uri: &str) -> String { [uri, PATH_V1, "oauth", "tokens"].join("/") @@ -418,6 +470,8 @@ pub(crate) fn oauth_params_from_props(props: &HashMap) -> HashMa #[derive(Debug)] struct RestClient { + /// Creates child authentication sessions for table-scoped requests. + auth_manager: Arc, /// Carries the session the auth manager derived from the merged /// configuration, so every request below is authenticated. http_client: HttpClient, @@ -444,7 +498,7 @@ impl RestClient { let init_session = auth_manager .init_session( &http_client.without_auth_session(), - &Self::auth_props(user_config), + &user_config.auth_props(), ) .await?; Self::load_config( @@ -464,13 +518,11 @@ impl RestClient { // The manager is handed an unauthenticated client: its own // requests must not be signed by the session it is deriving. let session = auth_manager - .catalog_session( - &http_client.without_auth_session(), - &Self::auth_props(&config), - ) + .catalog_session(&http_client.without_auth_session(), &config.auth_props()) .await?; Ok(Self { + auth_manager, config, http_client: http_client.with_auth_session(session), endpoints, @@ -488,25 +540,6 @@ impl RestClient { self.http_client.query_catalog(request).await } - /// The properties handed to the [`AuthManager`], with the catalog `uri` - /// and `warehouse` made explicit. - fn auth_props(config: &RestCatalogConfig) -> HashMap { - // `oauth2-server-uri` stays absent unless explicitly configured, so an - // injected manager keeps its own endpoint. The resolved `uri` and - // `warehouse` ARE passed: the builder moved them off the props, and - // the built-in manager recomputes its token endpoint from the URI. - let mut props = config.props.clone(); - props.insert(REST_CATALOG_PROP_URI.to_string(), config.uri.clone()); - if let Some(warehouse) = &config.warehouse { - // A fallback only: after the handshake the merged props hold - // the resolved warehouse, server override included. - props - .entry(REST_CATALOG_PROP_WAREHOUSE.to_string()) - .or_insert_with(|| warehouse.clone()); - } - props - } - /// Loads the runtime config from the server using `user_config`. /// /// It's required for a REST catalog to update its config after creation. @@ -758,56 +791,12 @@ impl RestSessionCatalog { } } - /// The configured auth scheme: explicit `rest.auth.type` (matched - /// case-insensitively) when set; otherwise `oauth2` when a `token`, - /// `credential` or `oauth2-server-uri` is configured (preserving - /// pre-`rest.auth.type` setups), `none` when none is. - fn auth_type(config: &RestCatalogConfig) -> String { - config - .props - .get(REST_CATALOG_PROP_AUTH_TYPE) - // Matched case-insensitively, as the other flag properties are. - .map(|auth_type| auth_type.to_ascii_lowercase()) - .unwrap_or_else(|| { - if config.token().is_some() - || config.credential().is_some() - || config.explicit_oauth2_server_uri().is_some() - { - AUTH_TYPE_OAUTH2.to_string() - } else { - AUTH_TYPE_NONE.to_string() - } - }) - } - /// Resolves the auth manager: a `with_auth_manager` override wins, /// otherwise one is built from the `rest.auth.type` configuration. fn resolve_auth_manager(&self) -> Result> { - if let Some(auth_manager) = &self.auth_manager { - return Ok(auth_manager.clone()); - } - let config = &self.user_config; - let auth_type = Self::auth_type(config); - // Java parity (`AuthManagers`): make the inference visible so users - // configure the type explicitly. - if auth_type == AUTH_TYPE_OAUTH2 && !config.props.contains_key(REST_CATALOG_PROP_AUTH_TYPE) - { - tracing::warn!( - "Inferring {REST_CATALOG_PROP_AUTH_TYPE}={AUTH_TYPE_OAUTH2} from the configured \ - OAuth properties; set it explicitly to avoid this warning" - ); - } - match auth_type.as_str() { - AUTH_TYPE_NONE => Ok(Arc::new(NoopAuthManager)), - AUTH_TYPE_OAUTH2 => Ok(Arc::new(OAuth2Manager::from_config(config)?)), - other => Err(Error::new( - ErrorKind::DataInvalid, - format!( - "unknown '{REST_CATALOG_PROP_AUTH_TYPE}': {other}; use \ - `RestSessionCatalogBuilder::with_auth_manager` or \ - `RestCatalogBuilder::with_auth_manager` to inject a custom auth manager" - ), - )), + match &self.auth_manager { + Some(auth_manager) => Ok(auth_manager.clone()), + None => load_auth_manager(&self.user_config), } } @@ -843,22 +832,27 @@ impl RestSessionCatalog { } } + /// Builds the FileIO for `table` from the catalog properties and, when + /// loaded from a table response, its `config` overridden by the user + /// properties. async fn load_file_io( &self, + table: &TableIdent, metadata_location: Option<&str>, - extra_config: Option>, + table_config: Option>, ) -> Result { - let mut props = self.client().await?.config.props.clone(); - if let Some(config) = extra_config { - props.extend(config); + let client = self.client().await?; + let mut props = client.config.props.clone(); + if let Some(table_config) = &table_config { + props.extend(table_config.clone()); + props.extend(self.user_config.props.clone()); } // If the warehouse is a logical identifier instead of a URL we don't want // to raise an exception - let warehouse_path = match self.client().await?.config.warehouse.as_deref() { + let warehouse_path = match client.config.warehouse.as_deref() { Some(url) if Url::parse(url).is_ok() => Some(url), - Some(_) => None, - None => None, + _ => None, }; if metadata_location.or(warehouse_path).is_none() { @@ -879,9 +873,29 @@ impl RestSessionCatalog { ) })?; - let file_io = FileIOBuilder::new(factory).with_props(props).build(); + // If the catalog vends refreshable credentials for this table's storage, + // attach a provider so the backend re-fetches them before they expire. + // Only catalog authentication resolved from the properties can be + // rebuilt after FileIO serialization. + let credential_provider = build_vended_credential_provider( + &client.http_client, + client.auth_manager.as_ref(), + RestVendedCredentialProviderFactory::new( + &client.config.uri, + table.clone(), + table_config.unwrap_or_default(), + ), + &props, + self.auth_manager.is_none(), + ) + .await?; + + let mut builder = FileIOBuilder::new(factory).with_props(props); + if let Some(provider) = credential_provider { + builder = builder.with_credential_provider(provider); + } - Ok(file_io) + Ok(builder.build()) } } @@ -1181,14 +1195,8 @@ impl SessionCatalog for RestSessionCatalog { "Metadata location missing in `create_table` response!", ))?; - let config = response - .config - .into_iter() - .chain(self.user_config.props.clone()) - .collect(); - let file_io = self - .load_file_io(Some(metadata_location), Some(config)) + .load_file_io(&table_ident, Some(metadata_location), Some(response.config)) .await?; let mut table_builder = Table::builder() @@ -1245,14 +1253,12 @@ impl SessionCatalog for RestSessionCatalog { } }; - let config = response - .config - .into_iter() - .chain(self.user_config.props.clone()) - .collect(); - let file_io = self - .load_file_io(response.metadata_location.as_deref(), Some(config)) + .load_file_io( + table_ident, + response.metadata_location.as_deref(), + Some(response.config), + ) .await?; let mut table_builder = Table::builder() @@ -1391,7 +1397,9 @@ impl SessionCatalog for RestSessionCatalog { "Metadata location missing in `register_table` response!", ))?; - let file_io = self.load_file_io(Some(metadata_location), None).await?; + let file_io = self + .load_file_io(table_ident, Some(metadata_location), Some(response.config)) + .await?; let mut table_builder = Table::builder() .identifier(table_ident.clone()) @@ -1470,7 +1478,7 @@ impl SessionCatalog for RestSessionCatalog { }; let file_io = self - .load_file_io(Some(&response.metadata_location), None) + .load_file_io(commit.identifier(), Some(&response.metadata_location), None) .await?; let mut table_builder = Table::builder() @@ -1680,7 +1688,7 @@ mod tests { use uuid::uuid; use super::*; - use crate::auth::AuthSession; + use crate::auth::{AuthSession, NoopAuthManager, OAuth2Manager}; use crate::request::HttpRequest; fn test_catalog(config: RestCatalogConfig) -> RestSessionCatalog { @@ -4146,6 +4154,54 @@ mod tests { load_table_mock.assert_async().await } + #[tokio::test] + async fn test_injected_auth_manager_file_io_serializes_without_provider_only() { + let mut server = Server::new_async().await; + let config_mock = create_config_mock(&mut server).await; + let mut load_response: serde_json::Value = serde_json::from_reader(BufReader::new( + File::open(format!( + "{}/testdata/load_table_response.json", + env!("CARGO_MANIFEST_DIR") + )) + .unwrap(), + )) + .unwrap(); + load_response["config"]["client.refresh-credentials-endpoint"] = + json!("/v1/namespaces/ns1/tables/test1/credentials"); + let load_table_mock = server + .mock("GET", "/v1/namespaces/ns1/tables/test1") + .with_status(200) + .with_body(load_response.to_string()) + .create_async() + .await; + + let catalog = RestCatalog::new( + SessionContext::empty(), + RestCatalogConfig::builder().uri(server.url()).build(), + Some(Box::new(NoopAuthManager)), + Some(Arc::new(LocalFsStorageFactory)), + Runtime::current(), + None, + ); + let table = catalog + .load_table(&TableIdent::from_strs(["ns1", "test1"]).unwrap()) + .await + .unwrap(); + + let error = table.file_io().serialize_all().unwrap_err(); + assert_eq!(error.kind(), ErrorKind::FeatureUnsupported); + assert!( + table + .file_io() + .without_credential_provider() + .serialize_all() + .is_ok() + ); + + config_mock.assert_async().await; + load_table_mock.assert_async().await; + } + #[tokio::test] async fn test_update_table_404() { let mut server = Server::new_async().await; diff --git a/crates/catalog/rest/src/client.rs b/crates/catalog/rest/src/client.rs index 6de740a523..7a9853fe45 100644 --- a/crates/catalog/rest/src/client.rs +++ b/crates/catalog/rest/src/client.rs @@ -24,8 +24,10 @@ use reqwest::header::HeaderMap; use reqwest::{Client, IntoUrl, Method, RequestBuilder}; use serde::de::DeserializeOwned; -use crate::RestCatalogConfig; use crate::auth::{AuthSession, NoopSession}; +use crate::catalog::{ + RestCatalogConfig, disable_header_redaction_from_props, explicit_headers_from_props, +}; use crate::request::HttpRequest; use crate::response::HttpResponse; @@ -141,10 +143,40 @@ impl HttpClient { }) } - /// Testing only: the session authenticating this client's requests. - #[cfg(test)] - pub(crate) fn auth_session(&self) -> &Arc { - &self.auth_session + /// Derives a client for table-scoped requests. + /// + /// The connection pool and catalog headers are inherited, headers in the + /// effective table properties override them, and `auth_session` replaces + /// the catalog session only when the auth manager selected a table-specific + /// child session. + pub(crate) fn for_table( + &self, + props: &HashMap, + auth_session: Arc, + ) -> Result { + let mut extra_headers = self.extra_headers.clone(); + let table_headers = explicit_headers_from_props(props)?; + let has_table_authorization = table_headers.contains_key(http::header::AUTHORIZATION); + + if !Arc::ptr_eq(&self.auth_session, &auth_session) && !has_table_authorization { + // An inherited catalog Authorization header would otherwise be + // applied after, and overwrite, the table session's authentication. + extra_headers.remove(http::header::AUTHORIZATION); + } + extra_headers.extend(table_headers); + + Ok(Self { + client: self.client.clone(), + extra_headers, + disable_header_redaction: disable_header_redaction_from_props(props) + .unwrap_or(self.disable_header_redaction), + auth_session, + }) + } + + /// Returns the session authenticating this client's requests. + pub(crate) fn auth_session(&self) -> Arc { + Arc::clone(&self.auth_session) } /// Testing only: the bearer token `session` would attach. @@ -260,6 +292,25 @@ pub(crate) fn format_headers_redacted(headers: &HeaderMap, disable_redaction: bo pub(crate) fn deserialize_unexpected_catalog_error( response: HttpResponse, disable_header_redaction: bool, +) -> Error { + unexpected_catalog_error(response, disable_header_redaction, true) +} + +/// Builds an unexpected catalog error without retaining the response body. +/// +/// Credential endpoints use this because even an unsuccessful response may +/// contain credential material that must not be surfaced through an error. +pub(crate) fn unexpected_catalog_error_without_body( + response: HttpResponse, + disable_header_redaction: bool, +) -> Error { + unexpected_catalog_error(response, disable_header_redaction, false) +} + +fn unexpected_catalog_error( + response: HttpResponse, + disable_header_redaction: bool, + include_body: bool, ) -> Error { let err = Error::new( ErrorKind::Unexpected, @@ -272,7 +323,7 @@ pub(crate) fn deserialize_unexpected_catalog_error( ); let bytes = response.body(); - if bytes.is_empty() { + if !include_body || bytes.is_empty() { return err; } err.with_context("json", String::from_utf8_lossy(bytes)) @@ -295,6 +346,19 @@ mod tests { } } + #[derive(Debug)] + struct TableSession; + + #[async_trait::async_trait] + impl AuthSession for TableSession { + async fn authenticate(&self, request: &mut HttpRequest) -> Result<()> { + request + .headers_mut() + .insert("authorization", "Bearer table-token".parse().unwrap()); + Ok(()) + } + } + #[tokio::test] async fn test_a_truncated_body_error_names_the_url() { // `bytes()` builds its error without a URL, so `read` attaches the one @@ -375,6 +439,23 @@ mod tests { assert!(!err.contains("leaked"), "{err}"); } + #[test] + fn test_unexpected_error_can_omit_a_secret_body() { + let response = HttpResponse::new( + http::StatusCode::INTERNAL_SERVER_ERROR, + HeaderMap::new(), + br#"{"secret": "do-not-expose"}"#.to_vec(), + ); + + let err = format!( + "{:?}", + unexpected_catalog_error_without_body(response, false) + ); + + assert!(err.contains("500"), "{err}"); + assert!(!err.contains("do-not-expose"), "{err}"); + } + #[tokio::test] async fn test_post_form_carries_the_session_until_it_is_removed() { // Every request a client sends carries its session; a caller that @@ -412,6 +493,66 @@ mod tests { unsigned.assert_async().await; } + #[tokio::test] + async fn test_table_client_does_not_inherit_conflicting_authorization_header() { + let mut server = mockito::Server::new_async().await; + let table_session = server + .mock("GET", "/session") + .match_header("authorization", "Bearer table-token") + .with_status(200) + .create_async() + .await; + let table_header = server + .mock("GET", "/header") + .match_header("authorization", "Bearer table-header") + .with_status(200) + .create_async() + .await; + + let config = RestCatalogConfig::builder() + .uri(server.url()) + .props(HashMap::from([( + "header.Authorization".to_string(), + "Bearer catalog-header".to_string(), + )])) + .build(); + let catalog_client = HttpClient::new(&config).unwrap(); + + let session_client = catalog_client + .for_table(&HashMap::new(), Arc::new(TableSession)) + .unwrap(); + session_client + .query_catalog( + HttpRequest::build( + session_client.request(Method::GET, format!("{}/session", server.url())), + ) + .unwrap(), + ) + .await + .unwrap(); + table_session.assert_async().await; + + let header_client = catalog_client + .for_table( + &HashMap::from([( + "header.Authorization".to_string(), + "Bearer table-header".to_string(), + )]), + Arc::new(TableSession), + ) + .unwrap(); + header_client + .query_catalog( + HttpRequest::build( + header_client.request(Method::GET, format!("{}/header", server.url())), + ) + .unwrap(), + ) + .await + .unwrap(); + table_header.assert_async().await; + } + #[test] fn test_format_headers_redacted_empty() { let headers = HeaderMap::new(); diff --git a/crates/catalog/rest/src/credential.rs b/crates/catalog/rest/src/credential.rs new file mode 100644 index 0000000000..ff8a8b0c0b --- /dev/null +++ b/crates/catalog/rest/src/credential.rs @@ -0,0 +1,2452 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +//! Vended storage credentials from a REST catalog. +//! +//! A REST catalog can vend short-lived storage credentials whose lifetime the +//! client does not control. [`RestVendedCredentialProvider`] implements the +//! core [`StorageCredentialProvider`] trait so storage backends re-fetch those +//! credentials from the catalog's table credentials endpoint before they +//! expire, keeping long-running jobs authenticated instead of failing with a +//! `403` once the initial token's TTL elapses. +//! +//! Unlike the Java client, which has one provider per cloud SDK, this is a +//! single backend-agnostic provider with an independent endpoint and cache for +//! each configured cloud. The path being accessed selects the cloud cache, and +//! the returned [`StorageCredential`] enum lets the storage adapter enforce the +//! expected backend-specific type. This preserves Java's per-cloud credential +//! selection and prefetch policies while supporting mixed-cloud tables through +//! a resolving FileIO. Unlike Java's scheduled refresh, which permanently stops +//! after a failed fetch, transient failures are retried here with jittered +//! exponential backoff while an unexpired credential remains available. +//! +//! Like Java's `VendedCredentialsProvider`, a provider is rebuilt from the +//! FileIO properties after [`FileIO`](iceberg::io::FileIO) serialization and +//! connects to the catalog lazily in the receiving process. This requires +//! catalog authentication that can be rebuilt from `rest.auth.type`. +//! +//! # Adding a cloud +//! +//! The refresh policy for each cloud lives in one [`CloudRefresh`] constant. To +//! add a backend, first add its credential type to Iceberg's storage API and +//! teach the storage adapter to consume it. Then write its `parse_*` function, +//! add a `CloudRefresh` constant, and list it in [`CloudRefresh::SUPPORTED`]. + +use std::collections::{HashMap, HashSet}; +use std::sync::Arc; +use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH}; + +use async_trait::async_trait; +use iceberg::io::{ + ADLS_REFRESH_CREDENTIALS_ENABLED, ADLS_REFRESH_CREDENTIALS_ENDPOINT, + ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX, ADLS_SAS_TOKEN_PREFIX, AWS_REFRESH_CREDENTIALS_ENABLED, + AWS_REFRESH_CREDENTIALS_ENDPOINT, AzdlsCredential, GCS_REFRESH_CREDENTIALS_ENABLED, + GCS_REFRESH_CREDENTIALS_ENDPOINT, GCS_TOKEN, GCS_TOKEN_EXPIRES_AT, GcsCredential, + S3_ACCESS_KEY_ID, S3_SECRET_ACCESS_KEY, S3_SESSION_TOKEN, S3_SESSION_TOKEN_EXPIRES_AT_MS, + S3Credential, StorageConfig, StorageCredential, StorageCredentialKind, + StorageCredentialProvider, StorageCredentialProviderFactory, storage_prefix_covers, +}; +use iceberg::{Error, ErrorKind, Result, TableIdent}; +use rand::Rng; +use reqwest::{Method, StatusCode, Url}; +use serde::{Deserialize, Serialize}; +use tokio::sync::{Mutex, OnceCell}; + +use crate::auth::{AuthManager, load_auth_manager}; +use crate::catalog::{REST_CATALOG_PROP_SCAN_PLAN_ID, RestCatalogConfig}; +use crate::client::{HttpClient, unexpected_catalog_error_without_body}; +use crate::request::HttpRequest; +use crate::types::LoadCredentialsResponse; + +type CredentialParser = + fn(config: &HashMap, prefix: Option) -> Result; +type KeyedSeedCredentialParser = + fn(config: &HashMap) -> HashMap; +type KeyedSeedPathResolver = fn(path: &str) -> Result; + +enum SeedStrategy { + /// One credential stored in the backend's flat properties. + Flat, + /// Credentials selected by a backend-specific key derived from each path. + Keyed(KeyedSeedStrategy), +} + +struct KeyedSeedStrategy { + parse_credentials: KeyedSeedCredentialParser, + resolve_path: KeyedSeedPathResolver, +} + +struct KeyedSeedPath { + key: String, + scope: String, +} + +/// Cloud-specific details regarding vended-credential refresh. +/// +/// It contains the location schemes it backs, the property keys it is configured +/// with, and how to parse its credential. The generic provider stays free of +/// any per-cloud knowledge. +struct CloudRefresh { + /// Location URL schemes this backend serves. + schemes: &'static [&'static str], + /// Table property naming the refresh endpoint (absolute or catalog-relative). + endpoint_key: &'static str, + /// Table property controlling refresh; only missing or case-insensitive `"true"` enables it. + enabled_key: &'static str, + /// Whether to jitter successful prefetch times like AWS `CachedSupplier`. + jitter_prefetch: bool, + /// Parse a complete credential from catalog-supplied properties. + parse_credential: CredentialParser, + /// How initial credentials in the table properties are selected. + seed_strategy: SeedStrategy, +} + +impl CloudRefresh { + /// S3 / AWS + const AWS: Self = Self { + schemes: &["s3", "s3a", "s3n"], + endpoint_key: AWS_REFRESH_CREDENTIALS_ENDPOINT, + enabled_key: AWS_REFRESH_CREDENTIALS_ENABLED, + jitter_prefetch: true, + parse_credential: parse_s3_credential, + seed_strategy: SeedStrategy::Flat, + }; + /// Google Cloud Storage + const GCP: Self = Self { + schemes: &["gs", "gcs"], + endpoint_key: GCS_REFRESH_CREDENTIALS_ENDPOINT, + enabled_key: GCS_REFRESH_CREDENTIALS_ENABLED, + jitter_prefetch: false, + parse_credential: parse_gcs_credential, + seed_strategy: SeedStrategy::Flat, + }; + /// Azure Data Lake Storage + const AZURE: Self = Self { + schemes: &["abfs", "abfss", "wasb", "wasbs"], + endpoint_key: ADLS_REFRESH_CREDENTIALS_ENDPOINT, + enabled_key: ADLS_REFRESH_CREDENTIALS_ENABLED, + jitter_prefetch: false, + parse_credential: parse_azdls_credential, + seed_strategy: SeedStrategy::Keyed(KeyedSeedStrategy { + parse_credentials: parse_azdls_account_seeds, + resolve_path: resolve_azdls_seed_path, + }), + }; + + /// Backends with refresh support + const SUPPORTED: &[Self] = &[Self::AWS, Self::GCP, Self::AZURE]; + + fn matches_location(&self, location: &str) -> bool { + self.schemes + .iter() + .any(|scheme| location.eq_ignore_ascii_case(scheme)) + || scheme_of(location).is_some_and(|scheme| self.schemes.contains(&scheme.as_str())) + } +} + +/// Re-fetch a credential once it is within this window of expiry, so a fresh +/// token is in hand before the object store would reject the old one. +const REFRESH_BUFFER: Duration = Duration::from_mins(5); + +/// AWS keeps at least one minute between its jittered prefetch time and expiry. +const MIN_REFRESH_BUFFER: Duration = Duration::from_mins(1); + +/// Initial ceiling for failure backoff. Equal jitter chooses from half this +/// value through the full value. +const INITIAL_FAILURE_BACKOFF: Duration = Duration::from_secs(1); + +/// Maximum delay between failed refresh attempts. +const MAX_FAILURE_BACKOFF: Duration = Duration::from_secs(30); + +/// A cached vended credential and its refresh schedule. +#[derive(Clone)] +struct CachedEntry { + credential: StorageCredential, + /// When this entry becomes eligible for prefetch. `None` means it does not + /// expire and therefore never needs proactive refresh. + refresh_at: Option, +} + +impl CachedEntry { + fn new(credential: StorageCredential, jitter_prefetch: bool) -> Self { + let refresh_at = credential + .expires_at() + .map(|expires_at| prefetch_time(SystemTime::now(), expires_at, jitter_prefetch)); + Self { + credential, + refresh_at, + } + } + + /// Seed entries that are already inside the nominal five-minute window are + /// immediately due. Otherwise AWS applies the same jitter as it does to a + /// freshly fetched value. + fn seed(credential: StorageCredential, jitter_prefetch: bool) -> Self { + let due = credential.expires_at().is_some_and(|expires_at| { + SystemTime::now() + .checked_add(REFRESH_BUFFER) + .is_none_or(|refresh_boundary| refresh_boundary >= expires_at) + }); + let mut entry = Self::new(credential, jitter_prefetch); + if due { + entry.refresh_at = Some(UNIX_EPOCH); + } + entry + } + + fn is_fresh(&self, now: SystemTime) -> bool { + self.refresh_at.is_none_or(|refresh_at| now < refresh_at) + } + + fn is_unexpired(&self, now: SystemTime) -> bool { + self.credential + .expires_at() + .is_none_or(|expires_at| now < expires_at) + } +} + +/// Parsed entries from one successful credentials response. +struct ParsedCredentials { + entries: Vec, + errors: Vec, +} + +struct CredentialError { + prefix: String, + error: Error, +} + +/// Cached credentials plus failure-backoff state. +/// +/// A fetch returns the credentials for every prefix of the table, so a path +/// the last response did not cover would not be covered by an immediate +/// re-fetch either. Backoff is therefore shared by the whole cloud. +struct CacheState { + entries: Vec, + consecutive_failures: u32, + retry_not_before: Option, +} + +impl CacheState { + fn new(entries: Vec) -> Self { + Self { + entries, + consecutive_failures: 0, + retry_not_before: None, + } + } + + fn record_success(&mut self) { + self.consecutive_failures = 0; + self.retry_not_before = None; + } + + fn record_failure(&mut self) { + self.consecutive_failures = self.consecutive_failures.saturating_add(1); + self.retry_not_before = + Instant::now().checked_add(failure_backoff(self.consecutive_failures)); + } + + /// Replace cached credentials with the fetched ones, prefix by prefix. + /// + /// An unexpired cached credential survives when the response carries no + /// valid replacement for its prefix, so an absent or malformed entry never + /// evicts a usable credential. Returns the prefixes the response replaced. + fn merge(&mut self, fetched: Vec, now: SystemTime) -> HashSet { + let fetched_prefixes = fetched + .iter() + .filter_map(|entry| entry.credential.prefix().map(str::to_owned)) + .collect::>(); + self.entries.retain(|entry| { + entry.is_unexpired(now) + && entry + .credential + .prefix() + .is_none_or(|prefix| !fetched_prefixes.contains(prefix)) + }); + self.entries.extend(fetched); + fetched_prefixes + } +} + +struct ConfiguredCloud { + cloud: &'static CloudRefresh, + endpoint: String, + cache: Mutex, + /// Initial property credentials whose scope cannot be represented by a + /// single leading URI prefix, keyed according to the cloud strategy. + keyed_seeds: HashMap, + /// Only one caller fetches at a time. The cache lock is deliberately + /// separate so other callers can keep using an unexpired credential while + /// the refresh is in flight. + refresh: Mutex<()>, +} + +impl ConfiguredCloud { + fn new( + cloud: &'static CloudRefresh, + endpoint: String, + entries: Vec, + keyed_seeds: HashMap, + ) -> Self { + Self { + cloud, + endpoint, + cache: Mutex::new(CacheState::new(entries)), + keyed_seeds, + refresh: Mutex::new(()), + } + } + + /// Configure `cloud` from the FileIO properties, or `None` when they + /// advertise no enabled refresh endpoint for it. + fn configure( + cloud: &'static CloudRefresh, + base_uri: &str, + props: &HashMap, + ) -> Option { + let enabled = props + .get(cloud.enabled_key) + .is_none_or(|value| value.eq_ignore_ascii_case("true")); + let endpoint = props + .get(cloud.endpoint_key) + .filter(|endpoint| enabled && !endpoint.is_empty()) + .map(|endpoint| resolve_endpoint(base_uri, endpoint))?; + + let mut entries = Vec::new(); + let mut keyed_seeds = HashMap::new(); + match &cloud.seed_strategy { + SeedStrategy::Flat => { + if let Ok(credential) = (cloud.parse_credential)(props, None) { + entries.push(CachedEntry::seed(credential, cloud.jitter_prefetch)); + } + } + SeedStrategy::Keyed(strategy) => { + keyed_seeds = (strategy.parse_credentials)(props) + .into_iter() + .map(|(key, credential)| { + (key, CachedEntry::seed(credential, cloud.jitter_prefetch)) + }) + .collect(); + } + } + Some(Self::new(cloud, endpoint, entries, keyed_seeds)) + } + + fn keyed_seed_for_path(&self, path: &str) -> Result> { + let SeedStrategy::Keyed(strategy) = &self.cloud.seed_strategy else { + return Ok(None); + }; + let resolved = (strategy.resolve_path)(path)?; + let Some(seed) = self.keyed_seeds.get(&resolved.key) else { + return Ok(None); + }; + Ok(Some(CachedEntry { + credential: seed.credential.clone().with_prefix(resolved.scope), + refresh_at: seed.refresh_at, + })) + } +} + +/// Serializable recipe for rebuilding a [`RestVendedCredentialProvider`] from +/// the FileIO properties, mirroring how Java rebuilds its provider on workers. +#[derive(Clone, Serialize, Deserialize)] +pub(crate) struct RestVendedCredentialProviderFactory { + /// Catalog URI, used to resolve relative refresh endpoints. + catalog_uri: String, + table: TableIdent, + /// The unmerged config returned by the table endpoint, from which the + /// auth manager derives a table session. Keeping it separate prevents + /// local FileIO overrides from masking table auth. + table_config: HashMap, +} + +impl RestVendedCredentialProviderFactory { + pub(crate) fn new( + catalog_uri: impl Into, + table: TableIdent, + table_config: HashMap, + ) -> Self { + Self { + catalog_uri: catalog_uri.into(), + table, + table_config, + } + } + + /// Connect to the catalog from the FileIO properties, as the catalog would. + async fn connect(&self, props: &HashMap) -> Result { + let config = RestCatalogConfig::builder() + .uri(self.catalog_uri.clone()) + .props(props.clone()) + .build(); + let auth_manager = load_auth_manager(&config)?; + let client = HttpClient::new(&config)?; + let session = auth_manager + .catalog_session(&client.without_auth_session(), &config.auth_props()) + .await?; + table_client( + &client.with_auth_session(session), + auth_manager.as_ref(), + &self.table, + props, + &self.table_config, + ) + .await + } +} + +impl std::fmt::Debug for RestVendedCredentialProviderFactory { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("RestVendedCredentialProviderFactory") + .field("catalog_uri", &self.catalog_uri) + .field("table", &self.table) + .finish_non_exhaustive() + } +} + +#[typetag::serde(name = "RestVendedCredentialProviderFactory")] +impl StorageCredentialProviderFactory for RestVendedCredentialProviderFactory { + fn build(&self, config: &StorageConfig) -> Result> { + let provider = + RestVendedCredentialProvider::configure(self, config.props()).ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + "FileIO configuration no longer configures vended credentials", + ) + })?; + Ok(Arc::new(RestVendedCredentialProvider { + factory: Some(self.clone()), + ..provider + })) + } +} + +/// Derive the table-scoped client used for credential requests. +async fn table_client( + catalog_client: &HttpClient, + auth_manager: &dyn AuthManager, + table: &TableIdent, + props: &HashMap, + table_config: &HashMap, +) -> Result { + let session = auth_manager + .table_session( + &catalog_client.without_auth_session(), + table, + table_config, + catalog_client.auth_session(), + ) + .await?; + catalog_client.for_table(props, session) +} + +/// Serves vended credentials for a table, refreshing them from the REST +/// catalog's table credentials endpoint. +/// +/// Each cloud cache is seeded with the credentials from the initial +/// table properties (when complete) and re-fetched from its endpoint as they +/// near expiry. +pub(crate) struct RestVendedCredentialProvider { + /// Table-scoped catalog client. Set when the catalog builds the provider, + /// and connected on first refresh after deserialization. + client: OnceCell, + /// Rebuilds this provider in another process. `None` when the catalog + /// authentication cannot be rebuilt from properties. + factory: Option, + /// Effective FileIO properties, which supply catalog connection settings. + props: HashMap, + /// Optional scan-plan identifier. + plan_id: Option, + /// Independently configured endpoint and cache for each backing cloud. + clouds: Vec, +} + +impl RestVendedCredentialProvider { + /// Configure a provider without a catalog connection, or `None` when no + /// supported cloud advertises an enabled refresh endpoint. + fn configure( + factory: &RestVendedCredentialProviderFactory, + props: &HashMap, + ) -> Option { + let clouds = CloudRefresh::SUPPORTED + .iter() + .filter_map(|cloud| ConfiguredCloud::configure(cloud, &factory.catalog_uri, props)) + .collect::>(); + (!clouds.is_empty()).then(|| Self { + client: OnceCell::new(), + factory: None, + props: props.clone(), + plan_id: props.get(REST_CATALOG_PROP_SCAN_PLAN_ID).cloned(), + clouds, + }) + } + + fn configured_cloud_for_location(&self, location: &str) -> Option<&ConfiguredCloud> { + self.clouds + .iter() + .find(|configured| configured.cloud.matches_location(location)) + } + + async fn client(&self) -> Result<&HttpClient> { + self.client + .get_or_try_init(|| async { + let factory = self.factory.as_ref().ok_or_else(|| { + Error::new( + ErrorKind::Unexpected, + "vended credential provider has no catalog connection", + ) + })?; + factory.connect(&self.props).await + }) + .await + } + + /// Fetch fresh credentials from the catalog's credentials endpoint. + async fn fetch(&self, configured: &ConfiguredCloud) -> Result { + let cloud = configured.cloud; + let client = self.client().await?; + let mut request = client.request(Method::GET, &configured.endpoint); + if let Some(plan_id) = &self.plan_id { + request = request.query(&[("planId", plan_id)]); + } + let request = HttpRequest::build(request)?; + let response = client.query_catalog(request).await?; + + if response.status() != StatusCode::OK { + return Err(unexpected_catalog_error_without_body( + response, + client.disable_header_redaction(), + )); + } + + // Credential responses contain secrets. Do not include the response + // body in a deserialization error. + let parsed: LoadCredentialsResponse = + serde_json::from_slice(response.body()).map_err(|error| { + Error::new( + ErrorKind::Unexpected, + "failed to parse vended credential response", + ) + .with_source(error) + })?; + let now = SystemTime::now(); + let mut entries = Vec::new(); + let mut errors = Vec::new(); + for credential in parsed + .storage_credentials + .into_iter() + .filter(|credential| cloud.matches_location(&credential.prefix)) + { + let prefix = credential.prefix; + let parsed = (cloud.parse_credential)(&credential.config, Some(prefix.clone())) + .map(|credential| CachedEntry::new(credential, cloud.jitter_prefetch)) + .and_then(|entry| { + if entry.is_unexpired(now) { + Ok(entry) + } else { + Err(Error::new( + ErrorKind::DataInvalid, + "invalid vended credential: credential is already expired", + )) + } + }); + match parsed { + Ok(entry) => entries.push(entry), + Err(error) => errors.push(CredentialError { prefix, error }), + } + } + Ok(ParsedCredentials { entries, errors }) + } + + async fn refresh_credential( + &self, + configured: &ConfiguredCloud, + path: &str, + fallback: Option, + ) -> Result { + let fetched = self.fetch(configured).await; + let mut cache = configured.cache.lock().await; + let now = SystemTime::now(); + + let (fetched_prefixes, failure) = match fetched { + Ok(ParsedCredentials { entries, errors }) => { + let failure = errors + .into_iter() + .filter(|error| storage_prefix_covers(&error.prefix, path)) + .max_by_key(|error| error.prefix.len()) + .map(|error| error.error); + (cache.merge(entries, now), failure) + } + Err(error) => (HashSet::new(), Some(error)), + }; + + let selected = longest_prefix_match(&cache.entries, path) + .filter(|entry| entry.is_unexpired(now)) + .cloned(); + let refreshed = selected.as_ref().is_some_and(|entry| { + entry + .credential + .prefix() + .is_some_and(|prefix| fetched_prefixes.contains(prefix)) + }); + if refreshed { + cache.record_success(); + } else { + // Graceful degradation: while a credential for this path remains + // usable, serve it and retry after jittered backoff. Expired + // credentials are never served. + cache.record_failure(); + } + + selected + .or_else(|| fallback.filter(|fallback| fallback.is_unexpired(now))) + .map(|entry| entry.credential) + .ok_or_else(|| { + failure.unwrap_or_else(|| { + Error::new( + ErrorKind::Unexpected, + format!("no unexpired vended credential matches storage location: {path}"), + ) + }) + }) + } +} + +impl std::fmt::Debug for RestVendedCredentialProvider { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("RestVendedCredentialProvider") + .field("configured_clouds", &self.clouds.len()) + .finish_non_exhaustive() + } +} + +enum CacheDecision { + Use(StorageCredential), + Refresh(Option), + Backoff, +} + +fn refresh_backoff_error(path: &str) -> Error { + Error::new( + ErrorKind::Unexpected, + format!("vended credential refresh is temporarily backed off for storage location: {path}"), + ) +} + +async fn cache_decision(configured: &ConfiguredCloud, path: &str) -> Result { + let keyed_seed = configured.keyed_seed_for_path(path)?; + let cache = configured.cache.lock().await; + let now = SystemTime::now(); + let current = match longest_prefix_match(&cache.entries, path).cloned() { + Some(cached) if cached.is_unexpired(now) => Some(cached), + Some(expired) => keyed_seed.or(Some(expired)), + None => keyed_seed, + }; + + if let Some(entry) = current.as_ref().filter(|entry| entry.is_fresh(now)) { + return Ok(CacheDecision::Use(entry.credential.clone())); + } + + if cache + .retry_not_before + .is_some_and(|retry_at| Instant::now() < retry_at) + { + return Ok(current + .filter(|entry| entry.is_unexpired(now)) + .map(|entry| CacheDecision::Use(entry.credential)) + .unwrap_or(CacheDecision::Backoff)); + } + + Ok(CacheDecision::Refresh(current)) +} + +#[async_trait] +impl StorageCredentialProvider for RestVendedCredentialProvider { + fn supports_path(&self, path: &str) -> bool { + self.configured_cloud_for_location(path).is_some() + } + + async fn load_credential(&self, path: &str) -> Result { + let configured = self.configured_cloud_for_location(path).ok_or_else(|| { + Error::new( + ErrorKind::FeatureUnsupported, + format!("vended credentials are not configured for storage location: {path}"), + ) + })?; + + let current = match cache_decision(configured, path).await? { + CacheDecision::Use(credential) => return Ok(credential), + CacheDecision::Refresh(current) => current, + CacheDecision::Backoff => return Err(refresh_backoff_error(path)), + }; + + // One caller refreshes, while concurrent callers immediately keep using the + // unexpired cached credential. With no usable credential, callers wait + // for the in-flight refresh instead. + let usable = current + .as_ref() + .filter(|entry| entry.is_unexpired(SystemTime::now())); + let _refresh_guard = if let Some(entry) = usable { + match configured.refresh.try_lock() { + Ok(guard) => guard, + Err(_) => return Ok(entry.credential.clone()), + } + } else { + configured.refresh.lock().await + }; + + // Another caller may have completed a refresh between our cache check + // and acquiring the single-flight guard. + let current = match cache_decision(configured, path).await? { + CacheDecision::Use(credential) => return Ok(credential), + CacheDecision::Refresh(current) => current, + CacheDecision::Backoff => return Err(refresh_backoff_error(path)), + }; + + self.refresh_credential(configured, path, current).await + } + + fn factory(&self) -> Result> { + match &self.factory { + Some(factory) => Ok(Arc::new(factory.clone())), + None => Err(Error::new( + ErrorKind::FeatureUnsupported, + "the vended credential provider cannot be serialized because the REST catalog \ + uses an injected AuthManager, which cannot be rebuilt in another process; use \ + FileIO::without_credential_provider to serialize without credential refresh", + )), + } + } +} + +/// Select the credential whose prefix is the longest match for `path`. +fn longest_prefix_match<'a>(entries: &'a [CachedEntry], path: &str) -> Option<&'a CachedEntry> { + entries + .iter() + .filter(|entry| entry.credential.covers(path)) + .max_by_key(|entry| entry.credential.prefix().map_or(0, str::len)) +} + +/// Compute the prefetch time of a credential obtained at `now`. +/// +/// Like Java, a credential is refreshed [`REFRESH_BUFFER`] before it expires. +/// Unlike Java, the buffer is capped at half the remaining lifetime: a +/// credential vended with a shorter lifetime would otherwise be due on +/// arrival, and every file operation would fetch again. +fn prefetch_time(now: SystemTime, expires_at: SystemTime, jitter: bool) -> SystemTime { + let lifetime = expires_at.duration_since(now).unwrap_or_default(); + let buffer = REFRESH_BUFFER.min(lifetime / 2); + let base = expires_at.checked_sub(buffer).unwrap_or(UNIX_EPOCH); + if !jitter { + return base; + } + + // The minimum distance from expiry shrinks with a capped buffer. + let min_buffer = + buffer.mul_f64(MIN_REFRESH_BUFFER.as_secs_f64() / REFRESH_BUFFER.as_secs_f64()); + let jitter_millis = buffer.saturating_sub(min_buffer).as_millis() as u64; + if jitter_millis == 0 { + return base; + } + base.checked_add(Duration::from_millis( + rand::rng().random_range(0..jitter_millis), + )) + .unwrap_or(base) +} + +/// Equal-jitter exponential backoff. The random lower half avoids both hot +/// retry loops and synchronized retries across clients. +fn failure_backoff(consecutive_failures: u32) -> Duration { + let exponent = consecutive_failures.saturating_sub(1).min(5); + let ceiling = INITIAL_FAILURE_BACKOFF + .checked_mul(1 << exponent) + .unwrap_or(MAX_FAILURE_BACKOFF) + .min(MAX_FAILURE_BACKOFF); + let ceiling_millis = ceiling.as_millis() as u64; + let floor_millis = ceiling_millis / 2; + Duration::from_millis(rand::rng().random_range(floor_millis..=ceiling_millis)) +} + +/// Build a credential provider for a table, or `None` when no supported cloud +/// advertises an enabled refresh endpoint. +/// +/// `catalog_client` carries the catalog session, from which `auth_manager` +/// derives a table session when the table config overrides authentication. +/// `props` contains the effective FileIO properties after applying local +/// overrides, and therefore supplies headers and client policy. `portable` +/// states whether the catalog authentication can be rebuilt from `props` in +/// another process, which makes the provider serializable. +pub(crate) async fn build_vended_credential_provider( + catalog_client: &HttpClient, + auth_manager: &dyn AuthManager, + factory: RestVendedCredentialProviderFactory, + props: &HashMap, + portable: bool, +) -> Result>> { + let Some(provider) = RestVendedCredentialProvider::configure(&factory, props) else { + return Ok(None); + }; + + let client = table_client( + catalog_client, + auth_manager, + &factory.table, + props, + &factory.table_config, + ) + .await?; + Ok(Some(Arc::new(RestVendedCredentialProvider { + client: OnceCell::new_with(Some(client)), + factory: portable.then_some(factory), + ..provider + }))) +} + +/// Resolve a possibly-relative refresh endpoint against the catalog base URI. +/// +/// Absolute endpoints are used as-is and receive the catalog credentials, as +/// in Java. The catalog is trusted to advertise only its own endpoints. +fn resolve_endpoint(base_uri: &str, endpoint: &str) -> String { + if endpoint.starts_with("http://") || endpoint.starts_with("https://") { + return endpoint.to_string(); + } + + let base = base_uri.trim_end_matches('/'); + let separator = if endpoint.starts_with('/') { "" } else { "/" }; + format!("{base}{separator}{endpoint}") +} + +/// The URL scheme of `location`, lowercased (e.g. `"s3"` for `s3://bucket/k`). +fn scheme_of(location: &str) -> Option { + Url::parse(location) + .ok() + .map(|url| url.scheme().to_string()) +} + +/// Parse a complete S3 credential supplied by the catalog. +fn parse_s3_credential( + config: &HashMap, + prefix: Option, +) -> Result { + let access_key_id = required_nonempty(config, S3_ACCESS_KEY_ID)?; + let secret_access_key = required_nonempty(config, S3_SECRET_ACCESS_KEY)?; + let session_token = required_nonempty(config, S3_SESSION_TOKEN)?; + let expires_at = required_epoch_millis(config, S3_SESSION_TOKEN_EXPIRES_AT_MS)?; + Ok(with_prefix( + StorageCredential::new(StorageCredentialKind::S3(S3Credential::new( + access_key_id, + secret_access_key, + Some(session_token), + ))) + .with_expiration(expires_at), + prefix, + )) +} + +/// Parse a complete GCS credential supplied by the catalog. +fn parse_gcs_credential( + config: &HashMap, + prefix: Option, +) -> Result { + let token = required_nonempty(config, GCS_TOKEN)?; + let expires_at = required_epoch_millis(config, GCS_TOKEN_EXPIRES_AT)?; + Ok(with_prefix( + StorageCredential::new(StorageCredentialKind::Gcs(GcsCredential::new(token))) + .with_expiration(expires_at), + prefix, + )) +} + +fn with_prefix(credential: StorageCredential, prefix: Option) -> StorageCredential { + match prefix { + Some(prefix) => credential.with_prefix(prefix), + None => credential, + } +} + +/// Parse a complete account-specific ADLS SAS credential supplied by the catalog. +fn parse_azdls_credential( + config: &HashMap, + prefix: Option, +) -> Result { + let prefix = prefix.ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + "invalid vended ADLS credential: storage prefix is missing", + ) + })?; + let account = azdls_account_name(&prefix)?; + let (suffix, sas_token) = azdls_sas_tokens(config) + .filter(|(_, token_account, _)| *token_account == account) + .map(|(suffix, _, token)| (suffix, token)) + // Prefer the host-keyed token when both forms name the account, as + // the storage backend does. + .max() + .ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + format!("invalid vended credential: no {ADLS_SAS_TOKEN_PREFIX}* token for account {account}"), + ) + })?; + let expires_at = required_epoch_millis( + config, + &format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}{suffix}"), + )?; + Ok( + StorageCredential::new(StorageCredentialKind::Azdls(AzdlsCredential::new( + sas_token, + ))) + .with_prefix(prefix) + .with_expiration(expires_at), + ) +} + +/// Parse every complete account-qualified ADLS credential from the initial +/// table properties. Unlike a URI prefix, an Azure account occurs after the +/// filesystem in a location, so these seeds are selected by account name. +fn parse_azdls_account_seeds( + config: &HashMap, +) -> HashMap { + let mut seeds = azdls_sas_tokens(config) + .filter_map(|(suffix, account, token)| { + let expires_at = required_epoch_millis( + config, + &format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}{suffix}"), + ) + .ok()?; + Some((suffix, account, token, expires_at)) + }) + .collect::>(); + // Deterministic choice when several keys name the same account: the last, + // host-keyed one wins, as in the storage backend. + seeds.sort_by(|left, right| left.0.cmp(right.0)); + seeds + .into_iter() + .map(|(_, account, token, expires_at)| { + ( + account.to_string(), + StorageCredential::new(StorageCredentialKind::Azdls(AzdlsCredential::new(token))) + .with_expiration(expires_at), + ) + }) + .collect() +} + +/// Account-specific SAS tokens as `(key suffix, account, token)`. Like Java, +/// keys may name the host (`account.dfs.core.windows.net`) or only the account. +fn azdls_sas_tokens(config: &HashMap) -> impl Iterator { + config.iter().filter_map(|(key, token)| { + let suffix = key.strip_prefix(ADLS_SAS_TOKEN_PREFIX)?; + let account = suffix.split('.').next().unwrap_or(suffix); + (!account.is_empty() && !token.is_empty()).then_some((suffix, account, token.as_str())) + }) +} + +fn resolve_azdls_seed_path(location: &str) -> Result { + let (mut url, account) = parse_azdls_location(location)?; + url.set_path("/"); + url.set_query(None); + url.set_fragment(None); + Ok(KeyedSeedPath { + key: account, + scope: url.to_string(), + }) +} + +fn azdls_account_name(location: &str) -> Result { + parse_azdls_location(location).map(|(_, account)| account) +} + +fn parse_azdls_location(location: &str) -> Result<(Url, String)> { + let url = Url::parse(location).map_err(|error| { + Error::new( + ErrorKind::DataInvalid, + format!("invalid ADLS storage location: {location}"), + ) + .with_source(error) + })?; + if !CloudRefresh::AZURE.schemes.contains(&url.scheme()) { + return Err(Error::new( + ErrorKind::DataInvalid, + format!("invalid ADLS storage location scheme: {}", url.scheme()), + )); + } + let account = url + .host_str() + .and_then(|host| host.split('.').next()) + .filter(|account| !account.is_empty()) + .map(str::to_owned) + .ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + format!("ADLS storage location has no account name: {location}"), + ) + })?; + Ok((url, account)) +} + +fn required_nonempty(config: &HashMap, key: &str) -> Result { + config + .get(key) + .filter(|value| !value.is_empty()) + .cloned() + .ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + format!("invalid vended credential: {key} is missing or empty"), + ) + }) +} + +fn required_epoch_millis(config: &HashMap, key: &str) -> Result { + let value = required_nonempty(config, key)?; + parse_epoch_millis(&value).ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + format!("invalid vended credential: {key} is not a valid epoch-millisecond timestamp"), + ) + }) +} + +/// Parse an epoch-millisecond timestamp into a [`SystemTime`]. +fn parse_epoch_millis(millis: &str) -> Option { + millis + .parse() + .ok() + .and_then(|millis| UNIX_EPOCH.checked_add(Duration::from_millis(millis))) +} + +#[cfg(test)] +mod tests { + use std::sync::{Barrier, mpsc}; + + use mockito::{Matcher, Server}; + + use super::*; + use crate::auth::{NoopAuthManager, OAuth2Manager}; + + fn epoch_millis(time: SystemTime) -> String { + time.duration_since(UNIX_EPOCH) + .unwrap() + .as_millis() + .to_string() + } + + fn s3_cred( + prefix: Option<&str>, + access_key_id: &str, + expires_at: Option, + ) -> StorageCredential { + let mut credential = StorageCredential::new(StorageCredentialKind::S3(S3Credential::new( + access_key_id, + "secret", + None, + ))); + if let Some(prefix) = prefix { + credential = credential.with_prefix(prefix); + } + if let Some(expires_at) = expires_at { + credential = credential.with_expiration(expires_at); + } + credential + } + + fn s3_access_key_id(credential: &StorageCredential) -> &str { + match credential.kind() { + StorageCredentialKind::S3(s3) => s3.access_key_id(), + other => panic!("expected S3 credential, got {other:?}"), + } + } + + /// A previously cached entry: like a seed, its age is unknown, so it is due + /// once inside the nominal refresh window. + fn cached_s3(prefix: &str, access_key_id: &str, expires_at: Option) -> CachedEntry { + CachedEntry::seed(s3_cred(Some(prefix), access_key_id, expires_at), false) + } + + fn test_client(base_uri: &str) -> HttpClient { + let config = RestCatalogConfig::builder() + .uri(base_uri.to_string()) + .build(); + HttpClient::new(&config).unwrap() + } + + fn test_factory( + base_uri: &str, + table_config: Option<&HashMap>, + ) -> RestVendedCredentialProviderFactory { + RestVendedCredentialProviderFactory::new( + base_uri, + test_table(), + table_config.cloned().unwrap_or_default(), + ) + } + + /// A provider whose AWS cache starts with `entries`. + fn provider_with_cached_s3( + base_uri: &str, + entries: Vec, + ) -> RestVendedCredentialProvider { + RestVendedCredentialProvider { + client: OnceCell::new_with(Some(test_client(base_uri))), + factory: None, + props: HashMap::new(), + plan_id: None, + clouds: vec![ConfiguredCloud::new( + &CloudRefresh::AWS, + format!("{base_uri}/v1/credentials"), + entries, + HashMap::new(), + )], + } + } + + fn aws_refresh_props(endpoint: &str) -> HashMap { + HashMap::from([( + CloudRefresh::AWS.endpoint_key.to_string(), + endpoint.to_string(), + )]) + } + + fn with_s3_seed( + mut props: HashMap, + access_key_id: &str, + expires_at: SystemTime, + ) -> HashMap { + props.extend([ + (S3_ACCESS_KEY_ID.to_string(), access_key_id.to_string()), + (S3_SECRET_ACCESS_KEY.to_string(), "SEED_SK".to_string()), + (S3_SESSION_TOKEN.to_string(), "SEED_TOK".to_string()), + ( + S3_SESSION_TOKEN_EXPIRES_AT_MS.to_string(), + epoch_millis(expires_at), + ), + ]); + props + } + + fn test_table() -> TableIdent { + TableIdent::from_strs(["namespace", "table"]).unwrap() + } + + async fn build_test_provider( + client: HttpClient, + base_uri: &str, + props: &HashMap, + table_auth_props: Option<&HashMap>, + ) -> Result>> { + build_vended_credential_provider( + &client, + &NoopAuthManager, + test_factory(base_uri, table_auth_props), + props, + true, + ) + .await + } + + async fn test_provider( + base_uri: &str, + props: &HashMap, + ) -> Arc { + build_test_provider(test_client(base_uri), base_uri, props, None) + .await + .expect("provider construction should succeed") + .expect("provider should be built") + } + + fn s3_response(prefix: &str, access_key_id: &str, expires_at: SystemTime) -> String { + let expires_at = epoch_millis(expires_at); + format!( + r#"{{"storage-credentials":[{{"prefix":"{prefix}","config":{{"s3.access-key-id":"{access_key_id}","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expires_at}"}}}}]}}"# + ) + } + + #[test] + fn resolve_endpoint_matches_java_semantics() { + assert_eq!( + resolve_endpoint("https://catalog/", "https://other/creds"), + "https://other/creds" + ); + assert_eq!( + resolve_endpoint("https://catalog", "http://other/creds"), + "http://other/creds" + ); + assert_eq!( + resolve_endpoint("https://catalog", "v1/creds"), + "https://catalog/v1/creds" + ); + assert_eq!( + resolve_endpoint("https://catalog/", "v1/creds"), + "https://catalog/v1/creds" + ); + // All trailing slashes stripped from the base (Java stripTrailingSlash). + assert_eq!( + resolve_endpoint("https://catalog///", "/v1/creds"), + "https://catalog/v1/creds" + ); + // Existing leading slashes on the endpoint are preserved, not collapsed + // (matches Java's resolveEndpoint, which only prepends when absent). + assert_eq!( + resolve_endpoint("https://catalog/", "//v1/creds"), + "https://catalog//v1/creds" + ); + } + + #[test] + fn cloud_selection_accepts_root_prefixes_and_url_schemes() { + let supported = |location: &str| { + CloudRefresh::SUPPORTED + .iter() + .any(|cloud| cloud.matches_location(location)) + }; + // Java uses these root prefixes for its fallback clients and accepts + // credentials scoped directly to them. + assert!(supported("s3")); + assert!(supported("gs")); + assert!(supported("s3://b/k")); + assert!(supported("S3://b/k")); + assert!(supported("s3a://b/k")); + assert!(supported("s3n://b/k")); + assert!(supported("gs://b/k")); + assert!(supported("gcs://b/k")); + assert!(supported("abfs")); + assert!(supported("abfss://fs@acct.dfs.core.windows.net/k")); + assert!(supported("wasb://fs@acct.blob.core.windows.net/k")); + assert!(supported("wasbs://fs@acct.blob.core.windows.net/k")); + assert!(!supported("not a url")); + assert!(CloudRefresh::AWS.matches_location("s3://bucket/path")); + assert!(!CloudRefresh::AWS.matches_location("s3evil://bucket/path")); + } + + #[test] + fn parses_s3_credential() { + let mut config = HashMap::new(); + assert!(parse_s3_credential(&config, None).is_err()); + config.insert(S3_ACCESS_KEY_ID.to_string(), "AK".to_string()); + assert!(parse_s3_credential(&config, None).is_err()); + config.insert(S3_SECRET_ACCESS_KEY.to_string(), "SK".to_string()); + assert!(parse_s3_credential(&config, None).is_err()); + config.insert(S3_SESSION_TOKEN.to_string(), "TOK".to_string()); + config.insert( + S3_SESSION_TOKEN_EXPIRES_AT_MS.to_string(), + "not-a-timestamp".to_string(), + ); + assert!(parse_s3_credential(&config, None).is_err()); + + config.insert( + S3_SESSION_TOKEN_EXPIRES_AT_MS.to_string(), + "1500".to_string(), + ); + let credential = parse_s3_credential(&config, Some("s3://bucket".to_string())).unwrap(); + assert_eq!(credential.prefix(), Some("s3://bucket")); + assert_eq!( + credential.expires_at(), + Some(UNIX_EPOCH + Duration::from_millis(1500)) + ); + match credential.kind() { + StorageCredentialKind::S3(s3) => assert_eq!(s3.session_token(), Some("TOK")), + other => panic!("expected S3, got {other:?}"), + } + } + + #[test] + fn parse_gcs_requires_token_and_expiry() { + let mut config = HashMap::new(); + assert!(parse_gcs_credential(&config, None).is_err()); + config.insert(GCS_TOKEN.to_string(), "ya29.token".to_string()); + assert!(parse_gcs_credential(&config, None).is_err()); + + config.insert(GCS_TOKEN_EXPIRES_AT.to_string(), "2000".to_string()); + let credential = parse_gcs_credential(&config, Some("gs://bucket".to_string())).unwrap(); + assert_eq!(credential.prefix(), Some("gs://bucket")); + match credential.kind() { + StorageCredentialKind::Gcs(gcs) => assert_eq!(gcs.token(), "ya29.token"), + other => panic!("expected GCS, got {other:?}"), + } + assert_eq!( + credential.expires_at(), + Some(UNIX_EPOCH + Duration::from_millis(2000)) + ); + } + + #[test] + fn parse_azdls_requires_account_specific_token_and_expiry() { + let prefix = "abfss://container@account1.dfs.core.windows.net/table"; + let token_key = format!("{ADLS_SAS_TOKEN_PREFIX}account1"); + let expiry_key = format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}account1"); + let mut config = HashMap::new(); + + assert!(parse_azdls_credential(&config, Some(prefix.to_string())).is_err()); + config.insert(token_key, "sv=2026&sig=secret".to_string()); + assert!(parse_azdls_credential(&config, Some(prefix.to_string())).is_err()); + config.insert(expiry_key, "2500".to_string()); + + let credential = parse_azdls_credential(&config, Some(prefix.to_string())).unwrap(); + assert_eq!(credential.prefix(), Some(prefix)); + assert_eq!( + credential.expires_at(), + Some(UNIX_EPOCH + Duration::from_millis(2500)) + ); + match credential.kind() { + StorageCredentialKind::Azdls(azdls) => { + assert_eq!(azdls.sas_token(), "sv=2026&sig=secret") + } + other => panic!("expected ADLS, got {other:?}"), + } + } + + #[test] + fn cached_entry_freshness() { + let no_expiry = cached_s3("", "a", None); + let far = cached_s3( + "", + "a", + Some(SystemTime::now() + REFRESH_BUFFER + Duration::from_secs(60)), + ); + // Within the buffer but not yet expired: stale for a fast-path read, but + // still usable for graceful degradation. + let soon = cached_s3("", "a", Some(SystemTime::now() + Duration::from_secs(60))); + let past = cached_s3("", "a", Some(SystemTime::now() - Duration::from_secs(60))); + + let now = SystemTime::now(); + assert!(no_expiry.is_fresh(now)); + assert!(no_expiry.is_unexpired(now)); + assert!(far.is_fresh(now)); + assert!(!soon.is_fresh(now)); + assert!(soon.is_unexpired(now)); + assert!(!past.is_unexpired(now)); + } + + #[test] + fn prefetch_time_matches_cloud_policy() { + let now = SystemTime::now(); + let expires_at = now + Duration::from_secs(3600); + assert_eq!( + prefetch_time(now, expires_at, false), + expires_at - REFRESH_BUFFER + ); + + for _ in 0..16 { + let refresh_at = prefetch_time(now, expires_at, true); + assert!(refresh_at >= expires_at - REFRESH_BUFFER); + assert!(refresh_at < expires_at - MIN_REFRESH_BUFFER); + } + } + + #[test] + fn prefetch_buffer_is_capped_at_half_the_lifetime() { + let now = SystemTime::now(); + let expires_at = now + Duration::from_secs(120); + assert_eq!( + prefetch_time(now, expires_at, false), + now + Duration::from_secs(60) + ); + + for _ in 0..16 { + let refresh_at = prefetch_time(now, expires_at, true); + assert!(refresh_at >= now + Duration::from_secs(60)); + assert!(refresh_at < expires_at - Duration::from_secs(12)); + } + + // An already expired credential is due immediately. + assert_eq!(prefetch_time(now, now, true), now); + } + + #[test] + fn seed_inside_nominal_window_is_immediately_due() { + let credential = s3_cred( + None, + "a", + Some(SystemTime::now() + REFRESH_BUFFER - Duration::from_secs(1)), + ); + let entry = CachedEntry::seed(credential, true); + let now = SystemTime::now(); + assert!(!entry.is_fresh(now)); + assert!(entry.is_unexpired(now)); + } + + #[test] + fn failure_backoff_is_jittered_and_capped() { + for (failures, ceiling) in [ + (1, Duration::from_secs(1)), + (2, Duration::from_secs(2)), + (5, Duration::from_secs(16)), + (6, MAX_FAILURE_BACKOFF), + (u32::MAX, MAX_FAILURE_BACKOFF), + ] { + for _ in 0..16 { + let backoff = failure_backoff(failures); + assert!(backoff >= ceiling / 2); + assert!(backoff <= ceiling); + } + } + } + + #[test] + fn longest_prefix_match_ignores_freshness() { + let far = Some(SystemTime::now() + REFRESH_BUFFER + Duration::from_secs(3600)); + let entries = vec![ + cached_s3("s3://bucket", "wide", far), + cached_s3("s3://bucket/warehouse/db", "narrow", far), + ]; + let got = longest_prefix_match(&entries, "s3://bucket/warehouse/db/t/f").unwrap(); + assert_eq!(s3_access_key_id(&got.credential), "narrow"); + assert_eq!(got.credential.prefix(), Some("s3://bucket/warehouse/db")); + assert!(longest_prefix_match(&entries, "s3://other/x").is_none()); + + let fresh = Some(SystemTime::now() + REFRESH_BUFFER + Duration::from_secs(60)); + let stale = Some(SystemTime::now() - Duration::from_secs(60)); + let entries = vec![ + cached_s3("s3://bucket", "wide", fresh), + cached_s3("s3://bucket/table", "narrow-stale", stale), + ]; + let selected = longest_prefix_match(&entries, "s3://bucket/table/f").unwrap(); + assert_eq!(s3_access_key_id(&selected.credential), "narrow-stale"); + assert!(!selected.is_fresh(SystemTime::now())); + } + + #[tokio::test] + async fn provider_support_tracks_configured_cloud_endpoints() { + let client = test_client("http://cat"); + let props = aws_refresh_props("/v1/creds"); + + // The provider is configured independently of the table metadata scheme, + // but advertises support only for clouds whose endpoint is present. + let provider = build_test_provider(client.clone(), "http://cat", &props, None) + .await + .unwrap() + .unwrap(); + assert!(provider.supports_path("s3://b/k")); + assert!(!provider.supports_path("abfss://fs@acct.dfs.core.windows.net/k")); + let enabled = HashMap::from([ + ( + CloudRefresh::AWS.endpoint_key.to_string(), + "/v1/creds".to_string(), + ), + ( + CloudRefresh::AWS.enabled_key.to_string(), + "True".to_string(), + ), + ]); + assert!( + build_test_provider(client.clone(), "http://cat", &enabled, None) + .await + .unwrap() + .is_some() + ); + // No endpoint advertised. + assert!( + build_test_provider(client.clone(), "http://cat", &HashMap::new(), None) + .await + .unwrap() + .is_none() + ); + // Explicitly disabled. + let disabled = HashMap::from([ + ( + CloudRefresh::AWS.endpoint_key.to_string(), + "/v1/creds".to_string(), + ), + ( + CloudRefresh::AWS.enabled_key.to_string(), + "false".to_string(), + ), + ]); + assert!( + build_test_provider(client, "http://cat", &disabled, None) + .await + .unwrap() + .is_none() + ); + // Java's `Strings.isNullOrEmpty` check treats an empty endpoint as absent. + let empty = HashMap::from([(CloudRefresh::AWS.endpoint_key.to_string(), String::new())]); + let client = test_client("http://cat"); + assert!( + build_test_provider(client, "http://cat", &empty, None) + .await + .unwrap() + .is_none() + ); + } + + #[tokio::test] + async fn refresh_includes_scan_plan_id() { + let mut server = Server::new_async().await; + let body = s3_response( + "s3://bucket", + "AK", + SystemTime::now() + Duration::from_secs(3600), + ); + let mock = server + .mock("GET", "/v1/credentials") + .match_query(Matcher::UrlEncoded( + "planId".to_string(), + "scan-plan-1".to_string(), + )) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + // No static creds -> no seed -> first load fetches from the endpoint. + let mut props = aws_refresh_props("/v1/credentials"); + props.insert( + REST_CATALOG_PROP_SCAN_PLAN_ID.to_string(), + "scan-plan-1".to_string(), + ); + + let provider = test_provider(&server.url(), &props).await; + let credential = provider + .load_credential("s3://bucket/warehouse/f") + .await + .unwrap(); + assert_eq!(credential.prefix(), Some("s3://bucket")); + + match credential.kind() { + StorageCredentialKind::S3(s3) => { + assert_eq!(s3.access_key_id(), "AK"); + assert_eq!(s3.session_token(), Some("TOK")); + } + other => panic!("expected S3, got {other:?}"), + } + mock.assert_async().await; + } + + #[tokio::test] + async fn successful_refresh_caches_entries_before_selecting_path() { + let mut server = Server::new_async().await; + let body = s3_response( + "s3://bucket/table-a", + "AK", + SystemTime::now() + Duration::from_secs(3600), + ); + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let props = aws_refresh_props("/v1/credentials"); + let provider = test_provider(&server.url(), &props).await; + + assert!( + provider + .load_credential("s3://bucket/table-b/file") + .await + .is_err() + ); + // A successful response with no matching credential is negatively + // cached instead of immediately hitting the endpoint again. + assert!( + provider + .load_credential("s3://bucket/table-b/file") + .await + .is_err() + ); + assert_eq!( + s3_access_key_id( + &provider + .load_credential("s3://bucket/table-a/file") + .await + .unwrap() + ), + "AK" + ); + mock.assert_async().await; + } + + #[tokio::test] + async fn malformed_entry_does_not_discard_valid_entries() { + let mut server = Server::new_async().await; + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let body = format!( + r#"{{"storage-credentials":[ + {{"prefix":"s3://bucket/invalid","config":{{"s3.access-key-id":"BAD"}}}}, + {{"prefix":"s3://bucket/valid","config":{{"s3.access-key-id":"AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expires}"}}}} + ]}}"# + ); + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let props = aws_refresh_props("/v1/credentials"); + let provider = test_provider(&server.url(), &props).await; + + assert!( + provider + .load_credential("s3://bucket/invalid/file") + .await + .is_err() + ); + assert!( + provider + .load_credential("s3://bucket/invalid/file") + .await + .is_err() + ); + assert_eq!( + s3_access_key_id( + &provider + .load_credential("s3://bucket/valid/file") + .await + .unwrap() + ), + "AK" + ); + mock.assert_async().await; + } + + #[tokio::test] + async fn fresh_seed_is_served_without_fetching() { + let mut server = Server::new_async().await; + // Any call to the server is a failure: the fresh seed must be reused. + let mock = server + .mock("GET", Matcher::Any) + .expect(0) + .create_async() + .await; + + let props = with_s3_seed( + aws_refresh_props("/v1/credentials"), + "SEED_AK", + SystemTime::now() + Duration::from_secs(3600), + ); + + let provider = test_provider(&server.url(), &props).await; + let credential = provider.load_credential("s3://bucket/x/f").await.unwrap(); + + assert_eq!(credential.prefix(), None); + assert_eq!(s3_access_key_id(&credential), "SEED_AK"); + mock.assert_async().await; + } + + #[tokio::test] + async fn fresh_azdls_seeds_are_scoped_and_served_without_fetching() { + let mut server = Server::new_async().await; + let mock = server + .mock("GET", Matcher::Any) + .expect(0) + .create_async() + .await; + let props = HashMap::from([ + ( + CloudRefresh::AZURE.endpoint_key.to_string(), + "/v1/credentials".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_PREFIX}account1"), + "sv=2026&sig=seed".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}account1"), + epoch_millis(SystemTime::now() + Duration::from_secs(3600)), + ), + ( + format!("{ADLS_SAS_TOKEN_PREFIX}account2"), + "sv=2026&sig=other-seed".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}account2"), + epoch_millis(SystemTime::now() + Duration::from_secs(3600)), + ), + ]); + let provider = test_provider(&server.url(), &props).await; + + let credential = provider + .load_credential("abfss://container@account1.dfs.core.windows.net/table/data/a.parquet") + .await + .unwrap(); + assert_eq!( + credential.prefix(), + Some("abfss://container@account1.dfs.core.windows.net/") + ); + match credential.kind() { + StorageCredentialKind::Azdls(azdls) => { + assert_eq!(azdls.sas_token(), "sv=2026&sig=seed") + } + other => panic!("expected ADLS credential, got {other:?}"), + } + + let credential = provider + .load_credential("abfss://other-container@account2.dfs.core.windows.net/data/b.parquet") + .await + .unwrap(); + assert_eq!( + credential.prefix(), + Some("abfss://other-container@account2.dfs.core.windows.net/") + ); + match credential.kind() { + StorageCredentialKind::Azdls(azdls) => { + assert_eq!(azdls.sas_token(), "sv=2026&sig=other-seed") + } + other => panic!("expected ADLS credential, got {other:?}"), + } + mock.assert_async().await; + } + + #[tokio::test] + async fn invalid_azdls_path_is_rejected_before_refresh() { + let mut server = Server::new_async().await; + let mock = server + .mock("GET", Matcher::Any) + .expect(0) + .create_async() + .await; + let props = HashMap::from([( + CloudRefresh::AZURE.endpoint_key.to_string(), + "/v1/credentials".to_string(), + )]); + let provider = test_provider(&server.url(), &props).await; + + let error = provider + .load_credential("abfss:///data.parquet") + .await + .unwrap_err(); + + assert_eq!(error.kind(), ErrorKind::DataInvalid); + assert!(error.message().contains("no account name")); + mock.assert_async().await; + } + + #[tokio::test] + async fn failed_refresh_is_backed_off_while_credential_is_unexpired() { + let mut server = Server::new_async().await; + // A refresh is due (the seed is within the buffer) but the catalog errors. + let mock = server + .mock("GET", Matcher::Any) + .expect(1) + .with_status(500) + .create_async() + .await; + + // Seeded credential is within the refresh buffer but not yet expired. + let props = with_s3_seed( + aws_refresh_props("/v1/credentials"), + "SEED_AK", + SystemTime::now() + Duration::from_secs(60), + ); + + let provider = test_provider(&server.url(), &props).await; + // The first refresh fails, but the still-valid seed is served. Immediate + // follow-up operations stay inside the first jittered backoff window. + for _ in 0..2 { + let credential = provider.load_credential("s3://bucket/x/f").await.unwrap(); + assert_eq!(s3_access_key_id(&credential), "SEED_AK"); + } + mock.assert_async().await; + } + + #[tokio::test] + async fn empty_refresh_preserves_unexpired_seed() { + let mut server = Server::new_async().await; + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(r#"{"storage-credentials":[]}"#) + .create_async() + .await; + + let props = with_s3_seed( + aws_refresh_props("/v1/credentials"), + "SEED_AK", + SystemTime::now() + Duration::from_secs(60), + ); + let provider = test_provider(&server.url(), &props).await; + + for _ in 0..2 { + let credential = provider.load_credential("s3://bucket/x/f").await.unwrap(); + assert_eq!(s3_access_key_id(&credential), "SEED_AK"); + } + mock.assert_async().await; + } + + #[tokio::test] + async fn malformed_specific_refresh_prefers_fallback_over_valid_broader_entry() { + let mut server = Server::new_async().await; + let refreshed_expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let body = format!( + r#"{{"storage-credentials":[ + {{"prefix":"s3://bucket/requested","config":{{"s3.access-key-id":"BAD"}}}}, + {{"prefix":"s3://bucket","config":{{"s3.access-key-id":"NEW_AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{refreshed_expires}"}}}} + ]}}"# + ); + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let provider = provider_with_cached_s3(&server.url(), vec![cached_s3( + "s3://bucket/requested", + "SEED_AK", + Some(SystemTime::now() + Duration::from_secs(60)), + )]); + + let fallback = provider + .load_credential("s3://bucket/requested/file") + .await + .unwrap(); + assert_eq!(s3_access_key_id(&fallback), "SEED_AK"); + + // The malformed specific replacement is a failed refresh for this path. + // Its still-valid fallback wins over the broader fetched credential and + // is backed off, so an immediate retry does not fetch again. + let fallback = provider + .load_credential("s3://bucket/requested/other-file") + .await + .unwrap(); + assert_eq!(s3_access_key_id(&fallback), "SEED_AK"); + + let refreshed = provider + .load_credential("s3://bucket/other/file") + .await + .unwrap(); + assert_eq!(s3_access_key_id(&refreshed), "NEW_AK"); + mock.assert_async().await; + } + + #[tokio::test] + async fn invalid_refresh_preserves_other_prefix_fallbacks() { + let mut server = Server::new_async().await; + let refreshed_expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let expired = epoch_millis(SystemTime::now() - Duration::from_secs(60)); + let body = format!( + r#"{{"storage-credentials":[ + {{"prefix":"s3://bucket/table-a","config":{{"s3.access-key-id":"NEW_AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{refreshed_expires}"}}}}, + {{"prefix":"s3://bucket/table-b","config":{{"s3.access-key-id":"BAD"}}}}, + {{"prefix":"s3://bucket/table-c","config":{{"s3.access-key-id":"EXPIRED_CK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expired}"}}}} + ]}}"# + ); + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let provider = provider_with_cached_s3(&server.url(), vec![ + cached_s3( + "s3://bucket/table-a", + "OLD_AK", + Some(SystemTime::now() + Duration::from_secs(60)), + ), + cached_s3( + "s3://bucket/table-b", + "OLDER_BK", + Some(SystemTime::now() + Duration::from_secs(3600)), + ), + cached_s3( + "s3://bucket/table-b", + "LATEST_BK", + Some(SystemTime::now() + Duration::from_secs(3600)), + ), + cached_s3( + "s3://bucket/table-c", + "VALID_CK", + Some(SystemTime::now() + Duration::from_secs(3600)), + ), + ]); + + let refreshed = provider + .load_credential("s3://bucket/table-a/file") + .await + .unwrap(); + assert_eq!(s3_access_key_id(&refreshed), "NEW_AK"); + + let fallback = provider + .load_credential("s3://bucket/table-b/file") + .await + .unwrap(); + assert_eq!(s3_access_key_id(&fallback), "LATEST_BK"); + + let fallback = provider + .load_credential("s3://bucket/table-c/file") + .await + .unwrap(); + assert_eq!(s3_access_key_id(&fallback), "VALID_CK"); + mock.assert_async().await; + } + + #[tokio::test] + async fn concurrent_prefetch_serves_unexpired_credential_without_waiting() { + let mut server = Server::new_async().await; + let body = s3_response( + "s3://bucket", + "NEW_AK", + SystemTime::now() + Duration::from_secs(3600), + ); + let (started_tx, started_rx) = mpsc::channel(); + let release = Arc::new(Barrier::new(2)); + let callback_release = Arc::clone(&release); + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_chunked_body(move |writer| { + started_tx.send(()).unwrap(); + callback_release.wait(); + writer.write_all(body.as_bytes()) + }) + .create_async() + .await; + + let props = with_s3_seed( + aws_refresh_props("/v1/credentials"), + "SEED_AK", + SystemTime::now() + Duration::from_secs(60), + ); + let provider = test_provider(&server.url(), &props).await; + + let first_provider = Arc::clone(&provider); + let first = + tokio::spawn(async move { first_provider.load_credential("s3://bucket/x/f").await }); + tokio::task::spawn_blocking(move || { + started_rx.recv_timeout(Duration::from_secs(5)).unwrap() + }) + .await + .unwrap(); + + let second_provider = Arc::clone(&provider); + let second = + tokio::spawn(async move { second_provider.load_credential("s3://bucket/x/f").await }); + for _ in 0..100 { + if second.is_finished() { + break; + } + tokio::task::yield_now().await; + } + let completed_without_waiting = second.is_finished(); + release.wait(); + + assert!( + completed_without_waiting, + "a concurrent prefetch waited instead of using the unexpired credential" + ); + assert_eq!(s3_access_key_id(&second.await.unwrap().unwrap()), "SEED_AK"); + assert_eq!(s3_access_key_id(&first.await.unwrap().unwrap()), "NEW_AK"); + mock.assert_async().await; + } + + #[tokio::test] + async fn refresh_uses_table_scoped_token() { + let mut server = Server::new_async().await; + let body = s3_response( + "s3://bucket", + "AK", + SystemTime::now() + Duration::from_secs(3600), + ); + let mock = server + .mock("GET", "/v1/credentials") + .match_header("authorization", "Bearer table-token") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let config = RestCatalogConfig::builder() + .uri(server.url()) + .props(HashMap::from([( + "token".to_string(), + "catalog-token".to_string(), + )])) + .build(); + let client = HttpClient::new(&config).unwrap(); + let props = HashMap::from([ + ( + CloudRefresh::AWS.endpoint_key.to_string(), + "/v1/credentials".to_string(), + ), + // Local properties win in the FileIO configuration merge, but must + // not mask auth returned by the table endpoint. + ("token".to_string(), "user-token".to_string()), + ]); + let table_auth = HashMap::from([("token".to_string(), "table-token".to_string())]); + let auth_manager = OAuth2Manager::new(format!("{}/v1/oauth/tokens", server.url())); + let provider = build_vended_credential_provider( + &client, + &auth_manager, + test_factory(&server.url(), Some(&table_auth)), + &props, + true, + ) + .await + .unwrap() + .unwrap(); + + provider.load_credential("s3://bucket/x/f").await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn effective_headers_override_raw_table_headers() { + let mut server = Server::new_async().await; + let body = s3_response( + "s3://bucket", + "AK", + SystemTime::now() + Duration::from_secs(3600), + ); + let mock = server + .mock("GET", "/v1/credentials") + .match_header("authorization", "Bearer local-header") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let client = test_client(&server.url()); + let props = HashMap::from([ + ( + CloudRefresh::AWS.endpoint_key.to_string(), + "/v1/credentials".to_string(), + ), + ( + "header.Authorization".to_string(), + "Bearer local-header".to_string(), + ), + ]); + let table_auth = HashMap::from([ + ("token".to_string(), "table-token".to_string()), + ( + "header.Authorization".to_string(), + "Bearer table-header".to_string(), + ), + ]); + let auth_manager = OAuth2Manager::new(format!("{}/v1/oauth/tokens", server.url())); + let provider = build_vended_credential_provider( + &client, + &auth_manager, + test_factory(&server.url(), Some(&table_auth)), + &props, + true, + ) + .await + .unwrap() + .unwrap(); + + provider.load_credential("s3://bucket/x/f").await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn one_provider_refreshes_multiple_clouds() { + let mut server = Server::new_async().await; + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let aws_body = s3_response( + "s3://bucket", + "AK", + SystemTime::now() + Duration::from_secs(3600), + ); + let gcp_body = format!( + r#"{{"storage-credentials":[{{"prefix":"gs://bucket","config":{{"gcs.oauth2.token":"GCS","gcs.oauth2.token-expires-at":"{expires}"}}}}]}}"# + ); + let aws_mock = server + .mock("GET", "/v1/aws-credentials") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(aws_body) + .create_async() + .await; + let gcp_mock = server + .mock("GET", "/v1/gcp-credentials") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(gcp_body) + .create_async() + .await; + + let props = HashMap::from([ + ( + CloudRefresh::AWS.endpoint_key.to_string(), + "/v1/aws-credentials".to_string(), + ), + ( + CloudRefresh::GCP.endpoint_key.to_string(), + "/v1/gcp-credentials".to_string(), + ), + ]); + let provider = test_provider(&server.url(), &props).await; + + assert!(provider.supports_path("s3://bucket/x")); + assert!(provider.supports_path("gs://bucket/x")); + assert_eq!( + s3_access_key_id(&provider.load_credential("s3://bucket/x").await.unwrap()), + "AK" + ); + match provider + .load_credential("gs://bucket/x") + .await + .unwrap() + .kind() + { + StorageCredentialKind::Gcs(gcs) => assert_eq!(gcs.token(), "GCS"), + other => panic!("expected GCS credential, got {other:?}"), + } + aws_mock.assert_async().await; + gcp_mock.assert_async().await; + } + + #[tokio::test] + async fn host_keyed_azdls_credentials_match_java() { + let mut server = Server::new_async().await; + let host = "account1.dfs.core.windows.net"; + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let prefix = format!("abfss://container@{host}/table"); + let body = format!( + r#"{{"storage-credentials":[{{"prefix":"{prefix}","config":{{"adls.sas-token.{host}":"sv=2026&sig=refreshed","adls.sas-token-expires-at-ms.{host}":"{expires}"}}}}]}}"# + ); + let mock = server + .mock("GET", "/v1/azure-credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + let props = HashMap::from([ + ( + CloudRefresh::AZURE.endpoint_key.to_string(), + "/v1/azure-credentials".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_PREFIX}{host}"), + "sv=2026&sig=seed".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}{host}"), + epoch_millis(SystemTime::now() + Duration::from_secs(3600)), + ), + ]); + let provider = test_provider(&server.url(), &props).await; + let sas_token = |credential: StorageCredential| match credential.into_kind() { + StorageCredentialKind::Azdls(azdls) => azdls.into_sas_token(), + other => panic!("expected ADLS credential, got {other:?}"), + }; + + // The host-keyed seed serves the account without fetching. + let seeded = provider + .load_credential(&format!("abfss://other@{host}/data.parquet")) + .await + .unwrap(); + assert_eq!(sas_token(seeded), "sv=2026&sig=seed"); + + // An account without a seed refreshes. The response is parsed with + // host-keyed tokens, and its entry then wins over the broader seed. + let uncovered = provider + .load_credential("abfss://container@account2.dfs.core.windows.net/table/a") + .await; + assert!(uncovered.is_err()); + let refreshed = provider + .load_credential(&format!("{prefix}/data.parquet")) + .await + .unwrap(); + assert_eq!(sas_token(refreshed), "sv=2026&sig=refreshed"); + mock.assert_async().await; + } + + #[tokio::test] + async fn refreshes_azdls_credentials() { + let mut server = Server::new_async().await; + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let prefix = "abfss://container@account1.dfs.core.windows.net/table"; + let body = format!( + r#"{{"storage-credentials":[{{"prefix":"{prefix}","config":{{"adls.sas-token.account1":"sv=2026&sig=secret","adls.sas-token-expires-at-ms.account1":"{expires}"}}}}]}}"# + ); + let mock = server + .mock("GET", "/v1/azure-credentials") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + let props = HashMap::from([( + CloudRefresh::AZURE.endpoint_key.to_string(), + "/v1/azure-credentials".to_string(), + )]); + let provider = test_provider(&server.url(), &props).await; + + let credential = provider + .load_credential(&format!("{prefix}/data.parquet")) + .await + .unwrap(); + match credential.kind() { + StorageCredentialKind::Azdls(azdls) => { + assert_eq!(azdls.sas_token(), "sv=2026&sig=secret") + } + other => panic!("expected ADLS credential, got {other:?}"), + } + mock.assert_async().await; + } + + #[tokio::test] + async fn refresh_failure_never_serves_expired_credential() { + let mut server = Server::new_async().await; + let mock = server + .mock("GET", Matcher::Any) + .expect(1) + .with_status(500) + .create_async() + .await; + + let props = with_s3_seed( + aws_refresh_props("/v1/credentials"), + "EXPIRED_AK", + SystemTime::now() - Duration::from_secs(60), + ); + + let provider = test_provider(&server.url(), &props).await; + + assert!(provider.load_credential("s3://bucket/x/f").await.is_err()); + assert!(provider.load_credential("s3://bucket/x/f").await.is_err()); + mock.assert_async().await; + } + + #[tokio::test] + async fn failed_azdls_refresh_uses_account_seed_until_expiry() { + let mut server = Server::new_async().await; + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(503) + .with_header("content-type", "application/json") + .with_body(r#"{"error":{"message":"unavailable","type":"test","code":503}}"#) + .create_async() + .await; + let props = HashMap::from([ + ( + CloudRefresh::AZURE.endpoint_key.to_string(), + "/v1/credentials".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_PREFIX}account2"), + "sv=2026&sig=seed".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}account2"), + epoch_millis(SystemTime::now() + Duration::from_secs(60)), + ), + ]); + let provider = test_provider(&server.url(), &props).await; + let path = "abfss://container@account2.dfs.core.windows.net/data/a.parquet"; + + for _ in 0..2 { + let credential = provider.load_credential(path).await.unwrap(); + assert_eq!( + credential.prefix(), + Some("abfss://container@account2.dfs.core.windows.net/") + ); + match credential.kind() { + StorageCredentialKind::Azdls(azdls) => { + assert_eq!(azdls.sas_token(), "sv=2026&sig=seed") + } + other => panic!("expected ADLS credential, got {other:?}"), + } + } + mock.assert_async().await; + } + + #[tokio::test] + async fn sub_buffer_ttl_is_reused_until_half_its_lifetime() { + let mut server = Server::new_async().await; + // The vended TTL (60s) is shorter than REFRESH_BUFFER. Java would treat + // it as due on arrival and fetch on every operation; the capped buffer + // keeps it for half its lifetime. + let body = s3_response( + "s3://bucket", + "AK", + SystemTime::now() + Duration::from_secs(60), + ); + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let provider = test_provider(&server.url(), &aws_refresh_props("/v1/credentials")).await; + + for _ in 0..2 { + let credential = provider.load_credential("s3://bucket/x/f").await.unwrap(); + assert_eq!(s3_access_key_id(&credential), "AK"); + } + mock.assert_async().await; + } + + #[tokio::test] + async fn refreshed_credentials_must_be_complete() { + let mut server = Server::new_async().await; + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body( + r#"{"storage-credentials":[{"prefix":"s3://bucket","config":{"s3.access-key-id":"AK","s3.secret-access-key":"SK"}}]}"#, + ) + .create_async() + .await; + + let provider = test_provider(&server.url(), &aws_refresh_props("/v1/credentials")).await; + let error = provider + .load_credential("s3://bucket/x/f") + .await + .unwrap_err(); + assert!(error.message().contains("is missing or empty"), "{error}"); + mock.assert_async().await; + } + + #[tokio::test] + async fn scheme_aliases_share_prefix_scoped_credentials() { + let mut server = Server::new_async().await; + let body = s3_response( + "s3://bucket/table", + "AK", + SystemTime::now() + Duration::from_secs(3600), + ); + let mock = server + .mock("GET", "/v1/credentials") + .expect(2) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let provider = test_provider(&server.url(), &aws_refresh_props("/v1/credentials")).await; + for path in ["s3a://bucket/table/f", "s3://bucket/table/f"] { + assert_eq!( + s3_access_key_id(&provider.load_credential(path).await.unwrap()), + "AK" + ); + } + // A prefix only covers whole path segments, so this path refreshes. + assert!( + provider + .load_credential("s3://bucket/table2/f") + .await + .is_err() + ); + mock.assert_async().await; + } + + #[tokio::test] + async fn serialized_provider_reconnects_from_properties() { + let mut server = Server::new_async().await; + let body = s3_response( + "s3://bucket", + "AK", + SystemTime::now() + Duration::from_secs(3600), + ); + let mock = server + .mock("GET", "/v1/credentials") + .match_header("authorization", "Bearer table-token") + .match_header("x-custom", "value") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let mut props = aws_refresh_props("/v1/credentials"); + props.extend([ + ("token".to_string(), "catalog-token".to_string()), + ("header.x-custom".to_string(), "value".to_string()), + ]); + let table_auth = HashMap::from([("token".to_string(), "table-token".to_string())]); + let provider = build_vended_credential_provider( + &test_client(&server.url()), + &NoopAuthManager, + test_factory(&server.url(), Some(&table_auth)), + &props, + true, + ) + .await + .unwrap() + .unwrap(); + + let factory: Arc = + serde_json::from_str(&serde_json::to_string(&provider.factory().unwrap()).unwrap()) + .unwrap(); + let config = StorageConfig::new().with_props(props); + let rebuilt = factory.build(&config).unwrap(); + + assert_eq!( + s3_access_key_id(&rebuilt.load_credential("s3://bucket/x/f").await.unwrap()), + "AK" + ); + // The rebuilt provider is itself serializable. + assert!(rebuilt.factory().is_ok()); + mock.assert_async().await; + } + + #[tokio::test] + async fn provider_with_injected_auth_manager_is_not_serializable() { + let provider = build_vended_credential_provider( + &test_client("http://cat"), + &NoopAuthManager, + test_factory("http://cat", None), + &aws_refresh_props("/v1/creds"), + false, + ) + .await + .unwrap() + .unwrap(); + + let error = provider.factory().unwrap_err(); + assert_eq!(error.kind(), ErrorKind::FeatureUnsupported); + assert!( + error.message().contains("without_credential_provider"), + "{error}" + ); + } + + #[test] + fn factory_debug_omits_table_auth() { + let factory = test_factory( + "http://cat", + Some(&HashMap::from([( + "token".to_string(), + "secret-token".to_string(), + )])), + ); + let debug = format!("{factory:?}"); + assert!(!debug.contains("secret-token"), "{debug}"); + } +} diff --git a/crates/catalog/rest/src/lib.rs b/crates/catalog/rest/src/lib.rs index 29a171de61..8bee2d2ef0 100644 --- a/crates/catalog/rest/src/lib.rs +++ b/crates/catalog/rest/src/lib.rs @@ -88,6 +88,7 @@ mod auth; mod catalog; mod client; +mod credential; pub use client::HttpClient; mod request; pub use request::{HttpRequest, HttpRequestBody}; diff --git a/crates/catalog/rest/src/types.rs b/crates/catalog/rest/src/types.rs index 390521229e..8265ceb40b 100644 --- a/crates/catalog/rest/src/types.rs +++ b/crates/catalog/rest/src/types.rs @@ -101,7 +101,7 @@ impl From for Error { } } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Serialize, Deserialize)] pub(super) struct TokenResponse { pub(super) access_token: String, pub(super) token_type: String, @@ -205,7 +205,7 @@ pub struct RenameTableRequest { pub destination: TableIdent, } -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +#[derive(Clone, Serialize, Deserialize, PartialEq, Eq)] #[serde(rename_all = "kebab-case")] /// Result returned when a table is successfully loaded or created. /// @@ -231,7 +231,18 @@ pub struct LoadTableResult { pub storage_credentials: Option>, } -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +impl std::fmt::Debug for LoadTableResult { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("LoadTableResult") + .field("metadata_location", &self.metadata_location) + .field("metadata", &self.metadata) + .field("config_keys", &self.config.keys().collect::>()) + .field("storage_credentials", &self.storage_credentials) + .finish_non_exhaustive() + } +} + +#[derive(Clone, Serialize, Deserialize, PartialEq, Eq)] /// Storage credential for a specific location prefix. /// /// Indicates a storage location prefix where the credential is relevant. Clients should @@ -244,6 +255,27 @@ pub struct StorageCredential { pub config: HashMap, } +impl std::fmt::Debug for StorageCredential { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("StorageCredential") + .field("prefix", &self.prefix) + .field("config_keys", &self.config.keys().collect::>()) + .finish_non_exhaustive() + } +} + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +#[serde(rename_all = "kebab-case")] +/// Response from the table credentials endpoint +/// (`GET /v1/{prefix}/namespaces/{namespace}/tables/{table}/credentials`). +/// +/// Returns freshly vended storage credentials so clients can refresh temporary +/// credentials before they expire. +pub struct LoadCredentialsResponse { + /// Storage credentials, one entry per location prefix. + pub storage_credentials: Vec, +} + #[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] #[serde(rename_all = "kebab-case")] /// Request to create a new table in a namespace. diff --git a/crates/iceberg/public-api.txt b/crates/iceberg/public-api.txt index 58c7db7dd6..ae3c2c4741 100644 --- a/crates/iceberg/public-api.txt +++ b/crates/iceberg/public-api.txt @@ -695,6 +695,14 @@ pub fn iceberg::inspect::SnapshotsTable<'a>::new(table: &'a iceberg::table::Tabl pub async fn iceberg::inspect::SnapshotsTable<'a>::scan(&self) -> iceberg::Result pub fn iceberg::inspect::SnapshotsTable<'a>::schema(&self) -> iceberg::spec::Schema pub mod iceberg::io +#[non_exhaustive] pub enum iceberg::io::StorageCredentialKind +pub iceberg::io::StorageCredentialKind::Azdls(iceberg::io::AzdlsCredential) +pub iceberg::io::StorageCredentialKind::Gcs(iceberg::io::GcsCredential) +pub iceberg::io::StorageCredentialKind::S3(iceberg::io::S3Credential) +impl core::clone::Clone for iceberg::io::StorageCredentialKind +pub fn iceberg::io::StorageCredentialKind::clone(&self) -> iceberg::io::StorageCredentialKind +impl core::fmt::Debug for iceberg::io::StorageCredentialKind +pub fn iceberg::io::StorageCredentialKind::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result pub struct iceberg::io::AzdlsConfig pub iceberg::io::AzdlsConfig::account_key: core::option::Option pub iceberg::io::AzdlsConfig::account_name: core::option::Option @@ -725,6 +733,15 @@ impl serde_core::ser::Serialize for iceberg::io::AzdlsConfig pub fn iceberg::io::AzdlsConfig::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::AzdlsConfig pub fn iceberg::io::AzdlsConfig::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> +pub struct iceberg::io::AzdlsCredential +impl iceberg::io::AzdlsCredential +pub fn iceberg::io::AzdlsCredential::into_sas_token(self) -> alloc::string::String +pub fn iceberg::io::AzdlsCredential::new(sas_token: impl core::convert::Into) -> Self +pub fn iceberg::io::AzdlsCredential::sas_token(&self) -> &str +impl core::clone::Clone for iceberg::io::AzdlsCredential +pub fn iceberg::io::AzdlsCredential::clone(&self) -> iceberg::io::AzdlsCredential +impl core::fmt::Debug for iceberg::io::AzdlsCredential +pub fn iceberg::io::AzdlsCredential::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result pub struct iceberg::io::FileIO impl iceberg::io::FileIO pub fn iceberg::io::FileIO::config(&self) -> &iceberg::io::StorageConfig @@ -738,6 +755,7 @@ pub fn iceberg::io::FileIO::new_output(&self, path: impl core::convert::AsRef Self pub fn iceberg::io::FileIO::new_with_memory() -> Self pub fn iceberg::io::FileIO::serialize_all(&self) -> iceberg::Result> +pub fn iceberg::io::FileIO::without_credential_provider(&self) -> Self impl core::clone::Clone for iceberg::io::FileIO pub fn iceberg::io::FileIO::clone(&self) -> iceberg::io::FileIO impl core::fmt::Debug for iceberg::io::FileIO @@ -747,6 +765,7 @@ impl iceberg::io::FileIOBuilder pub fn iceberg::io::FileIOBuilder::build(self) -> iceberg::io::FileIO pub fn iceberg::io::FileIOBuilder::config(&self) -> &iceberg::io::StorageConfig pub fn iceberg::io::FileIOBuilder::new(factory: alloc::sync::Arc) -> Self +pub fn iceberg::io::FileIOBuilder::with_credential_provider(self, provider: alloc::sync::Arc) -> Self pub fn iceberg::io::FileIOBuilder::with_prop(self, key: impl alloc::string::ToString, value: impl alloc::string::ToString) -> Self pub fn iceberg::io::FileIOBuilder::with_props(self, args: impl core::iter::traits::collect::IntoIterator) -> Self impl core::clone::Clone for iceberg::io::FileIOBuilder @@ -783,6 +802,15 @@ impl serde_core::ser::Serialize for iceberg::io::GcsConfig pub fn iceberg::io::GcsConfig::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::GcsConfig pub fn iceberg::io::GcsConfig::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> +pub struct iceberg::io::GcsCredential +impl iceberg::io::GcsCredential +pub fn iceberg::io::GcsCredential::into_token(self) -> alloc::string::String +pub fn iceberg::io::GcsCredential::new(token: impl core::convert::Into) -> Self +pub fn iceberg::io::GcsCredential::token(&self) -> &str +impl core::clone::Clone for iceberg::io::GcsCredential +pub fn iceberg::io::GcsCredential::clone(&self) -> iceberg::io::GcsCredential +impl core::fmt::Debug for iceberg::io::GcsCredential +pub fn iceberg::io::GcsCredential::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result pub struct iceberg::io::HfConfig pub iceberg::io::HfConfig::endpoint: core::option::Option pub iceberg::io::HfConfig::revision: core::option::Option @@ -850,6 +878,7 @@ impl core::fmt::Debug for iceberg::io::LocalFsStorageFactory pub fn iceberg::io::LocalFsStorageFactory::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result impl iceberg::io::StorageFactory for iceberg::io::LocalFsStorageFactory pub fn iceberg::io::LocalFsStorageFactory::build(&self, _config: &iceberg::io::StorageConfig) -> iceberg::Result> +pub fn iceberg::io::LocalFsStorageFactory::build_with_credentials(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> impl serde_core::ser::Serialize for iceberg::io::LocalFsStorageFactory pub fn iceberg::io::LocalFsStorageFactory::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::LocalFsStorageFactory @@ -888,6 +917,7 @@ impl core::fmt::Debug for iceberg::io::MemoryStorageFactory pub fn iceberg::io::MemoryStorageFactory::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result impl iceberg::io::StorageFactory for iceberg::io::MemoryStorageFactory pub fn iceberg::io::MemoryStorageFactory::build(&self, _config: &iceberg::io::StorageConfig) -> iceberg::Result> +pub fn iceberg::io::MemoryStorageFactory::build_with_credentials(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> impl serde_core::ser::Serialize for iceberg::io::MemoryStorageFactory pub fn iceberg::io::MemoryStorageFactory::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::MemoryStorageFactory @@ -963,6 +993,17 @@ impl serde_core::ser::Serialize for iceberg::io::S3Config pub fn iceberg::io::S3Config::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::S3Config pub fn iceberg::io::S3Config::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> +pub struct iceberg::io::S3Credential +impl iceberg::io::S3Credential +pub fn iceberg::io::S3Credential::access_key_id(&self) -> &str +pub fn iceberg::io::S3Credential::into_parts(self) -> (alloc::string::String, alloc::string::String, core::option::Option) +pub fn iceberg::io::S3Credential::new(access_key_id: impl core::convert::Into, secret_access_key: impl core::convert::Into, session_token: core::option::Option) -> Self +pub fn iceberg::io::S3Credential::secret_access_key(&self) -> &str +pub fn iceberg::io::S3Credential::session_token(&self) -> core::option::Option<&str> +impl core::clone::Clone for iceberg::io::S3Credential +pub fn iceberg::io::S3Credential::clone(&self) -> iceberg::io::S3Credential +impl core::fmt::Debug for iceberg::io::S3Credential +pub fn iceberg::io::S3Credential::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result pub struct iceberg::io::StorageConfig impl iceberg::io::StorageConfig pub fn iceberg::io::StorageConfig::from_props(props: std::collections::hash::map::HashMap) -> Self @@ -1000,14 +1041,34 @@ impl serde_core::ser::Serialize for iceberg::io::StorageConfig pub fn iceberg::io::StorageConfig::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::StorageConfig pub fn iceberg::io::StorageConfig::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> +pub struct iceberg::io::StorageCredential +impl iceberg::io::StorageCredential +pub fn iceberg::io::StorageCredential::covers(&self, location: &str) -> bool +pub fn iceberg::io::StorageCredential::expires_at(&self) -> core::option::Option +pub fn iceberg::io::StorageCredential::into_kind(self) -> iceberg::io::StorageCredentialKind +pub fn iceberg::io::StorageCredential::kind(&self) -> &iceberg::io::StorageCredentialKind +pub fn iceberg::io::StorageCredential::new(kind: iceberg::io::StorageCredentialKind) -> Self +pub fn iceberg::io::StorageCredential::prefix(&self) -> core::option::Option<&str> +pub fn iceberg::io::StorageCredential::with_expiration(self, expires_at: std::time::SystemTime) -> Self +pub fn iceberg::io::StorageCredential::with_prefix(self, prefix: impl core::convert::Into) -> Self +impl core::clone::Clone for iceberg::io::StorageCredential +pub fn iceberg::io::StorageCredential::clone(&self) -> iceberg::io::StorageCredential +impl core::fmt::Debug for iceberg::io::StorageCredential +pub fn iceberg::io::StorageCredential::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result pub const iceberg::io::ADLS_ACCOUNT_KEY: &str pub const iceberg::io::ADLS_ACCOUNT_NAME: &str pub const iceberg::io::ADLS_AUTHORITY_HOST: &str pub const iceberg::io::ADLS_CLIENT_ID: &str pub const iceberg::io::ADLS_CLIENT_SECRET: &str pub const iceberg::io::ADLS_CONNECTION_STRING: &str +pub const iceberg::io::ADLS_REFRESH_CREDENTIALS_ENABLED: &str +pub const iceberg::io::ADLS_REFRESH_CREDENTIALS_ENDPOINT: &str pub const iceberg::io::ADLS_SAS_TOKEN: &str +pub const iceberg::io::ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX: &str +pub const iceberg::io::ADLS_SAS_TOKEN_PREFIX: &str pub const iceberg::io::ADLS_TENANT_ID: &str +pub const iceberg::io::AWS_REFRESH_CREDENTIALS_ENABLED: &str +pub const iceberg::io::AWS_REFRESH_CREDENTIALS_ENDPOINT: &str pub const iceberg::io::CLIENT_REGION: &str pub const iceberg::io::GCS_ALLOW_ANONYMOUS: &str pub const iceberg::io::GCS_CREDENTIALS_JSON: &str @@ -1015,8 +1076,11 @@ pub const iceberg::io::GCS_DISABLE_CONFIG_LOAD: &str pub const iceberg::io::GCS_DISABLE_VM_METADATA: &str pub const iceberg::io::GCS_NO_AUTH: &str pub const iceberg::io::GCS_PROJECT_ID: &str +pub const iceberg::io::GCS_REFRESH_CREDENTIALS_ENABLED: &str +pub const iceberg::io::GCS_REFRESH_CREDENTIALS_ENDPOINT: &str pub const iceberg::io::GCS_SERVICE_HOST: &str pub const iceberg::io::GCS_TOKEN: &str +pub const iceberg::io::GCS_TOKEN_EXPIRES_AT: &str pub const iceberg::io::GCS_USER_PROJECT: &str pub const iceberg::io::HF_ENDPOINT: &str pub const iceberg::io::HF_REVISION: &str @@ -1036,6 +1100,7 @@ pub const iceberg::io::S3_PATH_STYLE_ACCESS: &str pub const iceberg::io::S3_REGION: &str pub const iceberg::io::S3_SECRET_ACCESS_KEY: &str pub const iceberg::io::S3_SESSION_TOKEN: &str +pub const iceberg::io::S3_SESSION_TOKEN_EXPIRES_AT_MS: &str pub const iceberg::io::S3_SSE_KEY: &str pub const iceberg::io::S3_SSE_MD5: &str pub const iceberg::io::S3_SSE_TYPE: &str @@ -1087,12 +1152,22 @@ pub fn iceberg::io::MemoryStorage::read<'life0, 'life1, 'async_trait>(&'life0 se pub fn iceberg::io::MemoryStorage::reader<'life0, 'life1, 'async_trait>(&'life0 self, path: &'life1 str) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait pub fn iceberg::io::MemoryStorage::write<'life0, 'life1, 'async_trait>(&'life0 self, path: &'life1 str, bs: bytes::bytes::Bytes) -> core::pin::Pin> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait pub fn iceberg::io::MemoryStorage::writer<'life0, 'life1, 'async_trait>(&'life0 self, path: &'life1 str) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait +pub trait iceberg::io::StorageCredentialProvider: core::fmt::Debug + core::marker::Send + core::marker::Sync +pub fn iceberg::io::StorageCredentialProvider::factory(&self) -> iceberg::Result> +pub fn iceberg::io::StorageCredentialProvider::load_credential<'life0, 'life1, 'async_trait>(&'life0 self, path: &'life1 str) -> core::pin::Pin> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait +pub fn iceberg::io::StorageCredentialProvider::supports_path(&self, _path: &str) -> bool +pub trait iceberg::io::StorageCredentialProviderFactory: core::fmt::Debug + core::marker::Send + core::marker::Sync + typetag::Serialize + typetag::Deserialize +pub fn iceberg::io::StorageCredentialProviderFactory::build(&self, config: &iceberg::io::StorageConfig) -> iceberg::Result> pub trait iceberg::io::StorageFactory: core::fmt::Debug + core::marker::Send + core::marker::Sync + typetag::Serialize + typetag::Deserialize pub fn iceberg::io::StorageFactory::build(&self, config: &iceberg::io::StorageConfig) -> iceberg::Result> +pub fn iceberg::io::StorageFactory::build_with_credentials(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> impl iceberg::io::StorageFactory for iceberg::io::LocalFsStorageFactory pub fn iceberg::io::LocalFsStorageFactory::build(&self, _config: &iceberg::io::StorageConfig) -> iceberg::Result> +pub fn iceberg::io::LocalFsStorageFactory::build_with_credentials(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> impl iceberg::io::StorageFactory for iceberg::io::MemoryStorageFactory pub fn iceberg::io::MemoryStorageFactory::build(&self, _config: &iceberg::io::StorageConfig) -> iceberg::Result> +pub fn iceberg::io::MemoryStorageFactory::build_with_credentials(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> +pub fn iceberg::io::storage_prefix_covers(prefix: &str, location: &str) -> bool pub mod iceberg::memory pub struct iceberg::memory::MemoryCatalog impl core::fmt::Debug for iceberg::memory::MemoryCatalog diff --git a/crates/iceberg/src/io/file_io.rs b/crates/iceberg/src/io/file_io.rs index 42eec4db56..300ea71f07 100644 --- a/crates/iceberg/src/io/file_io.rs +++ b/crates/iceberg/src/io/file_io.rs @@ -22,7 +22,8 @@ use bytes::Bytes; use futures::{Stream, StreamExt}; use super::storage::{ - LocalFsStorageFactory, MemoryStorageFactory, Storage, StorageConfig, StorageFactory, + LocalFsStorageFactory, MemoryStorageFactory, Storage, StorageConfig, StorageCredentialProvider, + StorageCredentialProviderFactory, StorageFactory, }; use crate::Result; @@ -65,6 +66,8 @@ pub struct FileIO { config: StorageConfig, /// Factory for creating storage instances factory: Arc, + /// Optional provider of refreshable, backend-specific credentials + credential_provider: Option>, /// Cached storage instance (lazily initialized) storage: Arc>>, } @@ -74,18 +77,22 @@ mod _serde { use serde::{Deserialize, Serialize}; - use super::{StorageConfig, StorageFactory}; + use super::{StorageConfig, StorageCredentialProviderFactory, StorageFactory}; #[derive(Serialize)] pub(super) struct SerializableFileIO<'a> { pub(super) config: &'a StorageConfig, pub(super) factory: &'a Arc, + #[serde(skip_serializing_if = "Option::is_none")] + pub(super) credential_provider: Option>, } #[derive(Deserialize)] pub(super) struct DeserializedFileIO { pub(super) config: StorageConfig, pub(super) factory: Arc, + #[serde(default)] + pub(super) credential_provider: Option>, } } @@ -97,6 +104,7 @@ impl FileIO { Self { config: StorageConfig::new(), factory: Arc::new(MemoryStorageFactory), + credential_provider: None, storage: Arc::new(OnceLock::new()), } } @@ -108,6 +116,7 @@ impl FileIO { Self { config: StorageConfig::new(), factory: Arc::new(LocalFsStorageFactory), + credential_provider: None, storage: Arc::new(OnceLock::new()), } } @@ -123,14 +132,28 @@ impl FileIO { /// /// All storage configuration properties are included in the serialized representation. These /// properties may contain credentials or other sensitive values, so the returned bytes must be - /// protected in transit and at rest by the application embedding this crate. + /// protected in transit and at rest by the application embedding this crate. A serialized + /// credential provider may likewise carry catalog authentication and vended credentials. /// /// Storage factories are serialized through [`typetag`](https://docs.rs/typetag). Third-party /// factories must use `#[typetag::serde]` on their [`StorageFactory`] implementation. + /// + /// A credential provider is serialized as the + /// [`StorageCredentialProviderFactory`] returned by + /// [`StorageCredentialProvider::factory`], and rebuilt on deserialization. Serialization fails + /// when the provider cannot be rebuilt in another process; use + /// [`FileIO::without_credential_provider`] to serialize without it. pub fn serialize_all(&self) -> Result> { + let credential_provider = self + .credential_provider + .as_ref() + .map(|provider| provider.factory()) + .transpose()?; + Ok(serde_json::to_vec(&_serde::SerializableFileIO { config: &self.config, factory: &self.factory, + credential_provider, })?) } @@ -140,14 +163,35 @@ impl FileIO { /// implementation so it is registered with `typetag`. Backend-specific requirements are /// documented by each storage factory implementation. pub fn deserialize_all(bytes: &[u8]) -> Result { - let _serde::DeserializedFileIO { config, factory } = serde_json::from_slice(bytes)?; + let _serde::DeserializedFileIO { + config, + factory, + credential_provider, + } = serde_json::from_slice(bytes)?; + let credential_provider = credential_provider + .map(|provider_factory| provider_factory.build(&config)) + .transpose()?; Ok(Self { config, factory, + credential_provider, storage: Arc::new(OnceLock::new()), }) } + /// Returns a copy of this `FileIO` without its credential provider. + /// + /// The copy uses only the credentials in its storage configuration, which are not refreshed. + /// Use this to serialize a `FileIO` whose credential provider cannot be serialized. + pub fn without_credential_provider(&self) -> Self { + Self { + config: self.config.clone(), + factory: Arc::clone(&self.factory), + credential_provider: None, + storage: Arc::new(OnceLock::new()), + } + } + /// Get the storage configuration. pub fn config(&self) -> &StorageConfig { &self.config @@ -163,8 +207,11 @@ impl FileIO { return Ok(storage.clone()); } - // Build the storage - let storage = self.factory.build(&self.config)?; + // Build the storage, passing any credential provider so backends that + // support refreshable credentials can wire it into their operators. + let storage = self + .factory + .build_with_credentials(&self.config, self.credential_provider.clone())?; // Try to set it (another thread might have set it first) let _ = self.storage.set(storage.clone()); @@ -247,6 +294,8 @@ pub struct FileIOBuilder { factory: Arc, /// Storage configuration config: StorageConfig, + /// Optional provider of refreshable, backend-specific credentials + credential_provider: Option>, } impl FileIOBuilder { @@ -255,6 +304,7 @@ impl FileIOBuilder { Self { factory, config: StorageConfig::new(), + credential_provider: None, } } @@ -280,11 +330,24 @@ impl FileIOBuilder { &self.config } + /// Attach a provider of refreshable, backend-specific credentials. + /// + /// Storage factories that cannot use the provider ignore it and use the + /// credentials in the configuration. + pub fn with_credential_provider( + mut self, + provider: Arc, + ) -> Self { + self.credential_provider = Some(provider); + self + } + /// Builds [`FileIO`]. pub fn build(self) -> FileIO { FileIO { config: self.config, factory: self.factory, + credential_provider: self.credential_provider, storage: Arc::new(OnceLock::new()), } } @@ -453,7 +516,53 @@ mod tests { use tempfile::TempDir; use super::{FileIO, FileIOBuilder}; - use crate::io::{LocalFsStorageFactory, MemoryStorageFactory}; + use crate::io::{ + GcsCredential, LocalFsStorageFactory, MemoryStorageFactory, StorageConfig, + StorageCredential, StorageCredentialKind, StorageCredentialProvider, + StorageCredentialProviderFactory, + }; + use crate::{ErrorKind, Result}; + + #[derive(Debug)] + struct TestCredentialProvider; + + #[async_trait::async_trait] + impl StorageCredentialProvider for TestCredentialProvider { + async fn load_credential(&self, _path: &str) -> Result { + unreachable!("unsupported factories must ignore the provider") + } + } + + /// A provider rebuilt from the `FileIO` configuration, like a catalog provider. + #[derive(Debug)] + struct PortableCredentialProvider { + endpoint: Option, + } + + #[async_trait::async_trait] + impl StorageCredentialProvider for PortableCredentialProvider { + async fn load_credential(&self, _path: &str) -> Result { + Ok(StorageCredential::new(StorageCredentialKind::Gcs( + GcsCredential::new(self.endpoint.clone().unwrap_or_default()), + ))) + } + + fn factory(&self) -> Result> { + Ok(Arc::new(PortableCredentialProviderFactory)) + } + } + + #[derive(Debug, serde::Serialize, serde::Deserialize)] + struct PortableCredentialProviderFactory; + + #[typetag::serde] + impl StorageCredentialProviderFactory for PortableCredentialProviderFactory { + fn build(&self, config: &StorageConfig) -> Result> { + Ok(Arc::new(PortableCredentialProvider { + endpoint: config.get("endpoint").cloned(), + })) + } + } fn create_local_file_io() -> FileIO { FileIO::new_with_fs() @@ -601,6 +710,63 @@ mod tests { assert_eq!(file_io.config().get("key2"), Some(&"value2".to_string())); } + #[tokio::test] + async fn test_file_io_ignores_credentials_for_unsupported_factory() { + let file_io = FileIOBuilder::new(Arc::new(MemoryStorageFactory)) + .with_credential_provider(Arc::new(TestCredentialProvider)) + .build(); + + file_io + .new_output("memory://file") + .unwrap() + .write("data".into()) + .await + .unwrap(); + assert!(file_io.exists("memory://file").await.unwrap()); + } + + #[test] + fn test_file_io_with_credential_provider_serialization_fails() { + let file_io = FileIOBuilder::new(Arc::new(MemoryStorageFactory)) + .with_credential_provider(Arc::new(TestCredentialProvider)) + .build(); + + let err = file_io.serialize_all().unwrap_err(); + assert_eq!(err.kind(), ErrorKind::FeatureUnsupported, "{err}"); + + let deserialized = FileIO::deserialize_all( + &file_io + .without_credential_provider() + .serialize_all() + .unwrap(), + ) + .unwrap(); + assert!(deserialized.credential_provider.is_none()); + } + + #[tokio::test] + async fn test_file_io_rebuilds_credential_provider_after_serialization() { + let file_io = FileIOBuilder::new(Arc::new(MemoryStorageFactory)) + .with_prop("endpoint", "https://catalog/credentials") + .with_credential_provider(Arc::new(PortableCredentialProvider { endpoint: None })) + .build(); + + let deserialized = FileIO::deserialize_all(&file_io.serialize_all().unwrap()).unwrap(); + let credential = deserialized + .credential_provider + .unwrap() + .load_credential("gs://bucket/file") + .await + .unwrap(); + // The provider is rebuilt from the deserialized configuration. + match credential.kind() { + StorageCredentialKind::Gcs(gcs) => { + assert_eq!(gcs.token(), "https://catalog/credentials") + } + other => panic!("expected GCS credential, got {other:?}"), + } + } + #[tokio::test] async fn test_memory_file_io_serialization_roundtrip() { let file_io = FileIOBuilder::new(Arc::new(MemoryStorageFactory)) diff --git a/crates/iceberg/src/io/storage/config/azdls.rs b/crates/iceberg/src/io/storage/config/azdls.rs index a9541e7791..7d49a354dc 100644 --- a/crates/iceberg/src/io/storage/config/azdls.rs +++ b/crates/iceberg/src/io/storage/config/azdls.rs @@ -36,6 +36,14 @@ pub const ADLS_ACCOUNT_NAME: &str = "adls.account-name"; pub const ADLS_ACCOUNT_KEY: &str = "adls.account-key"; /// The shared access signature. pub const ADLS_SAS_TOKEN: &str = "adls.sas-token"; +/// Prefix for account-specific shared access signatures vended by a REST catalog. +pub const ADLS_SAS_TOKEN_PREFIX: &str = "adls.sas-token."; +/// Prefix for the epoch-millisecond expiration of an account-specific SAS token. +pub const ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX: &str = "adls.sas-token-expires-at-ms."; +/// Table property naming the endpoint used to refresh vended ADLS credentials. +pub const ADLS_REFRESH_CREDENTIALS_ENDPOINT: &str = "adls.refresh-credentials-endpoint"; +/// Table property controlling whether vended ADLS credentials are refreshed. +pub const ADLS_REFRESH_CREDENTIALS_ENABLED: &str = "adls.refresh-credentials-enabled"; /// The tenant-id. pub const ADLS_TENANT_ID: &str = "adls.tenant-id"; /// The client-id. diff --git a/crates/iceberg/src/io/storage/config/gcs.rs b/crates/iceberg/src/io/storage/config/gcs.rs index 99867e5dfc..6d06ab6a34 100644 --- a/crates/iceberg/src/io/storage/config/gcs.rs +++ b/crates/iceberg/src/io/storage/config/gcs.rs @@ -41,6 +41,12 @@ pub const GCS_NO_AUTH: &str = "gcs.no-auth"; pub const GCS_CREDENTIALS_JSON: &str = "gcs.credentials-json"; /// Google Cloud Storage token. pub const GCS_TOKEN: &str = "gcs.oauth2.token"; +/// Epoch-millisecond timestamp at which the vended GCS OAuth2 token expires. +pub const GCS_TOKEN_EXPIRES_AT: &str = "gcs.oauth2.token-expires-at"; +/// Endpoint used to fetch and refresh vended GCS OAuth2 credentials. +pub const GCS_REFRESH_CREDENTIALS_ENDPOINT: &str = "gcs.oauth2.refresh-credentials-endpoint"; +/// Whether vended GCS OAuth2 credentials should be refreshed. Defaults to `true`. +pub const GCS_REFRESH_CREDENTIALS_ENABLED: &str = "gcs.oauth2.refresh-credentials-enabled"; /// Option to skip signing requests (e.g. for public buckets/folders). pub const GCS_ALLOW_ANONYMOUS: &str = "gcs.allow-anonymous"; /// Option to skip loading the credential from GCE metadata server. diff --git a/crates/iceberg/src/io/storage/config/mod.rs b/crates/iceberg/src/io/storage/config/mod.rs index d8d356de16..6b55aeda42 100644 --- a/crates/iceberg/src/io/storage/config/mod.rs +++ b/crates/iceberg/src/io/storage/config/mod.rs @@ -50,12 +50,22 @@ use serde::{Deserialize, Serialize}; /// This struct contains only configuration properties without specifying /// which storage backend to use. The storage type is determined by the /// explicit factory selection. -#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize, Default)] +#[derive(Clone, PartialEq, Eq, Serialize, Deserialize, Default)] pub struct StorageConfig { /// Configuration properties for the storage backend props: HashMap, } +impl std::fmt::Debug for StorageConfig { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + // Property values may hold vended credentials (e.g. `s3.secret-access-key`, + // `s3.session-token`) + f.debug_struct("StorageConfig") + .field("keys", &self.props.keys().collect::>()) + .finish_non_exhaustive() + } +} + impl StorageConfig { /// Create a new empty StorageConfig. pub fn new() -> Self { diff --git a/crates/iceberg/src/io/storage/config/s3.rs b/crates/iceberg/src/io/storage/config/s3.rs index ae3bc4ba3b..b0dcc115c1 100644 --- a/crates/iceberg/src/io/storage/config/s3.rs +++ b/crates/iceberg/src/io/storage/config/s3.rs @@ -36,10 +36,16 @@ pub const S3_ACCESS_KEY_ID: &str = "s3.access-key-id"; pub const S3_SECRET_ACCESS_KEY: &str = "s3.secret-access-key"; /// S3 session token (required when using temporary credentials). pub const S3_SESSION_TOKEN: &str = "s3.session-token"; +/// Epoch-millisecond timestamp at which the vended S3 session token expires. +pub const S3_SESSION_TOKEN_EXPIRES_AT_MS: &str = "s3.session-token-expires-at-ms"; /// S3 region. pub const S3_REGION: &str = "s3.region"; /// Region to use for the S3 client (takes precedence over [`S3_REGION`]). pub const CLIENT_REGION: &str = "client.region"; +/// Endpoint used to fetch and refresh vended AWS credentials. +pub const AWS_REFRESH_CREDENTIALS_ENDPOINT: &str = "client.refresh-credentials-endpoint"; +/// Whether vended AWS credentials should be refreshed. Defaults to `true`. +pub const AWS_REFRESH_CREDENTIALS_ENABLED: &str = "client.refresh-credentials-enabled"; /// S3 Path Style Access. pub const S3_PATH_STYLE_ACCESS: &str = "s3.path-style-access"; /// S3 Server Side Encryption Type. diff --git a/crates/iceberg/src/io/storage/mod.rs b/crates/iceberg/src/io/storage/mod.rs index 5276c7771f..a193d7af44 100644 --- a/crates/iceberg/src/io/storage/mod.rs +++ b/crates/iceberg/src/io/storage/mod.rs @@ -23,6 +23,7 @@ mod memory; use std::fmt::Debug; use std::sync::Arc; +use std::time::SystemTime; use async_trait::async_trait; use bytes::Bytes; @@ -32,7 +33,7 @@ pub use local_fs::{LocalFsStorage, LocalFsStorageFactory}; pub use memory::{MemoryStorage, MemoryStorageFactory}; use super::{FileMetadata, FileRead, FileWrite, InputFile, OutputFile}; -use crate::Result; +use crate::{Error, ErrorKind, Result}; /// Trait for storage operations in Iceberg. /// @@ -139,4 +140,358 @@ pub trait StorageFactory: Debug + Send + Sync { /// A `Result` containing an `Arc` on success, or an error /// if the storage could not be created. fn build(&self, config: &StorageConfig) -> Result>; + + /// Build a new Storage instance, optionally supplying a credential provider + /// that the backend can call to obtain and refresh short-lived credentials. + /// + /// Backends that cannot use the provider ignore it and use the credentials + /// in `config`, as they would without one. The default does exactly that. + #[allow(unused_variables)] + fn build_with_credentials( + &self, + config: &StorageConfig, + credential_provider: Option>, + ) -> Result> { + self.build(config) + } +} + +/// Supplies fresh, backend-specific storage credentials on demand. +/// +/// A catalog that vends temporary credentials implements this trait so that +/// storage backends can re-fetch credentials as they approach expiry instead +/// of failing once the initial token's TTL runs out. +/// +/// # Caching +/// +/// [`load_credential`](Self::load_credential) may be called very frequently. +/// Implementations must cache internally and only re-fetch when the current +/// credential is at or near expiry; otherwise every object-store request could +/// trigger a call back to the catalog. +#[async_trait] +pub trait StorageCredentialProvider: Debug + Send + Sync { + /// Return whether this provider has refresh configuration for `path`. + /// + /// Backends use this before replacing their normal credential chain. The + /// default is `true` for single-backend providers; multi-backend providers + /// should return `false` for schemes they do not configure. + fn supports_path(&self, _path: &str) -> bool { + true + } + + /// Load a fresh credential for the storage location identified by `path`. + /// + /// `path` is the absolute location being accessed (e.g. + /// `s3://bucket/warehouse/db/table/...`). Providers that vend distinct + /// credentials per location prefix use it to select the most specific + /// match. When the selected credential has a declared + /// [`StorageCredential::prefix`], it must [cover](StorageCredential::covers) `path`. + async fn load_credential(&self, path: &str) -> Result; + + /// Return a factory that rebuilds an equivalent provider in another process. + /// + /// [`FileIO::serialize_all`](crate::io::FileIO::serialize_all) serializes this + /// factory in place of the provider. The default reports that the provider + /// cannot be serialized. + fn factory(&self) -> Result> { + Err(Error::new( + ErrorKind::FeatureUnsupported, + "storage credential provider cannot be serialized", + )) + } +} + +/// Serializable recipe that rebuilds a [`StorageCredentialProvider`] after +/// [`FileIO`](crate::io::FileIO) deserialization. +/// +/// Factories are serialized through [`typetag`](https://docs.rs/typetag), so +/// implementations must use `#[typetag::serde]`, and the receiving binary must +/// link the concrete implementation. +#[typetag::serde(tag = "type")] +pub trait StorageCredentialProviderFactory: Debug + Send + Sync { + /// Build a provider for a `FileIO` with the given storage configuration. + fn build(&self, config: &StorageConfig) -> Result>; +} + +/// A vended storage credential together with its scope and expiry. +#[derive(Clone, Debug)] +pub struct StorageCredential { + /// Storage-location prefix this credential is scoped to. `None` represents a + /// credential without a declared scope, sourced from flat storage properties. + prefix: Option, + /// The backend-specific credential material. + kind: StorageCredentialKind, + /// When the credential expires, if known. `None` means non-expiring and + /// backends treat such a credential as always valid and never refresh it. + expires_at: Option, +} + +impl StorageCredential { + /// Create a storage credential with no declared scope or expiration. + pub fn new(kind: StorageCredentialKind) -> Self { + Self { + prefix: None, + kind, + expires_at: None, + } + } + + /// Set the storage-location prefix this credential is scoped to. + pub fn with_prefix(mut self, prefix: impl Into) -> Self { + self.prefix = Some(prefix.into()); + self + } + + /// Set when this credential expires. + pub fn with_expiration(mut self, expires_at: SystemTime) -> Self { + self.expires_at = Some(expires_at); + self + } + + /// Return the storage-location prefix this credential is scoped to. + pub fn prefix(&self) -> Option<&str> { + self.prefix.as_deref() + } + + /// Return whether this credential applies to `location`. + /// + /// A credential without a prefix covers every location. Otherwise the + /// prefix must match whole path segments of `location`, and scheme + /// aliases (`s3a`/`s3n` for `s3`, `gcs` for `gs`, and the plain-text + /// Azure schemes for their TLS variants) are treated as equal. A prefix + /// that is only a scheme, such as `s3`, covers every location with that + /// scheme. + pub fn covers(&self, location: &str) -> bool { + self.prefix + .as_deref() + .is_none_or(|prefix| storage_prefix_covers(prefix, location)) + } + + /// Return the backend-specific credential material. + pub fn kind(&self) -> &StorageCredentialKind { + &self.kind + } + + /// Consume this credential and return its backend-specific material. + pub fn into_kind(self) -> StorageCredentialKind { + self.kind + } + + /// Return when this credential expires. + pub fn expires_at(&self) -> Option { + self.expires_at + } +} + +/// Return whether the storage-location `prefix` covers `location`, with the +/// matching rules of [`StorageCredential::covers`]. +pub fn storage_prefix_covers(prefix: &str, location: &str) -> bool { + let Some((location_scheme, location_rest)) = location.split_once("://") else { + return false; + }; + let Some((prefix_scheme, prefix_rest)) = prefix.split_once("://") else { + return !prefix.is_empty() && canonical_scheme(prefix) == canonical_scheme(location_scheme); + }; + + canonical_scheme(prefix_scheme) == canonical_scheme(location_scheme) + && location_rest + .strip_prefix(prefix_rest) + .is_some_and(|remainder| { + prefix_rest.is_empty() + || prefix_rest.ends_with('/') + || remainder.is_empty() + || remainder.starts_with('/') + }) +} + +fn canonical_scheme(scheme: &str) -> String { + let scheme = scheme.to_ascii_lowercase(); + match scheme.as_str() { + "s3a" | "s3n" => "s3".to_string(), + "gcs" => "gs".to_string(), + "abfs" => "abfss".to_string(), + "wasb" => "wasbs".to_string(), + _ => scheme, + } +} + +/// Backend-specific credential material. +#[derive(Clone, Debug)] +#[non_exhaustive] +pub enum StorageCredentialKind { + /// Amazon S3 credentials. + S3(S3Credential), + /// Google Cloud Storage credentials. + Gcs(GcsCredential), + /// Azure Data Lake Storage credentials. + Azdls(AzdlsCredential), +} + +/// Temporary Azure Data Lake Storage credentials (a shared access signature). +#[derive(Clone)] +pub struct AzdlsCredential { + /// Shared access signature used to access Azure storage. + sas_token: String, +} + +impl AzdlsCredential { + /// Create an Azure Data Lake Storage credential. + pub fn new(sas_token: impl Into) -> Self { + Self { + sas_token: sas_token.into(), + } + } + + /// Return the Azure shared access signature. + pub fn sas_token(&self) -> &str { + &self.sas_token + } + + /// Consume this credential and return its shared access signature. + pub fn into_sas_token(self) -> String { + self.sas_token + } +} + +impl Debug for AzdlsCredential { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("AzdlsCredential").finish_non_exhaustive() + } +} + +/// Temporary Amazon S3 credentials. +#[derive(Clone)] +pub struct S3Credential { + /// AWS access key ID. + access_key_id: String, + /// AWS secret access key. + secret_access_key: String, + /// AWS session token, set for temporary (STS/vended) credentials. + session_token: Option, +} + +impl S3Credential { + /// Create temporary Amazon S3 credentials. + pub fn new( + access_key_id: impl Into, + secret_access_key: impl Into, + session_token: Option, + ) -> Self { + Self { + access_key_id: access_key_id.into(), + secret_access_key: secret_access_key.into(), + session_token, + } + } + + /// Return the AWS access key ID. + pub fn access_key_id(&self) -> &str { + &self.access_key_id + } + + /// Return the AWS secret access key. + pub fn secret_access_key(&self) -> &str { + &self.secret_access_key + } + + /// Return the AWS session token, if present. + pub fn session_token(&self) -> Option<&str> { + self.session_token.as_deref() + } + + /// Consume these credentials and return their component values. + pub fn into_parts(self) -> (String, String, Option) { + let Self { + access_key_id, + secret_access_key, + session_token, + } = self; + (access_key_id, secret_access_key, session_token) + } +} + +impl Debug for S3Credential { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("S3Credential").finish_non_exhaustive() + } +} + +/// Temporary Google Cloud Storage credentials (an OAuth2 access token). +#[derive(Clone)] +pub struct GcsCredential { + /// OAuth2 bearer token used to access GCS. + token: String, +} + +impl GcsCredential { + /// Create a Google Cloud Storage credential. + pub fn new(token: impl Into) -> Self { + Self { + token: token.into(), + } + } + + /// Return the OAuth2 bearer token used to access GCS. + pub fn token(&self) -> &str { + &self.token + } + + /// Consume this credential and return its OAuth2 bearer token. + pub fn into_token(self) -> String { + self.token + } +} + +impl Debug for GcsCredential { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("GcsCredential").finish_non_exhaustive() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn scoped(prefix: &str) -> StorageCredential { + StorageCredential::new(StorageCredentialKind::Gcs(GcsCredential::new("token"))) + .with_prefix(prefix) + } + + #[test] + fn credential_prefix_matches_whole_segments() { + let credential = scoped("s3://bucket/table"); + assert!(credential.covers("s3://bucket/table")); + assert!(credential.covers("s3://bucket/table/data/file.parquet")); + assert!(!credential.covers("s3://bucket/table2/data/file.parquet")); + assert!(!credential.covers("s3://bucket/tab")); + assert!(!credential.covers("s3://other/table/file.parquet")); + + assert!(scoped("s3://bucket/table/").covers("s3://bucket/table/file.parquet")); + assert!(scoped("s3://").covers("s3://any/file.parquet")); + } + + #[test] + fn credential_prefix_treats_scheme_aliases_as_equal() { + let credential = scoped("s3://bucket/table"); + assert!(credential.covers("s3a://bucket/table/file.parquet")); + assert!(credential.covers("S3N://bucket/table/file.parquet")); + assert!(scoped("gcs://bucket").covers("gs://bucket/file.parquet")); + assert!( + scoped("abfss://fs@account.dfs.core.windows.net/table") + .covers("abfs://fs@account.dfs.core.windows.net/table/file.parquet") + ); + assert!(!credential.covers("gs://bucket/table/file.parquet")); + } + + #[test] + fn credential_scheme_prefix_covers_the_whole_scheme() { + assert!(scoped("s3").covers("s3a://bucket/file.parquet")); + assert!(!scoped("s3").covers("gs://bucket/file.parquet")); + assert!(!scoped("").covers("s3://bucket/file.parquet")); + assert!(!scoped("s3").covers("not-a-url")); + assert!( + StorageCredential::new(StorageCredentialKind::Gcs(GcsCredential::new("token"))) + .covers("gs://bucket/file.parquet") + ); + } } diff --git a/crates/storage/opendal/Cargo.toml b/crates/storage/opendal/Cargo.toml index c6ea9f9b67..a837206a4e 100644 --- a/crates/storage/opendal/Cargo.toml +++ b/crates/storage/opendal/Cargo.toml @@ -39,9 +39,13 @@ opendal-all = [ "opendal-hf", ] -opendal-azdls = ["opendal/services-azdls"] +opendal-azdls = [ + "opendal/services-azdls", + "reqsign-azure-storage", + "reqsign-core", +] opendal-fs = ["opendal/services-fs"] -opendal-gcs = ["opendal/services-gcs"] +opendal-gcs = ["opendal/services-gcs", "reqsign-google", "reqsign-core"] opendal-hf = ["opendal/services-hf"] opendal-memory = ["opendal/services-memory"] opendal-oss = ["opendal/services-oss"] @@ -56,7 +60,9 @@ futures = { workspace = true } iceberg = { workspace = true } opendal = { workspace = true } reqsign-aws-v4 = { version = "3.0.0", optional = true } +reqsign-azure-storage = { version = "3.2.1", optional = true } reqsign-core = { version = "3.0.0", optional = true } +reqsign-google = { version = "3.0.0", optional = true } serde = { workspace = true } typetag = { workspace = true } url = { workspace = true } diff --git a/crates/storage/opendal/public-api.txt b/crates/storage/opendal/public-api.txt index d8c4ecdb38..ad6c4522b8 100644 --- a/crates/storage/opendal/public-api.txt +++ b/crates/storage/opendal/public-api.txt @@ -1,19 +1,23 @@ pub mod iceberg_storage_opendal pub use iceberg_storage_opendal::AwsCredential pub use iceberg_storage_opendal::ProvideCredential -pub enum iceberg_storage_opendal::OpenDalStorage -pub iceberg_storage_opendal::OpenDalStorage::Azdls +#[non_exhaustive] pub enum iceberg_storage_opendal::OpenDalStorage +#[non_exhaustive] pub iceberg_storage_opendal::OpenDalStorage::Azdls pub iceberg_storage_opendal::OpenDalStorage::Azdls::config: alloc::sync::Arc -pub iceberg_storage_opendal::OpenDalStorage::Gcs +pub iceberg_storage_opendal::OpenDalStorage::Azdls::credential_provider: core::option::Option> +pub iceberg_storage_opendal::OpenDalStorage::Azdls::sas_tokens: alloc::sync::Arc +#[non_exhaustive] pub iceberg_storage_opendal::OpenDalStorage::Gcs pub iceberg_storage_opendal::OpenDalStorage::Gcs::config: alloc::sync::Arc -pub iceberg_storage_opendal::OpenDalStorage::Hf +pub iceberg_storage_opendal::OpenDalStorage::Gcs::credential_provider: core::option::Option> +#[non_exhaustive] pub iceberg_storage_opendal::OpenDalStorage::Hf pub iceberg_storage_opendal::OpenDalStorage::Hf::config: alloc::sync::Arc pub iceberg_storage_opendal::OpenDalStorage::LocalFs pub iceberg_storage_opendal::OpenDalStorage::Memory(opendal_core::types::operator::operator::Operator) -pub iceberg_storage_opendal::OpenDalStorage::Oss +#[non_exhaustive] pub iceberg_storage_opendal::OpenDalStorage::Oss pub iceberg_storage_opendal::OpenDalStorage::Oss::config: alloc::sync::Arc -pub iceberg_storage_opendal::OpenDalStorage::S3 +#[non_exhaustive] pub iceberg_storage_opendal::OpenDalStorage::S3 pub iceberg_storage_opendal::OpenDalStorage::S3::config: alloc::sync::Arc +pub iceberg_storage_opendal::OpenDalStorage::S3::credential_provider: core::option::Option> pub iceberg_storage_opendal::OpenDalStorage::S3::customized_credential_load: core::option::Option impl core::clone::Clone for iceberg_storage_opendal::OpenDalStorage pub fn iceberg_storage_opendal::OpenDalStorage::clone(&self) -> iceberg_storage_opendal::OpenDalStorage @@ -50,10 +54,22 @@ impl core::fmt::Debug for iceberg_storage_opendal::OpenDalStorageFactory pub fn iceberg_storage_opendal::OpenDalStorageFactory::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result impl iceberg::io::storage::StorageFactory for iceberg_storage_opendal::OpenDalStorageFactory pub fn iceberg_storage_opendal::OpenDalStorageFactory::build(&self, config: &iceberg::io::storage::config::StorageConfig) -> iceberg::error::Result> +pub fn iceberg_storage_opendal::OpenDalStorageFactory::build_with_credentials(&self, config: &iceberg::io::storage::config::StorageConfig, credential_provider: core::option::Option>) -> iceberg::error::Result> impl serde_core::ser::Serialize for iceberg_storage_opendal::OpenDalStorageFactory pub fn iceberg_storage_opendal::OpenDalStorageFactory::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg_storage_opendal::OpenDalStorageFactory pub fn iceberg_storage_opendal::OpenDalStorageFactory::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> +pub struct iceberg_storage_opendal::AzdlsSasTokens(_) +impl core::clone::Clone for iceberg_storage_opendal::AzdlsSasTokens +pub fn iceberg_storage_opendal::AzdlsSasTokens::clone(&self) -> iceberg_storage_opendal::AzdlsSasTokens +impl core::default::Default for iceberg_storage_opendal::AzdlsSasTokens +pub fn iceberg_storage_opendal::AzdlsSasTokens::default() -> iceberg_storage_opendal::AzdlsSasTokens +impl core::fmt::Debug for iceberg_storage_opendal::AzdlsSasTokens +pub fn iceberg_storage_opendal::AzdlsSasTokens::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result +impl serde_core::ser::Serialize for iceberg_storage_opendal::AzdlsSasTokens +pub fn iceberg_storage_opendal::AzdlsSasTokens::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer +impl<'de> serde_core::de::Deserialize<'de> for iceberg_storage_opendal::AzdlsSasTokens +pub fn iceberg_storage_opendal::AzdlsSasTokens::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> pub struct iceberg_storage_opendal::CustomAwsCredentialLoader(_) impl iceberg_storage_opendal::CustomAwsCredentialLoader pub fn iceberg_storage_opendal::CustomAwsCredentialLoader::new(provider: impl reqsign_core::api::ProvideCredential + 'static) -> Self @@ -92,6 +108,7 @@ impl core::fmt::Debug for iceberg_storage_opendal::OpenDalResolvingStorageFactor pub fn iceberg_storage_opendal::OpenDalResolvingStorageFactory::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result impl iceberg::io::storage::StorageFactory for iceberg_storage_opendal::OpenDalResolvingStorageFactory pub fn iceberg_storage_opendal::OpenDalResolvingStorageFactory::build(&self, config: &iceberg::io::storage::config::StorageConfig) -> iceberg::error::Result> +pub fn iceberg_storage_opendal::OpenDalResolvingStorageFactory::build_with_credentials(&self, config: &iceberg::io::storage::config::StorageConfig, credential_provider: core::option::Option>) -> iceberg::error::Result> impl serde_core::ser::Serialize for iceberg_storage_opendal::OpenDalResolvingStorageFactory pub fn iceberg_storage_opendal::OpenDalResolvingStorageFactory::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg_storage_opendal::OpenDalResolvingStorageFactory diff --git a/crates/storage/opendal/src/azdls.rs b/crates/storage/opendal/src/azdls.rs index 6d39320df8..387915d54c 100644 --- a/crates/storage/opendal/src/azdls.rs +++ b/crates/storage/opendal/src/azdls.rs @@ -18,18 +18,25 @@ use std::collections::HashMap; use std::fmt::Display; use std::str::FromStr; +use std::sync::Arc; use iceberg::io::{ ADLS_ACCOUNT_KEY, ADLS_ACCOUNT_NAME, ADLS_AUTHORITY_HOST, ADLS_CLIENT_ID, ADLS_CLIENT_SECRET, - ADLS_CONNECTION_STRING, ADLS_SAS_TOKEN, ADLS_TENANT_ID, + ADLS_CONNECTION_STRING, ADLS_SAS_TOKEN, ADLS_SAS_TOKEN_PREFIX, ADLS_TENANT_ID, + StorageCredentialKind, StorageCredentialProvider, }; use iceberg::{Error, ErrorKind, Result}; use opendal::Configurator; use opendal::services::AzdlsConfig; +use reqsign_azure_storage::Credential as AzureCredential; +use reqsign_core::{ + Context, Error as ReqsignError, ProvideCredential, ProvideCredentialChain, + Result as ReqsignResult, +}; use serde::{Deserialize, Serialize}; use url::Url; -use crate::utils::from_opendal_error; +use crate::utils::{VendedCredentialSource, from_opendal_error}; /// Local version of `ensure_data_valid` macro since the iceberg crate's macro /// uses `$crate::error::Error` paths that don't resolve from external crates @@ -84,6 +91,53 @@ pub(crate) fn azdls_config_parse(mut properties: HashMap) -> Res Ok(config) } +/// Account-specific ADLS SAS tokens supplied through Java-compatible storage +/// properties. +#[derive(Clone, Default, Serialize, Deserialize)] +pub struct AzdlsSasTokens(HashMap); + +impl AzdlsSasTokens { + /// Collect `adls.sas-token.` properties by storage account. Like + /// Java, keys may name the host (`account.dfs.core.windows.net`) or only + /// the account. + pub(crate) fn from_properties(properties: &HashMap) -> Self { + let mut tokens = properties + .iter() + .filter_map(|(key, value)| { + let account = sas_token_account(key.strip_prefix(ADLS_SAS_TOKEN_PREFIX)?); + (!account.is_empty() && !value.is_empty()).then_some((key, account, value)) + }) + .collect::>(); + // Deterministic choice when several keys name the same account: the + // last, host-keyed one wins. + tokens.sort(); + Self( + tokens + .into_iter() + .map(|(_, account, value)| (account.to_string(), value.clone())) + .collect(), + ) + } + + fn for_path(&self, path: &AzureStoragePath) -> Option<&str> { + self.0.get(&path.account_name).map(String::as_str) + } +} + +/// The storage account named by the suffix of an account-specific SAS token +/// property: a host such as `account.dfs.core.windows.net`, or the account. +fn sas_token_account(suffix: &str) -> &str { + suffix.split('.').next().unwrap_or(suffix) +} + +impl std::fmt::Debug for AzdlsSasTokens { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("AzdlsSasTokens") + .field("account_count", &self.0.len()) + .finish_non_exhaustive() + } +} + /// Builds an OpenDAL operator from the AzdlsConfig and path. /// /// The path is expected to include the scheme in a format like: @@ -91,11 +145,21 @@ pub(crate) fn azdls_config_parse(mut properties: HashMap) -> Res pub(crate) fn azdls_create_operator<'a>( absolute_path: &'a str, config: &AzdlsConfig, + sas_tokens: &AzdlsSasTokens, + credential_provider: &Option>, + credential_location: Option<&str>, ) -> Result<(opendal::Operator, &'a str)> { let path = absolute_path.parse::()?; match_path_with_config(&path, config)?; - let op = azdls_config_build(config, &path)?; + let op = azdls_config_build( + config, + &path, + sas_tokens, + credential_provider, + absolute_path, + credential_location, + )?; // Paths to files in ADLS tend to be written in fully qualified form, // including their filesystem and account name. @@ -110,8 +174,9 @@ pub(crate) fn azdls_create_operator<'a>( /// Note that `abf[s]` and `wasb[s]` variants have different implications: /// - `abfs[s]` is used to refer to files in ADLS Gen2, backed by blob storage; /// paths are expected to contain the `dfs` storage service. -/// - `wasb[s]` is used to refer to files in Blob Storage directly; paths are -/// expected to contain the `blob` storage service. +/// - `wasb[s]` is accepted for compatibility with Blob Storage locations; +/// paths contain the `blob` storage service, but operations still use the +/// ADLS Gen2 `dfs` endpoint, matching Iceberg Java. #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] pub enum AzureStorageScheme { Abfs, @@ -121,12 +186,18 @@ pub enum AzureStorageScheme { } impl AzureStorageScheme { - // Returns the respective encrypted or plain-text HTTP scheme. + /// The HTTP scheme of an endpoint derived from a path. + /// + /// Iceberg Java accepts the non-secure aliases for compatibility but still + /// connects over TLS. SAS tokens are query parameters and must not be sent + /// over a plaintext connection unless the user configures such an endpoint. pub fn as_http_scheme(&self) -> &str { - match self { - AzureStorageScheme::Abfs | AzureStorageScheme::Wasb => "http", - AzureStorageScheme::Abfss | AzureStorageScheme::Wasbs => "https", - } + "https" + } + + /// Whether an explicitly configured endpoint must use TLS. + fn requires_tls(&self) -> bool { + matches!(self, AzureStorageScheme::Abfss | AzureStorageScheme::Wasbs) } } @@ -170,12 +241,13 @@ pub(crate) fn match_path_with_config(path: &AzureStoragePath, config: &AzdlsConf } if let Some(ref configured_endpoint) = config.endpoint { - let passed_http_scheme = path.scheme.as_http_scheme(); + // An explicit plaintext endpoint, such as a local emulator, remains + // valid for the non-secure schemes. ensure_data_valid!( - configured_endpoint.starts_with(passed_http_scheme), - "Storage::Azdls: Endpoint {} does not use the expected http scheme {}.", + !path.scheme.requires_tls() || configured_endpoint.starts_with("https://"), + "Storage::Azdls: Endpoint {} does not use https, which the {} scheme requires.", configured_endpoint, - passed_http_scheme + path.scheme ); let ends_with_expected_suffix = configured_endpoint @@ -192,7 +264,14 @@ pub(crate) fn match_path_with_config(path: &AzureStoragePath, config: &AzdlsConf Ok(()) } -fn azdls_config_build(config: &AzdlsConfig, path: &AzureStoragePath) -> Result { +fn azdls_config_build( + config: &AzdlsConfig, + path: &AzureStoragePath, + sas_tokens: &AzdlsSasTokens, + credential_provider: &Option>, + absolute_path: &str, + credential_location: Option<&str>, +) -> Result { let mut builder = config.clone().into_builder(); if config.endpoint.is_none() { @@ -201,9 +280,50 @@ fn azdls_config_build(config: &AzdlsConfig, path: &AzureStoragePath) -> Result ReqsignResult> { + match self.0.load("ADLS").await? { + (StorageCredentialKind::Azdls(azdls), expires_at) => { + let sas_token = azdls.into_sas_token(); + Ok(Some(match expires_at { + Some(expires_at) => { + AzureCredential::with_sas_token_expires_at(&sas_token, expires_at) + } + None => AzureCredential::with_sas_token(&sas_token), + })) + } + _ => Err(ReqsignError::unexpected( + "ADLS storage received a non-ADLS credential from the provider", + )), + } + } +} + /// Represents a fully qualified path to blob/ file in Azure Storage. #[derive(Debug, PartialEq)] pub(crate) struct AzureStoragePath { @@ -265,6 +385,14 @@ impl FromStr for AzureStoragePath { } } +pub(crate) fn azdls_batch_key(absolute_path: &str) -> Result { + let path = absolute_path.parse::()?; + Ok(format!( + "{}://{}@{}.{}", + path.scheme, path.filesystem, path.account_name, path.endpoint_suffix + )) +} + fn parse_azure_storage_endpoint(url: &Url) -> Result<(&str, &str, &str)> { let host = url.host_str().ok_or(Error::new( ErrorKind::DataInvalid, @@ -318,10 +446,33 @@ fn validate_storage_and_scheme( #[cfg(test)] mod tests { use std::collections::HashMap; - + use std::sync::Arc; + use std::time::{Duration, SystemTime}; + + use async_trait::async_trait; + use iceberg::Result; + use iceberg::io::{ + ADLS_SAS_TOKEN_PREFIX, AzdlsCredential, StorageCredential, StorageCredentialKind, + StorageCredentialProvider, + }; use opendal::services::AzdlsConfig; + use reqsign_azure_storage::Credential as AzureCredential; + use reqsign_core::{Context, ProvideCredential}; - use super::{AzureStoragePath, AzureStorageScheme, azdls_config_parse, azdls_create_operator}; + use super::{ + AzdlsSasTokens, AzureStoragePath, AzureStorageScheme, VendedAzdlsCredentialProvider, + VendedCredentialSource, azdls_batch_key, azdls_config_parse, azdls_create_operator, + }; + + #[derive(Debug)] + struct FixedCredentialProvider(StorageCredential); + + #[async_trait] + impl StorageCredentialProvider for FixedCredentialProvider { + async fn load_credential(&self, _path: &str) -> Result { + Ok(self.0.clone()) + } + } #[test] fn test_azdls_config_parse() { @@ -409,6 +560,18 @@ mod tests { ), None, ), + ( + "plaintext endpoint for a non-secure scheme", + ( + "abfs://myfs@myaccount.dfs.core.windows.net/path/to/file.parquet", + AzdlsConfig { + account_name: Some("myaccount".to_string()), + endpoint: Some("http://myaccount.dfs.core.windows.net".to_string()), + ..Default::default() + }, + ), + Some(("myfs", "/path/to/file.parquet")), + ), ( "incompatible scheme for endpoint", ( @@ -464,7 +627,8 @@ mod tests { ]; for (name, input, expected) in test_cases { - let result = azdls_create_operator(input.0, &input.1); + let result = + azdls_create_operator(input.0, &input.1, &AzdlsSasTokens::default(), &None, None); match expected { Some((expected_filesystem, expected_path)) => { assert!(result.is_ok(), "Test case {name} failed: {result:?}"); @@ -480,6 +644,106 @@ mod tests { } } + #[tokio::test] + async fn vended_provider_returns_expiring_sas_credential() { + let path = "abfss://container@account.dfs.core.windows.net/table/data.parquet"; + let expires_at = SystemTime::now() + Duration::from_secs(3600); + let credential = StorageCredential::new(StorageCredentialKind::Azdls( + AzdlsCredential::new("sv=2026&sig=secret"), + )) + .with_prefix("abfss://container@account.dfs.core.windows.net/table") + .with_expiration(expires_at); + let provider = VendedAzdlsCredentialProvider(VendedCredentialSource::new( + Arc::new(FixedCredentialProvider(credential)), + path.to_string(), + )); + + let credential = provider + .provide_credential(&Context::new()) + .await + .unwrap() + .unwrap(); + match credential { + AzureCredential::SasToken { + token, + expires_at: actual_expires_at, + } => { + assert_eq!(token, "sv=2026&sig=secret"); + assert_eq!( + actual_expires_at, + Some(crate::utils::system_time_to_timestamp(expires_at).unwrap()) + ); + } + other => panic!("expected SAS token, got {other:?}"), + } + } + + #[test] + fn account_specific_sas_tokens_are_selected_by_storage_account() { + let properties = HashMap::from([ + ( + format!("{ADLS_SAS_TOKEN_PREFIX}first"), + "sv=2026&sig=first".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_PREFIX}second"), + "sv=2026&sig=second".to_string(), + ), + ]); + let sas_tokens = AzdlsSasTokens::from_properties(&properties); + let path = "abfss://container@second.dfs.core.windows.net/table/data.parquet" + .parse::() + .unwrap(); + + assert_eq!(sas_tokens.for_path(&path), Some("sv=2026&sig=second")); + } + + #[test] + fn host_keyed_sas_tokens_match_java() { + let properties = HashMap::from([( + format!("{ADLS_SAS_TOKEN_PREFIX}account.dfs.core.windows.net"), + "sv=2026&sig=host".to_string(), + )]); + let sas_tokens = AzdlsSasTokens::from_properties(&properties); + + for location in [ + "abfss://container@account.dfs.core.windows.net/table/data.parquet", + "wasbs://container@account.blob.core.windows.net/table/data.parquet", + ] { + let path = location.parse::().unwrap(); + assert_eq!(sas_tokens.for_path(&path), Some("sv=2026&sig=host")); + } + } + + #[tokio::test] + async fn vended_provider_rejects_mismatched_prefix() { + let credential = StorageCredential::new(StorageCredentialKind::Azdls( + AzdlsCredential::new("sv=2026&sig=secret"), + )) + .with_prefix("abfss://other@account.dfs.core.windows.net/table") + .with_expiration(SystemTime::now() + Duration::from_secs(3600)); + let provider = VendedAzdlsCredentialProvider(VendedCredentialSource::new( + Arc::new(FixedCredentialProvider(credential)), + "abfss://container@account.dfs.core.windows.net/table/data.parquet".to_string(), + )); + + assert!(provider.provide_credential(&Context::new()).await.is_err()); + } + + #[test] + fn batch_key_distinguishes_azure_filesystems_and_schemes() { + let first = azdls_batch_key("abfss://first@account.dfs.core.windows.net/table/a.parquet"); + let second = azdls_batch_key("abfss://second@account.dfs.core.windows.net/table/b.parquet"); + let blob = azdls_batch_key("wasbs://first@account.blob.core.windows.net/table/c.parquet"); + + assert_ne!(first.unwrap(), second.unwrap()); + assert_ne!( + azdls_batch_key("abfss://first@account.dfs.core.windows.net/table/a.parquet").unwrap(), + blob.unwrap() + ); + assert!(azdls_batch_key("abfss:///no-account.parquet").is_err()); + } + #[test] fn test_azure_storage_path_parse() { let test_cases = vec![ @@ -540,7 +804,7 @@ mod tests { "https://myaccount.dfs.core.windows.net", ), ( - "abfs uses http", + "abfs uses https for Java compatibility and SAS security", AzureStoragePath { scheme: AzureStorageScheme::Abfs, filesystem: "myfs".to_string(), @@ -548,12 +812,12 @@ mod tests { endpoint_suffix: "core.windows.net".to_string(), path: "/path/to/file.parquet".to_string(), }, - "http://myaccount.dfs.core.windows.net", + "https://myaccount.dfs.core.windows.net", ), ( "wasbs uses https and dfs", AzureStoragePath { - scheme: AzureStorageScheme::Abfss, + scheme: AzureStorageScheme::Wasbs, filesystem: "myfs".to_string(), account_name: "myaccount".to_string(), endpoint_suffix: "core.windows.net".to_string(), diff --git a/crates/storage/opendal/src/gcs.rs b/crates/storage/opendal/src/gcs.rs index a00282cc22..f374b8034c 100644 --- a/crates/storage/opendal/src/gcs.rs +++ b/crates/storage/opendal/src/gcs.rs @@ -17,17 +17,20 @@ //! Google Cloud Storage properties use std::collections::HashMap; +use std::sync::Arc; use iceberg::io::{ GCS_ALLOW_ANONYMOUS, GCS_CREDENTIALS_JSON, GCS_DISABLE_CONFIG_LOAD, GCS_DISABLE_VM_METADATA, - GCS_NO_AUTH, GCS_SERVICE_HOST, GCS_TOKEN, + GCS_NO_AUTH, GCS_SERVICE_HOST, GCS_TOKEN, StorageCredentialKind, StorageCredentialProvider, }; use iceberg::{Error, ErrorKind, Result}; -use opendal::Operator; use opendal::services::GcsConfig; +use opendal::{Configurator, Operator}; +use reqsign_core::{Context, Error as ReqsignError, ProvideCredential, Result as ReqsignResult}; +use reqsign_google::{Credential as GoogleCredential, Token as GoogleToken}; use url::Url; -use crate::utils::{from_opendal_error, is_truthy}; +use crate::utils::{VendedCredentialSource, from_opendal_error, is_truthy}; /// Parse iceberg properties to [`GcsConfig`]. pub(crate) fn gcs_config_parse(mut m: HashMap) -> Result { @@ -45,7 +48,9 @@ pub(crate) fn gcs_config_parse(mut m: HashMap) -> Result) -> Result Result { +pub(crate) fn gcs_config_build( + cfg: &GcsConfig, + credential_provider: &Option>, + path: &str, + credential_location: Option<&str>, +) -> Result { let url = Url::parse(path)?; + if !matches!(url.scheme(), "gs" | "gcs") { + return Err(Error::new( + ErrorKind::DataInvalid, + format!("Invalid gcs url: {path}, expected gs:// or gcs://"), + )); + } let bucket = url.host_str().ok_or_else(|| { Error::new( ErrorKind::DataInvalid, @@ -82,5 +98,63 @@ pub(crate) fn gcs_config_build(cfg: &GcsConfig, path: &str) -> Result let mut cfg = cfg.clone(); cfg.bucket = bucket.to_string(); - Operator::from_config(cfg).map_err(from_opendal_error) + + // `reqsign_google` continues to the next provider even when a provider returns + // an error, and OpenDAL prepends custom providers to its default chain. Disable + // every other configured and ambient source so the catalog provider is + // effectively the sole source and refresh failures cannot silently fall back. + let credential_provider = credential_provider + .as_ref() + .filter(|provider| provider.supports_path(path)); + if credential_provider.is_some() { + if cfg.skip_signature { + return Err(Error::new( + ErrorKind::DataInvalid, + "Invalid GCS auth settings: anonymous access cannot be combined with refreshable credentials", + )); + } + cfg.token = None; + cfg.credential = None; + cfg.credential_path = None; + cfg.service_account = None; + cfg.disable_vm_metadata = true; + cfg.disable_config_load = true; + } + + let mut builder = cfg.into_builder(); + + // A catalog-supplied provider re-fetches the vended OAuth2 token as it nears expiry + if let Some(provider) = credential_provider { + builder = + builder.credential_provider(VendedGcsCredentialProvider(VendedCredentialSource::new( + Arc::clone(provider), + credential_location.unwrap_or(path).to_string(), + ))); + } + + Operator::new(builder).map_err(from_opendal_error) +} + +/// Adapts a generic [`StorageCredentialProvider`] into a `reqsign` +/// [`ProvideCredential`], so the GCS signer can obtain and refresh vended OAuth2 +/// tokens. +#[derive(Debug)] +struct VendedGcsCredentialProvider(VendedCredentialSource); + +impl ProvideCredential for VendedGcsCredentialProvider { + type Credential = GoogleCredential; + + async fn provide_credential(&self, _ctx: &Context) -> ReqsignResult> { + match self.0.load("GCS").await? { + (StorageCredentialKind::Gcs(gcs), expires_at) => { + Ok(Some(GoogleCredential::with_token(GoogleToken { + access_token: gcs.into_token(), + expires_at, + }))) + } + _ => Err(ReqsignError::unexpected( + "GCS storage received a non-GCS credential from the provider", + )), + } + } } diff --git a/crates/storage/opendal/src/lib.rs b/crates/storage/opendal/src/lib.rs index fe485cb915..67a71060c4 100644 --- a/crates/storage/opendal/src/lib.rs +++ b/crates/storage/opendal/src/lib.rs @@ -35,7 +35,7 @@ use futures::StreamExt; use futures::stream::BoxStream; use iceberg::io::{ FileMetadata, FileRead, FileWrite, InputFile, OutputFile, Storage, StorageConfig, - StorageFactory, + StorageCredentialProvider, StorageFactory, }; use iceberg::{Error, ErrorKind, Result}; use opendal::Operator; @@ -46,6 +46,7 @@ use utils::from_opendal_error; cfg_if! { if #[cfg(feature = "opendal-azdls")] { mod azdls; + pub use azdls::AzdlsSasTokens; use azdls::*; use opendal::services::AzdlsConfig; } @@ -161,8 +162,18 @@ where #[typetag::serde(name = "OpenDalStorageFactory")] impl StorageFactory for OpenDalStorageFactory { - #[allow(unused_variables)] fn build(&self, config: &StorageConfig) -> Result> { + self.build_with_credentials(config, None) + } + + #[allow(unused_variables)] + fn build_with_credentials( + &self, + config: &StorageConfig, + credential_provider: Option>, + ) -> Result> { + // Only S3, GCS and ADLS consume the provider; other backends use the + // credentials in `config`. match self { #[cfg(feature = "opendal-memory")] OpenDalStorageFactory::Memory => { @@ -176,10 +187,12 @@ impl StorageFactory for OpenDalStorageFactory { } => Ok(Arc::new(OpenDalStorage::S3 { config: s3_config_parse(config.props().clone())?.into(), customized_credential_load: customized_credential_load.clone(), + credential_provider, })), #[cfg(feature = "opendal-gcs")] OpenDalStorageFactory::Gcs => Ok(Arc::new(OpenDalStorage::Gcs { config: gcs_config_parse(config.props().clone())?.into(), + credential_provider, })), #[cfg(feature = "opendal-oss")] OpenDalStorageFactory::Oss => Ok(Arc::new(OpenDalStorage::Oss { @@ -188,6 +201,8 @@ impl StorageFactory for OpenDalStorageFactory { #[cfg(feature = "opendal-azdls")] OpenDalStorageFactory::Azdls => Ok(Arc::new(OpenDalStorage::Azdls { config: azdls_config_parse(config.props().clone())?.into(), + sas_tokens: Arc::new(AzdlsSasTokens::from_properties(config.props())), + credential_provider, })), #[cfg(feature = "opendal-hf")] OpenDalStorageFactory::Hf => Ok(Arc::new(OpenDalStorage::Hf { @@ -218,6 +233,7 @@ fn default_memory_operator() -> Operator { /// OpenDAL-based storage implementation. #[derive(Clone, Debug, Serialize, Deserialize)] +#[non_exhaustive] pub enum OpenDalStorage { /// Memory storage variant. #[cfg(feature = "opendal-memory")] @@ -229,6 +245,7 @@ pub enum OpenDalStorage { /// /// Accepts any S3-family URL (`s3://`, `s3a://`, `s3n://`); the scheme is /// derived from the path at call time. + #[non_exhaustive] #[cfg(feature = "opendal-s3")] S3 { /// S3 configuration. @@ -236,14 +253,22 @@ pub enum OpenDalStorage { /// Custom AWS credential loader. #[serde(skip)] customized_credential_load: Option, + /// Provider of refreshable vended credentials, supplied by the catalog. + #[serde(skip)] + credential_provider: Option>, }, /// GCS storage variant. + #[non_exhaustive] #[cfg(feature = "opendal-gcs")] Gcs { /// GCS configuration. config: Arc, + /// Provider of refreshable vended credentials, supplied by the catalog. + #[serde(skip)] + credential_provider: Option>, }, /// OSS storage variant. + #[non_exhaustive] #[cfg(feature = "opendal-oss")] Oss { /// OSS configuration. @@ -255,16 +280,24 @@ pub enum OpenDalStorage { /// `abfs[s]://@.dfs./` or /// `wasb[s]://@.blob./`. /// The scheme is derived from the path at call time. + #[non_exhaustive] #[cfg(feature = "opendal-azdls")] Azdls { /// Azure DLS configuration. config: Arc, + /// Account-specific SAS tokens supplied by Java-compatible properties. + #[serde(default)] + sas_tokens: Arc, + /// Provider of refreshable vended credentials, supplied by the catalog. + #[serde(skip)] + credential_provider: Option>, }, /// HuggingFace Hub storage variant. /// /// Accepts paths of the form /// `hf:////[@]/`, /// where `` must be one of `models`, `datasets`, `spaces`, or `buckets`. + #[non_exhaustive] #[cfg(feature = "opendal-hf")] Hf { /// HuggingFace Hub configuration (token + endpoint). @@ -272,6 +305,16 @@ pub enum OpenDalStorage { }, } +/// Groups bulk deletes by operator and credential scope. +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +struct DeleteBatchKey { + storage: String, + /// Location used to look up the dynamic credential shared by the batch: + /// the credential prefix, or the storage root for an unscoped credential. + /// `None` when the path is served by static credentials. + credential_location: Option, +} + impl OpenDalStorage { /// Creates operator from path. /// @@ -289,6 +332,20 @@ impl OpenDalStorage { pub(crate) fn create_operator<'a>( &self, path: &'a impl AsRef, + ) -> Result<(Operator, &'a str)> { + self.create_operator_with_credential_location(path, None) + } + + /// Creates an operator whose dynamic credentials are looked up for + /// `credential_location` instead of `path`. + /// + /// A bulk-delete batch passes its scope location, so every credential the + /// operator loads covers all paths in the batch. + #[allow(unreachable_code, unused_variables)] + fn create_operator_with_credential_location<'a>( + &self, + path: &'a impl AsRef, + credential_location: Option<&str>, ) -> Result<(Operator, &'a str)> { let path = path.as_ref(); let (operator, relative_path): (Operator, &str) = match self { @@ -313,8 +370,15 @@ impl OpenDalStorage { OpenDalStorage::S3 { config, customized_credential_load, + credential_provider, } => { - let op = s3_config_build(config, customized_credential_load, path)?; + let op = s3_config_build( + config, + customized_credential_load, + credential_provider, + path, + credential_location, + )?; let op_info = op.info(); // Use the URL scheme in the path for prefix matching. This enables @@ -336,9 +400,19 @@ impl OpenDalStorage { } } #[cfg(feature = "opendal-gcs")] - OpenDalStorage::Gcs { config } => { - let operator = gcs_config_build(config, path)?; - let prefix = format!("gs://{}/", operator.info().name()); + OpenDalStorage::Gcs { + config, + credential_provider, + } => { + let operator = + gcs_config_build(config, credential_provider, path, credential_location)?; + let url = url::Url::parse(path).map_err(|e| { + Error::new( + ErrorKind::DataInvalid, + format!("Invalid gcs url: {path}: {e}"), + ) + })?; + let prefix = format!("{}://{}/", url.scheme(), operator.info().name()); if path.starts_with(&prefix) { (operator, &path[prefix.len()..]) } else { @@ -362,7 +436,17 @@ impl OpenDalStorage { } } #[cfg(feature = "opendal-azdls")] - OpenDalStorage::Azdls { config } => azdls_create_operator(path, config)?, + OpenDalStorage::Azdls { + config, + sas_tokens, + credential_provider, + } => azdls_create_operator( + path, + config, + sas_tokens, + credential_provider, + credential_location, + )?, #[cfg(feature = "opendal-hf")] OpenDalStorage::Hf { config } => hf_config_build(config, path)?, #[cfg(all( @@ -399,17 +483,77 @@ impl OpenDalStorage { /// /// For most backends the URL host (bucket name) is sufficient. For HF the host /// encodes the repo type, not the repo identity, so a more specific key is used. - fn batch_key_for_path(&self, path: &str) -> String { + #[allow(unreachable_patterns)] + fn batch_key_for_path(&self, path: &str) -> Result { match self { #[cfg(feature = "opendal-hf")] - OpenDalStorage::Hf { .. } => hf_batch_key(path), - _ => url::Url::parse(path) + OpenDalStorage::Hf { .. } => Ok(hf_batch_key(path)), + #[cfg(feature = "opendal-azdls")] + OpenDalStorage::Azdls { .. } => azdls_batch_key(path), + _ => Ok(url::Url::parse(path) .ok() .and_then(|u| u.host_str().map(|s| s.to_string())) - .unwrap_or_default(), + .unwrap_or_default()), } } + /// Return the dynamic credential provider that serves `path`, if any. + fn credential_provider_for_path( + &self, + path: &str, + ) -> Option<&Arc> { + let provider: &Arc = (match self { + #[cfg(feature = "opendal-s3")] + OpenDalStorage::S3 { + customized_credential_load: None, + credential_provider: Some(provider), + .. + } => Some(provider), + #[cfg(feature = "opendal-gcs")] + OpenDalStorage::Gcs { + credential_provider: Some(provider), + .. + } => Some(provider), + #[cfg(feature = "opendal-azdls")] + OpenDalStorage::Azdls { + credential_provider: Some(provider), + .. + } => Some(provider), + _ => None, + })?; + provider.supports_path(path).then_some(provider) + } + + /// Returns a key that keeps bulk deletes within one operator and credential + /// scope. Loading the credential is normally a cache hit and avoids rebuilding + /// an operator for every path while preventing a batch from crossing prefixes. + async fn delete_batch_key_for_path(&self, path: &str) -> Result { + let credential_location = match self.credential_provider_for_path(path) { + Some(provider) => { + let credential = provider.load_credential(path).await?; + if !credential.covers(path) { + return Err(Error::new( + ErrorKind::DataInvalid, + format!( + "vended credential prefix {:?} does not cover storage location {path:?}", + credential.prefix() + ), + )); + } + Some(match credential.prefix() { + Some(prefix) => prefix.to_string(), + None => utils::storage_root(path)?, + }) + } + None => None, + }; + + Ok(DeleteBatchKey { + storage: self.batch_key_for_path(path)?, + credential_location, + }) + } + /// Extracts the relative path from an absolute path without building an operator. /// /// This is a lightweight alternative to [`create_operator`](Self::create_operator) for cases @@ -444,13 +588,19 @@ impl OpenDalStorage { #[cfg(feature = "opendal-gcs")] OpenDalStorage::Gcs { .. } => { let url = url::Url::parse(path)?; + if !matches!(url.scheme(), "gs" | "gcs") { + return Err(Error::new( + ErrorKind::DataInvalid, + format!("Invalid gcs url: {path}, expected gs:// or gcs://"), + )); + } let bucket = url.host_str().ok_or_else(|| { Error::new( ErrorKind::DataInvalid, format!("Invalid gcs url: {path}, missing bucket"), ) })?; - let prefix = format!("gs://{}/", bucket); + let prefix = format!("{}://{}/", url.scheme(), bucket); if path.starts_with(&prefix) { Ok(&path[prefix.len()..]) } else { @@ -480,7 +630,7 @@ impl OpenDalStorage { } } #[cfg(feature = "opendal-azdls")] - OpenDalStorage::Azdls { config } => { + OpenDalStorage::Azdls { config, .. } => { let azure_path = path.parse::()?; match_path_with_config(&azure_path, config)?; let relative_path_len = azure_path.path.len(); @@ -576,17 +726,20 @@ impl Storage for OpenDalStorage { } async fn delete_stream(&self, mut paths: BoxStream<'static, String>) -> Result<()> { - let mut deleters: HashMap = HashMap::new(); + let mut deleters: HashMap = HashMap::new(); while let Some(path) = paths.next().await { - let bucket = self.batch_key_for_path(&path); + let batch_key = self.delete_batch_key_for_path(&path).await?; - let (relative_path, deleter) = match deleters.entry(bucket) { + let (relative_path, deleter) = match deleters.entry(batch_key) { Entry::Occupied(entry) => { (self.relativize_path(&path)?.to_string(), entry.into_mut()) } Entry::Vacant(entry) => { - let (op, rel) = self.create_operator(&path)?; + let (op, rel) = self.create_operator_with_credential_location( + &path, + entry.key().credential_location.as_deref(), + )?; let rel = rel.to_string(); let deleter = op.deleter().await.map_err(from_opendal_error)?; (rel, entry.insert(deleter)) @@ -685,8 +838,39 @@ impl FileWrite for OpenDalWriter { #[cfg(test)] mod tests { + #[allow(unused_imports)] use super::*; + #[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + all(feature = "opendal-memory", feature = "opendal-azdls") + ))] + #[derive(Debug)] + struct AlwaysSupportedCredentialProvider; + + #[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + all(feature = "opendal-memory", feature = "opendal-azdls") + ))] + #[async_trait] + impl StorageCredentialProvider for AlwaysSupportedCredentialProvider { + async fn load_credential(&self, path: &str) -> Result { + let prefix = if path.contains("/table-a/") { + "s3://bucket/table-a" + } else { + "s3://bucket/table-b" + }; + Ok( + iceberg::io::StorageCredential::new(iceberg::io::StorageCredentialKind::S3( + iceberg::io::S3Credential::new("access-key", "secret-key", None), + )) + .with_prefix(prefix), + ) + } + } + #[cfg(feature = "opendal-s3")] #[derive(Debug)] struct EmptyCredentialLoader; @@ -725,6 +909,26 @@ mod tests { assert_eq!(op.info().scheme().to_string(), "memory"); } + #[cfg(all( + feature = "opendal-memory", + any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + ) + ))] + #[test] + fn test_factory_ignores_credentials_for_unsupported_backend() { + let storage = OpenDalStorageFactory::Memory + .build_with_credentials( + &StorageConfig::new(), + Some(Arc::new(AlwaysSupportedCredentialProvider)), + ) + .expect("memory must ignore a credential provider"); + + assert!(storage.new_input("memory:/key").is_ok()); + } + #[cfg(feature = "opendal-memory")] #[tokio::test] async fn test_writer_close_returns_stored_size() { @@ -796,6 +1000,7 @@ mod tests { let storage = OpenDalStorage::S3 { config: Arc::new(S3Config::default()), customized_credential_load: None, + credential_provider: None, }; // All S3-family schemes are accepted by the same storage instance. @@ -811,19 +1016,143 @@ mod tests { } } + #[cfg(feature = "opendal-s3")] + #[tokio::test] + async fn test_dynamic_credentials_batch_by_prefix() { + let storage = OpenDalStorage::S3 { + config: Arc::new(S3Config::default()), + customized_credential_load: None, + credential_provider: Some(Arc::new(AlwaysSupportedCredentialProvider)), + }; + let first = "s3://bucket/table-a/file.parquet"; + let same_scope = "s3://bucket/table-a/other.parquet"; + let other_scope = "s3://bucket/table-b/file.parquet"; + + let first_key = storage.delete_batch_key_for_path(first).await.unwrap(); + assert_eq!( + first_key.credential_location.as_deref(), + Some("s3://bucket/table-a") + ); + assert_eq!( + first_key, + storage.delete_batch_key_for_path(same_scope).await.unwrap() + ); + assert_ne!( + first_key, + storage + .delete_batch_key_for_path(other_scope) + .await + .unwrap() + ); + } + + #[cfg(feature = "opendal-s3")] + #[tokio::test] + async fn test_unscoped_dynamic_credentials_batch_by_storage_root() { + #[derive(Debug)] + struct UnscopedProvider; + + #[async_trait] + impl StorageCredentialProvider for UnscopedProvider { + async fn load_credential(&self, _path: &str) -> Result { + Ok(iceberg::io::StorageCredential::new( + iceberg::io::StorageCredentialKind::S3(iceberg::io::S3Credential::new( + "access-key", + "secret-key", + None, + )), + )) + } + } + + let storage = OpenDalStorage::S3 { + config: Arc::new(S3Config::default()), + customized_credential_load: None, + credential_provider: Some(Arc::new(UnscopedProvider)), + }; + + let key = storage + .delete_batch_key_for_path("s3://bucket/table-a/file.parquet") + .await + .unwrap(); + assert_eq!(key.credential_location.as_deref(), Some("s3://bucket/")); + assert_eq!( + key, + storage + .delete_batch_key_for_path("s3://bucket/table-b/file.parquet") + .await + .unwrap() + ); + } + + #[cfg(feature = "opendal-s3")] + #[tokio::test] + async fn test_custom_s3_credential_loader_ignores_dynamic_provider_for_batching() { + let storage = OpenDalStorage::S3 { + config: Arc::new(S3Config::default()), + customized_credential_load: Some(CustomAwsCredentialLoader::new( + reqsign_aws_v4::StaticCredentialProvider::new("access-key", "secret-key"), + )), + credential_provider: Some(Arc::new(AlwaysSupportedCredentialProvider)), + }; + + let key = storage + .delete_batch_key_for_path("s3://bucket/table-a/file.parquet") + .await + .unwrap(); + assert_eq!(key.credential_location, None); + } + + #[cfg(feature = "opendal-s3")] + #[test] + fn test_s3_rejects_anonymous_dynamic_credentials() { + let mut config = S3Config::default(); + config.skip_signature = true; + let storage = OpenDalStorage::S3 { + config: Arc::new(config), + customized_credential_load: None, + credential_provider: Some(Arc::new(AlwaysSupportedCredentialProvider)), + }; + + let error = storage + .create_operator(&"s3://bucket/file.parquet") + .unwrap_err(); + assert_eq!(error.kind(), ErrorKind::DataInvalid); + } + + #[cfg(feature = "opendal-gcs")] + #[test] + fn test_gcs_rejects_anonymous_dynamic_credentials() { + let mut config = GcsConfig::default(); + config.skip_signature = true; + let storage = OpenDalStorage::Gcs { + config: Arc::new(config), + credential_provider: Some(Arc::new(AlwaysSupportedCredentialProvider)), + }; + + let error = storage + .create_operator(&"gs://bucket/file.parquet") + .unwrap_err(); + assert_eq!(error.kind(), ErrorKind::DataInvalid); + } + #[cfg(feature = "opendal-gcs")] #[test] fn test_relativize_path_gcs() { let storage = OpenDalStorage::Gcs { config: Arc::new(GcsConfig::default()), + credential_provider: None, }; - assert_eq!( - storage - .relativize_path("gs://my-bucket/path/to/file.parquet") - .unwrap(), - "path/to/file.parquet" - ); + for scheme in ["gs", "gcs"] { + let path = format!("{scheme}://my-bucket/path/to/file.parquet"); + assert_eq!( + storage.relativize_path(&path).unwrap(), + "path/to/file.parquet" + ); + let (_, relative) = storage.create_operator(&path).unwrap(); + assert_eq!(relative, "path/to/file.parquet"); + } } #[cfg(feature = "opendal-gcs")] @@ -831,6 +1160,7 @@ mod tests { fn test_relativize_path_gcs_invalid_scheme() { let storage = OpenDalStorage::Gcs { config: Arc::new(GcsConfig::default()), + credential_provider: None, }; assert!( @@ -838,6 +1168,11 @@ mod tests { .relativize_path("s3://my-bucket/path/to/file.parquet") .is_err() ); + assert!( + storage + .create_operator(&"s3://my-bucket/path/to/file.parquet") + .is_err() + ); } #[cfg(feature = "opendal-oss")] @@ -878,6 +1213,8 @@ mod tests { endpoint: Some("https://myaccount.dfs.core.windows.net".to_string()), ..Default::default() }), + sas_tokens: Arc::new(AzdlsSasTokens::default()), + credential_provider: None, }; assert_eq!( diff --git a/crates/storage/opendal/src/resolving.rs b/crates/storage/opendal/src/resolving.rs index 3b99b08fc0..29a92ad9fa 100644 --- a/crates/storage/opendal/src/resolving.rs +++ b/crates/storage/opendal/src/resolving.rs @@ -27,7 +27,7 @@ use futures::StreamExt; use futures::stream::BoxStream; use iceberg::io::{ FileMetadata, FileRead, FileWrite, InputFile, OutputFile, Storage, StorageConfig, - StorageFactory, + StorageCredentialProvider, StorageFactory, }; use iceberg::{Error, ErrorKind, Result}; use serde::{Deserialize, Serialize}; @@ -81,10 +81,17 @@ fn extract_scheme(path: &str) -> Result<&'static str> { } /// Build an [`OpenDalStorage`] variant for the given scheme and config properties. +#[allow(unused_variables)] fn build_storage_for_scheme( scheme: &'static str, props: &HashMap, #[cfg(feature = "opendal-s3")] customized_credential_load: &Option, + #[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + ))] + credential_provider: &Option>, ) -> Result { match scheme { #[cfg(feature = "opendal-s3")] @@ -93,6 +100,7 @@ fn build_storage_for_scheme( Ok(OpenDalStorage::S3 { config: Arc::new(config), customized_credential_load: customized_credential_load.clone(), + credential_provider: credential_provider.clone(), }) } #[cfg(feature = "opendal-gcs")] @@ -100,6 +108,7 @@ fn build_storage_for_scheme( let config = crate::gcs::gcs_config_parse(props.clone())?; Ok(OpenDalStorage::Gcs { config: Arc::new(config), + credential_provider: credential_provider.clone(), }) } #[cfg(feature = "opendal-oss")] @@ -114,6 +123,8 @@ fn build_storage_for_scheme( let config = crate::azdls::azdls_config_parse(props.clone())?; Ok(OpenDalStorage::Azdls { config: Arc::new(config), + sas_tokens: Arc::new(crate::azdls::AzdlsSasTokens::from_properties(props)), + credential_provider: credential_provider.clone(), }) } #[cfg(feature = "opendal-fs")] @@ -196,11 +207,28 @@ impl OpenDalResolvingStorageFactory { #[typetag::serde] impl StorageFactory for OpenDalResolvingStorageFactory { fn build(&self, config: &StorageConfig) -> Result> { + self.build_with_credentials(config, None) + } + + #[allow(unused_variables)] + fn build_with_credentials( + &self, + config: &StorageConfig, + credential_provider: Option>, + ) -> Result> { + // Without a compatible backend the provider is ignored, and every + // backend uses the credentials in `config`. Ok(Arc::new(OpenDalResolvingStorage { props: config.props().clone(), storages: RwLock::new(HashMap::new()), #[cfg(feature = "opendal-s3")] customized_credential_load: self.customized_credential_load.clone(), + #[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + ))] + credential_provider, })) } } @@ -211,7 +239,7 @@ impl StorageFactory for OpenDalResolvingStorageFactory { /// Sub-storages are lazily created on first use for each scheme and cached /// for subsequent operations. Scheme aliases like `s3`/`s3a`/`s3n` map to /// the same canonical scheme, so they share a storage instance. -#[derive(Debug, Serialize, Deserialize)] +#[derive(Serialize, Deserialize)] pub struct OpenDalResolvingStorage { /// Configuration properties shared across all backends. props: HashMap, @@ -222,6 +250,24 @@ pub struct OpenDalResolvingStorage { #[cfg(feature = "opendal-s3")] #[serde(skip)] customized_credential_load: Option, + /// Provider of refreshable vended credentials. + #[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + ))] + #[serde(skip)] + credential_provider: Option>, +} + +impl std::fmt::Debug for OpenDalResolvingStorage { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + // `props` can contain storage secrets, and a custom credential provider + // may carry secret state in its own Debug implementation + f.debug_struct("OpenDalResolvingStorage") + .field("property_keys", &self.props.keys().collect::>()) + .finish_non_exhaustive() + } } impl OpenDalResolvingStorage { @@ -257,6 +303,12 @@ impl OpenDalResolvingStorage { &self.props, #[cfg(feature = "opendal-s3")] &self.customized_credential_load, + #[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + ))] + &self.credential_provider, )?; let storage = Arc::new(storage); cache.insert(scheme, storage.clone()); @@ -317,6 +369,7 @@ impl Storage for OpenDalResolvingStorage { Ok(()) } + #[allow(unreachable_code)] fn new_input(&self, path: &str) -> Result { Ok(InputFile::new( Arc::new(self.resolve(path)?.as_ref().clone()), @@ -324,6 +377,7 @@ impl Storage for OpenDalResolvingStorage { )) } + #[allow(unreachable_code)] fn new_output(&self, path: &str) -> Result { Ok(OutputFile::new( Arc::new(self.resolve(path)?.as_ref().clone()), @@ -334,8 +388,52 @@ impl Storage for OpenDalResolvingStorage { #[cfg(test)] mod tests { + #[allow(unused_imports)] use super::*; + #[cfg(any( + feature = "opendal-azdls", + not(any(feature = "opendal-s3", feature = "opendal-gcs")), + all( + feature = "opendal-memory", + any(feature = "opendal-s3", feature = "opendal-gcs") + ) + ))] + #[derive(Debug)] + struct AllPathsCredentialProvider; + + #[cfg(any( + feature = "opendal-azdls", + not(any(feature = "opendal-s3", feature = "opendal-gcs")), + all( + feature = "opendal-memory", + any(feature = "opendal-s3", feature = "opendal-gcs") + ) + ))] + #[async_trait] + impl StorageCredentialProvider for AllPathsCredentialProvider { + async fn load_credential(&self, _path: &str) -> Result { + unreachable!("unsupported backends must ignore the provider") + } + } + + #[cfg(not(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + )))] + #[test] + fn test_factory_ignores_credentials_without_compatible_backend() { + assert!( + OpenDalResolvingStorageFactory::new() + .build_with_credentials( + &StorageConfig::new(), + Some(Arc::new(AllPathsCredentialProvider)), + ) + .is_ok() + ); + } + #[cfg(feature = "opendal-s3")] #[derive(Debug)] struct EmptyCredentialLoader; @@ -368,12 +466,23 @@ mod tests { /// Builds a resolving storage with empty props, suitable for `resolve()` /// calls that don't actually hit any backend. + #[cfg(any( + feature = "opendal-s3", + feature = "opendal-azdls", + all(feature = "opendal-memory", feature = "opendal-gcs") + ))] fn empty_resolving_storage() -> OpenDalResolvingStorage { OpenDalResolvingStorage { props: HashMap::new(), storages: RwLock::new(HashMap::new()), #[cfg(feature = "opendal-s3")] customized_credential_load: None, + #[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + ))] + credential_provider: None, } } @@ -393,10 +502,27 @@ mod tests { assert!(Arc::ptr_eq(&a, &c), "s3 and s3n should share one instance"); } + #[cfg(all( + feature = "opendal-memory", + any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + ) + ))] + #[test] + fn test_resolver_ignores_credentials_for_unsupported_backend() { + let mut storage = empty_resolving_storage(); + storage.credential_provider = Some(Arc::new(AllPathsCredentialProvider)); + + assert!(storage.resolve("memory:/key").is_ok()); + } + #[cfg(feature = "opendal-azdls")] #[test] fn test_resolve_azdls_aliases_share_instance() { - let storage = empty_resolving_storage(); + let mut storage = empty_resolving_storage(); + storage.credential_provider = Some(Arc::new(AllPathsCredentialProvider)); let path_for = |scheme: &str| { format!("{scheme}://myfs@myaccount.dfs.core.windows.net/path/to/file.parquet") diff --git a/crates/storage/opendal/src/s3.rs b/crates/storage/opendal/src/s3.rs index 2997fa81d9..ebe0518221 100644 --- a/crates/storage/opendal/src/s3.rs +++ b/crates/storage/opendal/src/s3.rs @@ -22,7 +22,8 @@ use iceberg::io::{ CLIENT_REGION, S3_ACCESS_KEY_ID, S3_ALLOW_ANONYMOUS, S3_ASSUME_ROLE_ARN, S3_ASSUME_ROLE_EXTERNAL_ID, S3_ASSUME_ROLE_SESSION_NAME, S3_DISABLE_CONFIG_LOAD, S3_DISABLE_EC2_METADATA, S3_ENDPOINT, S3_PATH_STYLE_ACCESS, S3_REGION, S3_SECRET_ACCESS_KEY, - S3_SESSION_TOKEN, S3_SSE_KEY, S3_SSE_MD5, S3_SSE_TYPE, + S3_SESSION_TOKEN, S3_SSE_KEY, S3_SSE_MD5, S3_SSE_TYPE, StorageCredentialKind, + StorageCredentialProvider, }; use iceberg::{Error, ErrorKind, Result}; use opendal::services::S3Config; @@ -31,10 +32,13 @@ use opendal::{Configurator, Operator}; pub use reqsign_aws_v4::Credential as AwsCredential; /// Trait for types that can asynchronously supply [`AwsCredential`] to a [`CustomAwsCredentialLoader`]. pub use reqsign_core::ProvideCredential; -use reqsign_core::{ProvideCredentialChain, ProvideCredentialDyn}; +use reqsign_core::{ + Context, Error as ReqsignError, ProvideCredentialChain, ProvideCredentialDyn, + Result as ReqsignResult, +}; use url::Url; -use crate::utils::{from_opendal_error, is_truthy}; +use crate::utils::{VendedCredentialSource, from_opendal_error, is_truthy}; /// Parse iceberg props to s3 config. pub(crate) fn s3_config_parse(mut m: HashMap) -> Result { @@ -129,7 +133,9 @@ pub(crate) fn s3_config_parse(mut m: HashMap) -> Result, + credential_provider: &Option>, path: &str, + credential_location: Option<&str>, ) -> Result { let url = Url::parse(path)?; let bucket = url.host_str().ok_or_else(|| { @@ -139,6 +145,19 @@ pub(crate) fn s3_config_build( ) })?; + // Preserve the existing custom-loader precedence: an explicitly configured loader + // is the sole source, otherwise install the catalog provider as a replacement chain so + // refresh failures cannot fall through to broader ambient AWS credentials. + let credential_provider = credential_provider + .as_ref() + .filter(|provider| provider.supports_path(path)); + if customized_credential_load.is_none() && credential_provider.is_some() && cfg.skip_signature { + return Err(Error::new( + ErrorKind::DataInvalid, + "Invalid S3 auth settings: anonymous access cannot be combined with refreshable credentials", + )); + } + let mut builder = cfg .clone() .into_builder() @@ -148,11 +167,46 @@ pub(crate) fn s3_config_build( if let Some(loader) = customized_credential_load { let chain = ProvideCredentialChain::new().push(Arc::clone(&loader.0)); builder = builder.credential_provider_chain(chain); + } else if let Some(provider) = credential_provider { + let chain = ProvideCredentialChain::new().push(VendedS3CredentialProvider( + VendedCredentialSource::new( + Arc::clone(provider), + credential_location.unwrap_or(path).to_string(), + ), + )); + builder = builder.credential_provider_chain(chain); } Operator::new(builder).map_err(from_opendal_error) } +/// Adapts a generic [`StorageCredentialProvider`] into a reqsign +/// [`ProvideCredential`], so the S3 signer can obtain and refresh vended +/// credentials. +#[derive(Debug)] +struct VendedS3CredentialProvider(VendedCredentialSource); + +impl ProvideCredential for VendedS3CredentialProvider { + type Credential = AwsCredential; + + async fn provide_credential(&self, _ctx: &Context) -> ReqsignResult> { + match self.0.load("S3").await? { + (StorageCredentialKind::S3(s3), expires_in) => { + let (access_key_id, secret_access_key, session_token) = s3.into_parts(); + Ok(Some(AwsCredential { + access_key_id, + secret_access_key, + session_token, + expires_in, + })) + } + _ => Err(ReqsignError::unexpected( + "S3 storage received a non-S3 credential from the provider", + )), + } + } +} + /// Custom AWS credential loader. /// /// Wraps any [`ProvideCredential`] implementation for use with the S3 storage backend. @@ -176,7 +230,7 @@ impl std::fmt::Debug for CustomAwsCredentialLoader { impl CustomAwsCredentialLoader { /// Create a new custom AWS credential loader from any [`ProvideCredential`] implementation. pub fn new(provider: impl ProvideCredential + 'static) -> Self { - Self(Arc::new(provider) as Arc>) + Self(Arc::new(provider)) } } diff --git a/crates/storage/opendal/src/utils.rs b/crates/storage/opendal/src/utils.rs index 56f8c18059..d42170c71b 100644 --- a/crates/storage/opendal/src/utils.rs +++ b/crates/storage/opendal/src/utils.rs @@ -15,6 +15,7 @@ // specific language governing permissions and limitations // under the License. +#[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] pub(crate) fn is_truthy(value: &str) -> bool { ["true", "t", "1", "on"].contains(&value.to_lowercase().as_str()) } @@ -27,3 +28,203 @@ pub(crate) fn from_opendal_error(e: opendal::Error) -> iceberg::Error { ) .with_source(e) } + +/// Convert a [`SystemTime`](std::time::SystemTime) credential expiry into the +/// `reqsign` [`Timestamp`](reqsign_core::time::Timestamp) used on backend +/// credential types (e.g. `AwsCredential::expires_in`, `google::Token::expires_at`). +#[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" +))] +pub(crate) fn system_time_to_timestamp( + time: std::time::SystemTime, +) -> reqsign_core::Result { + let millis = time + .duration_since(std::time::UNIX_EPOCH) + .map_err(|e| { + reqsign_core::Error::unexpected(format!( + "credential expiry precedes the UNIX epoch: {e}" + )) + })? + .as_millis(); + let millis = i64::try_from(millis).map_err(|_| { + reqsign_core::Error::unexpected("credential expiry overflows i64 milliseconds") + })?; + reqsign_core::time::Timestamp::from_millisecond(millis) + .map_err(|e| reqsign_core::Error::unexpected(format!("invalid credential expiry: {e}"))) +} + +/// Validate that a provider's credential covers the location for which the +/// backend requested it. +#[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" +))] +pub(crate) fn validate_credential_prefix( + location: &str, + credential: &iceberg::io::StorageCredential, +) -> reqsign_core::Result<()> { + if credential.covers(location) { + Ok(()) + } else { + Err(reqsign_core::Error::unexpected(format!( + "vended credential prefix {:?} does not cover storage location {location:?}", + credential.prefix() + ))) + } +} + +/// The root of the storage location containing `path`, e.g. `s3://bucket/`. +pub(crate) fn storage_root(path: &str) -> iceberg::Result { + let url = url::Url::parse(path)?; + Ok(format!("{}://{}/", url.scheme(), url.authority())) +} + +/// Loads vended credentials for one storage location on behalf of a backend's +/// `reqsign` credential provider. +#[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" +))] +pub(crate) struct VendedCredentialSource { + provider: std::sync::Arc, + /// Location handed to the provider: the path an operator serves, or the + /// scope shared by a bulk-delete batch. + location: String, +} + +#[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" +))] +impl VendedCredentialSource { + pub(crate) fn new( + provider: std::sync::Arc, + location: String, + ) -> Self { + Self { provider, location } + } + + /// Load a credential covering the location, returning its backend-specific + /// material and expiry. + pub(crate) async fn load( + &self, + backend: &str, + ) -> reqsign_core::Result<( + iceberg::io::StorageCredentialKind, + Option, + )> { + let credential = self + .provider + .load_credential(&self.location) + .await + .map_err(|e| { + reqsign_core::Error::unexpected(format!( + "failed to load vended {backend} credential for {}", + self.location + )) + .with_source(e) + })?; + validate_credential_prefix(&self.location, &credential)?; + let expires_at = credential + .expires_at() + .map(system_time_to_timestamp) + .transpose()?; + Ok((credential.into_kind(), expires_at)) + } +} + +#[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" +))] +impl std::fmt::Debug for VendedCredentialSource { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("VendedCredentialSource") + .field("location", &self.location) + .finish_non_exhaustive() + } +} + +#[cfg(all( + test, + any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + ) +))] +mod tests { + use std::sync::{Arc, Mutex}; + + use async_trait::async_trait; + use iceberg::io::{ + GcsCredential, StorageCredential, StorageCredentialKind, StorageCredentialProvider, + }; + + use super::*; + + /// Returns a credential scoped to `prefix` and records requested locations. + #[derive(Debug)] + struct RecordingProvider { + prefix: Option<&'static str>, + requested: Mutex>, + } + + #[async_trait] + impl StorageCredentialProvider for RecordingProvider { + async fn load_credential(&self, path: &str) -> iceberg::Result { + self.requested.lock().unwrap().push(path.to_string()); + let credential = + StorageCredential::new(StorageCredentialKind::Gcs(GcsCredential::new("token"))); + Ok(match self.prefix { + Some(prefix) => credential.with_prefix(prefix), + None => credential, + }) + } + } + + async fn load(prefix: Option<&'static str>, location: &str) -> reqsign_core::Result<()> { + let provider = Arc::new(RecordingProvider { + prefix, + requested: Mutex::new(Vec::new()), + }); + let source = VendedCredentialSource::new(provider.clone(), location.to_string()); + let result = source.load("GCS").await.map(|_| ()); + assert_eq!(*provider.requested.lock().unwrap(), vec![location]); + result + } + + #[tokio::test] + async fn vended_source_requires_a_covering_credential() { + let location = "gs://bucket/table/data/file.parquet"; + assert!(load(None, location).await.is_ok()); + assert!(load(Some("gs://bucket/table"), location).await.is_ok()); + assert!(load(Some("gcs://bucket"), location).await.is_ok()); + assert!( + load(Some("gs://bucket/table/data/other"), location) + .await + .is_err() + ); + assert!(load(Some("gs://bucket/tab"), location).await.is_err()); + assert!(load(Some(""), location).await.is_err()); + } + + #[test] + fn storage_root_keeps_scheme_and_authority() { + assert_eq!( + storage_root("s3a://bucket/table/file.parquet").unwrap(), + "s3a://bucket/" + ); + assert_eq!( + storage_root("abfss://fs@account.dfs.core.windows.net/table/file.parquet").unwrap(), + "abfss://fs@account.dfs.core.windows.net/" + ); + assert!(storage_root("not a url").is_err()); + } +}