use std::convert::Infallible;
use std::sync::Arc;
use axum::extract::FromRequestParts;
use axum::http::header::{ACCEPT_LANGUAGE, CONTENT_LANGUAGE, VARY};
use axum::http::request::Parts;
use axum::http::{HeaderMap, HeaderValue, Request};
use axum::response::{IntoResponse, Response};
use tower::{Layer, Service};
use super::args::TranslationArgs;
use super::catalog::{Catalog, Catalogs};
use super::error::I18nError;
use super::locale::LocaleId;
const MAX_ACCEPT_LANGUAGE_LEN: usize = 512;
const MAX_CANDIDATES: usize = 16;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum LocaleSource {
Url,
Session,
Header,
Default,
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct Locale {
id: LocaleId,
source: LocaleSource,
catalogs: Catalogs,
}
impl Locale {
#[must_use]
pub fn id(&self) -> &LocaleId {
&self.id
}
#[must_use]
pub fn source(&self) -> LocaleSource {
self.source
}
#[must_use]
pub fn is_default(&self) -> bool {
self.source == LocaleSource::Default
}
#[must_use]
pub fn catalogs(&self) -> &Catalogs {
&self.catalogs
}
#[must_use]
pub fn catalog(&self) -> &Catalog {
self.catalogs
.catalog(&self.id)
.unwrap_or_else(|| self.catalogs.default_catalog())
}
pub fn message(&self, key: &str) -> Result<String, I18nError> {
self.catalogs.message(&self.id, key)
}
pub fn translate(&self, key: &str, args: &TranslationArgs) -> Result<String, I18nError> {
self.catalogs.translate(&self.id, key, args)
}
}
impl<S> FromRequestParts<S> for Locale
where
S: Send + Sync,
{
type Rejection = Response;
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
parts
.extensions
.get::<Locale>()
.cloned()
.ok_or_else(|| crate::Error::from(I18nError::NotNegotiated).into_response())
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct LocaleNegotiator {
catalogs: Catalogs,
query_parameter: Option<Arc<str>>,
session_key: Option<Arc<str>>,
}
impl LocaleNegotiator {
#[must_use]
pub fn new(catalogs: Catalogs) -> Self {
Self {
catalogs,
query_parameter: None,
session_key: None,
}
}
#[must_use]
pub fn query_parameter(mut self, name: impl Into<Arc<str>>) -> Self {
self.query_parameter = Some(name.into());
self
}
#[must_use]
pub fn session_key(mut self, key: impl Into<Arc<str>>) -> Self {
self.session_key = Some(key.into());
self
}
#[must_use]
pub fn catalogs(&self) -> &Catalogs {
&self.catalogs
}
#[must_use]
pub fn fallback(&self) -> Locale {
Locale {
id: self.catalogs.default_locale().clone(),
source: LocaleSource::Default,
catalogs: self.catalogs.clone(),
}
}
#[must_use]
pub fn resolve(
&self,
url: Option<&str>,
session: Option<&str>,
accept_language: Option<&str>,
) -> Locale {
if self.query_parameter.is_some()
&& let Some(id) = url.and_then(|tag| self.registered(tag))
{
return self.at(id, LocaleSource::Url);
}
if self.session_key.is_some()
&& let Some(id) = session.and_then(|tag| self.registered(tag))
{
return self.at(id, LocaleSource::Session);
}
if let Some(id) = accept_language.and_then(|header| self.match_accept_language(header)) {
return self.at(id, LocaleSource::Header);
}
self.fallback()
}
fn registered(&self, tag: &str) -> Option<LocaleId> {
let id = LocaleId::parse(tag).ok()?;
if self.catalogs.contains(&id) {
return Some(id);
}
self.catalogs
.locales()
.find(|registered| registered.language() == id.language())
.cloned()
}
fn at(&self, id: LocaleId, source: LocaleSource) -> Locale {
Locale {
id,
source,
catalogs: self.catalogs.clone(),
}
}
fn match_accept_language(&self, header: &str) -> Option<LocaleId> {
for (tag, _) in accept_language_candidates(header) {
if let Some(id) = self.registered(tag) {
return Some(id);
}
}
None
}
}
fn accept_language_candidates(header: &str) -> Vec<(&str, u16)> {
let header = &header[..header.len().min(MAX_ACCEPT_LANGUAGE_LEN)];
let mut candidates: Vec<(&str, u16)> = Vec::new();
for entry in header.split(',').take(MAX_CANDIDATES) {
let mut parts = entry.split(';');
let Some(range) = parts.next().map(str::trim) else {
continue;
};
if range.is_empty() || range == "*" {
continue;
}
let quality = parts
.filter_map(|parameter| {
let parameter = parameter.trim();
let value = parameter.strip_prefix("q=").or_else(|| {
parameter
.strip_prefix("Q=")
.or_else(|| parameter.strip_prefix("q ="))
})?;
let weight: f32 = value.trim().parse().ok()?;
if (0.0..=1.0).contains(&weight) {
#[expect(
clippy::cast_possible_truncation,
clippy::cast_sign_loss,
reason = "the range check above bounds the product to 0..=1000"
)]
Some((weight * 1000.0) as u16)
} else {
None
}
})
.next()
.unwrap_or(1000);
if quality == 0 {
continue;
}
candidates.push((range, quality));
}
candidates.sort_by_key(|candidate| std::cmp::Reverse(candidate.1));
candidates
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct LocaleLayer {
negotiator: Arc<LocaleNegotiator>,
}
impl LocaleLayer {
#[must_use]
pub fn new(negotiator: LocaleNegotiator) -> Self {
Self {
negotiator: Arc::new(negotiator),
}
}
}
impl<S> Layer<S> for LocaleLayer {
type Service = LocaleMiddleware<S>;
fn layer(&self, inner: S) -> Self::Service {
LocaleMiddleware {
inner,
negotiator: Arc::clone(&self.negotiator),
}
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct LocaleMiddleware<S> {
inner: S,
negotiator: Arc<LocaleNegotiator>,
}
impl<S, ReqBody> Service<Request<ReqBody>> for LocaleMiddleware<S>
where
S: Service<Request<ReqBody>, Response = Response, Error = Infallible> + Clone + Send + 'static,
S::Future: Send + 'static,
ReqBody: Send + 'static,
{
type Response = Response;
type Error = Infallible;
type Future = std::pin::Pin<
Box<dyn std::future::Future<Output = Result<Self::Response, Self::Error>> + Send>,
>;
fn poll_ready(
&mut self,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, request: Request<ReqBody>) -> Self::Future {
let negotiator = Arc::clone(&self.negotiator);
let (mut parts, body) = request.into_parts();
let url = negotiator
.query_parameter
.as_deref()
.and_then(|name| query_value(parts.uri.query(), name))
.map(str::to_owned);
let accept_language = parts
.headers
.get(ACCEPT_LANGUAGE)
.and_then(|value| value.to_str().ok())
.map(str::to_owned);
let mut inner = self.inner.clone();
Box::pin(async move {
let session = session_value(&parts, negotiator.session_key.as_deref()).await;
let locale = negotiator.resolve(
url.as_deref(),
session.as_deref(),
accept_language.as_deref(),
);
let tag = locale.id().clone();
parts.extensions.insert(locale);
let mut response = inner.call(Request::from_parts(parts, body)).await?;
annotate(response.headers_mut(), &tag);
Ok(response)
})
}
}
fn query_value<'q>(query: Option<&'q str>, name: &str) -> Option<&'q str> {
query?
.split('&')
.filter_map(|pair| {
let (key, value) = pair.split_once('=')?;
(key == name).then_some(value)
})
.next_back()
}
#[cfg(feature = "auth")]
async fn session_value(parts: &Parts, key: Option<&str>) -> Option<String> {
let key = key?;
let session = parts.extensions.get::<tower_sessions::Session>()?;
session.get::<String>(key).await.ok().flatten()
}
#[cfg(not(feature = "auth"))]
#[expect(
clippy::unused_async,
reason = "matches the `auth` signature so the call site needs no cfg"
)]
async fn session_value(_parts: &Parts, _key: Option<&str>) -> Option<String> {
None
}
fn annotate(headers: &mut HeaderMap, locale: &LocaleId) {
if !headers.contains_key(CONTENT_LANGUAGE)
&& let Ok(value) = HeaderValue::from_str(locale.as_str())
{
headers.insert(CONTENT_LANGUAGE, value);
}
ensure_vary_accept_language(headers);
}
fn ensure_vary_accept_language(headers: &mut HeaderMap) {
const ACCEPT_LANGUAGE_TOKEN: &str = "Accept-Language";
let mut tokens: Vec<String> = Vec::new();
for value in headers.get_all(VARY) {
let Ok(value) = value.to_str() else {
continue;
};
for token in value.split(',').map(str::trim).filter(|t| !t.is_empty()) {
if token == "*" {
return;
}
if token.eq_ignore_ascii_case(ACCEPT_LANGUAGE_TOKEN) {
return;
}
tokens.push(token.to_owned());
}
}
tokens.push(ACCEPT_LANGUAGE_TOKEN.to_owned());
if let Ok(value) = HeaderValue::from_str(&tokens.join(", ")) {
headers.remove(VARY);
headers.insert(VARY, value);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::i18n::Catalog;
fn catalogs() -> Catalogs {
Catalogs::new(Catalog::parse(id("en"), "hi = Hello").unwrap())
.with(Catalog::parse(id("fr"), "hi = Bonjour").unwrap())
.with(Catalog::parse(id("pt-BR"), "hi = Ola").unwrap())
}
fn id(tag: &str) -> LocaleId {
LocaleId::parse(tag).unwrap()
}
fn negotiator() -> LocaleNegotiator {
LocaleNegotiator::new(catalogs())
.query_parameter("lang")
.session_key("locale")
}
#[test]
fn a_hostile_locale_never_selects_anything() {
let negotiator = negotiator();
let hostile = [
"../../etc/passwd",
"../../../../../../etc/shadow",
"..\\..\\..\\windows\\win.ini",
"/etc/passwd",
"C:\\Windows\\System32",
"en/../../etc/passwd",
"fr/../../../root/.ssh/id_rsa",
"%2e%2e%2f%2e%2e%2fetc%2fpasswd",
"....//....//etc/passwd",
"en\0",
"en\0.ftl",
"\0/etc/passwd",
"fr\nSet-Cookie: session=stolen",
"fr\r\nX-Injected: 1",
"en; rm -rf /",
"$(cat /etc/passwd)",
"`id`",
"{{7*7}}",
"<script>alert(1)</script>",
"en\u{202e}",
"\u{feff}en",
];
for tag in hostile {
let mut outcomes = vec![
negotiator.resolve(Some(tag), None, None),
negotiator.resolve(None, Some(tag), None),
];
if !tag.contains([';', ',']) {
outcomes.push(negotiator.resolve(None, None, Some(tag)));
}
for locale in outcomes {
assert_eq!(locale.id(), &id("en"), "{tag:?} selected {}", locale.id());
assert_eq!(locale.source(), LocaleSource::Default, "{tag:?}");
}
assert!(LocaleId::parse(tag).is_err(), "{tag:?} parsed");
}
}
#[test]
fn a_header_parameter_is_not_part_of_the_language_range() {
let locale = negotiator().resolve(None, None, Some("en; rm -rf /"));
assert_eq!(locale.id(), &id("en"));
assert_eq!(locale.source(), LocaleSource::Header);
assert!(LocaleId::parse("en; rm -rf /").is_err());
assert_eq!(
negotiator()
.resolve(Some("en; rm -rf /"), None, None)
.source(),
LocaleSource::Default
);
assert_eq!(
negotiator().resolve(None, None, Some("de; fr")).source(),
LocaleSource::Default
);
}
#[test]
fn an_overlong_locale_never_selects_anything() {
let negotiator = negotiator();
for length in [36, 1024, 64 * 1024] {
let tag = "e".repeat(length);
let locale = negotiator.resolve(Some(&tag), None, None);
assert_eq!(locale.id(), &id("en"));
assert_eq!(locale.source(), LocaleSource::Default);
}
}
#[test]
fn a_hostile_candidate_is_skipped_and_the_next_one_is_tried() {
let locale = negotiator().resolve(
None,
None,
Some("../../etc/passwd;q=1.0,\0;q=0.95,fr;q=0.9"),
);
assert_eq!(locale.id(), &id("fr"));
assert_eq!(locale.source(), LocaleSource::Header);
}
#[test]
fn matching_is_not_a_prefix_match_on_the_raw_string() {
let negotiator = negotiator();
for tag in ["fr-", "fr..", "fr/x", "frx", "fr\u{0}"] {
assert_eq!(
negotiator.resolve(Some(tag), None, None).source(),
LocaleSource::Default,
"{tag:?}"
);
}
}
#[test]
fn an_unregistered_but_well_formed_locale_falls_back() {
let locale = negotiator().resolve(Some("de-DE"), None, None);
assert_eq!(locale.id(), &id("en"));
assert_eq!(locale.source(), LocaleSource::Default);
}
#[test]
fn the_url_beats_the_session_and_the_header() {
let locale = negotiator().resolve(Some("fr"), Some("pt-BR"), Some("en"));
assert_eq!(locale.id(), &id("fr"));
assert_eq!(locale.source(), LocaleSource::Url);
}
#[test]
fn the_session_beats_the_header() {
let locale = negotiator().resolve(None, Some("fr"), Some("en"));
assert_eq!(locale.id(), &id("fr"));
assert_eq!(locale.source(), LocaleSource::Session);
}
#[test]
fn the_header_beats_the_default() {
let locale = negotiator().resolve(None, None, Some("fr"));
assert_eq!(locale.id(), &id("fr"));
assert_eq!(locale.source(), LocaleSource::Header);
}
#[test]
fn nothing_at_all_is_the_default() {
let locale = negotiator().resolve(None, None, None);
assert_eq!(locale.id(), &id("en"));
assert!(locale.is_default());
}
#[test]
fn an_unconfigured_override_is_ignored() {
let bare = LocaleNegotiator::new(catalogs());
assert_eq!(
bare.resolve(Some("fr"), Some("fr"), None).source(),
LocaleSource::Default
);
assert_eq!(bare.resolve(None, None, Some("fr")).id(), &id("fr"));
}
#[test]
fn quality_values_order_the_candidates() {
let locale = negotiator().resolve(None, None, Some("de;q=1.0,fr;q=0.8,en;q=0.9"));
assert_eq!(locale.id(), &id("en"), "0.9 beats 0.8");
}
#[test]
fn a_weightless_entry_is_q_1() {
let locale = negotiator().resolve(None, None, Some("fr;q=0.9,en"));
assert_eq!(locale.id(), &id("en"));
}
#[test]
fn equal_weights_keep_the_order_they_were_sent_in() {
assert_eq!(
negotiator().resolve(None, None, Some("fr,en")).id(),
&id("fr")
);
assert_eq!(
negotiator().resolve(None, None, Some("en,fr")).id(),
&id("en")
);
}
#[test]
fn a_refused_language_is_not_selected() {
let locale = negotiator().resolve(None, None, Some("fr;q=0,pt-BR;q=0"));
assert_eq!(locale.source(), LocaleSource::Default);
}
#[test]
fn a_wildcard_is_not_a_locale() {
assert_eq!(
negotiator().resolve(None, None, Some("*")).source(),
LocaleSource::Default
);
}
#[test]
fn whitespace_and_casing_are_tolerated() {
let locale = negotiator().resolve(None, None, Some(" DE-de ;q=0.2 , FR ; q=0.9 "));
assert_eq!(locale.id(), &id("fr"));
}
#[test]
fn a_region_falls_back_to_its_language() {
let locale = negotiator().resolve(None, None, Some("fr-CA"));
assert_eq!(locale.id(), &id("fr"));
assert_eq!(locale.source(), LocaleSource::Header);
}
#[test]
fn a_language_falls_back_to_a_registered_region() {
let locale = negotiator().resolve(None, None, Some("pt"));
assert_eq!(locale.id(), &id("pt-BR"));
}
#[test]
fn an_exact_match_is_taken_before_a_language_match() {
let locale = negotiator().resolve(None, None, Some("pt-PT;q=0.9,pt-BR;q=0.9"));
assert_eq!(locale.id(), &id("pt-BR"));
}
#[test]
fn a_pathological_header_is_bounded() {
let header = "en;q=0.1,".repeat(10_000) + "fr";
let locale = negotiator().resolve(None, None, Some(&header));
assert_eq!(locale.id(), &id("en"));
let candidates = accept_language_candidates(&header);
assert!(candidates.len() <= MAX_CANDIDATES);
}
#[test]
fn an_empty_header_is_no_header() {
assert_eq!(
negotiator().resolve(None, None, Some("")).source(),
LocaleSource::Default
);
assert_eq!(
negotiator().resolve(None, None, Some(",,,")).source(),
LocaleSource::Default
);
}
#[test]
fn a_query_parameter_is_read_by_name() {
assert_eq!(query_value(Some("lang=fr"), "lang"), Some("fr"));
assert_eq!(query_value(Some("a=1&lang=fr&b=2"), "lang"), Some("fr"));
assert_eq!(query_value(Some("language=fr"), "lang"), None);
assert_eq!(query_value(Some("lang"), "lang"), None);
assert_eq!(query_value(None, "lang"), None);
}
#[test]
fn a_repeated_query_parameter_resolves_to_the_last_one() {
assert_eq!(query_value(Some("lang=fr&lang=en"), "lang"), Some("en"));
}
#[test]
fn a_response_says_what_language_it_is_in() {
let mut headers = HeaderMap::new();
annotate(&mut headers, &id("pt-BR"));
assert_eq!(headers[CONTENT_LANGUAGE], "pt-BR");
assert_eq!(headers[VARY], "Accept-Language");
}
#[test]
fn a_content_language_the_handler_set_is_left_alone() {
let mut headers = HeaderMap::new();
headers.insert(CONTENT_LANGUAGE, HeaderValue::from_static("mul"));
annotate(&mut headers, &id("fr"));
assert_eq!(headers[CONTENT_LANGUAGE], "mul");
}
#[test]
fn vary_is_merged_and_never_duplicated() {
let mut headers = HeaderMap::new();
headers.insert(VARY, HeaderValue::from_static("Accept-Encoding"));
annotate(&mut headers, &id("fr"));
assert_eq!(headers[VARY], "Accept-Encoding, Accept-Language");
annotate(&mut headers, &id("fr"));
assert_eq!(headers[VARY], "Accept-Encoding, Accept-Language");
}
#[test]
fn an_existing_vary_accept_language_is_recognised_whatever_its_casing() {
let mut headers = HeaderMap::new();
headers.insert(VARY, HeaderValue::from_static("accept-language"));
annotate(&mut headers, &id("fr"));
assert_eq!(headers[VARY], "accept-language");
}
#[test]
fn a_vary_star_is_left_alone() {
let mut headers = HeaderMap::new();
headers.insert(VARY, HeaderValue::from_static("*"));
annotate(&mut headers, &id("fr"));
assert_eq!(headers[VARY], "*");
}
async fn serve(uri: &str, headers: &[(&str, &str)]) -> Response {
use axum::Router;
use axum::routing::get;
use tower::ServiceExt as _;
let app: Router =
Router::new()
.route(
"/{*rest}",
get(|locale: Locale| async move {
format!("{}:{:?}", locale.id(), locale.source())
}),
)
.route(
"/",
get(|locale: Locale| async move {
format!("{}:{:?}", locale.id(), locale.source())
}),
)
.layer(LocaleLayer::new(negotiator()));
let mut request = Request::get(uri);
for (name, value) in headers {
request = request.header(*name, *value);
}
app.oneshot(request.body(axum::body::Body::empty()).unwrap())
.await
.unwrap()
}
async fn body_of(response: Response) -> String {
let bytes = axum::body::to_bytes(response.into_body(), 1 << 16)
.await
.unwrap();
String::from_utf8(bytes.to_vec()).unwrap()
}
#[tokio::test]
async fn the_layer_negotiates_from_the_header() {
let response = serve("/", &[("accept-language", "fr-CA,fr;q=0.9")]).await;
assert_eq!(response.headers()[CONTENT_LANGUAGE], "fr");
assert_eq!(response.headers()[VARY], "Accept-Language");
assert_eq!(body_of(response).await, "fr:Header");
}
#[tokio::test]
async fn the_layer_reads_the_url_override() {
let response = serve("/?lang=pt-BR", &[("accept-language", "fr")]).await;
assert_eq!(body_of(response).await, "pt-BR:Url");
}
#[tokio::test]
async fn the_layer_falls_back_on_a_hostile_url_override() {
let response = serve("/?lang=../../etc/passwd", &[]).await;
assert_eq!(response.headers()[CONTENT_LANGUAGE], "en");
assert_eq!(body_of(response).await, "en:Default");
}
#[tokio::test]
async fn an_unreadable_header_is_simply_absent() {
let response = serve("/", &[("accept-language", "fr\u{e9}")]).await;
assert_eq!(response.status(), axum::http::StatusCode::OK);
}
#[tokio::test]
async fn the_extractor_without_the_layer_is_a_500_that_leaks_nothing() {
use axum::Router;
use axum::routing::get;
use tower::ServiceExt as _;
let app: Router = Router::new().route(
"/",
get(|locale: Locale| async move { locale.id().to_string() }),
);
let response = app
.oneshot(Request::get("/").body(axum::body::Body::empty()).unwrap())
.await
.unwrap();
assert_eq!(
response.status(),
axum::http::StatusCode::INTERNAL_SERVER_ERROR
);
let body = body_of(response).await;
assert!(!body.contains("LocaleLayer"), "{body}");
assert!(!body.contains("negotiat"), "{body}");
}
}