use std::collections::HashMap;
use std::fmt;
use std::sync::Arc;
use std::time::Duration;
#[cfg(test)]
use mock_instant::thread_local::{SystemTime, UNIX_EPOCH};
#[cfg(not(test))]
use std::time::{SystemTime, UNIX_EPOCH};
use async_trait::async_trait;
use lance_namespace::LanceNamespace;
use lance_namespace::models::DescribeTableRequest;
use tokio::sync::RwLock;
use crate::{Error, Result};
pub const EXPIRES_AT_MILLIS_KEY: &str = "expires_at_millis";
pub const REFRESH_OFFSET_MILLIS_KEY: &str = "refresh_offset_millis";
const DEFAULT_REFRESH_OFFSET_MILLIS: u64 = 60_000;
#[async_trait]
pub trait StorageOptionsProvider: Send + Sync + fmt::Debug {
async fn fetch_storage_options(&self) -> Result<Option<HashMap<String, String>>>;
async fn force_fetch_storage_options(&self) -> Result<Option<HashMap<String, String>>> {
self.fetch_storage_options().await
}
fn provider_id(&self) -> String;
}
pub struct LanceNamespaceStorageOptionsProvider {
namespace_client: Arc<dyn LanceNamespace>,
table_id: Vec<String>,
}
impl fmt::Debug for LanceNamespaceStorageOptionsProvider {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.provider_id())
}
}
impl fmt::Display for LanceNamespaceStorageOptionsProvider {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.provider_id())
}
}
impl LanceNamespaceStorageOptionsProvider {
pub fn new(namespace_client: Arc<dyn LanceNamespace>, table_id: Vec<String>) -> Self {
Self {
namespace_client,
table_id,
}
}
}
#[async_trait]
impl StorageOptionsProvider for LanceNamespaceStorageOptionsProvider {
async fn fetch_storage_options(&self) -> Result<Option<HashMap<String, String>>> {
let request = DescribeTableRequest {
id: Some(self.table_id.clone()),
..Default::default()
};
let response = self
.namespace_client
.describe_table(request)
.await
.map_err(|e| {
Error::io_source(Box::new(std::io::Error::other(format!(
"Failed to fetch storage options: {}",
e
))))
})?;
Ok(response.storage_options)
}
fn provider_id(&self) -> String {
format!(
"LanceNamespaceStorageOptionsProvider {{ namespace_client: {}, table_id: {:?} }}",
self.namespace_client.namespace_id(),
self.table_id
)
}
}
pub const BASE_SCOPED_OPTION_PREFIX: &str = "base_";
pub fn parse_base_scoped_key(key: &str) -> Option<(u32, &str)> {
let rest = key.strip_prefix(BASE_SCOPED_OPTION_PREFIX)?;
let (id_str, scoped_key) = rest.split_once('.')?;
if scoped_key.is_empty() || id_str.is_empty() || !id_str.bytes().all(|b| b.is_ascii_digit()) {
return None;
}
let id = id_str.parse::<u32>().ok()?;
Some((id, scoped_key))
}
pub fn has_base_scoped_options(options: &HashMap<String, String>) -> bool {
options
.keys()
.any(|key| parse_base_scoped_key(key).is_some())
}
pub fn resolve_base_scoped_options(
options: &HashMap<String, String>,
base_id: Option<u32>,
) -> HashMap<String, String> {
let mut resolved = HashMap::with_capacity(options.len());
let mut overrides = Vec::new();
for (key, value) in options {
match parse_base_scoped_key(key) {
Some((id, scoped_key)) => {
if Some(id) == base_id {
overrides.push((scoped_key.to_string(), value.clone()));
}
}
None => {
resolved.insert(key.clone(), value.clone());
}
}
}
resolved.extend(overrides);
resolved
}
#[derive(Debug)]
pub struct BaseScopedStorageOptionsProvider {
inner: Arc<StorageOptionsAccessor>,
base_id: Option<u32>,
}
impl BaseScopedStorageOptionsProvider {
pub fn new(inner: Arc<StorageOptionsAccessor>, base_id: Option<u32>) -> Self {
Self { inner, base_id }
}
}
#[async_trait]
impl StorageOptionsProvider for BaseScopedStorageOptionsProvider {
async fn fetch_storage_options(&self) -> Result<Option<HashMap<String, String>>> {
let options = self.inner.get_storage_options().await?;
Ok(Some(resolve_base_scoped_options(&options.0, self.base_id)))
}
async fn force_fetch_storage_options(&self) -> Result<Option<HashMap<String, String>>> {
let options = self.inner.refresh_storage_options().await?;
Ok(Some(resolve_base_scoped_options(&options.0, self.base_id)))
}
fn provider_id(&self) -> String {
match self.base_id {
Some(id) => format!("base-scoped[base_id={}]({})", id, self.inner.accessor_id()),
None => format!("base-scoped[default]({})", self.inner.accessor_id()),
}
}
}
pub struct StorageOptionsAccessor {
initial_options: Option<HashMap<String, String>>,
provider: Option<Arc<dyn StorageOptionsProvider>>,
cache: Arc<RwLock<Option<CachedStorageOptions>>>,
refresh_offset: Duration,
scope_resolved: bool,
}
impl fmt::Debug for StorageOptionsAccessor {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("StorageOptionsAccessor")
.field("has_initial_options", &self.initial_options.is_some())
.field("has_provider", &self.provider.is_some())
.field("refresh_offset", &self.refresh_offset)
.finish()
}
}
#[derive(Debug, Clone)]
struct CachedStorageOptions {
options: HashMap<String, String>,
expires_at_millis: Option<u64>,
}
impl StorageOptionsAccessor {
fn extract_refresh_offset(options: &HashMap<String, String>) -> Duration {
options
.get(REFRESH_OFFSET_MILLIS_KEY)
.and_then(|s| s.parse::<u64>().ok())
.map(Duration::from_millis)
.unwrap_or(Duration::from_millis(DEFAULT_REFRESH_OFFSET_MILLIS))
}
fn effective_expires_at_millis(options: &HashMap<String, String>) -> Option<u64> {
options
.iter()
.filter(|(key, _)| {
key.as_str() == EXPIRES_AT_MILLIS_KEY
|| matches!(
parse_base_scoped_key(key),
Some((_, scoped_key)) if scoped_key == EXPIRES_AT_MILLIS_KEY
)
})
.filter_map(|(_, value)| value.parse::<u64>().ok())
.min()
}
pub fn with_static_options(options: HashMap<String, String>) -> Self {
let expires_at_millis = Self::effective_expires_at_millis(&options);
let refresh_offset = Self::extract_refresh_offset(&options);
Self {
initial_options: Some(options.clone()),
provider: None,
cache: Arc::new(RwLock::new(Some(CachedStorageOptions {
options,
expires_at_millis,
}))),
refresh_offset,
scope_resolved: false,
}
}
pub fn with_provider(provider: Arc<dyn StorageOptionsProvider>) -> Self {
Self {
initial_options: None,
provider: Some(provider),
cache: Arc::new(RwLock::new(None)),
refresh_offset: Duration::from_millis(DEFAULT_REFRESH_OFFSET_MILLIS),
scope_resolved: false,
}
}
pub fn with_initial_and_provider(
initial_options: HashMap<String, String>,
provider: Arc<dyn StorageOptionsProvider>,
) -> Self {
let expires_at_millis = Self::effective_expires_at_millis(&initial_options);
let refresh_offset = Self::extract_refresh_offset(&initial_options);
Self {
initial_options: Some(initial_options.clone()),
provider: Some(provider),
cache: Arc::new(RwLock::new(Some(CachedStorageOptions {
options: initial_options,
expires_at_millis,
}))),
refresh_offset,
scope_resolved: false,
}
}
pub async fn get_storage_options(&self) -> Result<super::StorageOptions> {
loop {
match self.do_get_storage_options().await? {
Some(options) => return Ok(options),
None => {
tokio::time::sleep(Duration::from_millis(10)).await;
continue;
}
}
}
}
pub(crate) async fn refresh_storage_options(&self) -> Result<super::StorageOptions> {
let Some(provider) = &self.provider else {
return self.get_storage_options().await;
};
log::debug!(
"Refreshing storage options from provider: {}",
provider.provider_id()
);
let storage_options_map = provider.force_fetch_storage_options().await.map_err(|e| {
Error::io_source(Box::new(std::io::Error::other(format!(
"Failed to fetch storage options: {}",
e
))))
})?;
let Some(options) = storage_options_map else {
if let Some(initial) = &self.initial_options {
return Ok(super::StorageOptions(initial.clone()));
}
log::debug!(
"Provider {} returned no storage options, using default credentials",
provider.provider_id()
);
return Ok(super::StorageOptions(HashMap::new()));
};
let expires_at_millis = Self::effective_expires_at_millis(&options);
let mut cache = self.cache.write().await;
*cache = Some(CachedStorageOptions {
options: options.clone(),
expires_at_millis,
});
Ok(super::StorageOptions(options))
}
async fn do_get_storage_options(&self) -> Result<Option<super::StorageOptions>> {
{
let cached = self.cache.read().await;
if !self.needs_refresh(&cached)
&& let Some(cached_opts) = &*cached
{
return Ok(Some(super::StorageOptions(cached_opts.options.clone())));
}
}
let Some(provider) = &self.provider else {
return if let Some(initial) = &self.initial_options {
Ok(Some(super::StorageOptions(initial.clone())))
} else {
Ok(Some(super::StorageOptions(HashMap::new())))
};
};
let Ok(mut cache) = self.cache.try_write() else {
return Ok(None);
};
if !self.needs_refresh(&cache)
&& let Some(cached_opts) = &*cache
{
return Ok(Some(super::StorageOptions(cached_opts.options.clone())));
}
log::debug!(
"Refreshing storage options from provider: {}",
provider.provider_id()
);
let storage_options_map = provider.fetch_storage_options().await.map_err(|e| {
Error::io_source(Box::new(std::io::Error::other(format!(
"Failed to fetch storage options: {}",
e
))))
})?;
let Some(options) = storage_options_map else {
if let Some(initial) = &self.initial_options {
return Ok(Some(super::StorageOptions(initial.clone())));
}
log::debug!(
"Provider {} returned no storage options, using default credentials",
provider.provider_id()
);
return Ok(Some(super::StorageOptions(HashMap::new())));
};
let expires_at_millis = Self::effective_expires_at_millis(&options);
if let Some(expires_at) = expires_at_millis {
let now_ms = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or(Duration::from_secs(0))
.as_millis() as u64;
let expires_in_secs = (expires_at.saturating_sub(now_ms)) / 1000;
log::debug!(
"Successfully refreshed storage options from provider: {}, options expire in {} seconds",
provider.provider_id(),
expires_in_secs
);
} else {
log::debug!(
"Successfully refreshed storage options from provider: {} (no expiration)",
provider.provider_id()
);
}
*cache = Some(CachedStorageOptions {
options: options.clone(),
expires_at_millis,
});
Ok(Some(super::StorageOptions(options)))
}
fn needs_refresh(&self, cached: &Option<CachedStorageOptions>) -> bool {
match cached {
None => true,
Some(cached_opts) => {
if let Some(expires_at_millis) = cached_opts.expires_at_millis {
let now_ms = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or(Duration::from_secs(0))
.as_millis() as u64;
let refresh_offset_millis = self.refresh_offset.as_millis() as u64;
now_ms + refresh_offset_millis >= expires_at_millis
} else {
false
}
}
}
}
pub fn initial_storage_options(&self) -> Option<&HashMap<String, String>> {
self.initial_options.as_ref()
}
pub fn accessor_id(&self) -> String {
if let Some(provider) = &self.provider {
provider.provider_id()
} else if let Some(initial) = &self.initial_options {
use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
let mut hasher = DefaultHasher::new();
let mut keys: Vec<_> = initial.keys().collect();
keys.sort();
for key in keys {
key.hash(&mut hasher);
initial.get(key).hash(&mut hasher);
}
format!("static_options_{:x}", hasher.finish())
} else {
"empty_accessor".to_string()
}
}
pub fn scoped_to_base(self: &Arc<Self>, base_id: Option<u32>) -> Arc<Self> {
if self.scope_resolved {
return self.clone();
}
if self.has_provider() {
let provider = Arc::new(BaseScopedStorageOptionsProvider::new(self.clone(), base_id));
let mut scoped = match self.initial_storage_options() {
Some(initial) => Self::with_initial_and_provider(
resolve_base_scoped_options(initial, base_id),
provider,
),
None => Self::with_provider(provider),
};
scoped.scope_resolved = true;
Arc::new(scoped)
} else {
match self.initial_storage_options() {
Some(initial) if has_base_scoped_options(initial) => {
let mut scoped =
Self::with_static_options(resolve_base_scoped_options(initial, base_id));
scoped.scope_resolved = true;
Arc::new(scoped)
}
_ => self.clone(),
}
}
}
pub fn has_provider(&self) -> bool {
self.provider.is_some()
}
pub fn refresh_offset(&self) -> Duration {
self.refresh_offset
}
pub fn provider(&self) -> Option<&Arc<dyn StorageOptionsProvider>> {
self.provider.as_ref()
}
}
#[cfg(test)]
mod tests {
use super::*;
use mock_instant::thread_local::MockClock;
#[derive(Debug)]
struct MockStorageOptionsProvider {
call_count: Arc<RwLock<usize>>,
expires_in_millis: Option<u64>,
}
impl MockStorageOptionsProvider {
fn new(expires_in_millis: Option<u64>) -> Self {
Self {
call_count: Arc::new(RwLock::new(0)),
expires_in_millis,
}
}
async fn get_call_count(&self) -> usize {
*self.call_count.read().await
}
}
#[async_trait]
impl StorageOptionsProvider for MockStorageOptionsProvider {
async fn fetch_storage_options(&self) -> Result<Option<HashMap<String, String>>> {
let count = {
let mut c = self.call_count.write().await;
*c += 1;
*c
};
let mut options = HashMap::from([
("aws_access_key_id".to_string(), format!("AKID_{}", count)),
(
"aws_secret_access_key".to_string(),
format!("SECRET_{}", count),
),
("aws_session_token".to_string(), format!("TOKEN_{}", count)),
]);
if let Some(expires_in) = self.expires_in_millis {
let now_ms = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_millis() as u64;
let expires_at = now_ms + expires_in;
options.insert(EXPIRES_AT_MILLIS_KEY.to_string(), expires_at.to_string());
}
Ok(Some(options))
}
fn provider_id(&self) -> String {
let ptr = Arc::as_ptr(&self.call_count) as usize;
format!("MockStorageOptionsProvider {{ id: {} }}", ptr)
}
}
#[tokio::test]
async fn test_static_options_only() {
let options = HashMap::from([
("key1".to_string(), "value1".to_string()),
("key2".to_string(), "value2".to_string()),
]);
let accessor = StorageOptionsAccessor::with_static_options(options.clone());
let result = accessor.get_storage_options().await.unwrap();
assert_eq!(result.0, options);
assert!(!accessor.has_provider());
assert_eq!(accessor.initial_storage_options(), Some(&options));
}
#[tokio::test]
async fn test_provider_only() {
MockClock::set_system_time(Duration::from_secs(100_000));
let mock_provider = Arc::new(MockStorageOptionsProvider::new(Some(600_000)));
let accessor = StorageOptionsAccessor::with_provider(mock_provider.clone());
let result = accessor.get_storage_options().await.unwrap();
assert!(result.0.contains_key("aws_access_key_id"));
assert_eq!(result.0.get("aws_access_key_id").unwrap(), "AKID_1");
assert!(accessor.has_provider());
assert_eq!(accessor.initial_storage_options(), None);
assert_eq!(mock_provider.get_call_count().await, 1);
}
#[tokio::test]
async fn test_initial_and_provider_uses_initial_first() {
MockClock::set_system_time(Duration::from_secs(100_000));
let now_ms = MockClock::system_time().as_millis() as u64;
let expires_at = now_ms + 600_000;
let initial = HashMap::from([
("aws_access_key_id".to_string(), "INITIAL_KEY".to_string()),
(
"aws_secret_access_key".to_string(),
"INITIAL_SECRET".to_string(),
),
(EXPIRES_AT_MILLIS_KEY.to_string(), expires_at.to_string()),
]);
let mock_provider = Arc::new(MockStorageOptionsProvider::new(Some(600_000)));
let accessor = StorageOptionsAccessor::with_initial_and_provider(
initial.clone(),
mock_provider.clone(),
);
let result = accessor.get_storage_options().await.unwrap();
assert_eq!(result.0.get("aws_access_key_id").unwrap(), "INITIAL_KEY");
assert_eq!(mock_provider.get_call_count().await, 0); }
#[tokio::test]
async fn test_caching_and_refresh() {
MockClock::set_system_time(Duration::from_secs(100_000));
let mock_provider = Arc::new(MockStorageOptionsProvider::new(Some(600_000))); let now_ms = MockClock::system_time().as_millis() as u64;
let expires_at = now_ms + 600_000; let initial = HashMap::from([
(EXPIRES_AT_MILLIS_KEY.to_string(), expires_at.to_string()),
(REFRESH_OFFSET_MILLIS_KEY.to_string(), "300000".to_string()), ]);
let accessor =
StorageOptionsAccessor::with_initial_and_provider(initial, mock_provider.clone());
let result = accessor.get_storage_options().await.unwrap();
assert!(result.0.contains_key(EXPIRES_AT_MILLIS_KEY));
assert_eq!(mock_provider.get_call_count().await, 0);
MockClock::set_system_time(Duration::from_secs(100_000 + 360));
let result = accessor.get_storage_options().await.unwrap();
assert_eq!(result.0.get("aws_access_key_id").unwrap(), "AKID_1");
assert_eq!(mock_provider.get_call_count().await, 1);
}
#[tokio::test]
async fn test_expired_initial_triggers_refresh() {
MockClock::set_system_time(Duration::from_secs(100_000));
let now_ms = MockClock::system_time().as_millis() as u64;
let expired_time = now_ms - 1_000;
let initial = HashMap::from([
("aws_access_key_id".to_string(), "EXPIRED_KEY".to_string()),
(EXPIRES_AT_MILLIS_KEY.to_string(), expired_time.to_string()),
]);
let mock_provider = Arc::new(MockStorageOptionsProvider::new(Some(600_000)));
let accessor =
StorageOptionsAccessor::with_initial_and_provider(initial, mock_provider.clone());
let result = accessor.get_storage_options().await.unwrap();
assert_eq!(result.0.get("aws_access_key_id").unwrap(), "AKID_1");
assert_eq!(mock_provider.get_call_count().await, 1);
}
#[tokio::test]
async fn test_accessor_id_with_provider() {
let mock_provider = Arc::new(MockStorageOptionsProvider::new(None));
let accessor = StorageOptionsAccessor::with_provider(mock_provider);
let id = accessor.accessor_id();
assert!(id.starts_with("MockStorageOptionsProvider"));
}
#[tokio::test]
async fn test_accessor_id_static() {
let options = HashMap::from([("key".to_string(), "value".to_string())]);
let accessor = StorageOptionsAccessor::with_static_options(options);
let id = accessor.accessor_id();
assert!(id.starts_with("static_options_"));
}
#[tokio::test]
async fn test_concurrent_access() {
let mock_provider = Arc::new(MockStorageOptionsProvider::new(Some(9999999999999)));
let accessor = Arc::new(StorageOptionsAccessor::with_provider(mock_provider.clone()));
let mut handles = vec![];
for i in 0..10 {
let acc = accessor.clone();
let handle = tokio::spawn(async move {
let result = acc.get_storage_options().await.unwrap();
assert_eq!(result.0.get("aws_access_key_id").unwrap(), "AKID_1");
i
});
handles.push(handle);
}
let results: Vec<_> = futures::future::join_all(handles)
.await
.into_iter()
.map(|r| r.unwrap())
.collect();
assert_eq!(results.len(), 10);
let call_count = mock_provider.get_call_count().await;
assert_eq!(
call_count, 1,
"Provider should be called exactly once despite concurrent access"
);
}
#[tokio::test]
async fn test_no_expiration_never_refreshes() {
MockClock::set_system_time(Duration::from_secs(100_000));
let mock_provider = Arc::new(MockStorageOptionsProvider::new(None)); let accessor = StorageOptionsAccessor::with_provider(mock_provider.clone());
accessor.get_storage_options().await.unwrap();
assert_eq!(mock_provider.get_call_count().await, 1);
MockClock::set_system_time(Duration::from_secs(200_000));
accessor.get_storage_options().await.unwrap();
assert_eq!(mock_provider.get_call_count().await, 1);
}
#[test]
fn test_parse_base_scoped_key() {
assert_eq!(
parse_base_scoped_key("base_1.account_key"),
Some((1, "account_key"))
);
assert_eq!(
parse_base_scoped_key("base_12.headers.x-ms-version"),
Some((12, "headers.x-ms-version"))
);
assert_eq!(parse_base_scoped_key("base_0.region"), Some((0, "region")));
assert_eq!(parse_base_scoped_key("account_key"), None);
assert_eq!(parse_base_scoped_key("base_url"), None);
assert_eq!(parse_base_scoped_key("base_hot.account_key"), None);
assert_eq!(parse_base_scoped_key("base_1x.account_key"), None);
assert_eq!(parse_base_scoped_key("base_+1.account_key"), None);
assert_eq!(parse_base_scoped_key("base_.account_key"), None);
assert_eq!(parse_base_scoped_key("base_1."), None);
assert_eq!(parse_base_scoped_key("base_1"), None);
assert_eq!(parse_base_scoped_key("base_4294967296.key"), None);
}
#[test]
fn test_resolve_base_scoped_options() {
let options = HashMap::from([
("region".to_string(), "us-east-1".to_string()),
("account_key".to_string(), "shared-key".to_string()),
("base_1.account_key".to_string(), "base1-key".to_string()),
("base_2.account_key".to_string(), "base2-key".to_string()),
("base_2.endpoint".to_string(), "http://b2".to_string()),
]);
assert!(has_base_scoped_options(&options));
let base1 = resolve_base_scoped_options(&options, Some(1));
assert_eq!(
base1,
HashMap::from([
("region".to_string(), "us-east-1".to_string()),
("account_key".to_string(), "base1-key".to_string()),
])
);
let base2 = resolve_base_scoped_options(&options, Some(2));
assert_eq!(
base2,
HashMap::from([
("region".to_string(), "us-east-1".to_string()),
("account_key".to_string(), "base2-key".to_string()),
("endpoint".to_string(), "http://b2".to_string()),
])
);
let base3 = resolve_base_scoped_options(&options, Some(3));
assert_eq!(
base3,
HashMap::from([
("region".to_string(), "us-east-1".to_string()),
("account_key".to_string(), "shared-key".to_string()),
])
);
let default = resolve_base_scoped_options(&options, None);
assert_eq!(default, base3);
assert!(!has_base_scoped_options(&HashMap::from([(
"account_key".to_string(),
"shared-key".to_string()
)])));
}
#[tokio::test]
async fn test_scoped_to_base_identity_and_idempotency() {
let accessor = Arc::new(StorageOptionsAccessor::with_static_options(HashMap::from(
[("account_key".to_string(), "shared-key".to_string())],
)));
assert!(Arc::ptr_eq(&accessor.scoped_to_base(Some(1)), &accessor));
assert!(Arc::ptr_eq(&accessor.scoped_to_base(None), &accessor));
let scoped = Arc::new(StorageOptionsAccessor::with_static_options(HashMap::from(
[
("account_key".to_string(), "shared-key".to_string()),
("base_1.account_key".to_string(), "base1-key".to_string()),
],
)))
.scoped_to_base(Some(1));
assert!(Arc::ptr_eq(&scoped.scoped_to_base(None), &scoped));
let provider_scoped = Arc::new(StorageOptionsAccessor::with_provider(Arc::new(
MockStorageOptionsProvider::new(None),
)))
.scoped_to_base(Some(1));
assert!(Arc::ptr_eq(
&provider_scoped.scoped_to_base(None),
&provider_scoped
));
}
#[tokio::test]
async fn test_scoped_to_base_provider_only_resolves_vended_options() {
MockClock::set_system_time(Duration::from_secs(100_000));
let provider = Arc::new(MockBaseScopedVendingProvider {
call_count: Arc::new(RwLock::new(0)),
expires_in_millis: 600_000,
});
let parent = Arc::new(StorageOptionsAccessor::with_provider(provider.clone()));
let base1 = parent.scoped_to_base(Some(1));
assert!(!Arc::ptr_eq(&base1, &parent));
let result = base1.get_storage_options().await.unwrap();
assert_eq!(result.0.get("account_key").unwrap(), "BASE1_1");
assert!(!result.0.contains_key("base_1.account_key"));
let default = parent.scoped_to_base(None);
let result = default.get_storage_options().await.unwrap();
assert_eq!(result.0.get("account_key").unwrap(), "SHARED_1");
assert!(!result.0.contains_key("base_1.account_key"));
assert_eq!(*provider.call_count.read().await, 1);
}
#[tokio::test]
async fn test_scoped_earlier_base_expiry_refreshes_parent() {
MockClock::set_system_time(Duration::from_secs(100_000));
let now_ms = MockClock::system_time().as_millis() as u64;
let provider = Arc::new(MockBaseScopedVendingProvider {
call_count: Arc::new(RwLock::new(0)),
expires_in_millis: 600_000,
});
let initial = HashMap::from([
("account_key".to_string(), "SHARED_0".to_string()),
("base_1.account_key".to_string(), "BASE1_0".to_string()),
(
EXPIRES_AT_MILLIS_KEY.to_string(),
(now_ms + 600_000).to_string(),
),
(
format!("base_1.{}", EXPIRES_AT_MILLIS_KEY),
(now_ms + 120_000).to_string(),
),
]);
let parent = Arc::new(StorageOptionsAccessor::with_initial_and_provider(
initial,
provider.clone(),
));
let base1 = parent.scoped_to_base(Some(1));
let result = base1.get_storage_options().await.unwrap();
assert_eq!(result.0.get("account_key").unwrap(), "BASE1_0");
assert_eq!(*provider.call_count.read().await, 0);
MockClock::set_system_time(Duration::from_secs(100_000 + 121));
let result = base1.get_storage_options().await.unwrap();
assert_eq!(result.0.get("account_key").unwrap(), "BASE1_1");
assert_eq!(*provider.call_count.read().await, 1);
}
#[tokio::test]
async fn test_scoped_to_base_static() {
let accessor = Arc::new(StorageOptionsAccessor::with_static_options(HashMap::from(
[
("account_key".to_string(), "shared-key".to_string()),
("base_1.account_key".to_string(), "base1-key".to_string()),
],
)));
let base1 = accessor.scoped_to_base(Some(1));
let result = base1.get_storage_options().await.unwrap();
assert_eq!(
result.0,
HashMap::from([("account_key".to_string(), "base1-key".to_string())])
);
assert!(!base1.has_provider());
let default = accessor.scoped_to_base(None);
let result = default.get_storage_options().await.unwrap();
assert_eq!(
result.0,
HashMap::from([("account_key".to_string(), "shared-key".to_string())])
);
assert_eq!(
accessor.scoped_to_base(Some(1)).accessor_id(),
base1.accessor_id()
);
assert_ne!(base1.accessor_id(), default.accessor_id());
assert_ne!(base1.accessor_id(), accessor.accessor_id());
}
#[derive(Debug)]
struct MockBaseScopedVendingProvider {
call_count: Arc<RwLock<usize>>,
expires_in_millis: u64,
}
#[async_trait]
impl StorageOptionsProvider for MockBaseScopedVendingProvider {
async fn fetch_storage_options(&self) -> Result<Option<HashMap<String, String>>> {
let count = {
let mut c = self.call_count.write().await;
*c += 1;
*c
};
let now_ms = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_millis() as u64;
Ok(Some(HashMap::from([
("account_key".to_string(), format!("SHARED_{}", count)),
("base_1.account_key".to_string(), format!("BASE1_{}", count)),
(
EXPIRES_AT_MILLIS_KEY.to_string(),
(now_ms + self.expires_in_millis).to_string(),
),
])))
}
fn provider_id(&self) -> String {
"MockBaseScopedVendingProvider".to_string()
}
}
#[tokio::test]
async fn test_scoped_to_base_refreshes_through_parent() {
MockClock::set_system_time(Duration::from_secs(100_000));
let now_ms = MockClock::system_time().as_millis() as u64;
let provider = Arc::new(MockBaseScopedVendingProvider {
call_count: Arc::new(RwLock::new(0)),
expires_in_millis: 600_000,
});
let initial = HashMap::from([
("account_key".to_string(), "SHARED_0".to_string()),
("base_1.account_key".to_string(), "BASE1_0".to_string()),
(
EXPIRES_AT_MILLIS_KEY.to_string(),
(now_ms + 600_000).to_string(),
),
]);
let parent = Arc::new(StorageOptionsAccessor::with_initial_and_provider(
initial,
provider.clone(),
));
let base1 = parent.scoped_to_base(Some(1));
let default = parent.scoped_to_base(None);
assert!(base1.has_provider());
let result = base1.get_storage_options().await.unwrap();
assert_eq!(result.0.get("account_key").unwrap(), "BASE1_0");
assert!(!result.0.contains_key("base_1.account_key"));
let result = default.get_storage_options().await.unwrap();
assert_eq!(result.0.get("account_key").unwrap(), "SHARED_0");
assert_eq!(*provider.call_count.read().await, 0);
MockClock::set_system_time(Duration::from_secs(100_000 + 601));
let result = base1.get_storage_options().await.unwrap();
assert_eq!(result.0.get("account_key").unwrap(), "BASE1_1");
assert_eq!(*provider.call_count.read().await, 1);
let result = default.get_storage_options().await.unwrap();
assert_eq!(result.0.get("account_key").unwrap(), "SHARED_1");
assert_eq!(*provider.call_count.read().await, 1);
}
#[tokio::test]
async fn test_scoped_forced_refresh_reaches_origin_provider() {
MockClock::set_system_time(Duration::from_secs(100_000));
let now_ms = MockClock::system_time().as_millis() as u64;
let provider = Arc::new(MockBaseScopedVendingProvider {
call_count: Arc::new(RwLock::new(0)),
expires_in_millis: 600_000,
});
let initial = HashMap::from([
("account_key".to_string(), "SHARED_0".to_string()),
("base_1.account_key".to_string(), "BASE1_0".to_string()),
(
EXPIRES_AT_MILLIS_KEY.to_string(),
(now_ms + 600_000).to_string(),
),
]);
let parent = Arc::new(StorageOptionsAccessor::with_initial_and_provider(
initial,
provider.clone(),
));
let base1 = parent.scoped_to_base(Some(1));
let result = base1.refresh_storage_options().await.unwrap();
assert_eq!(result.0.get("account_key").unwrap(), "BASE1_1");
assert_eq!(*provider.call_count.read().await, 1);
}
}