use super::documents::get_transaction_id;
use super::system::AppState;
use crate::{
error::DbError,
sdbql::{BodyClause, Query, QueryExecutor},
server::response::ApiResponse,
storage::{query_cache, StorageEngine},
};
use axum::{
extract::{Path, State},
http::{HeaderMap, StatusCode},
response::Json,
};
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::sync::Arc;
const QUERY_TIMEOUT_SECS: u64 = 30;
const SLOW_QUERY_THRESHOLD_MS: f64 = 100.0;
const MAX_BATCH_SIZE: usize = 10_000;
#[derive(Debug, Deserialize)]
pub struct ExecuteQueryRequest {
pub query: String,
#[serde(default, alias = "bindVars")]
pub bind_vars: std::collections::HashMap<String, Value>,
#[serde(default = "default_batch_size", alias = "batchSize")]
pub batch_size: usize,
#[serde(default = "default_cache")]
pub cache: bool,
}
fn default_cache() -> bool {
true
}
fn default_batch_size() -> usize {
1000
}
#[derive(Debug, Serialize)]
pub struct ExecuteQueryResponse {
pub result: Vec<Value>,
pub count: usize,
pub has_more: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub id: Option<String>,
pub cached: bool,
#[serde(rename = "executionTimeMs")]
pub execution_time_ms: f64,
#[serde(rename = "inserted")]
pub documents_inserted: usize,
#[serde(rename = "updated")]
pub documents_updated: usize,
#[serde(rename = "deleted")]
pub documents_removed: usize,
}
pub(crate) fn is_long_running_query(query: &Query) -> bool {
let own = query.body_clauses.iter().any(|clause| match clause {
BodyClause::Insert(_)
| BodyClause::Update(_)
| BodyClause::Remove(_)
| BodyClause::Upsert(_) => true,
BodyClause::For(_) => true,
_ => false,
});
own || query
.set_operations
.iter()
.any(|op| is_long_running_query(&op.query))
|| query.with_clause.as_ref().is_some_and(|with| {
with.ctes
.iter()
.any(|cte| is_long_running_query(&cte.query))
})
}
pub(crate) fn invalidate_collections(collections: &[String]) {
if collections.is_empty() {
query_cache::get_query_cache().invalidate_all();
} else {
for collection in collections {
query_cache::get_query_cache().invalidate_collection(collection);
}
}
}
pub(crate) fn mutated_collections(query: &Query) -> std::collections::HashSet<&str> {
let mut collections: std::collections::HashSet<&str> = query
.body_clauses
.iter()
.filter_map(|clause| match clause {
BodyClause::Insert(c) => Some(c.collection.as_str()),
BodyClause::Update(c) => Some(c.collection.as_str()),
BodyClause::Remove(c) => Some(c.collection.as_str()),
BodyClause::Upsert(c) => Some(c.collection.as_str()),
_ => None,
})
.collect();
for operand in query.set_operations.iter().map(|op| op.query.as_ref()) {
collections.extend(mutated_collections(operand));
}
if let Some(with) = &query.with_clause {
for cte in &with.ctes {
collections.extend(mutated_collections(&cte.query));
}
}
collections
}
#[allow(clippy::too_many_arguments)]
fn log_slow_query(
storage: Arc<StorageEngine>,
db_name: String,
query_text: String,
execution_time_ms: f64,
results_count: usize,
documents_inserted: usize,
documents_updated: usize,
documents_removed: usize,
origin: Option<String>,
cf_ops_during: u64,
cf_ops_ms_during: f64,
) {
if execution_time_ms < SLOW_QUERY_THRESHOLD_MS {
return;
}
if query_text.contains("_slow_queries") {
return;
}
tokio::spawn(async move {
let slow_query_coll = format!("{}:_slow_queries", db_name);
let collection = match storage.get_collection(&slow_query_coll) {
Ok(coll) => coll,
Err(_) => {
if let Ok(db) = storage.get_database(&db_name) {
let _ = db.create_collection("_slow_queries".to_string(), None);
} else {
return;
}
let mut last_err = None;
let mut resolved = None;
for _ in 0..10 {
match storage.get_collection(&slow_query_coll) {
Ok(coll) => {
resolved = Some(coll);
break;
}
Err(e) => {
last_err = Some(e);
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
}
}
if let Some(coll) = resolved {
coll
} else {
tracing::warn!(
"Failed to get _slow_queries collection after retries: {}",
last_err
.map(|e| e.to_string())
.unwrap_or_else(|| "unknown error".to_string())
);
return;
}
}
};
let log_entry = serde_json::json!({
"query": query_text,
"execution_time_ms": execution_time_ms,
"timestamp": chrono::Utc::now().to_rfc3339(),
"results_count": results_count,
"documents_inserted": documents_inserted,
"documents_updated": documents_updated,
"documents_removed": documents_removed,
"origin": origin,
"cf_ops_during": cf_ops_during,
"cf_ops_ms_during": cf_ops_ms_during
});
if let Err(e) = collection.insert(log_entry) {
tracing::warn!("Failed to log slow query: {}", e);
}
});
}
pub(crate) fn principal_from_claims(
claims: &crate::server::auth::Claims,
) -> crate::sdbql::QueryPrincipal {
let roles = claims.roles.clone().unwrap_or_default();
let lower: Vec<String> = roles.iter().map(|r| r.to_ascii_lowercase()).collect();
let can_admin = lower.iter().any(|r| r == "admin");
let can_write = can_admin || lower.iter().any(|r| r == "editor" || r == "write");
let can_read = can_write || lower.iter().any(|r| r == "viewer" || r == "read");
crate::sdbql::QueryPrincipal {
user: claims.sub.clone(),
roles,
can_read,
can_write,
can_admin,
}
}
pub async fn execute_query(
State(state): State<AppState>,
Path(db_name): Path<String>,
headers: HeaderMap,
axum::Extension(claims): axum::Extension<crate::server::auth::Claims>,
Json(req): Json<ExecuteQueryRequest>,
) -> Result<ApiResponse<ExecuteQueryResponse>, DbError> {
state
.query_counter
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
{
let prepared = crate::sdbql::get_prepared_statement_cache().parse_if_needed(&req.query)?;
if prepared.query.has_mutations() {
crate::server::authz_middleware::enforce(
&claims,
&state,
crate::server::authorization::PermissionAction::Write,
Some(&db_name),
)
.await?;
}
}
if let Some(tx_id) = get_transaction_id(&headers) {
use crate::sdbql::ast::BodyClause;
let prepared = crate::sdbql::get_prepared_statement_cache().parse_if_needed(&req.query)?;
let query = prepared.query.as_ref();
let tx_manager = state.storage.transaction_manager()?;
let tx_arc = tx_manager.get(tx_id)?;
let mut tx = tx_arc
.write()
.map_err(|_| DbError::InternalError("Transaction lock poisoned".into()))?;
let wal = tx_manager.wal().clone();
let lock_manager = tx_manager.lock_manager().clone();
let has_mutations = query.has_mutations();
if !has_mutations {
let executor = if req.bind_vars.is_empty() {
QueryExecutor::with_database(&state.storage, db_name)
} else {
QueryExecutor::with_database_and_bind_vars(&state.storage, db_name, req.bind_vars)
};
let results = executor.execute(query)?;
return Ok(ApiResponse::new(
ExecuteQueryResponse {
result: results.clone(),
count: results.len(),
has_more: false,
id: None,
cached: false,
execution_time_ms: 0.0,
documents_inserted: 0,
documents_updated: 0,
documents_removed: 0,
},
&headers,
));
}
let executor = if req.bind_vars.is_empty() {
QueryExecutor::with_database(&state.storage, db_name.clone())
} else {
QueryExecutor::with_database_and_bind_vars(
&state.storage,
db_name.clone(),
req.bind_vars.clone(),
)
};
let mut initial_bindings = std::collections::HashMap::new();
for (key, value) in &req.bind_vars {
initial_bindings.insert(format!("@{}", key), value.clone());
}
for let_clause in &query.let_clauses {
let value =
executor.evaluate_expr_with_context(&let_clause.expression, &initial_bindings)?;
initial_bindings.insert(let_clause.variable.clone(), value);
}
let mut rows: Vec<std::collections::HashMap<String, Value>> =
vec![initial_bindings.clone()];
let mut mutation_count = 0;
for clause in &query.body_clauses {
match clause {
BodyClause::For(for_clause) => {
let mut new_rows = Vec::new();
for ctx in &rows {
let docs = if let Some(ref expr) = for_clause.source_expression {
let value = executor.evaluate_expr_with_context(expr, ctx)?;
match value {
Value::Array(arr) => arr,
other => vec![other],
}
} else {
let source_name = for_clause
.source_variable
.as_ref()
.unwrap_or(&for_clause.collection);
if let Some(value) = ctx.get(source_name) {
match value {
Value::Array(arr) => arr.clone(),
other => vec![other.clone()],
}
} else {
let full_coll_name =
format!("{}:{}", db_name, for_clause.collection);
let collection = state.storage.get_collection(&full_coll_name)?;
let shard_config = collection.get_shard_config();
if let (Some(config), Some(coordinator)) =
(shard_config, &state.shard_coordinator)
{
let coordinator_clone = coordinator.clone();
let db_name_owned = db_name.to_string();
let coll_name_owned = for_clause.collection.clone();
let config_clone = config.clone();
match tokio::task::block_in_place(|| {
tokio::runtime::Handle::current().block_on(async {
coordinator_clone
.scan_all_shards(
&db_name_owned,
&coll_name_owned,
&config_clone,
)
.await
})
}) {
Ok(docs) => {
docs.into_iter().map(|d| d.to_value()).collect()
}
Err(e) => {
eprintln!("Scatter-gather failed: {:?}, using local shards only", e);
collection
.scan(None)
.into_iter()
.map(|d| d.to_value())
.collect()
}
}
} else {
collection
.scan(None)
.into_iter()
.map(|d| d.to_value())
.collect()
}
}
};
for doc in docs {
let mut new_ctx = ctx.clone();
new_ctx.insert(for_clause.variable.clone(), doc);
new_rows.push(new_ctx);
}
}
rows = new_rows;
}
BodyClause::Let(let_clause) => {
for ctx in &mut rows {
let value =
executor.evaluate_expr_with_context(&let_clause.expression, ctx)?;
ctx.insert(let_clause.variable.clone(), value);
}
}
BodyClause::Filter(filter_clause) | BodyClause::Search(filter_clause) => {
rows.retain(|ctx| {
executor
.evaluate_filter_with_context(&filter_clause.expression, ctx)
.unwrap_or(false)
});
}
BodyClause::Insert(insert_clause) => {
let full_coll_name = format!("{}:{}", db_name, insert_clause.collection);
let collection = state.storage.get_collection(&full_coll_name)?;
for ctx in &rows {
let doc_value =
executor.evaluate_expr_with_context(&insert_clause.document, ctx)?;
collection.insert_tx(&mut tx, &wal, &lock_manager, doc_value)?;
mutation_count += 1;
}
}
BodyClause::Update(update_clause) => {
let full_coll_name = format!("{}:{}", db_name, update_clause.collection);
let collection = state.storage.get_collection(&full_coll_name)?;
for ctx in &rows {
let selector_value =
executor.evaluate_expr_with_context(&update_clause.selector, ctx)?;
let key = match &selector_value {
Value::String(s) => s.clone(),
Value::Object(obj) => obj.get("_key")
.and_then(|v| v.as_str())
.map(|s| s.to_string())
.ok_or_else(|| DbError::ExecutionError(
"UPDATE: selector object must have a _key field".to_string()
))?,
_ => return Err(DbError::ExecutionError(
"UPDATE: selector must be a string key or an object with _key field".to_string()
)),
};
let changes_value =
executor.evaluate_expr_with_context(&update_clause.changes, ctx)?;
collection.update_tx(&mut tx, &wal, &lock_manager, &key, changes_value)?;
mutation_count += 1;
}
}
BodyClause::Remove(remove_clause) => {
let full_coll_name = format!("{}:{}", db_name, remove_clause.collection);
let collection = state.storage.get_collection(&full_coll_name)?;
for ctx in &rows {
let selector_value =
executor.evaluate_expr_with_context(&remove_clause.selector, ctx)?;
let key = match &selector_value {
Value::String(s) => s.clone(),
Value::Object(obj) => obj.get("_key")
.and_then(|v| v.as_str())
.map(|s| s.to_string())
.ok_or_else(|| DbError::ExecutionError(
"REMOVE: selector object must have a _key field".to_string()
))?,
_ => return Err(DbError::ExecutionError(
"REMOVE: selector must be a string key or an object with _key field".to_string()
)),
};
collection.delete_tx(&mut tx, &wal, &lock_manager, &key)?;
mutation_count += 1;
}
}
_ => {}
}
}
return Ok(ApiResponse::new(
ExecuteQueryResponse {
result: vec![serde_json::json!({
"mutationCount": mutation_count,
"message": format!("{} operation(s) staged in transaction. Commit to apply changes.", mutation_count)
})],
count: 1,
has_more: false,
id: None,
cached: false,
execution_time_ms: 0.0,
documents_inserted: 0, documents_updated: 0,
documents_removed: 0,
},
&headers,
));
}
let prepared = crate::sdbql::get_prepared_statement_cache().parse_if_needed(&req.query)?;
let query = prepared.query.as_ref();
let is_read_only = !query.has_mutations();
let cache_key = if is_read_only && req.cache {
Some(query_cache::hash_query(
&db_name,
&req.query,
&req.bind_vars,
))
} else {
None
};
if let Some(ref key) = cache_key {
let cached_result = query_cache::get_query_cache().get(key);
if let Some(result) = cached_result {
tracing::debug!(
"Query cache hit for: {}",
&req.query[..req.query.len().min(50)]
);
return Ok(ApiResponse::new(
ExecuteQueryResponse {
result: result.as_ref().clone(),
count: result.len(),
has_more: false,
id: None,
cached: true,
execution_time_ms: 0.0,
documents_inserted: 0,
documents_updated: 0,
documents_removed: 0,
},
&headers,
));
}
}
if let Some(ref _create_stream) = query.create_stream_clause {
if let Some(manager) = &state.stream_manager {
match manager.create_stream(&db_name, (*query).clone()) {
Ok(_name) => {
return Ok(ApiResponse::new(
ExecuteQueryResponse {
result: Vec::new(),
count: 0,
has_more: false,
id: None,
cached: false,
execution_time_ms: 0.0,
documents_inserted: 0,
documents_updated: 0,
documents_removed: 0,
},
&headers,
));
}
Err(e) => return Err(e),
}
} else {
return Err(DbError::OperationNotSupported(
"Stream processing not enabled".to_string(),
));
}
}
let batch_size = req.batch_size.min(MAX_BATCH_SIZE);
let db_name_for_logging = db_name.clone();
let db_name_for_cursor = db_name.clone();
let query_text_for_logging = req.query.clone();
let cf_ops_before = crate::storage::cf_ops::snapshot();
let mutates = query.has_mutations();
let (query_result, execution_time_ms) = if is_long_running_query(query) {
let storage = state.storage.clone();
let bind_vars = req.bind_vars.clone();
let replication_log = state.replication_log.clone();
let shard_coordinator = state.shard_coordinator.clone();
let is_scatter_gather = headers.contains_key("X-Scatter-Gather");
let query = (*query).clone();
let principal = principal_from_claims(&claims);
let invalidated: Vec<String> = if mutates {
mutated_collections(&query)
.into_iter()
.map(|c| c.to_string())
.collect()
} else {
Vec::new()
};
let mut task = tokio::task::spawn_blocking(move || {
let mut executor = if bind_vars.is_empty() {
QueryExecutor::with_database(&storage, db_name)
} else {
QueryExecutor::with_database_and_bind_vars(&storage, db_name, bind_vars)
}
.with_principal(principal);
if let Some(ref log) = replication_log {
executor = executor.with_replication(log);
}
if !is_scatter_gather {
if let Some(coord) = shard_coordinator {
executor = executor.with_shard_coordinator(coord);
}
}
let start = std::time::Instant::now();
let result = executor.execute_with_stats(&query)?;
let execution_time_ms = start.elapsed().as_secs_f64() * 1000.0;
Ok::<_, DbError>((result, execution_time_ms))
});
match tokio::time::timeout(
std::time::Duration::from_secs(QUERY_TIMEOUT_SECS),
&mut task,
)
.await
{
Ok(join_result) => join_result
.map_err(|e| DbError::InternalError(format!("Task join error: {}", e)))??,
Err(_) => {
if mutates {
invalidate_collections(&invalidated);
tokio::spawn(async move {
let _ = task.await;
invalidate_collections(&invalidated);
});
}
return Err(DbError::BadRequest(format!(
"Query execution timeout: exceeded {} seconds",
QUERY_TIMEOUT_SECS
)));
}
}
} else {
let mut executor = if req.bind_vars.is_empty() {
QueryExecutor::with_database(&state.storage, db_name)
} else {
QueryExecutor::with_database_and_bind_vars(&state.storage, db_name, req.bind_vars)
}
.with_principal(principal_from_claims(&claims));
if let Some(ref log) = state.replication_log {
executor = executor.with_replication(log);
}
if !headers.contains_key("X-Scatter-Gather") {
if let Some(coordinator) = state.shard_coordinator.clone() {
executor = executor.with_shard_coordinator(coordinator);
}
}
let start = std::time::Instant::now();
let result = executor.execute_with_stats(query)?;
let execution_time_ms = start.elapsed().as_secs_f64() * 1000.0;
(result, execution_time_ms)
};
let total_count = query_result.results.len();
let mutations = &query_result.mutations;
if mutations.has_mutations() {
state
.write_counter
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
if let Some(key) = cache_key {
let result_clone: Vec<serde_json::Value> = query_result.results.clone();
query_cache::get_query_cache().put(key, result_clone);
}
if mutations.has_mutations() {
let collections: Vec<String> = mutated_collections(query)
.into_iter()
.map(|c| c.to_string())
.collect();
invalidate_collections(&collections);
}
let cf_ops_after = crate::storage::cf_ops::snapshot();
log_slow_query(
state.storage.clone(),
db_name_for_logging,
query_text_for_logging,
execution_time_ms,
total_count,
mutations.documents_inserted,
mutations.documents_updated,
mutations.documents_removed,
Some(claims.sub.clone()),
cf_ops_before.ops_since(&cf_ops_after),
cf_ops_before.ms_since(&cf_ops_after),
);
let (cursor_id, result_batch, has_more) = state.cursor_store.store_and_get_first_batch(
db_name_for_cursor,
query_result.results,
batch_size,
);
Ok(ApiResponse::new(
ExecuteQueryResponse {
result: result_batch,
count: total_count,
has_more,
id: cursor_id,
cached: false,
execution_time_ms,
documents_inserted: mutations.documents_inserted,
documents_updated: mutations.documents_updated,
documents_removed: mutations.documents_removed,
},
&headers,
))
}
pub async fn explain_query(
State(state): State<AppState>,
Path(db_name): Path<String>,
headers: HeaderMap,
axum::Extension(claims): axum::Extension<crate::server::auth::Claims>,
Json(req): Json<ExecuteQueryRequest>,
) -> Result<Json<crate::sdbql::QueryExplain>, DbError> {
let prepared = crate::sdbql::get_prepared_statement_cache().parse_if_needed(&req.query)?;
let query = (*prepared.query).clone();
let bind_vars = req.bind_vars.clone();
let storage = state.storage.clone();
let shard_coordinator = state.shard_coordinator.clone();
let is_scatter_gather = headers.contains_key("X-Scatter-Gather");
let principal = principal_from_claims(&claims);
let explain = {
let storage = storage.clone();
match tokio::time::timeout(
std::time::Duration::from_secs(QUERY_TIMEOUT_SECS),
tokio::task::spawn_blocking(move || {
let mut executor = if bind_vars.is_empty() {
QueryExecutor::with_database(&storage, db_name)
} else {
QueryExecutor::with_database_and_bind_vars(&storage, db_name, bind_vars)
}
.with_principal(principal);
if !is_scatter_gather {
if let Some(coordinator) = shard_coordinator {
executor = executor.with_shard_coordinator(coordinator);
}
}
executor.explain(&query)
}),
)
.await
{
Ok(join_result) => join_result
.map_err(|e| DbError::InternalError(format!("Task join error: {}", e)))??,
Err(_) => {
return Err(DbError::BadRequest(format!(
"Explain timeout: exceeded {} seconds",
QUERY_TIMEOUT_SECS
)))
}
}
};
Ok(Json(explain))
}
pub async fn get_next_batch(
State(state): State<AppState>,
Path(cursor_id): Path<String>,
headers: HeaderMap,
axum::Extension(claims): axum::Extension<crate::server::auth::Claims>,
) -> Result<ApiResponse<ExecuteQueryResponse>, DbError> {
if let Some(db) = state.cursor_store.db_name(&cursor_id) {
crate::server::authz_middleware::enforce(
&claims,
&state,
crate::server::authorization::PermissionAction::Read,
Some(&db),
)
.await?;
}
if let Some((batch, has_more)) = state.cursor_store.get_next_batch(&cursor_id) {
let count = batch.len();
Ok(ApiResponse::new(
ExecuteQueryResponse {
result: batch,
count,
has_more,
id: if has_more { Some(cursor_id) } else { None },
cached: true,
execution_time_ms: 0.0, documents_inserted: 0, documents_updated: 0,
documents_removed: 0,
},
&headers,
))
} else {
Err(DbError::DocumentNotFound(format!(
"Cursor not found or expired: {}",
cursor_id
)))
}
}
pub async fn delete_cursor(
State(state): State<AppState>,
Path(cursor_id): Path<String>,
axum::Extension(claims): axum::Extension<crate::server::auth::Claims>,
) -> Result<StatusCode, DbError> {
if let Some(db) = state.cursor_store.db_name(&cursor_id) {
crate::server::authz_middleware::enforce(
&claims,
&state,
crate::server::authorization::PermissionAction::Read,
Some(&db),
)
.await?;
}
if state.cursor_store.delete(&cursor_id) {
Ok(StatusCode::NO_CONTENT)
} else {
Err(DbError::DocumentNotFound(format!(
"Cursor not found: {}",
cursor_id
)))
}
}