use super::*;
pub(crate) fn keys_equal(
table: &DynamoTable,
a: &HashMap<String, AttributeValue>,
b: &HashMap<String, AttributeValue>,
) -> bool {
use super::partiql::values_equal;
let hash_key = table.hash_key_name();
if !(values_equal(a.get(hash_key), b.get(hash_key))
&& a.get(hash_key).is_some()
&& b.get(hash_key).is_some())
{
return false;
}
match table.range_key_name() {
Some(rk) => values_equal(a.get(rk), b.get(rk)),
None => true,
}
}
pub(crate) fn extract_key(
table: &DynamoTable,
item: &HashMap<String, AttributeValue>,
) -> HashMap<String, AttributeValue> {
let mut key = HashMap::new();
let hash_key = table.hash_key_name();
if let Some(v) = item.get(hash_key) {
key.insert(hash_key.to_string(), v.clone());
}
if let Some(range_key) = table.range_key_name() {
if let Some(v) = item.get(range_key) {
key.insert(range_key.to_string(), v.clone());
}
}
key
}
pub(crate) fn parse_key_map(value: &Value) -> Option<HashMap<String, AttributeValue>> {
let obj = value.as_object()?;
if obj.is_empty() {
return None;
}
Some(obj.iter().map(|(k, v)| (k.clone(), v.clone())).collect())
}
pub(crate) fn item_matches_key(
item: &HashMap<String, AttributeValue>,
key: &HashMap<String, AttributeValue>,
hash_key_name: &str,
range_key_name: Option<&str>,
) -> bool {
let hash_match = match (item.get(hash_key_name), key.get(hash_key_name)) {
(Some(a), Some(b)) => a == b,
_ => false,
};
if !hash_match {
return false;
}
match range_key_name {
Some(rk) => match (item.get(rk), key.get(rk)) {
(Some(a), Some(b)) => a == b,
(None, None) => true,
_ => false,
},
None => true,
}
}
pub(crate) fn extract_key_for_schema(
item: &HashMap<String, AttributeValue>,
hash_key_name: &str,
range_key_name: Option<&str>,
) -> HashMap<String, AttributeValue> {
let mut key = HashMap::new();
if let Some(v) = item.get(hash_key_name) {
key.insert(hash_key_name.to_string(), v.clone());
}
if let Some(rk) = range_key_name {
if let Some(v) = item.get(rk) {
key.insert(rk.to_string(), v.clone());
}
}
key
}
pub(crate) fn validate_item_attribute_values(
item: &HashMap<String, AttributeValue>,
) -> Result<(), AwsServiceError> {
for v in item.values() {
validate_attribute_value(v)?;
}
Ok(())
}
fn validate_attribute_value(v: &Value) -> Result<(), AwsServiceError> {
let Some((tag, val)) = v.as_object().and_then(|o| o.iter().next()) else {
return Ok(());
};
let bad_number = |n: &str| {
AwsServiceError::aws_error(
StatusCode::BAD_REQUEST,
"ValidationException",
format!("The parameter cannot be converted to a numeric value: {n}"),
)
};
match tag.as_str() {
"N" => {
let s = val.as_str().unwrap_or_default();
if !is_valid_number(s) {
return Err(bad_number(s));
}
}
"NS" => {
for el in val.as_array().into_iter().flatten() {
let s = el.as_str().unwrap_or_default();
if !is_valid_number(s) {
return Err(bad_number(s));
}
}
}
"L" => {
for el in val.as_array().into_iter().flatten() {
validate_attribute_value(el)?;
}
}
"M" => {
if let Some(m) = val.as_object() {
for el in m.values() {
validate_attribute_value(el)?;
}
}
}
_ => {}
}
Ok(())
}
pub(crate) fn validate_key_in_item(
table: &DynamoTable,
item: &HashMap<String, AttributeValue>,
) -> Result<(), AwsServiceError> {
let hash_key = table.hash_key_name();
if !item.contains_key(hash_key) {
return Err(AwsServiceError::aws_error(
StatusCode::BAD_REQUEST,
"ValidationException",
format!("Missing the key {hash_key} in the item"),
));
}
if let Some(range_key) = table.range_key_name() {
if !item.contains_key(range_key) {
return Err(AwsServiceError::aws_error(
StatusCode::BAD_REQUEST,
"ValidationException",
format!("Missing the key {range_key} in the item"),
));
}
}
check_key_type(table, item, hash_key)?;
if let Some(range_key) = table.range_key_name() {
check_key_type(table, item, range_key)?;
}
Ok(())
}
pub(crate) fn validate_key_attributes_in_key(
table: &DynamoTable,
key: &HashMap<String, AttributeValue>,
) -> Result<(), AwsServiceError> {
let hash_key = table.hash_key_name();
if !key.contains_key(hash_key) {
return Err(AwsServiceError::aws_error(
StatusCode::BAD_REQUEST,
"ValidationException",
format!("Missing the key {hash_key} in the item"),
));
}
if let Some(range_key) = table.range_key_name() {
if !key.contains_key(range_key) {
return Err(AwsServiceError::aws_error(
StatusCode::BAD_REQUEST,
"ValidationException",
format!("Missing the key {range_key} in the item"),
));
}
}
check_key_type(table, key, hash_key)?;
if let Some(range_key) = table.range_key_name() {
check_key_type(table, key, range_key)?;
}
Ok(())
}
fn check_key_type(
table: &DynamoTable,
attrs: &HashMap<String, AttributeValue>,
name: &str,
) -> Result<(), AwsServiceError> {
let Some(val) = attrs.get(name) else {
return Ok(());
};
let Some(expected) = table
.attribute_definitions
.iter()
.find(|d| d.attribute_name == name)
.map(|d| d.attribute_type.as_str())
else {
return Ok(());
};
let actual = val
.as_object()
.and_then(|o| o.keys().next().map(|k| k.as_str()));
if actual != Some(expected) {
return Err(AwsServiceError::aws_error(
StatusCode::BAD_REQUEST,
"ValidationException",
format!(
"One or more parameter values were invalid: Type mismatch for key {name} expected: {expected} actual: {}",
actual.unwrap_or("NULL"),
),
));
}
Ok(())
}
#[cfg(test)]
mod attr_value_validation_tests {
use super::*;
use serde_json::json;
#[test]
fn rejects_invalid_number_attribute() {
let item: HashMap<String, AttributeValue> =
HashMap::from([("n".to_string(), json!({"N": "abc"}))]);
assert!(validate_item_attribute_values(&item).is_err());
}
#[test]
fn rejects_invalid_number_nested_in_list_and_map() {
let item: HashMap<String, AttributeValue> = HashMap::from([(
"l".to_string(),
json!({"L": [{"N": "1"}, {"M": {"x": {"N": "oops"}}}]}),
)]);
assert!(validate_item_attribute_values(&item).is_err());
}
#[test]
fn rejects_invalid_number_set_member() {
let item: HashMap<String, AttributeValue> =
HashMap::from([("ns".to_string(), json!({"NS": ["1", "x"]}))]);
assert!(validate_item_attribute_values(&item).is_err());
}
#[test]
fn accepts_valid_values() {
let item: HashMap<String, AttributeValue> = HashMap::from([
("n".to_string(), json!({"N": "3.14"})),
("s".to_string(), json!({"S": "hi"})),
("ns".to_string(), json!({"NS": ["1", "2.5"]})),
]);
assert!(validate_item_attribute_values(&item).is_ok());
}
}