1use base64::Engine;
9use base64::engine::general_purpose::URL_SAFE_NO_PAD;
10use serde::Deserialize;
11
12#[derive(Deserialize)]
13struct Claims {
14 #[serde(default)]
15 tid: Option<String>,
16}
17
18#[derive(Debug, Clone, PartialEq, Eq)]
23pub struct IdTokenClaims {
24 pub tid: String,
25 pub oid: String,
26}
27
28#[derive(Deserialize)]
29struct IdentityClaims {
30 #[serde(default)]
31 tid: Option<String>,
32 #[serde(default)]
33 oid: Option<String>,
34}
35
36pub fn extract_id_claims(jwt: &str) -> Option<IdTokenClaims> {
41 let mid = jwt.split('.').nth(1)?;
42 let bytes = URL_SAFE_NO_PAD.decode(mid).ok()?;
43 let claims: IdentityClaims = serde_json::from_slice(&bytes).ok()?;
44 match (claims.tid, claims.oid) {
45 (Some(tid), Some(oid)) if !tid.is_empty() && !oid.is_empty() => {
46 Some(IdTokenClaims { tid, oid })
47 }
48 _ => None,
49 }
50}
51
52pub fn extract_tenant_id(jwt: &str) -> Option<String> {
55 let mid = jwt.split('.').nth(1)?;
56 let bytes = URL_SAFE_NO_PAD.decode(mid).ok()?;
57 let claims: Claims = serde_json::from_slice(&bytes).ok()?;
58 claims.tid
59}
60
61#[cfg(test)]
62mod tests {
63 use super::*;
64
65 fn make_jwt(payload_json: &str) -> String {
66 let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"RS256","typ":"JWT"}"#);
67 let payload = URL_SAFE_NO_PAD.encode(payload_json);
68 let signature = URL_SAFE_NO_PAD.encode("dummy-signature");
69 format!("{header}.{payload}.{signature}")
70 }
71
72 #[test]
73 fn extracts_tid_from_valid_jwt() {
74 let jwt = make_jwt(
75 r#"{"tid":"11111111-2222-3333-4444-555555555555","iss":"https://login.microsoftonline.com/.."}"#,
76 );
77 assert_eq!(
78 extract_tenant_id(&jwt),
79 Some("11111111-2222-3333-4444-555555555555".to_string())
80 );
81 }
82
83 #[test]
84 fn returns_none_when_tid_is_missing() {
85 let jwt = make_jwt(r#"{"iss":"https://login.microsoftonline.com/.."}"#);
86 assert_eq!(extract_tenant_id(&jwt), None);
87 }
88
89 #[test]
90 fn extracts_tid_and_oid_together() {
91 let jwt = make_jwt(r#"{"tid":"t-1","oid":"o-1","email":"x@example.com"}"#);
92 assert_eq!(
93 extract_id_claims(&jwt),
94 Some(IdTokenClaims {
95 tid: "t-1".into(),
96 oid: "o-1".into()
97 })
98 );
99 }
100
101 #[test]
102 fn id_claims_need_both_tid_and_oid() {
103 assert_eq!(extract_id_claims(&make_jwt(r#"{"tid":"t-1"}"#)), None);
104 assert_eq!(extract_id_claims(&make_jwt(r#"{"oid":"o-1"}"#)), None);
105 assert_eq!(
106 extract_id_claims(&make_jwt(r#"{"tid":"","oid":"o-1"}"#)),
107 None
108 );
109 assert_eq!(extract_id_claims("not-a-jwt"), None);
110 }
111
112 #[test]
113 fn returns_none_for_malformed_jwt() {
114 assert_eq!(extract_tenant_id("not-a-jwt"), None);
115 }
116
117 #[test]
118 fn returns_none_for_non_base64_middle_segment() {
119 assert_eq!(
120 extract_tenant_id("a.this-isn't-base64-because-of-?-char.b"),
121 None
122 );
123 }
124}