use std::{
collections::HashMap,
net::Ipv4Addr,
time::{SystemTime, UNIX_EPOCH},
};
use jsonwebtoken::{Algorithm, DecodingKey, EncodingKey, Header, Validation, decode, encode};
use kynos::{
error::rejection::AuthRejection,
prelude::*,
security::{
Authenticates, Authenticator,
auth::{Auth, MaybeAuth, Scoped, Scopes},
carrier::BearerToken,
constant_time_eq,
schemes::{Basic, Credentials},
},
server::Server,
};
use serde::{Deserialize, Serialize};
const ISSUER: &str = "https://auth.example.com";
const TOKEN_LIFETIME: u64 = 900;
const AUDIENCE: &str = "https://api.example.com";
#[derive(Clone, Debug, Serialize, Deserialize)]
struct Claims {
sub: String,
iss: String,
aud: String,
exp: u64,
nbf: u64,
scope: String,
}
impl Claims {
fn scopes(&self) -> impl Iterator<Item = &str> {
self.scope.split(' ').filter(|scope| !scope.is_empty())
}
}
#[derive(Schema, Serialize)]
#[expect(
clippy::struct_field_names,
reason = "RFC 6749 section 5.1 fixes these names"
)]
struct Token {
access_token: String,
token_type: String,
expires_in: u64,
}
#[derive(SecurityScheme)]
#[security(bearer(format = "JWT"))]
#[security(credential = Claims, description = "A short-lived access token")]
struct AccessToken;
struct ReadReports;
impl Scopes for ReadReports {
const SCOPES: &'static [&'static str] = &["reports:read"];
}
struct Key {
encoding: EncodingKey,
decoding: DecodingKey,
}
struct Keys {
current: String,
by_id: HashMap<String, Key>,
}
impl Keys {
fn seeded() -> Self {
let mut by_id = HashMap::new();
for (id, secret) in [
("k1", &b"the-previous-secret"[..]),
("k2", &b"the-current-secret"[..]),
] {
by_id.insert(
id.to_owned(),
Key {
encoding: EncodingKey::from_secret(secret),
decoding: DecodingKey::from_secret(secret),
},
);
}
Self {
current: "k2".to_owned(),
by_id,
}
}
fn issue(&self, subject: &str, scopes: &str) -> String {
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.expect("the clock is after 1970")
.as_secs();
let mut header = Header::new(Algorithm::HS256);
header.kid = Some(self.current.clone());
let claims = Claims {
sub: subject.to_owned(),
iss: ISSUER.to_owned(),
aud: AUDIENCE.to_owned(),
exp: now + TOKEN_LIFETIME,
nbf: now,
scope: scopes.to_owned(),
};
let key = self
.by_id
.get(&self.current)
.expect("the current key is one this store holds");
encode(&header, &claims, &key.encoding).expect("HMAC signing does not fail")
}
}
impl<C: Sync> Authenticator<AccessToken, C> for Keys {
async fn authenticate(
&self,
presented: BearerToken,
context: &C,
) -> Result<Claims, AuthRejection> {
let _ = context;
let header =
jsonwebtoken::decode_header(presented.as_str()).map_err(unauthenticated_whatever)?;
let key = header
.kid
.and_then(|kid| self.by_id.get(&kid))
.ok_or_else(AuthRejection::unauthenticated)?;
let mut validation = Validation::new(Algorithm::HS256);
validation.set_issuer(&[ISSUER]);
validation.set_audience(&[AUDIENCE]);
validation.validate_exp = true;
validation.validate_nbf = true;
validation.leeway = 30;
decode::<Claims>(presented.as_str(), &key.decoding, &validation)
.map(|token| token.claims)
.map_err(unauthenticated_whatever)
}
async fn authorize(
&self,
credential: &Claims,
scopes: &'static [&'static str],
context: &C,
) -> Result<(), AuthRejection> {
let _ = context;
if scopes
.iter()
.all(|demanded| credential.scopes().any(|held| held == *demanded))
{
Ok(())
} else {
Err(AuthRejection::forbidden())
}
}
}
#[allow(clippy::needless_pass_by_value)]
fn unauthenticated_whatever<E>(error: E) -> AuthRejection {
let _ = error;
AuthRejection::unauthenticated()
}
struct Passwords;
impl<C: Sync> Authenticator<Basic<Credentials>, C> for Passwords {
async fn authenticate(
&self,
presented: Credentials,
context: &C,
) -> Result<Credentials, AuthRejection> {
let _ = context;
let known = presented.username == "reporter" || presented.username == "reader";
let correct = constant_time_eq(presented.password.as_bytes(), b"correct-horse");
if known && correct {
Ok(presented)
} else {
Err(AuthRejection::unauthenticated())
}
}
async fn authorize(
&self,
_: &Credentials,
_: &'static [&'static str],
_: &C,
) -> Result<(), AuthRejection> {
Ok(())
}
}
struct App {
keys: std::sync::Arc<Keys>,
passwords: Passwords,
}
impl kynos::di::Provides<std::sync::Arc<Keys>> for App {
fn provide(&self) -> std::sync::Arc<Keys> {
std::sync::Arc::clone(&self.keys)
}
}
impl Authenticates<AccessToken> for App {
type Authenticator = Keys;
fn authenticator(&self) -> &Self::Authenticator {
&self.keys
}
}
impl Authenticates<Basic<Credentials>> for App {
type Authenticator = Passwords;
fn authenticator(&self) -> &Self::Authenticator {
&self.passwords
}
}
#[kynos::post("/session")]
async fn sign_in(
Auth(credentials): Auth<Basic<Credentials>>,
Inject(keys): Inject<std::sync::Arc<Keys>>,
) -> Json<Token> {
let scopes = if credentials.username == "reporter" {
"reports:read"
} else {
""
};
Json(Token {
access_token: keys.issue(&credentials.username, scopes),
token_type: "Bearer".to_owned(),
expires_in: TOKEN_LIFETIME,
})
}
#[kynos::get("/me")]
async fn me(Auth(claims): Auth<AccessToken>) -> Json<Subject> {
Json(Subject {
subject: claims.sub,
})
}
#[derive(Schema, Serialize)]
struct Subject {
subject: String,
}
#[kynos::get("/reports")]
async fn reports(caller: Scoped<AccessToken, ReadReports>) -> NoContent {
let _ = caller.into_inner();
NoContent
}
#[kynos::get("/feed")]
async fn feed(caller: MaybeAuth<AccessToken>) -> Json<Subject> {
Json(Subject {
subject: caller
.into_inner()
.map_or_else(|| "anonymous".to_owned(), |claims| claims.sub),
})
}
#[tokio::main]
async fn main() -> kynos::Result<()> {
let router = Router::<App>::new().mount(kynos::routes![sign_in, me, reports, feed]);
println!("{}", router.openapi()?.to_json()?);
let context = App {
keys: std::sync::Arc::new(Keys::seeded()),
passwords: Passwords,
};
Server::new(router.build(context)?)
.bind((Ipv4Addr::UNSPECIFIED, 3000))
.serve()
.await
}