use std::collections::HashMap;
use serde::Deserialize;
use serde::Serialize;
fn bool_from_int_or_bool<'de, D>(deserializer: D) -> Result<bool, D::Error>
where
D: serde::Deserializer<'de>,
{
use serde::de;
struct BoolVisitor;
impl<'de> de::Visitor<'de> for BoolVisitor {
type Value = bool;
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("a boolean or an integer (0/1)")
}
fn visit_bool<E: de::Error>(self, v: bool) -> Result<bool, E> {
Ok(v)
}
fn visit_i64<E: de::Error>(self, v: i64) -> Result<bool, E> {
Ok(v != 0)
}
fn visit_u64<E: de::Error>(self, v: u64) -> Result<bool, E> {
Ok(v != 0)
}
}
deserializer.deserialize_any(BoolVisitor)
}
fn string_or_json<'de, D>(deserializer: D) -> Result<String, D::Error>
where
D: serde::Deserializer<'de>,
{
let value = serde_json::Value::deserialize(deserializer)?;
match value {
serde_json::Value::String(s) => Ok(s),
other => Ok(other.to_string()),
}
}
#[derive(Debug, Default, Deserialize, Serialize, PartialEq)]
#[serde(default)]
#[serde(rename_all = "camelCase")]
#[allow(dead_code)] pub(crate) struct QueryPlan {
pub partitioned_query_execution_info_version: usize,
#[serde(default)]
pub query_info: Option<QueryInfo>,
pub query_ranges: Vec<QueryRange>,
pub hybrid_search_query_info: Option<HybridSearchQueryInfo>,
}
#[derive(Debug, Deserialize, Serialize, PartialEq)]
#[serde(rename_all = "camelCase")]
#[allow(dead_code)] pub(crate) struct HybridSearchQueryInfo {
pub global_statistics_query: String,
pub component_query_infos: Vec<QueryInfo>,
#[serde(default)]
pub component_weights: Vec<f64>,
pub skip: Option<u64>,
pub take: Option<u64>,
#[serde(deserialize_with = "bool_from_int_or_bool")]
pub requires_global_statistics: bool,
}
#[derive(Debug, Clone, Deserialize, Default, PartialEq, Eq, Serialize)]
pub(crate) enum DistinctType {
#[default]
None,
Ordered,
Unordered,
}
#[derive(Debug, Deserialize, Default, Serialize, PartialEq)]
#[serde(default)]
#[serde(rename_all = "camelCase")]
pub(crate) struct QueryInfo {
pub distinct_type: DistinctType,
pub top: Option<u64>,
pub offset: Option<u64>,
pub limit: Option<u64>,
pub order_by: Vec<SortOrder>,
pub order_by_expressions: Vec<String>,
pub group_by_expressions: Vec<String>,
pub group_by_aliases: Vec<String>,
pub aggregates: Vec<String>,
#[serde(default)]
pub group_by_alias_to_aggregate_type: HashMap<String, serde_json::Value>,
#[serde(default)]
pub rewritten_query: Option<String>,
#[serde(deserialize_with = "bool_from_int_or_bool")]
pub has_select_value: bool,
#[serde(deserialize_with = "bool_from_int_or_bool")]
pub has_non_streaming_order_by: bool,
}
#[derive(Debug, Deserialize, Clone, Copy, PartialEq, Eq, Serialize)]
pub(crate) enum SortOrder {
Ascending,
Descending,
}
#[derive(Debug, Deserialize, Serialize, PartialEq)]
#[serde(rename_all = "camelCase")]
#[allow(dead_code)] pub(crate) struct QueryRange {
#[serde(deserialize_with = "string_or_json")]
pub min: String,
#[serde(deserialize_with = "string_or_json")]
pub max: String,
#[serde(deserialize_with = "bool_from_int_or_bool")]
pub is_min_inclusive: bool,
#[serde(deserialize_with = "bool_from_int_or_bool")]
pub is_max_inclusive: bool,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn deserializes_minimal_query_plan() {
let json = r#"{
"partitionedQueryExecutionInfoVersion": 1,
"queryRanges": [
{
"min": "",
"max": "FF",
"isMinInclusive": true,
"isMaxInclusive": false
}
]
}"#;
let plan: QueryPlan = serde_json::from_str(json).unwrap();
assert_eq!(plan.partitioned_query_execution_info_version, 1);
assert!(plan.query_info.is_none());
assert!(plan.hybrid_search_query_info.is_none());
assert_eq!(plan.query_ranges.len(), 1);
assert_eq!(plan.query_ranges[0].min, "");
assert_eq!(plan.query_ranges[0].max, "FF");
assert!(plan.query_ranges[0].is_min_inclusive);
assert!(!plan.query_ranges[0].is_max_inclusive);
}
#[test]
fn deserializes_query_plan_with_order_by() {
let json = r#"{
"partitionedQueryExecutionInfoVersion": 2,
"queryInfo": {
"orderBy": ["Ascending", "Descending"],
"orderByExpressions": ["c.name", "c.age"],
"rewrittenQuery": "SELECT c.name, c.age FROM c ORDER BY c.name ASC, c.age DESC"
},
"queryRanges": []
}"#;
let plan: QueryPlan = serde_json::from_str(json).unwrap();
let info = plan.query_info.unwrap();
assert_eq!(
info.order_by,
vec![SortOrder::Ascending, SortOrder::Descending]
);
assert_eq!(info.order_by_expressions, vec!["c.name", "c.age"]);
}
#[test]
fn deserializes_query_plan_with_top_and_aggregates() {
let json = r#"{
"partitionedQueryExecutionInfoVersion": 1,
"queryInfo": {
"top": 10,
"aggregates": ["Count"],
"distinctType": "Ordered"
},
"queryRanges": []
}"#;
let plan: QueryPlan = serde_json::from_str(json).unwrap();
let info = plan.query_info.unwrap();
assert_eq!(info.top, Some(10));
assert_eq!(info.aggregates, vec!["Count"]);
assert_eq!(info.distinct_type, DistinctType::Ordered);
}
#[test]
fn deserializes_query_plan_with_hybrid_search() {
let json = r#"{
"partitionedQueryExecutionInfoVersion": 1,
"queryRanges": [],
"hybridSearchQueryInfo": {
"globalStatisticsQuery": "SELECT COUNT(1) FROM c",
"componentQueryInfos": [],
"componentWeights": [0.5, 0.5],
"skip": null,
"take": 10,
"requiresGlobalStatistics": true
}
}"#;
let plan: QueryPlan = serde_json::from_str(json).unwrap();
let hybrid = plan.hybrid_search_query_info.unwrap();
assert_eq!(hybrid.global_statistics_query, "SELECT COUNT(1) FROM c");
assert_eq!(hybrid.component_weights, vec![0.5, 0.5]);
assert_eq!(hybrid.take, Some(10));
assert!(hybrid.requires_global_statistics);
}
#[test]
fn deserializes_query_plan_with_offset_limit() {
let json = r#"{
"partitionedQueryExecutionInfoVersion": 1,
"queryInfo": {
"offset": 5,
"limit": 20
},
"queryRanges": []
}"#;
let plan: QueryPlan = serde_json::from_str(json).unwrap();
let info = plan.query_info.unwrap();
assert_eq!(info.offset, Some(5));
assert_eq!(info.limit, Some(20));
}
#[test]
fn deserializes_multiple_query_ranges() {
let json = r#"{
"partitionedQueryExecutionInfoVersion": 1,
"queryRanges": [
{ "min": "", "max": "40", "isMinInclusive": true, "isMaxInclusive": false },
{ "min": "80", "max": "FF", "isMinInclusive": true, "isMaxInclusive": false }
]
}"#;
let plan: QueryPlan = serde_json::from_str(json).unwrap();
assert_eq!(plan.query_ranges.len(), 2);
assert_eq!(plan.query_ranges[0].max, "40");
assert_eq!(plan.query_ranges[1].min, "80");
}
}