use azure_core::http::headers::HeaderName;
use azure_core::http::Request;
use percent_encoding::percent_decode_str;
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) enum OperationType {
ReadAccount,
CreateDatabase,
ReadDatabase,
DeleteDatabase,
CreateContainer,
ReadContainer,
DeleteContainer,
ReadPKRanges,
Create,
Read,
Replace,
Upsert,
Delete,
Query,
Unsupported(String),
BadRequestPath(String),
}
#[derive(Debug, Clone)]
pub(crate) struct ParsedRequest {
pub operation: OperationType,
pub db_id: Option<String>,
pub coll_id: Option<String>,
pub doc_id: Option<String>,
pub partition_key_header: Option<String>,
pub if_match: Option<String>,
pub if_none_match: Option<String>,
pub session_token: Option<String>,
pub activity_id: Option<String>,
pub content_response_on_write: bool,
pub offer_throughput: Option<u32>,
#[allow(dead_code)]
pub is_upsert: bool, }
static IS_UPSERT: HeaderName = HeaderName::from_static("x-ms-documentdb-is-upsert");
static PARTITION_KEY: HeaderName = HeaderName::from_static("x-ms-documentdb-partitionkey");
static IF_MATCH: HeaderName = HeaderName::from_static("if-match");
static IF_NONE_MATCH: HeaderName = HeaderName::from_static("if-none-match");
static SESSION_TOKEN: HeaderName = HeaderName::from_static("x-ms-session-token");
static ACTIVITY_ID: HeaderName = HeaderName::from_static("x-ms-activity-id");
static CONTENT_RESPONSE: HeaderName =
HeaderName::from_static("x-ms-cosmos-populate-content-response-on-write");
static PREFER: HeaderName = HeaderName::from_static("prefer");
static IS_QUERY: HeaderName = HeaderName::from_static("x-ms-documentdb-query");
static OFFER_THROUGHPUT: HeaderName = HeaderName::from_static("x-ms-offer-throughput");
pub(crate) fn parse_request(request: &Request) -> ParsedRequest {
let url = request.url();
let method = request.method();
let headers = request.headers();
let partition_key_header = headers
.get_optional_str(&PARTITION_KEY)
.map(|s| s.to_string());
let if_match = headers.get_optional_str(&IF_MATCH).map(|s| s.to_string());
let if_none_match = headers
.get_optional_str(&IF_NONE_MATCH)
.map(|s| s.to_string());
let session_token = headers
.get_optional_str(&SESSION_TOKEN)
.map(|s| s.to_string());
let activity_id = headers
.get_optional_str(&ACTIVITY_ID)
.map(|s| s.to_string());
let content_response_on_write = if let Some(val) = headers.get_optional_str(&CONTENT_RESPONSE) {
val.eq_ignore_ascii_case("true")
} else if let Some(prefer) = headers.get_optional_str(&PREFER) {
!prefer.contains("return=minimal")
} else {
true
};
let is_upsert = headers
.get_optional_str(&IS_UPSERT)
.map(|s| s.eq_ignore_ascii_case("true"))
.unwrap_or(false);
let is_query = headers
.get_optional_str(&IS_QUERY)
.map(|v| v.eq_ignore_ascii_case("true"))
.unwrap_or(false);
let offer_throughput = headers
.get_optional_str(&OFFER_THROUGHPUT)
.and_then(|s| s.trim().parse::<u32>().ok());
let path = url.path();
let has_trailing_slash = path.len() > 1 && path.ends_with('/');
let segments = parse_path_segments(path);
let operation = if has_trailing_slash {
OperationType::BadRequestPath(format!(
"{} {} (trailing slash rejected)",
method.as_ref(),
path
))
} else {
resolve_operation(method.as_ref(), &segments, is_upsert, is_query)
};
let db_id = segment_after_keyword(&segments, 0, "dbs");
let coll_id = segment_after_keyword(&segments, 2, "colls");
let doc_id = segment_after_keyword(&segments, 4, "docs");
ParsedRequest {
operation,
db_id,
coll_id,
doc_id,
partition_key_header,
if_match,
if_none_match,
session_token,
activity_id,
content_response_on_write,
offer_throughput,
is_upsert,
}
}
fn parse_path_segments(path: &str) -> Vec<String> {
path.split('/')
.filter(|s| !s.is_empty())
.map(|s| {
percent_decode_str(s)
.decode_utf8()
.map(|cow| cow.into_owned())
.unwrap_or_else(|_| s.to_string())
})
.collect()
}
fn segment_after_keyword(
segments: &[String],
keyword_index: usize,
keyword: &str,
) -> Option<String> {
if segments.get(keyword_index).map(String::as_str) != Some(keyword) {
return None;
}
segments.get(keyword_index + 1).cloned()
}
fn resolve_operation(
method: &str,
segments: &[String],
is_upsert: bool,
is_query: bool,
) -> OperationType {
let depth = segments.len();
match (method, depth) {
("GET", 0) => OperationType::ReadAccount,
("POST", 1) if segments[0] == "dbs" => OperationType::CreateDatabase,
("GET", 2) if segments[0] == "dbs" => OperationType::ReadDatabase,
("DELETE", 2) if segments[0] == "dbs" => OperationType::DeleteDatabase,
("POST", 3) if segments[0] == "dbs" && segments[2] == "colls" => {
OperationType::CreateContainer
}
("GET", 4) if segments[0] == "dbs" && segments[2] == "colls" => {
OperationType::ReadContainer
}
("DELETE", 4) if segments[0] == "dbs" && segments[2] == "colls" => {
OperationType::DeleteContainer
}
("GET", 5)
if segments[0] == "dbs" && segments[2] == "colls" && segments[4] == "pkranges" =>
{
OperationType::ReadPKRanges
}
("POST", 5) if segments[0] == "dbs" && segments[2] == "colls" && segments[4] == "docs" => {
if is_query {
OperationType::Query
} else if is_upsert {
OperationType::Upsert
} else {
OperationType::Create
}
}
("GET", 6) if segments[0] == "dbs" && segments[2] == "colls" && segments[4] == "docs" => {
OperationType::Read
}
("PUT", 6) if segments[0] == "dbs" && segments[2] == "colls" && segments[4] == "docs" => {
OperationType::Replace
}
("DELETE", 6)
if segments[0] == "dbs" && segments[2] == "colls" && segments[4] == "docs" =>
{
OperationType::Delete
}
_ => OperationType::Unsupported(format!("{} {}", method, segments.join("/"))),
}
}
pub(crate) fn resolve_region<'a>(
url: &azure_core::http::Url,
config: &'a super::config::VirtualAccountConfig,
) -> Option<&'a str> {
config.region_for_url(url)
}
#[cfg(test)]
mod tests {
use super::*;
use azure_core::http::{Method, Request, Url};
fn make_request(method: &str, path: &str) -> Request {
let url = format!("https://test.emulator.local{}", path);
let method = match method {
"GET" => Method::Get,
"POST" => Method::Post,
"PUT" => Method::Put,
"DELETE" => Method::Delete,
_ => Method::Get,
};
Request::new(url.parse().unwrap(), method)
}
#[test]
fn read_account() {
let req = make_request("GET", "/");
let parsed = parse_request(&req);
assert_eq!(parsed.operation, OperationType::ReadAccount);
}
#[test]
fn segment_after_keyword_anchors_by_position() {
let segments = vec![
"dbs".to_string(),
"colls".to_string(),
"colls".to_string(),
"mycoll".to_string(),
"docs".to_string(),
"d1".to_string(),
];
assert_eq!(
segment_after_keyword(&segments, 0, "dbs"),
Some("colls".to_string())
);
assert_eq!(
segment_after_keyword(&segments, 2, "colls"),
Some("mycoll".to_string())
);
assert_eq!(
segment_after_keyword(&segments, 4, "docs"),
Some("d1".to_string())
);
}
#[test]
fn segment_after_keyword_returns_none_when_keyword_absent_at_position() {
let segments = vec!["dbs".to_string(), "mydb".to_string()];
assert_eq!(segment_after_keyword(&segments, 2, "colls"), None);
}
#[test]
fn parse_request_resolves_container_named_after_keyword() {
let req = make_request("GET", "/dbs/colls/colls/mycoll");
let parsed = parse_request(&req);
assert_eq!(parsed.operation, OperationType::ReadContainer);
assert_eq!(parsed.db_id.as_deref(), Some("colls"));
assert_eq!(parsed.coll_id.as_deref(), Some("mycoll"));
}
#[test]
fn parse_request_resolves_document_when_container_named_docs() {
let req = make_request("GET", "/dbs/mydb/colls/docs/docs/d1");
let parsed = parse_request(&req);
assert_eq!(parsed.operation, OperationType::Read);
assert_eq!(parsed.db_id.as_deref(), Some("mydb"));
assert_eq!(parsed.coll_id.as_deref(), Some("docs"));
assert_eq!(parsed.doc_id.as_deref(), Some("d1"));
}
#[test]
fn create_database() {
let req = make_request("POST", "/dbs");
let parsed = parse_request(&req);
assert_eq!(parsed.operation, OperationType::CreateDatabase);
}
#[test]
fn read_database() {
let req = make_request("GET", "/dbs/mydb");
let parsed = parse_request(&req);
assert_eq!(parsed.operation, OperationType::ReadDatabase);
assert_eq!(parsed.db_id.as_deref(), Some("mydb"));
}
#[test]
fn create_container() {
let req = make_request("POST", "/dbs/mydb/colls");
let parsed = parse_request(&req);
assert_eq!(parsed.operation, OperationType::CreateContainer);
assert_eq!(parsed.db_id.as_deref(), Some("mydb"));
}
#[test]
fn read_document() {
let req = make_request("GET", "/dbs/mydb/colls/mycoll/docs/doc1");
let parsed = parse_request(&req);
assert_eq!(parsed.operation, OperationType::Read);
assert_eq!(parsed.db_id.as_deref(), Some("mydb"));
assert_eq!(parsed.coll_id.as_deref(), Some("mycoll"));
assert_eq!(parsed.doc_id.as_deref(), Some("doc1"));
}
#[test]
fn create_document() {
let req = make_request("POST", "/dbs/mydb/colls/mycoll/docs");
let parsed = parse_request(&req);
assert_eq!(parsed.operation, OperationType::Create);
}
#[test]
fn upsert_document() {
let url: Url = "https://test.emulator.local/dbs/mydb/colls/mycoll/docs"
.parse()
.unwrap();
let mut req = Request::new(url, Method::Post);
req.headers_mut().insert(
IS_UPSERT.clone(),
azure_core::http::headers::HeaderValue::from("True".to_string()),
);
let parsed = parse_request(&req);
assert_eq!(parsed.operation, OperationType::Upsert);
}
#[test]
fn pkranges() {
let req = make_request("GET", "/dbs/mydb/colls/mycoll/pkranges");
let parsed = parse_request(&req);
assert_eq!(parsed.operation, OperationType::ReadPKRanges);
}
#[test]
fn trailing_slash_on_docs_collection_is_rejected() {
let req = make_request("POST", "/dbs/mydb/colls/mycoll/docs/");
let parsed = parse_request(&req);
assert!(matches!(parsed.operation, OperationType::BadRequestPath(_)));
}
#[test]
fn trailing_slash_on_document_is_rejected() {
let req = make_request("GET", "/dbs/mydb/colls/mycoll/docs/d1/");
let parsed = parse_request(&req);
assert!(matches!(parsed.operation, OperationType::BadRequestPath(_)));
}
#[test]
fn root_path_is_still_read_account() {
let req = make_request("GET", "/");
let parsed = parse_request(&req);
assert_eq!(parsed.operation, OperationType::ReadAccount);
}
}