use crate::errors::Error;
use crate::Result;
use percent_encoding::percent_encode_byte;
use url::{form_urlencoded, Url};
fn query_escape(s: &str) -> String {
let mut escaped = String::with_capacity(s.len());
for &byte in s.as_bytes() {
match byte {
b'-' | b'.' | b'0'..=b'9' | b'A'..=b'Z' | b'_' | b'a'..=b'z' | b'~' => {
escaped.push(byte as char)
}
b' ' => escaped.push('+'),
_ => escaped.push_str(percent_encode_byte(byte)),
}
}
escaped
}
pub fn filter_query_params(url: &str, filtered_query_params: &[String]) -> Result<String> {
if filtered_query_params.is_empty() {
return Ok(url.to_string());
}
let mut url =
Url::parse(url).map_err(|err| Error::InvalidArgument(format!("invalid url: {err}")))?;
let mut query_pairs: Vec<(String, String)> = url
.query()
.unwrap_or_default()
.split('&')
.filter(|segment| !segment.contains(';'))
.flat_map(|segment| form_urlencoded::parse(segment.as_bytes()))
.filter(|(k, _)| {
!filtered_query_params
.iter()
.any(|param| param.as_str() == k.as_ref())
})
.map(|(k, v)| (k.into_owned(), v.into_owned()))
.collect();
query_pairs.sort_by(|(a, _), (b, _)| a.cmp(b));
if query_pairs.is_empty() {
url.set_query(None);
} else {
let query = query_pairs
.iter()
.map(|(k, v)| format!("{}={}", query_escape(k), query_escape(v)))
.collect::<Vec<_>>()
.join("&");
url.set_query(Some(&query));
}
let filtered = url.to_string();
if url.path() == "/" && filtered.ends_with('/') {
return Ok(filtered.trim_end_matches('/').to_string());
}
Ok(filtered)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn should_escape_query() {
let test_cases = vec![
("a b", "a+b"),
("x*y", "x%2Ay"),
("c~d", "c~d"),
("1+1", "1%2B1"),
("δΈ", "%E4%B8%AD"),
];
for (s, expected) in test_cases {
assert_eq!(query_escape(s), expected);
}
}
#[test]
fn should_filter_query_params() {
let test_cases = vec![
(
"https://example.com/file.txt?z=9&b=2&a=1",
vec!["z".to_string()],
"https://example.com/file.txt?a=1&b=2",
),
(
"https://example.com/file.txt?b=2&a=1&b=1",
vec!["c".to_string()],
"https://example.com/file.txt?a=1&b=2&b=1",
),
(
"https://example.com?foo=foo",
vec!["foo".to_string()],
"https://example.com",
),
(
"https://example.com/file.txt?k=a b&m=x*y&n=c~d",
vec!["none".to_string()],
"https://example.com/file.txt?k=a+b&m=x%2Ay&n=c~d",
),
(
"http://www.xx.yy/path?u=f&x=y&m=z&x=s#size",
vec!["x".to_string(), "m".to_string()],
"http://www.xx.yy/path?u=f#size",
),
(
"http://www.xx.yy/path?u=f&x=y&m=z&x=s#size",
vec![],
"http://www.xx.yy/path?u=f&x=y&m=z&x=s#size",
),
(
"https://example.com/file.txt?a=1;x&b=2",
vec!["none".to_string()],
"https://example.com/file.txt?b=2",
),
];
for (url, filtered_query_params, expected) in test_cases {
assert_eq!(
filter_query_params(url, &filtered_query_params).unwrap(),
expected
);
}
assert!(filter_query_params(":error_url", &["x".to_string()]).is_err());
}
}