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
140pub 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#[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 let state = make_state(1).await;
354 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 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}