use std::collections::HashMap;
use std::fmt::Debug;
use std::hash::{Hash, Hasher};
use std::sync::Arc;
use anyhow::Result;
use async_graphql::dynamic::Schema;
use tokio::sync::RwLock;
use super::error::GqlError;
use super::schema::generate_schema;
use crate::catalog::providers::{AuthorisationProvider, DatabaseProvider, TableProvider};
use crate::catalog::{
DatabaseId, GraphQLConfig, GraphQLFunctionsConfig, GraphQLTablesConfig, NamespaceId,
};
use crate::dbs::Session;
use crate::kvs::{Datastore, Transaction};
type CacheKey = (String, String, GraphQLConfig, u64);
const SCHEMA_CACHE_MAX_ENTRIES: usize = 256;
#[derive(Clone, Default)]
pub struct GraphQLSchemaCache {
ns_db_schema_cache: Arc<RwLock<HashMap<CacheKey, Schema>>>,
}
impl Debug for GraphQLSchemaCache {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SchemaCache").field("ns_db_schema_cache", &self.ns_db_schema_cache).finish()
}
}
impl GraphQLSchemaCache {
pub async fn get_schema(
&self,
datastore: &Arc<Datastore>,
session: &Session,
) -> Result<Schema, GqlError> {
use crate::kvs::{LockType, TransactionType};
let ns = session.ns.as_ref().ok_or(GqlError::UnspecifiedNamespace)?;
let db = session.db.as_ref().ok_or(GqlError::UnspecifiedDatabase)?;
let kvs = datastore;
let tx = kvs.transaction(TransactionType::Read, LockType::Optimistic).await?;
let db_def = match tx.get_db_by_name(ns, db, None).await? {
Some(db) => db,
None => return Err(GqlError::NotConfigured),
};
let cg = tx
.expect_db_config(db_def.namespace_id, db_def.database_id, "graphql")
.await
.map_err(|e| {
if matches!(e.downcast_ref(), Some(crate::err::Error::CgNotFound { .. })) {
GqlError::NotConfigured
} else {
GqlError::DbError(e)
}
})?;
let gql_config = (*cg).clone().try_into_graphql()?;
let fingerprint =
compute_schema_fingerprint(&tx, db_def.namespace_id, db_def.database_id, &gql_config)
.await
.map_err(GqlError::DbError)?;
let cache_key = (ns.to_owned(), db.to_owned(), gql_config.clone(), fingerprint);
{
let guard = self.ns_db_schema_cache.read().await;
if let Some(cand) = guard.get(&cache_key) {
return Ok(cand.clone());
}
};
let schema = match generate_schema(datastore, session, gql_config).await {
Ok(s) => s,
Err(e) => {
if matches!(e, GqlError::DbError(_) | GqlError::SchemaError(_)) {
let mut guard = self.ns_db_schema_cache.write().await;
guard.remove(&cache_key);
}
return Err(e);
}
};
{
let mut guard = self.ns_db_schema_cache.write().await;
let (ns_key, db_key, cfg_key, _) = &cache_key;
guard.retain(|(n, d, c, _), _| !(n == ns_key && d == db_key && c == cfg_key));
while guard.len() >= SCHEMA_CACHE_MAX_ENTRIES {
if let Some(k) = guard.keys().next().cloned() {
guard.remove(&k);
} else {
break;
}
}
guard.insert(cache_key, schema.clone());
}
Ok(schema)
}
}
async fn compute_schema_fingerprint(
tx: &Transaction,
ns: NamespaceId,
db: DatabaseId,
gql_config: &GraphQLConfig,
) -> Result<u64> {
let mut hasher = std::collections::hash_map::DefaultHasher::new();
let tbs = tx.all_tb(ns, db, None).await?;
let mut tables_to_hash: Vec<&crate::catalog::TableDefinition> = match &gql_config.tables {
GraphQLTablesConfig::None => Vec::new(),
GraphQLTablesConfig::Auto => tbs.iter().collect(),
GraphQLTablesConfig::Include(inc) => tbs.iter().filter(|t| inc.contains(&t.name)).collect(),
GraphQLTablesConfig::Exclude(exc) => {
tbs.iter().filter(|t| !exc.contains(&t.name)).collect()
}
};
tables_to_hash.sort_by(|a, b| a.name.as_str().cmp(b.name.as_str()));
for tb in &tables_to_hash {
tb.hash(&mut hasher);
}
let fns = tx.all_db_functions(ns, db, None).await?;
let mut fns_to_hash: Vec<&crate::catalog::FunctionDefinition> = match &gql_config.functions {
GraphQLFunctionsConfig::None => Vec::new(),
GraphQLFunctionsConfig::Auto => fns.iter().collect(),
GraphQLFunctionsConfig::Include(inc) => {
fns.iter().filter(|f| inc.iter().any(|n| n.as_str() == f.name.as_str())).collect()
}
GraphQLFunctionsConfig::Exclude(exc) => {
fns.iter().filter(|f| !exc.iter().any(|n| n.as_str() == f.name.as_str())).collect()
}
};
fns_to_hash.sort_by(|a, b| a.name.as_str().cmp(b.name.as_str()));
for f in &fns_to_hash {
f.hash(&mut hasher);
}
let accesses = tx.all_db_accesses(ns, db, None).await?;
let mut accesses_sorted: Vec<&crate::catalog::AccessDefinition> = accesses.iter().collect();
accesses_sorted.sort_by(|a, b| a.name.as_str().cmp(b.name.as_str()));
for a in &accesses_sorted {
a.hash(&mut hasher);
}
Ok(hasher.finish())
}