use std::collections::{HashMap, HashSet};
use std::sync::Arc;
use corium_db::Db;
use corium_peer::Connection;
use corium_pgwire::{CatalogError, DbCatalog};
use tokio::sync::Mutex;
use crate::ClientFlags;
pub struct PeerCatalog {
client: ClientFlags,
whitelist: Option<HashSet<String>>,
connections: Mutex<HashMap<String, Arc<Connection>>>,
}
impl PeerCatalog {
pub fn new(client: ClientFlags, databases: Vec<String>) -> Self {
let whitelist = if databases.is_empty() {
None
} else {
Some(databases.into_iter().collect())
};
Self {
client,
whitelist,
connections: Mutex::new(HashMap::new()),
}
}
fn allowed(&self, name: &str) -> bool {
self.whitelist
.as_ref()
.is_none_or(|allowed| allowed.contains(name))
}
}
#[async_trait::async_trait]
impl DbCatalog for PeerCatalog {
async fn list(&self) -> Result<Vec<String>, CatalogError> {
let mut admin = crate::admin_client(&self.client)
.await
.map_err(CatalogError::Unavailable)?;
let mut names = admin
.list_databases()
.await
.map_err(|error| CatalogError::Unavailable(error.to_string()))?;
if let Some(allowed) = &self.whitelist {
names.retain(|name| allowed.contains(name));
}
names.sort();
Ok(names)
}
async fn db(&self, name: &str) -> Result<Db, CatalogError> {
if !self.allowed(name) {
return Err(CatalogError::NotFound(name.to_owned()));
}
{
let cache = self.connections.lock().await;
if let Some(connection) = cache.get(name) {
return Ok(connection.db());
}
}
let config = self
.client
.connect_config(name.to_owned())
.await
.map_err(CatalogError::Unavailable)?;
let connection = Arc::new(
Connection::connect(config)
.await
.map_err(|error| CatalogError::Unavailable(error.to_string()))?,
);
let mut cache = self.connections.lock().await;
let entry = cache.entry(name.to_owned()).or_insert(connection);
Ok(entry.db())
}
}