use std::cmp::Ordering::Equal;
use std::str::FromStr;
use anyhow::Result;
use http::header::{ACCEPT, CONTENT_TYPE};
use http::{HeaderMap, HeaderValue};
use mime::{APPLICATION_JSON, APPLICATION_OCTET_STREAM, Mime, Name, TEXT_PLAIN};
use crate::api::err::ApiError;
use crate::api::format as api_format;
use crate::api::middleware::common::{
APPLICATION_CBOR, APPLICATION_SDB_FB, APPLICATION_SDB_NATIVE, BodyStrategy,
};
use crate::api::response::ApiResponse;
use crate::rpc::format;
use crate::sql::expression::convert_public_value_to_internal;
use crate::types::{PublicBytes, PublicValue};
use crate::val::Bytes;
pub fn output_body_strategy(headers: &HeaderMap, strategy: BodyStrategy) -> Option<BodyStrategy> {
let Some(accepted) = headers.get(ACCEPT) else {
return Some(strategy);
};
let accepted = parse_accept(accepted);
if accepted.is_empty() {
return None;
}
let supported: &[_] = match strategy {
BodyStrategy::Json => &[(BodyStrategy::Json, &APPLICATION_JSON)],
BodyStrategy::Cbor => &[(BodyStrategy::Cbor, &*APPLICATION_CBOR)],
BodyStrategy::Flatbuffers => &[(BodyStrategy::Flatbuffers, &*APPLICATION_SDB_FB)],
BodyStrategy::Plain => &[(BodyStrategy::Plain, &TEXT_PLAIN)],
BodyStrategy::Bytes => &[(BodyStrategy::Bytes, &APPLICATION_OCTET_STREAM)],
BodyStrategy::Native => &[(BodyStrategy::Native, &*APPLICATION_SDB_NATIVE)],
BodyStrategy::Auto => &[
(BodyStrategy::Json, &APPLICATION_JSON),
(BodyStrategy::Cbor, &*APPLICATION_CBOR),
(BodyStrategy::Flatbuffers, &*APPLICATION_SDB_FB),
(BodyStrategy::Plain, &TEXT_PLAIN),
(BodyStrategy::Bytes, &APPLICATION_OCTET_STREAM),
(BodyStrategy::Native, &*APPLICATION_SDB_NATIVE),
],
};
for range in accepted.iter() {
for (strategy, mime) in supported.iter() {
if range.matches(mime) {
return Some(*strategy);
}
}
}
None
}
pub fn convert_response_value(response: &mut ApiResponse, strategy: BodyStrategy) -> Result<()> {
match strategy {
BodyStrategy::Auto | BodyStrategy::Json => {
response.body = PublicValue::Bytes(PublicBytes::from(
format::json::encode(response.body.clone())
.map_err(|_| ApiError::BodyEncodeFailure)?,
));
response.headers.insert(CONTENT_TYPE, api_format::JSON.try_into()?);
}
BodyStrategy::Cbor => {
response.body = PublicValue::Bytes(PublicBytes::from(
format::cbor::encode(response.body.clone())
.map_err(|_| ApiError::BodyEncodeFailure)?,
));
response.headers.insert(CONTENT_TYPE, api_format::CBOR.try_into()?);
}
BodyStrategy::Flatbuffers => {
response.body = PublicValue::Bytes(PublicBytes::from(
format::flatbuffers::encode(&response.body)
.map_err(|_| ApiError::BodyEncodeFailure)?,
));
response.headers.insert(CONTENT_TYPE, api_format::FLATBUFFERS.try_into()?);
}
BodyStrategy::Bytes => {
let bytes = convert_public_value_to_internal(response.body.clone())
.cast_to::<Bytes>()
.map_err(|_| ApiError::BodyEncodeFailure)?
.0;
response.body = PublicValue::Bytes(PublicBytes::from(bytes));
response.headers.insert(CONTENT_TYPE, api_format::OCTET_STREAM.try_into()?);
}
BodyStrategy::Plain => {
let text = convert_public_value_to_internal(response.body.clone())
.cast_to::<String>()
.map_err(|_| ApiError::BodyEncodeFailure)?;
response.body = PublicValue::Bytes(PublicBytes::from(text.into_bytes()));
response.headers.insert(CONTENT_TYPE, api_format::PLAIN.try_into()?);
}
BodyStrategy::Native => {
response.headers.insert(CONTENT_TYPE, api_format::NATIVE.try_into()?);
}
}
Ok(())
}
fn parse_accept(value: &HeaderValue) -> Vec<AcceptRange> {
let s = value.to_str().unwrap_or("").trim();
if s.is_empty() {
return Vec::new();
}
let mut accepted: Vec<(f32, AcceptRange, usize)> = Vec::new();
for (i, part) in s.split(',').enumerate() {
let part = part.trim();
if part.is_empty() {
continue;
}
let Ok(mime) = Mime::from_str(part) else {
continue;
};
let q = mime
.get_param("q")
.and_then(|x: Name<'_>| x.as_str().parse::<f32>().ok())
.unwrap_or(1.0);
if q <= 0.0 {
continue;
}
accepted.push((q, mime.into(), i));
}
accepted.sort_by(|a, b| {
b.0.partial_cmp(&a.0)
.unwrap_or(Equal)
.then_with(|| a.1.partial_cmp(&b.1).unwrap_or(Equal))
.then_with(|| a.2.partial_cmp(&b.2).unwrap_or(Equal))
});
accepted.into_iter().map(|x| x.1).collect()
}
#[derive(Clone, Debug)]
enum AcceptRange {
Exact(Mime),
TypeWildcard(String),
Any,
}
impl From<Mime> for AcceptRange {
fn from(mime: Mime) -> Self {
if mime.subtype() == mime::STAR {
match mime.type_() {
mime::STAR => Self::Any,
x => Self::TypeWildcard(x.to_string()),
}
} else {
Self::Exact(mime)
}
}
}
impl PartialEq for AcceptRange {
fn eq(&self, other: &Self) -> bool {
match (self, other) {
(Self::Exact(a), Self::Exact(b)) => a.eq(b),
(Self::TypeWildcard(a), Self::TypeWildcard(b)) => a.eq(b),
(Self::Any, Self::Any) => true,
_ => false,
}
}
}
impl PartialOrd for AcceptRange {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
self.specifity().partial_cmp(&other.specifity())
}
}
impl AcceptRange {
pub(super) fn specifity(&self) -> u8 {
match self {
Self::Exact(_) => 0,
Self::TypeWildcard(_) => 1,
Self::Any => 2,
}
}
pub(super) fn matches(&self, mime: &Mime) -> bool {
match self {
Self::Exact(x) => x.type_() == mime.type_() && x.subtype() == mime.subtype(),
Self::TypeWildcard(x) => mime.type_().as_str().eq_ignore_ascii_case(x),
Self::Any => true,
}
}
}
#[cfg(test)]
mod tests {
use http::HeaderMap;
use http::header::ACCEPT;
use crate::api::middleware::common::BodyStrategy;
use crate::api::middleware::res::output_body_strategy;
macro_rules! case {
($in:ident => None, $header:expr) => {{
let mut headers = HeaderMap::new();
headers.insert(ACCEPT, $header.parse().unwrap());
let out = output_body_strategy(&headers, BodyStrategy::$in);
assert!(out.is_none());
}};
($in:ident => $out:ident, $header:expr) => {{
let mut headers = HeaderMap::new();
headers.insert(ACCEPT, $header.parse().unwrap());
let out = output_body_strategy(&headers, BodyStrategy::$in).unwrap();
assert_eq!(out, BodyStrategy::$out);
}};
}
#[test]
fn tests() {
case!(Auto => Json, "application/*;q=0.9, application/json;q=0.9");
case!(Auto => Plain, "text/plain, application/json");
case!(Auto => Json, "*/*;q=1.0, application/json;q=1.0");
case!(Auto => Plain, "*/*;q=0.9, text/*;q=0.9");
case!(Auto => Json, "*/*");
case!(Bytes => Bytes, "application/octet-stream, */*;q=0.1");
case!(Plain => Plain, "text/*;q=0.2, application/*;q=0.9");
case!(Auto => Json, "text/plain;q=0, application/json;q=0.1");
case!(Auto => Json, "application/json; charset=utf-8");
case!(Auto => Plain, "TeXt/*");
case!(Auto => Json, "application/json;q=0.2, application/json;q=0.8");
case!(Auto => Json, "application/json, application/cbor");
case!(Auto => None, "");
case!(Auto => None, "not/a-mime, also-bad");
case!(Auto => None, "application/json;q=0, */*;q=0");
}
}