diff --git a/src/Common/ProfileEvents.cpp b/src/Common/ProfileEvents.cpp index 7e7d268a88ad..124c4767b0e6 100644 --- a/src/Common/ProfileEvents.cpp +++ b/src/Common/ProfileEvents.cpp @@ -1697,6 +1697,9 @@ The server successfully detected this situation and will download merged part fr M(AIRowsProcessed, "Number of rows that received an AI result.", ValueType::Number) \ M(AIRowsSkipped, "Number of rows that received a default value due to quota or error.", ValueType::Number) \ \ + M(DataLakeRestCatalogCredentialsVended, "Number of table metadata requests to REST catalog asking to vend storage credentials.", ValueType::Number) \ + M(DataLakeRestCatalogCredentialsCacheHits, "Number of table metadata requests to REST catalog reusing cached storage credentials.", ValueType::Number) \ + \ M(StatelessWorkerRequested, "Number of stateless workers requested by queries for distributed query execution.", ValueType::Number) \ M(StatelessWorkerProvided, "Number of stateless workers provided to queries for distributed query execution.", ValueType::Number) \ M(StatelessWorkerProvisioningMicroseconds, "Total time queries spent waiting for stateless workers to be provisioned.", ValueType::Microseconds) \ diff --git a/src/Databases/DataLake/DatabaseDataLake.cpp b/src/Databases/DataLake/DatabaseDataLake.cpp index d0c25a2cb838..77330de7440c 100644 --- a/src/Databases/DataLake/DatabaseDataLake.cpp +++ b/src/Databases/DataLake/DatabaseDataLake.cpp @@ -2,6 +2,7 @@ #include #include +#include #include #include #include @@ -72,6 +73,7 @@ namespace DatabaseDataLakeSetting extern const DatabaseDataLakeSettingsString oauth_server_uri; extern const DatabaseDataLakeSettingsBool oauth_server_use_request_body; extern const DatabaseDataLakeSettingsBool vended_credentials; + extern const DatabaseDataLakeSettingsUInt64 vended_credentials_cache_ttl; extern const DatabaseDataLakeSettingsString object_storage_cluster; extern const DatabaseDataLakeSettingsString aws_access_key_id; extern const DatabaseDataLakeSettingsString aws_secret_access_key; @@ -468,6 +470,11 @@ std::shared_ptr DatabaseDataLake::getCatalog() const #endif catalog_unavailable_reason); } + + const auto settings_version = database_settings.get(); + catalog_impl->setVendedCredentialsCacheTTL( + std::chrono::seconds((*settings_version)[DatabaseDataLakeSetting::vended_credentials_cache_ttl].value)); + return catalog_impl; } @@ -1744,6 +1751,7 @@ The following settings are supported: | `storage_endpoint` | Endpoint URL for the underlying storage | | `oauth_server_uri` | URI of the OAuth2 authorization server for authentication | | `vended_credentials` | Boolean indicating whether to use vended credentials from the catalog (supports AWS S3 and Azure ADLS Gen2) | +| `vended_credentials_cache_ttl` | Maximum cache entry lifetime (in seconds) for vended credentials (REST catalogs only). Default `300`; `0` disables caching. | | `aws_access_key_id` | AWS access key ID for S3/Glue access (if not using vended credentials) | | `aws_secret_access_key` | AWS secret access key for S3/Glue access (if not using vended credentials) | | `aws_role_arn` | ARN of the IAM role to assume for AWS/Glue access. When set, ClickHouse uses AWS STS `AssumeRole` with base credentials from `aws_access_key_id` and `aws_secret_access_key` when both are provided, or from the default AWS credential chain otherwise (the role must trust the identity the server runs under). | diff --git a/src/Databases/DataLake/DatabaseDataLakeSettings.cpp b/src/Databases/DataLake/DatabaseDataLakeSettings.cpp index 866d124fbe0f..5cdc07b7e6d9 100644 --- a/src/Databases/DataLake/DatabaseDataLakeSettings.cpp +++ b/src/Databases/DataLake/DatabaseDataLakeSettings.cpp @@ -21,6 +21,7 @@ namespace ErrorCodes DECLARE(DatabaseDataLakeCatalogType, catalog_type, DatabaseDataLakeCatalogType::NONE, "Catalog type", 0) \ DECLARE(String, catalog_credential, "", "", 0) \ DECLARE(Bool, vended_credentials, true, "Use vended credentials (storage credentials) from catalog", 0) \ + DECLARE(UInt64, vended_credentials_cache_ttl, 300, "Maximum cache entry lifetime (in seconds) for vended credentials. '0' disables caching.", 0) \ DECLARE(String, auth_scope, "PRINCIPAL_ROLE:ALL", "Authorization scope for client credentials or token exchange", 0) \ DECLARE(String, oauth_server_uri, "", "OAuth server uri", 0) \ DECLARE(Bool, oauth_server_use_request_body, true, "Put parameters into request body or query params", 0) \ diff --git a/src/Databases/DataLake/ICatalog.h b/src/Databases/DataLake/ICatalog.h index 6aa932f26a3f..3e5d4050f6e3 100644 --- a/src/Databases/DataLake/ICatalog.h +++ b/src/Databases/DataLake/ICatalog.h @@ -1,4 +1,5 @@ #pragma once +#include #include #include #include @@ -292,6 +293,8 @@ class ICatalog return std::nullopt; } + virtual void setVendedCredentialsCacheTTL(std::chrono::seconds /*ttl*/) {} + /// Result of `prepareSettingsChanges`: the new catalog state built off to the side, /// ready to be published by `commitSettingsChanges`. struct PreparedSettingsChanges diff --git a/src/Databases/DataLake/RestCatalog.cpp b/src/Databases/DataLake/RestCatalog.cpp index ccd48047d39e..6c69ab9c95ce 100644 --- a/src/Databases/DataLake/RestCatalog.cpp +++ b/src/Databases/DataLake/RestCatalog.cpp @@ -54,6 +54,11 @@ #include #include #include +#include +#include +#include +#include +#include #include #include @@ -82,6 +87,8 @@ namespace DB::FailPoints namespace ProfileEvents { + extern const Event DataLakeRestCatalogCredentialsVended; + extern const Event DataLakeRestCatalogCredentialsCacheHits; extern const Event OneLakeAccessTokenRequests; extern const Event OneLakeAccessTokenRequestFailures; extern const Event OneLakeAccessTokenRequestMicroseconds; @@ -133,6 +140,13 @@ static constexpr auto UNKNOWN_EXPIRATION_TOKEN_LIFETIME = std::chrono::minutes(1 namespace { +String parseTableUuid(const Poco::JSON::Object::Ptr & metadata_object) +{ + if (metadata_object && metadata_object->has("table-uuid")) + return metadata_object->get("table-uuid").extract(); + return {}; +} + std::pair parseCatalogCredential(const std::string & catalog_credential) { /// Parse a string of format ":" @@ -508,6 +522,19 @@ void RestCatalog::commitSettingsChanges(ICatalog::PreparedSettingsChangesPtr pre state.set(std::move(prepared_auth->new_state)); if (prepared_auth->new_access_token) access_token.set(std::move(prepared_auth->new_access_token)); + + std::lock_guard lock(credentials_cache_mutex); + credentials_cache.clear(); +} + +void RestCatalog::setVendedCredentialsCacheTTL(std::chrono::seconds ttl) +{ + std::lock_guard lock(credentials_cache_mutex); + if (ttl != vended_credentials_cache_ttl) + { + vended_credentials_cache_ttl = ttl; + credentials_cache.clear(); + } } void RestCatalog::applySettingsChangesToState( @@ -1765,11 +1792,92 @@ void RestCatalog::getTableMetadata( throw DB::Exception(DB::ErrorCodes::DATALAKE_DATABASE_ERROR, "No response from iceberg catalog"); } +namespace +{ + +/// Overlay `config` with the `storage-credentials` entry whose prefix best matches the table location. +Poco::JSON::Object::Ptr effectiveVendedConfig(const Poco::JSON::Object::Ptr & load_table_result, const std::string & location) +{ + static constexpr auto storage_credentials_str = "storage-credentials"; + + Poco::JSON::Object::Ptr config_object; + if (load_table_result->has("config")) + { + config_object = load_table_result->getObject("config"); + if (!config_object) + throw DB::Exception(DB::ErrorCodes::DATALAKE_DATABASE_ERROR, "Cannot parse config result"); + } + + const auto entries + = load_table_result->isArray(storage_credentials_str) ? load_table_result->getArray(storage_credentials_str) : nullptr; + + Poco::JSON::Object::Ptr best_config; + size_t best_prefix_size = 0; + for (size_t i = 0; entries && i < entries->size(); ++i) + { + const auto entry = entries->getObject(static_cast(i)); + if (!entry) + continue; + const auto prefix_var = entry->get("prefix"); + if (!prefix_var.isString()) + continue; + const auto & prefix = prefix_var.extract(); + if (!location.starts_with(prefix)) + continue; + const auto entry_config = entry->getObject("config"); + if (!entry_config) + continue; + if (!best_config || prefix.size() > best_prefix_size) + { + best_config = entry_config; + best_prefix_size = prefix.size(); + } + } + + if (!best_config) + return config_object; + if (!config_object) + config_object = new Poco::JSON::Object(); + + Poco::JSON::Object::Ptr merged = new Poco::JSON::Object(*config_object); + std::vector names; + best_config->getNames(names); + + /// An entry that supplies any key of a credential group replaces that whole group. + static const std::vector> credential_groups = { + {"s3.access-key-id", "s3.secret-access-key", "s3.session-token", "s3.session-token-expires-at-ms"}, + {"gcs.oauth2.token", "gcs.oauth2.token-expires-at"}, + }; + for (const auto & group : credential_groups) + if (std::any_of(group.begin(), group.end(), [&](const auto & key) { return best_config->has(key); })) + for (const auto & key : group) + merged->remove(key); + + /// Azure SAS tokens form one group across all `adls.sas-token.*` keys. + static constexpr auto sas_prefix = "adls.sas-token."; + if (std::any_of(names.begin(), names.end(), [](const auto & name) { return name.starts_with(sas_prefix); })) + { + std::vector base_names; + merged->getNames(base_names); + for (const auto & name : base_names) + if (name.starts_with(sas_prefix)) + merged->remove(name); + } + + for (const auto & name : names) + merged->set(name, best_config->get(name)); + + return merged; +} + +} + bool RestCatalog::getTableMetadataImpl( const std::string & namespace_name, const std::string & table_name, DB::ContextPtr context_, - TableMetadata & result) const + TableMetadata & result, + bool allow_credentials_cache) const { LOG_DEBUG(log, "Checking table {} in namespace {}", table_name, namespace_name); @@ -1778,15 +1886,26 @@ bool RestCatalog::getTableMetadataImpl( "Namespace {} is filtered by `namespaces` database parameter", namespace_name); DB::HTTPHeaderEntries headers; - if (result.requiresCredentials()) + + const bool want_credentials = result.requiresCredentials(); + + std::optional cached_credentials; + if (want_credentials) { + if (allow_credentials_cache) + cached_credentials = tryGetCachedCredentials(namespace_name, table_name); + /// Header `X-Iceberg-Access-Delegation` tells catalog to include storage credentials in LoadTableResponse. /// Value can be one of the two: /// 1. `vended-credentials` /// 2. `remote-signing` /// Currently we support only the first. /// https://github.com/apache/iceberg/blob/3badfe0c1fcf0c0adfc7aa4a10f0b50365c48cf9/open-api/rest-catalog-open-api.yaml#L1832 - headers.emplace_back("X-Iceberg-Access-Delegation", "vended-credentials"); + if (!cached_credentials) + { + ProfileEvents::increment(ProfileEvents::DataLakeRestCatalogCredentialsVended); + headers.emplace_back("X-Iceberg-Access-Delegation", "vended-credentials"); + } } const auto state_snapshot = state.get(); @@ -1822,20 +1941,18 @@ bool RestCatalog::getTableMetadataImpl( if (!metadata_object) throw DB::Exception(DB::ErrorCodes::LOGICAL_ERROR, "Cannot parse result"); + const std::string table_uuid = parseTableUuid(metadata_object); + std::string location; - if (result.requiresLocation()) + if (metadata_object->has("location")) { - if (metadata_object->has("location")) - { - location = metadata_object->get("location").extract(); + location = metadata_object->get("location").extract(); + if (result.requiresLocation()) result.setLocation(location); - LOG_DEBUG(log, "Location for table {}: {}", table_name, location); - } - else - { - result.setTableIsNotReadable(fmt::format("Cannot read table {}, because no 'location' in response", table_name)); - } + LOG_DEBUG(log, "Location for table {}: {}", table_name, location); } + else if (result.requiresLocation()) + result.setTableIsNotReadable(fmt::format("Cannot read table {}, because no 'location' in response", table_name)); if (result.requiresSchema()) { @@ -1847,16 +1964,36 @@ bool RestCatalog::getTableMetadataImpl( result.setSchema(*schema); } - if (result.isDefaultReadableTable() && result.requiresCredentials() && object->has("config")) + if (want_credentials && result.isDefaultReadableTable()) { - auto config_object = object->get("config").extract(); - if (!config_object) - throw DB::Exception(DB::ErrorCodes::LOGICAL_ERROR, "Cannot parse config result"); - auto [parsed_credentials, parsed_endpoint] = getCredentialsAndEndpoint(config_object, location); - if (parsed_credentials) - result.setStorageCredentials(parsed_credentials); - if (!parsed_endpoint.empty()) - result.setEndpoint(parsed_endpoint); + if (cached_credentials) + { + if (table_uuid.empty() || cached_credentials->table_uuid != table_uuid || cached_credentials->location != location) + { + { + std::lock_guard lock(credentials_cache_mutex); + credentials_cache.erase({namespace_name, table_name}); + } + return getTableMetadataImpl(namespace_name, table_name, context_, result, /* allow_credentials_cache */ false); + } + ProfileEvents::increment(ProfileEvents::DataLakeRestCatalogCredentialsCacheHits); + result.setStorageCredentials(cached_credentials->credentials); + if (!cached_credentials->endpoint.empty()) + result.setEndpoint(cached_credentials->endpoint); + } + else if (const auto config_object = effectiveVendedConfig(object, location)) + { + auto parsed = getCredentialsAndEndpoint(config_object, location); + parsed.table_uuid = table_uuid; + parsed.location = location; + if (parsed.credentials) + { + result.setStorageCredentials(parsed.credentials); + cacheCredentials(namespace_name, table_name, parsed, state_snapshot); + } + if (!parsed.endpoint.empty()) + result.setEndpoint(parsed.endpoint); + } } if (result.requiresDataLakeSpecificProperties()) @@ -1868,8 +2005,8 @@ bool RestCatalog::getTableMetadataImpl( } } - if (metadata_object->has("table-uuid")) - result.setTableUUID(metadata_object->get("table-uuid").extract()); + if (!table_uuid.empty()) + result.setTableUUID(table_uuid); return true; } @@ -2218,7 +2355,53 @@ void RestCatalog::dropTable(const String & namespace_name, const String & table_ } } -std::pair, String> RestCatalog::getCredentialsAndEndpoint(Poco::JSON::Object::Ptr object, const String & location) const +namespace +{ +std::optional +parseExpiresAtMs(const Poco::JSON::Object::Ptr & object, const std::string & key) +{ + if (!object->has(key)) + return std::nullopt; + + static constexpr Int64 max_representable_sec + = std::chrono::duration_cast(std::chrono::system_clock::duration::max()).count(); + const Int64 expires_at_ms = object->get(key).convert(); + if (expires_at_ms <= 0) + return std::chrono::system_clock::time_point{}; + if (expires_at_ms / 1000 < max_representable_sec) + return std::chrono::system_clock::from_time_t(static_cast(expires_at_ms / 1000)); + return std::chrono::system_clock::time_point::max(); +} + +std::chrono::system_clock::time_point parseSasTokenExpiry(const std::string & sas_token) +{ + std::string token = sas_token; + if (!token.empty() && token.front() == '?') + token.erase(0, 1); + + Poco::StringTokenizer params(token, "&", Poco::StringTokenizer::TOK_IGNORE_EMPTY | Poco::StringTokenizer::TOK_TRIM); + for (const auto & param : params) + { + if (!param.starts_with("se=")) + continue; + + std::string decoded; + Poco::URI::decode(param.substr(3), decoded); + + int time_zone_differential = 0; + Poco::DateTime date_time; + if (!Poco::DateTimeParser::tryParse(Poco::DateTimeFormat::ISO8601_FORMAT, decoded, date_time, time_zone_differential)) + throw DB::Exception(DB::ErrorCodes::DATALAKE_DATABASE_ERROR, "Cannot parse Azure SAS token expiration `{}`", decoded); + + date_time.makeUTC(time_zone_differential); + return std::chrono::system_clock::from_time_t(date_time.timestamp().epochTime()); + } + + return std::chrono::system_clock::time_point{}; +} +} + +VendedStorageCredentials RestCatalog::getCredentialsAndEndpoint(Poco::JSON::Object::Ptr object, const String & location) const { auto storage_type = parseStorageTypeFromLocation(location); switch (storage_type) @@ -2226,22 +2409,27 @@ std::pair, String> RestCatalog::getCredenti case StorageType::S3: { static constexpr auto gcs_token_str = "gcs.oauth2.token"; + static constexpr auto gcs_token_expires_at_str = "gcs.oauth2.token-expires-at"; static constexpr auto access_key_id_str = "s3.access-key-id"; static constexpr auto secret_access_key_str = "s3.secret-access-key"; static constexpr auto session_token_str = "s3.session-token"; static constexpr auto storage_endpoint_str = "s3.endpoint"; + static constexpr auto session_token_expires_at_ms_str = "s3.session-token-expires-at-ms"; - if (object->has(gcs_token_str)) + /// `gs://` also maps to `StorageType::S3`; the merged config may contain both GCS and S3 credentials. + if (location.starts_with("gs://") && object->has(gcs_token_str)) { auto gcs_token = object->get(gcs_token_str).extract(); LOG_DEBUG(log, "Using GCS OAuth2 token for location {}", location); - return {std::make_shared(gcs_token), ""}; + auto expires_at = parseExpiresAtMs(object, gcs_token_expires_at_str).value_or(std::chrono::system_clock::time_point{}); + return {std::make_shared(gcs_token), "", expires_at}; } std::string access_key_id; std::string secret_access_key; std::string session_token; std::string storage_endpoint; + std::optional expires_at; if (object->has(access_key_id_str)) access_key_id = object->get(access_key_id_str).extract(); if (object->has(secret_access_key_str)) @@ -2250,39 +2438,96 @@ std::pair, String> RestCatalog::getCredenti session_token = object->get(session_token_str).extract(); if (object->has(storage_endpoint_str)) storage_endpoint = object->get(storage_endpoint_str).extract(); + expires_at = parseExpiresAtMs(object, session_token_expires_at_ms_str); + if (!expires_at.has_value() && !session_token.empty()) + expires_at = std::chrono::system_clock::time_point{}; LOG_DEBUG(log, "get tokens for location {}", location); - return {std::make_shared(access_key_id, secret_access_key, session_token), storage_endpoint}; + return {std::make_shared(access_key_id, secret_access_key, session_token), storage_endpoint, expires_at}; } case StorageType::Azure: { /// Azure ADLS Gen2 vended credentials use SAS tokens. /// The config keys follow the pattern: adls.sas-token. /// or adls.sas-token..dfs.core.windows.net - /// We look for any key starting with "adls.sas-token." and use the first one found. + /// Select the key matching the account of `location` (abfss://container@account.dfs.core.windows.net/path). String sas_token; - std::vector names; - object->getNames(names); - for (const auto & name : names) + const auto authority_start = location.find("://") + std::strlen("://"); + const auto authority = location.substr(authority_start, location.find('/', authority_start) - authority_start); + const auto at_pos = authority.find('@'); + if (at_pos != String::npos) { - if (name.starts_with("adls.sas-token.")) + const auto account_host = authority.substr(at_pos + 1); + const auto account_name = account_host.substr(0, account_host.find('.')); + for (const auto & key : {"adls.sas-token." + account_host, "adls.sas-token." + account_name}) { - sas_token = object->get(name).extract(); - LOG_DEBUG(log, "Found Azure SAS token with key: {}", name); - break; + if (object->has(key)) + { + sas_token = object->get(key).extract(); + LOG_DEBUG(log, "Found Azure SAS token with key: {}", key); + break; + } } } if (!sas_token.empty()) - { - return {std::make_shared(sas_token), ""}; - } + return {std::make_shared(sas_token), "", parseSasTokenExpiry(sas_token)}; break; } default: break; } - return {nullptr, ""}; + return {nullptr, "", std::nullopt}; +} + +std::optional RestCatalog::tryGetCachedCredentials( + const std::string & namespace_name, const std::string & table_name) const +{ + std::lock_guard lock(credentials_cache_mutex); + if (vended_credentials_cache_ttl <= std::chrono::seconds::zero()) + return std::nullopt; + + auto it = credentials_cache.find({namespace_name, table_name}); + if (it == credentials_cache.end()) + return std::nullopt; + if (std::chrono::system_clock::now() >= it->second.expires_at.value()) + { + credentials_cache.erase(it); + return std::nullopt; + } + + return it->second; +} + +void RestCatalog::cacheCredentials( + const std::string & namespace_name, + const std::string & table_name, + const VendedStorageCredentials & parsed, + const CatalogStateVersion & state_snapshot) const +{ + if (!parsed.credentials || parsed.credentials->isEmpty()) + return; + + const auto now = std::chrono::system_clock::now(); + + std::lock_guard lock(credentials_cache_mutex); + if (vended_credentials_cache_ttl <= std::chrono::seconds::zero() || state.get() != state_snapshot) + return; + + auto refresh_after = now + vended_credentials_cache_ttl; + if (parsed.expires_at) + { + const auto safe_expiry = parsed.expires_at.value() - credentials_expiry_safety_window; + if (safe_expiry < refresh_after) + refresh_after = safe_expiry; + } + if (refresh_after <= now) + return; + + if (credentials_cache.size() >= credentials_cache_cleanup_threshold) + std::erase_if(credentials_cache, [&now](const auto & entry) { return now >= entry.second.expires_at.value(); }); + credentials_cache[{namespace_name, table_name}] + = VendedStorageCredentials{parsed.credentials, parsed.endpoint, refresh_after, parsed.table_uuid, parsed.location}; } ICatalog::CredentialsRefreshCallback RestCatalog::getCredentialsConfigurationCallback(const DB::StorageID & storage_id) @@ -2319,32 +2564,40 @@ ICatalog::CredentialsRefreshCallback RestCatalog::getCredentialsConfigurationCal Poco::Dynamic::Var json = parser.parse(json_str); const Poco::JSON::Object::Ptr & object = json.extract(); - if (!object->has("config")) - { - LOG_DEBUG(log, "No 'config' in response for table {} – catalog does not support credential vending", table_name); - return nullptr; - } - - auto config_object = object->get("config").extract(); - if (!config_object) - { - LOG_DEBUG(log, "Empty 'config' in response for table {}", table_name); - return nullptr; - } + Poco::JSON::Object::Ptr metadata_object; + if (object->has("metadata")) + metadata_object = object->get("metadata").extract(); + /// Prefix matching uses the table location; metadata may live outside it with a custom `write.metadata.path`. std::string location; - if (object->has("metadata-location")) - { + if (metadata_object && metadata_object->has("location")) + location = metadata_object->get("location").extract(); + else if (object->has("metadata-location")) location = object->get("metadata-location").extract(); - LOG_DEBUG(log, "Location for table {}: {}", table_name, location); - } else + throw DB::Exception(DB::ErrorCodes::DATALAKE_DATABASE_ERROR, "Cannot read table {}, because no location in response", table_name); + LOG_DEBUG(log, "Location for table {}: {}", table_name, location); + + const auto table_uuid = parseTableUuid(metadata_object); + DB::UUID parsed_table_uuid = DB::UUIDHelpers::Nil; + if (storage_id.hasUUID() && (!DB::tryParse(parsed_table_uuid, table_uuid) || parsed_table_uuid != storage_id.uuid)) + throw DB::Exception( + DB::ErrorCodes::DATALAKE_DATABASE_ERROR, + "Cannot refresh credentials for table {} because its identity changed", + table_name); + + const auto config_object = effectiveVendedConfig(object, location); + if (!config_object) { - throw DB::Exception(DB::ErrorCodes::BAD_ARGUMENTS, "Cannot read table {}, because no 'metadata-location' in response", table_name); + LOG_DEBUG(log, "No credentials in response for table {} – catalog does not support credential vending", table_name); + return nullptr; } - auto [new_credentials, _] = getCredentialsAndEndpoint(config_object, location); - return new_credentials; + auto parsed = getCredentialsAndEndpoint(config_object, location); + parsed.table_uuid = table_uuid; + parsed.location = location; + cacheCredentials(namespace_name, table_name, parsed, state_snapshot); + return parsed.credentials; }; } diff --git a/src/Databases/DataLake/RestCatalog.h b/src/Databases/DataLake/RestCatalog.h index 077f186395ca..669547be4cc2 100644 --- a/src/Databases/DataLake/RestCatalog.h +++ b/src/Databases/DataLake/RestCatalog.h @@ -8,7 +8,12 @@ #include #include #include +#include +#include #include +#include +#include +#include #include #include @@ -33,6 +38,15 @@ struct AccessToken } }; +struct VendedStorageCredentials +{ + std::shared_ptr credentials; + std::string endpoint; + std::optional expires_at; + std::string table_uuid{}; + std::string location{}; +}; + class RestCatalog : public ICatalog, public DB::WithContext { public: @@ -93,6 +107,8 @@ class RestCatalog : public ICatalog, public DB::WithContext ICatalog::CredentialsRefreshCallback getCredentialsConfigurationCallback(const DB::StorageID & storage_id) override; + void setVendedCredentialsCacheTTL(std::chrono::seconds ttl) override; + struct Config { /// Prefix is a path of the catalog endpoint, @@ -180,6 +196,15 @@ class RestCatalog : public ICatalog, public DB::WithContext protected: AllowedNamespaces allowed_namespaces; + static constexpr size_t credentials_cache_cleanup_threshold = 1000; + + static constexpr std::chrono::seconds credentials_expiry_safety_window{60}; + mutable std::mutex credentials_cache_mutex; + std::chrono::seconds vended_credentials_cache_ttl TSA_GUARDED_BY(credentials_cache_mutex){std::chrono::seconds::zero()}; + + mutable std::map, VendedStorageCredentials> credentials_cache + TSA_GUARDED_BY(credentials_cache_mutex); + Poco::Net::HTTPBasicCredentials credentials{}; /// `catalog_state` is the snapshot the caller derived the endpoint from, so that one @@ -228,7 +253,8 @@ class RestCatalog : public ICatalog, public DB::WithContext const std::string & namespace_name, const std::string & table_name, DB::ContextPtr context_, - TableMetadata & result) const; + TableMetadata & result, + bool allow_credentials_cache = true) const; /// Load catalog config (special http handler) utilizing information from catalog_state and auth_headers. Config loadConfig(const CatalogState & catalog_state, const std::optional & auth_headers = std::nullopt); @@ -248,7 +274,16 @@ class RestCatalog : public ICatalog, public DB::WithContext const String & method, bool ignore_result) const; - std::pair, String> getCredentialsAndEndpoint(Poco::JSON::Object::Ptr object, const String & location) const; + VendedStorageCredentials getCredentialsAndEndpoint(Poco::JSON::Object::Ptr object, const String & location) const; + + std::optional tryGetCachedCredentials( + const std::string & namespace_name, const std::string & table_name) const; + + void cacheCredentials( + const std::string & namespace_name, + const std::string & table_name, + const VendedStorageCredentials & parsed, + const CatalogStateVersion & state_snapshot) const; AccessToken retrieveAccessToken(const std::string & client_id, const std::string & client_secret) const; diff --git a/src/Databases/DataLake/StorageCredentials.h b/src/Databases/DataLake/StorageCredentials.h index 7e73d66ff673..f0d63b1a2887 100644 --- a/src/Databases/DataLake/StorageCredentials.h +++ b/src/Databases/DataLake/StorageCredentials.h @@ -18,6 +18,7 @@ class IStorageCredentials virtual ~IStorageCredentials() = default; virtual void addCredentialsToEngineArgs(DB::ASTs & engine_args) const = 0; + virtual bool isEmpty() const = 0; }; class S3Credentials final : public IStorageCredentials @@ -32,7 +33,7 @@ class S3Credentials final : public IStorageCredentials , session_token(session_token_) {} - bool isEmpty() const { return access_key_id.empty() || secret_access_key.empty(); } + bool isEmpty() const override { return access_key_id.empty() || secret_access_key.empty(); } void addCredentialsToEngineArgs(DB::ASTs & engine_args) const override { @@ -89,6 +90,8 @@ class GCSCredentials final : public IStorageCredentials DB::make_intrusive("Bearer " + oauth_token)))); } + bool isEmpty() const override { return oauth_token.empty(); } + const std::string & getToken() const { return oauth_token; } private: @@ -111,6 +114,10 @@ class AzureCredentials final : public IStorageCredentials engine_args.push_back(DB::make_intrusive(sas_token)); } + bool isEmpty() const override { return sas_token.empty(); } + + const std::string & getToken() const { return sas_token; } + private: std::string sas_token; }; diff --git a/src/Databases/DataLake/UnityCatalog.cpp b/src/Databases/DataLake/UnityCatalog.cpp index 9d1859343c19..de1e4c5d967d 100644 --- a/src/Databases/DataLake/UnityCatalog.cpp +++ b/src/Databases/DataLake/UnityCatalog.cpp @@ -537,7 +537,6 @@ bool UnityCatalog::isNamespaceAllowed(const std::string & namespace_) const return allowed_namespaces.contains("*") || allowed_namespaces.contains(namespace_); } -/// getCredentialsConfigurationCallback method is supported only for S3 storage ICatalog::CredentialsRefreshCallback UnityCatalog::getCredentialsConfigurationCallback(const DB::StorageID & table_id) { if (!table_id.hasUUID()) @@ -548,10 +547,14 @@ ICatalog::CredentialsRefreshCallback UnityCatalog::getCredentialsConfigurationCa const String unity_table_id = toString(table_id.uuid); - return [this, unity_table_id] () -> std::shared_ptr { + return [this, unity_table_id] () -> std::shared_ptr + { LOG_DEBUG(log, "Update credentials in the catalog"); - return parseS3Credentials(requestReadCredentials(unity_table_id)); + auto response = requestReadCredentials(unity_table_id); + if (hasValueAndItsNotNone("azure_user_delegation_sas", response)) + return parseAzureCredentials(response); + return parseS3Credentials(response); }; } diff --git a/src/Disks/DiskObjectStorage/ObjectStorages/AzureBlobStorage/AzureBlobStorageCommon.h b/src/Disks/DiskObjectStorage/ObjectStorages/AzureBlobStorage/AzureBlobStorageCommon.h index 79ff8caeaedb..a5a96311a56e 100644 --- a/src/Disks/DiskObjectStorage/ObjectStorages/AzureBlobStorage/AzureBlobStorageCommon.h +++ b/src/Disks/DiskObjectStorage/ObjectStorages/AzureBlobStorage/AzureBlobStorageCommon.h @@ -154,6 +154,9 @@ struct ConnectionParams std::unique_ptr createForContainer() const; }; +using ConnectionParamsRefreshCallback = std::function()>; +using ContainerClientRefreshCallback = std::function()>; + void processURL(const String & url, const String & container_name, Endpoint & endpoint, AuthMethod & auth_method); std::unique_ptr getContainerClient(const ConnectionParams & params, bool readonly); diff --git a/src/Disks/DiskObjectStorage/ObjectStorages/AzureBlobStorage/AzureObjectStorage.cpp b/src/Disks/DiskObjectStorage/ObjectStorages/AzureBlobStorage/AzureObjectStorage.cpp index a36ed00fbfc6..4c3a95c73762 100644 --- a/src/Disks/DiskObjectStorage/ObjectStorages/AzureBlobStorage/AzureObjectStorage.cpp +++ b/src/Disks/DiskObjectStorage/ObjectStorages/AzureBlobStorage/AzureObjectStorage.cpp @@ -18,6 +18,7 @@ #include #include #include +#include #include #include @@ -125,6 +126,37 @@ class AzureIteratorAsync final : public IObjectStorageIteratorAsync Azure::Storage::Blobs::ListBlobsOptions options; }; +/// Capture by value: buffers may outlive the storage. +AzureBlobStorage::ContainerClientRefreshCallback makeContainerClientRefresher( + const AzureObjectStorage::AzureCredentialsRefreshCallback & credentials_refresh_callback) +{ + if (!credentials_refresh_callback) + return {}; + + return [credentials_refresh_callback]() -> std::unique_ptr + { + auto params = credentials_refresh_callback(); + if (!params) + return nullptr; + return params->createForContainer(); + }; +} + +WriteBufferFromAzureDataLakeStorage::FileClientRefreshCallback makeDataLakeFileClientRefresher( + const AzureObjectStorage::AzureCredentialsRefreshCallback & credentials_refresh_callback, const String & blob_path) +{ + if (!credentials_refresh_callback) + return {}; + + return [credentials_refresh_callback, blob_path]() -> std::optional + { + auto params = credentials_refresh_callback(); + if (!params) + return {}; + return makeAdlsGen2FileClient(params->endpoint, params->auth_method, params->client_options, blob_path); + }; +} + } @@ -136,7 +168,8 @@ AzureObjectStorage::AzureObjectStorage( const AzureBlobStorage::ConnectionParams & connection_params_, const String & object_namespace_, const String & description_, - const String & common_key_prefix_) + const String & common_key_prefix_, + AzureCredentialsRefreshCallback credentials_refresh_callback_) : name(name_) , auth_method(std::move(auth_method_)) , client(std::move(client_)) @@ -145,10 +178,39 @@ AzureObjectStorage::AzureObjectStorage( , description(description_) , common_key_prefix(common_key_prefix_) , connection_params(connection_params_) + , credentials_refresh_callback(std::move(credentials_refresh_callback_)) , log(getLogger("AzureObjectStorage")) { } +bool AzureObjectStorage::tryRefreshClient(const Azure::Core::RequestFailedException & e) const +{ + if (!credentials_refresh_callback || !isAzureAccessTokenExpiredError(e)) + return false; + + auto params = credentials_refresh_callback(); + if (!params) + return false; + + client.set(params->createForContainer()); + LOG_DEBUG(log, "Refreshed Azure credentials after an authentication failure"); + return true; +} + +std::optional +AzureObjectStorage::tryRefreshDataLakeFileClient(const Azure::Core::RequestFailedException & e, const String & blob_path) const +{ + if (!credentials_refresh_callback || !isAzureAccessTokenExpiredError(e)) + return {}; + + auto params = credentials_refresh_callback(); + if (!params) + return {}; + + LOG_DEBUG(log, "Refreshed Azure credentials for `{}` after an authentication failure", blob_path); + return makeAdlsGen2FileClient(params->endpoint, params->auth_method, params->client_options, blob_path); +} + ObjectStorageKeyGeneratorPtr AzureObjectStorage::createKeyGenerator() const { return createObjectStorageKeyGeneratorByTemplate("[a-z]{32}"); @@ -156,23 +218,28 @@ ObjectStorageKeyGeneratorPtr AzureObjectStorage::createKeyGenerator() const bool AzureObjectStorage::exists(const StoredObject & object) const { - auto client_ptr = client.get(); + for (size_t attempt = 0; ; ++attempt) + { + auto client_ptr = client.get(); - ProfileEvents::increment(ProfileEvents::AzureGetProperties); - if (client_ptr->IsClientForDisk()) - ProfileEvents::increment(ProfileEvents::DiskAzureGetProperties); + ProfileEvents::increment(ProfileEvents::AzureGetProperties); + if (client_ptr->IsClientForDisk()) + ProfileEvents::increment(ProfileEvents::DiskAzureGetProperties); - try - { - auto blob_client = client_ptr->GetBlobClient(object.remote_path); - blob_client.GetProperties(); - return true; - } - catch (const Azure::Storage::StorageException & e) - { - if (e.StatusCode == Azure::Core::Http::HttpStatusCode::NotFound) - return false; - throw; + try + { + auto blob_client = client_ptr->GetBlobClient(object.remote_path); + blob_client.GetProperties(); + return true; + } + catch (const Azure::Storage::StorageException & e) + { + if (e.StatusCode == Azure::Core::Http::HttpStatusCode::NotFound) + return false; + if (attempt == 0 && tryRefreshClient(e)) + continue; + throw; + } } } @@ -191,8 +258,6 @@ ObjectStorageIteratorPtr AzureObjectStorage::iterate( void AzureObjectStorage::listObjects(const std::string & path, RelativePathsWithMetadata & children, size_t max_keys) const { - auto client_ptr = client.get(); - Azure::Storage::Blobs::ListBlobsOptions options; options.Prefix = path; if (max_keys) @@ -200,6 +265,31 @@ void AzureObjectStorage::listObjects(const std::string & path, RelativePathsWith else options.PageSizeHint = settings.get()->list_object_keys_size; + /// The continuation token stays valid across a credentials refresh, so a page can be re-fetched without restarting the listing. + auto list_blobs = [&] + { + for (size_t attempt = 0; ; ++attempt) + { + auto client_ptr = client.get(); + try + { + auto response = client_ptr->ListBlobs(options); + + ProfileEvents::increment(ProfileEvents::AzureListObjects); + if (client_ptr->IsClientForDisk()) + ProfileEvents::increment(ProfileEvents::DiskAzureListObjects); + + return response; + } + catch (const Azure::Core::RequestFailedException & e) + { + if (attempt == 0 && tryRefreshClient(e)) + continue; + throw; + } + } + }; + /// Re-issue ListBlobs per page through the client wrapper (which strips the endpoint prefix); the SDK's /// MoveToNextPage refetches pages 2..N directly and would leave the raw Azure prefix on their blob names. while (true) @@ -207,13 +297,9 @@ void AzureObjectStorage::listObjects(const std::string & path, RelativePathsWith AzureBlobStorage::ListBlobsPagedResponse blob_list_response; { ProfileEventTimeIncrement watch(ProfileEvents::AzureListObjectsMicroseconds); - blob_list_response = client_ptr->ListBlobs(options); + blob_list_response = list_blobs(); } - ProfileEvents::increment(ProfileEvents::AzureListObjects); - if (client_ptr->IsClientForDisk()) - ProfileEvents::increment(ProfileEvents::DiskAzureListObjects); - for (const auto & blob : blob_list_response.Blobs) { children.emplace_back(std::make_shared( @@ -266,7 +352,8 @@ std::unique_ptr AzureObjectStorage::readObject( /// NOLI restrict_seek, /* read_until_position */0, std::move(blob_storage_log), - connection_params.getContainer()); + connection_params.getContainer(), + makeContainerClientRefresher(credentials_refresh_callback)); } SmallObjectDataWithMetadata AzureObjectStorage::readSmallObjectAndGetObjectMetadata( /// NOLINT @@ -331,7 +418,8 @@ std::unique_ptr AzureObjectStorage::writeObject( /// NO patchSettings(write_settings), settings.get(), connection_params.getContainer(), - std::move(blob_storage_log)); + std::move(blob_storage_log), + makeDataLakeFileClientRefresher(credentials_refresh_callback, object.remote_path)); } return std::make_unique( @@ -342,7 +430,8 @@ std::unique_ptr AzureObjectStorage::writeObject( /// NO settings.get(), connection_params.getContainer(), std::move(blob_storage_log), - std::move(scheduler)); + std::move(scheduler), + makeContainerClientRefresher(credentials_refresh_callback)); } void AzureObjectStorage::removeObjectImpl( @@ -351,86 +440,113 @@ void AzureObjectStorage::removeObjectImpl( bool if_exists, BlobStorageLogWriterPtr blob_storage_log) { - ProfileEvents::increment(ProfileEvents::AzureDeleteObjects); - if (client_ptr->IsClientForDisk()) - ProfileEvents::increment(ProfileEvents::DiskAzureDeleteObjects); - const auto & path = object.remote_path; LOG_TEST(log, "Removing single object: {}", path); - Stopwatch watch; - Int32 error_code = 0; - String error_message; - bool success = false; - try + const bool is_adls_gen2 = isAdlsGen2Endpoint(connection_params.endpoint); + auto current_client = client_ptr; + std::optional refreshed_file_client; + + for (size_t attempt = 0; ; ++attempt) { - if (isAdlsGen2Endpoint(connection_params.endpoint)) + ProfileEvents::increment(ProfileEvents::AzureDeleteObjects); + if (current_client->IsClientForDisk()) + ProfileEvents::increment(ProfileEvents::DiskAzureDeleteObjects); + + Stopwatch watch; + Int32 error_code = 0; + String error_message; + bool success = false; + try { - buildDataLakeFileClient(path)->Delete(); - success = true; + if (is_adls_gen2) + { + if (refreshed_file_client) + refreshed_file_client->Delete(); + else + buildDataLakeFileClient(path)->Delete(); + success = true; + } + else + { + auto delete_info = current_client->GetBlobClient(path).Delete(); + success = delete_info.Value.Deleted; + if (!if_exists && !delete_info.Value.Deleted) + throw Exception( + ErrorCodes::AZURE_BLOB_STORAGE_ERROR, "Failed to delete file (path: {}) in AzureBlob Storage, reason: {}", + path, delete_info.RawResponse ? delete_info.RawResponse->GetReasonPhrase() : "Unknown"); + } } - else + catch (const Azure::Storage::StorageException & e) { - auto delete_info = client_ptr->GetBlobClient(path).Delete(); - success = delete_info.Value.Deleted; - if (!if_exists && !delete_info.Value.Deleted) - throw Exception( - ErrorCodes::AZURE_BLOB_STORAGE_ERROR, "Failed to delete file (path: {}) in AzureBlob Storage, reason: {}", - path, delete_info.RawResponse ? delete_info.RawResponse->GetReasonPhrase() : "Unknown"); - } - } - catch (const Azure::Storage::StorageException & e) - { - error_code = static_cast(e.StatusCode); - error_message = e.Message; + error_code = static_cast(e.StatusCode); + error_message = e.Message; - if (!if_exists) - { - if (blob_storage_log) - blob_storage_log->addEvent( - BlobStorageLogElement::EventType::Delete, - /* bucket */ connection_params.getContainer(), - /* remote_path */ path, - object.local_path, - object.bytes_size, - watch.elapsedMicroseconds(), - error_code, - error_message); - throw; - } + if (attempt == 0) + { + if (is_adls_gen2) + { + refreshed_file_client = tryRefreshDataLakeFileClient(e, path); + if (refreshed_file_client) + continue; + } + else if (tryRefreshClient(e)) + { + current_client = client.get(); + continue; + } + } - /// If object doesn't exist. - if (e.StatusCode == Azure::Core::Http::HttpStatusCode::NotFound) - { - auto elapsed = watch.elapsedMicroseconds(); - if (blob_storage_log) - blob_storage_log->addEvent( - BlobStorageLogElement::EventType::Delete, - /* bucket */ connection_params.getContainer(), - /* remote_path */ path, - object.local_path, - object.bytes_size, - elapsed, - error_code, - error_message); - return; + if (!if_exists) + { + if (blob_storage_log) + blob_storage_log->addEvent( + BlobStorageLogElement::EventType::Delete, + /* bucket */ connection_params.getContainer(), + /* remote_path */ path, + object.local_path, + object.bytes_size, + watch.elapsedMicroseconds(), + error_code, + error_message); + throw; + } + + /// If object doesn't exist. + if (e.StatusCode == Azure::Core::Http::HttpStatusCode::NotFound) + { + auto elapsed = watch.elapsedMicroseconds(); + if (blob_storage_log) + blob_storage_log->addEvent( + BlobStorageLogElement::EventType::Delete, + /* bucket */ connection_params.getContainer(), + /* remote_path */ path, + object.local_path, + object.bytes_size, + elapsed, + error_code, + error_message); + + return; + } + + tryLogCurrentException(__PRETTY_FUNCTION__); + throw; } + auto elapsed = watch.elapsedMicroseconds(); - tryLogCurrentException(__PRETTY_FUNCTION__); - throw; + if (blob_storage_log) + blob_storage_log->addEvent( + BlobStorageLogElement::EventType::Delete, + /* bucket */ connection_params.getContainer(), + /* remote_path */ path, + object.local_path, + object.bytes_size, + elapsed, + success ? 0 : error_code, + success ? "" : error_message); + return; } - auto elapsed = watch.elapsedMicroseconds(); - - if (blob_storage_log) - blob_storage_log->addEvent( - BlobStorageLogElement::EventType::Delete, - /* bucket */ connection_params.getContainer(), - /* remote_path */ path, - object.local_path, - object.bytes_size, - elapsed, - success ? 0 : error_code, - success ? "" : error_message); } void AzureObjectStorage::removeObjectIfExists(const StoredObject & object) @@ -440,7 +556,7 @@ void AzureObjectStorage::removeObjectIfExists(const StoredObject & object) } void AzureObjectStorage::removeObjectsBatchIfExists( - const StoredObjects & objects, + StoredObjectsSpan & rest_objects, const std::shared_ptr & client_ptr, BlobStorageLogWriterPtr blob_storage_log) { @@ -467,11 +583,9 @@ void AzureObjectStorage::removeObjectsBatchIfExists( } }; - StoredObjectsSpan rest_objects = objects; while (!rest_objects.empty()) { auto object_batch = rest_objects.first(std::min(rest_objects.size(), AZURE_BATCH_MAX_SUBREQUESTS)); - SCOPE_EXIT({ rest_objects = rest_objects.last(rest_objects.size() - object_batch.size()); }); Stopwatch watch; AzureBlobStorage::BlobContainerBatch requests = client_ptr->CreateBatch(); @@ -536,6 +650,8 @@ void AzureObjectStorage::removeObjectsBatchIfExists( if (throw_at_end) std::rethrow_exception(throw_at_end); + + rest_objects = rest_objects.last(rest_objects.size() - object_batch.size()); } } @@ -544,17 +660,31 @@ void AzureObjectStorage::removeObjectsIfExist(const StoredObjects & objects) if (objects.empty()) return; - auto client_ptr = client.get(); auto blob_storage_log = BlobStorageLogWriter::create(name); if (isAdlsGen2Endpoint(connection_params.endpoint)) { + auto client_ptr = client.get(); for (const auto & object : objects) removeObjectImpl(object, client_ptr, /*if_exists=*/ true, blob_storage_log); return; } - removeObjectsBatchIfExists(objects, client_ptr, blob_storage_log); + StoredObjectsSpan rest_objects = objects; + for (size_t attempt = 0; ; ++attempt) + { + try + { + removeObjectsBatchIfExists(rest_objects, client.get(), blob_storage_log); + return; + } + catch (const Azure::Core::RequestFailedException & e) + { + if (attempt == 0 && tryRefreshClient(e)) + continue; + throw; + } + } } static void setAzureBlobTag( @@ -591,25 +721,37 @@ void AzureObjectStorage::tagObjects(const StoredObjects & objects, const std::st ObjectMetadata AzureObjectStorage::getObjectMetadata(const std::string & path, bool) const { - auto client_ptr = client.get(); - auto blob_client = client_ptr->GetBlobClient(path); - auto properties = blob_client.GetProperties().Value; + for (size_t attempt = 0; ; ++attempt) + { + auto client_ptr = client.get(); + try + { + auto blob_client = client_ptr->GetBlobClient(path); + auto properties = blob_client.GetProperties().Value; - ProfileEvents::increment(ProfileEvents::AzureGetProperties); - if (client_ptr->IsClientForDisk()) - ProfileEvents::increment(ProfileEvents::DiskAzureGetProperties); + ProfileEvents::increment(ProfileEvents::AzureGetProperties); + if (client_ptr->IsClientForDisk()) + ProfileEvents::increment(ProfileEvents::DiskAzureGetProperties); - ObjectMetadata result; - result.size_bytes = properties.BlobSize; - result.etag = properties.ETag.ToString(); - if (!properties.Metadata.empty()) - { - result.attributes.emplace(); - for (const auto & [key, value] : properties.Metadata) - result.attributes[key] = value; + ObjectMetadata result; + result.size_bytes = properties.BlobSize; + result.etag = properties.ETag.ToString(); + if (!properties.Metadata.empty()) + { + result.attributes.emplace(); + for (const auto & [key, value] : properties.Metadata) + result.attributes[key] = value; + } + result.last_modified = static_cast(properties.LastModified).time_since_epoch().count(); + return result; + } + catch (const Azure::Core::RequestFailedException & e) + { + if (attempt == 0 && tryRefreshClient(e)) + continue; + throw; + } } - result.last_modified = static_cast(properties.LastModified).time_since_epoch().count(); - return result; } std::optional AzureObjectStorage::tryGetObjectMetadata(const std::string & path, bool with_tags) const diff --git a/src/Disks/DiskObjectStorage/ObjectStorages/AzureBlobStorage/AzureObjectStorage.h b/src/Disks/DiskObjectStorage/ObjectStorages/AzureBlobStorage/AzureObjectStorage.h index 88420b30472f..39e99f5dc5e6 100644 --- a/src/Disks/DiskObjectStorage/ObjectStorages/AzureBlobStorage/AzureObjectStorage.h +++ b/src/Disks/DiskObjectStorage/ObjectStorages/AzureBlobStorage/AzureObjectStorage.h @@ -4,6 +4,7 @@ #if USE_AZURE_BLOB_STORAGE #include +#include #include #include #include @@ -25,6 +26,7 @@ class AzureObjectStorage : public IObjectStorage public: using ClientPtr = std::unique_ptr; using SettingsPtr = std::unique_ptr; + using AzureCredentialsRefreshCallback = AzureBlobStorage::ConnectionParamsRefreshCallback; AzureObjectStorage( const String & name_, @@ -34,7 +36,8 @@ class AzureObjectStorage : public IObjectStorage const AzureBlobStorage::ConnectionParams & connection_params_, const String & object_namespace_, const String & description_, - const String & common_key_prefix_); + const String & common_key_prefix_, + AzureCredentialsRefreshCallback credentials_refresh_callback_ = {}); bool supportsListObjectsCache() override { return true; } @@ -138,17 +141,23 @@ class AzureObjectStorage : public IObjectStorage bool if_exists, BlobStorageLogWriterPtr blob_storage_log); + /// Advances `rest_objects` past each fully-deleted batch, so a caller can resume after refreshing credentials. void removeObjectsBatchIfExists( - const StoredObjects & objects, + StoredObjectsSpan & rest_objects, const std::shared_ptr & client_ptr, BlobStorageLogWriterPtr blob_storage_log); std::unique_ptr buildDataLakeFileClient(const String & blob_path) const; + bool tryRefreshClient(const Azure::Core::RequestFailedException & e) const; + + std::optional + tryRefreshDataLakeFileClient(const Azure::Core::RequestFailedException & e, const String & blob_path) const; + const String name; AzureBlobStorage::AuthMethod auth_method; - /// client used to access the files in the Blob Storage cloud - MultiVersion client; + /// Client used to access the files in the Blob Storage cloud. + mutable MultiVersion client; MultiVersion settings; const String object_namespace; /// container + prefix @@ -159,6 +168,8 @@ class AzureObjectStorage : public IObjectStorage const AzureBlobStorage::ConnectionParams connection_params; + const AzureCredentialsRefreshCallback credentials_refresh_callback; + LoggerPtr log; }; diff --git a/src/Disks/IO/ReadBufferFromAzureBlobStorage.cpp b/src/Disks/IO/ReadBufferFromAzureBlobStorage.cpp index 9cbed6ada133..05651f3c0d01 100644 --- a/src/Disks/IO/ReadBufferFromAzureBlobStorage.cpp +++ b/src/Disks/IO/ReadBufferFromAzureBlobStorage.cpp @@ -49,9 +49,11 @@ ReadBufferFromAzureBlobStorage::ReadBufferFromAzureBlobStorage( bool restricted_seek_, size_t read_until_position_, BlobStorageLogWriterPtr blob_storage_log_, - String container_for_logging_) + String container_for_logging_, + AzureClientRefreshCallback credentials_refresh_callback_) : ReadBufferFromFileBase() , blob_container_client(blob_container_client_) + , credentials_refresh_callback(std::move(credentials_refresh_callback_)) , path(path_) , max_single_read_retries(max_single_read_retries_) , max_single_download_retries(max_single_download_retries_) @@ -72,6 +74,36 @@ ReadBufferFromAzureBlobStorage::ReadBufferFromAzureBlobStorage( } } +std::pair +ReadBufferFromAzureBlobStorage::tryGetRefreshedClient(const Azure::Core::RequestFailedException & e) const +{ + if (!credentials_refresh_callback || !isAzureAccessTokenExpiredError(e)) + return {}; + + auto new_container = credentials_refresh_callback(); + if (!new_container) + return {}; + + BlobClientPtr new_blob = std::make_unique(new_container->GetBlobClient(path)); + return {std::move(new_container), std::move(new_blob)}; +} + +bool ReadBufferFromAzureBlobStorage::tryRefreshCredentials(const Azure::Core::RequestFailedException & e) +{ + if (credentials_refreshed) + return false; + + auto [new_container, new_blob] = tryGetRefreshedClient(e); + if (!new_container) + return false; + + blob_container_client = std::move(new_container); + blob_client = std::move(new_blob); + credentials_refreshed = true; + LOG_DEBUG(log, "Refreshed Azure credentials for {} after an authentication failure", path); + return true; +} + void ReadBufferFromAzureBlobStorage::setReadUntilEnd() { if (read_until_position) @@ -134,6 +166,14 @@ bool ReadBufferFromAzureBlobStorage::nextImpl() ProfileEvents::increment(ProfileEvents::ReadBufferFromAzureRequestsErrors); LOG_DEBUG(log, "Exception caught during Azure Read for file {} at attempt {}/{}: {}", path, i + 1, max_single_read_retries, e.Message); + if (tryRefreshCredentials(e)) + { + initialized = false; + initialize(i + 1); + --i; /// Don't count the refreshed retry against the budget (refresh happens at most once). + continue; + } + if (i + 1 == max_single_read_retries || !isRetryableAzureException(e)) throw; @@ -293,6 +333,12 @@ void ReadBufferFromAzureBlobStorage::initialize(size_t attempt) ProfileEvents::increment(ProfileEvents::ReadBufferFromAzureRequestsErrors); LOG_DEBUG(log, "Exception caught during Azure Download for file {} at offset {} at attempt {}/{}: {}", path, offset, i + 1, max_single_download_retries, e.Message); + if (tryRefreshCredentials(e)) + { + --i; /// Don't count the refreshed retry against the budget (refresh happens at most once). + continue; + } + if (i + 1 == max_single_download_retries || !isRetryableAzureException(e)) throw; @@ -335,24 +381,62 @@ void ReadBufferFromAzureBlobStorage::initialize(size_t attempt) std::optional ReadBufferFromAzureBlobStorage::tryGetFileSize() { + if (file_size) + return file_size; + if (!blob_client) blob_client = std::make_unique(blob_container_client->GetBlobClient(path)); - if (!file_size) - file_size = blob_client->GetProperties().Value.BlobSize; - - return file_size; + for (size_t attempt = 0; ; ++attempt) + { + try + { + file_size = blob_client->GetProperties().Value.BlobSize; + return file_size; + } + catch (const Azure::Core::RequestFailedException & e) + { + if (attempt == 0 && tryRefreshCredentials(e)) + continue; + throw; + } + } } std::optional ReadBufferFromAzureBlobStorage::getRemoteFileMetadata() const { - const auto properties = blob_container_client->GetBlobClient(path).GetProperties().Value; - const auto last_modification_time = std::chrono::duration_cast( - static_cast(properties.LastModified).time_since_epoch()) - .count(); - return RemoteFileMetadata{ - .size = static_cast(properties.BlobSize), - .last_modification_time = static_cast(last_modification_time)}; + ContainerClientPtr refreshed_container_client; + BlobClientPtr refreshed_blob_client; + auto initial_blob_client = blob_container_client->GetBlobClient(path); + const AzureBlobStorage::BlobClient * current_blob_client = &initial_blob_client; + + for (size_t attempt = 0; ; ++attempt) + { + try + { + const auto properties = current_blob_client->GetProperties().Value; + const auto last_modification_time = std::chrono::duration_cast( + static_cast(properties.LastModified).time_since_epoch()) + .count(); + return RemoteFileMetadata{ + .size = static_cast(properties.BlobSize), + .last_modification_time = static_cast(last_modification_time)}; + } + catch (const Azure::Core::RequestFailedException & e) + { + if (attempt == 0) + { + if (auto [new_container, new_blob] = tryGetRefreshedClient(e); new_container) + { + refreshed_container_client = std::move(new_container); + refreshed_blob_client = std::move(new_blob); + current_blob_client = refreshed_blob_client.get(); + continue; + } + } + throw; + } + } } size_t ReadBufferFromAzureBlobStorage::readBigAt(char * to, size_t n, size_t range_begin, const std::function & /*progress_callback*/) const @@ -362,6 +446,11 @@ size_t ReadBufferFromAzureBlobStorage::readBigAt(char * to, size_t n, size_t ran ProfileEventTimeIncrement watch(ProfileEvents::ReadBufferFromAzureMicroseconds); + ContainerClientPtr refreshed_container_client; + BlobClientPtr refreshed_blob_client; + const AzureBlobStorage::BlobClient * current_blob_client = blob_client.get(); + bool credentials_refreshed_locally = false; + for (size_t i = 0; i < max_single_download_retries && n > 0; ++i) { size_t bytes_copied = 0; @@ -377,7 +466,7 @@ size_t ReadBufferFromAzureBlobStorage::readBigAt(char * to, size_t n, size_t ran download_options.Range = {static_cast(range_begin), n}; Azure::Core::Context azure_context = Azure::Core::Context().WithValue(PocoAzureHTTPClient::getSDKContextKeyForBufferRetry(), size_t{0}); - auto download_response = blob_client->Download(download_options, azure_context); + auto download_response = current_blob_client->Download(download_options, azure_context); if (blob_storage_log) { blob_storage_log->addEvent( @@ -413,6 +502,20 @@ size_t ReadBufferFromAzureBlobStorage::readBigAt(char * to, size_t n, size_t ran ProfileEvents::increment(ProfileEvents::ReadBufferFromAzureRequestsErrors); LOG_DEBUG(log, "Exception caught during Azure Download for file {} at offset {} at attempt {}/{}: {}", path, offset, i + 1, max_single_download_retries, e.Message); + if (!credentials_refreshed_locally) + { + if (auto [new_container, new_blob] = tryGetRefreshedClient(e); new_container) + { + refreshed_container_client = std::move(new_container); + refreshed_blob_client = std::move(new_blob); + current_blob_client = refreshed_blob_client.get(); + credentials_refreshed_locally = true; + LOG_DEBUG(log, "Refreshed Azure credentials for {} after an authentication failure in readBigAt", path); + --i; /// Don't count the refreshed retry against the budget. + continue; + } + } + if (i + 1 == max_single_download_retries || !isRetryableAzureException(e)) throw; diff --git a/src/Disks/IO/ReadBufferFromAzureBlobStorage.h b/src/Disks/IO/ReadBufferFromAzureBlobStorage.h index cc0a75e02c58..3b5d30b05615 100644 --- a/src/Disks/IO/ReadBufferFromAzureBlobStorage.h +++ b/src/Disks/IO/ReadBufferFromAzureBlobStorage.h @@ -1,5 +1,6 @@ #pragma once +#include #include #include "config.h" @@ -23,6 +24,7 @@ class ReadBufferFromAzureBlobStorage : public ReadBufferFromFileBase public: using ContainerClientPtr = std::shared_ptr; using BlobClientPtr = std::unique_ptr; + using AzureClientRefreshCallback = AzureBlobStorage::ContainerClientRefreshCallback; ReadBufferFromAzureBlobStorage( ContainerClientPtr blob_container_client_, @@ -34,7 +36,8 @@ class ReadBufferFromAzureBlobStorage : public ReadBufferFromFileBase bool restricted_seek_ = false, size_t read_until_position_ = 0, BlobStorageLogWriterPtr blob_storage_log_ = {}, - String container_for_logging_ = {}); + String container_for_logging_ = {}, + AzureClientRefreshCallback credentials_refresh_callback_ = {}); off_t seek(off_t off, int whence) override; @@ -70,9 +73,16 @@ class ReadBufferFromAzureBlobStorage : public ReadBufferFromFileBase void initialize(size_t attempt); void setMetadataFromResponse(const Azure::Storage::Blobs::Models::DownloadBlobDetails & details, size_t blob_size) const; + std::pair tryGetRefreshedClient(const Azure::Core::RequestFailedException & e) const; + + /// Sequential reads only; refreshes at most once per buffer. + bool tryRefreshCredentials(const Azure::Core::RequestFailedException & e); + std::unique_ptr data_stream; ContainerClientPtr blob_container_client; BlobClientPtr blob_client; + const AzureClientRefreshCallback credentials_refresh_callback; + bool credentials_refreshed = false; const String path; size_t max_single_read_retries; diff --git a/src/Disks/IO/WriteBufferFromAzureBlobStorage.cpp b/src/Disks/IO/WriteBufferFromAzureBlobStorage.cpp index 13649b2c21d6..c8807127883b 100644 --- a/src/Disks/IO/WriteBufferFromAzureBlobStorage.cpp +++ b/src/Disks/IO/WriteBufferFromAzureBlobStorage.cpp @@ -60,7 +60,8 @@ WriteBufferFromAzureBlobStorage::WriteBufferFromAzureBlobStorage( std::shared_ptr settings_, const String & container_for_logging_, BlobStorageLogWriterPtr blob_log_, - ThreadPoolCallbackRunnerUnsafe schedule_) + ThreadPoolCallbackRunnerUnsafe schedule_, + AzureBlobStorage::ContainerClientRefreshCallback credentials_refresh_callback_) : WriteBufferFromFileBase(std::min(buf_size_, static_cast(DBMS_DEFAULT_BUFFER_SIZE)), nullptr, 0) , log(getLogger("WriteBufferFromAzureBlobStorage")) , buffer_allocation_policy(createBufferAllocationPolicy(*settings_)) @@ -69,6 +70,7 @@ WriteBufferFromAzureBlobStorage::WriteBufferFromAzureBlobStorage( , blob_path(blob_path_) , write_settings(write_settings_) , blob_container_client(blob_container_client_) + , credentials_refresh_callback(std::move(credentials_refresh_callback_)) , task_tracker( std::make_unique( std::move(schedule_), @@ -112,20 +114,59 @@ WriteBufferFromAzureBlobStorage::~WriteBufferFromAzureBlobStorage() task_tracker->safeWaitAll(); } -void WriteBufferFromAzureBlobStorage::execWithRetry(std::function func, size_t num_tries, size_t cost) +WriteBufferFromAzureBlobStorage::AzureClientPtr WriteBufferFromAzureBlobStorage::getClient() const +{ + std::lock_guard lock(client_mutex); + return blob_container_client; +} + +bool WriteBufferFromAzureBlobStorage::tryRefreshCredentials( + const Azure::Core::RequestFailedException & e, const AzureClientPtr & used_client) +{ + if (!credentials_refresh_callback || !isAzureAccessTokenExpiredError(e)) + return false; + + std::lock_guard lock(client_mutex); + + /// Another part upload already refreshed the credentials while this attempt was in flight. + if (blob_container_client != used_client) + return true; + + if (credentials_refreshed) + return false; + + auto new_client = credentials_refresh_callback(); + if (!new_client) + return false; + + blob_container_client = std::move(new_client); + credentials_refreshed = true; + LOG_DEBUG(log, "Refreshed Azure credentials for blob `{}` after an authentication failure", blob_path); + return true; +} + +void WriteBufferFromAzureBlobStorage::execWithRetry(std::function func, size_t num_tries, size_t cost) { size_t sleep_time_with_backoff_milliseconds = 100; for (size_t i = 0; i < num_tries; ++i) { + auto client_ptr = getClient(); try { ResourceGuard rlock(ResourceGuard::Metrics::getIOWrite(), write_settings.io_scheduling.write_resource_link, cost); // Note that zero-cost requests are ignored - func(i); + func(i, client_ptr); rlock.unlock(cost); break; } catch (const Azure::Core::RequestFailedException & e) { + /// Credentials are refreshed at most once, so the retry after it costs no attempt. + if (tryRefreshCredentials(e, client_ptr)) + { + --i; + continue; + } + if (i == num_tries - 1 || !isRetryableAzureException(e)) throw; @@ -166,11 +207,9 @@ void WriteBufferFromAzureBlobStorage::preFinalize() if (block_ids.empty()) { ProfileEvents::increment(ProfileEvents::AzureUpload); - if (blob_container_client->IsClientForDisk()) + if (getClient()->IsClientForDisk()) ProfileEvents::increment(ProfileEvents::DiskAzureUpload); - auto block_blob_client = blob_container_client->GetBlockBlobClient(blob_path); - /// If there is only one block and size is less than or equal to max_single_part_upload_size /// then we use single part upload instead of multi part upload if (detached_part_data.size() == 1 && detached_part_data.front().data_size <= max_single_part_upload_size) @@ -185,7 +224,7 @@ void WriteBufferFromAzureBlobStorage::preFinalize() try { execWithRetry( - [&](size_t retry_attempt) + [&](size_t retry_attempt, const AzureClientPtr & client_ptr) { Azure::Storage::Blobs::UploadBlockBlobOptions options; @@ -195,7 +234,7 @@ void WriteBufferFromAzureBlobStorage::preFinalize() if (!write_settings.object_storage_write_if_match.empty()) options.AccessConditions.IfMatch = Azure::ETag(write_settings.object_storage_write_if_match); - block_blob_client.Upload( + client_ptr->GetBlockBlobClient(blob_path).Upload( memory_stream, options, azure_context.WithValue(PocoAzureHTTPClient::getSDKContextKeyForBufferRetry(), retry_attempt)); @@ -248,7 +287,7 @@ void WriteBufferFromAzureBlobStorage::preFinalize() try { execWithRetry( - [&](size_t retry_attempt) + [&](size_t retry_attempt, const AzureClientPtr & client_ptr) { Azure::Storage::Blobs::UploadBlockBlobOptions options; @@ -258,7 +297,7 @@ void WriteBufferFromAzureBlobStorage::preFinalize() if (!write_settings.object_storage_write_if_match.empty()) options.AccessConditions.IfMatch = Azure::ETag(write_settings.object_storage_write_if_match); - block_blob_client.Upload( + client_ptr->GetBlockBlobClient(blob_path).Upload( memory_stream, options, azure_context.WithValue(PocoAzureHTTPClient::getSDKContextKeyForBufferRetry(), retry_attempt)); @@ -319,9 +358,8 @@ void WriteBufferFromAzureBlobStorage::finalizeImpl() if (!block_ids.empty()) { - auto block_blob_client = blob_container_client->GetBlockBlobClient(blob_path); ProfileEvents::increment(ProfileEvents::AzureCommitBlockList); - if (blob_container_client->IsClientForDisk()) + if (getClient()->IsClientForDisk()) ProfileEvents::increment(ProfileEvents::DiskAzureCommitBlockList); Stopwatch watch; @@ -330,7 +368,7 @@ void WriteBufferFromAzureBlobStorage::finalizeImpl() try { execWithRetry( - [&](size_t retry_attetmpt) + [&](size_t retry_attetmpt, const AzureClientPtr & client_ptr) { Azure::Storage::Blobs::CommitBlockListOptions options; @@ -341,7 +379,7 @@ void WriteBufferFromAzureBlobStorage::finalizeImpl() options.AccessConditions.IfMatch = Azure::ETag(write_settings.object_storage_write_if_match); - block_blob_client.CommitBlockList( + client_ptr->GetBlockBlobClient(blob_path).CommitBlockList( block_ids, options, azure_context.WithValue(PocoAzureHTTPClient::getSDKContextKeyForBufferRetry(), retry_attetmpt)); @@ -384,8 +422,12 @@ void WriteBufferFromAzureBlobStorage::finalizeImpl() { try { - auto blob_client = blob_container_client->GetBlobClient(blob_path); - blob_client.GetProperties(); + execWithRetry( + [&](size_t, const AzureClientPtr & client_ptr) + { + client_ptr->GetBlobClient(blob_path).GetProperties(); + }, + max_unexpected_write_error_retries); } catch (const Azure::Storage::StorageException & e) { @@ -507,10 +549,9 @@ void WriteBufferFromAzureBlobStorage::writePart(WriteBufferFromAzureBlobStorage: { auto & data_size = std::get<1>(*worker_data).data_size; auto & data_block_id = std::get<0>(*worker_data); - auto block_blob_client = blob_container_client->GetBlockBlobClient(blob_path); ProfileEvents::increment(ProfileEvents::AzureStageBlock); - if (blob_container_client->IsClientForDisk()) + if (getClient()->IsClientForDisk()) ProfileEvents::increment(ProfileEvents::DiskAzureStageBlock); Azure::Core::IO::MemoryBodyStream memory_stream(reinterpret_cast(std::get<1>(*worker_data).memory.data()), data_size); @@ -521,9 +562,9 @@ void WriteBufferFromAzureBlobStorage::writePart(WriteBufferFromAzureBlobStorage: try { execWithRetry( - [&](size_t retry_attempt) + [&](size_t retry_attempt, const AzureClientPtr & client_ptr) { - block_blob_client.StageBlock( + client_ptr->GetBlockBlobClient(blob_path).StageBlock( data_block_id, memory_stream, Azure::Storage::Blobs::StageBlockOptions{}, diff --git a/src/Disks/IO/WriteBufferFromAzureBlobStorage.h b/src/Disks/IO/WriteBufferFromAzureBlobStorage.h index 5ddda389d467..1324d9fbaafa 100644 --- a/src/Disks/IO/WriteBufferFromAzureBlobStorage.h +++ b/src/Disks/IO/WriteBufferFromAzureBlobStorage.h @@ -5,7 +5,9 @@ #if USE_AZURE_BLOB_STORAGE #include +#include +#include #include #include #include @@ -39,7 +41,8 @@ class WriteBufferFromAzureBlobStorage : public WriteBufferFromFileBase std::shared_ptr settings_, const String & container_for_logging_ = {}, BlobStorageLogWriterPtr blob_log_ = {}, - ThreadPoolCallbackRunnerUnsafe schedule_ = {}); + ThreadPoolCallbackRunnerUnsafe schedule_ = {}, + AzureBlobStorage::ContainerClientRefreshCallback credentials_refresh_callback_ = {}); ~WriteBufferFromAzureBlobStorage() override; @@ -60,9 +63,14 @@ class WriteBufferFromAzureBlobStorage : public WriteBufferFromFileBase void setFakeBufferWhenPreFinalized(); void finalizeImpl() override; - void execWithRetry(std::function func, size_t num_tries, size_t cost = 0); + /// `func` must use the supplied client, which may change between attempts. + void execWithRetry(std::function func, size_t num_tries, size_t cost = 0); void uploadBlock(const char * data, size_t size); + AzureClientPtr getClient() const; + + bool tryRefreshCredentials(const Azure::Core::RequestFailedException & e, const AzureClientPtr & used_client); + /// Returns true if not a single byte was written to the buffer bool isEmpty() const { return total_size == 0 && count() == 0 && hidden_size == 0 && offset() == 0; } @@ -81,7 +89,12 @@ class WriteBufferFromAzureBlobStorage : public WriteBufferFromFileBase /// Track that prefinalize() is called only once bool is_prefinalized = false; - AzureClientPtr blob_container_client; + mutable std::mutex client_mutex; + AzureClientPtr blob_container_client TSA_GUARDED_BY(client_mutex); + bool credentials_refreshed TSA_GUARDED_BY(client_mutex) = false; + + const AzureBlobStorage::ContainerClientRefreshCallback credentials_refresh_callback; + std::vector block_ids; using MemoryBufferPtr = std::unique_ptr>; diff --git a/src/Disks/IO/WriteBufferFromAzureDataLakeStorage.cpp b/src/Disks/IO/WriteBufferFromAzureDataLakeStorage.cpp index 6102225ee53c..9be7df094836 100644 --- a/src/Disks/IO/WriteBufferFromAzureDataLakeStorage.cpp +++ b/src/Disks/IO/WriteBufferFromAzureDataLakeStorage.cpp @@ -121,13 +121,15 @@ WriteBufferFromAzureDataLakeStorage::WriteBufferFromAzureDataLakeStorage( const WriteSettings & write_settings_, std::shared_ptr settings_, const String & container_for_logging_, - BlobStorageLogWriterPtr blob_log_) + BlobStorageLogWriterPtr blob_log_, + FileClientRefreshCallback credentials_refresh_callback_) : WriteBufferFromFileBase(buf_size_, nullptr, 0) , log(getLogger("WriteBufferFromAzureDataLakeStorage")) , file_client(makeAdlsGen2FileClient(endpoint_, auth_method_, blob_client_options_, blob_path_)) , blob_path(blob_path_) , write_settings(write_settings_) , max_unexpected_write_error_retries(settings_->max_unexpected_write_error_retries) + , credentials_refresh_callback(std::move(credentials_refresh_callback_)) , container_for_logging(container_for_logging_) , blob_log(std::move(blob_log_)) { @@ -145,6 +147,21 @@ WriteBufferFromAzureDataLakeStorage::~WriteBufferFromAzureDataLakeStorage() } } +bool WriteBufferFromAzureDataLakeStorage::tryRefreshCredentials(const Azure::Core::RequestFailedException & e) +{ + if (credentials_refreshed || !credentials_refresh_callback || !isAzureAccessTokenExpiredError(e)) + return false; + + auto new_file_client = credentials_refresh_callback(); + if (!new_file_client) + return false; + + file_client = std::move(*new_file_client); + credentials_refreshed = true; + LOG_DEBUG(log, "Refreshed Azure credentials for `{}` after an authentication failure", blob_path); + return true; +} + void WriteBufferFromAzureDataLakeStorage::runWithRetries( const std::function & op, const char * what, @@ -179,6 +196,13 @@ void WriteBufferFromAzureDataLakeStorage::runWithRetries( } catch (const Azure::Core::RequestFailedException & e) { + /// Credentials are refreshed at most once, so the retry after it costs no attempt. + if (tryRefreshCredentials(e)) + { + --attempt; + continue; + } + const bool retryable = isRetryableAzureException(e); if (!retryable || attempt >= max_unexpected_write_error_retries) { @@ -301,7 +325,11 @@ void WriteBufferFromAzureDataLakeStorage::cancelImpl() noexcept try { LOG_INFO(log, "Deleting incomplete ADLS Gen2 file `{}` after cancel", blob_path); - file_client.DeleteIfExists(); + runWithRetries( + [&]() { file_client.DeleteIfExists(); }, + "DeleteIfExists", + BlobStorageLogElement::EventType::Delete, + /*data_size=*/ 0); } catch (...) { diff --git a/src/Disks/IO/WriteBufferFromAzureDataLakeStorage.h b/src/Disks/IO/WriteBufferFromAzureDataLakeStorage.h index be4acef948b2..8f625fcd4071 100644 --- a/src/Disks/IO/WriteBufferFromAzureDataLakeStorage.h +++ b/src/Disks/IO/WriteBufferFromAzureDataLakeStorage.h @@ -5,6 +5,7 @@ #if USE_AZURE_BLOB_STORAGE #include +#include #include #include @@ -21,6 +22,8 @@ namespace DB class WriteBufferFromAzureDataLakeStorage : public WriteBufferFromFileBase { public: + using FileClientRefreshCallback = std::function()>; + WriteBufferFromAzureDataLakeStorage( const AzureBlobStorage::Endpoint & endpoint_, const AzureBlobStorage::AuthMethod & auth_method_, @@ -30,7 +33,8 @@ class WriteBufferFromAzureDataLakeStorage : public WriteBufferFromFileBase const WriteSettings & write_settings_, std::shared_ptr settings_, const String & container_for_logging_ = {}, - BlobStorageLogWriterPtr blob_log_ = {}); + BlobStorageLogWriterPtr blob_log_ = {}, + FileClientRefreshCallback credentials_refresh_callback_ = {}); ~WriteBufferFromAzureDataLakeStorage() override; @@ -50,6 +54,8 @@ class WriteBufferFromAzureDataLakeStorage : public WriteBufferFromFileBase BlobStorageLogElement::EventType event_type, size_t data_size); + bool tryRefreshCredentials(const Azure::Core::RequestFailedException & e); + LoggerPtr log; Azure::Storage::Files::DataLake::DataLakeFileClient file_client; @@ -57,6 +63,9 @@ class WriteBufferFromAzureDataLakeStorage : public WriteBufferFromFileBase const WriteSettings write_settings; const size_t max_unexpected_write_error_retries; + const FileClientRefreshCallback credentials_refresh_callback; + bool credentials_refreshed = false; + bool file_created = false; bool is_prefinalized = false; int64_t bytes_appended = 0; diff --git a/src/IO/AzureBlobStorage/isRetryableAzureException.cpp b/src/IO/AzureBlobStorage/isRetryableAzureException.cpp index 923ef5d3f4a4..9f82b4f2fa53 100644 --- a/src/IO/AzureBlobStorage/isRetryableAzureException.cpp +++ b/src/IO/AzureBlobStorage/isRetryableAzureException.cpp @@ -24,6 +24,12 @@ bool isRetryableAzureException(const Azure::Core::RequestFailedException & e) return e.StatusCode >= Azure::Core::Http::HttpStatusCode::InternalServerError; } +bool isAzureAccessTokenExpiredError(const Azure::Core::RequestFailedException & e) +{ + return e.StatusCode == Azure::Core::Http::HttpStatusCode::Unauthorized + || e.StatusCode == Azure::Core::Http::HttpStatusCode::Forbidden; +} + } #endif diff --git a/src/IO/AzureBlobStorage/isRetryableAzureException.h b/src/IO/AzureBlobStorage/isRetryableAzureException.h index dfd13e4c98a0..aa11d2d16473 100644 --- a/src/IO/AzureBlobStorage/isRetryableAzureException.h +++ b/src/IO/AzureBlobStorage/isRetryableAzureException.h @@ -9,6 +9,8 @@ namespace DB bool isRetryableAzureException(const Azure::Core::RequestFailedException & e); +bool isAzureAccessTokenExpiredError(const Azure::Core::RequestFailedException & e); + } #endif diff --git a/src/Storages/ObjectStorage/Azure/Configuration.cpp b/src/Storages/ObjectStorage/Azure/Configuration.cpp index abb2f3d13ece..53753fb0b5bc 100644 --- a/src/Storages/ObjectStorage/Azure/Configuration.cpp +++ b/src/Storages/ObjectStorage/Azure/Configuration.cpp @@ -11,6 +11,7 @@ #include #include #include +#include #include #include #include @@ -90,7 +91,7 @@ StorageObjectStorageQuerySettings StorageAzureConfiguration::getQuerySettings(co }; } -ObjectStoragePtr StorageAzureConfiguration::createObjectStorage(ContextPtr context, bool is_readonly, CredentialsConfigurationCallback /*refresh_credentials_callback*/) /// NOLINT +ObjectStoragePtr StorageAzureConfiguration::createObjectStorage(ContextPtr context, bool is_readonly, CredentialsConfigurationCallback refresh_credentials_callback) /// NOLINT { assertInitialized(); check(context); @@ -98,6 +99,30 @@ ObjectStoragePtr StorageAzureConfiguration::createObjectStorage(ContextPtr conte auto settings = AzureBlobStorage::getRequestSettings(context->getSettingsRef()); auto client = AzureBlobStorage::getContainerClient(connection_params, is_readonly); + AzureObjectStorage::AzureCredentialsRefreshCallback credentials_refresher; + if (refresh_credentials_callback) + { + credentials_refresher = [refresh_credentials_callback, params = connection_params]() -> std::optional + { + auto new_credentials = (*refresh_credentials_callback)(); + if (!new_credentials) + return {}; + + auto azure_credentials = std::dynamic_pointer_cast(new_credentials); + if (!azure_credentials) + throw Exception(ErrorCodes::BAD_ARGUMENTS, "Unexpected credentials type for Azure storage"); + if (azure_credentials->isEmpty()) + return {}; + + auto new_params = params; + std::string sas = azure_credentials->getToken(); + if (!sas.empty() && sas.front() == '?') + sas.erase(0, 1); + new_params.endpoint.sas_auth = std::move(sas); + return new_params; + }; + } + return std::make_unique( "AzureBlobStorage", connection_params.auth_method, @@ -106,7 +131,8 @@ ObjectStoragePtr StorageAzureConfiguration::createObjectStorage(ContextPtr conte connection_params, connection_params.getContainer(), connection_params.getConnectionURL(), - /*common_key_prefix*/ ""); + /*common_key_prefix*/ "", + std::move(credentials_refresher)); } AzureBlobStorage::ConnectionParams getAzureConnectionParams( diff --git a/tests/integration/compose/docker_compose_iceberg_lakekeeper_catalog.yml b/tests/integration/compose/docker_compose_iceberg_lakekeeper_catalog.yml index a2f68b606c77..7f2951ddc986 100644 --- a/tests/integration/compose/docker_compose_iceberg_lakekeeper_catalog.yml +++ b/tests/integration/compose/docker_compose_iceberg_lakekeeper_catalog.yml @@ -1,6 +1,6 @@ services: lakekeeper: - image: vakamo/lakekeeper:v0.9.4 + image: vakamo/lakekeeper:v0.13.1 environment: - LAKEKEEPER__PG_ENCRYPTION_KEY=This-is-NOT-Secure! - LAKEKEEPER__PG_DATABASE_URL_READ=postgresql://postgres:postgres@db:5432/postgres @@ -23,7 +23,7 @@ services: cpus: 3 migrate: - image: vakamo/lakekeeper:v0.9.4 + image: vakamo/lakekeeper:v0.13.1 environment: - LAKEKEEPER__PG_ENCRYPTION_KEY=This-is-NOT-Secure! - LAKEKEEPER__PG_DATABASE_URL_READ=postgresql://postgres:postgres@db:5432/postgres diff --git a/tests/integration/test_database_iceberg_lakekeeper_catalog/test.py b/tests/integration/test_database_iceberg_lakekeeper_catalog/test.py index d8ac831a78d8..77a5c1f6f492 100644 --- a/tests/integration/test_database_iceberg_lakekeeper_catalog/test.py +++ b/tests/integration/test_database_iceberg_lakekeeper_catalog/test.py @@ -547,3 +547,125 @@ def test_auth_token_profile_events(started_cluster): refreshed, cache_hits = get_auth_token_profile_events(node, qid2) assert refreshed == 0 and cache_hits >= 1 + +def get_credentials_profile_events(node, query_id): + node.query("SYSTEM FLUSH LOGS") + result = node.query( + "SELECT ProfileEvents['DataLakeRestCatalogCredentialsVended'], " + "ProfileEvents['DataLakeRestCatalogCredentialsCacheHits'] " + f"FROM system.query_log WHERE query_id = '{query_id}' AND type = 'QueryFinish'" + ) + return tuple(int(value) for value in result.split()) + + +def create_int_table(catalog, namespace, table_name, rows=1): + if namespace not in catalog.list_namespaces(): + catalog.create_namespace(namespace) + table = catalog.create_table( + namespace + (table_name,), + schema=Schema(NestedField(field_id=1, name="id", field_type=IntegerType(), required=False)), + properties={"write.metadata.compression-codec": "none"}, + ) + table.append(pa.table({"id": pa.array(range(rows), type=pa.int32())})) + + +def test_vended_credentials_cache_disabled(started_cluster): + node = started_cluster.instances["node1"] + catalog = load_catalog_impl(started_cluster) + + test_ref = f"test_vended_credentials_cache_disabled_{uuid.uuid4().hex[:8]}" + namespace = (f"{test_ref}_namespace",) + table_name = f"{test_ref}_table" + db_name = f"{test_ref}_database" + + create_int_table(catalog, namespace, table_name) + create_clickhouse_iceberg_database( + started_cluster, node, db_name, + additional_settings={"vended_credentials_cache_ttl": 0}, + ) + query = f"SELECT count() FROM {db_name}.`{namespace[0]}.{table_name}`" + + for attempt in range(2): + qid = f"{test_ref}-{attempt}" + node.query(query, query_id=qid) + vended, hits = get_credentials_profile_events(node, qid) + assert vended >= 1 and hits == 0 + + +def test_vended_credentials_cache_invalidated_on_table_replace(started_cluster): + node = started_cluster.instances["node1"] + catalog = load_catalog_impl(started_cluster) + + test_ref = f"test_vended_credentials_cache_replace_{uuid.uuid4().hex[:8]}" + namespace = (f"{test_ref}_namespace",) + table_name = f"{test_ref}_table" + db_name = f"{test_ref}_database" + + create_int_table(catalog, namespace, table_name) + create_clickhouse_iceberg_database(started_cluster, node, db_name) + query = f"SELECT count() FROM {db_name}.`{namespace[0]}.{table_name}`" + + node.query(query) + qid = f"{test_ref}-2" + node.query(query, query_id=qid) + vended, hits = get_credentials_profile_events(node, qid) + assert vended == 0 and hits >= 1 + + catalog.drop_table(namespace + (table_name,)) + create_int_table(catalog, namespace, table_name, rows=2) + + qid = f"{test_ref}-3" + assert node.query(query, query_id=qid).strip() == "2" + vended, _ = get_credentials_profile_events(node, qid) + assert vended >= 1 + + qid = f"{test_ref}-4" + node.query(query, query_id=qid) + vended, hits = get_credentials_profile_events(node, qid) + assert vended == 0 and hits >= 1 + + +def test_vended_credentials_cache_cleared_on_auth_change(started_cluster): + node = started_cluster.instances["node1"] + catalog = load_catalog_impl(started_cluster) + + test_ref = f"test_vended_credentials_cache_auth_{uuid.uuid4().hex[:8]}" + namespace = (f"{test_ref}_namespace",) + table_name = f"{test_ref}_table" + db_name = f"{test_ref}_database" + + create_int_table(catalog, namespace, table_name) + + node.query(f"DROP DATABASE IF EXISTS {db_name}") + node.query( + f""" + CREATE DATABASE {db_name} + ENGINE = DataLakeCatalog('{BASE_URL}') + SETTINGS + catalog_type = 'rest', + warehouse = 'demo', + storage_endpoint = 'http://minio1:9001/warehouse-rest', + auth_header = 'Authorization: Bearer initial_dummy' + """, + settings={"allow_experimental_database_iceberg": 1}, + ) + + query = f"SELECT count() FROM {db_name}.`{namespace[0]}.{table_name}`" + + node.query(query) + qid = f"{test_ref}-2" + node.query(query, query_id=qid) + vended, hits = get_credentials_profile_events(node, qid) + assert vended == 0 and hits >= 1 + + node.query( + f"ALTER DATABASE {db_name} MODIFY SETTING auth_header = 'Authorization: Bearer altered_dummy'" + ) + + qid = f"{test_ref}-3" + node.query(query, query_id=qid) + vended, _ = get_credentials_profile_events(node, qid) + assert vended >= 1 + + node.query(f"DROP DATABASE IF EXISTS {db_name}") +