pub mod convert;
pub mod exec;
pub mod session;
pub mod validate;
use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::Arc;
use std::time::{Duration, Instant};
use axum::extract::State;
use axum::{Extension, Json};
use serde::{Deserialize, Serialize};
use crate::config::{Config, S3Config, StorageType};
use crate::error::AppError;
use crate::handlers::AppState;
pub enum SqlTarget {
S3(S3Config),
Local { root_path: PathBuf },
}
pub struct SqlState {
pub cfg: crate::config::SqlConfig,
targets: HashMap<String, SqlTarget>,
disabled: HashMap<String, String>,
default_name: String,
}
impl SqlState {
pub fn from_config(cfg: &Config) -> Self {
let mut targets = HashMap::new();
let mut disabled = HashMap::new();
for s in &cfg.storages {
let target = match (s.r#type, &s.s3, &s.local) {
(StorageType::S3, Some(s3), _) => SqlTarget::S3(s3.clone()),
(StorageType::Local, _, Some(local)) => {
if !local.follow_symlinks {
disabled.insert(
s.name.clone(),
"storage has follow_symlinks = false, which the SQL sandbox \
cannot enforce (DuckDB follows symlinks inside the root)"
.into(),
);
continue;
}
SqlTarget::Local {
root_path: local.root_path.clone(),
}
}
_ => continue,
};
targets.insert(s.name.clone(), target);
}
let default_name = cfg
.active_storage()
.map(|s| s.name.clone())
.unwrap_or_default();
Self {
cfg: cfg.sql.clone(),
targets,
disabled,
default_name,
}
}
fn resolve(&self, name: Option<&str>) -> Result<&SqlTarget, AppError> {
let name = name.unwrap_or(&self.default_name);
if let Some(t) = self.targets.get(name) {
return Ok(t);
}
if let Some(reason) = self.disabled.get(name) {
return Err(AppError::Unsupported(format!(
"SQL is not available on storage '{name}': {reason}"
)));
}
Err(AppError::NotFound(format!("unknown storage: {name}")))
}
}
#[derive(Debug, Deserialize)]
pub struct QueryRequest {
pub sql: String,
#[serde(default)]
pub storage: Option<String>,
}
#[derive(Debug, Serialize)]
pub struct ColumnInfo {
pub name: String,
pub r#type: String,
}
#[derive(Debug, Serialize)]
pub struct QueryResponse {
pub columns: Vec<ColumnInfo>,
pub rows: Vec<Vec<serde_json::Value>>,
pub row_count: usize,
pub truncated: bool,
pub elapsed_ms: u64,
}
pub async fn query_handler(
State(state): State<AppState>,
Extension(sql_state): Extension<Arc<SqlState>>,
Json(req): Json<QueryRequest>,
) -> Result<Json<QueryResponse>, AppError> {
if !state.sql_enabled() {
return Err(AppError::Forbidden(
"SQL endpoint disabled: requires auth.enabled = true (and [sql] enabled)".into(),
));
}
let target = sql_state.resolve(req.storage.as_deref())?;
validate::validate_readonly(&req.sql)?;
let setup = session::setup_statements(&sql_state.cfg, target)?;
let timeout_secs = sql_state.cfg.query_timeout_secs;
let max_rows = sql_state.cfg.max_rows as usize;
let user_sql = req.sql.clone();
let start = Instant::now();
let (tx, rx) = tokio::sync::oneshot::channel();
let join = tokio::task::spawn_blocking(move || {
let conn = match duckdb::Connection::open_in_memory() {
Ok(c) => c,
Err(e) => return (Err(AppError::Backend(format!("open duckdb: {e}"))), None),
};
let _ = tx.send(conn.interrupt_handle());
if let Err(e) = conn.execute_batch(&setup) {
return (
Err(AppError::Backend(format!("sql session setup: {e}"))),
Some(conn),
);
}
(exec::run_query(&conn, &user_sql, max_rows), Some(conn))
});
let watchdog = tokio::spawn(async move {
if let Ok(handle) = rx.await {
tokio::time::sleep(Duration::from_secs(timeout_secs)).await;
handle.interrupt();
}
});
let joined = join.await;
watchdog.abort();
let _ = watchdog.await;
let (result, conn) = joined.map_err(|e| AppError::Backend(format!("query task failed: {e}")))?;
drop(conn);
let elapsed = start.elapsed();
let (columns, rows, truncated) = match result {
Ok(ok) => ok,
Err(AppError::Query(_)) if elapsed >= Duration::from_secs(timeout_secs) => {
return Err(AppError::QueryTimeout(timeout_secs));
}
Err(e) => return Err(e),
};
Ok(Json(QueryResponse {
row_count: rows.len(),
columns,
rows,
truncated,
elapsed_ms: elapsed.as_millis() as u64,
}))
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use super::*;
use crate::config::{LocalConfig, SqlConfig, StorageConfig};
use crate::storage::factory::BackendRegistry;
fn app_state(sql_enabled: bool) -> AppState {
let reg = BackendRegistry {
backends: HashMap::new(),
invalid: HashMap::new(),
order: vec![],
default_name: "t".into(),
};
AppState::new(
reg,
None,
std::sync::Arc::new("test".into()),
true,
true,
sql_enabled,
)
}
fn sql_state(root: &std::path::Path, follow_symlinks: bool) -> Arc<SqlState> {
let cfg = Config {
server: Default::default(),
storages: vec![StorageConfig {
name: "t".into(),
r#type: StorageType::Local,
active: true,
s3: None,
local: Some(LocalConfig {
root_path: root.to_path_buf(),
follow_symlinks,
}),
}],
auth: Default::default(),
thumbnails: Default::default(),
sql: SqlConfig::default(),
};
Arc::new(SqlState::from_config(&cfg))
}
fn req(sql: &str) -> Json<QueryRequest> {
Json(QueryRequest {
sql: sql.into(),
storage: Some("t".into()),
})
}
fn tmp_root() -> std::path::PathBuf {
let dir = std::env::temp_dir().join("omni-sql-handler-test");
std::fs::create_dir_all(&dir).unwrap();
dir
}
#[tokio::test]
async fn rejects_when_sql_disabled() {
let res = query_handler(
State(app_state(false)),
Extension(sql_state(&tmp_root(), true)),
req("SELECT 1"),
)
.await;
assert!(matches!(res, Err(AppError::Forbidden(_))), "{res:?}");
}
#[tokio::test]
async fn rejects_invalid_sql_before_execution() {
let res = query_handler(
State(app_state(true)),
Extension(sql_state(&tmp_root(), true)),
req("DROP TABLE t"),
)
.await;
assert!(matches!(res, Err(AppError::QueryRejected(_))), "{res:?}");
}
#[tokio::test]
async fn rejects_copy_statement() {
let res = query_handler(
State(app_state(true)),
Extension(sql_state(&tmp_root(), true)),
req("COPY (SELECT 1) TO 'out.parquet' (FORMAT PARQUET)"),
)
.await;
assert!(matches!(res, Err(AppError::QueryRejected(_))), "{res:?}");
}
#[tokio::test]
async fn rejects_unknown_storage() {
let res = query_handler(
State(app_state(true)),
Extension(sql_state(&tmp_root(), true)),
Json(QueryRequest {
sql: "SELECT 1".into(),
storage: Some("nope".into()),
}),
)
.await;
assert!(matches!(res, Err(AppError::NotFound(_))), "{res:?}");
}
#[tokio::test]
async fn rejects_no_follow_symlinks_storage_with_reason() {
let res = query_handler(
State(app_state(true)),
Extension(sql_state(&tmp_root(), false)),
req("SELECT 1"),
)
.await;
match res {
Err(AppError::Unsupported(msg)) => {
assert!(msg.contains("follow_symlinks"), "{msg}");
}
other => panic!("expected Unsupported, got {other:?}"),
}
}
#[tokio::test]
async fn executes_select_end_to_end() {
let res = query_handler(
State(app_state(true)),
Extension(sql_state(&tmp_root(), true)),
req("SELECT 1 AS a, 'x' AS b"),
)
.await
.expect("query should succeed");
let body = res.0;
assert_eq!(body.row_count, 1);
assert_eq!(body.columns.len(), 2);
assert_eq!(body.columns[0].name, "a");
assert!(!body.truncated);
assert_eq!(body.rows[0][0], serde_json::json!(1));
assert_eq!(body.rows[0][1], serde_json::json!("x"));
}
#[tokio::test]
async fn select_query_runs_without_auth_token() {
let res = query_handler(
State(app_state(true)),
Extension(sql_state(&tmp_root(), true)),
req("SELECT 1 AS a"),
)
.await
.expect("read-only query should succeed");
assert_eq!(res.0.row_count, 1);
}
}