use bytesize::ByteSize;
use dragonfly_api::common::v2::{Priority, SchedulingPolicy};
use reqwest::header::HeaderMap;
use std::{fmt, str::FromStr};
use tracing::error;
pub const DRAGONFLY_TAG_HEADER: &str = "X-Dragonfly-Tag";
pub const DRAGONFLY_APPLICATION_HEADER: &str = "X-Dragonfly-Application";
pub const DRAGONFLY_PRIORITY_HEADER: &str = "X-Dragonfly-Priority";
pub const DRAGONFLY_REGISTRY_HEADER: &str = "X-Dragonfly-Registry";
pub const DRAGONFLY_FILTERED_QUERY_PARAMS_HEADER: &str = "X-Dragonfly-Filtered-Query-Params";
pub const DRAGONFLY_USE_P2P_HEADER: &str = "X-Dragonfly-Use-P2P";
pub const DRAGONFLY_PREFETCH_HEADER: &str = "X-Dragonfly-Prefetch";
pub const DRAGONFLY_OUTPUT_PATH_HEADER: &str = "X-Dragonfly-Output-Path";
pub const DRAGONFLY_FORCE_HARD_LINK_HEADER: &str = "X-Dragonfly-Force-Hard-Link";
pub const DRAGONFLY_PIECE_LENGTH_HEADER: &str = "X-Dragonfly-Piece-Length";
pub const DRAGONFLY_CONTENT_FOR_CALCULATING_TASK_ID_HEADER: &str =
"X-Dragonfly-Content-For-Calculating-Task-ID";
pub const DRAGONFLY_ENABLE_TASK_ID_BASED_BLOB_DIGEST: &str =
"X-Dragonfly-Enable-Task-ID-Based-Blob-Digest";
pub const DRAGONFLY_SCHEDULING_POLICY_HEADER: &str = "X-Dragonfly-Scheduling-Policy";
pub const DRAGONFLY_TASK_DOWNLOAD_FINISHED_HEADER: &str = "X-Dragonfly-Task-Download-Finished";
pub const DRAGONFLY_TASK_ID_HEADER: &str = "X-Dragonfly-Task-ID";
pub const DRAGONFLY_SERVER_IP_HEADER: &str = "X-Dragonfly-Server-IP";
pub const DRAGONFLY_ERROR_TYPE_HEADER: &str = "X-Dragonfly-Error-Type";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ErrorType {
Backend,
Proxy,
Dfdaemon,
}
impl ErrorType {
pub fn as_str(&self) -> &'static str {
match self {
ErrorType::Backend => "backend",
ErrorType::Proxy => "proxy",
ErrorType::Dfdaemon => "dfdaemon",
}
}
}
impl fmt::Display for ErrorType {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.as_str())
}
}
impl FromStr for ErrorType {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s {
"backend" => Ok(ErrorType::Backend),
"proxy" => Ok(ErrorType::Proxy),
"dfdaemon" => Ok(ErrorType::Dfdaemon),
_ => Err(format!("invalid error type: {s}")),
}
}
}
pub fn get_tag(header: &HeaderMap) -> Option<String> {
header
.get(DRAGONFLY_TAG_HEADER)
.and_then(|tag| tag.to_str().ok())
.map(|tag| tag.to_string())
}
pub fn get_application(header: &HeaderMap) -> Option<String> {
header
.get(DRAGONFLY_APPLICATION_HEADER)
.and_then(|application| application.to_str().ok())
.map(|application| application.to_string())
}
pub fn get_priority(header: &HeaderMap) -> i32 {
let default_priority = Priority::Level6 as i32;
match header.get(DRAGONFLY_PRIORITY_HEADER) {
Some(priority) => match priority.to_str() {
Ok(priority) => match priority.parse::<i32>() {
Ok(priority) => priority,
Err(err) => {
error!("parse priority from header failed: {}", err);
default_priority
}
},
Err(err) => {
error!("get priority from header failed: {}", err);
default_priority
}
},
None => default_priority,
}
}
pub fn get_registry(header: &HeaderMap) -> Option<String> {
header
.get(DRAGONFLY_REGISTRY_HEADER)
.and_then(|registry| registry.to_str().ok())
.map(|registry| registry.to_string())
}
pub fn get_filtered_query_params(
header: &HeaderMap,
default_filtered_query_params: &[String],
) -> Vec<String> {
match header.get(DRAGONFLY_FILTERED_QUERY_PARAMS_HEADER) {
Some(filters) => match filters.to_str() {
Ok(filters) => filters.split(',').map(|s| s.trim().to_string()).collect(),
Err(err) => {
error!("get filters from header failed: {}", err);
default_filtered_query_params.to_vec()
}
},
None => default_filtered_query_params.to_vec(),
}
}
pub fn get_use_p2p(header: &HeaderMap) -> bool {
match header.get(DRAGONFLY_USE_P2P_HEADER) {
Some(value) => match value.to_str() {
Ok(value) => value.eq_ignore_ascii_case("true"),
Err(err) => {
error!("get use p2p from header failed: {}", err);
false
}
},
None => false,
}
}
pub fn get_prefetch(header: &HeaderMap) -> Option<bool> {
match header.get(DRAGONFLY_PREFETCH_HEADER) {
Some(value) => match value.to_str() {
Ok(value) => Some(value.eq_ignore_ascii_case("true")),
Err(err) => {
error!("get use p2p from header failed: {}", err);
None
}
},
None => None,
}
}
pub fn get_output_path(header: &HeaderMap) -> Option<String> {
header
.get(DRAGONFLY_OUTPUT_PATH_HEADER)
.and_then(|output_path| output_path.to_str().ok())
.map(|output_path| output_path.to_string())
}
pub fn get_force_hard_link(header: &HeaderMap) -> bool {
match header.get(DRAGONFLY_FORCE_HARD_LINK_HEADER) {
Some(value) => match value.to_str() {
Ok(value) => value.eq_ignore_ascii_case("true"),
Err(err) => {
error!("get force hard link from header failed: {}", err);
false
}
},
None => false,
}
}
pub fn get_piece_length(header: &HeaderMap) -> Option<ByteSize> {
match header.get(DRAGONFLY_PIECE_LENGTH_HEADER) {
Some(piece_length) => match piece_length.to_str() {
Ok(piece_length) => match piece_length.parse::<ByteSize>() {
Ok(piece_length) => Some(piece_length),
Err(err) => {
error!("parse piece length from header failed: {}", err);
None
}
},
Err(err) => {
error!("get piece length from header failed: {}", err);
None
}
},
None => None,
}
}
pub fn get_content_for_calculating_task_id(header: &HeaderMap) -> Option<String> {
header
.get(DRAGONFLY_CONTENT_FOR_CALCULATING_TASK_ID_HEADER)
.and_then(|content| content.to_str().ok())
.map(|content| content.to_string())
}
pub fn get_enable_task_id_based_blob_digest(header: &HeaderMap, default: bool) -> bool {
match header.get(DRAGONFLY_ENABLE_TASK_ID_BASED_BLOB_DIGEST) {
Some(value) => match value.to_str() {
Ok(value) => value.eq_ignore_ascii_case("true"),
Err(err) => {
error!(
"get enable task id based blob digest from header failed: {}",
err
);
default
}
},
None => default,
}
}
pub fn get_scheduling_policy(header: &HeaderMap, default: SchedulingPolicy) -> SchedulingPolicy {
match header.get(DRAGONFLY_SCHEDULING_POLICY_HEADER) {
Some(value) => match value.to_str() {
Ok(value) if value.eq_ignore_ascii_case("auto") => SchedulingPolicy::Auto,
Ok(value) if value.eq_ignore_ascii_case("always") => SchedulingPolicy::Always,
Ok(value) => {
error!("invalid scheduling policy from header: {}", value);
default
}
Err(err) => {
error!("get scheduling policy from header failed: {}", err);
default
}
},
None => default,
}
}
#[cfg(test)]
mod tests {
#![allow(clippy::type_complexity)]
use super::*;
use reqwest::header::HeaderValue;
#[test]
fn error_type_parses_and_formats_its_name() {
let test_cases = vec![
("backend", Ok(ErrorType::Backend), Some("backend")),
("proxy", Ok(ErrorType::Proxy), Some("proxy")),
("dfdaemon", Ok(ErrorType::Dfdaemon), Some("dfdaemon")),
(
"Backend",
Err("invalid error type: Backend".to_string()),
None,
),
];
for (name, expected, expected_display) in test_cases {
let error_type = name.parse::<ErrorType>();
assert_eq!(error_type, expected);
assert_eq!(
error_type.ok().map(|error_type| error_type.to_string()),
expected_display.map(str::to_string)
);
}
}
#[test]
fn string_getters_return_the_header_value() {
let test_cases: Vec<(&'static str, fn(&HeaderMap) -> Option<String>)> = vec![
(DRAGONFLY_TAG_HEADER, get_tag),
(DRAGONFLY_APPLICATION_HEADER, get_application),
(DRAGONFLY_REGISTRY_HEADER, get_registry),
(DRAGONFLY_OUTPUT_PATH_HEADER, get_output_path),
(
DRAGONFLY_CONTENT_FOR_CALCULATING_TASK_ID_HEADER,
get_content_for_calculating_task_id,
),
];
for (name, getter) in test_cases {
let mut headers = HeaderMap::new();
headers.insert(name, HeaderValue::from_str("value").unwrap());
assert_eq!(getter(&headers), Some("value".to_string()));
assert_eq!(getter(&HeaderMap::new()), None);
headers.insert(name, HeaderValue::from_str("é").unwrap());
assert_eq!(getter(&headers), None);
}
}
#[test]
fn get_priority_parses_the_header_or_falls_back_to_level6() {
let test_cases = vec![
(Some("5"), 5),
(Some("invalid"), Priority::Level6 as i32),
(Some("é"), Priority::Level6 as i32),
(None, Priority::Level6 as i32),
];
for (value, expected) in test_cases {
let mut headers = HeaderMap::new();
if let Some(value) = value {
headers.insert(
DRAGONFLY_PRIORITY_HEADER,
HeaderValue::from_str(value).unwrap(),
);
}
assert_eq!(get_priority(&headers), expected);
}
}
#[test]
fn get_filtered_query_params_splits_the_header_or_uses_defaults() {
let default_filtered_query_params = vec!["default".to_string()];
let test_cases = vec![
(Some("param1,param2"), vec!["param1", "param2"]),
(
Some("param1, param2 ,param3"),
vec!["param1", "param2", "param3"],
),
(Some("é"), vec!["default"]),
(None, vec!["default"]),
];
for (value, expected) in test_cases {
let expected: Vec<String> = expected.iter().map(|param| param.to_string()).collect();
let mut headers = HeaderMap::new();
if let Some(value) = value {
headers.insert(
DRAGONFLY_FILTERED_QUERY_PARAMS_HEADER,
HeaderValue::from_str(value).unwrap(),
);
}
assert_eq!(
get_filtered_query_params(&headers, &default_filtered_query_params),
expected
);
}
}
#[test]
fn bool_getters_are_true_only_for_a_true_header() {
let test_cases: Vec<(&'static str, fn(&HeaderMap) -> bool)> = vec![
(DRAGONFLY_USE_P2P_HEADER, get_use_p2p),
(DRAGONFLY_FORCE_HARD_LINK_HEADER, get_force_hard_link),
];
for (name, getter) in test_cases {
let mut headers = HeaderMap::new();
headers.insert(name, HeaderValue::from_str("true").unwrap());
assert!(getter(&headers));
headers.insert(name, HeaderValue::from_str("TRUE").unwrap());
assert!(getter(&headers));
headers.insert(name, HeaderValue::from_str("false").unwrap());
assert!(!getter(&headers));
headers.insert(name, HeaderValue::from_str("é").unwrap());
assert!(!getter(&headers));
assert!(!getter(&HeaderMap::new()));
}
}
#[test]
fn get_prefetch_returns_the_flag_or_none() {
let test_cases = vec![
(Some("true"), Some(true)),
(Some("false"), Some(false)),
(Some("é"), None),
(None, None),
];
for (value, expected) in test_cases {
let mut headers = HeaderMap::new();
if let Some(value) = value {
headers.insert(
DRAGONFLY_PREFETCH_HEADER,
HeaderValue::from_str(value).unwrap(),
);
}
assert_eq!(get_prefetch(&headers), expected);
}
}
#[test]
fn get_piece_length_parses_human_readable_sizes() {
let test_cases = vec![
(Some("4mib"), Some(ByteSize::mib(4))),
(Some("0"), Some(ByteSize::b(0))),
(Some("invalid"), None),
(Some("é"), None),
(None, None),
];
for (value, expected) in test_cases {
let mut headers = HeaderMap::new();
if let Some(value) = value {
headers.insert(
DRAGONFLY_PIECE_LENGTH_HEADER,
HeaderValue::from_str(value).unwrap(),
);
}
assert_eq!(get_piece_length(&headers), expected);
}
}
#[test]
fn get_enable_task_id_based_blob_digest_falls_back_to_default() {
let test_cases = vec![
(Some("true"), false, true),
(Some("false"), true, false),
(Some("é"), true, true),
(None, true, true),
(None, false, false),
];
for (value, default, expected) in test_cases {
let mut headers = HeaderMap::new();
if let Some(value) = value {
headers.insert(
DRAGONFLY_ENABLE_TASK_ID_BASED_BLOB_DIGEST,
HeaderValue::from_str(value).unwrap(),
);
}
assert_eq!(
get_enable_task_id_based_blob_digest(&headers, default),
expected
);
}
}
#[test]
fn get_scheduling_policy_parses_case_insensitively_or_uses_default() {
let test_cases = vec![
(
Some("always"),
SchedulingPolicy::Auto,
SchedulingPolicy::Always,
),
(
Some("AUTO"),
SchedulingPolicy::Always,
SchedulingPolicy::Auto,
),
(
Some("invalid"),
SchedulingPolicy::Always,
SchedulingPolicy::Always,
),
(Some("é"), SchedulingPolicy::Auto, SchedulingPolicy::Auto),
(None, SchedulingPolicy::Always, SchedulingPolicy::Always),
(None, SchedulingPolicy::Auto, SchedulingPolicy::Auto),
];
for (value, default, expected) in test_cases {
let mut headers = HeaderMap::new();
if let Some(value) = value {
headers.insert(
DRAGONFLY_SCHEDULING_POLICY_HEADER,
HeaderValue::from_str(value).unwrap(),
);
}
assert_eq!(get_scheduling_policy(&headers, default), expected);
}
}
}