use std::sync::Arc;
use axum::{
Router,
extract::{Request, State},
http::{HeaderValue, StatusCode, header},
middleware::{self, Next},
response::{IntoResponse, Response},
routing::get,
};
use security_rust::{Decision, MemoryStore, RequestContext, SessionConfig, SessionGuard};
use tower::ServiceExt;
type GuardState = Arc<SessionGuard<MemoryStore>>;
const NOW: u64 = 1_700_000_000;
const TOKEN_HEADER: &str = "x-session-token";
const FINGERPRINT_HEADER: &str = "x-client-fingerprint";
const REGION_HEADER: &str = "x-client-region";
#[tokio::main(flavor = "current_thread")]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
self_check().await
}
#[tokio::test]
async fn middleware_branches_on_every_decision() {
self_check().await.expect("自检失败");
}
async fn self_check() -> Result<(), Box<dyn std::error::Error>> {
let guard: GuardState = Arc::new(SessionGuard::new(
MemoryStore::new(),
SessionConfig::default(),
));
let app = build_app(Arc::clone(&guard));
let login = RequestContext {
token: "tok-alice-1",
subject: "alice",
fingerprint: "ip=10.0.0.1|ua=Firefox/128",
location: Some("CN-BJ"),
coords: None,
signature: None,
at: Some(NOW),
};
guard.bind(&login, NOW)?;
println!("security-rust · axum 中间件接入示例");
println!("已种会话:alice @ CN-BJ · tok-alice-1\n");
let cases = [
(
"ALLOW 正常会话",
"tok-alice-1",
"ip=10.0.0.1|ua=Firefox/128",
"CN-BJ",
Expect::Pass,
),
(
"CHALLENGE 异地(位置变了,指纹没变)",
"tok-alice-1",
"ip=10.0.0.1|ua=Firefox/128",
"US-NY",
Expect::StepUp,
),
(
"BLOCK 指纹不符(token 被盗)",
"tok-alice-1",
"ip=203.0.113.9|ua=curl/8",
"CN-BJ",
Expect::Reject,
),
(
"BLOCK 未知 token",
"tok-forged",
"ip=10.0.0.1|ua=Firefox/128",
"CN-BJ",
Expect::Reject,
),
];
for (label, token, fp, region, expect) in cases {
let (status, step_up) = probe(app.clone(), token, fp, region).await?;
let seen = match (status, step_up) {
(StatusCode::OK, _) => Expect::Pass,
(_, true) => Expect::StepUp,
_ => Expect::Reject,
};
println!(
" {label} → {status}{}",
match step_up {
true => " · X-Step-Up-Auth: required",
false => "",
}
);
assert_eq!(seen, expect, "{label}");
}
println!("\n三种 Decision 分支均按预期工作。");
Ok(())
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Expect {
Pass,
StepUp,
Reject,
}
async fn probe(
app: Router,
token: &str,
fingerprint: &str,
region: &str,
) -> Result<(StatusCode, bool), Box<dyn std::error::Error>> {
let req = Request::builder()
.uri("/protected")
.header(TOKEN_HEADER, token)
.header(FINGERPRINT_HEADER, fingerprint)
.header(REGION_HEADER, region)
.body(axum::body::Body::empty())?;
let res = app.oneshot(req).await?;
let step_up = res.headers().get("x-step-up-auth").is_some();
Ok((res.status(), step_up))
}
fn build_app(guard: GuardState) -> Router {
Router::new()
.route("/protected", get(protected))
.layer(middleware::from_fn_with_state(guard, session_gate))
}
async fn protected() -> &'static str {
"ok"
}
async fn session_gate(State(guard): State<GuardState>, req: Request, next: Next) -> Response {
let decision = {
let headers = req.headers();
let ctx = RequestContext {
token: headers
.get(TOKEN_HEADER)
.and_then(|v| v.to_str().ok())
.unwrap_or(""),
subject: "",
fingerprint: headers
.get(FINGERPRINT_HEADER)
.and_then(|v| v.to_str().ok())
.unwrap_or(""),
location: headers.get(REGION_HEADER).and_then(|v| v.to_str().ok()),
coords: None,
signature: None,
at: Some(NOW),
};
guard.verify(&ctx, NOW).decision
};
match decision {
Decision::Allow => next.run(req).await,
Decision::Challenge => {
let mut res = (StatusCode::UNAUTHORIZED, "step-up required").into_response();
res.headers_mut()
.insert("x-step-up-auth", HeaderValue::from_static("required"));
res
}
Decision::Block => (
StatusCode::UNAUTHORIZED,
[(header::WWW_AUTHENTICATE, "Session")],
"session rejected",
)
.into_response(),
}
}