1use std::{
5 sync::{Arc, Mutex, MutexGuard, PoisonError},
6 time::Duration,
7};
8
9use axum::{
10 extract::FromRequestParts,
11 http::{HeaderMap, HeaderValue, header, request::Parts},
12};
13use cookie::{Cookie, CookieJar, SameSite};
14
15use crate::{
16 Error, Result,
17 session::{Keys, Rejection, SESSION_COOKIE, reject},
18};
19
20#[derive(Clone)]
53pub struct Cookies(Arc<Mutex<State>>);
54
55struct State {
56 jar: CookieJar,
57 keys: std::result::Result<Keys, String>,
58 secure: bool,
59}
60
61impl Cookies {
62 pub(crate) fn from_headers(headers: &HeaderMap, keys: std::result::Result<Keys, String>, secure: bool) -> Self {
64 let mut jar = CookieJar::new();
65 let cookies = headers.get_all(header::COOKIE).iter().filter_map(|value| value.to_str().ok());
66 for cookie in cookies.flat_map(Cookie::split_parse_encoded).filter_map(std::result::Result::ok) {
67 jar.add_original(cookie.into_owned());
68 }
69 Self(Arc::new(Mutex::new(State { jar, keys, secure })))
70 }
71
72 fn state(&self) -> MutexGuard<'_, State> {
73 self.0.lock().unwrap_or_else(PoisonError::into_inner)
74 }
75
76 pub fn get(&self, name: &str) -> Option<String> {
87 self.state().jar.get(name).map(|cookie| cookie.value().to_owned())
88 }
89
90 pub fn signed(&self, name: &str) -> Result<Option<String>> {
96 let state = self.state();
97 let keys = keys(&state)?;
98 Ok(keys.iter().find_map(|key| state.jar.signed(key).get(name)).map(|cookie| cookie.value().to_owned()))
99 }
100
101 pub fn encrypted(&self, name: &str) -> Result<Option<String>> {
107 let state = self.state();
108 let keys = keys(&state)?;
109 Ok(keys.iter().find_map(|key| state.jar.private(key).get(name)).map(|cookie| cookie.value().to_owned()))
110 }
111
112 pub fn set(&self, name: &str, value: &str, max_age: Option<Duration>) -> Result<()> {
118 let cookie = self.build(name, value, max_age)?;
119 self.state().jar.add(cookie);
120 Ok(())
121 }
122
123 pub fn set_signed(&self, name: &str, value: &str, max_age: Option<Duration>) -> Result<()> {
130 let cookie = self.build(name, value, max_age)?;
131 let mut state = self.state();
132 let key = keys(&state)?[0].clone();
133 state.jar.signed_mut(&key).add(cookie);
134 Ok(())
135 }
136
137 pub fn set_encrypted(&self, name: &str, value: &str, max_age: Option<Duration>) -> Result<()> {
144 let cookie = self.build(name, value, max_age)?;
145 let mut state = self.state();
146 let key = keys(&state)?[0].clone();
147 state.jar.private_mut(&key).add(cookie);
148 Ok(())
149 }
150
151 pub fn remove(&self, name: &str) {
153 let mut cookie = Cookie::build((name.to_owned(), "")).path("/").build();
156 cookie.set_max_age(cookie::time::Duration::ZERO);
157 cookie.set_expires(cookie::time::OffsetDateTime::UNIX_EPOCH);
158 self.state().jar.add(cookie);
159 }
160
161 fn build(&self, name: &str, value: &str, max_age: Option<Duration>) -> Result<Cookie<'static>> {
162 if name == SESSION_COOKIE {
163 return Err(Error::internal(format!(
164 "`{SESSION_COOKIE}` is the session cookie. Fix: store the value in the `Session`, or pick another cookie name"
165 )));
166 }
167 let mut cookie = Cookie::build((name.to_owned(), value.to_owned()))
168 .path("/")
169 .http_only(true)
170 .same_site(SameSite::Lax)
171 .secure(self.state().secure)
172 .build();
173 if let Some(max_age) = max_age {
174 cookie.set_max_age(cookie::time::Duration::seconds(i64::try_from(max_age.as_secs()).unwrap_or(i64::MAX)));
175 }
176 Ok(cookie)
177 }
178
179 pub(crate) fn set_cookies(&self) -> Vec<HeaderValue> {
181 let state = self.state();
182 let values = state.jar.delta().map(|cookie| HeaderValue::from_str(&cookie.encoded().to_string()));
183 values.filter_map(std::result::Result::ok).collect()
184 }
185}
186
187fn keys(state: &State) -> Result<Vec<cookie::Key>> {
189 let keys = state.keys.as_ref().map_err(|err| Error::internal(err.clone()))?;
190 Ok(std::iter::once(&keys.current).chain(&keys.previous).cloned().collect())
191}
192
193impl<S: Sync> FromRequestParts<S> for Cookies {
194 type Rejection = Rejection;
195
196 async fn from_request_parts(parts: &mut Parts, _state: &S) -> std::result::Result<Self, Rejection> {
197 parts.extensions.get::<Cookies>().cloned().ok_or_else(|| {
198 reject(Error::internal(
199 "no cookies on this request. Fix: serve the app with `ocre::serve`, which adds them",
200 ))
201 })
202 }
203}
204
205#[cfg(test)]
206#[path = "../tests/cookies.rs"]
207mod tests;