use std::collections::HashMap;
use std::sync::{Arc, Condvar, Mutex, OnceLock};
use std::time::{Duration, Instant};
use super::{KeyStore, KeyStoreError};
use crate::credential_registry::{env_var_for, provider_for_env_var};
pub const STORE_READ_TIMEOUT: Duration = Duration::from_secs(3);
pub const STORE_ERROR_CACHE_TTL: Duration = Duration::from_secs(45);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum StoreErrorKind {
Timeout,
Keyring,
Io,
Toml,
HomeUnavailable,
}
impl StoreErrorKind {
#[must_use]
pub fn as_str(self) -> &'static str {
match self {
Self::Timeout => "timeout",
Self::Keyring => "keyring-backend",
Self::Io => "io",
Self::Toml => "toml",
Self::HomeUnavailable => "home-unavailable",
}
}
}
impl std::fmt::Display for StoreErrorKind {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
impl From<&KeyStoreError> for StoreErrorKind {
fn from(e: &KeyStoreError) -> Self {
match e {
KeyStoreError::Io { .. } => Self::Io,
KeyStoreError::Toml { .. } => Self::Toml,
KeyStoreError::HomeUnavailable => Self::HomeUnavailable,
KeyStoreError::Keyring(_) => Self::Keyring,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct StoreFailure {
pub kind: StoreErrorKind,
pub cached: bool,
}
impl StoreFailure {
#[must_use]
fn fresh(kind: StoreErrorKind) -> Self {
Self {
kind,
cached: false,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum SecretResolveError {
#[error(
"`{var}` is not a registered credential name — register it in \
credential_registry::REGISTRY before it can be resolved"
)]
Unregistered {
var: String,
},
#[error(
"no value configured for `{var}` in the environment, `.env.local`, or the credential store"
)]
Absent {
var: String,
},
#[error(
"reading `{var}` from the credential store timed out after {waited_ms} ms \
(a Keychain approval dialog may be waiting on screen)"
)]
Timeout {
var: String,
waited_ms: u64,
cached: bool,
},
#[error("the credential store could not supply `{var}`: {kind}")]
Store {
var: String,
kind: StoreErrorKind,
cached: bool,
},
}
impl SecretResolveError {
#[must_use]
pub fn kind(&self) -> &'static str {
match self {
Self::Unregistered { .. } => "unregistered",
Self::Absent { .. } => "absent",
Self::Timeout { .. } => StoreErrorKind::Timeout.as_str(),
Self::Store { kind, .. } => kind.as_str(),
}
}
#[must_use]
pub fn var(&self) -> &str {
match self {
Self::Unregistered { var }
| Self::Absent { var }
| Self::Timeout { var, .. }
| Self::Store { var, .. } => var,
}
}
}
struct Flight {
outcome: Mutex<Option<Result<Option<String>, StoreErrorKind>>>,
done: Condvar,
}
static INFLIGHT: OnceLock<Mutex<HashMap<String, Arc<Flight>>>> = OnceLock::new();
static ERROR_CACHE: OnceLock<Mutex<HashMap<String, (Instant, StoreErrorKind)>>> = OnceLock::new();
fn map<V: 'static>(
cell: &'static OnceLock<Mutex<HashMap<String, V>>>,
) -> &'static Mutex<HashMap<String, V>> {
cell.get_or_init(|| Mutex::new(HashMap::new()))
}
pub fn store_get_bounded(
store: Arc<dyn KeyStore>,
provider: &str,
timeout: Duration,
) -> Result<Option<String>, StoreFailure> {
if let Some(kind) = cached_error(provider) {
return Err(StoreFailure { kind, cached: true });
}
let (flight, is_leader) = join_or_start(provider);
if is_leader {
spawn_reader(store, provider.to_string(), Arc::clone(&flight));
}
let mut guard = flight
.outcome
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let deadline = Instant::now() + timeout;
while guard.is_none() {
let remaining = deadline.saturating_duration_since(Instant::now());
if remaining.is_zero() {
record_error(provider, StoreErrorKind::Timeout);
drop(guard);
return Err(StoreFailure::fresh(StoreErrorKind::Timeout));
}
let (next, _) = flight
.done
.wait_timeout(guard, remaining)
.unwrap_or_else(std::sync::PoisonError::into_inner);
guard = next;
}
match guard.clone() {
Some(Ok(value)) => Ok(value),
Some(Err(kind)) => Err(StoreFailure::fresh(kind)),
None => Err(StoreFailure::fresh(StoreErrorKind::Timeout)),
}
}
fn cached_error(provider: &str) -> Option<StoreErrorKind> {
let mut cache = map(&ERROR_CACHE)
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
match cache.get(provider) {
Some((at, kind)) if at.elapsed() < STORE_ERROR_CACHE_TTL => Some(*kind),
Some(_) => {
cache.remove(provider);
None
}
None => None,
}
}
fn record_error(provider: &str, kind: StoreErrorKind) {
map(&ERROR_CACHE)
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.insert(provider.to_string(), (Instant::now(), kind));
}
fn clear_error(provider: &str) {
map(&ERROR_CACHE)
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.remove(provider);
}
fn join_or_start(provider: &str) -> (Arc<Flight>, bool) {
let mut inflight = map(&INFLIGHT)
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(existing) = inflight.get(provider) {
return (Arc::clone(existing), false);
}
let flight = Arc::new(Flight {
outcome: Mutex::new(None),
done: Condvar::new(),
});
inflight.insert(provider.to_string(), Arc::clone(&flight));
(flight, true)
}
fn spawn_reader(store: Arc<dyn KeyStore>, provider: String, flight: Arc<Flight>) {
let read_provider = provider.clone();
let read_flight = Arc::clone(&flight);
let spawned = std::thread::Builder::new()
.name("cred-store-read".to_string())
.spawn(move || {
let outcome = store
.try_get(&read_provider)
.map_err(|e| StoreErrorKind::from(&e));
finish(&read_provider, &read_flight, outcome);
});
if spawned.is_err() {
finish(&provider, &flight, Err(StoreErrorKind::Io));
}
}
fn finish(provider: &str, flight: &Flight, outcome: Result<Option<String>, StoreErrorKind>) {
let mut slot = flight
.outcome
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
match &outcome {
Err(kind) => record_error(provider, *kind),
Ok(_) => clear_error(provider),
}
park_before_publish(provider);
map(&INFLIGHT)
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.remove(provider);
*slot = Some(outcome);
drop(slot);
flight.done.notify_all();
}
#[cfg(test)]
type ParkHook = Arc<dyn Fn(&str) + Send + Sync>;
#[cfg(test)]
static PARK_BEFORE_PUBLISH: Mutex<Option<ParkHook>> = Mutex::new(None);
#[cfg(test)]
fn park_before_publish(provider: &str) {
let hook = PARK_BEFORE_PUBLISH
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.clone();
if let Some(hook) = hook {
hook(provider);
}
}
#[cfg(not(test))]
#[inline]
fn park_before_publish(_provider: &str) {}
pub fn resolve_env_var_bounded(var: &str) -> Result<String, SecretResolveError> {
let Some(provider) = provider_for_env_var(var) else {
return Err(SecretResolveError::Unregistered {
var: var.to_string(),
});
};
super::dotenv::load_env_local_once();
resolve_provider_bounded_with(
provider,
Arc::from(super::default_store()),
STORE_READ_TIMEOUT,
)
}
pub fn resolve_provider_bounded_with(
provider: &str,
store: Arc<dyn KeyStore>,
timeout: Duration,
) -> Result<String, SecretResolveError> {
let var = env_var_for(provider)
.map(str::to_string)
.unwrap_or_else(|| provider.to_string());
if let Ok(value) = std::env::var(&var)
&& !value.is_empty()
{
return Ok(value);
}
match store_get_bounded(store, provider, timeout) {
Ok(Some(value)) if !value.is_empty() => Ok(value),
Ok(_) => Err(SecretResolveError::Absent { var }),
Err(StoreFailure {
kind: StoreErrorKind::Timeout,
cached,
}) => Err(SecretResolveError::Timeout {
var,
waited_ms: u64::try_from(timeout.as_millis()).unwrap_or(u64::MAX),
cached,
}),
Err(StoreFailure { kind, cached }) => Err(SecretResolveError::Store { var, kind, cached }),
}
}
#[cfg(test)]
#[path = "bounded_store/tests.rs"]
mod tests;