use crate::{Error, HttpRequest};
use bytes::Bytes;
use serde::de::DeserializeOwned;
use std::ops::Deref;
use std::sync::Arc;
pub trait FromRequest: Sized {
fn from_request(request: &HttpRequest) -> Result<Self, Error>;
}
#[derive(Debug)]
pub struct State<T: Send + Sync + 'static>(pub Arc<T>);
impl<T: Send + Sync + 'static> State<T> {
#[inline]
pub fn new(value: Arc<T>) -> Self {
Self(value)
}
#[inline]
pub fn into_inner(self) -> Arc<T> {
self.0
}
}
impl<T: Send + Sync + 'static> Clone for State<T> {
#[inline]
fn clone(&self) -> Self {
Self(Arc::clone(&self.0))
}
}
impl<T: Send + Sync + 'static> Deref for State<T> {
type Target = T;
#[inline]
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl<T: Send + Sync + 'static> AsRef<T> for State<T> {
#[inline]
fn as_ref(&self) -> &T {
&self.0
}
}
impl<T: Send + Sync + 'static> FromRequest for State<T> {
#[inline]
fn from_request(request: &HttpRequest) -> Result<Self, Error> {
request.extensions.get_arc::<T>().map(State).ok_or_else(|| {
Error::ProviderNotFound(format!(
"State<{}> not found in request extensions. \
Did you forget to register it with `app.with_state()`?",
std::any::type_name::<T>()
))
})
}
}
pub trait FromRequestNamed: Sized {
fn from_request(request: &HttpRequest, name: &str) -> Result<Self, Error>;
}
#[derive(Debug, Clone)]
pub struct Body<T>(pub T);
impl<T> Body<T> {
pub fn new(value: T) -> Self {
Self(value)
}
pub fn into_inner(self) -> T {
self.0
}
}
impl<T> Deref for Body<T> {
type Target = T;
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl<T: DeserializeOwned> FromRequest for Body<T> {
fn from_request(request: &HttpRequest) -> Result<Self, Error> {
let value: T = request.json()?;
Ok(Body(value))
}
}
#[derive(Debug, Clone)]
pub struct Query<T>(pub T);
impl<T> Query<T> {
pub fn new(value: T) -> Self {
Self(value)
}
pub fn into_inner(self) -> T {
self.0
}
}
impl<T> Deref for Query<T> {
type Target = T;
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl<T: DeserializeOwned> FromRequest for Query<T> {
fn from_request(request: &HttpRequest) -> Result<Self, Error> {
let value: T = serde_urlencoded::from_str(request.query_string().unwrap_or(""))
.map_err(|e| Error::Validation(format!("Invalid query parameters: {}", e)))?;
Ok(Query(value))
}
}
#[derive(Debug, Clone)]
pub struct Path<T>(pub T);
impl<T> Path<T> {
pub fn new(value: T) -> Self {
Self(value)
}
pub fn into_inner(self) -> T {
self.0
}
}
impl<T> Deref for Path<T> {
type Target = T;
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl<T: std::str::FromStr> FromRequestNamed for Path<T>
where
T::Err: std::fmt::Display,
{
fn from_request(request: &HttpRequest, name: &str) -> Result<Self, Error> {
let value_str = request
.param(name)
.ok_or_else(|| Error::Validation(format!("Missing path parameter: {}", name)))?;
let value: T = value_str.parse().map_err(|e: T::Err| {
Error::Validation(format!("Invalid path parameter '{}': {}", name, e))
})?;
Ok(Path(value))
}
}
#[derive(Debug, Clone)]
pub struct PathParams<T>(pub T);
impl<T> PathParams<T> {
pub fn new(value: T) -> Self {
Self(value)
}
pub fn into_inner(self) -> T {
self.0
}
}
impl<T> Deref for PathParams<T> {
type Target = T;
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl<T: DeserializeOwned> FromRequest for PathParams<T> {
fn from_request(request: &HttpRequest) -> Result<Self, Error> {
let pairs: Vec<(&str, &str)> = request
.path_params
.iter()
.filter_map(|(k, v)| std::str::from_utf8(v).ok().map(|v| (*k, v)))
.collect();
let params_string = serde_urlencoded::to_string(&pairs)
.map_err(|e| Error::Validation(format!("Invalid path parameters: {}", e)))?;
let value: T = serde_urlencoded::from_str(¶ms_string)
.map_err(|e| Error::Validation(format!("Invalid path parameters: {}", e)))?;
Ok(PathParams(value))
}
}
#[derive(Debug, Clone)]
pub struct Header {
name: String,
value: String,
}
impl Header {
pub fn new(name: impl Into<String>, value: impl Into<String>) -> Self {
Self {
name: name.into(),
value: value.into(),
}
}
pub fn name(&self) -> &str {
&self.name
}
pub fn value(&self) -> &str {
&self.value
}
pub fn into_value(self) -> String {
self.value
}
pub fn optional(request: &HttpRequest, name: &str) -> Option<Self> {
request.headers.get(name).map(|v| Header::new(name, v))
}
}
impl FromRequestNamed for Header {
fn from_request(request: &HttpRequest, name: &str) -> Result<Self, Error> {
let value = request
.headers
.get(name)
.ok_or_else(|| Error::Validation(format!("Missing header: {}", name)))?;
Ok(Header::new(name, value))
}
}
impl Deref for Header {
type Target = str;
fn deref(&self) -> &Self::Target {
&self.value
}
}
#[derive(Debug, Clone)]
pub struct Headers(pub std::collections::HashMap<String, String>);
impl Headers {
pub fn get(&self, name: &str) -> Option<&String> {
if let Some(value) = self.0.get(name) {
return Some(value);
}
self.0
.iter()
.find(|(k, _)| k.eq_ignore_ascii_case(name))
.map(|(_, v)| v)
}
pub fn contains(&self, name: &str) -> bool {
self.get(name).is_some()
}
pub fn iter(&self) -> impl Iterator<Item = (&String, &String)> {
self.0.iter()
}
}
impl FromRequest for Headers {
fn from_request(request: &HttpRequest) -> Result<Self, Error> {
Ok(Headers(request.headers.clone().into()))
}
}
impl Deref for Headers {
type Target = std::collections::HashMap<String, String>;
fn deref(&self) -> &Self::Target {
&self.0
}
}
#[derive(Debug, Clone)]
pub struct RawBody(pub Bytes);
impl RawBody {
pub fn new(data: impl Into<Bytes>) -> Self {
Self(data.into())
}
pub fn len(&self) -> usize {
self.0.len()
}
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
pub fn to_string_lossy(&self) -> String {
String::from_utf8_lossy(&self.0).to_string()
}
pub fn to_string(&self) -> Result<String, std::string::FromUtf8Error> {
String::from_utf8(self.0.to_vec())
}
pub fn into_inner(self) -> Bytes {
self.0
}
}
impl FromRequest for RawBody {
fn from_request(request: &HttpRequest) -> Result<Self, Error> {
Ok(RawBody(request.body.clone()))
}
}
impl Deref for RawBody {
type Target = [u8];
fn deref(&self) -> &Self::Target {
&self.0
}
}
#[derive(Debug, Clone)]
pub struct Form<T>(pub T);
impl<T> Form<T> {
pub fn new(value: T) -> Self {
Self(value)
}
pub fn into_inner(self) -> T {
self.0
}
}
impl<T> Deref for Form<T> {
type Target = T;
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl<T: DeserializeOwned> FromRequest for Form<T> {
fn from_request(request: &HttpRequest) -> Result<Self, Error> {
let value: T = request.form()?;
Ok(Form(value))
}
}
#[derive(Debug, Clone)]
pub struct ContentType(pub String);
impl ContentType {
pub fn is_json(&self) -> bool {
self.0.contains("application/json")
}
pub fn is_form(&self) -> bool {
self.0.contains("application/x-www-form-urlencoded")
}
pub fn is_multipart(&self) -> bool {
self.0.contains("multipart/form-data")
}
pub fn into_inner(self) -> String {
self.0
}
}
impl FromRequest for ContentType {
fn from_request(request: &HttpRequest) -> Result<Self, Error> {
let value = request
.headers
.get("content-type")
.map(str::to_owned)
.unwrap_or_default();
Ok(ContentType(value))
}
}
impl Deref for ContentType {
type Target = str;
fn deref(&self) -> &Self::Target {
&self.0
}
}
#[derive(Debug, Clone)]
pub struct MethodExtractor(pub crate::Method);
impl MethodExtractor {
pub fn is_get(&self) -> bool {
self.0 == "GET"
}
pub fn is_post(&self) -> bool {
self.0 == "POST"
}
pub fn is_put(&self) -> bool {
self.0 == "PUT"
}
pub fn is_delete(&self) -> bool {
self.0 == "DELETE"
}
pub fn is_patch(&self) -> bool {
self.0 == "PATCH"
}
}
impl FromRequest for MethodExtractor {
fn from_request(request: &HttpRequest) -> Result<Self, Error> {
Ok(MethodExtractor(request.method.clone()))
}
}
impl Deref for MethodExtractor {
type Target = crate::Method;
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl FromRequest for HttpRequest {
fn from_request(request: &HttpRequest) -> Result<Self, Error> {
Ok(request.clone())
}
}
#[macro_export]
macro_rules! body {
($request:expr, $type:ty) => {
<$crate::extractors::Body<$type> as $crate::extractors::FromRequest>::from_request(
&$request,
)
.map(|b| b.into_inner())
};
}
#[macro_export]
macro_rules! query {
($request:expr, $type:ty) => {
<$crate::extractors::Query<$type> as $crate::extractors::FromRequest>::from_request(
&$request,
)
.map(|q| q.into_inner())
};
}
#[macro_export]
macro_rules! path {
($request:expr, $name:expr, $type:ty) => {
<$crate::extractors::Path<$type> as $crate::extractors::FromRequestNamed>::from_request(
&$request, $name,
)
.map(|p| p.into_inner())
};
}
#[macro_export]
macro_rules! header {
($request:expr, $name:expr) => {
<$crate::extractors::Header as $crate::extractors::FromRequestNamed>::from_request(
&$request, $name,
)
.map(|h| h.into_value())
};
}
#[cfg(test)]
mod tests {
use super::*;
use serde::Deserialize;
fn create_request() -> HttpRequest {
let mut req = HttpRequest::new("GET", "/users/123?page=1&limit=10");
req.push_param("id", "123");
req.headers
.insert("Authorization", "Bearer token123".to_string());
req.headers
.insert("Content-Type", "application/json".to_string());
req
}
#[test]
fn test_path_extraction() {
let request = create_request();
let id: Path<u32> = Path::from_request(&request, "id").unwrap();
assert_eq!(*id, 123);
}
#[test]
fn test_path_missing() {
let request = create_request();
let result: Result<Path<u32>, _> = Path::from_request(&request, "missing");
assert!(result.is_err());
}
#[test]
fn test_header_extraction() {
let request = create_request();
let auth: Header = Header::from_request(&request, "Authorization").unwrap();
assert_eq!(auth.value(), "Bearer token123");
}
#[test]
fn test_header_optional() {
let request = create_request();
let auth = Header::optional(&request, "Authorization");
assert!(auth.is_some());
let missing = Header::optional(&request, "X-Missing");
assert!(missing.is_none());
}
#[test]
fn test_headers_extraction() {
let request = create_request();
let headers: Headers = Headers::from_request(&request).unwrap();
assert!(headers.contains("Authorization"));
assert!(headers.contains("Content-Type"));
assert!(!headers.contains("X-Missing"));
}
#[test]
fn test_query_extraction() {
let request = create_request();
#[derive(Debug, Deserialize, PartialEq)]
struct Pagination {
page: u32,
limit: u32,
}
let query: Query<Pagination> = Query::from_request(&request).unwrap();
assert_eq!(query.page, 1);
assert_eq!(query.limit, 10);
}
#[test]
fn test_query_extraction_with_ampersand_and_equals_in_value() {
let request = HttpRequest::new("GET", "/items?note=1%26b%3D2&name=a%3Db%26c");
#[derive(Debug, Deserialize, PartialEq)]
struct Filters {
note: String,
name: String,
}
let query: Query<Filters> = Query::from_request(&request).unwrap();
assert_eq!(query.note, "1&b=2");
assert_eq!(query.name, "a=b&c");
}
#[test]
fn test_path_params_extraction_with_ampersand_and_equals_in_value() {
let mut request = HttpRequest::new("GET", "/items/x");
request.push_param("slug", Bytes::from_static(b"a&b=c"));
#[derive(Debug, Deserialize, PartialEq)]
struct Params {
slug: String,
}
let params: PathParams<Params> = PathParams::from_request(&request).unwrap();
assert_eq!(params.slug, "a&b=c");
}
#[test]
fn test_body_extraction() {
let mut request = create_request();
request.body = Bytes::from(
serde_json::to_vec(&serde_json::json!({
"name": "Test",
"email": "test@example.com"
}))
.unwrap(),
);
#[derive(Debug, Deserialize)]
struct CreateUser {
name: String,
email: String,
}
let body: Body<CreateUser> = Body::from_request(&request).unwrap();
assert_eq!(body.name, "Test");
assert_eq!(body.email, "test@example.com");
}
#[test]
fn test_raw_body() {
let mut request = create_request();
request.body = Bytes::from_static(b"raw content");
let raw: RawBody = RawBody::from_request(&request).unwrap();
assert_eq!(raw.len(), 11);
assert_eq!(raw.to_string_lossy(), "raw content");
}
#[test]
fn test_content_type() {
let request = create_request();
let ct: ContentType = ContentType::from_request(&request).unwrap();
assert!(ct.is_json());
assert!(!ct.is_form());
assert!(!ct.is_multipart());
}
#[test]
fn test_method() {
let request = create_request();
let method: MethodExtractor = MethodExtractor::from_request(&request).unwrap();
assert!(method.is_get());
assert!(!method.is_post());
}
#[test]
fn test_request_extraction() {
let request = create_request();
let extracted: HttpRequest = HttpRequest::from_request(&request).unwrap();
assert_eq!(extracted.method, request.method);
assert_eq!(extracted.path, request.path);
}
}