use async_trait::async_trait;
use reinhardt_http::Request;
use std::fmt::{self, Debug};
use std::marker::PhantomData;
use std::ops::Deref;
use super::{ParamContext, ParamError, ParamResult, extract::FromRequest};
pub trait CookieName {
const NAME: &'static str;
}
pub struct SessionId;
impl CookieName for SessionId {
const NAME: &'static str = "sessionid";
}
pub struct CsrfToken;
impl CookieName for CsrfToken {
const NAME: &'static str = "csrftoken";
}
pub struct CookieNamed<N: CookieName, T> {
value: T,
_phantom: PhantomData<N>,
}
impl<N: CookieName, T> CookieNamed<N, T> {
pub fn into_inner(self) -> T {
self.value
}
pub fn new(value: T) -> Self {
CookieNamed {
value,
_phantom: PhantomData,
}
}
pub const fn name() -> &'static str {
N::NAME
}
}
impl<N: CookieName, T> Deref for CookieNamed<N, T> {
type Target = T;
fn deref(&self) -> &Self::Target {
&self.value
}
}
impl<N: CookieName, T: Debug> Debug for CookieNamed<N, T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("CookieNamed")
.field("name", &N::NAME)
.field("value", &self.value)
.finish()
}
}
use super::cookie_util::parse_cookies;
#[async_trait]
impl<N> FromRequest for CookieNamed<N, String>
where
N: CookieName + Send,
{
async fn from_request(req: &Request, _ctx: &ParamContext) -> ParamResult<Self> {
let cookie_header = req
.headers
.get(http::header::COOKIE)
.and_then(|h| h.to_str().ok())
.unwrap_or("");
let cookies_map = parse_cookies(cookie_header);
let value = cookies_map
.get(N::NAME)
.ok_or_else(|| ParamError::MissingParameter(N::NAME.to_owned()))?;
Ok(CookieNamed::new(value.clone()))
}
}
#[async_trait]
impl<N> FromRequest for CookieNamed<N, Option<String>>
where
N: CookieName + Send,
{
async fn from_request(req: &Request, _ctx: &ParamContext) -> ParamResult<Self> {
let cookie_header = req
.headers
.get(http::header::COOKIE)
.and_then(|h| h.to_str().ok())
.unwrap_or("");
let cookies_map = parse_cookies(cookie_header);
Ok(CookieNamed::new(cookies_map.get(N::NAME).cloned()))
}
}
#[cfg(feature = "validation")]
impl<N: CookieName, T> super::validation::WithValidation for CookieNamed<N, T> {}
#[cfg(test)]
mod tests {
use super::*;
use crate::params::extract::FromRequest;
struct OAuthState;
impl CookieName for OAuthState {
const NAME: &'static str = "oauth_state";
}
#[test]
fn test_cookie_named_new() {
let cookie = CookieNamed::<SessionId, String>::new("12345".to_string());
assert_eq!(*cookie, "12345");
assert_eq!(CookieNamed::<SessionId, String>::name(), "sessionid");
}
#[test]
fn test_cookie_named_into_inner() {
let cookie = CookieNamed::<CsrfToken, String>::new("abc-def-ghi".to_string());
let value = cookie.into_inner();
assert_eq!(value, "abc-def-ghi");
}
#[test]
fn test_cookie_named_deref() {
let cookie = CookieNamed::<SessionId, String>::new("session123".to_string());
assert_eq!(&*cookie, "session123");
}
#[test]
fn test_cookie_named_optional() {
let cookie1 = CookieNamed::<CsrfToken, Option<String>>::new(Some("dark".to_string()));
assert_eq!(*cookie1, Some("dark".to_string()));
let cookie2 = CookieNamed::<CsrfToken, Option<String>>::new(None);
assert_eq!(*cookie2, None);
}
#[test]
fn test_parse_cookies() {
let cookies = parse_cookies("sessionid=abc123; csrftoken=xyz789; user=john");
assert_eq!(cookies.get("sessionid"), Some(&"abc123".to_string()));
assert_eq!(cookies.get("csrftoken"), Some(&"xyz789".to_string()));
assert_eq!(cookies.get("user"), Some(&"john".to_string()));
}
#[test]
fn test_parse_cookies_with_encoding() {
let cookies = parse_cookies("name=value%20with%20spaces");
assert_eq!(cookies.get("name"), Some(&"value with spaces".to_string()));
}
#[tokio::test]
async fn test_cookie_named_extracts_custom_required_cookie() {
let request = Request::builder()
.uri("/")
.header("Cookie", "sessionid=abc123; oauth_state=state-456")
.build()
.expect("request should build");
let ctx = ParamContext::new();
let cookie = CookieNamed::<OAuthState, String>::from_request(&request, &ctx)
.await
.expect("custom named cookie should be extracted");
assert_eq!(cookie.into_inner(), "state-456");
}
#[tokio::test]
async fn test_cookie_named_reports_custom_required_cookie_name_when_missing() {
let request = Request::builder()
.uri("/")
.header("Cookie", "sessionid=abc123")
.build()
.expect("request should build");
let ctx = ParamContext::new();
let error = CookieNamed::<OAuthState, String>::from_request(&request, &ctx)
.await
.expect_err("missing custom named cookie should fail");
assert!(matches!(
error,
ParamError::MissingParameter(name) if name == "oauth_state"
));
}
#[tokio::test]
async fn test_cookie_named_extracts_custom_optional_cookie() {
let request = Request::builder()
.uri("/")
.header("Cookie", "oauth_state=value%20with%20spaces")
.build()
.expect("request should build");
let ctx = ParamContext::new();
let cookie = CookieNamed::<OAuthState, Option<String>>::from_request(&request, &ctx)
.await
.expect("optional custom named cookie should extract");
assert_eq!(cookie.into_inner(), Some("value with spaces".to_string()));
}
#[tokio::test]
async fn test_cookie_named_custom_optional_cookie_is_none_when_missing() {
let request = Request::builder()
.uri("/")
.header("Cookie", "sessionid=abc123")
.build()
.expect("request should build");
let ctx = ParamContext::new();
let cookie = CookieNamed::<OAuthState, Option<String>>::from_request(&request, &ctx)
.await
.expect("missing optional custom named cookie should not fail");
assert_eq!(cookie.into_inner(), None);
}
}