Skip to main content

allowthem_saas/manage/
mod.rs

1pub mod applications;
2pub mod domains;
3pub mod logs;
4pub mod permissions;
5pub mod roles;
6pub mod sessions;
7pub mod usage;
8pub mod users;
9
10use std::num::NonZeroU32;
11use std::path::PathBuf;
12use std::sync::Arc;
13
14use axum::body::Body;
15use axum::extract::{FromRequestParts, State};
16use axum::http::{Request, StatusCode, header, request::Parts};
17use axum::middleware::Next;
18use axum::response::{IntoResponse, Response};
19use governor::Quota;
20use governor::RateLimiter;
21use governor::clock::{DefaultClock, QuantaInstant};
22use governor::middleware::NoOpMiddleware;
23use governor::state::keyed::DashMapStateStore;
24use serde::Serialize;
25use sha2::{Digest, Sha256};
26
27use crate::api_keys::{ApiKey, ApiKeyScope};
28use crate::cache::HandleCache;
29use crate::control_db::ControlDb;
30use crate::dns::DnsResolver;
31use crate::router::build_handle;
32use crate::tenants::TenantBuilderConfig;
33
34type KeyedLimiter =
35    RateLimiter<Vec<u8>, DashMapStateStore<Vec<u8>>, DefaultClock, NoOpMiddleware<QuantaInstant>>;
36
37#[derive(Debug, thiserror::Error)]
38pub enum ManageError {
39    #[error("unauthorized")]
40    Unauthorized,
41    #[error("forbidden")]
42    Forbidden,
43    #[error("rate limited")]
44    RateLimited(u64),
45    #[error("tenant not found")]
46    TenantNotFound,
47    #[error("not found")]
48    NotFound,
49    #[error("conflict")]
50    Conflict,
51    #[error("invalid request: {0}")]
52    InvalidRequest(String),
53    #[error("internal error: {0}")]
54    Internal(String),
55}
56
57impl IntoResponse for ManageError {
58    fn into_response(self) -> Response {
59        match self {
60            ManageError::Unauthorized => StatusCode::UNAUTHORIZED.into_response(),
61            ManageError::Forbidden => StatusCode::FORBIDDEN.into_response(),
62            ManageError::RateLimited(retry) => {
63                let mut res = StatusCode::TOO_MANY_REQUESTS.into_response();
64                if let Ok(val) = axum::http::HeaderValue::from_str(&retry.to_string()) {
65                    res.headers_mut().insert("retry-after", val);
66                }
67                res
68            }
69            ManageError::TenantNotFound => StatusCode::NOT_FOUND.into_response(),
70            ManageError::NotFound => StatusCode::NOT_FOUND.into_response(),
71            ManageError::Conflict => StatusCode::CONFLICT.into_response(),
72            ManageError::InvalidRequest(_) => StatusCode::BAD_REQUEST.into_response(),
73            ManageError::Internal(_) => StatusCode::INTERNAL_SERVER_ERROR.into_response(),
74        }
75    }
76}
77
78#[derive(Serialize)]
79pub struct ListResponse<T: Serialize> {
80    pub items: Vec<T>,
81    pub next_cursor: Option<String>,
82}
83
84#[derive(Clone)]
85pub struct ManageRateLimiter {
86    inner: Arc<KeyedLimiter>,
87}
88
89impl ManageRateLimiter {
90    pub fn new(quota: Quota) -> Self {
91        Self {
92            inner: Arc::new(RateLimiter::keyed(quota)),
93        }
94    }
95
96    fn check(&self, key_hash: &[u8]) -> Result<(), ManageError> {
97        match self.inner.check_key(&key_hash.to_vec()) {
98            Ok(_) => Ok(()),
99            Err(not_until) => {
100                use governor::clock::Clock;
101                let wait = not_until.wait_time_from(DefaultClock::default().now());
102                Err(ManageError::RateLimited(wait.as_secs().saturating_add(1)))
103            }
104        }
105    }
106}
107
108#[derive(Clone)]
109pub struct ManageState {
110    pub control_db: Arc<ControlDb>,
111    pub handle_cache: HandleCache,
112    pub tenant_data_dir: PathBuf,
113    pub config: Arc<TenantBuilderConfig>,
114    pub rate_limiter: ManageRateLimiter,
115    pub dns_resolver: Arc<dyn DnsResolver>,
116}
117
118impl ManageState {
119    pub fn new(
120        control_db: Arc<ControlDb>,
121        handle_cache: HandleCache,
122        tenant_data_dir: PathBuf,
123        config: Arc<TenantBuilderConfig>,
124        requests_per_minute: u32,
125        dns_resolver: Arc<dyn DnsResolver>,
126    ) -> Self {
127        let rpm =
128            NonZeroU32::new(requests_per_minute).unwrap_or_else(|| NonZeroU32::new(60).unwrap());
129        Self {
130            control_db,
131            handle_cache,
132            tenant_data_dir,
133            config,
134            rate_limiter: ManageRateLimiter::new(Quota::per_minute(rpm)),
135            dns_resolver,
136        }
137    }
138}
139
140/// Authenticates the request via `Authorization: Bearer <key>`.
141///
142/// On success, inserts `ApiKey` and `AllowThem` handle into request extensions.
143pub async fn api_key_auth_middleware(
144    State(state): State<ManageState>,
145    mut req: Request<Body>,
146    next: Next,
147) -> Response {
148    let raw_key = match extract_bearer(req.headers()) {
149        Some(k) => k,
150        None => return ManageError::Unauthorized.into_response(),
151    };
152
153    let key_hash = Sha256::digest(raw_key.as_bytes()).to_vec();
154    if let Err(e) = state.rate_limiter.check(&key_hash) {
155        return e.into_response();
156    }
157
158    let api_key = match state.control_db.verify_api_key(&raw_key).await {
159        Ok(Some(k)) => k,
160        Ok(None) => return ManageError::Unauthorized.into_response(),
161        Err(e) => {
162            tracing::error!(error = %e, "api key verification failed");
163            return ManageError::Internal(e.to_string()).into_response();
164        }
165    };
166
167    let tenant_id = api_key.tenant_id;
168    let slug = match state.control_db.tenant_by_id(&tenant_id).await {
169        Ok(Some(t)) => t.slug,
170        Ok(None) => return ManageError::TenantNotFound.into_response(),
171        Err(e) => {
172            tracing::error!(error = %e, "tenant lookup failed");
173            return ManageError::Internal(e.to_string()).into_response();
174        }
175    };
176
177    let ctrl = state.control_db.clone();
178    let dir = state.tenant_data_dir.clone();
179    let cfg = state.config.clone();
180    let handle = match state
181        .handle_cache
182        .get_or_init(tenant_id, async move {
183            build_handle(ctrl, dir, cfg, tenant_id, &slug).await
184        })
185        .await
186    {
187        Ok(h) => h,
188        Err(e) => {
189            tracing::error!(error = %e, tenant_id = %tenant_id.as_uuid(), "tenant handle init failed");
190            return ManageError::TenantNotFound.into_response();
191        }
192    };
193
194    req.extensions_mut().insert(api_key);
195    req.extensions_mut().insert(handle);
196    next.run(req).await
197}
198
199fn extract_bearer(headers: &axum::http::HeaderMap) -> Option<String> {
200    let val = headers.get(header::AUTHORIZATION)?;
201    let s = val.to_str().ok()?;
202    let token = s.strip_prefix("Bearer ")?;
203    if token.is_empty() {
204        return None;
205    }
206    Some(token.to_owned())
207}
208
209/// Extracts the `ApiKey` from extensions and requires `Admin` scope.
210#[derive(Debug)]
211pub struct AdminKey(pub ApiKey);
212
213impl<S: Send + Sync> FromRequestParts<S> for AdminKey {
214    type Rejection = ManageError;
215
216    async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, ManageError> {
217        let key = parts
218            .extensions
219            .get::<ApiKey>()
220            .ok_or(ManageError::Unauthorized)?;
221        if !key.scope.contains(&ApiKeyScope::Admin) {
222            return Err(ManageError::Forbidden);
223        }
224        Ok(AdminKey(key.clone()))
225    }
226}
227
228pub fn manage_router(state: ManageState) -> axum::Router {
229    axum::Router::<ManageState>::new()
230        .nest("/applications", applications::application_routes())
231        .nest("/domains", domains::domain_routes())
232        .nest("/users", users::user_routes())
233        .nest("/sessions", sessions::session_routes())
234        .nest("/roles", roles::role_routes())
235        .nest("/permissions", permissions::permission_routes())
236        .nest("/logs", logs::log_routes())
237        .nest("/usage", usage::usage_routes())
238        .route_layer(axum::middleware::from_fn_with_state(
239            state.clone(),
240            api_key_auth_middleware,
241        ))
242        .with_state(state)
243}
244
245#[cfg(test)]
246mod tests {
247    use super::*;
248    use axum::Router;
249    use axum::http::{Request, StatusCode};
250    use axum::routing::get;
251    use tower::ServiceExt;
252    use uuid::Uuid;
253
254    use crate::api_keys::ApiKeyScope;
255    use crate::tenants::TenantId;
256
257    #[test]
258    fn manage_error_conflict_is_409() {
259        let resp = ManageError::Conflict.into_response();
260        assert_eq!(resp.status(), StatusCode::CONFLICT);
261    }
262
263    #[test]
264    fn manage_error_invalid_request_is_400() {
265        let resp = ManageError::InvalidRequest("bad input".into()).into_response();
266        assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
267    }
268
269    #[test]
270    fn list_response_serializes_correctly() {
271        let r: ListResponse<i32> = ListResponse {
272            items: vec![1, 2, 3],
273            next_cursor: Some("tok".into()),
274        };
275        let json = serde_json::to_value(&r).unwrap();
276        assert_eq!(json["items"], serde_json::json!([1, 2, 3]));
277        assert_eq!(json["next_cursor"], "tok");
278    }
279
280    async fn make_state(rpm: u32) -> ManageState {
281        use std::str::FromStr;
282        let opts = sqlx::sqlite::SqliteConnectOptions::from_str("sqlite::memory:")
283            .unwrap()
284            .pragma("foreign_keys", "ON");
285        let pool = sqlx::SqlitePool::connect_with(opts).await.unwrap();
286        let ctrl = Arc::new(ControlDb::new(pool).await.unwrap());
287        ManageState::new(
288            ctrl,
289            HandleCache::new(10),
290            PathBuf::from("/tmp"),
291            Arc::new(TenantBuilderConfig {
292                mfa_key: [0u8; 32],
293                signing_key: [0u8; 32],
294                csrf_key: [0u8; 32],
295                base_domain: "example.com".into(),
296                is_production: false,
297                email_sender: None,
298                event_sink: None,
299                event_sink_factory: None,
300                mau_sink: None,
301                email_sender_factory: None,
302            }),
303            rpm,
304            Arc::new(crate::dns::MockDnsResolver::new()),
305        )
306    }
307
308    fn test_app(state: ManageState) -> Router {
309        Router::new()
310            .route("/test", get(|| async { StatusCode::OK }))
311            .route_layer(axum::middleware::from_fn_with_state(
312                state.clone(),
313                api_key_auth_middleware,
314            ))
315            .with_state(state)
316    }
317
318    async fn status_of(app: Router, req: Request<Body>) -> StatusCode {
319        let resp = app.oneshot(req).await.unwrap();
320        resp.status()
321    }
322
323    #[tokio::test]
324    async fn missing_auth_header_returns_401() {
325        let app = test_app(make_state(60).await);
326        let req = Request::get("/test").body(Body::empty()).unwrap();
327        assert_eq!(status_of(app, req).await, StatusCode::UNAUTHORIZED);
328    }
329
330    #[tokio::test]
331    async fn malformed_bearer_returns_401() {
332        let app = test_app(make_state(60).await);
333        let req = Request::get("/test")
334            .header("Authorization", "Basic abc123")
335            .body(Body::empty())
336            .unwrap();
337        assert_eq!(status_of(app, req).await, StatusCode::UNAUTHORIZED);
338    }
339
340    #[tokio::test]
341    async fn unknown_key_returns_401() {
342        let app = test_app(make_state(60).await);
343        let req = Request::get("/test")
344            .header("Authorization", "Bearer sak_aGVsbG8gd29ybGQ")
345            .body(Body::empty())
346            .unwrap();
347        assert_eq!(status_of(app, req).await, StatusCode::UNAUTHORIZED);
348    }
349
350    #[tokio::test]
351    async fn rate_limit_triggers_429() {
352        // quota of 1 request per minute
353        let state = make_state(1).await;
354        // First request: passes rate limit, fails at DB (unknown key → 401)
355        let app1 = test_app(state.clone());
356        let req1 = Request::get("/test")
357            .header("Authorization", "Bearer sak_aGVsbG8gd29ybGQ")
358            .body(Body::empty())
359            .unwrap();
360        assert_eq!(status_of(app1, req1).await, StatusCode::UNAUTHORIZED);
361
362        // Second request: same key hash, rate limited → 429
363        let app2 = test_app(state.clone());
364        let req2 = Request::get("/test")
365            .header("Authorization", "Bearer sak_aGVsbG8gd29ybGQ")
366            .body(Body::empty())
367            .unwrap();
368        assert_eq!(status_of(app2, req2).await, StatusCode::TOO_MANY_REQUESTS);
369    }
370
371    fn make_api_key(scopes: Vec<ApiKeyScope>) -> ApiKey {
372        use crate::api_keys::ApiKeyId;
373        use chrono::Utc;
374        ApiKey {
375            id: ApiKeyId::from_uuid(Uuid::nil()),
376            tenant_id: TenantId::from(Uuid::nil()),
377            name: "test-key".into(),
378            scope: scopes,
379            created_at: Utc::now(),
380            expires_at: None,
381            last_used_at: None,
382        }
383    }
384
385    #[tokio::test]
386    async fn admin_key_extractor_passes_admin_scope() {
387        let api_key = make_api_key(vec![ApiKeyScope::Admin]);
388        let mut req = axum::http::Request::new(());
389        req.extensions_mut().insert(api_key);
390        let (mut parts, _) = req.into_parts();
391        let result = AdminKey::from_request_parts(&mut parts, &()).await;
392        assert!(result.is_ok());
393    }
394
395    #[tokio::test]
396    async fn admin_key_extractor_rejects_missing_key() {
397        let req = axum::http::Request::new(());
398        let (mut parts, _) = req.into_parts();
399        let result = AdminKey::from_request_parts(&mut parts, &()).await;
400        assert!(matches!(result.unwrap_err(), ManageError::Unauthorized));
401    }
402}