use crate::partiql::ast::*;
use crate::partiql::validator::{DynamoDBValidator, QueryType};
use crate::{Error, Key, Result};
use bytes::Bytes;
pub struct PartiQLTranslator;
impl PartiQLTranslator {
pub fn translate_select(stmt: &SelectStatement) -> Result<SelectTranslation> {
let query_type = DynamoDBValidator::validate_select(stmt)?;
match query_type {
QueryType::Query { pk_condition, sk_condition } => {
let (pk_bytes, multiple_pks) = Self::extract_pk_bytes(&pk_condition)?;
if multiple_pks {
Ok(SelectTranslation::MultiGet {
keys: pk_bytes,
index_name: stmt.index_name.clone(),
})
} else {
let pk = pk_bytes.into_iter().next().unwrap();
let sk_condition_translated = sk_condition
.as_ref()
.map(Self::translate_sk_condition)
.transpose()?;
Ok(SelectTranslation::Query {
pk,
sk_condition: sk_condition_translated,
index_name: stmt.index_name.clone(),
forward: stmt.order_by.as_ref().map_or(true, |o| o.ascending),
})
}
}
QueryType::Scan => {
Ok(SelectTranslation::Scan {
filter_conditions: stmt
.where_clause
.as_ref()
.map(|wc| wc.conditions.clone())
.unwrap_or_default(),
})
}
}
}
fn extract_pk_bytes(condition: &Condition) -> Result<(Vec<Bytes>, bool)> {
match &condition.operator {
CompareOp::Equal => {
let bytes = Self::value_to_bytes(&condition.value)?;
Ok((vec![bytes], false))
}
CompareOp::In => {
match &condition.value {
SqlValue::List(values) => {
let bytes_vec: Result<Vec<Bytes>> = values
.iter()
.map(Self::value_to_bytes)
.collect();
Ok((bytes_vec?, true))
}
_ => Err(Error::InvalidQuery("IN value must be a list".into())),
}
}
_ => Err(Error::InvalidQuery(
"Partition key must use = or IN operator".into(),
)),
}
}
fn value_to_bytes(value: &SqlValue) -> Result<Bytes> {
match value {
SqlValue::String(s) => Ok(Bytes::copy_from_slice(s.as_bytes())),
SqlValue::Number(n) => Ok(Bytes::copy_from_slice(n.as_bytes())),
_ => Err(Error::InvalidQuery(format!(
"Unsupported key value type: {:?}",
value
))),
}
}
fn translate_sk_condition(condition: &Condition) -> Result<SortKeyConditionType> {
let sk_bytes = Self::value_to_bytes(&condition.value)?;
match condition.operator {
CompareOp::Equal => Ok(SortKeyConditionType::Equal(sk_bytes)),
CompareOp::LessThan => Ok(SortKeyConditionType::LessThan(sk_bytes)),
CompareOp::LessThanOrEqual => Ok(SortKeyConditionType::LessThanOrEqual(sk_bytes)),
CompareOp::GreaterThan => Ok(SortKeyConditionType::GreaterThan(sk_bytes)),
CompareOp::GreaterThanOrEqual => Ok(SortKeyConditionType::GreaterThanOrEqual(sk_bytes)),
CompareOp::Between => {
match &condition.value {
SqlValue::List(values) if values.len() == 2 => {
let low = Self::value_to_bytes(&values[0])?;
let high = Self::value_to_bytes(&values[1])?;
Ok(SortKeyConditionType::Between(low, high))
}
_ => Err(Error::InvalidQuery("BETWEEN requires exactly 2 values".into())),
}
}
_ => Err(Error::InvalidQuery(format!(
"Unsupported sort key operator: {:?}",
condition.operator
))),
}
}
pub fn translate_insert(stmt: &InsertStatement) -> Result<InsertTranslation> {
DynamoDBValidator::validate_insert(stmt)?;
let value_map = match &stmt.value {
SqlValue::Map(map) => map,
_ => return Err(Error::InvalidQuery("INSERT value must be a map".into())),
};
let pk_value = value_map
.get("pk")
.ok_or_else(|| Error::InvalidQuery("INSERT value must contain 'pk'".into()))?;
let pk_bytes = Self::value_to_bytes(pk_value)?;
let sk_bytes = value_map
.get("sk")
.map(Self::value_to_bytes)
.transpose()?;
let key = if let Some(sk) = sk_bytes {
Key::with_sk(pk_bytes.to_vec(), sk.to_vec())
} else {
Key::new(pk_bytes.to_vec())
};
let mut item = std::collections::HashMap::new();
for (attr_name, attr_value) in value_map {
if attr_name != "pk" && attr_name != "sk" {
item.insert(attr_name.clone(), attr_value.to_kstone_value());
}
}
Ok(InsertTranslation { key, item })
}
pub fn translate_update(stmt: &UpdateStatement) -> Result<UpdateTranslation> {
DynamoDBValidator::validate_update(stmt)?;
let pk_cond = stmt
.where_clause
.get_condition("pk")
.ok_or_else(|| Error::InvalidQuery("UPDATE must specify pk in WHERE clause".into()))?;
let pk_bytes = Self::value_to_bytes(&pk_cond.value)?;
let sk_bytes = stmt
.where_clause
.get_condition("sk")
.map(|c| Self::value_to_bytes(&c.value))
.transpose()?;
let key = if let Some(sk) = sk_bytes {
Key::with_sk(pk_bytes.to_vec(), sk.to_vec())
} else {
Key::new(pk_bytes.to_vec())
};
let mut expression_parts = Vec::new();
let mut values = std::collections::HashMap::new();
let mut value_counter = 1;
if !stmt.set_assignments.is_empty() {
let mut set_exprs = Vec::new();
for assignment in &stmt.set_assignments {
match &assignment.value {
SetValue::Literal(sql_value) => {
let placeholder = format!(":v{}", value_counter);
set_exprs.push(format!("{} = {}", assignment.attribute, placeholder));
values.insert(placeholder, sql_value.to_kstone_value());
value_counter += 1;
}
SetValue::Add { attribute, value } => {
let placeholder = format!(":v{}", value_counter);
set_exprs.push(format!(
"{} = {} + {}",
assignment.attribute, attribute, placeholder
));
values.insert(placeholder, value.to_kstone_value());
value_counter += 1;
}
SetValue::Subtract { attribute, value } => {
let placeholder = format!(":v{}", value_counter);
set_exprs.push(format!(
"{} = {} - {}",
assignment.attribute, attribute, placeholder
));
values.insert(placeholder, value.to_kstone_value());
value_counter += 1;
}
}
}
expression_parts.push(format!("SET {}", set_exprs.join(", ")));
}
if !stmt.remove_attributes.is_empty() {
expression_parts.push(format!("REMOVE {}", stmt.remove_attributes.join(", ")));
}
let expression = expression_parts.join(" ");
Ok(UpdateTranslation {
key,
expression,
values,
})
}
pub fn translate_delete(stmt: &DeleteStatement) -> Result<DeleteTranslation> {
DynamoDBValidator::validate_delete(stmt)?;
let pk_cond = stmt
.where_clause
.get_condition("pk")
.ok_or_else(|| Error::InvalidQuery("DELETE must specify pk in WHERE clause".into()))?;
let pk_bytes = Self::value_to_bytes(&pk_cond.value)?;
let sk_bytes = stmt
.where_clause
.get_condition("sk")
.map(|c| Self::value_to_bytes(&c.value))
.transpose()?;
let key = if let Some(sk) = sk_bytes {
Key::with_sk(pk_bytes.to_vec(), sk.to_vec())
} else {
Key::new(pk_bytes.to_vec())
};
Ok(DeleteTranslation { key })
}
}
#[derive(Debug)]
pub enum SelectTranslation {
Query {
pk: Bytes,
sk_condition: Option<SortKeyConditionType>,
index_name: Option<String>,
forward: bool,
},
MultiGet {
keys: Vec<Bytes>,
index_name: Option<String>,
},
Scan {
filter_conditions: Vec<Condition>,
},
}
#[derive(Debug, Clone)]
pub enum SortKeyConditionType {
Equal(Bytes),
LessThan(Bytes),
LessThanOrEqual(Bytes),
GreaterThan(Bytes),
GreaterThanOrEqual(Bytes),
Between(Bytes, Bytes),
}
#[derive(Debug)]
pub struct InsertTranslation {
pub key: Key,
pub item: crate::Item,
}
#[derive(Debug)]
pub struct UpdateTranslation {
pub key: Key,
pub expression: String,
pub values: std::collections::HashMap<String, crate::Value>,
}
#[derive(Debug)]
pub struct DeleteTranslation {
pub key: Key,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_translate_select_query() {
let stmt = SelectStatement {
table_name: "users".to_string(),
index_name: None,
select_list: SelectList::All,
where_clause: Some(WhereClause {
conditions: vec![Condition {
attribute: "pk".to_string(),
operator: CompareOp::Equal,
value: SqlValue::String("user#123".to_string()),
}],
}),
order_by: None,
limit: None,
offset: None,
};
let translation = PartiQLTranslator::translate_select(&stmt).unwrap();
match translation {
SelectTranslation::Query { pk, .. } => {
assert_eq!(pk, Bytes::from("user#123"));
}
_ => panic!("Expected Query translation"),
}
}
#[test]
fn test_translate_select_scan() {
let stmt = SelectStatement {
table_name: "users".to_string(),
index_name: None,
select_list: SelectList::All,
where_clause: None,
order_by: None,
limit: None,
offset: None,
};
let translation = PartiQLTranslator::translate_select(&stmt).unwrap();
match translation {
SelectTranslation::Scan { .. } => {}
_ => panic!("Expected Scan translation"),
}
}
#[test]
fn test_translate_insert() {
let mut map = std::collections::HashMap::new();
map.insert("pk".to_string(), SqlValue::String("user#123".to_string()));
map.insert("name".to_string(), SqlValue::String("Alice".to_string()));
map.insert("age".to_string(), SqlValue::Number("30".to_string()));
let stmt = InsertStatement {
table_name: "users".to_string(),
value: SqlValue::Map(map),
};
let translation = PartiQLTranslator::translate_insert(&stmt).unwrap();
assert_eq!(translation.key.pk.as_ref(), "user#123".as_bytes());
assert_eq!(translation.item.len(), 2); }
#[test]
fn test_translate_delete() {
let stmt = DeleteStatement {
table_name: "users".to_string(),
where_clause: WhereClause {
conditions: vec![Condition {
attribute: "pk".to_string(),
operator: CompareOp::Equal,
value: SqlValue::String("user#123".to_string()),
}],
},
};
let translation = PartiQLTranslator::translate_delete(&stmt).unwrap();
assert_eq!(translation.key.pk.as_ref(), "user#123".as_bytes());
}
#[test]
fn test_translate_update_simple() {
let stmt = UpdateStatement {
table_name: "users".to_string(),
where_clause: WhereClause {
conditions: vec![Condition {
attribute: "pk".to_string(),
operator: CompareOp::Equal,
value: SqlValue::String("user#123".to_string()),
}],
},
set_assignments: vec![
SetAssignment {
attribute: "name".to_string(),
value: SetValue::Literal(SqlValue::String("Alice".to_string())),
},
SetAssignment {
attribute: "age".to_string(),
value: SetValue::Literal(SqlValue::Number("30".to_string())),
},
],
remove_attributes: vec![],
};
let translation = PartiQLTranslator::translate_update(&stmt).unwrap();
assert_eq!(translation.key.pk.as_ref(), "user#123".as_bytes());
assert!(translation.expression.contains("SET"));
assert_eq!(translation.values.len(), 2); }
#[test]
fn test_translate_update_with_arithmetic() {
let stmt = UpdateStatement {
table_name: "users".to_string(),
where_clause: WhereClause {
conditions: vec![Condition {
attribute: "pk".to_string(),
operator: CompareOp::Equal,
value: SqlValue::String("user#123".to_string()),
}],
},
set_assignments: vec![
SetAssignment {
attribute: "age".to_string(),
value: SetValue::Add {
attribute: "age".to_string(),
value: SqlValue::Number("1".to_string()),
},
},
SetAssignment {
attribute: "count".to_string(),
value: SetValue::Subtract {
attribute: "count".to_string(),
value: SqlValue::Number("5".to_string()),
},
},
],
remove_attributes: vec![],
};
let translation = PartiQLTranslator::translate_update(&stmt).unwrap();
assert!(translation.expression.contains("age = age + :v1"));
assert!(translation.expression.contains("count = count - :v2"));
assert_eq!(translation.values.len(), 2);
}
#[test]
fn test_translate_update_with_remove() {
let stmt = UpdateStatement {
table_name: "users".to_string(),
where_clause: WhereClause {
conditions: vec![Condition {
attribute: "pk".to_string(),
operator: CompareOp::Equal,
value: SqlValue::String("user#123".to_string()),
}],
},
set_assignments: vec![SetAssignment {
attribute: "name".to_string(),
value: SetValue::Literal(SqlValue::String("Alice".to_string())),
}],
remove_attributes: vec!["tags".to_string(), "metadata".to_string()],
};
let translation = PartiQLTranslator::translate_update(&stmt).unwrap();
assert!(translation.expression.contains("SET"));
assert!(translation.expression.contains("REMOVE tags, metadata"));
assert_eq!(translation.values.len(), 1); }
#[test]
fn test_translate_update_remove_only() {
let stmt = UpdateStatement {
table_name: "users".to_string(),
where_clause: WhereClause {
conditions: vec![
Condition {
attribute: "pk".to_string(),
operator: CompareOp::Equal,
value: SqlValue::String("user#123".to_string()),
},
Condition {
attribute: "sk".to_string(),
operator: CompareOp::Equal,
value: SqlValue::String("profile".to_string()),
},
],
},
set_assignments: vec![],
remove_attributes: vec!["tags".to_string(), "metadata".to_string()],
};
let translation = PartiQLTranslator::translate_update(&stmt).unwrap();
assert_eq!(translation.key.pk.as_ref(), "user#123".as_bytes());
assert_eq!(translation.key.sk.as_ref().map(|b| b.as_ref()), Some("profile".as_bytes()));
assert_eq!(translation.expression, "REMOVE tags, metadata");
assert_eq!(translation.values.len(), 0); }
}