1use std::collections::HashMap;
2use std::fmt;
3use std::sync::LazyLock;
4
5use anyhow::{Result, bail};
6use jsonwebtoken::{Algorithm, Header};
7use serde::{Deserialize, Serialize};
8use surrealdb_types::SurrealValue;
9
10use crate::dbs::Session;
11use crate::err::Error;
12use crate::kvs::Datastore;
13use crate::sql::expression::convert_public_value_to_internal;
14use crate::val::{Object, Value, convert_object_to_public_map};
15use crate::{iam, syn};
16pub static HEADER: LazyLock<Header> = LazyLock::new(|| Header::new(Algorithm::HS512));
17
18fn decode_access_token_claims(token: &str) -> Result<jsonwebtoken::TokenData<Claims>> {
25 Ok(jsonwebtoken::dangerous::insecure_decode::<Claims>(token)?)
26}
27
28#[derive(Clone, Eq, PartialEq, PartialOrd, SurrealValue, Hash)]
68#[surreal(crate = "surrealdb_types")]
69#[surreal(untagged)]
70#[cfg_attr(feature = "arbitrary", derive(arbitrary::Arbitrary))]
71pub enum Token {
72 Access(String),
77 WithRefresh {
83 access: String,
85 refresh: String,
87 },
88}
89
90impl Token {
91 pub async fn refresh(self, kvs: &Datastore, session: &mut Session) -> Result<Self> {
150 match self {
151 Token::Access(_) => bail!(Error::InvalidFunctionArguments {
152 name: "refresh".into(),
153 message: "Token is an access token, cannot refresh".into(),
154 }),
155 Token::WithRefresh {
156 access,
157 refresh,
158 } => {
159 let token_data = decode_access_token_claims(&access)?;
165 let claims = token_data.claims.into_claims_object();
166 let mut vars = convert_object_to_public_map(claims)?;
170 vars.insert("refresh".to_string(), refresh.into_value());
173 iam::signin::signin(kvs, session, vars.into()).await
179 }
180 }
181 }
182
183 pub async fn revoke_refresh_token(self, kvs: &Datastore) -> Result<()> {
184 match self {
185 Token::Access(_) => bail!(Error::InvalidFunctionArguments {
186 name: "refresh".into(),
187 message: "Token is an access token, cannot revoke refresh token".into(),
188 }),
189 Token::WithRefresh {
190 access,
191 refresh,
192 } => {
193 let grant_id = iam::signin::validate_grant_bearer(&refresh)?;
194 let token_data = decode_access_token_claims(&access)?;
195 let ns = token_data.claims.ns.ok_or_else(|| Error::InvalidFunctionArguments {
196 name: "ns".into(),
197 message: "Token does not contain a namespace".into(),
198 })?;
199 let db = token_data.claims.db.ok_or_else(|| Error::InvalidFunctionArguments {
200 name: "db".into(),
201 message: "Token does not contain a database".into(),
202 })?;
203 let ac = token_data.claims.ac.ok_or_else(|| Error::InvalidFunctionArguments {
204 name: "ac".into(),
205 message: "Token does not contain an access name".into(),
206 })?;
207 iam::access::revoke_refresh_token_record(kvs, grant_id, ac, &ns, &db).await?;
208 Ok(())
209 }
210 }
211 }
212}
213
214impl fmt::Debug for Token {
215 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
216 match self {
217 Token::Access(_) => write!(f, "Token::Access(REDACTED)"),
218 Token::WithRefresh {
219 ..
220 } => write!(f, "Token::WithRefresh {{ access: REDACTED, refresh: REDACTED }}"),
221 }
222 }
223}
224
225#[derive(Debug, Serialize, Deserialize, Clone)]
226#[serde(untagged)]
227pub enum Audience {
228 Single(String),
229 Multiple(Vec<String>),
230}
231
232#[derive(Debug, Default, Serialize, Deserialize, Clone)]
233pub struct Claims {
234 #[serde(skip_serializing_if = "Option::is_none")]
235 pub iat: Option<i64>,
236 #[serde(skip_serializing_if = "Option::is_none")]
237 pub nbf: Option<i64>,
238 #[serde(skip_serializing_if = "Option::is_none")]
239 pub exp: Option<i64>,
240 #[serde(skip_serializing_if = "Option::is_none")]
241 pub iss: Option<String>,
242 #[serde(skip_serializing_if = "Option::is_none")]
243 pub sub: Option<String>,
244 #[serde(skip_serializing_if = "Option::is_none")]
245 pub aud: Option<Audience>,
246 #[serde(skip_serializing_if = "Option::is_none")]
247 pub jti: Option<String>,
248 #[serde(alias = "ns")]
249 #[serde(alias = "NS")]
250 #[serde(rename = "NS")]
251 #[serde(alias = "https://surrealdb.com/ns")]
252 #[serde(alias = "https://surrealdb.com/namespace")]
253 #[serde(skip_serializing_if = "Option::is_none")]
254 pub ns: Option<String>,
255 #[serde(alias = "db")]
256 #[serde(alias = "DB")]
257 #[serde(rename = "DB")]
258 #[serde(alias = "https://surrealdb.com/db")]
259 #[serde(alias = "https://surrealdb.com/database")]
260 #[serde(skip_serializing_if = "Option::is_none")]
261 pub db: Option<String>,
262 #[serde(alias = "ac")]
263 #[serde(alias = "AC")]
264 #[serde(rename = "AC")]
265 #[serde(alias = "https://surrealdb.com/ac")]
266 #[serde(alias = "https://surrealdb.com/access")]
267 #[serde(skip_serializing_if = "Option::is_none")]
268 pub ac: Option<String>,
269 #[serde(alias = "id")]
270 #[serde(alias = "ID")]
271 #[serde(rename = "ID")]
272 #[serde(alias = "https://surrealdb.com/id")]
273 #[serde(alias = "https://surrealdb.com/record")]
274 #[serde(skip_serializing_if = "Option::is_none")]
275 pub id: Option<String>,
276 #[serde(alias = "rl")]
277 #[serde(alias = "RL")]
278 #[serde(rename = "RL")]
279 #[serde(alias = "https://surrealdb.com/rl")]
280 #[serde(alias = "https://surrealdb.com/roles")]
281 #[serde(skip_serializing_if = "Option::is_none")]
282 pub roles: Option<Vec<String>>,
283
284 #[serde(flatten)]
285 #[serde(skip_serializing_if = "Option::is_none")]
286 pub custom_claims: Option<HashMap<String, serde_json::Value>>,
287}
288
289impl Claims {
290 pub(crate) fn into_claims_object(self) -> Object {
291 let mut out = Object::default();
293 if let Some(iss) = self.iss {
295 out.insert("iss", iss.into());
296 }
297 if let Some(sub) = self.sub {
299 out.insert("sub", sub.into());
300 }
301 if let Some(aud) = self.aud {
303 match aud {
304 Audience::Single(v) => out.insert("aud", Value::String(v.into())),
305 Audience::Multiple(v) => {
306 out.insert("aud", v.into_iter().map(Value::from).collect::<Vec<_>>().into())
307 }
308 };
309 }
310 if let Some(iat) = self.iat {
312 out.insert("iat", iat.into());
313 }
314 if let Some(nbf) = self.nbf {
316 out.insert("nbf", nbf.into());
317 }
318 if let Some(exp) = self.exp {
320 out.insert("exp", exp.into());
321 }
322 if let Some(jti) = self.jti {
324 out.insert("jti", jti.into());
325 }
326 if let Some(ns) = self.ns {
328 out.insert("NS", ns.into());
329 }
330 if let Some(db) = self.db {
332 out.insert("DB", db.into());
333 }
334 if let Some(ac) = self.ac {
336 out.insert("AC", ac.into());
337 }
338 if let Some(id) = self.id {
340 out.insert("ID", id.into());
341 }
342 if let Some(role) = self.roles {
344 out.insert("RL", role.into_iter().map(Value::from).collect::<Vec<_>>().into());
345 }
346 if let Some(custom_claims) = self.custom_claims {
348 for (claim, value) in custom_claims {
349 let claim_json = match serde_json::to_string(&value) {
351 Ok(claim_json) => claim_json,
352 Err(err) => {
353 debug!("Failed to serialize token claim '{}': {}", claim, err);
354 continue;
355 }
356 };
357 let claim_value = match syn::json(&claim_json) {
359 Ok(claim_value) => claim_value,
360 Err(err) => {
361 debug!("Failed to parse token claim '{}': {}", claim, err);
362 continue;
363 }
364 };
365 let claim_value = convert_public_value_to_internal(claim_value);
366 out.insert(claim.clone(), claim_value);
367 }
368 }
369 out
371 }
372}