use std::{fmt, marker::PhantomData, sync::Arc};
use async_trait::async_trait;
use axum::{
extract::FromRequestParts,
http::{StatusCode, request::Parts},
};
use axum_session_auth::Authentication;
use surrealdb::{Connection, Surreal};
use surrealdb_types::{RecordId, SurrealValue, ToSql};
use super::SurrealSessionPool;
pub trait AuthUser: Send + Sync + 'static {
const TABLE: &'static str = "user";
const NAME_FIELD: &'static str = "display_name";
const FILTER: Option<&'static str> = None;
}
pub struct SessionUser<U> {
pub id: String,
pub anonymous: bool,
pub username: String,
user: PhantomData<fn() -> U>,
}
impl<U> SessionUser<U> {
pub fn signed_in(id: impl Into<String>, username: impl Into<String>) -> Self {
Self {
id: id.into(),
anonymous: false,
username: username.into(),
user: PhantomData,
}
}
pub fn anonymous() -> Self {
Self {
id: String::new(),
anonymous: true,
username: String::new(),
user: PhantomData,
}
}
}
impl<U: AuthUser> SessionUser<U> {
pub fn record_id(&self) -> RecordId {
RecordId::new(U::TABLE, self.id.clone())
}
}
impl<U> Default for SessionUser<U> {
fn default() -> Self {
Self::anonymous()
}
}
impl<U> Clone for SessionUser<U> {
fn clone(&self) -> Self {
Self {
id: self.id.clone(),
anonymous: self.anonymous,
username: self.username.clone(),
user: PhantomData,
}
}
}
impl<U> fmt::Debug for SessionUser<U> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SessionUser")
.field("id", &self.id)
.field("anonymous", &self.anonymous)
.field("username", &self.username)
.finish()
}
}
#[derive(SurrealValue)]
struct LoadedUser {
id: RecordId,
username: Option<String>,
}
#[async_trait]
impl<U, C> Authentication<SessionUser<U>, String, Arc<Surreal<C>>> for SessionUser<U>
where
U: AuthUser,
C: Connection,
{
async fn load_user(
userid: String,
db: Option<&Arc<Surreal<C>>>,
) -> anyhow::Result<SessionUser<U>> {
let db = db.ok_or_else(|| anyhow::anyhow!("Database connection not provided"))?;
let filter = U::FILTER
.map(|filter| format!(" WHERE {filter}"))
.unwrap_or_default();
let user = db
.query(format!(
"SELECT id, {name} AS username FROM $record{filter}",
name = U::NAME_FIELD,
))
.bind(("record", RecordId::new(U::TABLE, userid)))
.await?
.take::<Option<LoadedUser>>(0)?
.ok_or_else(|| anyhow::anyhow!("User not found"))?;
Ok(SessionUser::signed_in(
user.id.key.to_sql(),
user.username.unwrap_or_default(),
))
}
fn is_authenticated(&self) -> bool {
!self.anonymous
}
fn is_active(&self) -> bool {
!self.anonymous
}
fn is_anonymous(&self) -> bool {
self.anonymous
}
}
pub type AuthSession<U, C> =
axum_session_auth::AuthSession<SessionUser<U>, String, SurrealSessionPool<C>, Arc<Surreal<C>>>;
pub type AuthSessionLayer<U, C> = axum_session_auth::AuthSessionLayer<
SessionUser<U>,
String,
SurrealSessionPool<C>,
Arc<Surreal<C>>,
>;
pub struct SessionContext<U: AuthUser, C: Connection> {
pub db: Arc<Surreal<C>>,
pub auth_session: AuthSession<U, C>,
pub session_user: SessionUser<U>,
}
impl<S, U, C> FromRequestParts<S> for SessionContext<U, C>
where
S: Send + Sync,
U: AuthUser,
C: Connection,
{
type Rejection = (StatusCode, &'static str);
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
let db = parts
.extensions
.get::<Arc<Surreal<C>>>()
.ok_or((
StatusCode::INTERNAL_SERVER_ERROR,
"Database missing. Is `Extension(Arc<Surreal<_>>)` on the router?",
))?
.clone();
let auth_session = parts
.extensions
.get::<AuthSession<U, C>>()
.ok_or((
StatusCode::INTERNAL_SERVER_ERROR,
"Auth session missing. Is `AuthSessionLayer` installed?",
))?
.clone();
let session_user = auth_session.current_user.clone().unwrap_or_default();
Ok(Self {
db,
auth_session,
session_user,
})
}
}