Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions src/iceberg/catalog/rest/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@ set(ICEBERG_REST_SOURCES
auth/auth_managers.cc
auth/auth_properties.cc
auth/auth_session.cc
auth/auth_session_cache.cc
auth/oauth2_util.cc
auth/sigv4_manager.cc
auth/token_refresh_scheduler.cc
Expand Down
95 changes: 78 additions & 17 deletions src/iceberg/catalog/rest/auth/auth_manager.cc
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@
#include "iceberg/catalog/rest/auth/auth_manager_internal.h"
#include "iceberg/catalog/rest/auth/auth_properties.h"
#include "iceberg/catalog/rest/auth/auth_session.h"
#include "iceberg/catalog/rest/auth/auth_session_cache_internal.h"
#include "iceberg/catalog/rest/auth/auth_session_internal.h"
#include "iceberg/catalog/rest/auth/oauth2_util.h"
#include "iceberg/catalog/session_context.h"
Expand Down Expand Up @@ -176,6 +177,24 @@ class OAuth2Manager : public AuthManager {
ICEBERG_PRECHECK(shared_client != nullptr,
"OAuth2 catalog session HTTP client must not be null");
refresh_client_ = std::move(shared_client);

// Initialize catalog-level configuration
keep_refreshed_ = config.keep_refreshed();
exchange_enabled_ = config.exchange_enabled();
session_timeout_ =
std::chrono::milliseconds(config.Get(AuthProperties::kSessionTimeoutMs));

// Create session cache
ICEBERG_ASSIGN_OR_RAISE(
session_cache_,
internal::AuthSessionCache::Make(
session_timeout_, [](std::shared_ptr<AuthSession> session) {
if (auto oauth2 =
std::dynamic_pointer_cast<internal::OAuth2Session>(session)) {
oauth2->StopRefreshing();
}
}));

// Reuse the token response and start time from the init phase.
if (auth_response_.has_value()) {
return internal::MakeOAuth2Session(
Expand Down Expand Up @@ -217,21 +236,31 @@ class OAuth2Manager : public AuthManager {

Result<std::shared_ptr<AuthSession>> ContextualSession(
const SessionContext& context, std::shared_ptr<AuthSession> parent) override {
// TODO(lishuxu): Add child-session caching and refresh, matching Java
// AuthSessionCache.
// Use session_id as cache key for contextual sessions
std::string cache_key = context.session_id.empty() ? "" : "ctx:" + context.session_id;
return MaybeCreateChildSession(context.credentials, /*allow_credential=*/true,
std::move(parent));
std::move(parent), cache_key);
}

Result<std::shared_ptr<AuthSession>> TableSession(
[[maybe_unused]] const TableIdentifier& table,
const std::unordered_map<std::string, std::string>& properties,
std::shared_ptr<AuthSession> parent) override {
// Use token value as cache key for table sessions
auto token_it = properties.find(AuthProperties::kToken.key());
std::string cache_key;
if (token_it != properties.end() && !token_it->second.empty()) {
cache_key = "tbl:" + token_it->second;
}
return MaybeCreateChildSession(FilterTableSessionProperties(properties),
/*allow_credential=*/false, std::move(parent));
/*allow_credential=*/false, std::move(parent),
cache_key);
}

Status Close() override {
if (session_cache_) {
session_cache_->Close();
}
refresh_client_.reset();
return {};
}
Expand Down Expand Up @@ -267,7 +296,8 @@ class OAuth2Manager : public AuthManager {

Result<std::shared_ptr<AuthSession>> MaybeCreateChildSession(
const std::unordered_map<std::string, std::string>& credentials,
bool allow_credential, std::shared_ptr<AuthSession> parent) {
bool allow_credential, std::shared_ptr<AuthSession> parent,
const std::string& cache_key) {
auto token_it = credentials.find(AuthProperties::kToken.key());
auto credential_it = credentials.find(AuthProperties::kCredential.key());
auto typed_token = FindPreferredTypedToken(credentials);
Expand All @@ -283,42 +313,73 @@ class OAuth2Manager : public AuthManager {
ICEBERG_PRECHECK(parent_info.has_value(),
"OAuth2 child session requires OAuth2 parent metadata");

// Use cache if we have a valid cache_key
if (!cache_key.empty() && session_cache_) {
return session_cache_->Get(cache_key,
[&]() -> Result<std::shared_ptr<AuthSession>> {
return CreateChildSessionUncached(
credentials, allow_credential, *parent_info);
});
}

// No cache, create session directly
return CreateChildSessionUncached(credentials, allow_credential, *parent_info);
}

Result<std::shared_ptr<AuthSession>> CreateChildSessionUncached(
const std::unordered_map<std::string, std::string>& credentials,
bool allow_credential, const OAuth2SessionInfo& parent_info) {
auto token_it = credentials.find(AuthProperties::kToken.key());
auto credential_it = credentials.find(AuthProperties::kCredential.key());
auto typed_token = FindPreferredTypedToken(credentials);

if (token_it != credentials.end()) {
ICEBERG_ASSIGN_OR_RAISE(auto config,
ChildConfig(*parent_info, parent_info->credential));
ChildConfig(parent_info, parent_info.credential));
return MakeSession(AccessTokenResponse(token_it->second), config,
/*keep_refreshed=*/false);
}

if (allow_credential && credential_it != credentials.end()) {
ICEBERG_ASSIGN_OR_RAISE(auto config,
ChildConfig(*parent_info, credential_it->second));
ICEBERG_ASSIGN_OR_RAISE(auto response,
OAuth2Util::FetchToken(*refresh_client_, *parent, config));
return MakeSession(response, config, /*keep_refreshed=*/false);
ChildConfig(parent_info, credential_it->second));
// Create a temporary parent session for FetchToken
auto temp_parent = AuthSession::MakeDefault(OAuth2Util::AuthHeaders(""));
ICEBERG_ASSIGN_OR_RAISE(
auto response, OAuth2Util::FetchToken(*refresh_client_, *temp_parent, config));
// Credential child sessions use keep_refreshed from catalog level
return MakeSession(response, config, keep_refreshed_);
}

std::optional<std::string> actor_token;
std::optional<std::string> actor_token_type;
if (!parent_info->token.empty()) {
actor_token = parent_info->token;
actor_token_type = parent_info->issued_token_type;
if (!parent_info.token.empty()) {
actor_token = parent_info.token;
actor_token_type = parent_info.issued_token_type;
}
// Create a temporary parent session for ExchangeToken
auto temp_parent = AuthSession::MakeDefault(OAuth2Util::AuthHeaders(""));
ICEBERG_ASSIGN_OR_RAISE(
auto response,
OAuth2Util::ExchangeToken(*refresh_client_, *parent, {}, typed_token->second,
OAuth2Util::ExchangeToken(*refresh_client_, *temp_parent, {}, typed_token->second,
typed_token->first, actor_token, actor_token_type,
parent_info->scope, parent_info->oauth2_server_uri,
parent_info->optional_oauth_params));
parent_info.scope, parent_info.oauth2_server_uri,
parent_info.optional_oauth_params));
ICEBERG_ASSIGN_OR_RAISE(auto config,
ChildConfig(*parent_info, parent_info->credential));
ChildConfig(parent_info, parent_info.credential));
return MakeSession(response, config, /*keep_refreshed=*/false);
}

/// Token response and start time captured by InitSession.
std::optional<OAuthTokenResponse> auth_response_;
std::optional<std::chrono::steady_clock::time_point> start_time_;
std::shared_ptr<HttpClient> refresh_client_;

// Catalog-level configuration initialized in CatalogSession
bool keep_refreshed_ = true;
bool exchange_enabled_ = true;
std::chrono::milliseconds session_timeout_{3'600'000};
std::shared_ptr<internal::AuthSessionCache> session_cache_;
};

Result<std::unique_ptr<AuthManager>> MakeOAuth2Manager(
Expand Down
4 changes: 4 additions & 0 deletions src/iceberg/catalog/rest/auth/auth_properties.h
Original file line number Diff line number Diff line change
Expand Up @@ -82,6 +82,10 @@ class ICEBERG_REST_EXPORT AuthProperties : public ConfigBase<AuthProperties> {
inline static Entry<std::string> kAudience{"audience", ""};
inline static Entry<std::string> kResource{"resource", ""};

/// Session cache timeout in milliseconds. Sessions will be eligible for eviction
/// after this duration of inactivity. Default is 1 hour (3,600,000 ms).
inline static Entry<int64_t> kSessionTimeoutMs{"auth.session-timeout-ms", 3'600'000};

// ---- OAuth2 token type constants ----

inline static const std::string kAccessTokenType =
Expand Down
Loading
Loading