use std::collections::HashMap;
use parse_rust_core::ParseError;
use parse_rust_rest::{parse_include, parse_where, FindOptions, ParsedWhere};
use parse_rust_storage::{QueryOptions, DEFAULT_LIMIT};
use serde_json::Value as Json;
#[derive(Debug, Clone, Default)]
pub struct Params(HashMap<String, String>);
const FIND_KEYS: [&str; 16] = [
"skip",
"limit",
"order",
"count",
"keys",
"excludeKeys",
"include",
"includeAll",
"redirectClassNameForKey",
"where",
"readPreference",
"includeReadPreference",
"subqueryReadPreference",
"hint",
"explain",
"comment",
];
const GET_KEYS: [&str; 6] = [
"keys",
"include",
"excludeKeys",
"readPreference",
"includeReadPreference",
"subqueryReadPreference",
];
const UNIMPLEMENTED_KEYS: [&str; 5] = [
"includeAll",
"redirectClassNameForKey",
"hint",
"explain",
"comment",
];
impl Params {
pub fn from_map(map: HashMap<String, String>) -> Self {
Self(map)
}
pub fn from_json(value: Option<&Json>) -> Self {
let mut map = HashMap::new();
if let Some(Json::Object(object)) = value {
for (key, value) in object {
let text = match value {
Json::String(s) => s.clone(),
other => other.to_string(),
};
map.insert(key.clone(), text);
}
}
Self(map)
}
pub fn get(&self, key: &str) -> Option<&str> {
self.0.get(key).map(String::as_str)
}
pub fn reject_unknown_find_keys(&self) -> Result<(), ParseError> {
for key in self.0.keys() {
if !FIND_KEYS.contains(&key.as_str()) {
return Err(ParseError::invalid_query(format!(
"Invalid parameter for query: {key}"
)));
}
}
self.reject_unimplemented()
}
pub fn reject_unknown_get_keys(&self) -> Result<(), ParseError> {
for key in self.0.keys() {
if !GET_KEYS.contains(&key.as_str()) {
return Err(ParseError::invalid_query("Improper encode of parameter"));
}
}
self.reject_unimplemented()
}
fn reject_unimplemented(&self) -> Result<(), ParseError> {
for key in UNIMPLEMENTED_KEYS {
if self.0.contains_key(key) {
return Err(ParseError::new(
parse_rust_core::ErrorCode::CommandUnavailable,
format!("The {key} query parameter is not supported yet."),
));
}
}
Ok(())
}
pub fn parse_where(&self) -> Result<ParsedWhere, ParseError> {
let Some(raw) = self.get("where") else {
return Ok(ParsedWhere::default());
};
let value: Json = serde_json::from_str(raw)
.map_err(|_| ParseError::invalid_json("where parameter is not valid JSON"))?;
parse_where(&value)
}
pub fn wants_count(&self) -> bool {
!matches!(
self.get("count"),
None | Some("0") | Some("false") | Some("")
)
}
pub fn find_options(&self) -> Result<FindOptions, ParseError> {
Ok(FindOptions {
limit: Some(
self.get("limit")
.and_then(|v| v.parse::<u32>().ok())
.unwrap_or(DEFAULT_LIMIT),
),
skip: self.get("skip").and_then(|v| v.parse().ok()),
order: self
.get("order")
.map(QueryOptions::parse_order)
.unwrap_or_default(),
keys: self.csv("keys"),
exclude_keys: self.csv("excludeKeys"),
include: match self.get("include") {
Some(raw) => parse_include(raw)?,
None => Vec::new(),
},
})
}
pub fn get_options(&self) -> Result<FindOptions, ParseError> {
Ok(FindOptions {
limit: Some(1),
skip: None,
order: Vec::new(),
keys: self.csv("keys"),
exclude_keys: self.csv("excludeKeys"),
include: match self.get("include") {
Some(raw) => parse_include(raw)?,
None => Vec::new(),
},
})
}
fn csv(&self, key: &str) -> Option<Vec<String>> {
self.get(key).map(|raw| {
raw.split(',')
.map(str::trim)
.filter(|s| !s.is_empty())
.map(str::to_string)
.collect()
})
}
}
#[cfg(test)]
mod tests {
use super::*;
fn params(pairs: &[(&str, &str)]) -> Params {
Params::from_map(
pairs
.iter()
.map(|(k, v)| ((*k).to_string(), (*v).to_string()))
.collect(),
)
}
#[test]
fn an_unknown_find_parameter_is_named_in_the_error() {
let e = params(&[("nonsense", "1")])
.reject_unknown_find_keys()
.unwrap_err();
assert_eq!(e.code, parse_rust_core::ErrorCode::InvalidQuery);
assert_eq!(e.message, "Invalid parameter for query: nonsense");
}
#[test]
fn an_unknown_get_parameter_is_not_named() {
let e = params(&[("limit", "1")])
.reject_unknown_get_keys()
.unwrap_err();
assert_eq!(e.message, "Improper encode of parameter");
}
#[test]
fn unimplemented_parameters_are_refused_rather_than_ignored() {
for key in UNIMPLEMENTED_KEYS {
let e = params(&[(key, "1")])
.reject_unknown_find_keys()
.unwrap_err();
assert_eq!(
e.code,
parse_rust_core::ErrorCode::CommandUnavailable,
"{key}"
);
}
assert!(params(&[("readPreference", "SECONDARY")])
.reject_unknown_find_keys()
.is_ok());
}
#[test]
fn a_batch_sub_request_carries_its_parameters_as_json() {
let body: Json = serde_json::from_str(r#"{"where":{"a":1},"limit":5}"#).expect("literal");
let p = Params::from_json(Some(&body));
assert_eq!(p.get("where"), Some(r#"{"a":1}"#));
assert_eq!(p.get("limit"), Some("5"));
assert_eq!(p.find_options().expect("options").limit, Some(5));
}
#[test]
fn a_malformed_where_reports_upstreams_message() {
let e = params(&[("where", "{oops")]).parse_where().unwrap_err();
assert_eq!(e.code, parse_rust_core::ErrorCode::InvalidJson);
assert_eq!(e.message, "where parameter is not valid JSON");
}
#[test]
fn count_is_a_truthiness_test() {
assert!(params(&[("count", "1")]).wants_count());
assert!(params(&[("count", "true")]).wants_count());
assert!(!params(&[("count", "0")]).wants_count());
assert!(!params(&[]).wants_count());
}
#[test]
fn limit_falls_back_to_the_parse_default_and_zero_is_honoured() {
assert_eq!(params(&[]).find_options().expect("o").limit, Some(100));
assert_eq!(
params(&[("limit", "0")]).find_options().expect("o").limit,
Some(0)
);
}
}