tuwunel_api/client/session/
jwt.rs1use std::str::FromStr;
2
3use jwt::{Algorithm, DecodingKey, Validation, dangerous::insecure_decode, decode};
4use ruma::{
5 OwnedUserId, UserId,
6 api::client::session::login::v3::{Request, Token},
7};
8use serde::Deserialize;
9use tuwunel_core::{Err, Result, at, config::JwtConfig, debug, err, jwt, warn};
10use tuwunel_service::Services;
11
12use crate::Ruma;
13
14#[derive(Debug, Deserialize)]
15struct Claim {
16 sub: String,
18}
19
20pub(super) async fn handle_login(
21 services: &Services,
22 _body: &Ruma<Request>,
23 info: &Token,
24) -> Result<OwnedUserId> {
25 let user_id = validate_user(services, &info.token)?;
26 if !services.users.exists(&user_id).await {
27 let config = &services.config.jwt;
28 if !config.register_user {
29 return Err!(Request(NotFound("User {user_id} is not registered on this server.")));
30 }
31
32 services
33 .users
34 .create(&user_id, Some("*"), Some("jwt"))
35 .await?;
36 }
37
38 Ok(user_id)
39}
40
41pub(crate) fn validate_user(services: &Services, token: &str) -> Result<OwnedUserId> {
42 let config = &services.config.jwt;
43 if !config.enable {
44 return Err!(Request(Unauthorized("JWT login is not enabled.")));
45 }
46
47 let claim = validate(config, token)?;
48 let local = claim.sub.to_lowercase();
49 let server = &services.server.name;
50 let user_id = UserId::parse_with_server_name(local, server).map_err(|e| {
51 err!(Request(InvalidUsername("JWT subject is not a valid user MXID: {e}")))
52 })?;
53
54 Ok(user_id)
55}
56
57fn validate(config: &JwtConfig, token: &str) -> Result<Claim> {
58 let token_data = if cfg!(debug_assertions) && !config.validate_signature {
59 warn!("JWT signature validation is disabled!");
60 insecure_decode::<Claim>(token)
61 } else {
62 let verifier = init_verifier(config)?;
63 let validator = init_validator(config)?;
64 decode::<Claim>(token, &verifier, &validator)
65 };
66
67 token_data
68 .map(|decoded| (decoded.header, decoded.claims))
69 .inspect(|(head, claim)| debug!(?head, ?claim, "JWT token decoded"))
70 .map_err(|e| err!(Request(Forbidden("Invalid JWT token: {e}"))))
71 .map(at!(1))
72}
73
74fn init_verifier(config: &JwtConfig) -> Result<DecodingKey> {
75 let key = &config.key;
76 let format = config.format.to_uppercase();
77
78 Ok(match format.as_str() {
79 | "HMAC" => DecodingKey::from_secret(key.as_bytes()),
80
81 | "HMACB64" => DecodingKey::from_base64_secret(key.as_str())
82 .map_err(|e| err!(Config("jwt.key", "JWT key is not valid base64: {e}")))?,
83
84 | "ECDSA" => DecodingKey::from_ec_pem(key.as_bytes())
85 .map_err(|e| err!(Config("jwt.key", "JWT key is not valid ECDSA PEM: {e}")))?,
86
87 | "EDDSA" => DecodingKey::from_ed_pem(key.as_bytes())
88 .map_err(|e| err!(Config("jwt.key", "JWT key is not valid EDDSA PEM: {e}")))?,
89
90 | _ => return Err!(Config("jwt.format", "Key format {format:?} is not supported.")),
91 })
92}
93
94fn init_validator(config: &JwtConfig) -> Result<Validation> {
95 let alg = config.algorithm.as_str();
96 let alg = Algorithm::from_str(alg).map_err(|e| {
97 err!(Config("jwt.algorithm", "JWT algorithm is not recognized or configured {e}"))
98 })?;
99
100 let mut validator = Validation::new(alg);
101 let mut required_spec_claims: Vec<_> = ["sub"].into();
102
103 validator.validate_exp = config.validate_exp;
104 if config.require_exp {
105 required_spec_claims.push("exp");
106 }
107
108 validator.validate_nbf = config.validate_nbf;
109 if config.require_nbf {
110 required_spec_claims.push("nbf");
111 }
112
113 if !config.audience.is_empty() {
114 required_spec_claims.push("aud");
115 validator.set_audience(&config.audience);
116 }
117
118 if !config.issuer.is_empty() {
119 required_spec_claims.push("iss");
120 validator.set_issuer(&config.issuer);
121 }
122
123 validator.set_required_spec_claims(&required_spec_claims);
124 debug!(?validator, "JWT configured");
125
126 Ok(validator)
127}