tako_rs_extractors/
basic.rs1use base64::Engine;
34use base64::engine::general_purpose::STANDARD;
35use http::StatusCode;
36use http::request::Parts;
37use tako_rs_core::extractors::FromRequest;
38use tako_rs_core::extractors::FromRequestParts;
39use tako_rs_core::responder::Responder;
40use tako_rs_core::types::Request;
41
42pub struct Basic {
49 pub username: String,
51 pub password: String,
53 pub raw: String,
55}
56
57#[derive(Debug)]
59pub enum BasicAuthError {
60 MissingAuthHeader,
62 InvalidAuthHeader,
64 InvalidBasicFormat,
66 InvalidBase64,
68 InvalidUtf8,
70 InvalidCredentialsFormat,
72}
73
74impl Responder for BasicAuthError {
75 fn into_response(self) -> tako_rs_core::types::Response {
77 let (status, message) = match self {
78 BasicAuthError::MissingAuthHeader => {
79 (StatusCode::UNAUTHORIZED, "Missing Authorization header")
80 }
81 BasicAuthError::InvalidAuthHeader => {
82 (StatusCode::UNAUTHORIZED, "Invalid Authorization header")
83 }
84 BasicAuthError::InvalidBasicFormat => (
85 StatusCode::UNAUTHORIZED,
86 "Authorization header is not Basic auth",
87 ),
88 BasicAuthError::InvalidBase64 => (
89 StatusCode::UNAUTHORIZED,
90 "Invalid Base64 encoding in Basic auth",
91 ),
92 BasicAuthError::InvalidUtf8 => (
93 StatusCode::UNAUTHORIZED,
94 "Invalid UTF-8 in Basic auth credentials",
95 ),
96 BasicAuthError::InvalidCredentialsFormat => (
97 StatusCode::UNAUTHORIZED,
98 "Invalid credentials format in Basic auth",
99 ),
100 };
101 (status, message).into_response()
102 }
103}
104
105impl Basic {
106 fn extract_from_headers(headers: &http::HeaderMap) -> Result<Self, BasicAuthError> {
108 let auth_header = headers
109 .get("Authorization")
110 .ok_or(BasicAuthError::MissingAuthHeader)?;
111
112 let auth_str = auth_header
113 .to_str()
114 .map_err(|_| BasicAuthError::InvalidAuthHeader)?;
115
116 if !auth_str.starts_with("Basic ") {
117 return Err(BasicAuthError::InvalidBasicFormat);
118 }
119
120 let encoded = &auth_str[6..];
121 let decoded = STANDARD
122 .decode(encoded)
123 .map_err(|_| BasicAuthError::InvalidBase64)?;
124
125 let decoded_str = std::str::from_utf8(&decoded).map_err(|_| BasicAuthError::InvalidUtf8)?;
126
127 let parts: Vec<&str> = decoded_str.splitn(2, ':').collect();
128 if parts.len() != 2 {
129 return Err(BasicAuthError::InvalidCredentialsFormat);
130 }
131
132 Ok(Basic {
133 username: parts[0].to_string(),
134 password: parts[1].to_string(),
135 raw: auth_str.to_string(),
136 })
137 }
138}
139
140impl<'a> FromRequest<'a> for Basic {
141 type Error = BasicAuthError;
142
143 fn from_request(
144 req: &'a mut Request,
145 ) -> impl core::future::Future<Output = core::result::Result<Self, Self::Error>> + Send + 'a {
146 futures_util::future::ready(Self::extract_from_headers(req.headers()))
147 }
148}
149
150impl<'a> FromRequestParts<'a> for Basic {
151 type Error = BasicAuthError;
152
153 fn from_request_parts(
154 parts: &'a mut Parts,
155 ) -> impl core::future::Future<Output = core::result::Result<Self, Self::Error>> + Send + 'a {
156 futures_util::future::ready(Self::extract_from_headers(&parts.headers))
157 }
158}