prax_query/tenant/
task_local.rs1use std::cell::Cell;
36use std::future::Future;
37
38use super::context::{TenantContext, TenantId};
39
40tokio::task_local! {
41 static TENANT_CONTEXT: TenantContext;
43}
44
45thread_local! {
46 static SYNC_TENANT_ID: Cell<Option<TenantId>> = const { Cell::new(None) };
49}
50
51pub async fn with_tenant<F, T>(tenant_id: impl Into<TenantId>, f: F) -> T
68where
69 F: Future<Output = T>,
70{
71 let ctx = TenantContext::new(tenant_id);
72 TENANT_CONTEXT.scope(ctx, f).await
73}
74
75pub async fn with_context<F, T>(ctx: TenantContext, f: F) -> T
77where
78 F: Future<Output = T>,
79{
80 TENANT_CONTEXT.scope(ctx, f).await
81}
82
83#[inline]
97pub fn current_tenant() -> Option<TenantContext> {
98 TENANT_CONTEXT.try_with(|ctx| ctx.clone()).ok()
99}
100
101#[inline]
105pub fn current_tenant_id() -> Option<TenantId> {
106 TENANT_CONTEXT.try_with(|ctx| ctx.id.clone()).ok()
107}
108
109#[deprecated(
114 since = "0.11.0",
115 note = "always returns an empty string; use current_tenant_id() instead"
116)]
117#[inline]
118pub fn current_tenant_id_str() -> &'static str {
119 ""
122}
123
124#[inline]
126pub fn has_tenant() -> bool {
127 TENANT_CONTEXT.try_with(|_| ()).is_ok()
128}
129
130#[inline]
134pub fn with_current_tenant<F, T>(f: F) -> Option<T>
135where
136 F: FnOnce(&TenantContext) -> T,
137{
138 TENANT_CONTEXT.try_with(f).ok()
139}
140
141#[inline]
143pub fn require_tenant() -> Result<TenantContext, TenantNotSetError> {
144 current_tenant().ok_or(TenantNotSetError)
145}
146
147#[derive(Debug, Clone, Copy)]
149pub struct TenantNotSetError;
150
151impl std::fmt::Display for TenantNotSetError {
152 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
153 write!(f, "tenant context not set")
154 }
155}
156
157impl std::error::Error for TenantNotSetError {}
158
159pub fn set_sync_tenant(tenant_id: impl Into<TenantId>) -> SyncTenantGuard {
177 let id = tenant_id.into();
178 let previous = SYNC_TENANT_ID.with(|cell| cell.replace(Some(id)));
179 SyncTenantGuard { previous }
180}
181
182#[inline]
184pub fn sync_tenant_id() -> Option<TenantId> {
185 SYNC_TENANT_ID.with(|cell| {
186 unsafe { &*cell.as_ptr() }.clone()
188 })
189}
190
191pub struct SyncTenantGuard {
193 previous: Option<TenantId>,
194}
195
196impl Drop for SyncTenantGuard {
197 fn drop(&mut self) {
198 SYNC_TENANT_ID.with(|cell| cell.set(self.previous.take()));
199 }
200}
201
202#[derive(Debug, Clone)]
225pub struct TenantScope {
226 context: TenantContext,
227}
228
229impl TenantScope {
230 pub fn new(tenant_id: impl Into<TenantId>) -> Self {
232 Self {
233 context: TenantContext::new(tenant_id),
234 }
235 }
236
237 pub fn from_context(context: TenantContext) -> Self {
239 Self { context }
240 }
241
242 pub fn tenant_id(&self) -> &TenantId {
244 &self.context.id
245 }
246
247 pub fn context(&self) -> &TenantContext {
249 &self.context
250 }
251
252 pub async fn run<F, T>(&self, f: F) -> T
254 where
255 F: Future<Output = T>,
256 {
257 TENANT_CONTEXT.scope(self.context.clone(), f).await
258 }
259
260 pub fn run_sync<F, T>(&self, f: F) -> T
262 where
263 F: FnOnce() -> T,
264 {
265 let _guard = set_sync_tenant(self.context.id.clone());
266 f()
267 }
268}
269
270pub trait TenantExtractor: Send + Sync {
276 fn extract(&self, headers: &[(String, String)]) -> Option<TenantId>;
278}
279
280#[derive(Debug, Clone)]
282pub struct HeaderExtractor {
283 header_name: String,
284}
285
286impl HeaderExtractor {
287 pub fn new(header_name: impl Into<String>) -> Self {
289 Self {
290 header_name: header_name.into(),
291 }
292 }
293
294 pub fn default_header() -> Self {
296 Self::new("X-Tenant-ID")
297 }
298}
299
300impl TenantExtractor for HeaderExtractor {
301 fn extract(&self, headers: &[(String, String)]) -> Option<TenantId> {
302 headers
303 .iter()
304 .find(|(k, _)| k.eq_ignore_ascii_case(&self.header_name))
305 .map(|(_, v)| TenantId::new(v.clone()))
306 }
307}
308
309#[derive(Debug, Clone)]
318pub struct UnverifiedJwtClaimExtractor {
319 claim_name: String,
320}
321
322impl UnverifiedJwtClaimExtractor {
323 pub fn new(claim_name: impl Into<String>) -> Self {
325 Self {
326 claim_name: claim_name.into(),
327 }
328 }
329
330 pub fn default_claim() -> Self {
332 Self::new("tenant_id")
333 }
334
335 pub fn claim_name(&self) -> &str {
337 &self.claim_name
338 }
339}
340
341pub type JwtClaimExtractor = UnverifiedJwtClaimExtractor;
350
351impl TenantExtractor for UnverifiedJwtClaimExtractor {
352 fn extract(&self, headers: &[(String, String)]) -> Option<TenantId> {
359 use base64::Engine as _;
360
361 let auth = headers
362 .iter()
363 .find(|(k, _)| k.eq_ignore_ascii_case("authorization"))
364 .map(|(_, v)| v.as_str())?;
365
366 let Some(token) = auth
368 .get(..7)
369 .filter(|scheme| scheme.eq_ignore_ascii_case("bearer "))
370 .map(|_| &auth[7..])
371 else {
372 tracing::debug!("Authorization header present but not a Bearer token");
373 return None;
374 };
375
376 let Some(payload) = token.split('.').nth(1) else {
378 tracing::debug!("Authorization header present but JWT has no payload segment");
379 return None;
380 };
381 let decoded = match base64::engine::general_purpose::URL_SAFE_NO_PAD.decode(payload) {
382 Ok(decoded) => decoded,
383 Err(error) => {
384 tracing::debug!(%error, "Authorization header present but JWT payload undecodable");
385 return None;
386 }
387 };
388 let claims: serde_json::Value = match serde_json::from_slice(&decoded) {
389 Ok(claims) => claims,
390 Err(error) => {
391 tracing::debug!(%error, "Authorization header present but JWT payload is not JSON");
392 return None;
393 }
394 };
395
396 claims
397 .get(self.claim_name.as_str())?
398 .as_str()
399 .map(TenantId::new)
400 }
401}
402
403pub struct CompositeExtractor {
405 extractors: Vec<Box<dyn TenantExtractor>>,
406}
407
408impl CompositeExtractor {
409 pub fn new() -> Self {
411 Self {
412 extractors: Vec::new(),
413 }
414 }
415
416 pub fn add<E: TenantExtractor + 'static>(mut self, extractor: E) -> Self {
418 self.extractors.push(Box::new(extractor));
419 self
420 }
421}
422
423impl Default for CompositeExtractor {
424 fn default() -> Self {
425 Self::new()
426 }
427}
428
429impl TenantExtractor for CompositeExtractor {
430 fn extract(&self, headers: &[(String, String)]) -> Option<TenantId> {
431 for extractor in &self.extractors {
432 if let Some(id) = extractor.extract(headers) {
433 return Some(id);
434 }
435 }
436 None
437 }
438}
439
440#[cfg(test)]
441mod tests {
442 use super::*;
443
444 #[tokio::test]
445 async fn test_with_tenant() {
446 let result = with_tenant("test-tenant", async { current_tenant_id() }).await;
447
448 assert_eq!(result.unwrap().as_str(), "test-tenant");
449 }
450
451 #[tokio::test]
452 async fn test_no_tenant() {
453 assert!(current_tenant().is_none());
454 assert!(!has_tenant());
455 }
456
457 #[tokio::test]
458 async fn test_nested_tenant() {
459 with_tenant("outer", async {
460 assert_eq!(current_tenant_id().unwrap().as_str(), "outer");
461
462 with_tenant("inner", async {
463 assert_eq!(current_tenant_id().unwrap().as_str(), "inner");
464 })
465 .await;
466
467 assert_eq!(current_tenant_id().unwrap().as_str(), "outer");
469 })
470 .await;
471 }
472
473 #[tokio::test]
474 async fn test_tenant_scope() {
475 let scope = TenantScope::new("scoped-tenant");
476
477 let result = scope
478 .run(async { current_tenant_id().map(|id| id.as_str().to_string()) })
479 .await;
480
481 assert_eq!(result, Some("scoped-tenant".to_string()));
482 }
483
484 #[test]
485 fn test_sync_tenant() {
486 {
487 let _guard = set_sync_tenant("sync-tenant");
488 assert_eq!(sync_tenant_id().unwrap().as_str(), "sync-tenant");
489 }
490
491 assert!(sync_tenant_id().is_none());
493 }
494
495 #[test]
496 fn test_header_extractor() {
497 let extractor = HeaderExtractor::new("X-Tenant-ID");
498
499 let headers = vec![
500 ("Content-Type".to_string(), "application/json".to_string()),
501 ("X-Tenant-ID".to_string(), "tenant-from-header".to_string()),
502 ];
503
504 let id = extractor.extract(&headers);
505 assert_eq!(id.unwrap().as_str(), "tenant-from-header");
506 }
507
508 #[test]
509 fn test_jwt_claim_extractor() {
510 let extractor = UnverifiedJwtClaimExtractor::new("tenant_id");
511
512 let jwt = "eyJhbGciOiJub25lIiwidHlwIjoiSldUIn0.eyJ0ZW5hbnRfaWQiOiJ0ZW5hbnQtZnJvbS1qd3QifQ.";
516 let headers = vec![
517 ("Content-Type".to_string(), "application/json".to_string()),
518 ("Authorization".to_string(), format!("Bearer {jwt}")),
519 ];
520
521 let id = extractor.extract(&headers);
522 assert_eq!(id.unwrap().as_str(), "tenant-from-jwt");
523
524 let _alias: JwtClaimExtractor = UnverifiedJwtClaimExtractor::default_claim();
526
527 assert!(extractor.extract(&[]).is_none());
529
530 for bad in ["not-a-jwt", "Bearer !!!.@@@.###", "Bearer a.b.c"] {
532 let headers = vec![("Authorization".to_string(), bad.to_string())];
533 assert!(extractor.extract(&headers).is_none());
534 }
535
536 let other = UnverifiedJwtClaimExtractor::new("org_id");
538 assert!(other.extract(&headers).is_none());
539 }
540
541 #[test]
542 fn test_composite_extractor() {
543 let extractor = CompositeExtractor::new()
544 .add(HeaderExtractor::new("X-Organization-ID"))
545 .add(HeaderExtractor::new("X-Tenant-ID"));
546
547 let headers = vec![("X-Tenant-ID".to_string(), "fallback-tenant".to_string())];
548
549 let id = extractor.extract(&headers);
550 assert_eq!(id.unwrap().as_str(), "fallback-tenant");
551 }
552}