use crate::models::{PartitionKeyKind, PartitionKeyValue, PartitionKeyVersion};
pub use crate::models::effective_partition_key::EffectivePartitionKey as Epk;
pub(crate) type PartitionKeyComponent = PartitionKeyValue;
pub(crate) fn compute_epk(
components: &[PartitionKeyValue],
kind: PartitionKeyKind,
version: PartitionKeyVersion,
) -> Epk {
Epk::compute(components, kind, version)
}
pub(crate) fn parse_partition_key_header(
header: &str,
) -> crate::error::Result<Vec<PartitionKeyValue>> {
let trimmed = header.trim();
if trimmed.is_empty() || trimmed == "[]" {
return Ok(Vec::new());
}
let value: serde_json::Value = serde_json::from_str(trimmed).map_err(|e| {
crate::error::CosmosError::builder()
.with_status(crate::error::CosmosStatus::new(
azure_core::http::StatusCode::BadRequest,
))
.with_message(format!("invalid partition key header: {e}"))
.build()
})?;
let arr = value.as_array().ok_or_else(|| {
crate::error::CosmosError::builder()
.with_status(crate::error::CosmosStatus::new(
azure_core::http::StatusCode::BadRequest,
))
.with_message("partition key header must be a JSON array")
.build()
})?;
arr.iter().map(json_to_pk_component).collect()
}
pub(crate) fn extract_pk_from_body(
body: &serde_json::Value,
pk_paths: &[impl AsRef<str>],
) -> crate::error::Result<Vec<PartitionKeyValue>> {
if !body.is_object() {
return Err(crate::error::CosmosError::builder()
.with_status(crate::error::CosmosStatus::new(
azure_core::http::StatusCode::BadRequest,
))
.with_message("document body must be a JSON object to extract a partition key")
.build());
}
pk_paths
.iter()
.map(|path| extract_pk_at_path(body, path.as_ref()))
.collect()
}
fn extract_pk_at_path(
body: &serde_json::Value,
path: &str,
) -> crate::error::Result<PartitionKeyValue> {
let path_str = path.trim_start_matches('/');
if path_str.is_empty() {
return json_to_pk_component(body);
}
let segments: Vec<&str> = path_str.split('/').collect();
let last_idx = segments.len() - 1;
let mut current = body;
for (i, segment) in segments.iter().enumerate() {
let obj = current.as_object().ok_or_else(|| {
crate::error::CosmosError::builder()
.with_status(crate::error::CosmosStatus::new(
azure_core::http::StatusCode::BadRequest,
))
.with_message(format!(
"partition key path component '{segment}' encountered a non-object intermediate"
))
.build()
})?;
match obj.get(*segment) {
Some(next) if i == last_idx => return json_to_pk_component(next),
Some(next) => current = next,
None => return Ok(PartitionKeyValue::UNDEFINED),
}
}
Ok(PartitionKeyValue::UNDEFINED)
}
fn json_to_pk_component(value: &serde_json::Value) -> crate::error::Result<PartitionKeyValue> {
match value {
serde_json::Value::Null => Ok(Option::<&str>::None.into()),
serde_json::Value::Bool(b) => Ok(PartitionKeyValue::from(*b)),
serde_json::Value::String(s) => Ok(PartitionKeyValue::from(s.clone())),
serde_json::Value::Number(n) => {
let f = n.as_f64().ok_or_else(|| {
crate::error::CosmosError::builder()
.with_status(crate::error::CosmosStatus::new(
azure_core::http::StatusCode::BadRequest,
))
.with_message("partition key number is not representable as f64")
.build()
})?;
if !f.is_finite() {
return Err(crate::error::CosmosError::builder()
.with_status(crate::error::CosmosStatus::new(
azure_core::http::StatusCode::BadRequest,
))
.with_message(
"partition key numbers must be finite (NaN and Infinity are not allowed)",
)
.build());
}
Ok(PartitionKeyValue::from(f))
}
serde_json::Value::Object(_) | serde_json::Value::Array(_) => {
Err(crate::error::CosmosError::builder()
.with_status(crate::error::CosmosStatus::new(
azure_core::http::StatusCode::BadRequest,
))
.with_message(
"partition key components must be scalar (null, bool, number, or string)",
)
.build())
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_pk_header_string() {
let components = parse_partition_key_header(r#"["hello"]"#).unwrap();
assert_eq!(
components,
vec![PartitionKeyValue::from("hello".to_string())]
);
}
#[test]
fn parse_pk_header_number() {
let components = parse_partition_key_header("[42]").unwrap();
assert_eq!(components, vec![PartitionKeyValue::from(42.0_f64)]);
}
#[test]
fn parse_pk_header_null() {
let components = parse_partition_key_header("[null]").unwrap();
let expected: PartitionKeyValue = Option::<&str>::None.into();
assert_eq!(components, vec![expected]);
}
#[test]
fn parse_pk_header_hierarchical() {
let components = parse_partition_key_header(r#"["tenant1", "user1"]"#).unwrap();
assert_eq!(
components,
vec![
PartitionKeyValue::from("tenant1".to_string()),
PartitionKeyValue::from("user1".to_string()),
]
);
}
#[test]
fn parse_pk_header_empty_is_cross_partition() {
assert!(parse_partition_key_header("[]").unwrap().is_empty());
assert!(parse_partition_key_header("").unwrap().is_empty());
}
#[test]
fn parse_pk_header_invalid_json_errors() {
assert!(parse_partition_key_header("not json").is_err());
assert!(parse_partition_key_header("{\"x\":1}").is_err());
}
#[test]
fn parse_pk_header_object_errors() {
let err = parse_partition_key_header(r#"[{"x":1}]"#).unwrap_err();
assert!(err.to_string().contains("scalar"));
}
#[test]
fn extract_pk_from_json_body() {
let body = serde_json::json!({"id": "doc1", "pk": "value1", "nested": {"key": 42}});
let components = extract_pk_from_body(&body, &["/pk"]).unwrap();
assert_eq!(
components,
vec![PartitionKeyValue::from("value1".to_string())]
);
}
#[test]
fn extract_pk_nested() {
let body = serde_json::json!({"id": "doc1", "nested": {"key": 42}});
let components = extract_pk_from_body(&body, &["/nested/key"]).unwrap();
assert_eq!(components, vec![PartitionKeyValue::from(42.0_f64)]);
}
#[test]
fn extract_pk_missing_path_is_undefined() {
let body = serde_json::json!({"id": "doc1"});
let components = extract_pk_from_body(&body, &["/missing"]).unwrap();
assert_eq!(components, vec![PartitionKeyValue::UNDEFINED]);
}
#[test]
fn extract_pk_object_value_errors() {
let body = serde_json::json!({"pk": {"nested": "object"}});
assert!(extract_pk_from_body(&body, &["/pk"]).is_err());
}
#[test]
fn extract_pk_non_object_body_errors() {
let body = serde_json::json!([1, 2, 3]);
assert!(extract_pk_from_body(&body, &["/pk"]).is_err());
let body = serde_json::json!("string");
assert!(extract_pk_from_body(&body, &["/pk"]).is_err());
}
#[test]
fn extract_pk_non_object_intermediate_errors() {
let body = serde_json::json!({"a": 42});
assert!(extract_pk_from_body(&body, &["/a/b"]).is_err());
let body = serde_json::json!({"a": [1, 2, 3]});
assert!(extract_pk_from_body(&body, &["/a/b"]).is_err());
}
#[test]
fn compute_epk_uses_product_v2_vector() {
let epk = compute_epk(
&[PartitionKeyValue::from("customer42".to_string())],
PartitionKeyKind::Hash,
PartitionKeyVersion::V2,
);
assert_eq!(epk.as_str(), "19819C94CE42A1654CCC8110539D9589");
}
#[test]
fn header_and_body_extraction_produce_identical_epk() {
let cases: Vec<(&str, serde_json::Value, &str)> = vec![
("integer 42", serde_json::json!({"pk": 42}), "[42]"),
("integer 0", serde_json::json!({"pk": 0}), "[0]"),
("negative integer", serde_json::json!({"pk": -7}), "[-7]"),
(
"integer-as-f64 42.0",
serde_json::json!({"pk": 42.0}),
"[42.0]",
),
("fractional 1.5", serde_json::json!({"pk": 1.5}), "[1.5]"),
("large 1e10", serde_json::json!({"pk": 1e10}), "[1e10]"),
(
"u64 boundary 2^53",
serde_json::json!({"pk": 9_007_199_254_740_992_u64}),
"[9007199254740992]",
),
("string", serde_json::json!({"pk": "abc"}), r#"["abc"]"#),
("bool true", serde_json::json!({"pk": true}), "[true]"),
("bool false", serde_json::json!({"pk": false}), "[false]"),
("null", serde_json::json!({"pk": null}), "[null]"),
];
for (label, body, header) in cases {
let from_body = extract_pk_from_body(&body, &["/pk"]).unwrap_or_else(|e| {
panic!("body extraction failed for {}: {}", label, e);
});
let from_header = parse_partition_key_header(header).unwrap_or_else(|e| {
panic!("header parsing failed for {}: {}", label, e);
});
assert_eq!(
from_body, from_header,
"components diverge for {}: body={:?} header={:?}",
label, from_body, from_header
);
let epk_body = compute_epk(&from_body, PartitionKeyKind::Hash, PartitionKeyVersion::V2);
let epk_header = compute_epk(
&from_header,
PartitionKeyKind::Hash,
PartitionKeyVersion::V2,
);
assert_eq!(
epk_body,
epk_header,
"EPK diverges for {}: body={} header={}",
label,
epk_body.as_str(),
epk_header.as_str(),
);
let epk_body_v1 =
compute_epk(&from_body, PartitionKeyKind::Hash, PartitionKeyVersion::V1);
let epk_header_v1 = compute_epk(
&from_header,
PartitionKeyKind::Hash,
PartitionKeyVersion::V1,
);
assert_eq!(
epk_body_v1,
epk_header_v1,
"V1 EPK diverges for {}: body={} header={}",
label,
epk_body_v1.as_str(),
epk_header_v1.as_str(),
);
}
}
}