use async_trait::async_trait;
use sea_orm::{EntityTrait, PrimaryKeyTrait};
use std::marker::PhantomData;
use std::sync::Arc;
use super::authenticatable::Authenticatable;
use crate::database::DB;
use crate::error::FrameworkError;
#[async_trait]
pub trait UserProvider: Send + Sync + 'static {
async fn retrieve_by_id(
&self,
id: i64,
) -> Result<Option<Arc<dyn Authenticatable>>, FrameworkError>;
async fn retrieve_by_credentials(
&self,
_credentials: &serde_json::Value,
) -> Result<Option<Arc<dyn Authenticatable>>, FrameworkError> {
Ok(None)
}
async fn validate_credentials(
&self,
_user: &dyn Authenticatable,
_credentials: &serde_json::Value,
) -> Result<bool, FrameworkError> {
Ok(false)
}
}
pub struct ModelUserProvider<E: EntityTrait> {
_marker: PhantomData<fn() -> E>,
}
impl<E: EntityTrait> ModelUserProvider<E> {
pub fn new() -> Self {
Self {
_marker: PhantomData,
}
}
}
impl<E: EntityTrait> Default for ModelUserProvider<E> {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl<E> UserProvider for ModelUserProvider<E>
where
E: EntityTrait + 'static,
E::Model: Authenticatable + Clone,
<E::PrimaryKey as PrimaryKeyTrait>::ValueType: TryFrom<i64> + Send,
{
async fn retrieve_by_id(
&self,
id: i64,
) -> Result<Option<Arc<dyn Authenticatable>>, FrameworkError> {
let pk = <<E as EntityTrait>::PrimaryKey as PrimaryKeyTrait>::ValueType::try_from(id)
.map_err(|_| {
FrameworkError::internal(format!(
"authenticated id {id} is out of range for {}'s primary key",
std::any::type_name::<E>()
))
})?;
let db = DB::connection()?;
let model = E::find_by_id(pk)
.one(db.inner())
.await
.map_err(|e| FrameworkError::database(e.to_string()))?;
Ok(model.map(|m| Arc::new(m) as Arc<dyn Authenticatable>))
}
}