use std::fmt;
use std::hash::{Hash, Hasher};
use std::sync::{Arc, Mutex};
use datafusion_common::{HashMap, Result, ScalarValue, TableReference, internal_err};
#[derive(Clone, Debug, Default)]
pub struct PhysicalPlanningContext {
indexes: Arc<HashMap<crate::logical_plan::Subquery, SubqueryIndex>>,
results: ScalarSubqueryResults,
lambda_variable_qualifier: HashMap<String, TableReference>,
}
impl PhysicalPlanningContext {
pub fn new(
indexes: HashMap<crate::logical_plan::Subquery, SubqueryIndex>,
results: ScalarSubqueryResults,
) -> Self {
Self {
indexes: Arc::new(indexes),
results,
lambda_variable_qualifier: HashMap::new(),
}
}
pub fn index_of(
&self,
subquery: &crate::logical_plan::Subquery,
) -> Option<SubqueryIndex> {
self.indexes.get(subquery).copied()
}
pub fn results(&self) -> &ScalarSubqueryResults {
&self.results
}
pub fn with_qualified_lambda_variables(
mut self,
qualifier: &TableReference,
variables: &[String],
) -> Self {
for var in variables {
self.lambda_variable_qualifier
.entry_ref(var)
.insert(qualifier.clone());
}
self
}
pub fn lambda_variable_qualifier(&self, name: &str) -> Option<&TableReference> {
self.lambda_variable_qualifier.get(name)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub struct SubqueryIndex(usize);
impl SubqueryIndex {
pub const fn new(index: usize) -> Self {
Self(index)
}
pub const fn as_usize(self) -> usize {
self.0
}
}
#[derive(Clone, Default)]
pub struct ScalarSubqueryResults {
slots: Arc<Vec<Mutex<Option<ScalarValue>>>>,
}
impl ScalarSubqueryResults {
pub fn new(n: usize) -> Self {
Self {
slots: Arc::new((0..n).map(|_| Mutex::new(None)).collect()),
}
}
pub fn get(&self, index: SubqueryIndex) -> Option<ScalarValue> {
let slot = self.slots.get(index.as_usize())?;
slot.lock().unwrap().clone()
}
pub fn set(&self, index: SubqueryIndex, value: ScalarValue) -> Result<()> {
let Some(slot) = self.slots.get(index.as_usize()) else {
return internal_err!(
"ScalarSubqueryResults: result index {} is out of bounds",
index.as_usize()
);
};
let mut slot = slot.lock().unwrap();
if slot.is_some() {
return internal_err!(
"ScalarSubqueryResults: result for index {} was already populated",
index.as_usize()
);
}
*slot = Some(value);
Ok(())
}
pub fn clear(&self) {
for slot in self.slots.iter() {
*slot.lock().unwrap() = None;
}
}
pub fn ptr_eq(this: &Self, other: &Self) -> bool {
Arc::ptr_eq(&this.slots, &other.slots)
}
}
impl fmt::Debug for ScalarSubqueryResults {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_list()
.entries(self.slots.iter().map(|slot| slot.lock().unwrap().clone()))
.finish()
}
}
impl PartialEq for ScalarSubqueryResults {
fn eq(&self, other: &Self) -> bool {
Self::ptr_eq(self, other)
}
}
impl Eq for ScalarSubqueryResults {}
impl Hash for ScalarSubqueryResults {
fn hash<H: Hasher>(&self, state: &mut H) {
Arc::as_ptr(&self.slots).hash(state);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn scalar_subquery_results_set_and_get() -> Result<()> {
let results = ScalarSubqueryResults::new(1);
assert_eq!(results.get(SubqueryIndex::new(0)), None);
results.set(SubqueryIndex::new(0), ScalarValue::Int32(Some(42)))?;
assert_eq!(
results.get(SubqueryIndex::new(0)),
Some(ScalarValue::Int32(Some(42)))
);
assert!(
results
.set(SubqueryIndex::new(0), ScalarValue::Int32(Some(7)))
.is_err()
);
Ok(())
}
#[test]
fn lambda_variables_shadow_outer_scope() {
let outer = TableReference::bare("lambda_1");
let inner = TableReference::bare("lambda_2");
let ctx = PhysicalPlanningContext::default()
.with_qualified_lambda_variables(&outer, &["x".to_string(), "y".to_string()])
.with_qualified_lambda_variables(&inner, &["y".to_string()]);
assert_eq!(ctx.lambda_variable_qualifier("x"), Some(&outer));
assert_eq!(ctx.lambda_variable_qualifier("y"), Some(&inner));
assert_eq!(ctx.lambda_variable_qualifier("z"), None);
}
#[test]
fn scalar_subquery_results_clear() -> Result<()> {
let results = ScalarSubqueryResults::new(1);
results.set(SubqueryIndex::new(0), ScalarValue::Int32(Some(42)))?;
results.clear();
assert_eq!(results.get(SubqueryIndex::new(0)), None);
results.set(SubqueryIndex::new(0), ScalarValue::Int32(Some(7)))?;
assert_eq!(
results.get(SubqueryIndex::new(0)),
Some(ScalarValue::Int32(Some(7)))
);
Ok(())
}
}