#[cfg(not(feature = "sync"))]
use std::future::Future;
use std::{
collections::HashMap,
io::BufReader,
path::Path,
sync::{Arc, RwLock},
};
use typedb_protocol::migration::Item;
use super::Database;
use crate::{
Error,
common::{Result, address::Address, error::ConnectionError},
connection::server_connection::ServerConnection,
database::migration::{ProtoMessageIterator, try_open_import_file},
info::DatabaseInfo,
resolve,
};
#[derive(Debug)]
pub struct DatabaseManager {
server_connections: HashMap<Address, ServerConnection>,
databases_cache: RwLock<HashMap<String, Arc<Database>>>,
}
impl DatabaseManager {
pub(crate) fn new(
server_connections: HashMap<Address, ServerConnection>,
database_info: Vec<DatabaseInfo>,
) -> Result<Self> {
let mut databases = HashMap::new();
for info in database_info {
let database = Database::new(info, server_connections.clone())?;
databases.insert(database.name().to_owned(), Arc::new(database));
}
Ok(Self { server_connections, databases_cache: RwLock::new(databases) })
}
#[cfg_attr(feature = "sync", doc = "driver.databases().all();")]
#[cfg_attr(not(feature = "sync"), doc = "driver.databases().all().await;")]
#[cfg_attr(feature = "sync", maybe_async::must_be_sync)]
pub async fn all(&self) -> Result<Vec<Arc<Database>>> {
let mut error_buffer = Vec::with_capacity(self.server_connections.len());
for (server_id, server_connection) in self.server_connections.iter() {
match server_connection.all_databases().await {
Ok(list) => {
let mut new_databases: Vec<Arc<Database>> = Vec::new();
for db_info in list {
new_databases.push(Arc::new(Database::new(db_info, self.server_connections.clone())?));
}
let mut databases = self.databases_cache.write().unwrap();
databases.clear();
databases
.extend(new_databases.iter().map(|database| (database.name().to_owned(), database.clone())));
return Ok(new_databases);
}
Err(err) => error_buffer.push(format!("- {}: {}", server_id, err)),
}
}
Err(ConnectionError::ServerConnectionFailedWithError { error: error_buffer.join("\n") })?
}
#[cfg_attr(feature = "sync", doc = "driver.databases().get(name);")]
#[cfg_attr(not(feature = "sync"), doc = "driver.databases().get(name).await;")]
#[cfg_attr(feature = "sync", maybe_async::must_be_sync)]
pub async fn get(&self, name: impl AsRef<str>) -> Result<Arc<Database>> {
let name = name.as_ref();
if !self.contains(name.to_owned()).await? {
self.databases_cache.write().unwrap().remove(name);
return Err(ConnectionError::DatabaseNotFound { name: name.to_owned() }.into());
}
if let Some(cached_database) = self.try_get_cached(name) {
return Ok(cached_database);
}
self.cache_insert(Database::get(name.to_owned(), self.server_connections.clone()).await?);
Ok(self.try_get_cached(name).unwrap())
}
#[cfg_attr(feature = "sync", doc = "driver.databases().contains(name);")]
#[cfg_attr(not(feature = "sync"), doc = "driver.databases().contains(name).await;")]
#[cfg_attr(feature = "sync", maybe_async::must_be_sync)]
pub async fn contains(&self, name: impl Into<String>) -> Result<bool> {
let name = name.into();
self.run_failsafe(
name,
|server_connection, name| async move { server_connection.contains_database(name).await },
)
.await
}
#[cfg_attr(feature = "sync", doc = "driver.databases().create(name);")]
#[cfg_attr(not(feature = "sync"), doc = "driver.databases().create(name).await;")]
#[cfg_attr(feature = "sync", maybe_async::must_be_sync)]
pub async fn create(&self, name: impl Into<String>) -> Result {
let name = name.into();
let database_info = self
.run_failsafe(name, |server_connection, name| async move { server_connection.create_database(name).await }) .await?;
self.cache_insert(Database::new(database_info, self.server_connections.clone())?);
Ok(())
}
#[cfg_attr(feature = "sync", doc = "driver.databases().import_from_file(name, schema, data_path);")]
#[cfg_attr(not(feature = "sync"), doc = "driver.databases().import_from_file(name, schema, data_path).await;")]
#[cfg_attr(feature = "sync", maybe_async::must_be_sync)]
pub async fn import_from_file(
&self,
name: impl Into<String>,
schema: impl Into<String>,
data_file_path: impl AsRef<Path>,
) -> Result {
const ITEM_BATCH_SIZE: usize = 250;
let name = name.into();
let schema: String = schema.into();
let schema_ref: &str = schema.as_ref();
let data_file_path = data_file_path.as_ref();
self.run_failsafe(name, |server_connection, name| async move {
let file = try_open_import_file(data_file_path)?;
let mut import_stream = server_connection.import_database(name, schema_ref.to_string()).await?;
let mut item_buffer = Vec::with_capacity(ITEM_BATCH_SIZE);
for item in ProtoMessageIterator::<Item, _>::new(BufReader::new(file)) {
let item = item?;
item_buffer.push(item);
if item_buffer.len() >= ITEM_BATCH_SIZE {
import_stream.send_items(item_buffer.split_off(0))?;
}
}
if !item_buffer.is_empty() {
import_stream.send_items(item_buffer)?;
}
resolve!(import_stream.done())
})
.await
}
#[cfg_attr(feature = "sync", maybe_async::must_be_sync)]
pub(crate) async fn get_cached_or_fetch(&self, name: &str) -> Result<Arc<Database>> {
match self.try_get_cached(name) {
Some(cached_database) => Ok(cached_database),
None => self.get(name).await,
}
}
fn try_get_cached(&self, name: &str) -> Option<Arc<Database>> {
self.databases_cache.read().unwrap().get(name).cloned()
}
fn cache_insert(&self, database: Database) {
self.databases_cache.write().unwrap().insert(database.name().to_owned(), Arc::new(database));
}
#[cfg_attr(feature = "sync", maybe_async::must_be_sync)]
async fn run_failsafe<F, P, R>(&self, name: String, task: F) -> Result<R>
where
F: Fn(ServerConnection, String) -> P,
P: Future<Output = Result<R>>,
{
let mut error_buffer = Vec::with_capacity(self.server_connections.len());
for (server_id, server_connection) in self.server_connections.iter() {
match task(server_connection.clone(), name.clone()).await {
Ok(res) => return Ok(res),
err @ Err(Error::Connection(ConnectionError::ServerConnectionIsClosed)) => return err,
Err(err) => error_buffer.push(format!("- {}: {}", server_id, err)),
}
}
Err(ConnectionError::ServerConnectionFailedWithError { error: error_buffer.join("\n") })?
}
}