use crate::driver::protocol::{DriverError, Response};
use crate::driver::DriverHandler;
use crate::sdbql::QueryExecutor;
use crate::server::handlers::query::{
invalidate_collections, is_long_running_query, mutated_collections,
};
use crate::storage::query_cache;
use std::collections::HashMap;
const QUERY_TIMEOUT_SECS: u64 = 30;
pub async fn handle_query(
handler: &DriverHandler,
database: String,
sdbql: String,
bind_vars: Option<HashMap<String, serde_json::Value>>,
cache: bool,
) -> Response {
let bind_vars = bind_vars.unwrap_or_default();
let prepared = match crate::sdbql::get_prepared_statement_cache().parse_if_needed(&sdbql) {
Ok(p) => p,
Err(e) => {
return Response::error(DriverError::DatabaseError(format!("Parse error: {}", e)))
}
};
let query = prepared.query.as_ref();
let mutates = query.has_mutations();
if mutates {
if let Err(e) = crate::server::AuthorizationService::check_permission_raw(
&handler.session_permissions,
crate::server::PermissionAction::Write,
Some(&database),
handler.session_scoped_databases.as_deref(),
) {
return Response::error(DriverError::AuthError(e.to_string()));
}
}
let cache_key = if mutates || !cache {
None
} else {
Some(query_cache::hash_query(&database, &sdbql, &bind_vars))
};
if let Some(ref key) = cache_key {
if let Some(hit) = query_cache::get_query_cache().get(key) {
return Response::ok(serde_json::json!(hit.as_ref().clone()));
}
}
let principal = handler.query_principal(&database);
let invalidated: Vec<String> = if mutates {
mutated_collections(query)
.into_iter()
.map(|s| s.to_string())
.collect()
} else {
Vec::new()
};
let storage = handler.storage.clone();
let replication = handler.replication.clone();
if !is_long_running_query(query) {
let mut executor = if bind_vars.is_empty() {
QueryExecutor::with_database(&storage, database)
} else {
QueryExecutor::with_database_and_bind_vars(&storage, database, bind_vars)
}
.with_principal(principal);
if let Some(ref log) = replication {
executor = executor.with_replication(log);
}
return match executor.execute(query) {
Ok(results) => {
if let Some(key) = cache_key {
query_cache::get_query_cache().put(key, results.clone());
}
if mutates {
invalidate_collections(&invalidated);
}
Response::ok(serde_json::json!(results))
}
Err(e) => Response::error(DriverError::DatabaseError(e.to_string())),
};
}
let exec_query = (*query).clone();
let mut task = tokio::task::spawn_blocking(move || {
let mut executor = if bind_vars.is_empty() {
QueryExecutor::with_database(&storage, database)
} else {
QueryExecutor::with_database_and_bind_vars(&storage, database, bind_vars)
}
.with_principal(principal);
if let Some(ref log) = replication {
executor = executor.with_replication(log);
}
executor.execute(&exec_query)
});
match tokio::time::timeout(
std::time::Duration::from_secs(QUERY_TIMEOUT_SECS),
&mut task,
)
.await
{
Ok(join_result) => match join_result {
Ok(Ok(results)) => {
if let Some(key) = cache_key {
query_cache::get_query_cache().put(key, results.clone());
}
if mutates {
invalidate_collections(&invalidated);
}
Response::ok(serde_json::json!(results))
}
Ok(Err(e)) => Response::error(DriverError::DatabaseError(e.to_string())),
Err(e) => Response::error(DriverError::DatabaseError(format!(
"Task join error: {}",
e
))),
},
Err(_) => {
if mutates {
invalidate_collections(&invalidated);
tokio::spawn(async move {
let _ = task.await;
invalidate_collections(&invalidated);
});
}
Response::error(DriverError::DatabaseError(format!(
"Query execution timeout: exceeded {} seconds",
QUERY_TIMEOUT_SECS
)))
}
}
}
pub async fn handle_explain(
handler: &DriverHandler,
database: String,
sdbql: String,
bind_vars: Option<HashMap<String, serde_json::Value>>,
) -> Response {
let bind_vars = bind_vars.unwrap_or_default();
let prepared = match crate::sdbql::get_prepared_statement_cache().parse_if_needed(&sdbql) {
Ok(p) => p,
Err(e) => {
return Response::error(DriverError::DatabaseError(format!("Parse error: {}", e)))
}
};
let principal = handler.query_principal(&database);
let storage = handler.storage.clone();
let exec_query = (*prepared.query).clone();
let task = tokio::task::spawn_blocking(move || {
let executor = if bind_vars.is_empty() {
QueryExecutor::with_database(&storage, database)
} else {
QueryExecutor::with_database_and_bind_vars(&storage, database, bind_vars)
}
.with_principal(principal);
executor.explain(&exec_query)
});
match tokio::time::timeout(std::time::Duration::from_secs(QUERY_TIMEOUT_SECS), task).await {
Ok(join_result) => match join_result {
Ok(Ok(explanation)) => {
Response::ok(serde_json::to_value(explanation).unwrap_or_default())
}
Ok(Err(e)) => Response::error(DriverError::DatabaseError(e.to_string())),
Err(e) => Response::error(DriverError::DatabaseError(format!(
"Task join error: {}",
e
))),
},
Err(_) => Response::error(DriverError::DatabaseError(format!(
"Explain timeout: exceeded {} seconds",
QUERY_TIMEOUT_SECS
))),
}
}