use std::collections::HashMap;
use std::future::{Future, ready};
use std::pin::Pin;
use std::sync::{Arc, Mutex, MutexGuard};
use super::AuthProxyError;
use super::registry::Grant;
pub type GrantFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, AuthProxyError>> + Send + 'a>>;
pub trait GrantBackend: Send + Sync {
fn insert(&self, grant: Grant) -> GrantFuture<'_, ()>;
fn get(&self, id: &str) -> GrantFuture<'_, Option<Grant>>;
fn remove(&self, id: &str) -> GrantFuture<'_, bool>;
fn remove_run(&self, run_id: &str) -> GrantFuture<'_, usize>;
fn purge_expired(&self, now: u64) -> GrantFuture<'_, usize>;
fn len(&self) -> GrantFuture<'_, usize>;
fn is_empty(&self) -> GrantFuture<'_, bool> {
Box::pin(async move { Ok(self.len().await? == 0) })
}
}
#[derive(Clone, Default)]
pub struct MemoryGrantBackend {
grants: Arc<Mutex<HashMap<String, Grant>>>,
}
impl MemoryGrantBackend {
fn grants(&self) -> MutexGuard<'_, HashMap<String, Grant>> {
self.grants.lock().unwrap_or_else(|e| e.into_inner())
}
}
impl GrantBackend for MemoryGrantBackend {
fn insert(&self, grant: Grant) -> GrantFuture<'_, ()> {
self.grants().insert(grant.id.clone(), grant);
Box::pin(ready(Ok(())))
}
fn get(&self, id: &str) -> GrantFuture<'_, Option<Grant>> {
let grant = self.grants().get(id).cloned();
Box::pin(ready(Ok(grant)))
}
fn remove(&self, id: &str) -> GrantFuture<'_, bool> {
let removed = self.grants().remove(id).is_some();
Box::pin(ready(Ok(removed)))
}
fn remove_run(&self, run_id: &str) -> GrantFuture<'_, usize> {
let removed = {
let mut grants = self.grants();
let before = grants.len();
grants.retain(|_, grant| grant.run_id != run_id);
before - grants.len()
};
Box::pin(ready(Ok(removed)))
}
fn purge_expired(&self, now: u64) -> GrantFuture<'_, usize> {
let purged = {
let mut grants = self.grants();
let before = grants.len();
grants.retain(|_, grant| grant.expires_at > now);
before - grants.len()
};
Box::pin(ready(Ok(purged)))
}
fn len(&self) -> GrantFuture<'_, usize> {
let len = self.grants().len();
Box::pin(ready(Ok(len)))
}
}
#[cfg(test)]
mod tests {
use super::super::registry::{AuthProxyRegistry, TokenRequest, token_id};
use super::super::{CredentialKind, ProxyCredential};
use super::*;
const NOW: u64 = 1_700_000_000;
fn grant(id: &str, run_id: &str, expires_at: u64) -> Grant {
Grant {
run_id: run_id.to_string(),
step: "review".to_string(),
expires_at,
credential: ProxyCredential::new(
CredentialKind::OauthToken,
"sk-ant-oat01-test".to_string(),
),
id: id.to_string(),
}
}
#[tokio::test]
async fn insert_then_get_returns_grant() {
let backend = MemoryGrantBackend::default();
backend.insert(grant("a", "run-1", NOW)).await.unwrap();
let stored = backend.get("a").await.unwrap().unwrap();
assert_eq!(stored.id, "a");
assert_eq!(stored.run_id, "run-1");
assert_eq!(stored.expires_at, NOW);
assert_eq!(stored.credential.expose(), "sk-ant-oat01-test");
assert!(backend.get("b").await.unwrap().is_none());
}
#[tokio::test]
async fn remove_reports_whether_it_existed() {
let backend = MemoryGrantBackend::default();
backend.insert(grant("a", "run-1", NOW)).await.unwrap();
assert!(!backend.is_empty().await.unwrap());
assert!(backend.remove("a").await.unwrap());
assert!(!backend.remove("a").await.unwrap());
assert_eq!(backend.len().await.unwrap(), 0);
assert!(backend.is_empty().await.unwrap());
}
#[tokio::test]
async fn remove_run_only_drops_that_run() {
let backend = MemoryGrantBackend::default();
backend.insert(grant("a1", "run-a", NOW)).await.unwrap();
backend.insert(grant("a2", "run-a", NOW)).await.unwrap();
backend.insert(grant("b", "run-b", NOW)).await.unwrap();
assert_eq!(backend.remove_run("run-a").await.unwrap(), 2);
assert_eq!(backend.remove_run("run-a").await.unwrap(), 0);
assert_eq!(backend.len().await.unwrap(), 1);
assert!(backend.get("b").await.unwrap().is_some());
}
#[tokio::test]
async fn purge_expired_drops_grants_at_or_before_now() {
let backend = MemoryGrantBackend::default();
backend.insert(grant("a", "run-1", NOW + 10)).await.unwrap();
backend.insert(grant("b", "run-1", NOW + 20)).await.unwrap();
backend.insert(grant("c", "run-1", NOW + 21)).await.unwrap();
assert_eq!(backend.purge_expired(NOW + 20).await.unwrap(), 2);
assert_eq!(backend.len().await.unwrap(), 1);
assert!(backend.get("c").await.unwrap().is_some());
assert_eq!(backend.purge_expired(NOW + 20).await.unwrap(), 0);
}
#[tokio::test]
async fn clones_share_grants() {
let backend = MemoryGrantBackend::default();
let clone = backend.clone();
backend.insert(grant("a", "run-1", NOW)).await.unwrap();
assert_eq!(clone.len().await.unwrap(), 1);
assert!(clone.remove("a").await.unwrap());
assert_eq!(backend.len().await.unwrap(), 0);
}
#[tokio::test]
async fn registry_never_stores_the_token() {
let backend = Arc::new(MemoryGrantBackend::default());
let registry = AuthProxyRegistry::with_backend(backend.clone());
let issued = registry
.issue(
TokenRequest {
run_id: "run-1".to_string(),
step: "review".to_string(),
expires_at: NOW + 600,
credential: ProxyCredential::new(
CredentialKind::OauthToken,
"sk-ant-oat01-test".to_string(),
),
},
NOW,
)
.await
.unwrap();
let grants = backend.grants();
let keys: Vec<&String> = grants.keys().collect();
assert_eq!(keys, vec![&token_id(&issued.token)]);
assert_ne!(keys[0], &issued.token);
let stored = format!("{:?}", grants.values().collect::<Vec<_>>());
assert!(!stored.contains(&issued.token), "{stored}");
assert!(!stored.contains("sk-ant-oat01-test"), "{stored}");
}
}