use std::sync::Arc;
use anyhow::Result;
use chrono::Utc;
use serde::{Deserialize, Serialize};
use surrealdb_types::ToSql;
use uuid::Uuid;
use crate::dbs::{DurableSession, NewPlannerStrategy};
use crate::iam::{Auth, Level, Role};
use crate::types::{PublicValue, PublicVariables};
use crate::val::Value;
#[derive(Clone, Debug)]
pub struct SessionData {
public: Arc<PublicValue>,
internal: Arc<Value>,
digest: PayloadDigest,
}
type PayloadDigest = [u8; 32];
fn payload_digest(value: &PublicValue) -> PayloadDigest {
use sha2::{Digest, Sha256};
fn fold_unordered<'a>(
hasher: &mut sha2::Sha256,
tag: u8,
items: impl ExactSizeIterator<Item = &'a PublicValue>,
) {
hasher.update([tag]);
hasher.update((items.len() as u64).to_be_bytes());
let mut digests: Vec<PayloadDigest> = items.map(payload_digest).collect();
digests.sort_unstable();
for digest in digests {
hasher.update(digest);
}
}
let mut hasher = Sha256::new();
match value {
PublicValue::Array(items) => fold_unordered(&mut hasher, b'A', items.iter()),
PublicValue::Set(items) => fold_unordered(&mut hasher, b'S', items.iter()),
PublicValue::Object(fields) => {
hasher.update([b'O']);
hasher.update((fields.len() as u64).to_be_bytes());
for (key, field) in fields.iter() {
hasher.update((key.len() as u64).to_be_bytes());
hasher.update(key.as_bytes());
hasher.update(payload_digest(field));
}
}
scalar => {
hasher.update([b'V']);
let rendered = scalar.to_sql();
hasher.update((rendered.len() as u64).to_be_bytes());
hasher.update(rendered.as_bytes());
}
}
hasher.finalize().into()
}
impl SessionData {
pub fn new(value: PublicValue) -> Self {
use crate::val::convert_public::convert_public_value_to_internal;
let internal = Arc::new(convert_public_value_to_internal(value.clone()));
let digest = payload_digest(&value);
Self {
public: Arc::new(value),
internal,
digest,
}
}
pub fn public(&self) -> &PublicValue {
&self.public
}
pub(crate) fn internal(&self) -> &Value {
&self.internal
}
}
impl PartialEq for SessionData {
fn eq(&self, other: &Self) -> bool {
self.digest == other.digest
}
}
impl Eq for SessionData {}
impl Serialize for SessionData {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
self.public.serialize(serializer)
}
}
impl<'de> Deserialize<'de> for SessionData {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
Ok(Self::new(PublicValue::deserialize(deserializer)?))
}
}
#[derive(Clone, Debug, Default, Eq, PartialEq, Serialize, Deserialize)]
pub struct Session {
pub au: Arc<Auth>,
pub rt: bool,
pub ip: Option<String>,
pub or: Option<String>,
pub id: Option<Uuid>,
pub ns: Option<String>,
pub db: Option<String>,
pub ac: Option<String>,
pub tk: Option<PublicValue>,
pub rd: Option<PublicValue>,
pub data: Option<SessionData>,
pub exp: Option<i64>,
pub variables: PublicVariables,
pub new_planner_strategy: NewPlannerStrategy,
pub redact_volatile_explain_attrs: bool,
}
impl Session {
pub fn with_ns(mut self, ns: &str) -> Session {
self.ns = Some(ns.to_owned());
self
}
pub fn with_db(mut self, db: &str) -> Session {
self.db = Some(db.to_owned());
self
}
pub fn with_ac(mut self, ac: &str) -> Session {
self.ac = Some(ac.to_owned());
self
}
pub fn with_rt(mut self, rt: bool) -> Session {
self.rt = rt;
self
}
pub fn new_planner_strategy(mut self, strategy: NewPlannerStrategy) -> Session {
self.new_planner_strategy = strategy;
self
}
pub(crate) fn ns(&self) -> Option<Arc<str>> {
self.ns.as_deref().map(Into::into)
}
pub(crate) fn db(&self) -> Option<Arc<str>> {
self.db.as_deref().map(Into::into)
}
pub(crate) fn live(&self) -> bool {
self.rt
}
pub fn expired(&self) -> bool {
match self.exp {
Some(exp) => Utc::now().timestamp() > exp,
None => false,
}
}
pub(crate) fn values(&self) -> Vec<(&'static str, Arc<Value>)> {
use crate::val::convert_public::convert_public_value_to_internal;
let access = self.ac.as_deref().map(Value::from).unwrap_or(Value::None);
let auth = self.rd.clone().map(convert_public_value_to_internal).unwrap_or(Value::None);
let token = self.tk.clone().map(convert_public_value_to_internal).unwrap_or(Value::None);
let session = Value::from(map! {
"ac" => access.clone(),
"exp" => self.exp.map(Value::from).unwrap_or(Value::None),
"db" => self.db.as_deref().map(Value::from).unwrap_or(Value::None),
"id" => self.id.map(Value::from).unwrap_or(Value::None),
"ip" => self.ip.as_deref().map(Value::from).unwrap_or(Value::None),
"ns" => self.ns.as_deref().map(Value::from).unwrap_or(Value::None),
"or" => self.or.as_deref().map(Value::from).unwrap_or(Value::None),
"rd" => auth.clone(),
"tk" => token.clone(),
"data", if let Some(v) = self.data.as_ref() => v.internal().clone(),
});
vec![
("access", Arc::new(access)),
("auth", Arc::new(auth)),
("token", Arc::new(token)),
("session", Arc::new(session)),
]
}
pub fn for_level(level: Level, role: Role) -> Session {
let mut sess = Session::default();
match level {
Level::Root => {
sess.au = Arc::new(Auth::for_root(role));
}
Level::Namespace(ns) => {
sess.au = Arc::new(Auth::for_ns(role, &ns));
sess.ns = Some(ns);
}
Level::Database(ns, db) => {
sess.au = Arc::new(Auth::for_db(role, &ns, &db));
sess.ns = Some(ns);
sess.db = Some(db);
}
_ => {}
}
sess
}
pub fn for_record(ns: &str, db: &str, ac: &str, rid: PublicValue) -> Session {
Session {
ac: Some(ac.to_owned()),
au: Arc::new(Auth::for_record(rid.to_sql(), ns, db, ac)),
rt: false,
ip: None,
or: None,
id: None,
ns: Some(ns.to_owned()),
db: Some(db.to_owned()),
tk: None,
rd: Some(rid),
data: None,
exp: None,
variables: Default::default(),
new_planner_strategy: NewPlannerStrategy::default(),
redact_volatile_explain_attrs: false,
}
}
pub fn owner() -> Session {
Session::for_level(Level::Root, Role::Owner)
}
pub fn editor() -> Session {
Session::for_level(Level::Root, Role::Editor)
}
pub fn viewer() -> Session {
Session::for_level(Level::Root, Role::Viewer)
}
}
#[derive(Clone, Debug)]
pub struct AuthPrincipalSnapshot {
id: String,
level: Level,
data: Option<SessionData>,
}
impl AuthPrincipalSnapshot {
pub fn capture(session: &Session) -> Self {
Self {
id: session.au.id().to_string(),
level: session.au.level().clone(),
data: session.data.clone(),
}
}
pub fn differs_from(&self, session: &Session) -> bool {
session.au.id() != self.id || session.au.level() != &self.level || session.data != self.data
}
}
pub(crate) fn durable_session(session: &Session, expires_at: u64) -> DurableSession {
use crate::val::convert_public::convert_public_value_to_internal;
DurableSession {
expires_at,
au: (*session.au).clone(),
rt: session.rt,
ip: session.ip.clone(),
or: session.or.clone(),
id: session.id,
ns: session.ns.clone(),
db: session.db.clone(),
ac: session.ac.clone(),
tk: session.tk.clone().map(convert_public_value_to_internal),
rd: session.rd.clone().map(convert_public_value_to_internal),
data: session.data.as_ref().map(|d| d.internal().clone()),
exp: session.exp,
variables: session
.variables
.clone()
.into_iter()
.map(|(k, v)| (k, convert_public_value_to_internal(v)))
.collect(),
new_planner_strategy: session.new_planner_strategy,
redact_volatile_explain_attrs: session.redact_volatile_explain_attrs,
}
}
pub(crate) fn restore_session(durable: DurableSession) -> Result<Session> {
use crate::val::convert_value_to_public_value;
Ok(Session {
au: Arc::new(durable.au),
rt: durable.rt,
ip: durable.ip,
or: durable.or,
id: durable.id,
ns: durable.ns,
db: durable.db,
ac: durable.ac,
tk: durable.tk.map(convert_value_to_public_value).transpose()?,
rd: durable.rd.map(convert_value_to_public_value).transpose()?,
data: durable.data.map(convert_value_to_public_value).transpose()?.map(SessionData::new),
exp: durable.exp,
variables: durable
.variables
.into_iter()
.map(|(k, v)| Ok((k.into_string(), convert_value_to_public_value(v)?)))
.collect::<Result<PublicVariables>>()?,
new_planner_strategy: durable.new_planner_strategy,
redact_volatile_explain_attrs: durable.redact_volatile_explain_attrs,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn json_round_trip_preserves_auth_and_context() {
let original = Session {
id: Some(Uuid::from_u128(1)),
exp: Some(1_700_000_000),
..Session::owner().with_ns("app").with_db("app")
};
let json = serde_json::to_string(&original).expect("serialize");
let restored: Session = serde_json::from_str(&json).expect("deserialize");
assert_eq!(original, restored);
assert!(restored.au.is_root());
assert_eq!(restored.ns.as_deref(), Some("app"));
assert_eq!(restored.db.as_deref(), Some("app"));
}
#[test]
fn durable_session_round_trip_preserves_all_fields() {
use std::collections::BTreeMap;
use surrealdb_types::{Number, Value as PV};
let mut variables = PublicVariables::default();
variables.insert("str", PV::String("hello".to_owned()));
variables.insert("dec", PV::Number(Number::Decimal("1.5".parse().unwrap())));
variables.insert("rid", PV::RecordId(surrealdb_types::RecordId::new("person", "tobie")));
variables.insert("dt", PV::Datetime(surrealdb_types::Datetime::now()));
variables.insert(
"obj",
PV::Object(surrealdb_types::Object::from(BTreeMap::from([("nested", PV::Bool(true))]))),
);
let original = Session {
rt: true,
ip: Some("10.0.0.1".to_owned()),
or: Some("example.com".to_owned()),
id: Some(Uuid::from_u128(7)),
exp: Some(1_700_000_000),
tk: Some(PV::Object(surrealdb_types::Object::from(BTreeMap::from([(
"iss",
PV::String("surrealdb".to_owned()),
)])))),
rd: Some(PV::RecordId(surrealdb_types::RecordId::new("person", "tobie"))),
data: Some(SessionData::new(PV::Object(surrealdb_types::Object::from(
BTreeMap::from([(
"orgs",
PV::Array(surrealdb_types::Array::from(vec![
PV::String("acme".to_owned()),
PV::String("globex".to_owned()),
])),
)]),
)))),
variables,
redact_volatile_explain_attrs: true,
..Session::for_record(
"app",
"app",
"account",
PublicValue::RecordId(surrealdb_types::RecordId::new("person", "tobie")),
)
};
let durable = durable_session(&original, 123_456_789);
assert_eq!(durable.expires_at, 123_456_789);
let bytes = revision::to_vec(&durable).expect("revision encode");
let decoded: DurableSession = revision::from_slice(&bytes).expect("revision decode");
assert_eq!(durable, decoded);
let restored = restore_session(decoded).expect("convert back");
assert_eq!(original, restored);
}
#[test]
fn session_values_only_carry_data_when_a_context_clause_populated_it() {
use surrealdb_types::Value as PV;
fn session_keys(session: &Session) -> Vec<String> {
let values = session.values();
let (_, session_value) =
values.iter().find(|(k, _)| *k == "session").expect("a `session` entry");
let obj = match session_value.as_ref() {
Value::Object(o) => o,
other => panic!("expected an object, got {other:?}"),
};
let mut keys: Vec<String> = obj.keys().map(|k| k.as_str().to_owned()).collect();
keys.sort();
keys
}
let mut session = Session::for_record(
"app",
"app",
"account",
PublicValue::RecordId(surrealdb_types::RecordId::new("person", "tobie")),
);
assert_eq!(
session_keys(&session),
vec!["ac", "db", "exp", "id", "ip", "ns", "or", "rd", "tk"]
);
session.data = Some(SessionData::new(PV::String("payload".to_owned())));
assert_eq!(
session_keys(&session),
vec!["ac", "data", "db", "exp", "id", "ip", "ns", "or", "rd", "tk"]
);
let values = session.values();
let (_, session_value) = values.iter().find(|(k, _)| *k == "session").unwrap();
let Value::Object(obj) = session_value.as_ref() else {
panic!("expected an object")
};
assert_eq!(obj.get("data"), Some(&Value::from("payload")));
}
#[test]
fn a_reordered_payload_is_the_same_payload() {
use surrealdb_types::{Array, Object, Value as PV};
fn array(items: &[&str]) -> PV {
PV::Array(Array::from(
items.iter().map(|s| PV::String((*s).to_owned())).collect::<Vec<_>>(),
))
}
assert_eq!(
SessionData::new(array(&["acme", "globex"])),
SessionData::new(array(&["globex", "acme"]))
);
let nested = |orgs: PV| {
let mut obj = Object::default();
obj.insert("orgs", orgs);
SessionData::new(PV::Object(obj))
};
assert_eq!(nested(array(&["acme", "globex"])), nested(array(&["globex", "acme"])));
assert_ne!(
SessionData::new(array(&["acme", "globex"])),
SessionData::new(array(&["acme", "initech"]))
);
assert_ne!(SessionData::new(array(&["acme"])), SessionData::new(array(&["acme", "acme"])));
assert_ne!(
SessionData::new(array(&["acme"])),
SessionData::new(PV::String("acme".to_owned()))
);
assert_ne!(
SessionData::new(PV::String("1".to_owned())),
SessionData::new(PV::Number(surrealdb_types::Number::Int(1)))
);
}
#[test]
fn auth_principal_snapshot_tracks_the_context_payload() {
use surrealdb_types::Value as PV;
let mut session = Session::for_record(
"app",
"app",
"account",
PublicValue::RecordId(surrealdb_types::RecordId::new("person", "tobie")),
);
session.data = Some(SessionData::new(PV::String("acme".to_owned())));
let snapshot = AuthPrincipalSnapshot::capture(&session);
assert!(!snapshot.differs_from(&session));
session.data = Some(SessionData::new(PV::String("acme".to_owned())));
assert!(!snapshot.differs_from(&session));
session.data = Some(SessionData::new(PV::String("globex".to_owned())));
assert!(snapshot.differs_from(&session));
session.data = None;
assert!(snapshot.differs_from(&session));
}
}