use crate::provider_credential::domain::{ProviderAuthStatus, ProviderSlug, StoredCredential};
use std::collections::HashMap;
use std::sync::Arc;
#[async_trait::async_trait]
pub trait ProviderCredentialStore: Send + Sync {
async fn get(&self, slug: &ProviderSlug) -> Option<StoredCredential>;
async fn set(&self, slug: &ProviderSlug, credential: StoredCredential);
async fn remove(&self, slug: &ProviderSlug);
async fn list_slugs(&self) -> Vec<ProviderSlug>;
}
pub struct InMemoryCredentialStore {
inner: parking_lot::Mutex<HashMap<String, StoredCredential>>,
}
impl InMemoryCredentialStore {
pub fn new() -> Self {
Self {
inner: parking_lot::Mutex::new(HashMap::new()),
}
}
pub fn from_map(map: HashMap<String, StoredCredential>) -> Self {
Self {
inner: parking_lot::Mutex::new(map),
}
}
pub fn snapshot(&self) -> HashMap<String, StoredCredential> {
self.inner.lock().clone()
}
pub fn status_from_credential(cred: Option<&StoredCredential>) -> ProviderAuthStatus {
match cred {
None => ProviderAuthStatus::NotConfigured,
Some(StoredCredential::ApiKey(_)) => ProviderAuthStatus::ConfiguredApiKey,
Some(StoredCredential::OAuthBearer(tokens)) => {
let now = std::time::SystemTime::now();
match tokens.expires_at {
Some(expires_at) if now >= expires_at => ProviderAuthStatus::Expired,
_ => ProviderAuthStatus::ConnectedOAuth {
expires_at: tokens.expires_at,
},
}
}
}
}
}
impl Default for InMemoryCredentialStore {
fn default() -> Self {
Self::new()
}
}
#[async_trait::async_trait]
impl ProviderCredentialStore for InMemoryCredentialStore {
async fn get(&self, slug: &ProviderSlug) -> Option<StoredCredential> {
self.inner.lock().get(slug.as_str()).cloned()
}
async fn set(&self, slug: &ProviderSlug, credential: StoredCredential) {
self.inner
.lock()
.insert(slug.as_str().to_string(), credential);
}
async fn remove(&self, slug: &ProviderSlug) {
self.inner.lock().remove(slug.as_str());
}
async fn list_slugs(&self) -> Vec<ProviderSlug> {
self.inner
.lock()
.keys()
.cloned()
.map(ProviderSlug::new)
.collect()
}
}
pub struct NullCredentialStore;
#[async_trait::async_trait]
impl ProviderCredentialStore for NullCredentialStore {
async fn get(&self, _slug: &ProviderSlug) -> Option<StoredCredential> {
None
}
async fn set(&self, _slug: &ProviderSlug, _credential: StoredCredential) {
}
async fn remove(&self, _slug: &ProviderSlug) {
}
async fn list_slugs(&self) -> Vec<ProviderSlug> {
Vec::new()
}
}
pub type DynCredentialStore = Arc<dyn ProviderCredentialStore>;
#[cfg(test)]
mod tests {
use super::*;
use crate::provider_credential::domain::{OAuthTokenSet, ProviderAuthStatus};
#[tokio::test]
async fn in_memory_store_roundtrip() {
let store = InMemoryCredentialStore::new();
let slug = ProviderSlug::new("kimi-code");
let cred = StoredCredential::ApiKey("sk-test".into());
assert!(store.get(&slug).await.is_none());
store.set(&slug, cred.clone()).await;
let got = store.get(&slug).await;
assert_eq!(got, Some(cred));
store.remove(&slug).await;
assert!(store.get(&slug).await.is_none());
}
#[tokio::test]
async fn in_memory_store_list_slugs() {
let store = InMemoryCredentialStore::new();
let slug_a = ProviderSlug::new("kimi-code");
let slug_b = ProviderSlug::new("codex");
store
.set(&slug_a, StoredCredential::ApiKey("sk-a".into()))
.await;
store
.set(&slug_b, StoredCredential::ApiKey("sk-b".into()))
.await;
let mut slugs = store.list_slugs().await;
slugs.sort_by(|a, b| a.as_str().cmp(b.as_str()));
assert_eq!(slugs, vec![slug_b, slug_a]);
}
#[tokio::test]
async fn null_store_always_empty() {
let store = NullCredentialStore;
let slug = ProviderSlug::new("codex");
assert!(store.get(&slug).await.is_none());
assert!(store.list_slugs().await.is_empty());
}
#[test]
fn status_from_credential_api_key() {
let cred = StoredCredential::ApiKey("sk-test".into());
assert_eq!(
InMemoryCredentialStore::status_from_credential(Some(&cred)),
ProviderAuthStatus::ConfiguredApiKey
);
}
#[test]
fn status_from_credential_oauth_valid() {
let cred = StoredCredential::OAuthBearer(OAuthTokenSet {
access_token: "at".into(),
refresh_token: "rt".into(),
expires_at: Some(std::time::SystemTime::now() + std::time::Duration::from_secs(3600)),
id_token: None,
});
assert!(matches!(
InMemoryCredentialStore::status_from_credential(Some(&cred)),
ProviderAuthStatus::ConnectedOAuth { .. }
));
}
#[test]
fn status_from_credential_oauth_expired() {
let cred = StoredCredential::OAuthBearer(OAuthTokenSet {
access_token: "at".into(),
refresh_token: "rt".into(),
expires_at: Some(std::time::SystemTime::now() - std::time::Duration::from_secs(1)),
id_token: None,
});
assert_eq!(
InMemoryCredentialStore::status_from_credential(Some(&cred)),
ProviderAuthStatus::Expired
);
}
#[test]
fn status_from_credential_none() {
assert_eq!(
InMemoryCredentialStore::status_from_credential(None),
ProviderAuthStatus::NotConfigured
);
}
}