#[cfg(all(feature = "sessions", feature = "test-support"))]
mod tests {
use std::sync::Arc;
use std::time::Duration;
use bytes::Bytes;
use hyper::{HeaderMap, Method, Version};
use reinhardt_di::params::{ParamContext, ParamError, extract::FromRequest};
use reinhardt_di::{InjectionContext, SingletonScope};
use reinhardt_http::Request;
use reinhardt_middleware::session::{
OptionalSessionValue, OptionalSessionValueNamed, SessionData, SessionKey, SessionStore,
SessionValue, SessionValueNamed, USER_ID_SESSION_KEY, test_support::TenantIdKey,
};
fn build_request_with_session(
store: Arc<SessionStore>,
session_cookie: Option<&str>,
) -> Request {
let build = || -> Request {
let mut headers = HeaderMap::new();
if let Some(cookie) = session_cookie {
headers.insert(
hyper::header::COOKIE,
hyper::header::HeaderValue::from_str(&format!("sessionid={cookie}")).unwrap(),
);
}
Request::builder()
.method(Method::GET)
.uri("/test")
.version(Version::HTTP_11)
.headers(headers)
.body(Bytes::new())
.build()
.unwrap()
};
let request_for_ctx = build();
let request = build();
let singleton: Arc<SingletonScope> = Arc::new(SingletonScope::new());
singleton.set_arc(store);
let ctx = InjectionContext::builder(singleton).build();
ctx.set_request(request_for_ctx);
request.extensions.insert(Arc::new(ctx));
request
}
#[tokio::test]
async fn session_value_from_request_returns_stored_value() {
let store = Arc::new(SessionStore::new());
let mut session = SessionData::new(Duration::from_secs(3600));
session.set(USER_ID_SESSION_KEY.to_string(), 7i64).unwrap();
let session_id = session.id.clone();
store.save(session);
let request = build_request_with_session(Arc::clone(&store), Some(&session_id));
let ctx = ParamContext::new();
let SessionValue(user_id) =
<SessionValue<i64> as FromRequest>::from_request(&request, &ctx)
.await
.expect("SessionValue should extract when the user-id key is present");
assert_eq!(user_id, 7);
}
#[tokio::test]
async fn session_value_from_request_fails_authentication_when_session_missing() {
let store = Arc::new(SessionStore::new());
let request = build_request_with_session(store, None);
let ctx = ParamContext::new();
let err = <SessionValue<i64> as FromRequest>::from_request(&request, &ctx)
.await
.expect_err("SessionValue must fail when no session is present");
assert!(
matches!(err, ParamError::Authentication(_)),
"expected ParamError::Authentication, got {err:?}"
);
}
#[tokio::test]
async fn optional_session_value_from_request_yields_none_when_session_missing() {
let store = Arc::new(SessionStore::new());
let request = build_request_with_session(store, None);
let ctx = ParamContext::new();
let OptionalSessionValue(maybe_id) =
<OptionalSessionValue<i64> as FromRequest>::from_request(&request, &ctx)
.await
.expect("OptionalSessionValue must never fail extraction");
assert_eq!(maybe_id, None);
}
#[tokio::test]
async fn optional_session_value_from_request_returns_some_when_key_present() {
let store = Arc::new(SessionStore::new());
let mut session = SessionData::new(Duration::from_secs(3600));
session.set(USER_ID_SESSION_KEY.to_string(), 42i64).unwrap();
let session_id = session.id.clone();
store.save(session);
let request = build_request_with_session(Arc::clone(&store), Some(&session_id));
let ctx = ParamContext::new();
let OptionalSessionValue(maybe_id) =
<OptionalSessionValue<i64> as FromRequest>::from_request(&request, &ctx)
.await
.expect("OptionalSessionValue must succeed when the key is present");
assert_eq!(maybe_id, Some(42));
}
#[tokio::test]
async fn session_value_named_from_request_reads_custom_key() {
let store = Arc::new(SessionStore::new());
let mut session = SessionData::new(Duration::from_secs(3600));
session.set(TenantIdKey::KEY.to_string(), 9001i64).unwrap();
let session_id = session.id.clone();
store.save(session);
let request = build_request_with_session(Arc::clone(&store), Some(&session_id));
let ctx = ParamContext::new();
let extracted =
<SessionValueNamed<TenantIdKey, i64> as FromRequest>::from_request(&request, &ctx)
.await
.expect("SessionValueNamed should extract the configured key");
assert_eq!(*extracted, 9001);
}
#[tokio::test]
async fn session_value_named_from_request_fails_when_key_missing() {
let store = Arc::new(SessionStore::new());
let session = SessionData::new(Duration::from_secs(3600));
let session_id = session.id.clone();
store.save(session);
let request = build_request_with_session(Arc::clone(&store), Some(&session_id));
let ctx = ParamContext::new();
let err =
<SessionValueNamed<TenantIdKey, i64> as FromRequest>::from_request(&request, &ctx)
.await
.expect_err("SessionValueNamed must fail when the named key is absent");
assert!(
matches!(err, ParamError::Authentication(_)),
"expected ParamError::Authentication, got {err:?}"
);
}
#[tokio::test]
async fn optional_session_value_named_from_request_returns_some_when_key_present() {
let store = Arc::new(SessionStore::new());
let mut session = SessionData::new(Duration::from_secs(3600));
session.set(TenantIdKey::KEY.to_string(), 4242i64).unwrap();
let session_id = session.id.clone();
store.save(session);
let request = build_request_with_session(Arc::clone(&store), Some(&session_id));
let ctx = ParamContext::new();
let extracted = <OptionalSessionValueNamed<TenantIdKey, i64> as FromRequest>::from_request(
&request, &ctx,
)
.await
.expect("OptionalSessionValueNamed must succeed when the key is present");
assert_eq!(*extracted, Some(4242));
}
#[tokio::test]
async fn optional_session_value_named_from_request_yields_none_when_key_missing() {
let store = Arc::new(SessionStore::new());
let session = SessionData::new(Duration::from_secs(3600));
let session_id = session.id.clone();
store.save(session);
let request = build_request_with_session(Arc::clone(&store), Some(&session_id));
let ctx = ParamContext::new();
let extracted = <OptionalSessionValueNamed<TenantIdKey, i64> as FromRequest>::from_request(
&request, &ctx,
)
.await
.expect("OptionalSessionValueNamed must never fail extraction");
assert_eq!(*extracted, None);
}
#[tokio::test]
async fn optional_session_value_named_from_request_yields_none_when_session_missing() {
let store = Arc::new(SessionStore::new());
let request = build_request_with_session(store, None);
let ctx = ParamContext::new();
let extracted = <OptionalSessionValueNamed<TenantIdKey, i64> as FromRequest>::from_request(
&request, &ctx,
)
.await
.expect("OptionalSessionValueNamed must tolerate a missing session");
assert_eq!(*extracted, None);
}
#[tokio::test]
async fn optional_session_value_named_from_request_yields_none_on_deserialisation_mismatch() {
let store = Arc::new(SessionStore::new());
let mut session = SessionData::new(Duration::from_secs(3600));
session
.set(TenantIdKey::KEY.to_string(), "not-an-i64".to_string())
.unwrap();
let session_id = session.id.clone();
store.save(session);
let request = build_request_with_session(Arc::clone(&store), Some(&session_id));
let ctx = ParamContext::new();
let extracted = <OptionalSessionValueNamed<TenantIdKey, i64> as FromRequest>::from_request(
&request, &ctx,
)
.await
.expect("OptionalSessionValueNamed must absorb deserialisation failures");
assert_eq!(*extracted, None);
}
#[tokio::test]
async fn optional_session_value_named_from_request_yields_none_without_di_context() {
let request = Request::builder()
.method(Method::GET)
.uri("/test")
.version(Version::HTTP_11)
.headers(HeaderMap::new())
.body(Bytes::new())
.build()
.unwrap();
let ctx = ParamContext::new();
let extracted = <OptionalSessionValueNamed<TenantIdKey, i64> as FromRequest>::from_request(
&request, &ctx,
)
.await
.expect("OptionalSessionValueNamed must tolerate a missing DI context");
assert_eq!(*extracted, None);
}
#[tokio::test]
async fn optional_session_value_from_request_yields_none_without_di_context() {
let request = Request::builder()
.method(Method::GET)
.uri("/test")
.version(Version::HTTP_11)
.headers(HeaderMap::new())
.body(Bytes::new())
.build()
.unwrap();
let ctx = ParamContext::new();
let OptionalSessionValue(maybe_id) =
<OptionalSessionValue<i64> as FromRequest>::from_request(&request, &ctx)
.await
.expect("OptionalSessionValue must tolerate a missing DI context");
assert_eq!(maybe_id, None);
}
#[tokio::test]
async fn session_value_from_request_fails_internal_without_di_context() {
let request = Request::builder()
.method(Method::GET)
.uri("/test")
.version(Version::HTTP_11)
.headers(HeaderMap::new())
.body(Bytes::new())
.build()
.unwrap();
let ctx = ParamContext::new();
let err = <SessionValue<i64> as FromRequest>::from_request(&request, &ctx)
.await
.expect_err("SessionValue must fail without a DI context");
assert!(
matches!(err, ParamError::Internal(_)),
"expected ParamError::Internal, got {err:?}"
);
}
}