Skip to main content

ares_types/models/
tenant.rs

1use serde::{Deserialize, Serialize};
2
3#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
4#[serde(rename_all = "lowercase")]
5pub enum TenantTier {
6    Free,
7    Dev,
8    Pro,
9    Enterprise,
10}
11
12impl TenantTier {
13    #[allow(clippy::should_implement_trait)]
14    pub fn from_str(s: &str) -> Option<Self> {
15        match s.to_lowercase().as_str() {
16            "free" => Some(TenantTier::Free),
17            "dev" => Some(TenantTier::Dev),
18            "pro" => Some(TenantTier::Pro),
19            "enterprise" => Some(TenantTier::Enterprise),
20            _ => None,
21        }
22    }
23
24    pub fn as_str(&self) -> &'static str {
25        match self {
26            TenantTier::Free => "free",
27            TenantTier::Dev => "dev",
28            TenantTier::Pro => "pro",
29            TenantTier::Enterprise => "enterprise",
30        }
31    }
32}
33
34impl std::str::FromStr for TenantTier {
35    type Err = String;
36    fn from_str(s: &str) -> Result<Self, Self::Err> {
37        TenantTier::from_str(s).ok_or_else(|| format!("unknown tenant tier: {}", s))
38    }
39}
40
41#[derive(Debug, Clone, Serialize, Deserialize)]
42pub struct TenantQuota {
43    pub tier: TenantTier,
44    pub requests_per_month: u64,
45    pub tokens_per_month: u64,
46    pub max_agents: u32,
47    pub requests_per_day: u64,
48}
49
50impl Default for TenantQuota {
51    fn default() -> Self {
52        Self::free()
53    }
54}
55
56impl TenantQuota {
57    pub fn free() -> Self {
58        Self {
59            tier: TenantTier::Free,
60            requests_per_month: 1_000,
61            tokens_per_month: 100_000,
62            max_agents: 1,
63            requests_per_day: 50,
64        }
65    }
66
67    pub fn dev() -> Self {
68        Self {
69            tier: TenantTier::Dev,
70            requests_per_month: 50_000,
71            tokens_per_month: 5_000_000,
72            max_agents: 10,
73            requests_per_day: 2_000,
74        }
75    }
76
77    pub fn pro() -> Self {
78        Self {
79            tier: TenantTier::Pro,
80            requests_per_month: 500_000,
81            tokens_per_month: 50_000_000,
82            max_agents: u32::MAX,
83            requests_per_day: 20_000,
84        }
85    }
86
87    pub fn enterprise() -> Self {
88        Self {
89            tier: TenantTier::Enterprise,
90            requests_per_month: u64::MAX,
91            tokens_per_month: u64::MAX,
92            max_agents: u32::MAX,
93            requests_per_day: u64::MAX,
94        }
95    }
96
97    pub fn from_tier(tier: &TenantTier) -> Self {
98        match tier {
99            TenantTier::Free => Self::free(),
100            TenantTier::Dev => Self::dev(),
101            TenantTier::Pro => Self::pro(),
102            TenantTier::Enterprise => Self::enterprise(),
103        }
104    }
105}
106
107#[derive(Debug, Clone, Serialize, Deserialize)]
108pub struct Tenant {
109    pub id: String,
110    pub name: String,
111    pub tier: TenantTier,
112    pub created_at: i64,
113    pub updated_at: i64,
114}
115
116impl Tenant {
117    pub fn new(id: String, name: String, tier: TenantTier) -> Self {
118        let now = chrono::Utc::now().timestamp();
119        Self {
120            id,
121            name,
122            tier,
123            created_at: now,
124            updated_at: now,
125        }
126    }
127}
128
129#[derive(Debug, Clone, Serialize, Deserialize)]
130pub struct ApiKey {
131    pub id: String,
132    pub tenant_id: String,
133    pub key_hash: String,
134    pub key_prefix: String,
135    pub name: String,
136    pub is_active: bool,
137    pub created_at: i64,
138    pub expires_at: Option<i64>,
139}
140
141impl ApiKey {
142    pub fn new(
143        id: String,
144        tenant_id: String,
145        key_hash: String,
146        key_prefix: String,
147        name: String,
148    ) -> Self {
149        Self {
150            id,
151            tenant_id,
152            key_hash,
153            key_prefix,
154            name,
155            is_active: true,
156            created_at: chrono::Utc::now().timestamp(),
157            expires_at: None,
158        }
159    }
160}
161
162#[derive(Debug, Clone, Copy, PartialEq, Eq)]
163pub enum QuotaExceeded {
164    Monthly,
165    Daily,
166}
167
168impl QuotaExceeded {
169    pub fn message(self) -> &'static str {
170        match self {
171            Self::Monthly => "Monthly request quota exceeded",
172            Self::Daily => "Daily rate limit exceeded",
173        }
174    }
175}
176
177impl From<QuotaExceeded> for crate::types::AppError {
178    fn from(exceeded: QuotaExceeded) -> Self {
179        crate::types::AppError::RateLimited(exceeded.message().to_string())
180    }
181}
182
183#[derive(Debug, Clone, Serialize, Deserialize)]
184pub struct TenantContext {
185    pub tenant_id: String,
186    pub tier: TenantTier,
187    pub quota: TenantQuota,
188}
189
190impl TenantContext {
191    pub fn new(tenant_id: String, tier: TenantTier) -> Self {
192        Self {
193            tenant_id,
194            tier,
195            quota: TenantQuota::from_tier(&tier),
196        }
197    }
198
199    pub fn admit(&self, monthly_requests: u64, daily_requests: u64) -> Result<(), QuotaExceeded> {
200        if monthly_requests >= self.quota.requests_per_month {
201            return Err(QuotaExceeded::Monthly);
202        }
203        if daily_requests >= self.quota.requests_per_day {
204            return Err(QuotaExceeded::Daily);
205        }
206        Ok(())
207    }
208
209    pub fn can_make_request(&self, monthly_requests: u64, daily_requests: u64) -> bool {
210        self.admit(monthly_requests, daily_requests).is_ok()
211    }
212
213    pub fn can_use_tokens(&self, monthly_tokens: u64, additional_tokens: u64) -> bool {
214        let Some(new_total) = monthly_tokens.checked_add(additional_tokens) else {
215            return false;
216        };
217        new_total <= self.quota.tokens_per_month
218    }
219}
220
221// Cordis Service impl — makes TenantContext a valid intercept key so per-request
222// tenant scope flows through the Cordis context via ctx.with_intercept(tenant_ctx).
223impl cordis::Service for TenantContext {
224    fn name(&self) -> &'static str { "tenant_context" }
225    fn init(&self, _ctx: &std::sync::Arc<cordis::Context>) -> cordis::ServiceInitFuture<'_> {
226        Box::pin(async { Ok(None) })
227    }
228    fn check(&self) -> bool { true }
229}
230
231#[cfg(test)]
232mod tests {
233    use super::*;
234
235    #[test]
236    fn test_tier_from_str() {
237        assert_eq!(TenantTier::from_str("free"), Some(TenantTier::Free));
238        assert_eq!(TenantTier::from_str("dev"), Some(TenantTier::Dev));
239        assert_eq!(TenantTier::from_str("pro"), Some(TenantTier::Pro));
240        assert_eq!(
241            TenantTier::from_str("enterprise"),
242            Some(TenantTier::Enterprise)
243        );
244        assert_eq!(TenantTier::from_str("unknown"), None);
245    }
246
247    #[test]
248    fn test_tier_as_str() {
249        assert_eq!(TenantTier::Free.as_str(), "free");
250        assert_eq!(TenantTier::Dev.as_str(), "dev");
251        assert_eq!(TenantTier::Pro.as_str(), "pro");
252        assert_eq!(TenantTier::Enterprise.as_str(), "enterprise");
253    }
254
255    #[test]
256    fn test_free_quota() {
257        let quota = TenantQuota::free();
258        assert_eq!(quota.tier, TenantTier::Free);
259        assert_eq!(quota.requests_per_month, 1_000);
260        assert_eq!(quota.tokens_per_month, 100_000);
261        assert_eq!(quota.max_agents, 1);
262        assert_eq!(quota.requests_per_day, 50);
263    }
264
265    #[test]
266    fn test_dev_quota() {
267        let quota = TenantQuota::dev();
268        assert_eq!(quota.tier, TenantTier::Dev);
269        assert_eq!(quota.requests_per_month, 50_000);
270        assert_eq!(quota.tokens_per_month, 5_000_000);
271        assert_eq!(quota.max_agents, 10);
272        assert_eq!(quota.requests_per_day, 2_000);
273    }
274
275    #[test]
276    fn test_pro_quota() {
277        let quota = TenantQuota::pro();
278        assert_eq!(quota.tier, TenantTier::Pro);
279        assert_eq!(quota.requests_per_month, 500_000);
280        assert_eq!(quota.tokens_per_month, 50_000_000);
281        assert_eq!(quota.max_agents, u32::MAX);
282        assert_eq!(quota.requests_per_day, 20_000);
283    }
284
285    #[test]
286    fn test_enterprise_quota() {
287        let quota = TenantQuota::enterprise();
288        assert_eq!(quota.tier, TenantTier::Enterprise);
289        assert_eq!(quota.requests_per_month, u64::MAX);
290        assert_eq!(quota.tokens_per_month, u64::MAX);
291    }
292
293    #[test]
294    fn test_quota_from_tier() {
295        assert_eq!(
296            TenantQuota::from_tier(&TenantTier::Free).requests_per_month,
297            1_000
298        );
299        assert_eq!(
300            TenantQuota::from_tier(&TenantTier::Dev).requests_per_month,
301            50_000
302        );
303        assert_eq!(
304            TenantQuota::from_tier(&TenantTier::Pro).requests_per_month,
305            500_000
306        );
307    }
308
309    #[test]
310    fn test_tenant_context_can_make_request() {
311        let ctx = TenantContext::new("test".to_string(), TenantTier::Free);
312        assert!(ctx.can_make_request(0, 0));
313        assert!(ctx.can_make_request(999, 0));
314        assert!(ctx.can_make_request(0, 49));
315        assert!(!ctx.can_make_request(1000, 0));
316        assert!(!ctx.can_make_request(0, 50));
317    }
318
319    #[test]
320    fn test_tenant_context_admit_ok() {
321        let ctx = TenantContext::new("test".to_string(), TenantTier::Free);
322        assert!(ctx.admit(0, 0).is_ok());
323        assert!(ctx.admit(999, 49).is_ok());
324    }
325
326    #[test]
327    fn test_tenant_context_admit_monthly() {
328        let ctx = TenantContext::new("test".to_string(), TenantTier::Free);
329        assert_eq!(ctx.admit(1000, 0), Err(QuotaExceeded::Monthly));
330        assert_eq!(ctx.admit(1000, 50), Err(QuotaExceeded::Monthly));
331    }
332
333    #[test]
334    fn test_tenant_context_admit_daily() {
335        let ctx = TenantContext::new("test".to_string(), TenantTier::Free);
336        assert_eq!(ctx.admit(0, 50), Err(QuotaExceeded::Daily));
337    }
338
339    #[test]
340    fn test_tenant_context_admit_enterprise() {
341        let ctx = TenantContext::new("ent".to_string(), TenantTier::Enterprise);
342        assert!(ctx.admit(1_000_000, 1_000_000).is_ok());
343        let err: crate::types::AppError = QuotaExceeded::Monthly.into();
344        match err {
345            crate::types::AppError::RateLimited(msg) => {
346                assert_eq!(msg, "Monthly request quota exceeded");
347            }
348            other => panic!("expected RateLimited, got {other:?}"),
349        }
350        let err: crate::types::AppError = QuotaExceeded::Daily.into();
351        match err {
352            crate::types::AppError::RateLimited(msg) => {
353                assert_eq!(msg, "Daily rate limit exceeded");
354            }
355            other => panic!("expected RateLimited, got {other:?}"),
356        }
357    }
358
359    #[test]
360    fn test_tenant_context_can_use_tokens() {
361        let ctx = TenantContext::new("test".to_string(), TenantTier::Free);
362        assert!(ctx.can_use_tokens(0, 100_000));
363        assert!(ctx.can_use_tokens(50_000, 50_000));
364        assert!(!ctx.can_use_tokens(50_000, 50_001));
365        assert!(!ctx.can_use_tokens(100_000, 1));
366    }
367
368    #[test]
369    fn test_tenant_creation() {
370        let tenant = Tenant::new("t1".to_string(), "Test Tenant".to_string(), TenantTier::Dev);
371        assert_eq!(tenant.id, "t1");
372        assert_eq!(tenant.name, "Test Tenant");
373        assert_eq!(tenant.tier, TenantTier::Dev);
374        assert!(tenant.created_at > 0);
375    }
376
377    #[test]
378    fn test_api_key_creation() {
379        let key = ApiKey::new(
380            "k1".to_string(),
381            "t1".to_string(),
382            "hash123".to_string(),
383            "ares_abc".to_string(),
384            "Test Key".to_string(),
385        );
386        assert_eq!(key.id, "k1");
387        assert_eq!(key.tenant_id, "t1");
388        assert!(key.is_active);
389        assert!(key.created_at > 0);
390    }
391
392    #[test]
393    fn test_tenant_tier_serde_roundtrip() {
394        for tier in [
395            TenantTier::Free,
396            TenantTier::Dev,
397            TenantTier::Pro,
398            TenantTier::Enterprise,
399        ] {
400            let json = serde_json::to_string(&tier).unwrap();
401            let parsed: TenantTier = serde_json::from_str(&json).unwrap();
402            assert_eq!(parsed, tier);
403            assert_eq!(parsed.as_str(), tier.as_str());
404        }
405    }
406
407    #[test]
408    fn test_tenant_tier_partial_eq_and_copy() {
409        let a = TenantTier::Pro;
410        let b = a;
411        assert_eq!(a, b);
412        assert_ne!(a, TenantTier::Free);
413        assert!(format!("{:?}", a).contains("Pro"));
414    }
415
416    #[test]
417    fn test_tenant_tier_from_str_case_insensitive() {
418        assert_eq!(TenantTier::from_str("FREE"), Some(TenantTier::Free));
419        assert_eq!(TenantTier::from_str("Enterprise"), Some(TenantTier::Enterprise));
420        assert_eq!(TenantTier::from_str(""), None);
421    }
422
423    #[test]
424    fn test_tenant_quota_default_is_free() {
425        let default = TenantQuota::default();
426        let free = TenantQuota::free();
427        assert_eq!(default.tier, free.tier);
428        assert_eq!(default.requests_per_month, free.requests_per_month);
429        assert_eq!(default.tokens_per_month, free.tokens_per_month);
430    }
431
432    #[test]
433    fn test_tenant_quota_serde_roundtrip() {
434        let quota = TenantQuota::enterprise();
435        let parsed: TenantQuota =
436            serde_json::from_str(&serde_json::to_string(&quota).unwrap()).unwrap();
437        assert_eq!(parsed.tier, TenantTier::Enterprise);
438        assert_eq!(parsed.requests_per_month, u64::MAX);
439        assert_eq!(parsed.max_agents, u32::MAX);
440    }
441
442    #[test]
443    fn test_tenant_serde_roundtrip_unicode_name() {
444        let tenant = Tenant::new(
445            "t-unicode".into(),
446            "租户 🏢".into(),
447            TenantTier::Pro,
448        );
449        let parsed: Tenant =
450            serde_json::from_str(&serde_json::to_string(&tenant).unwrap()).unwrap();
451        assert_eq!(parsed.name, "租户 🏢");
452        assert_eq!(parsed.tier, TenantTier::Pro);
453    }
454
455    #[test]
456    fn test_api_key_serde_roundtrip_with_expiry() {
457        let key = ApiKey {
458            id: "k2".into(),
459            tenant_id: "t1".into(),
460            key_hash: "hash".into(),
461            key_prefix: "ares_".into(),
462            name: String::new(),
463            is_active: false,
464            created_at: 0,
465            expires_at: Some(i64::MAX),
466        };
467        let parsed: ApiKey =
468            serde_json::from_str(&serde_json::to_string(&key).unwrap()).unwrap();
469        assert!(!parsed.is_active);
470        assert_eq!(parsed.expires_at, Some(i64::MAX));
471        assert!(parsed.name.is_empty());
472    }
473
474    #[test]
475    fn test_tenant_context_serde_roundtrip() {
476        let ctx = TenantContext::new("tenant-1".into(), TenantTier::Dev);
477        let parsed: TenantContext =
478            serde_json::from_str(&serde_json::to_string(&ctx).unwrap()).unwrap();
479        assert_eq!(parsed.tenant_id, "tenant-1");
480        assert_eq!(parsed.tier, TenantTier::Dev);
481        assert_eq!(parsed.quota.requests_per_day, 2_000);
482    }
483
484    #[test]
485    fn test_tenant_context_token_overflow_boundary() {
486        let ctx = TenantContext::new("t".into(), TenantTier::Free);
487        assert!(!ctx.can_use_tokens(u64::MAX, 1));
488        assert!(ctx.can_use_tokens(0, 0));
489        assert!(!ctx.can_make_request(u64::MAX, 0));
490    }
491
492    #[test]
493    fn test_enterprise_context_unlimited_requests() {
494        let ctx = TenantContext::new("ent".into(), TenantTier::Enterprise);
495        assert!(ctx.can_make_request(u64::MAX - 1, u64::MAX - 1));
496        assert!(ctx.can_use_tokens(u64::MAX - 1, 1));
497    }
498
499    #[test]
500    fn test_enterprise_quota_full_limits() {
501        let quota = TenantQuota::enterprise();
502        assert_eq!(quota.max_agents, u32::MAX);
503        assert_eq!(quota.requests_per_day, u64::MAX);
504    }
505
506    #[test]
507    fn test_quota_from_tier_enterprise() {
508        let quota = TenantQuota::from_tier(&TenantTier::Enterprise);
509        assert_eq!(quota.tier, TenantTier::Enterprise);
510        assert_eq!(quota.requests_per_month, u64::MAX);
511        assert_eq!(quota.tokens_per_month, u64::MAX);
512        assert_eq!(quota.max_agents, u32::MAX);
513        assert_eq!(quota.requests_per_day, u64::MAX);
514    }
515
516    #[test]
517    fn test_tenant_tier_serde_json_is_lowercase() {
518        let json = serde_json::to_string(&TenantTier::Pro).unwrap();
519        assert_eq!(json, "\"pro\"");
520        assert_eq!(
521            serde_json::from_str::<TenantTier>("\"pro\"").unwrap(),
522            TenantTier::Pro
523        );
524    }
525
526    #[test]
527    fn test_tenant_tier_serde_rejects_unknown_variant() {
528        let err = serde_json::from_str::<TenantTier>("\"platinum\"").unwrap_err();
529        assert!(err.is_data());
530    }
531
532    #[test]
533    fn test_tier_from_str_rejects_whitespace_and_garbage() {
534        assert_eq!(TenantTier::from_str(" free"), None);
535        assert_eq!(TenantTier::from_str("free "), None);
536        assert_eq!(TenantTier::from_str("free\n"), None);
537        assert_eq!(TenantTier::from_str("pro "), None);
538    }
539
540    #[test]
541    fn test_api_key_serde_roundtrip_from_new() {
542        let key = ApiKey::new(
543            "k-new".into(),
544            "tenant".into(),
545            "hash".into(),
546            "ares_xyz".into(),
547            "Primary".into(),
548        );
549        let parsed: ApiKey =
550            serde_json::from_str(&serde_json::to_string(&key).unwrap()).unwrap();
551        assert_eq!(parsed.id, key.id);
552        assert_eq!(parsed.tenant_id, key.tenant_id);
553        assert_eq!(parsed.key_hash, key.key_hash);
554        assert_eq!(parsed.key_prefix, key.key_prefix);
555        assert_eq!(parsed.name, key.name);
556        assert!(parsed.is_active);
557        assert_eq!(parsed.created_at, key.created_at);
558        assert_eq!(parsed.expires_at, None);
559    }
560
561    #[test]
562    fn test_tenant_serde_preserves_timestamps_and_id() {
563        let tenant = Tenant {
564            id: "fixed-id".into(),
565            name: "Acme".into(),
566            tier: TenantTier::Dev,
567            created_at: 1_700_000_000,
568            updated_at: 1_700_000_001,
569        };
570        let parsed: Tenant =
571            serde_json::from_str(&serde_json::to_string(&tenant).unwrap()).unwrap();
572        assert_eq!(parsed.id, "fixed-id");
573        assert_eq!(parsed.created_at, 1_700_000_000);
574        assert_eq!(parsed.updated_at, 1_700_000_001);
575    }
576
577    #[test]
578    fn test_tenant_quota_serde_rejects_missing_field() {
579        let err = serde_json::from_str::<TenantQuota>(r#"{"tier":"free"}"#).unwrap_err();
580        assert!(err.is_data());
581    }
582
583    #[test]
584    fn test_tenant_quota_serde_all_tiers_roundtrip() {
585        for factory in [
586            TenantQuota::free,
587            TenantQuota::dev,
588            TenantQuota::pro,
589            TenantQuota::enterprise,
590        ] {
591            let quota = factory();
592            let parsed: TenantQuota =
593                serde_json::from_str(&serde_json::to_string(&quota).unwrap()).unwrap();
594            assert_eq!(parsed.tier, quota.tier);
595            assert_eq!(parsed.requests_per_month, quota.requests_per_month);
596            assert_eq!(parsed.tokens_per_month, quota.tokens_per_month);
597            assert_eq!(parsed.max_agents, quota.max_agents);
598            assert_eq!(parsed.requests_per_day, quota.requests_per_day);
599        }
600    }
601
602    #[test]
603    fn test_can_make_request_exact_monthly_and_daily_boundaries() {
604        let ctx = TenantContext::new("t".into(), TenantTier::Free);
605        assert!(ctx.can_make_request(999, 49));
606        assert!(!ctx.can_make_request(1_000, 0));
607        assert!(!ctx.can_make_request(0, 50));
608        assert!(!ctx.can_make_request(1_000, 50));
609    }
610
611    #[test]
612    fn test_can_use_tokens_exact_monthly_ceiling() {
613        let ctx = TenantContext::new("t".into(), TenantTier::Free);
614        assert!(ctx.can_use_tokens(99_999, 1));
615        assert!(ctx.can_use_tokens(100_000, 0));
616        assert!(!ctx.can_use_tokens(100_000, 1));
617        assert!(!ctx.can_use_tokens(99_999, 2));
618    }
619
620    #[test]
621    fn test_dev_context_daily_request_boundary() {
622        let ctx = TenantContext::new("dev-tenant".into(), TenantTier::Dev);
623        assert!(ctx.can_make_request(0, 1_999));
624        assert!(!ctx.can_make_request(0, 2_000));
625    }
626
627    #[test]
628    fn test_tenant_context_clone_matches_original() {
629        let ctx = TenantContext::new("clone-me".into(), TenantTier::Pro);
630        let cloned = ctx.clone();
631        assert_eq!(cloned.tenant_id, ctx.tenant_id);
632        assert_eq!(cloned.tier, ctx.tier);
633        assert_eq!(cloned.quota.requests_per_month, ctx.quota.requests_per_month);
634    }
635
636    #[test]
637    fn test_enterprise_context_rejects_at_hard_ceiling() {
638        let ctx = TenantContext::new("ent".into(), TenantTier::Enterprise);
639        assert!(!ctx.can_make_request(u64::MAX, 0));
640        assert!(!ctx.can_make_request(0, u64::MAX));
641        assert!(!ctx.can_use_tokens(u64::MAX, 1));
642    }
643
644    #[test]
645    fn tenant_context_readable_via_cordis_intercept() {
646        // Cordis design: per-request tenant scope should flow via
647        // ctx.with_intercept(TenantContext) so downstream services read it
648        // from the context (ctx.get::<TenantContext>()) instead of Axum
649        // request extensions.
650        use std::sync::Arc;
651        let root: Arc<cordis::Context> = cordis::Context::new_root();
652
653        // Before intercept — no TenantContext in context.
654        assert!(root.get::<TenantContext>().is_none());
655
656        // After intercept — TenantContext is readable.
657        let tc = TenantContext::new("acme".into(), TenantTier::Pro);
658        let child = root.with_intercept(tc);
659        let retrieved = child.get::<TenantContext>().expect("intercept must make TenantContext readable");
660        assert_eq!(retrieved.tenant_id, "acme");
661        assert_eq!(retrieved.tier, TenantTier::Pro);
662    }
663}