use std::path::{Path, PathBuf};
use async_trait::async_trait;
use super::file::FileTokenStore;
#[cfg(feature = "keyring")]
use super::keyring::KeyringTokenStore;
use super::{PersistedTokens, TokenKey, TokenStore, TokenStoreError};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum AutoBackend {
#[cfg(feature = "keyring")]
Keyring,
File,
}
pub struct AutoTokenStore {
file: FileTokenStore,
#[cfg(feature = "keyring")]
keyring: KeyringTokenStore,
#[cfg(not(feature = "keyring"))]
_service_name: String,
}
impl AutoTokenStore {
#[cfg(feature = "keyring")]
pub fn new(root: impl Into<PathBuf>, service_name: impl Into<String>) -> Self {
Self {
file: FileTokenStore::new(root),
keyring: KeyringTokenStore::new(service_name),
}
}
#[cfg(not(feature = "keyring"))]
pub fn new(root: impl Into<PathBuf>, service_name: impl Into<String>) -> Self {
Self {
file: FileTokenStore::new(root),
_service_name: service_name.into(),
}
}
pub fn root(&self) -> &Path {
self.file.root()
}
#[cfg(feature = "keyring")]
async fn resolve_backend(&self, key: &TokenKey) -> Result<AutoBackend, TokenStoreError> {
match self.keyring.load(key).await {
Ok(Some(_)) => return Ok(AutoBackend::Keyring),
Ok(None) => {}
Err(TokenStoreError::KeyringUnavailable(_)) => return Ok(AutoBackend::File),
Err(e) => return Err(e),
}
if self.file.load(key).await?.is_some() {
return Ok(AutoBackend::File);
}
Ok(AutoBackend::Keyring)
}
#[cfg(not(feature = "keyring"))]
async fn resolve_backend(&self, _key: &TokenKey) -> Result<AutoBackend, TokenStoreError> {
Ok(AutoBackend::File)
}
}
#[async_trait]
impl TokenStore for AutoTokenStore {
async fn load(&self, key: &TokenKey) -> Result<Option<PersistedTokens>, TokenStoreError> {
match self.resolve_backend(key).await? {
#[cfg(feature = "keyring")]
AutoBackend::Keyring => self.keyring.load(key).await,
AutoBackend::File => self.file.load(key).await,
}
}
async fn save(&self, key: &TokenKey, tokens: &PersistedTokens) -> Result<(), TokenStoreError> {
match self.resolve_backend(key).await? {
#[cfg(feature = "keyring")]
AutoBackend::Keyring => {
self.keyring.save(key, tokens).await?;
self.file.clear(key).await?;
Ok(())
}
AutoBackend::File => {
self.file.save(key, tokens).await?;
#[cfg(feature = "keyring")]
match self.keyring.clear(key).await {
Ok(()) | Err(TokenStoreError::KeyringUnavailable(_)) => {}
Err(e) => return Err(e),
}
Ok(())
}
}
}
async fn clear(&self, key: &TokenKey) -> Result<(), TokenStoreError> {
#[cfg(feature = "keyring")]
{
match self.keyring.clear(key).await {
Ok(()) | Err(TokenStoreError::KeyringUnavailable(_)) => {}
Err(e) => return Err(e),
}
}
self.file.clear(key).await
}
async fn list(&self) -> Result<Vec<TokenKey>, TokenStoreError> {
let mut combined = Vec::new();
#[cfg(feature = "keyring")]
{
match self.keyring.list().await {
Ok(mut v) => combined.append(&mut v),
Err(TokenStoreError::KeyringUnavailable(_) | TokenStoreError::Unavailable(_)) => {}
Err(e) => return Err(e),
}
}
let mut file_keys = self.file.list().await?;
combined.append(&mut file_keys);
combined.sort();
combined.dedup();
Ok(combined)
}
fn backend_name(&self) -> &'static str {
"auto"
}
}
#[cfg(all(test, not(feature = "keyring")))]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod tests {
use super::*;
fn sample_tokens(secret: &str) -> PersistedTokens {
PersistedTokens {
auth_mode: meerkat_core::auth::token_store::PersistedAuthMode::ApiKey,
primary_secret: Some(secret.to_string()),
refresh_token: None,
id_token: None,
expires_at: None,
last_refresh: None,
scopes: Vec::new(),
account_id: None,
metadata: serde_json::Value::Null,
}
}
#[tokio::test]
async fn save_resave_targets_single_backend_no_divergent_copies() {
let dir = std::env::temp_dir().join(format!("rkat-auto-{}", uuid::Uuid::new_v4()));
let store = AutoTokenStore::new(&dir, "rkat-test-auto");
let key = TokenKey::parse("dev", "default_openai").unwrap();
store.save(&key, &sample_tokens("first")).await.unwrap();
let loaded = store.load(&key).await.unwrap().unwrap();
assert_eq!(loaded.primary_secret.as_deref(), Some("first"));
store.save(&key, &sample_tokens("second")).await.unwrap();
let reloaded = store.load(&key).await.unwrap().unwrap();
assert_eq!(reloaded.primary_secret.as_deref(), Some("second"));
let listed = store.list().await.unwrap();
assert_eq!(
listed.iter().filter(|k| **k == key).count(),
1,
"key must not have divergent copies across backends after re-save"
);
store.clear(&key).await.unwrap();
assert!(store.load(&key).await.unwrap().is_none());
let _ = std::fs::remove_dir_all(&dir);
}
}