use azure_core::fmt::SafeDebug;
use serde::{Deserialize, Serialize};
use crate::query::ast::{
SqlBinaryOp, SqlCollectionExpression, SqlLimitSpec, SqlLiteral, SqlOffsetSpec, SqlQuery,
SqlScalarExpression, SqlSelectClause, SqlSelectSpec, SqlSortOrder, SqlTopSpec,
};
use crate::query::common::get_root_alias;
#[derive(SafeDebug, Clone, PartialEq, Serialize)]
#[serde(rename_all = "camelCase")]
pub(crate) struct QueryPlan {
pub(crate) pk_filters: PartitionKeyFilter,
pub(crate) query_info: LocalQueryInfo,
}
#[derive(SafeDebug, Clone, PartialEq, Default, Serialize)]
#[serde(rename_all = "camelCase")]
pub(crate) struct LocalQueryInfo {
#[serde(default)]
pub(crate) distinct_type: DistinctType,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub(crate) top: Option<i64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub(crate) offset: Option<i64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub(crate) limit: Option<i64>,
#[serde(default)]
pub(crate) order_by: Vec<SortOrder>,
#[serde(default)]
pub(crate) order_by_expressions: Vec<String>,
#[serde(default)]
pub(crate) group_by_expressions: Vec<String>,
#[serde(default)]
pub(crate) aggregates: Vec<AggregateKind>,
#[serde(default)]
pub(crate) has_select_value: bool,
#[serde(default)]
pub(crate) has_join: bool,
#[serde(default)]
pub(crate) has_subquery: bool,
#[serde(default)]
pub(crate) has_where: bool,
#[serde(default)]
pub(crate) has_udf: bool,
}
#[derive(SafeDebug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[non_exhaustive]
pub(crate) enum DistinctType {
#[default]
None,
Ordered,
Unordered,
}
#[derive(SafeDebug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub(crate) enum SortOrder {
Ascending,
Descending,
}
#[derive(SafeDebug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub(crate) enum AggregateKind {
Count,
Sum,
Avg,
Min,
Max,
}
#[derive(SafeDebug, Clone, PartialEq, Serialize)]
#[serde(rename_all = "camelCase")]
#[non_exhaustive]
pub(crate) enum PartitionKeyFilter {
Equality(Vec<PartitionKeyValue>),
InList(Vec<Vec<PartitionKeyValue>>),
Unconstrained,
Contradictory,
NotEvaluated,
}
#[derive(SafeDebug, Clone, PartialEq, Serialize)]
#[serde(tag = "type", content = "value", rename_all = "camelCase")]
#[non_exhaustive]
pub(crate) enum PartitionKeyValue {
String(String),
Number(f64),
Bool(bool),
Null,
Undefined,
UnboundParameter(String),
InvalidParameter {
name: String,
reason: String,
},
}
impl PartitionKeyValue {
pub(crate) fn try_number(n: f64) -> Option<Self> {
if n.is_finite() {
Some(PartitionKeyValue::Number(n))
} else {
None
}
}
}
impl Eq for PartitionKeyValue {}
impl std::hash::Hash for PartitionKeyValue {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
std::mem::discriminant(self).hash(state);
match self {
PartitionKeyValue::String(s) => s.hash(state),
PartitionKeyValue::Number(n) => n.to_bits().hash(state),
PartitionKeyValue::Bool(b) => b.hash(state),
PartitionKeyValue::Null | PartitionKeyValue::Undefined => {}
PartitionKeyValue::UnboundParameter(s) => s.hash(state),
PartitionKeyValue::InvalidParameter { name, reason } => {
name.hash(state);
reason.hash(state);
}
}
}
}
pub(crate) fn generate_query_plan(
query: &SqlQuery,
pk_paths: &[&str],
) -> crate::error::Result<QueryPlan> {
generate_query_plan_with_parameters(query, pk_paths, &[])
}
pub(crate) use crate::query::common::Params;
pub(crate) fn generate_query_plan_with_parameters(
query: &SqlQuery,
pk_paths: &[&str],
parameters: &Params,
) -> crate::error::Result<QueryPlan> {
let query_info = analyze_query(query, parameters)?;
let root_alias = get_root_alias(query);
let pk_filters = if pk_paths.is_empty() {
PartitionKeyFilter::NotEvaluated
} else {
let pk_segments: Vec<Vec<&str>> = pk_paths
.iter()
.map(|p| p.strip_prefix('/').unwrap_or(p).split('/').collect())
.collect();
if let Some(where_clause) = &query.where_clause {
extract_pk_from_expression(
&where_clause.expression,
&pk_segments,
root_alias.as_deref(),
parameters,
)
} else {
PartitionKeyFilter::Unconstrained
}
};
Ok(QueryPlan {
pk_filters,
query_info,
})
}
fn resolve_integer_parameter(name: &str, parameters: &Params) -> crate::error::Result<i64> {
crate::query::common::resolve_non_negative_integer_parameter(parameters, name).map_err(|msg| {
crate::error::CosmosError::builder()
.with_status(crate::error::CosmosStatus::CLIENT_QUERY_PLAN_INVALID_TOP_OFFSET_LIMIT)
.with_message(format!("{msg} (TOP/OFFSET/LIMIT clause)"))
.build()
})
}
fn is_constant_expression(expr: &SqlScalarExpression) -> bool {
match expr {
SqlScalarExpression::Literal(_) => true,
SqlScalarExpression::ArrayCreate(items) => items.iter().all(is_constant_expression),
SqlScalarExpression::ObjectCreate(props) => {
props.iter().all(|p| is_constant_expression(&p.expression))
}
SqlScalarExpression::Unary { operand, .. } => is_constant_expression(operand),
SqlScalarExpression::Binary { left, right, .. } => {
is_constant_expression(left) && is_constant_expression(right)
}
_ => false,
}
}
fn analyze_query(query: &SqlQuery, parameters: &Params) -> crate::error::Result<LocalQueryInfo> {
let mut info = LocalQueryInfo {
has_select_value: matches!(query.select.spec, SqlSelectSpec::Value(_)),
has_where: query.where_clause.is_some(),
..Default::default()
};
if query.select.distinct {
let is_constant_select = match &query.select.spec {
SqlSelectSpec::Value(expr) => is_constant_expression(expr),
_ => false,
};
if is_constant_select {
info.distinct_type = DistinctType::None;
} else if query.order_by.is_some() {
info.distinct_type = DistinctType::Ordered;
} else {
info.distinct_type = DistinctType::Unordered;
}
}
info.top = match &query.select.top {
Some(SqlTopSpec::Literal(n)) => Some(*n),
Some(SqlTopSpec::Parameter(name)) => Some(resolve_integer_parameter(name, parameters)?),
None => None,
};
if let Some(ol) = &query.offset_limit {
info.offset = match &ol.offset {
SqlOffsetSpec::Literal(n) => Some(*n),
SqlOffsetSpec::Parameter(name) => Some(resolve_integer_parameter(name, parameters)?),
};
info.limit = match &ol.limit {
SqlLimitSpec::Literal(n) => Some(*n),
SqlLimitSpec::Parameter(name) => Some(resolve_integer_parameter(name, parameters)?),
};
}
if let Some(order_by) = &query.order_by {
for item in &order_by.items {
let sort = match item.order {
SqlSortOrder::Descending => SortOrder::Descending,
_ => SortOrder::Ascending,
};
info.order_by.push(sort);
info.order_by_expressions
.push(expr_to_path_string(&item.expression)?);
}
}
if let Some(group_by) = &query.group_by {
for expr in &group_by.expressions {
info.group_by_expressions.push(expr_to_path_string(expr)?);
}
}
if let Some(from) = &query.from {
info.has_join = has_join(&from.collection);
}
visit_select_for_info(&query.select, &mut info);
if let Some(w) = &query.where_clause {
visit_expr_for_info(&w.expression, &mut info);
}
if let Some(ob) = &query.order_by {
for item in &ob.items {
visit_expr_for_info(&item.expression, &mut info);
}
}
if let Some(gb) = &query.group_by {
for expr in &gb.expressions {
visit_expr_for_info(expr, &mut info);
}
}
Ok(info)
}
fn expr_to_path_string(expr: &SqlScalarExpression) -> crate::error::Result<String> {
let mut parts = Vec::new();
if collect_path_parts(expr, &mut parts) {
Ok(parts.join("."))
} else {
Err(crate::error::CosmosError::builder().with_status(crate::error::CosmosStatus::CLIENT_QUERY_PLAN_COMPLEX_PROJECTION_UNSUPPORTED).with_message(format!(
"{} GROUP BY / ORDER BY expression is not a property path; local plan generation cannot reproduce the Gateway's rewrite. Fall back to the Gateway query-plan endpoint. expression: {expr:?}",
LocalPlanFallbackError::NEEDS_GATEWAY_FALLBACK
)).build())
}
}
pub(crate) struct LocalPlanFallbackError;
impl LocalPlanFallbackError {
pub(crate) const NEEDS_GATEWAY_FALLBACK: &'static str = "[NEEDS_GATEWAY_FALLBACK]";
}
fn is_unresolved_pk_value(v: &PartitionKeyValue) -> bool {
matches!(
v,
PartitionKeyValue::UnboundParameter(_) | PartitionKeyValue::InvalidParameter { .. }
)
}
#[allow(clippy::collapsible_match)] fn collect_path_parts(expr: &SqlScalarExpression, parts: &mut Vec<String>) -> bool {
match expr {
SqlScalarExpression::PropertyRef(name) => {
parts.push(name.clone());
true
}
SqlScalarExpression::MemberRef { source, member } => {
if collect_path_parts(source, parts) {
parts.push(member.clone());
true
} else {
false
}
}
_ => false,
}
}
fn has_join(coll: &SqlCollectionExpression) -> bool {
matches!(coll, SqlCollectionExpression::Join { .. })
}
fn visit_select_for_info(select: &SqlSelectClause, info: &mut LocalQueryInfo) {
match &select.spec {
SqlSelectSpec::List(items) => {
for item in items {
visit_expr_for_info(&item.expression, info);
}
}
SqlSelectSpec::Value(expr) => visit_expr_for_info(expr.as_ref(), info),
SqlSelectSpec::Star => {}
}
}
fn visit_expr_for_info(expr: &SqlScalarExpression, info: &mut LocalQueryInfo) {
walk_expr_for_info(expr, info, false);
}
fn visit_expr_for_info_no_aggregates(expr: &SqlScalarExpression, info: &mut LocalQueryInfo) {
walk_expr_for_info(expr, info, true);
}
fn walk_expr_for_info(
root: &SqlScalarExpression,
info: &mut LocalQueryInfo,
no_aggregates_root: bool,
) {
let mut stack: Vec<(&SqlScalarExpression, bool)> = vec![(root, no_aggregates_root)];
while let Some((expr, no_aggregates)) = stack.pop() {
match expr {
SqlScalarExpression::FunctionCall {
name, args, is_udf, ..
} => {
if *is_udf {
info.has_udf = true;
for arg in args.iter().rev() {
stack.push((arg, true));
}
} else {
if !no_aggregates {
let upper = name.to_ascii_uppercase();
match upper.as_str() {
"COUNT" => info.aggregates.push(AggregateKind::Count),
"SUM" => info.aggregates.push(AggregateKind::Sum),
"AVG" => info.aggregates.push(AggregateKind::Avg),
"MIN" => info.aggregates.push(AggregateKind::Min),
"MAX" => info.aggregates.push(AggregateKind::Max),
_ => {}
}
}
for arg in args.iter().rev() {
stack.push((arg, no_aggregates));
}
}
}
SqlScalarExpression::Exists(_)
| SqlScalarExpression::Subquery(_)
| SqlScalarExpression::Array(_) => {
info.has_subquery = true;
}
SqlScalarExpression::Binary { left, right, .. } => {
stack.push((right, no_aggregates));
stack.push((left, no_aggregates));
}
SqlScalarExpression::Unary { operand, .. } => {
stack.push((operand, no_aggregates));
}
SqlScalarExpression::Conditional {
condition,
if_true,
if_false,
} => {
stack.push((if_false, no_aggregates));
stack.push((if_true, no_aggregates));
stack.push((condition, no_aggregates));
}
SqlScalarExpression::Coalesce { left, right } => {
stack.push((right, no_aggregates));
stack.push((left, no_aggregates));
}
SqlScalarExpression::In {
expression, items, ..
} => {
for item in items.iter().rev() {
stack.push((item, no_aggregates));
}
stack.push((expression, no_aggregates));
}
SqlScalarExpression::Between {
expression,
low,
high,
..
} => {
stack.push((high, no_aggregates));
stack.push((low, no_aggregates));
stack.push((expression, no_aggregates));
}
SqlScalarExpression::Like {
expression,
pattern,
..
} => {
stack.push((pattern, no_aggregates));
stack.push((expression, no_aggregates));
}
SqlScalarExpression::ArrayCreate(items) => {
for item in items.iter().rev() {
stack.push((item, no_aggregates));
}
}
SqlScalarExpression::ObjectCreate(props) => {
for prop in props.iter().rev() {
stack.push((&prop.expression, no_aggregates));
}
}
_ => {}
}
}
}
fn extract_pk_from_expression(
expr: &SqlScalarExpression,
pk_segments: &[Vec<&str>],
root_alias: Option<&str>,
parameters: &Params,
) -> PartitionKeyFilter {
if pk_segments.len() == 1 {
return extract_single_pk(expr, &pk_segments[0], root_alias, parameters);
}
extract_hierarchical_pk(expr, pk_segments, root_alias, parameters)
}
fn extract_single_pk(
expr: &SqlScalarExpression,
pk_path: &[&str],
root_alias: Option<&str>,
parameters: &Params,
) -> PartitionKeyFilter {
match expr {
SqlScalarExpression::Binary {
op: SqlBinaryOp::Equal,
left,
right,
} => {
if is_pk_reference(left, pk_path, root_alias) {
if let Some(val) = extract_literal_value(right, parameters) {
return PartitionKeyFilter::Equality(vec![val]);
}
}
if is_pk_reference(right, pk_path, root_alias) {
if let Some(val) = extract_literal_value(left, parameters) {
return PartitionKeyFilter::Equality(vec![val]);
}
}
PartitionKeyFilter::Unconstrained
}
SqlScalarExpression::In {
expression,
items,
not: false,
} => {
if is_pk_reference(expression, pk_path, root_alias) {
let values: Vec<Vec<PartitionKeyValue>> = items
.iter()
.filter_map(|item| extract_literal_value(item, parameters).map(|v| vec![v]))
.collect();
if values.len() == items.len() {
return PartitionKeyFilter::InList(values);
}
}
PartitionKeyFilter::Unconstrained
}
SqlScalarExpression::Binary {
op: SqlBinaryOp::And,
..
} => {
let mut conjuncts = Vec::new();
flatten_and(expr, &mut conjuncts);
conjuncts
.into_iter()
.map(|c| extract_single_pk(c, pk_path, root_alias, parameters))
.reduce(intersect_pk_filters)
.unwrap_or(PartitionKeyFilter::Unconstrained)
}
SqlScalarExpression::Binary {
op: SqlBinaryOp::Or,
..
} => {
let mut disjuncts = Vec::new();
flatten_or(expr, &mut disjuncts);
disjuncts
.into_iter()
.map(|d| extract_single_pk(d, pk_path, root_alias, parameters))
.reduce(union_pk_filters)
.unwrap_or(PartitionKeyFilter::Unconstrained)
}
_ => PartitionKeyFilter::Unconstrained,
}
}
fn union_pk_filters(a: PartitionKeyFilter, b: PartitionKeyFilter) -> PartitionKeyFilter {
match (a, b) {
(PartitionKeyFilter::Equality(a), PartitionKeyFilter::Equality(b)) => {
normalize_pk_union(vec![a, b])
}
(PartitionKeyFilter::Equality(a), PartitionKeyFilter::InList(mut list))
| (PartitionKeyFilter::InList(mut list), PartitionKeyFilter::Equality(a)) => {
list.push(a);
normalize_pk_union(list)
}
(PartitionKeyFilter::InList(mut a), PartitionKeyFilter::InList(b)) => {
a.extend(b);
normalize_pk_union(a)
}
(PartitionKeyFilter::Contradictory, other) | (other, PartitionKeyFilter::Contradictory) => {
other
}
_ => PartitionKeyFilter::Unconstrained,
}
}
fn normalize_pk_union(values: Vec<Vec<PartitionKeyValue>>) -> PartitionKeyFilter {
let mut seen: std::collections::HashSet<Vec<PartitionKeyValue>> =
std::collections::HashSet::with_capacity(values.len());
let mut deduped: Vec<Vec<PartitionKeyValue>> = Vec::with_capacity(values.len());
for value in values {
if !seen.contains(&value) {
seen.insert(value.clone());
deduped.push(value);
}
}
match deduped.len() {
0 => PartitionKeyFilter::Unconstrained,
1 => PartitionKeyFilter::Equality(deduped.into_iter().next().unwrap()),
_ => PartitionKeyFilter::InList(deduped),
}
}
fn intersect_pk_filters(a: PartitionKeyFilter, b: PartitionKeyFilter) -> PartitionKeyFilter {
match (a, b) {
(PartitionKeyFilter::Unconstrained, other) | (other, PartitionKeyFilter::Unconstrained) => {
other
}
(PartitionKeyFilter::Contradictory, _) | (_, PartitionKeyFilter::Contradictory) => {
PartitionKeyFilter::Contradictory
}
(PartitionKeyFilter::Equality(a), PartitionKeyFilter::Equality(b)) => {
let a_unresolved = a.iter().any(is_unresolved_pk_value);
let b_unresolved = b.iter().any(is_unresolved_pk_value);
match (a_unresolved, b_unresolved) {
(true, true) => PartitionKeyFilter::Unconstrained,
(true, false) => PartitionKeyFilter::Equality(b),
(false, true) => PartitionKeyFilter::Equality(a),
(false, false) => {
if a == b {
PartitionKeyFilter::Equality(a)
} else {
PartitionKeyFilter::Contradictory
}
}
}
}
(PartitionKeyFilter::Equality(eq), PartitionKeyFilter::InList(list))
| (PartitionKeyFilter::InList(list), PartitionKeyFilter::Equality(eq)) => {
if eq.iter().any(is_unresolved_pk_value) {
normalize_pk_union(list)
} else if list.contains(&eq) {
PartitionKeyFilter::Equality(eq)
} else {
PartitionKeyFilter::Contradictory
}
}
(PartitionKeyFilter::InList(a), PartitionKeyFilter::InList(b)) => {
let intersection: Vec<Vec<PartitionKeyValue>> =
a.into_iter().filter(|item| b.contains(item)).collect();
match intersection.len() {
0 => PartitionKeyFilter::Contradictory,
1 => PartitionKeyFilter::Equality(intersection.into_iter().next().unwrap()),
_ => PartitionKeyFilter::InList(intersection),
}
}
(PartitionKeyFilter::NotEvaluated, other) | (other, PartitionKeyFilter::NotEvaluated) => {
other
}
}
}
fn extract_hierarchical_pk(
expr: &SqlScalarExpression,
pk_segments: &[Vec<&str>],
root_alias: Option<&str>,
parameters: &Params,
) -> PartitionKeyFilter {
if let SqlScalarExpression::Binary {
op: SqlBinaryOp::Or,
..
} = expr
{
let mut disjuncts = Vec::new();
flatten_or(expr, &mut disjuncts);
return disjuncts
.into_iter()
.map(|d| extract_hierarchical_pk(d, pk_segments, root_alias, parameters))
.reduce(union_pk_filters)
.unwrap_or(PartitionKeyFilter::Unconstrained);
}
let mut conjuncts = Vec::new();
flatten_and(expr, &mut conjuncts);
const MAX_HPK_TUPLES: usize = 1024;
let mut per_component: Vec<Vec<PartitionKeyValue>> = Vec::with_capacity(pk_segments.len());
for pk_path in pk_segments {
let mut equal_value: Option<PartitionKeyValue> = None;
let mut in_values: Option<Vec<PartitionKeyValue>> = None;
for conjunct in &conjuncts {
match conjunct {
SqlScalarExpression::Binary {
op: SqlBinaryOp::Equal,
left,
right,
} => {
let val = if is_pk_reference(left, pk_path, root_alias) {
extract_literal_value(right, parameters)
} else if is_pk_reference(right, pk_path, root_alias) {
extract_literal_value(left, parameters)
} else {
None
};
if let Some(v) = val {
match &equal_value {
None => equal_value = Some(v),
Some(existing) if *existing == v => {} Some(_) => return PartitionKeyFilter::Contradictory,
}
}
}
SqlScalarExpression::In {
expression,
items,
not: false,
} if is_pk_reference(expression, pk_path, root_alias) => {
let mut vs: Vec<PartitionKeyValue> = Vec::with_capacity(items.len());
let mut all_literal = true;
for item in items {
match extract_literal_value(item, parameters) {
Some(v) => vs.push(v),
None => {
all_literal = false;
break;
}
}
}
if !all_literal {
continue;
}
in_values = Some(match in_values {
None => vs,
Some(existing) => existing.into_iter().filter(|v| vs.contains(v)).collect(),
});
if matches!(in_values.as_ref(), Some(v) if v.is_empty()) {
return PartitionKeyFilter::Contradictory;
}
}
_ => {}
}
}
let component_values: Vec<PartitionKeyValue> = match (equal_value, in_values) {
(Some(eq), Some(list)) => {
if list.contains(&eq) {
vec![eq]
} else {
return PartitionKeyFilter::Contradictory;
}
}
(Some(eq), None) => vec![eq],
(None, Some(list)) => list,
(None, None) => return PartitionKeyFilter::Unconstrained,
};
per_component.push(component_values);
}
let total: usize = per_component.iter().map(|v| v.len()).product();
if total == 0 {
return PartitionKeyFilter::Contradictory;
}
if total > MAX_HPK_TUPLES {
return PartitionKeyFilter::Unconstrained;
}
let mut tuples: Vec<Vec<PartitionKeyValue>> = vec![Vec::with_capacity(per_component.len())];
for component in &per_component {
let mut next: Vec<Vec<PartitionKeyValue>> =
Vec::with_capacity(tuples.len() * component.len());
for prefix in &tuples {
for v in component {
let mut t = prefix.clone();
t.push(v.clone());
next.push(t);
}
}
tuples = next;
}
if tuples.len() == 1 {
PartitionKeyFilter::Equality(tuples.into_iter().next().unwrap())
} else {
PartitionKeyFilter::InList(tuples)
}
}
fn flatten_and<'a>(expr: &'a SqlScalarExpression, out: &mut Vec<&'a SqlScalarExpression>) {
flatten_chain(expr, SqlBinaryOp::And, out);
}
fn flatten_or<'a>(expr: &'a SqlScalarExpression, out: &mut Vec<&'a SqlScalarExpression>) {
flatten_chain(expr, SqlBinaryOp::Or, out);
}
fn flatten_chain<'a>(
root: &'a SqlScalarExpression,
op: SqlBinaryOp,
out: &mut Vec<&'a SqlScalarExpression>,
) {
let mut stack: Vec<&'a SqlScalarExpression> = vec![root];
while let Some(node) = stack.pop() {
match node {
SqlScalarExpression::Binary {
op: node_op,
left,
right,
} if *node_op == op => {
stack.push(right);
stack.push(left);
}
other => out.push(other),
}
}
}
fn is_pk_reference(expr: &SqlScalarExpression, pk_path: &[&str], root_alias: Option<&str>) -> bool {
let mut resolved_path = Vec::new();
if !resolve_property_path(expr, &mut resolved_path) {
return false;
}
if let Some(alias) = root_alias {
if resolved_path.first().map(String::as_str) == Some(alias) {
return resolved_path[1..]
.iter()
.map(String::as_str)
.collect::<Vec<_>>()
== pk_path;
}
}
resolved_path.iter().map(String::as_str).collect::<Vec<_>>() == pk_path
}
#[allow(clippy::collapsible_match)] fn resolve_property_path(expr: &SqlScalarExpression, path: &mut Vec<String>) -> bool {
match expr {
SqlScalarExpression::PropertyRef(name) => {
path.push(name.clone());
true
}
SqlScalarExpression::MemberRef { source, member } => {
if resolve_property_path(source, path) {
path.push(member.clone());
true
} else {
false
}
}
_ => false,
}
}
fn extract_literal_value(
expr: &SqlScalarExpression,
parameters: &Params,
) -> Option<PartitionKeyValue> {
match expr {
SqlScalarExpression::Literal(lit) => match lit {
SqlLiteral::String(s) => Some(PartitionKeyValue::String(s.clone())),
SqlLiteral::Number(n) => PartitionKeyValue::try_number(*n),
SqlLiteral::Integer(n) => PartitionKeyValue::try_number(*n as f64),
SqlLiteral::Boolean(b) => Some(PartitionKeyValue::Bool(*b)),
SqlLiteral::Null => Some(PartitionKeyValue::Null),
SqlLiteral::Undefined => Some(PartitionKeyValue::Undefined),
},
SqlScalarExpression::ParameterRef(name) => {
Some(resolve_pk_parameter(name, parameters))
}
_ => None,
}
}
fn resolve_pk_parameter(name: &str, parameters: &Params) -> PartitionKeyValue {
let needle = name.trim_start_matches('@');
let entry = parameters
.iter()
.find(|(n, _)| n.trim_start_matches('@') == needle);
let value = match entry {
Some((_, v)) => v,
None => return PartitionKeyValue::UnboundParameter(needle.to_string()),
};
match value {
serde_json::Value::String(s) => PartitionKeyValue::String(s.clone()),
serde_json::Value::Number(n) => {
n.as_f64()
.and_then(PartitionKeyValue::try_number)
.unwrap_or_else(|| PartitionKeyValue::InvalidParameter {
name: needle.to_string(),
reason: format!("number value `{n}` is not a finite f64"),
})
}
serde_json::Value::Bool(b) => PartitionKeyValue::Bool(*b),
serde_json::Value::Null => PartitionKeyValue::Null,
serde_json::Value::Array(_) => PartitionKeyValue::InvalidParameter {
name: needle.to_string(),
reason: "array values cannot be used as a partition key".to_string(),
},
serde_json::Value::Object(_) => PartitionKeyValue::InvalidParameter {
name: needle.to_string(),
reason: "object values cannot be used as a partition key".to_string(),
},
}
}
#[cfg(any(test, feature = "__internal_testing"))]
#[doc(hidden)]
pub fn __test_only_generate_query_plan_for_pk_paths(
sql: &str,
pk_paths: &[&str],
parameters: &[(String, serde_json::Value)],
) -> crate::error::Result<serde_json::Value> {
let program = crate::query::parse(sql).map_err(|e| {
crate::error::CosmosError::builder()
.with_status(crate::error::CosmosStatus::SERIALIZATION_RESPONSE_BODY_INVALID)
.with_message("failed to parse query")
.with_source(e)
.build()
})?;
let raw_plan = generate_query_plan_with_parameters(&program.query, pk_paths, parameters)?;
serde_json::to_value(&raw_plan).map_err(|e| {
crate::error::CosmosError::builder()
.with_status(crate::error::CosmosStatus::SERIALIZATION_RESPONSE_BODY_INVALID)
.with_message("failed to serialize query plan")
.with_source(e)
.build()
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::query::parse;
fn plan(sql: &str) -> QueryPlan {
let p = parse(sql).unwrap();
generate_query_plan(&p.query, &["/pk"]).unwrap()
}
#[path = "query_plan_comparison.rs"]
mod query_plan_comparison;
#[test]
fn pk_equality() {
assert_eq!(
plan("SELECT * FROM c WHERE c.pk = 'hello'").pk_filters,
PartitionKeyFilter::Equality(vec![PartitionKeyValue::String("hello".into())])
);
}
#[test]
fn pk_with_and() {
assert_eq!(
plan("SELECT * FROM c WHERE c.pk = 'x' AND c.age > 21").pk_filters,
PartitionKeyFilter::Equality(vec![PartitionKeyValue::String("x".into())])
);
}
#[test]
fn pk_in_list() {
match plan("SELECT * FROM c WHERE c.pk IN ('a', 'b')").pk_filters {
PartitionKeyFilter::InList(list) => assert_eq!(list.len(), 2),
other => panic!("expected InList, got {other:?}"),
}
}
#[test]
fn pk_or_with_contradictory_disjunct_preserves_other_side() {
let qp = plan("SELECT * FROM c WHERE (c.pk = 'a' AND c.pk = 'b') OR c.pk = 'c'");
assert_eq!(
qp.pk_filters,
PartitionKeyFilter::Equality(vec![PartitionKeyValue::String("c".into())])
);
}
#[test]
fn pk_and_with_unbound_parameter_keeps_literal_side() {
let p = parse("SELECT * FROM c WHERE c.pk = 'a' AND c.pk = @unbound").unwrap();
let qp = generate_query_plan_with_parameters(&p.query, &["/pk"], &[]).unwrap();
assert_eq!(
qp.pk_filters,
PartitionKeyFilter::Equality(vec![PartitionKeyValue::String("a".into())])
);
}
#[test]
fn try_number_rejects_non_finite() {
assert!(PartitionKeyValue::try_number(f64::NAN).is_none());
assert!(PartitionKeyValue::try_number(f64::INFINITY).is_none());
assert!(PartitionKeyValue::try_number(f64::NEG_INFINITY).is_none());
assert!(PartitionKeyValue::try_number(0.0).is_some());
assert!(PartitionKeyValue::try_number(1.5).is_some());
}
#[test]
fn aggregate_inside_udf_arg_not_advertised() {
let p = parse("SELECT udf.foo(COUNT(c.x)) FROM c").unwrap();
let qp = generate_query_plan(&p.query, &["/pk"]).unwrap();
assert!(qp.query_info.has_udf);
assert!(
qp.query_info.aggregates.is_empty(),
"aggregates inside UDF args must be skipped; got {:?}",
qp.query_info.aggregates
);
}
#[test]
fn numeric_pk_canonicalization_int_and_float_match() {
let int_form = generate_query_plan(
&parse("SELECT * FROM c WHERE c.pk = 1").unwrap().query,
&["/pk"],
)
.unwrap();
let float_form = generate_query_plan(
&parse("SELECT * FROM c WHERE c.pk = 1.0").unwrap().query,
&["/pk"],
)
.unwrap();
assert_eq!(int_form.pk_filters, float_form.pk_filters);
assert_eq!(int_form.query_info, float_form.query_info);
}
#[test]
fn bracket_paths_fall_back_to_gateway() {
for sql in [
"SELECT * FROM c ORDER BY c['name'] ASC",
"SELECT * FROM c ORDER BY c[\"name\"] ASC",
"SELECT * FROM c ORDER BY c.scores[0] ASC",
] {
let p = parse(sql).unwrap();
let err = generate_query_plan(&p.query, &["/pk"])
.expect_err(&format!("bracket path must surface fallback: {sql}"));
assert!(
format!("{err}").contains(LocalPlanFallbackError::NEEDS_GATEWAY_FALLBACK),
"fallback sentinel missing for {sql}: {err}"
);
}
}
#[test]
fn non_path_group_by_errors() {
let p = parse("SELECT c.x & 1 AS parity, COUNT(1) FROM c GROUP BY c.x & 1").unwrap();
let err = generate_query_plan(&p.query, &["/pk"]).expect_err(
"non-path GROUP BY must surface an error so callers can fall back to Gateway",
);
let msg = format!("{err}");
assert!(
msg.contains("GROUP BY / ORDER BY"),
"unexpected error message: {msg}"
);
assert!(
msg.contains(LocalPlanFallbackError::NEEDS_GATEWAY_FALLBACK),
"error must carry the fallback sentinel; got: {msg}"
);
}
#[test]
fn no_pk_filter() {
assert_eq!(
plan("SELECT * FROM c WHERE c.age > 21").pk_filters,
PartitionKeyFilter::Unconstrained
);
}
#[test]
fn no_where_clause() {
assert_eq!(
plan("SELECT * FROM c").pk_filters,
PartitionKeyFilter::Unconstrained
);
}
#[test]
fn distinct_unordered() {
let qp = plan("SELECT DISTINCT c.name FROM c");
assert_eq!(qp.query_info.distinct_type, DistinctType::Unordered);
}
#[test]
fn distinct_ordered() {
let qp = plan("SELECT DISTINCT c.name FROM c ORDER BY c.name");
assert_eq!(qp.query_info.distinct_type, DistinctType::Ordered);
}
#[test]
fn no_distinct() {
let qp = plan("SELECT c.name FROM c");
assert_eq!(qp.query_info.distinct_type, DistinctType::None);
}
#[test]
fn top_value() {
assert_eq!(plan("SELECT TOP 10 * FROM c").query_info.top, Some(10));
}
#[test]
fn offset_limit() {
let qp = plan("SELECT * FROM c OFFSET 5 LIMIT 20");
assert_eq!(qp.query_info.offset, Some(5));
assert_eq!(qp.query_info.limit, Some(20));
}
#[test]
fn order_by_single_asc() {
let qp = plan("SELECT * FROM c ORDER BY c.name ASC");
assert_eq!(qp.query_info.order_by, vec![SortOrder::Ascending]);
assert_eq!(qp.query_info.order_by_expressions, vec!["c.name"]);
}
#[test]
fn order_by_single_desc() {
let qp = plan("SELECT * FROM c ORDER BY c.name DESC");
assert_eq!(qp.query_info.order_by, vec![SortOrder::Descending]);
}
#[test]
fn order_by_multiple() {
let qp = plan("SELECT * FROM c ORDER BY c.name ASC, c.age DESC");
assert_eq!(
qp.query_info.order_by,
vec![SortOrder::Ascending, SortOrder::Descending]
);
assert_eq!(qp.query_info.order_by_expressions, vec!["c.name", "c.age"]);
}
#[test]
fn group_by_single() {
let qp = plan("SELECT c.city, COUNT(1) FROM c GROUP BY c.city");
assert_eq!(qp.query_info.group_by_expressions, vec!["c.city"]);
assert!(qp.query_info.aggregates.contains(&AggregateKind::Count));
}
#[test]
fn group_by_multiple() {
let qp = plan("SELECT c.city, c.state, COUNT(1) FROM c GROUP BY c.city, c.state");
assert_eq!(
qp.query_info.group_by_expressions,
vec!["c.city", "c.state"]
);
}
#[test]
fn aggregate_count() {
let qp = plan("SELECT COUNT(1) FROM c");
assert_eq!(qp.query_info.aggregates, vec![AggregateKind::Count]);
}
#[test]
fn aggregate_sum() {
let qp = plan("SELECT SUM(c.price) FROM c");
assert_eq!(qp.query_info.aggregates, vec![AggregateKind::Sum]);
}
#[test]
fn aggregate_avg() {
let qp = plan("SELECT AVG(c.score) FROM c");
assert_eq!(qp.query_info.aggregates, vec![AggregateKind::Avg]);
}
#[test]
fn aggregate_min_max() {
let qp = plan("SELECT MIN(c.age), MAX(c.age) FROM c");
assert!(qp.query_info.aggregates.contains(&AggregateKind::Min));
assert!(qp.query_info.aggregates.contains(&AggregateKind::Max));
}
#[test]
fn multiple_aggregates() {
let qp = plan("SELECT COUNT(1), SUM(c.price), AVG(c.score) FROM c");
assert_eq!(qp.query_info.aggregates.len(), 3);
assert!(qp.query_info.aggregates.contains(&AggregateKind::Count));
assert!(qp.query_info.aggregates.contains(&AggregateKind::Sum));
assert!(qp.query_info.aggregates.contains(&AggregateKind::Avg));
}
#[test]
fn no_aggregates() {
let qp = plan("SELECT * FROM c");
assert!(qp.query_info.aggregates.is_empty());
}
#[test]
fn select_value_detected() {
assert!(
plan("SELECT VALUE c.name FROM c")
.query_info
.has_select_value
);
}
#[test]
fn select_star_not_value() {
assert!(!plan("SELECT * FROM c").query_info.has_select_value);
}
#[test]
fn join_detected() {
assert!(plan("SELECT * FROM c JOIN t IN c.tags").query_info.has_join);
}
#[test]
fn no_join() {
assert!(!plan("SELECT * FROM c").query_info.has_join);
}
#[test]
fn exists_subquery_detected() {
assert!(
plan("SELECT * FROM c WHERE EXISTS(SELECT VALUE t FROM t IN c.tags)")
.query_info
.has_subquery
);
}
#[test]
fn array_subquery_detected() {
assert!(
plan("SELECT ARRAY(SELECT t FROM t IN c.tags) FROM c")
.query_info
.has_subquery
);
}
#[test]
fn udf_detected() {
assert!(
plan("SELECT * FROM c WHERE udf.myFunc(c.x) > 0")
.query_info
.has_udf
);
}
#[test]
fn builtin_function_not_udf() {
assert!(
!plan("SELECT * FROM c WHERE CONTAINS(c.name, 'x')")
.query_info
.has_udf
);
}
#[test]
fn has_where() {
assert!(plan("SELECT * FROM c WHERE c.x = 1").query_info.has_where);
}
#[test]
fn no_where() {
assert!(!plan("SELECT * FROM c").query_info.has_where);
}
#[test]
fn aggregate_with_pk_and_group_by() {
let qp = plan(
"SELECT c.city, COUNT(1) AS cnt, SUM(c.revenue) AS total \
FROM c WHERE c.pk = 'x' GROUP BY c.city",
);
assert_eq!(
qp.pk_filters,
PartitionKeyFilter::Equality(vec![PartitionKeyValue::String("x".into())])
);
assert_eq!(qp.query_info.group_by_expressions, vec!["c.city"]);
assert!(qp.query_info.aggregates.contains(&AggregateKind::Count));
assert!(qp.query_info.aggregates.contains(&AggregateKind::Sum));
}
#[test]
fn order_by_with_pk_and_top() {
let qp = plan("SELECT TOP 5 * FROM c WHERE c.pk = 'x' ORDER BY c.name DESC");
assert_eq!(
qp.pk_filters,
PartitionKeyFilter::Equality(vec![PartitionKeyValue::String("x".into())])
);
assert_eq!(qp.query_info.top, Some(5));
assert_eq!(qp.query_info.order_by, vec![SortOrder::Descending]);
}
#[test]
fn cross_partition_aggregate_with_order_by() {
let qp = plan("SELECT c.city, COUNT(1) FROM c GROUP BY c.city ORDER BY c.city ASC");
assert_eq!(qp.pk_filters, PartitionKeyFilter::Unconstrained);
assert!(!qp.query_info.group_by_expressions.is_empty());
assert!(!qp.query_info.order_by.is_empty());
assert!(!qp.query_info.aggregates.is_empty());
}
#[test]
fn and_contradictory_equality_is_contradictory() {
assert_eq!(
plan("SELECT * FROM c WHERE c.pk = 'a' AND c.pk = 'b'").pk_filters,
PartitionKeyFilter::Contradictory
);
}
#[test]
fn and_redundant_equality_is_ok() {
assert_eq!(
plan("SELECT * FROM c WHERE c.pk = 'a' AND c.pk = 'a'").pk_filters,
PartitionKeyFilter::Equality(vec![PartitionKeyValue::String("a".into())])
);
}
#[test]
fn and_equality_narrows_in_list() {
assert_eq!(
plan("SELECT * FROM c WHERE c.pk = 'a' AND c.pk IN ('a', 'b')").pk_filters,
PartitionKeyFilter::Equality(vec![PartitionKeyValue::String("a".into())])
);
}
#[test]
fn and_equality_not_in_list_is_contradictory() {
assert_eq!(
plan("SELECT * FROM c WHERE c.pk = 'c' AND c.pk IN ('a', 'b')").pk_filters,
PartitionKeyFilter::Contradictory
);
}
#[test]
fn and_in_list_narrows_in_list() {
let qp = plan("SELECT * FROM c WHERE c.pk IN ('a', 'b', 'c') AND c.pk IN ('b', 'c', 'd')");
match qp.pk_filters {
PartitionKeyFilter::InList(ref list) => {
assert_eq!(list.len(), 2);
assert!(list.contains(&vec![PartitionKeyValue::String("b".into())]));
assert!(list.contains(&vec![PartitionKeyValue::String("c".into())]));
}
_ => panic!("expected InList, got {:?}", qp.pk_filters),
}
}
#[test]
fn and_in_list_intersection_single_becomes_equality() {
assert_eq!(
plan("SELECT * FROM c WHERE c.pk IN ('a', 'b') AND c.pk IN ('b', 'c')").pk_filters,
PartitionKeyFilter::Equality(vec![PartitionKeyValue::String("b".into())])
);
}
#[test]
fn and_in_list_empty_intersection_is_contradictory() {
assert_eq!(
plan("SELECT * FROM c WHERE c.pk IN ('a', 'b') AND c.pk IN ('c', 'd')").pk_filters,
PartitionKeyFilter::Contradictory
);
}
#[test]
fn and_pk_with_non_pk_keeps_pk() {
assert_eq!(
plan("SELECT * FROM c WHERE c.pk = 'a' AND c.other > 5").pk_filters,
PartitionKeyFilter::Equality(vec![PartitionKeyValue::String("a".into())])
);
}
#[test]
fn and_non_pk_with_pk_keeps_pk() {
assert_eq!(
plan("SELECT * FROM c WHERE c.other > 5 AND c.pk = 'a'").pk_filters,
PartitionKeyFilter::Equality(vec![PartitionKeyValue::String("a".into())])
);
}
#[test]
fn and_chain_multiple_consistent() {
assert_eq!(
plan("SELECT * FROM c WHERE c.pk = 'a' AND c.x > 1 AND c.pk = 'a' AND c.y < 10")
.pk_filters,
PartitionKeyFilter::Equality(vec![PartitionKeyValue::String("a".into())])
);
}
#[test]
fn and_chain_contradictory() {
assert_eq!(
plan("SELECT * FROM c WHERE c.pk = 'a' AND c.x > 1 AND c.pk = 'b'").pk_filters,
PartitionKeyFilter::Contradictory
);
}
#[test]
fn and_in_list_with_non_pk() {
match plan("SELECT * FROM c WHERE c.pk IN ('a', 'b') AND c.other > 5").pk_filters {
PartitionKeyFilter::InList(list) => assert_eq!(list.len(), 2),
other => panic!("expected InList, got {other:?}"),
}
}
fn plan_hpk(sql: &str) -> QueryPlan {
let p = parse(sql).unwrap();
generate_query_plan(&p.query, &["/tenant", "/userId"]).unwrap()
}
#[test]
fn hpk_contradictory_first_component() {
assert_eq!(
plan_hpk("SELECT * FROM c WHERE c.tenant = 'a' AND c.tenant = 'b' AND c.userId = 'u1'")
.pk_filters,
PartitionKeyFilter::Contradictory
);
}
#[test]
fn hpk_contradictory_second_component() {
assert_eq!(
plan_hpk(
"SELECT * FROM c WHERE c.tenant = 'a' AND c.userId = 'u1' AND c.userId = 'u2'"
)
.pk_filters,
PartitionKeyFilter::Contradictory
);
}
#[test]
fn hpk_redundant_constraints_ok() {
assert_eq!(
plan_hpk("SELECT * FROM c WHERE c.tenant = 'a' AND c.userId = 'u1' AND c.tenant = 'a'")
.pk_filters,
PartitionKeyFilter::Equality(vec![
PartitionKeyValue::String("a".into()),
PartitionKeyValue::String("u1".into()),
])
);
}
#[test]
fn contradictory_pk_equality_is_distinct_from_unconstrained() {
assert_eq!(
plan("SELECT * FROM c WHERE c.pk = 'a' AND c.pk = 'b'").pk_filters,
PartitionKeyFilter::Contradictory
);
assert_eq!(
plan("SELECT * FROM c").pk_filters,
PartitionKeyFilter::Unconstrained
);
assert_eq!(
plan("SELECT * FROM c WHERE c.age > 18").pk_filters,
PartitionKeyFilter::Unconstrained
);
}
#[test]
fn contradictory_is_absorbing_under_and() {
assert_eq!(
plan("SELECT * FROM c WHERE c.pk = 'a' AND c.pk = 'b' AND c.age > 18").pk_filters,
PartitionKeyFilter::Contradictory
);
}
#[test]
fn unbound_pk_parameter_is_distinct_variant() {
let p = parse("SELECT * FROM c WHERE c.pk = @missing").unwrap();
let qp = generate_query_plan_with_parameters(&p.query, &["/pk"], &[]).unwrap();
match qp.pk_filters {
PartitionKeyFilter::Equality(values) => {
assert_eq!(values.len(), 1);
match &values[0] {
PartitionKeyValue::UnboundParameter(name) => assert_eq!(name, "missing"),
other => panic!("expected UnboundParameter, got {other:?}"),
}
}
other => panic!("expected Equality(UnboundParameter), got {other:?}"),
}
}
#[test]
fn invalid_pk_parameter_carries_reason() {
let p = parse("SELECT * FROM c WHERE c.pk = @bad").unwrap();
let params = vec![("bad".to_string(), serde_json::json!([1, 2, 3]))];
let qp = generate_query_plan_with_parameters(&p.query, &["/pk"], ¶ms).unwrap();
match qp.pk_filters {
PartitionKeyFilter::Equality(values) => match &values[0] {
PartitionKeyValue::InvalidParameter { name, reason } => {
assert_eq!(name, "bad");
assert!(reason.contains("array"), "reason was: {reason}");
}
other => panic!("expected InvalidParameter, got {other:?}"),
},
other => panic!("expected Equality, got {other:?}"),
}
}
}