use std::collections::HashMap;
use crate::velesql::{
Condition, DmlStatement, HavingClause, Query, SelectColumns, Subquery, Value,
};
use crate::{Error, Result};
use super::Database;
const MAX_SUBQUERY_DEPTH: u32 = 8;
impl Database {
pub(super) fn resolve_subqueries(
&self,
query: &Query,
params: &HashMap<String, serde_json::Value>,
) -> Result<Option<Query>> {
if !query_has_resolvable_subquery(query) {
return Ok(None);
}
let mut rewritten = query.clone();
let scope = query.select.outer_table_scope();
self.rewrite_query_subqueries(&mut rewritten, params, 0, &scope)?;
Ok(Some(rewritten))
}
fn rewrite_query_subqueries(
&self,
query: &mut Query,
params: &HashMap<String, serde_json::Value>,
depth: u32,
outer_tables: &[&str],
) -> Result<()> {
if let Some(cond) = query.select.where_clause.as_mut() {
self.rewrite_condition(cond, params, depth, outer_tables)?;
}
if let Some(having) = query.select.having.as_mut() {
self.rewrite_having(having, params, depth, outer_tables)?;
}
if let Some(dml) = query.dml.as_mut() {
self.rewrite_dml(dml, params, depth)?;
}
Ok(())
}
fn rewrite_condition(
&self,
cond: &mut Condition,
params: &HashMap<String, serde_json::Value>,
depth: u32,
outer_tables: &[&str],
) -> Result<()> {
match cond {
Condition::And(l, r) | Condition::Or(l, r) => {
self.rewrite_condition(l, params, depth, outer_tables)?;
self.rewrite_condition(r, params, depth, outer_tables)
}
Condition::Group(inner) | Condition::Not(inner) => {
self.rewrite_condition(inner, params, depth, outer_tables)
}
other => self.rewrite_leaf_condition(other, params, depth, outer_tables),
}
}
fn rewrite_leaf_condition(
&self,
cond: &mut Condition,
params: &HashMap<String, serde_json::Value>,
depth: u32,
outer_tables: &[&str],
) -> Result<()> {
match cond {
Condition::Comparison(c) => {
self.rewrite_value(&mut c.value, params, depth, outer_tables)
}
Condition::Between(c) => {
self.rewrite_value(&mut c.low, params, depth, outer_tables)?;
self.rewrite_value(&mut c.high, params, depth, outer_tables)
}
Condition::In(c) => self.rewrite_values(&mut c.values, params, depth, outer_tables),
Condition::Contains(c) => {
self.rewrite_values(&mut c.values, params, depth, outer_tables)
}
_ => Ok(()),
}
}
fn rewrite_having(
&self,
having: &mut HavingClause,
params: &HashMap<String, serde_json::Value>,
depth: u32,
outer_tables: &[&str],
) -> Result<()> {
for cond in &mut having.conditions {
self.rewrite_value(&mut cond.value, params, depth, outer_tables)?;
}
Ok(())
}
fn rewrite_dml(
&self,
dml: &mut DmlStatement,
params: &HashMap<String, serde_json::Value>,
depth: u32,
) -> Result<()> {
match dml {
DmlStatement::Insert(s) | DmlStatement::Upsert(s) => {
for row in &mut s.rows {
self.rewrite_values(row, params, depth, &[s.table.as_str()])?;
}
Ok(())
}
DmlStatement::Update(s) => {
let scope = [s.table.as_str()];
for assignment in &mut s.assignments {
self.rewrite_value(&mut assignment.value, params, depth, &scope)?;
}
if let Some(cond) = s.where_clause.as_mut() {
self.rewrite_condition(cond, params, depth, &scope)?;
}
Ok(())
}
DmlStatement::Delete(s) => {
self.rewrite_condition(&mut s.where_clause, params, depth, &[s.table.as_str()])
}
_ => Ok(()),
}
}
fn rewrite_value(
&self,
value: &mut Value,
params: &HashMap<String, serde_json::Value>,
depth: u32,
outer_tables: &[&str],
) -> Result<()> {
if let Value::Subquery(subquery) = value {
if !subquery.references_outer_table(outer_tables) {
*value = self.execute_scalar_subquery(subquery, params, depth)?;
}
}
Ok(())
}
fn rewrite_values(
&self,
values: &mut [Value],
params: &HashMap<String, serde_json::Value>,
depth: u32,
outer_tables: &[&str],
) -> Result<()> {
for value in values {
self.rewrite_value(value, params, depth, outer_tables)?;
}
Ok(())
}
fn execute_scalar_subquery(
&self,
subquery: &Subquery,
params: &HashMap<String, serde_json::Value>,
depth: u32,
) -> Result<Value> {
if depth >= MAX_SUBQUERY_DEPTH {
return Err(Error::Query(format!(
"scalar subquery nesting exceeds the maximum depth of {MAX_SUBQUERY_DEPTH}"
)));
}
let mut inner = Query::from_select(subquery.select.clone());
let inner_scope: Vec<String> = inner
.select
.outer_table_scope()
.iter()
.map(|s| (*s).to_string())
.collect();
let inner_scope: Vec<&str> = inner_scope.iter().map(String::as_str).collect();
self.rewrite_query_subqueries(&mut inner, params, depth + 1, &inner_scope)?;
self.run_scalar_subquery(&inner, params)
}
fn run_scalar_subquery(
&self,
inner: &Query,
params: &HashMap<String, serde_json::Value>,
) -> Result<Value> {
if inner.select.is_aggregation_query() {
let json = self.execute_aggregate(inner, params)?;
return scalar_from_aggregate(&json);
}
let column = projected_subquery_column(&inner.select)?;
let results = self.execute_query(inner, params)?;
scalar_from_rows(&results, &column)
}
}
fn query_has_resolvable_subquery(query: &Query) -> bool {
let scope = query.select.outer_table_scope();
let where_resolvable = query
.select
.where_clause
.as_ref()
.is_some_and(|c| condition_has_resolvable_subquery(c, &scope));
where_resolvable
|| having_has_resolvable_subquery(query)
|| dml_has_resolvable_subquery(query.dml.as_ref())
}
fn having_has_resolvable_subquery(query: &Query) -> bool {
query.has_having_subquery() && !query.has_correlated_having_subquery()
}
fn condition_has_resolvable_subquery(cond: &Condition, outer_tables: &[&str]) -> bool {
cond.has_subquery() && !cond.has_correlated_subquery(outer_tables)
}
fn dml_has_resolvable_subquery(dml: Option<&DmlStatement>) -> bool {
match dml {
Some(DmlStatement::Insert(s) | DmlStatement::Upsert(s)) => {
s.rows.iter().any(|row| row.iter().any(Value::is_subquery))
}
Some(DmlStatement::Update(s)) => {
let scope = [s.table.as_str()];
s.assignments.iter().any(|a| a.value.is_subquery())
|| s.where_clause
.as_ref()
.is_some_and(|c| condition_has_resolvable_subquery(c, &scope))
}
Some(DmlStatement::Delete(s)) => {
condition_has_resolvable_subquery(&s.where_clause, &[s.table.as_str()])
}
_ => false,
}
}
fn projected_subquery_column(select: &crate::velesql::SelectStatement) -> Result<String> {
match &select.columns {
SelectColumns::Columns(cols) if cols.len() == 1 => Ok(cols[0].name.clone()),
_ => Err(Error::Query(
"scalar subquery must select exactly one column (e.g. \
(SELECT amount FROM t WHERE ...) or an aggregate like (SELECT AVG(amount) FROM t))"
.to_string(),
)),
}
}
fn scalar_from_aggregate(json: &serde_json::Value) -> Result<Value> {
let obj = json
.as_object()
.ok_or_else(|| Error::Query("scalar subquery aggregate did not return an object".into()))?;
if obj.len() != 1 {
return Err(Error::Query(
"scalar subquery must return exactly one column".to_string(),
));
}
let only = obj
.values()
.next()
.ok_or_else(|| Error::Query("scalar subquery aggregate returned no column".to_string()))?;
json_to_value(only)
}
fn scalar_from_rows(results: &[crate::SearchResult], column: &str) -> Result<Value> {
match results {
[] => Ok(Value::Null),
[single] => {
let cell = single
.point
.payload
.as_ref()
.and_then(|p| nested_payload_value(p, column))
.unwrap_or(&serde_json::Value::Null);
json_to_value(cell)
}
_ => Err(Error::Query(format!(
"scalar subquery returned {} rows but must return at most one row",
results.len()
))),
}
}
fn nested_payload_value<'a>(
payload: &'a serde_json::Value,
path: &str,
) -> Option<&'a serde_json::Value> {
let mut current = payload;
for part in path.split('.') {
current = current.as_object()?.get(part)?;
}
Some(current)
}
fn json_to_value(json: &serde_json::Value) -> Result<Value> {
match json {
serde_json::Value::Null => Ok(Value::Null),
serde_json::Value::Bool(b) => Ok(Value::Boolean(*b)),
serde_json::Value::String(s) => Ok(Value::String(s.clone())),
serde_json::Value::Number(n) => number_to_value(n),
serde_json::Value::Array(_) | serde_json::Value::Object(_) => Err(Error::Query(
"scalar subquery column resolved to a non-scalar (array/object) value".to_string(),
)),
}
}
fn number_to_value(n: &serde_json::Number) -> Result<Value> {
if let Some(i) = n.as_i64() {
return Ok(Value::Integer(i));
}
if let Some(u) = n.as_u64() {
return Ok(Value::UnsignedInteger(u));
}
if let Some(f) = n.as_f64() {
return Ok(Value::Float(f));
}
Err(Error::Query(
"scalar subquery returned an unrepresentable number".to_string(),
))
}