use std::sync::Arc;
use axum::extract::{Path, State};
use axum::http::{header, HeaderMap, StatusCode};
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use axum::routing::{get, post};
use axum::{Form, Json, Router};
use serde::Deserialize;
use serde_json::json;
use uuid::Uuid;
use backbone_rate_limit::{InMemoryStorage, RateLimitMiddleware};
use crate::application::service::trace_click_ports::TraceClickSlot;
use crate::application::service::trace_route_service::{TraceRouteError, TraceRouteService};
pub const TRACE_RATE_LIMIT_MAX: u64 = 120;
pub const TRACE_RATE_LIMIT_WINDOW_SECS: u64 = 60;
#[derive(Clone)]
pub struct TraceRouteState {
service: Arc<TraceRouteService>,
limiter: Arc<RateLimitMiddleware<InMemoryStorage>>,
}
const PIXEL_GIF: &[u8] = &[
0x47, 0x49, 0x46, 0x38, 0x39, 0x61, 0x01, 0x00, 0x01, 0x00, 0x80, 0x00, 0x00, 0xff, 0xff,
0xff, 0x00, 0x00, 0x00, 0x21, 0xf9, 0x04, 0x01, 0x00, 0x00, 0x00, 0x00, 0x2c, 0x00, 0x00,
0x00, 0x00, 0x01, 0x00, 0x01, 0x00, 0x00, 0x02, 0x02, 0x44, 0x01, 0x00, 0x3b,
];
fn not_found() -> Response {
(StatusCode::NOT_FOUND, "not found").into_response()
}
fn trace_route_err(context: &str, e: TraceRouteError) -> Response {
match e {
TraceRouteError::UnknownTrace | TraceRouteError::Inconsistent => not_found(),
other => {
tracing::warn!(context, error = %other, "trace route refused (internal class)");
not_found()
}
}
}
fn visitor_ip(headers: &HeaderMap) -> Option<String> {
headers
.get("x-forwarded-for")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.split(',').next())
.map(|v| v.trim().to_string())
.filter(|v| !v.is_empty())
}
fn visitor_country(headers: &HeaderMap) -> Option<String> {
headers
.get("x-visitor-country")
.and_then(|v| v.to_str().ok())
.map(|v| v.trim().to_string())
}
async fn click_redirect(
State(state): State<TraceRouteState>,
headers: HeaderMap,
Path((code, trace_raw)): Path<(String, String)>,
) -> Response {
let Ok(trace_id) = Uuid::parse_str(&trace_raw) else {
return not_found();
};
match state
.service
.record_click(
&code,
trace_id,
visitor_ip(&headers).as_deref(),
visitor_country(&headers).as_deref(),
)
.await
{
Ok(target) => (
StatusCode::MOVED_PERMANENTLY,
[(header::LOCATION, target)],
)
.into_response(),
Err(e) => trace_route_err("click", e),
}
}
async fn open_pixel(
State(state): State<TraceRouteState>,
Path((code, trace_raw)): Path<(String, String)>,
) -> Response {
let gif = || {
(
StatusCode::OK,
[(header::CONTENT_TYPE, "image/gif")],
PIXEL_GIF,
)
.into_response()
};
let Ok(trace_id) = Uuid::parse_str(&trace_raw) else {
return gif();
};
match state.service.record_open(&code, trace_id).await {
Ok(()) => gif(),
Err(e) => {
trace_route_err("pixel", e);
gif()
}
}
}
#[derive(Debug, Deserialize)]
pub struct UnsubscribeForm {
pub reason: Option<Uuid>,
}
async fn unsubscribe(
State(state): State<TraceRouteState>,
Path((code, trace_raw)): Path<(String, String)>,
form: Option<Form<UnsubscribeForm>>,
) -> Response {
let Ok(trace_id) = Uuid::parse_str(&trace_raw) else {
return not_found();
};
let reason = form.and_then(|Form(f)| f.reason);
match state.service.unsubscribe(&code, trace_id, reason).await {
Ok(outcome) => {
let audiences: Vec<_> = outcome
.iter()
.map(|a| {
json!({
"audience": a.audience_id,
"changed": a.changed,
"optedOut": a.opted_out,
})
})
.collect();
(
StatusCode::OK,
Json(json!({
"trace": trace_id,
"audiences": audiences,
"count": outcome.len(),
})),
)
.into_response()
}
Err(e) => trace_route_err("unsubscribe", e),
}
}
pub fn code_throttle_key(path: &str) -> String {
let mut segments = path.split('/');
while let Some(segment) = segments.next() {
if segment == "r" {
if let Some(code) = segments.next() {
if !code.is_empty() {
return format!("mailing-trace:{code}");
}
}
break;
}
}
"mailing-trace:other".to_string()
}
fn no_store(mut res: Response) -> Response {
res.headers_mut().insert(
header::CACHE_CONTROL,
header::HeaderValue::from_static("no-store, private"),
);
res
}
async fn per_code_throttle(
State(limiter): State<Arc<RateLimitMiddleware<InMemoryStorage>>>,
req: axum::extract::Request,
next: Next,
) -> Result<Response, Response> {
let key = code_throttle_key(req.uri().path());
match limiter.check(&key).await {
Ok(resp) if resp.allowed => Ok(no_store(next.run(req).await)),
Ok(resp) => {
let mut res = Json(resp).into_response();
*res.status_mut() = StatusCode::TOO_MANY_REQUESTS;
Err(no_store(res))
}
Err(e) => {
tracing::error!("trace-route rate limit check failed for key {key}: {e}");
Ok(no_store(next.run(req).await))
}
}
}
pub fn public_composer(service: Arc<TraceRouteService>) -> Router {
let state = TraceRouteState {
service,
limiter: backbone_rate_limit::middleware(
TRACE_RATE_LIMIT_MAX,
TRACE_RATE_LIMIT_WINDOW_SECS,
),
};
Router::new()
.route("/r/:code/m/:trace", get(click_redirect))
.route("/r/:code/m/:trace/pixel.gif", get(open_pixel))
.route("/r/:code/m/:trace/unsubscribe", post(unsubscribe))
.fallback(|| async { not_found() })
.route_layer(axum::middleware::from_fn_with_state(
state.limiter.clone(),
per_code_throttle,
))
.with_state(state)
}
pub fn public_trace_routes(pool: sqlx::PgPool, clicks: TraceClickSlot) -> Router {
public_composer(Arc::new(TraceRouteService::new(pool, clicks)))
}