use std::collections::HashMap;
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::expr::Error;
use crate::kvs::Datastore;
use crate::val::convert_public::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)?)
}
use surrealdb_rpc::Token;
pub async fn refresh(token: Token, kvs: &Datastore, session: &mut Session) -> Result<Token> {
match token {
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(token: Token, kvs: &Datastore) -> Result<()> {
match token {
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::verify_and_revoke_refresh_token_record(
kvs, grant_id, ac, refresh, &ns, &db,
)
.await?;
Ok(())
}
}
}
#[derive(Debug, Serialize, Deserialize, Clone)]
#[serde(untagged)]
pub enum Audience {
Single(String),
Multiple(Vec<String>),
}
impl Audience {
pub(crate) fn from_configured(list: &[String]) -> Self {
match list {
[single] => Audience::Single(single.clone()),
many => Audience::Multiple(many.to_vec()),
}
}
}
#[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
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use super::*;
use crate::dbs::Session;
use crate::iam::signin::db_access;
use crate::kvs::Datastore;
use crate::types::PublicVariables;
async fn refresh_token_fixture() -> (Arc<Datastore>, String, String) {
let ds = Datastore::new("memory").await.unwrap();
let sess = Session::owner().with_ns("test").with_db("test");
ds.execute(
r#"
DEFINE ACCESS user ON DATABASE TYPE RECORD
SIGNIN (
SELECT * FROM user WHERE name = $user AND crypto::argon2::compare(pass, $pass)
)
WITH REFRESH
DURATION FOR GRANT 1w, FOR SESSION 2h
;
CREATE user:test CONTENT {
name: 'user',
pass: crypto::argon2::generate('pass')
}
"#,
&sess,
None,
)
.await
.unwrap();
let mut sess = Session {
ns: Some("test".to_string()),
db: Some("test".to_string()),
..Default::default()
};
let mut vars = PublicVariables::new();
vars.insert("user", "user");
vars.insert("pass", "pass");
let token = db_access(
&ds,
&mut sess,
"test".to_string(),
"test".to_string(),
"user".to_string(),
vars,
)
.await
.expect("signin with credentials should succeed");
match token {
Token::WithRefresh {
access,
refresh,
} => (ds, access, refresh),
Token::Access(_) => panic!("a WITH REFRESH access method must return a refresh token"),
}
}
fn with_wrong_key(refresh: &str) -> String {
let parts: Vec<&str> = refresh.split('-').collect();
assert_eq!(parts.len(), 4, "a bearer token has four dash-separated parts");
format!("{}-{}-{}-{}", parts[0], parts[1], parts[2], "A".repeat(parts[3].len()))
}
#[tokio::test]
async fn revoke_refresh_token_requires_the_grant_key() {
let (ds, access, refresh) = refresh_token_fixture().await;
let forged = with_wrong_key(&refresh);
assert_ne!(forged, refresh, "the forged token must differ from the real one");
let res = revoke_refresh_token(
Token::WithRefresh {
access: access.clone(),
refresh: forged,
},
&ds,
)
.await;
assert!(res.is_err(), "revoking with a wrong key must fail");
let mut sess = Session {
ns: Some("test".to_string()),
db: Some("test".to_string()),
..Default::default()
};
let mut vars = PublicVariables::new();
vars.insert("refresh", refresh.clone());
let res = db_access(
&ds,
&mut sess,
"test".to_string(),
"test".to_string(),
"user".to_string(),
vars,
)
.await;
assert!(res.is_ok(), "the victim's refresh token must still be usable: {res:?}");
}
#[tokio::test]
async fn revoke_refresh_token_succeeds_for_the_key_holder() {
let (ds, access, refresh) = refresh_token_fixture().await;
revoke_refresh_token(
Token::WithRefresh {
access,
refresh: refresh.clone(),
},
&ds,
)
.await
.expect("the key holder must be able to revoke their own refresh token");
let mut sess = Session {
ns: Some("test".to_string()),
db: Some("test".to_string()),
..Default::default()
};
let mut vars = PublicVariables::new();
vars.insert("refresh", refresh);
let res = db_access(
&ds,
&mut sess,
"test".to_string(),
"test".to_string(),
"user".to_string(),
vars,
)
.await;
assert!(res.is_err(), "a revoked refresh token must no longer be usable");
}
}