use std::cell::RefCell;
use std::collections::HashMap;
use std::future::Future;
use std::sync::{Arc, RwLock};
use crate::token::{TokenInfo, TokenValue};
#[derive(Debug, Default)]
pub struct SaTokenContextInner {
pub token: Option<TokenValue>,
pub token_info: Option<Arc<TokenInfo>>,
pub login_id: Option<String>,
pub switch_login_id: Option<String>,
pub auth_meta: RequestAuthMeta,
}
#[derive(Debug, Clone, Default)]
pub struct RequestAuthMeta {
pub authorization: Option<String>,
pub same_token: Option<String>,
}
impl RequestAuthMeta {
pub fn from_request<R: sa_token_adapter::context::SaRequest>(
req: &R,
same_token_header: &str,
) -> Self {
let authorization = req
.get_header("Authorization")
.or_else(|| req.get_header("authorization"));
let same_token = req.get_header(same_token_header).or_else(|| {
if same_token_header.eq_ignore_ascii_case("SA-SAME-TOKEN") {
req.get_header("sa-same-token")
} else {
None
}
});
Self {
authorization,
same_token,
}
}
}
thread_local! {
static TLS_CTX: RefCell<Option<SaTokenContext>> = const { RefCell::new(None) };
static TLS_GRANTS: RefCell<Option<GrantScope>> = const { RefCell::new(None) };
}
tokio::task_local! {
static TASK_CTX: SaTokenContext;
static TASK_GRANTS: GrantScope;
}
#[derive(Debug, Clone, Default)]
pub struct GrantScope {
entries: Arc<RwLock<HashMap<String, Arc<[String]>>>>,
}
impl GrantScope {
pub fn new() -> Self {
Self::default()
}
pub fn get(&self, key: &str) -> Option<Arc<[String]>> {
let guard = self.entries.read().ok()?;
guard.get(key).map(Arc::clone)
}
pub fn put(&self, key: String, value: Arc<[String]>) {
if let Ok(mut guard) = self.entries.write() {
guard.insert(key, value);
}
}
pub fn remove(&self, key: &str) {
if let Ok(mut guard) = self.entries.write() {
guard.remove(key);
}
}
pub fn clear(&self) {
if let Ok(mut guard) = self.entries.write() {
guard.clear();
}
}
pub async fn run<F, T>(scope: Self, future: F) -> T
where
F: std::future::Future<Output = T>,
{
TASK_GRANTS.scope(scope, future).await
}
}
#[derive(Clone)]
pub struct SaTokenContext {
inner: Arc<RwLock<SaTokenContextInner>>,
}
impl std::fmt::Debug for SaTokenContext {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self.inner.read() {
Ok(guard) => f
.debug_struct("SaTokenContext")
.field("inner", &*guard)
.finish(),
Err(_) => f
.debug_struct("SaTokenContext")
.field("inner", &"<poisoned>")
.finish(),
}
}
}
impl SaTokenContext {
pub fn new() -> Self {
Self {
inner: Arc::new(RwLock::new(SaTokenContextInner::default())),
}
}
pub fn builder() -> SaTokenContextBuilder {
SaTokenContextBuilder::new()
}
pub fn token(&self) -> Option<TokenValue> {
Self::read_inner(&self.inner).token.clone()
}
pub fn token_info(&self) -> Option<Arc<TokenInfo>> {
Self::read_inner(&self.inner).token_info.clone()
}
pub fn login_id(&self) -> Option<String> {
Self::read_inner(&self.inner).login_id.clone()
}
pub fn switch_login_id(&self) -> Option<String> {
Self::read_inner(&self.inner).switch_login_id.clone()
}
pub fn auth_meta(&self) -> RequestAuthMeta {
Self::read_inner(&self.inner).auth_meta.clone()
}
pub async fn scope<F, R>(ctx: SaTokenContext, fut: F) -> R
where
F: Future<Output = R>,
{
TASK_CTX
.scope(ctx, TASK_GRANTS.scope(GrantScope::new(), fut))
.await
}
pub fn try_current() -> Option<SaTokenContext> {
match TASK_CTX.try_with(|c| c.clone()) {
Ok(c) => Some(c),
Err(_) => TLS_CTX.with(|c| c.borrow().clone()),
}
}
pub fn get_current() -> Option<SaTokenContext> {
Self::try_current()
}
pub fn set_current(ctx: SaTokenContext) {
if TASK_CTX.try_with(|_| ()).is_ok() {
let _ = Self::with_current_mut(|inner| {
let snap = Self::read_inner(&ctx.inner);
inner.token.clone_from(&snap.token);
inner.token_info.clone_from(&snap.token_info);
inner.login_id.clone_from(&snap.login_id);
inner.switch_login_id.clone_from(&snap.switch_login_id);
inner.auth_meta = snap.auth_meta.clone();
});
return;
}
TLS_CTX.with(|c| {
*c.borrow_mut() = Some(ctx);
});
TLS_GRANTS.with(|g| {
*g.borrow_mut() = Some(GrantScope::new());
});
}
pub fn clear() {
TLS_CTX.with(|c| {
*c.borrow_mut() = None;
});
TLS_GRANTS.with(|g| {
*g.borrow_mut() = None;
});
}
pub fn with_current_mut<F, R>(f: F) -> Option<R>
where
F: FnOnce(&mut SaTokenContextInner) -> R,
{
if let Ok(handle) = TASK_CTX.try_with(|c| c.clone()) {
let mut guard = Self::write_inner(&handle.inner);
return Some(f(&mut guard));
}
TLS_CTX.with(|cell| {
let mut opt = cell.borrow_mut();
if opt.is_none() {
let auto_create = crate::util::StpUtil::try_get_config()
.map(|c| c.context_auto_create)
.unwrap_or(false);
if !auto_create {
return None;
}
*opt = Some(SaTokenContext::new());
}
let handle = opt.as_ref()?;
let mut guard = Self::write_inner(&handle.inner);
Some(f(&mut guard))
})
}
pub fn current_grant_scope() -> Option<GrantScope> {
match TASK_GRANTS.try_with(|s| s.clone()) {
Ok(s) => Some(s),
Err(_) => TLS_GRANTS.with(|s| s.borrow().clone()),
}
}
pub fn current_login_type() -> Option<String> {
Self::try_current()
.and_then(|ctx| ctx.token_info())
.map(|info| info.login_type.to_string())
.filter(|lt| !lt.is_empty())
}
fn read_inner(
inner: &Arc<RwLock<SaTokenContextInner>>,
) -> std::sync::RwLockReadGuard<'_, SaTokenContextInner> {
inner.read().unwrap_or_else(|e| e.into_inner())
}
fn write_inner(
inner: &Arc<RwLock<SaTokenContextInner>>,
) -> std::sync::RwLockWriteGuard<'_, SaTokenContextInner> {
inner.write().unwrap_or_else(|e| e.into_inner())
}
}
impl Default for SaTokenContext {
fn default() -> Self {
Self::new()
}
}
pub struct SaTokenContextBuilder {
inner: SaTokenContextInner,
}
impl std::fmt::Debug for SaTokenContextBuilder {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("SaTokenContextBuilder { .. }")
}
}
impl SaTokenContextBuilder {
pub fn new() -> Self {
Self {
inner: SaTokenContextInner::default(),
}
}
pub fn token(mut self, token: TokenValue) -> Self {
self.inner.token = Some(token);
self
}
pub fn token_info(mut self, info: Arc<TokenInfo>) -> Self {
self.inner.token_info = Some(info);
self
}
pub fn login_id(mut self, login_id: impl Into<String>) -> Self {
self.inner.login_id = Some(login_id.into());
self
}
pub fn switch_login_id(mut self, login_id: impl Into<String>) -> Self {
self.inner.switch_login_id = Some(login_id.into());
self
}
pub fn auth_meta(mut self, meta: RequestAuthMeta) -> Self {
self.inner.auth_meta = meta;
self
}
pub fn build(self) -> SaTokenContext {
SaTokenContext {
inner: Arc::new(RwLock::new(self.inner)),
}
}
}
impl Default for SaTokenContextBuilder {
fn default() -> Self {
Self::new()
}
}