use crate::control::planner::procedural::executor::bindings::RowBindings;
use crate::control::planner::procedural::executor::core::StatementExecutor;
use crate::control::planner::procedural::executor::fuel::ExecutionBudget;
use crate::control::security::catalog::procedure_types::ParamDirection;
use crate::control::security::identity::AuthenticatedIdentity;
use crate::control::server::response_shape::types::ShapedRows;
use crate::control::state::SharedState;
use super::super::super::result::{DdlError, DdlResult};
use super::parens::find_matching_paren;
pub async fn call_procedure(
state: &SharedState,
identity: &AuthenticatedIdentity,
sql: &str,
) -> Result<Vec<DdlResult>, DdlError> {
let (name, args) = parse_call(sql)?;
let tenant_id = identity.tenant_id;
let catalog = state.credentials.catalog();
let proc = catalog
.get_procedure(tenant_id.as_u64(), &name)
.map_err(|e| DdlError {
sqlstate: "XX000".to_string(),
message: e.to_string(),
})?
.ok_or_else(|| DdlError {
sqlstate: "42883".to_string(),
message: format!("procedure '{name}' does not exist"),
})?;
let in_params: Vec<_> = proc
.parameters
.iter()
.filter(|p| matches!(p.direction, ParamDirection::In | ParamDirection::InOut))
.collect();
if args.len() != in_params.len() {
return Err(DdlError {
sqlstate: "42601".to_string(),
message: format!(
"procedure '{}' expects {} argument(s), got {}",
name,
in_params.len(),
args.len()
),
});
}
let mut param_map = std::collections::HashMap::new();
for (param, arg) in in_params.iter().zip(args.iter()) {
param_map.insert(param.name.clone(), arg.clone());
}
let bindings = RowBindings::with_params(param_map);
let block =
crate::control::planner::procedural::parse_block(&proc.body_sql).map_err(|e| DdlError {
sqlstate: "42601".to_string(),
message: format!("procedure body parse error: {e}"),
})?;
let mut budget = ExecutionBudget::new(proc.max_iterations, proc.timeout_secs);
let executor =
StatementExecutor::new(state, identity.clone(), tenant_id, 0).with_transaction_context();
executor
.execute_block_with_budget(&block, &bindings, &mut budget)
.await
.map_err(|e| DdlError {
sqlstate: "P0001".to_string(),
message: e.to_string(),
})?;
let out_params: Vec<_> = proc
.parameters
.iter()
.filter(|p| matches!(p.direction, ParamDirection::Out | ParamDirection::InOut))
.collect();
if out_params.is_empty() {
return Ok(vec![DdlResult::Status {
command: "CALL".to_string(),
rows_affected: None,
}]);
}
let out_values = executor.take_out_values();
Ok(build_out_response(&out_params, &out_values))
}
fn build_out_response(
out_params: &[&crate::control::security::catalog::procedure_types::ProcedureParam],
out_values: &std::collections::HashMap<String, nodedb_types::Value>,
) -> Vec<DdlResult> {
let columns: Vec<String> = out_params.iter().map(|p| p.name.clone()).collect();
let mut row = serde_json::Map::new();
for param in out_params {
let value = out_values
.get(¶m.name)
.or_else(|| {
if out_params.len() == 1 {
out_values.get("__return")
} else {
None
}
});
let text = match value {
Some(nodedb_types::Value::Null) | None => String::new(),
Some(nodedb_types::Value::String(s)) => s.clone(),
Some(v) => v.to_sql_literal(),
};
row.insert(param.name.clone(), serde_json::Value::String(text));
}
let column_types = ShapedRows::text_types(columns.len());
vec![DdlResult::Rows(ShapedRows {
columns,
column_types,
rows: vec![row],
notice: None,
})]
}
fn parse_call(sql: &str) -> Result<(String, Vec<String>), DdlError> {
let trimmed = sql.trim().trim_end_matches(';').trim();
let upper = trimmed.to_uppercase();
if !upper.starts_with("CALL ") {
return Err(DdlError {
sqlstate: "42601".to_string(),
message: "expected CALL <procedure>(...)".to_string(),
});
}
let after_call = &trimmed["CALL ".len()..].trim();
let paren_pos = after_call.find('(').ok_or_else(|| DdlError {
sqlstate: "42601".to_string(),
message: "expected '(' after procedure name in CALL".to_string(),
})?;
let name = after_call[..paren_pos].trim().to_lowercase();
if name.is_empty() {
return Err(DdlError {
sqlstate: "42601".to_string(),
message: "procedure name required in CALL".to_string(),
});
}
let close_paren = find_matching_paren(after_call, paren_pos).ok_or_else(|| DdlError {
sqlstate: "42601".to_string(),
message: "unmatched '(' in CALL".to_string(),
})?;
let args_str = &after_call[paren_pos + 1..close_paren];
let args = if args_str.trim().is_empty() {
Vec::new()
} else {
split_call_args(args_str)
};
Ok((name, args))
}
fn split_call_args(s: &str) -> Vec<String> {
let mut args = Vec::new();
let mut current = String::new();
let mut depth = 0i32;
let mut in_string = false;
for ch in s.chars() {
if in_string {
current.push(ch);
if ch == '\'' {
in_string = false;
}
continue;
}
match ch {
'\'' => {
in_string = true;
current.push(ch);
}
'(' => {
depth += 1;
current.push(ch);
}
')' => {
depth -= 1;
current.push(ch);
}
',' if depth == 0 => {
args.push(current.trim().to_string());
current.clear();
}
_ => current.push(ch),
}
}
let last = current.trim().to_string();
if !last.is_empty() {
args.push(last);
}
args
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_call_basic() {
let (name, args) = parse_call("CALL archive(90)").unwrap();
assert_eq!(name, "archive");
assert_eq!(args, vec!["90"]);
}
#[test]
fn parse_call_multiple_args() {
let (name, args) = parse_call("CALL migrate('users', 100)").unwrap();
assert_eq!(name, "migrate");
assert_eq!(args, vec!["'users'", "100"]);
}
#[test]
fn parse_call_no_args() {
let (name, args) = parse_call("CALL cleanup()").unwrap();
assert_eq!(name, "cleanup");
assert!(args.is_empty());
}
#[test]
fn parse_call_nested_parens() {
let (_, args) = parse_call("CALL p(func(1, 2), 3)").unwrap();
assert_eq!(args, vec!["func(1, 2)", "3"]);
}
#[test]
fn parse_call_with_semicolon() {
let (name, _) = parse_call("CALL cleanup();").unwrap();
assert_eq!(name, "cleanup");
}
}