1use base64ct::{Base64UrlUnpadded, Encoding};
2use chrono::{DateTime, Utc};
3use rand::TryRngCore;
4use rand::rngs::OsRng;
5use sha2::{Digest, Sha256};
6use sqlx::Row;
7use subtle::ConstantTimeEq;
8use uuid::Uuid;
9
10use crate::control_db::ControlDb;
11use crate::error::SaasError;
12use crate::tenants::TenantId;
13
14const KEY_PREFIX: &str = "sak_";
15
16#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
18pub struct ApiKeyId(Uuid);
19
20impl ApiKeyId {
21 pub fn new() -> Self {
22 Self(Uuid::now_v7())
23 }
24
25 pub fn from_uuid(id: Uuid) -> Self {
26 Self(id)
27 }
28
29 pub fn as_bytes(&self) -> &[u8] {
30 self.0.as_bytes()
31 }
32}
33
34impl Default for ApiKeyId {
35 fn default() -> Self {
36 Self::new()
37 }
38}
39
40impl std::fmt::Display for ApiKeyId {
41 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
42 self.0.fmt(f)
43 }
44}
45
46#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
47#[serde(rename_all = "snake_case")]
48pub enum ApiKeyScope {
49 Admin,
50}
51
52#[derive(Debug, Clone)]
53pub struct ApiKey {
54 pub id: ApiKeyId,
55 pub tenant_id: TenantId,
56 pub name: String,
57 pub scope: Vec<ApiKeyScope>,
58 pub created_at: DateTime<Utc>,
59 pub expires_at: Option<DateTime<Utc>>,
60 pub last_used_at: Option<DateTime<Utc>>,
61}
62
63pub struct ApiKeyMintResult {
64 pub api_key: ApiKey,
65 pub raw_key: String,
67}
68
69fn generate_raw_key() -> Result<([u8; 32], String), SaasError> {
70 let mut bytes = [0u8; 32];
71 OsRng
72 .try_fill_bytes(&mut bytes)
73 .map_err(|e| SaasError::ProvisionFailed(e.to_string()))?;
74 let encoded = format!("{}{}", KEY_PREFIX, Base64UrlUnpadded::encode_string(&bytes));
75 Ok((bytes, encoded))
76}
77
78fn hash_key_bytes(bytes: &[u8]) -> Vec<u8> {
79 Sha256::digest(bytes).to_vec()
80}
81
82fn decode_raw_key(raw_key: &str) -> Option<Vec<u8>> {
83 let encoded = raw_key.strip_prefix(KEY_PREFIX)?;
84 Base64UrlUnpadded::decode_vec(encoded).ok()
85}
86
87impl ControlDb {
88 pub async fn mint_api_key(
89 &self,
90 tenant_id: &TenantId,
91 name: &str,
92 scopes: Vec<ApiKeyScope>,
93 expires_at: Option<DateTime<Utc>>,
94 ) -> Result<ApiKeyMintResult, SaasError> {
95 let (raw_bytes, raw_key) = generate_raw_key()?;
96 let key_hash = hash_key_bytes(&raw_bytes);
97 let key_id = ApiKeyId::new();
98 let scope_json = serde_json::to_string(&scopes)
99 .map_err(|e| SaasError::ProvisionFailed(e.to_string()))?;
100
101 sqlx::query(
102 "INSERT INTO tenant_api_keys (id, tenant_id, name, key_hash, scope, expires_at) \
103 VALUES (?1, ?2, ?3, ?4, ?5, ?6)",
104 )
105 .bind(key_id.as_bytes())
106 .bind(tenant_id.as_bytes())
107 .bind(name)
108 .bind(&key_hash)
109 .bind(&scope_json)
110 .bind(expires_at)
111 .execute(self.pool())
112 .await?;
113
114 let api_key = ApiKey {
115 id: key_id,
116 tenant_id: *tenant_id,
117 name: name.to_owned(),
118 scope: scopes,
119 created_at: Utc::now(),
120 expires_at,
121 last_used_at: None,
122 };
123
124 Ok(ApiKeyMintResult { api_key, raw_key })
125 }
126
127 pub async fn verify_api_key(&self, raw_key: &str) -> Result<Option<ApiKey>, SaasError> {
130 let Some(raw_bytes) = decode_raw_key(raw_key) else {
131 return Ok(None);
132 };
133 let candidate_hash = hash_key_bytes(&raw_bytes);
134
135 let row = sqlx::query(
137 "SELECT id, tenant_id, name, scope, key_hash, created_at, expires_at, \
138 revoked_at, last_used_at \
139 FROM tenant_api_keys WHERE key_hash = ?1",
140 )
141 .bind(&candidate_hash)
142 .fetch_optional(self.pool())
143 .await?;
144
145 let Some(row) = row else {
146 return Ok(None);
147 };
148
149 let stored_hash: Vec<u8> = row.try_get("key_hash")?;
151 if !bool::from(candidate_hash.ct_eq(&stored_hash)) {
152 return Ok(None);
153 }
154
155 let revoked_at: Option<DateTime<Utc>> = row.try_get("revoked_at")?;
156 if revoked_at.is_some() {
157 return Ok(None);
158 }
159
160 let expires_at: Option<DateTime<Utc>> = row.try_get("expires_at")?;
161 if expires_at.is_some_and(|exp| exp <= Utc::now()) {
162 return Ok(None);
163 }
164
165 let id_bytes: Vec<u8> = row.try_get("id")?;
166 let tenant_bytes: Vec<u8> = row.try_get("tenant_id")?;
167 let scope_json: String = row.try_get("scope")?;
168 let scopes: Vec<ApiKeyScope> = serde_json::from_str(&scope_json)
169 .map_err(|e| SaasError::ProvisionFailed(e.to_string()))?;
170 let key_id = Uuid::from_slice(&id_bytes).map_err(|_| SaasError::TenantNotFound)?;
171 let tenant_id = Uuid::from_slice(&tenant_bytes).map_err(|_| SaasError::TenantNotFound)?;
172 let created_at: DateTime<Utc> = row.try_get("created_at")?;
173 let last_used_at: Option<DateTime<Utc>> = row.try_get("last_used_at")?;
174 let name: String = row.try_get("name")?;
175
176 sqlx::query("UPDATE tenant_api_keys SET last_used_at = ?1 WHERE key_hash = ?2")
177 .bind(Utc::now())
178 .bind(&candidate_hash)
179 .execute(self.pool())
180 .await?;
181
182 Ok(Some(ApiKey {
183 id: ApiKeyId::from_uuid(key_id),
184 tenant_id: TenantId::from(tenant_id),
185 name,
186 scope: scopes,
187 created_at,
188 expires_at,
189 last_used_at,
190 }))
191 }
192
193 pub async fn revoke_api_key(
194 &self,
195 key_id: &ApiKeyId,
196 tenant_id: &TenantId,
197 ) -> Result<(), SaasError> {
198 sqlx::query(
199 "UPDATE tenant_api_keys SET revoked_at = ?1 \
200 WHERE id = ?2 AND tenant_id = ?3 AND revoked_at IS NULL",
201 )
202 .bind(Utc::now())
203 .bind(key_id.as_bytes())
204 .bind(tenant_id.as_bytes())
205 .execute(self.pool())
206 .await?;
207 Ok(())
208 }
209
210 pub async fn list_api_keys_for_tenant(
211 &self,
212 tenant_id: &TenantId,
213 ) -> Result<Vec<ApiKey>, SaasError> {
214 let rows = sqlx::query(
215 "SELECT id, tenant_id, name, scope, created_at, expires_at, last_used_at \
216 FROM tenant_api_keys \
217 WHERE tenant_id = ?1 AND revoked_at IS NULL \
218 ORDER BY created_at DESC",
219 )
220 .bind(tenant_id.as_bytes())
221 .fetch_all(self.pool())
222 .await?;
223
224 let mut result = Vec::with_capacity(rows.len());
225 for row in rows {
226 let id_bytes: Vec<u8> = row.try_get("id")?;
227 let tenant_bytes: Vec<u8> = row.try_get("tenant_id")?;
228 let scope_json: String = row.try_get("scope")?;
229 let scopes: Vec<ApiKeyScope> = serde_json::from_str(&scope_json)
230 .map_err(|e| SaasError::ProvisionFailed(e.to_string()))?;
231 let key_id = Uuid::from_slice(&id_bytes).map_err(|_| SaasError::TenantNotFound)?;
232 let t_id = Uuid::from_slice(&tenant_bytes).map_err(|_| SaasError::TenantNotFound)?;
233
234 result.push(ApiKey {
235 id: ApiKeyId::from_uuid(key_id),
236 tenant_id: TenantId::from(t_id),
237 name: row.try_get("name")?,
238 scope: scopes,
239 created_at: row.try_get("created_at")?,
240 expires_at: row.try_get("expires_at")?,
241 last_used_at: row.try_get("last_used_at")?,
242 });
243 }
244 Ok(result)
245 }
246}
247
248#[cfg(test)]
249mod tests {
250 use super::*;
251 use crate::control_db::tests::test_pool;
252
253 async fn make_db() -> ControlDb {
254 let pool = test_pool().await;
255 ControlDb::new(pool).await.unwrap()
256 }
257
258 async fn make_tenant(db: &ControlDb) -> TenantId {
259 let plan_id: Vec<u8> = sqlx::query("SELECT id FROM tenant_plans LIMIT 1")
260 .fetch_one(db.pool())
261 .await
262 .unwrap()
263 .get("id");
264 let id = TenantId::new();
265 sqlx::query(
266 "INSERT INTO tenants (id, name, slug, owner_email, plan_id, status, db_path) \
267 VALUES (?, 'Test', 'test-slug', 'test@test.com', ?, 'active', 'test.db')",
268 )
269 .bind(id.as_bytes())
270 .bind(&plan_id)
271 .execute(db.pool())
272 .await
273 .unwrap();
274 id
275 }
276
277 #[tokio::test]
278 async fn mint_returns_raw_key_with_prefix() {
279 let db = make_db().await;
280 let tid = make_tenant(&db).await;
281 let result = db
282 .mint_api_key(&tid, "test-key", vec![ApiKeyScope::Admin], None)
283 .await
284 .unwrap();
285 assert!(result.raw_key.starts_with(KEY_PREFIX));
286 }
287
288 #[tokio::test]
289 async fn verify_valid_key_returns_some() {
290 let db = make_db().await;
291 let tid = make_tenant(&db).await;
292 let result = db
293 .mint_api_key(&tid, "valid", vec![ApiKeyScope::Admin], None)
294 .await
295 .unwrap();
296 let verified = db.verify_api_key(&result.raw_key).await.unwrap();
297 assert!(verified.is_some());
298 assert_eq!(verified.unwrap().name, "valid");
299 }
300
301 #[tokio::test]
302 async fn verify_garbage_key_returns_none() {
303 let db = make_db().await;
304 let _tid = make_tenant(&db).await;
305 let result = db.verify_api_key("sak_notavalidkey!!!").await.unwrap();
306 assert!(result.is_none());
307 }
308
309 #[tokio::test]
310 async fn verify_revoked_key_returns_none() {
311 let db = make_db().await;
312 let tid = make_tenant(&db).await;
313 let result = db
314 .mint_api_key(&tid, "to-revoke", vec![ApiKeyScope::Admin], None)
315 .await
316 .unwrap();
317 db.revoke_api_key(&result.api_key.id, &tid).await.unwrap();
318 let verified = db.verify_api_key(&result.raw_key).await.unwrap();
319 assert!(verified.is_none());
320 }
321
322 #[tokio::test]
323 async fn verify_expired_key_returns_none() {
324 let db = make_db().await;
325 let tid = make_tenant(&db).await;
326 let past = Utc::now() - chrono::Duration::hours(1);
327 let result = db
328 .mint_api_key(&tid, "expired", vec![ApiKeyScope::Admin], Some(past))
329 .await
330 .unwrap();
331 let verified = db.verify_api_key(&result.raw_key).await.unwrap();
332 assert!(verified.is_none());
333 }
334
335 #[tokio::test]
336 async fn list_excludes_revoked_keys() {
337 let db = make_db().await;
338 let tid = make_tenant(&db).await;
339 let r1 = db
340 .mint_api_key(&tid, "keep", vec![ApiKeyScope::Admin], None)
341 .await
342 .unwrap();
343 let r2 = db
344 .mint_api_key(&tid, "revoke-me", vec![ApiKeyScope::Admin], None)
345 .await
346 .unwrap();
347 db.revoke_api_key(&r2.api_key.id, &tid).await.unwrap();
348 let list = db.list_api_keys_for_tenant(&tid).await.unwrap();
349 assert_eq!(list.len(), 1);
350 assert_eq!(list[0].id, r1.api_key.id);
351 }
352
353 #[tokio::test]
354 async fn revoke_is_idempotent() {
355 let db = make_db().await;
356 let tid = make_tenant(&db).await;
357 let result = db
358 .mint_api_key(&tid, "idem", vec![ApiKeyScope::Admin], None)
359 .await
360 .unwrap();
361 db.revoke_api_key(&result.api_key.id, &tid).await.unwrap();
362 db.revoke_api_key(&result.api_key.id, &tid).await.unwrap();
363 let verified = db.verify_api_key(&result.raw_key).await.unwrap();
364 assert!(verified.is_none());
365 }
366
367 #[tokio::test]
368 async fn verify_updates_last_used_at() {
369 let db = make_db().await;
370 let tid = make_tenant(&db).await;
371 let result = db
372 .mint_api_key(&tid, "track", vec![ApiKeyScope::Admin], None)
373 .await
374 .unwrap();
375 let before = db.verify_api_key(&result.raw_key).await.unwrap().unwrap();
376 let _ = db.verify_api_key(&result.raw_key).await.unwrap().unwrap();
378 let after = db.verify_api_key(&result.raw_key).await.unwrap();
380 assert!(after.is_some());
381 let _ = before; }
383}