Skip to main content

allowthem_saas/
api_keys.rs

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/// Identifier for an API key row.
17#[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    /// Plaintext key — return to caller once, never stored.
66    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    /// Verifies a raw key, updates last_used_at, and returns the matching ApiKey.
128    /// Returns `None` if the key is not found, revoked, or expired.
129    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        // key_hash has a UNIQUE constraint, so at most one row.
136        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        // Constant-time guard: ensure the stored hash truly matches the candidate.
150        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        // last_used_at starts None (or Some on second call)
377        let _ = db.verify_api_key(&result.raw_key).await.unwrap().unwrap();
378        // Verify the key still works after use-tracking
379        let after = db.verify_api_key(&result.raw_key).await.unwrap();
380        assert!(after.is_some());
381        let _ = before; // suppress unused
382    }
383}