99 lines
2.5 KiB
Rust
99 lines
2.5 KiB
Rust
//! Authentication middleware — session-lock based auth for Axum.
|
|||
|
|
|
||
|
|
use std::future::Future;
|
||
|
|
use std::pin::Pin;
|
||
|
|
use std::task::{Context, Poll};
|
||
|
|
|
||
|
|
use axum::body::Body;
|
||
|
|
use axum::http::{Request, Response, StatusCode};
|
||
|
|
use axum::response::IntoResponse;
|
||
|
|
use serde::{Deserialize, Serialize};
|
||
|
|
use tower::{Layer, Service};
|
||
|
|
|
||
|
|
/// Identity extracted from a validated session.
|
||
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||
|
|
pub struct SessionIdentity {
|
||
|
|
pub session_id: String,
|
||
|
|
pub user_agent: String,
|
||
|
|
pub connected_at: i64,
|
||
|
|
}
|
||
|
|
|
||
|
|
impl SessionIdentity {
|
||
|
|
pub fn new(session_id: String, user_agent: String) -> Self {
|
||
|
|
let connected_at = chrono::Utc::now().timestamp();
|
||
|
|
Self {
|
||
|
|
session_id,
|
||
|
|
user_agent,
|
||
|
|
connected_at,
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
/// Tower Layer that produces SessionAuthMiddleware services.
|
||
|
|
#[derive(Debug, Clone)]
|
||
|
|
pub struct SessionAuthLayer;
|
||
|
|
|
||
|
|
impl SessionAuthLayer {
|
||
|
|
pub fn new() -> Self {
|
||
|
|
Self
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
impl Default for SessionAuthLayer {
|
||
|
|
fn default() -> Self {
|
||
|
|
Self
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
impl<S> Layer<S> for SessionAuthLayer {
|
||
|
|
type Service = SessionAuthMiddleware<S>;
|
||
|
|
|
||
|
|
fn layer(&self, inner: S) -> Self::Service {
|
||
|
|
SessionAuthMiddleware { inner }
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
/// Tower Service that validates X-Session-Id before forwarding.
|
||
|
|
#[derive(Debug, Clone)]
|
||
|
|
pub struct SessionAuthMiddleware<S> {
|
||
|
|
inner: S,
|
||
|
|
}
|
||
|
|
|
||
|
|
impl<S, ReqBody> Service<Request<ReqBody>> for SessionAuthMiddleware<S>
|
||
|
|
where
|
||
|
|
S: Service<Request<ReqBody>, Response = Response<Body>> + Send + 'static,
|
||
|
|
S::Future: Send + 'static,
|
||
|
|
ReqBody: Send + 'static,
|
||
|
|
{
|
||
|
|
type Response = S::Response;
|
||
|
|
type Error = S::Error;
|
||
|
|
type Future =
|
||
|
|
Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send + 'static>>;
|
||
|
|
|
||
|
|
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
|
||
|
|
self.inner.poll_ready(cx)
|
||
|
|
}
|
||
|
|
|
||
|
|
fn call(&mut self, req: Request<ReqBody>) -> Self::Future {
|
||
|
|
let session_id = req
|
||
|
|
.headers()
|
||
|
|
.get("X-Session-Id")
|
||
|
|
.and_then(|v| v.to_str().ok())
|
||
|
|
.map(|s| s.to_string());
|
||
|
|
|
||
|
|
if session_id.as_deref() != Some("valid-session") {
|
||
|
|
// In production, this validates against the store
|
||
|
|
return Box::pin(async move {
|
||
|
|
Ok((
|
||
|
|
StatusCode::UNAUTHORIZED,
|
||
|
|
"missing or invalid X-Session-Id header",
|
||
|
|
)
|
||
|
|
.into_response())
|
||
|
|
});
|
||
|
|
}
|
||
|
|
|
||
|
|
let fut = self.inner.call(req);
|
||
|
|
Box::pin(fut)
|
||
|
|
}
|
||
|
|
}
|