use std::collections::HashMap;
use std::fmt;
use std::sync::LazyLock;
use anyhow::{Result, bail};
use jsonwebtoken::{Algorithm, Header};
use serde::{Deserialize, Serialize};
use surrealdb_types::SurrealValue;
use crate::dbs::Session;
use crate::err::Error;
use crate::kvs::Datastore;
use crate::sql::expression::convert_public_value_to_internal;
use crate::val::{Object, Value, convert_object_to_public_map};
use crate::{iam, syn};
pub static HEADER: LazyLock<Header> = LazyLock::new(|| Header::new(Algorithm::HS512));
fn decode_access_token_claims(token: &str) -> Result<jsonwebtoken::TokenData<Claims>> {
Ok(jsonwebtoken::dangerous::insecure_decode::<Claims>(token)?)
}
#[derive(Clone, Eq, PartialEq, PartialOrd, SurrealValue, Hash)]
#[surreal(crate = "surrealdb_types")]
#[surreal(untagged)]
#[cfg_attr(feature = "arbitrary", derive(arbitrary::Arbitrary))]
pub enum Token {
Access(String),
WithRefresh {
access: String,
refresh: String,
},
}
impl Token {
pub async fn refresh(self, kvs: &Datastore, session: &mut Session) -> Result<Self> {
match self {
Token::Access(_) => bail!(Error::InvalidFunctionArguments {
name: "refresh".into(),
message: "Token is an access token, cannot refresh".into(),
}),
Token::WithRefresh {
access,
refresh,
} => {
let token_data = decode_access_token_claims(&access)?;
let claims = token_data.claims.into_claims_object();
let mut vars = convert_object_to_public_map(claims)?;
vars.insert("refresh".to_string(), refresh.into_value());
iam::signin::signin(kvs, session, vars.into()).await
}
}
}
pub async fn revoke_refresh_token(self, kvs: &Datastore) -> Result<()> {
match self {
Token::Access(_) => bail!(Error::InvalidFunctionArguments {
name: "refresh".into(),
message: "Token is an access token, cannot revoke refresh token".into(),
}),
Token::WithRefresh {
access,
refresh,
} => {
let grant_id = iam::signin::validate_grant_bearer(&refresh)?;
let token_data = decode_access_token_claims(&access)?;
let ns = token_data.claims.ns.ok_or_else(|| Error::InvalidFunctionArguments {
name: "ns".into(),
message: "Token does not contain a namespace".into(),
})?;
let db = token_data.claims.db.ok_or_else(|| Error::InvalidFunctionArguments {
name: "db".into(),
message: "Token does not contain a database".into(),
})?;
let ac = token_data.claims.ac.ok_or_else(|| Error::InvalidFunctionArguments {
name: "ac".into(),
message: "Token does not contain an access name".into(),
})?;
iam::access::revoke_refresh_token_record(kvs, grant_id, ac, &ns, &db).await?;
Ok(())
}
}
}
}
impl fmt::Debug for Token {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Token::Access(_) => write!(f, "Token::Access(REDACTED)"),
Token::WithRefresh {
..
} => write!(f, "Token::WithRefresh {{ access: REDACTED, refresh: REDACTED }}"),
}
}
}
#[derive(Debug, Serialize, Deserialize, Clone)]
#[serde(untagged)]
pub enum Audience {
Single(String),
Multiple(Vec<String>),
}
#[derive(Debug, Default, Serialize, Deserialize, Clone)]
pub struct Claims {
#[serde(skip_serializing_if = "Option::is_none")]
pub iat: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub nbf: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub exp: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub iss: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub sub: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub aud: Option<Audience>,
#[serde(skip_serializing_if = "Option::is_none")]
pub jti: Option<String>,
#[serde(alias = "ns")]
#[serde(alias = "NS")]
#[serde(rename = "NS")]
#[serde(alias = "https://surrealdb.com/ns")]
#[serde(alias = "https://surrealdb.com/namespace")]
#[serde(skip_serializing_if = "Option::is_none")]
pub ns: Option<String>,
#[serde(alias = "db")]
#[serde(alias = "DB")]
#[serde(rename = "DB")]
#[serde(alias = "https://surrealdb.com/db")]
#[serde(alias = "https://surrealdb.com/database")]
#[serde(skip_serializing_if = "Option::is_none")]
pub db: Option<String>,
#[serde(alias = "ac")]
#[serde(alias = "AC")]
#[serde(rename = "AC")]
#[serde(alias = "https://surrealdb.com/ac")]
#[serde(alias = "https://surrealdb.com/access")]
#[serde(skip_serializing_if = "Option::is_none")]
pub ac: Option<String>,
#[serde(alias = "id")]
#[serde(alias = "ID")]
#[serde(rename = "ID")]
#[serde(alias = "https://surrealdb.com/id")]
#[serde(alias = "https://surrealdb.com/record")]
#[serde(skip_serializing_if = "Option::is_none")]
pub id: Option<String>,
#[serde(alias = "rl")]
#[serde(alias = "RL")]
#[serde(rename = "RL")]
#[serde(alias = "https://surrealdb.com/rl")]
#[serde(alias = "https://surrealdb.com/roles")]
#[serde(skip_serializing_if = "Option::is_none")]
pub roles: Option<Vec<String>>,
#[serde(flatten)]
#[serde(skip_serializing_if = "Option::is_none")]
pub custom_claims: Option<HashMap<String, serde_json::Value>>,
}
impl Claims {
pub(crate) fn into_claims_object(self) -> Object {
let mut out = Object::default();
if let Some(iss) = self.iss {
out.insert("iss", iss.into());
}
if let Some(sub) = self.sub {
out.insert("sub", sub.into());
}
if let Some(aud) = self.aud {
match aud {
Audience::Single(v) => out.insert("aud", Value::String(v.into())),
Audience::Multiple(v) => {
out.insert("aud", v.into_iter().map(Value::from).collect::<Vec<_>>().into())
}
};
}
if let Some(iat) = self.iat {
out.insert("iat", iat.into());
}
if let Some(nbf) = self.nbf {
out.insert("nbf", nbf.into());
}
if let Some(exp) = self.exp {
out.insert("exp", exp.into());
}
if let Some(jti) = self.jti {
out.insert("jti", jti.into());
}
if let Some(ns) = self.ns {
out.insert("NS", ns.into());
}
if let Some(db) = self.db {
out.insert("DB", db.into());
}
if let Some(ac) = self.ac {
out.insert("AC", ac.into());
}
if let Some(id) = self.id {
out.insert("ID", id.into());
}
if let Some(role) = self.roles {
out.insert("RL", role.into_iter().map(Value::from).collect::<Vec<_>>().into());
}
if let Some(custom_claims) = self.custom_claims {
for (claim, value) in custom_claims {
let claim_json = match serde_json::to_string(&value) {
Ok(claim_json) => claim_json,
Err(err) => {
debug!("Failed to serialize token claim '{}': {}", claim, err);
continue;
}
};
let claim_value = match syn::json(&claim_json) {
Ok(claim_value) => claim_value,
Err(err) => {
debug!("Failed to parse token claim '{}': {}", claim, err);
continue;
}
};
let claim_value = convert_public_value_to_internal(claim_value);
out.insert(claim.clone(), claim_value);
}
}
out
}
}