use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::Mutex;
use crate::config::DatabaseConfig;
use crate::{Database, Result};
pub const DEFAULT_BUDGET: usize = 80;
struct Entry {
database: Database,
used_at: u64,
}
#[derive(Clone)]
pub struct Connections {
inner: Arc<Mutex<HashMap<String, Entry>>>,
budget: usize,
tick: Arc<std::sync::atomic::AtomicU64>,
}
impl Default for Connections {
fn default() -> Self {
Connections::with_budget(DEFAULT_BUDGET)
}
}
impl Connections {
pub fn new() -> Connections {
Connections::default()
}
pub fn with_budget(budget: usize) -> Connections {
Connections {
inner: Arc::new(Mutex::new(HashMap::new())),
budget: budget.max(1),
tick: Arc::new(std::sync::atomic::AtomicU64::new(0)),
}
}
pub fn budget(&self) -> usize {
self.budget
}
pub async fn insert(&self, name: impl Into<String>, database: Database) {
let used_at = self.next_tick();
self.inner.lock().await.insert(name.into(), Entry { database, used_at });
}
pub async fn get(&self, name: &str) -> Option<Database> {
let used_at = self.next_tick();
let mut held = self.inner.lock().await;
let entry = held.get_mut(name)?;
entry.used_at = used_at;
Some(entry.database.clone())
}
pub async fn get_or_open(&self, name: &str, config: DatabaseConfig) -> Result<Database> {
if let Some(database) = self.get(name).await {
return Ok(database);
}
self.make_room(config.max_connections).await;
let database = Database::with_config(config).await?;
self.insert(name, database.clone()).await;
Ok(database)
}
pub async fn forget(&self, name: &str) -> bool {
let Some(entry) = self.inner.lock().await.remove(name) else { return false };
let open = entry.database.pool().open_count().await;
entry.database.pool().close_idle(open).await;
true
}
pub async fn names(&self) -> Vec<String> {
let mut names: Vec<String> = self.inner.lock().await.keys().cloned().collect();
names.sort();
names
}
pub async fn open_count(&self) -> usize {
let entries: Vec<Database> =
self.inner.lock().await.values().map(|entry| entry.database.clone()).collect();
let mut total = 0;
for database in entries {
total += database.pool().open_count().await;
}
total
}
async fn make_room(&self, wanted: usize) {
let mut open = self.open_count().await;
if open + wanted <= self.budget {
return;
}
let mut candidates: Vec<(u64, Database)> = self
.inner
.lock()
.await
.values()
.map(|entry| (entry.used_at, entry.database.clone()))
.collect();
candidates.sort_by_key(|(used_at, _)| *used_at);
for (_, database) in candidates {
if open + wanted <= self.budget {
return;
}
let over = (open + wanted) - self.budget;
open -= database.pool().close_idle(over).await;
}
}
fn next_tick(&self) -> u64 {
self.tick.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn connections_are_held_by_name() {
let registry = Connections::new();
assert!(registry.get("central").await.is_none());
assert_eq!(registry.names().await, Vec::<String>::new());
assert!(!registry.forget("central").await);
}
#[test]
fn the_default_budget_leaves_the_server_something() {
const { assert!(DEFAULT_BUDGET < 100) };
assert_eq!(Connections::new().budget(), DEFAULT_BUDGET);
}
#[test]
fn a_budget_is_at_least_one() {
assert_eq!(Connections::with_budget(0).budget(), 1);
}
#[tokio::test]
async fn an_empty_registry_holds_nothing_open() {
assert_eq!(Connections::new().open_count().await, 0);
}
#[tokio::test]
async fn making_room_when_there_is_nothing_to_free_is_harmless() {
let registry = Connections::with_budget(4);
registry.make_room(10).await;
assert_eq!(registry.open_count().await, 0);
}
}