steid

@jamesgill /

1//! SQLite repository implementations.
2//!
3//! Rows are reassembled with `from_trusted`: they were validated on the way in, and
4//! re-validating them would make a tightened rule turn old rows unreadable.
5
6use std::time::{Duration, SystemTime};
7
8use sqlx::{Row, SqlitePool, sqlite::SqliteRow};
9
10use crate::domain::{
11 Email, Membership, MembershipId, OrgId, OrgName, Organization, PasswordHash, Role, Session,
12 SessionTokenHash, User, UserId,
13 repository::{
14 MembershipRepository, OrgRepository, RepositoryError, RepositoryResult, SessionRepository,
15 UserRepository,
16 },
17};
18
19fn backend(error: sqlx::Error) -> RepositoryError {
20 RepositoryError::backend(error)
21}
22
23#[derive(Debug, Clone)]
24pub struct SqliteUserRepo {
25 pool: SqlitePool,
26}
27
28impl SqliteUserRepo {
29 pub fn new(pool: SqlitePool) -> Self {
30 Self { pool }
31 }
32
33 fn map(row: &SqliteRow) -> User {
34 User::new(
35 UserId::from_trusted(row.get::<String, _>("id")),
36 Email::from_trusted(row.get::<String, _>("email")),
37 PasswordHash::from_trusted(row.get::<String, _>("password_hash")),
38 OrgId::from_trusted(row.get::<String, _>("personal_org_id")),
39 )
40 }
41}
42
43impl UserRepository for SqliteUserRepo {
44 async fn find_by_id(&self, id: &UserId) -> RepositoryResult<Option<User>> {
45 let row = sqlx::query("select * from users where id = ?")
46 .bind(id.as_str())
47 .fetch_optional(&self.pool)
48 .await
49 .map_err(backend)?;
50
51 Ok(row.as_ref().map(Self::map))
52 }
53
54 async fn find_by_email(&self, email: &Email) -> RepositoryResult<Option<User>> {
55 let row = sqlx::query("select * from users where email = ?")
56 .bind(email.as_str())
57 .fetch_optional(&self.pool)
58 .await
59 .map_err(backend)?;
60
61 Ok(row.as_ref().map(Self::map))
62 }
63
64 async fn save(&self, user: &User) -> RepositoryResult<()> {
65 sqlx::query(
66 "insert into users (id, email, password_hash, personal_org_id)
67 values (?, ?, ?, ?)
68 on conflict (id) do update set
69 email = excluded.email,
70 password_hash = excluded.password_hash,
71 personal_org_id = excluded.personal_org_id",
72 )
73 .bind(user.id.as_str())
74 .bind(user.email.as_str())
75 .bind(user.password_hash.as_str())
76 .bind(user.personal_org_id.as_str())
77 .execute(&self.pool)
78 .await
79 .map_err(backend)?;
80
81 Ok(())
82 }
83
84 async fn any_exist(&self) -> RepositoryResult<bool> {
85 let count: i64 = sqlx::query_scalar("select exists (select 1 from users)")
86 .fetch_one(&self.pool)
87 .await
88 .map_err(backend)?;
89
90 Ok(count != 0)
91 }
92}
93
94#[derive(Debug, Clone)]
95pub struct SqliteOrgRepo {
96 pool: SqlitePool,
97}
98
99impl SqliteOrgRepo {
100 pub fn new(pool: SqlitePool) -> Self {
101 Self { pool }
102 }
103
104 fn map(row: &SqliteRow) -> Organization {
105 Organization::from_trusted(
106 OrgId::from_trusted(row.get::<String, _>("id")),
107 OrgName::from_trusted(row.get::<String, _>("name")),
108 row.get::<Option<String>, _>("display_name"),
109 row.get::<Option<String>, _>("bio"),
110 )
111 }
112}
113
114impl OrgRepository for SqliteOrgRepo {
115 async fn find_by_id(&self, id: &OrgId) -> RepositoryResult<Option<Organization>> {
116 let row = sqlx::query("select * from orgs where id = ?")
117 .bind(id.as_str())
118 .fetch_optional(&self.pool)
119 .await
120 .map_err(backend)?;
121
122 Ok(row.as_ref().map(Self::map))
123 }
124
125 async fn find_by_name(&self, name: &OrgName) -> RepositoryResult<Option<Organization>> {
126 let row = sqlx::query("select * from orgs where name = ?")
127 .bind(name.as_str())
128 .fetch_optional(&self.pool)
129 .await
130 .map_err(backend)?;
131
132 Ok(row.as_ref().map(Self::map))
133 }
134
135 async fn save(&self, org: &Organization) -> RepositoryResult<()> {
136 sqlx::query(
137 "insert into orgs (id, name, display_name, bio)
138 values (?, ?, ?, ?)
139 on conflict (id) do update set
140 name = excluded.name,
141 display_name = excluded.display_name,
142 bio = excluded.bio",
143 )
144 .bind(org.id.as_str())
145 .bind(org.name.as_str())
146 .bind(org.display_name.as_deref())
147 .bind(org.bio.as_deref())
148 .execute(&self.pool)
149 .await
150 .map_err(backend)?;
151
152 Ok(())
153 }
154}
155
156#[derive(Debug, Clone)]
157pub struct SqliteMembershipRepo {
158 pool: SqlitePool,
159}
160
161impl SqliteMembershipRepo {
162 pub fn new(pool: SqlitePool) -> Self {
163 Self { pool }
164 }
165
166 /// An unparseable role is a storage fault, not a missing membership, so it
167 /// surfaces rather than silently downgrading the member's access.
168 fn map(row: &SqliteRow) -> RepositoryResult<Membership> {
169 let raw: String = row.get("role");
170 let role: Role = raw
171 .parse()
172 .map_err(|error| RepositoryError::backend(format!("{error}")))?;
173
174 Ok(Membership::new(
175 MembershipId::from_trusted(row.get::<String, _>("id")),
176 OrgId::from_trusted(row.get::<String, _>("org_id")),
177 UserId::from_trusted(row.get::<String, _>("user_id")),
178 role,
179 ))
180 }
181}
182
183impl MembershipRepository for SqliteMembershipRepo {
184 async fn find(&self, org_id: &OrgId, user_id: &UserId) -> RepositoryResult<Option<Membership>> {
185 let row = sqlx::query("select * from memberships where org_id = ? and user_id = ?")
186 .bind(org_id.as_str())
187 .bind(user_id.as_str())
188 .fetch_optional(&self.pool)
189 .await
190 .map_err(backend)?;
191
192 row.as_ref().map(Self::map).transpose()
193 }
194
195 async fn list_for_user(&self, user_id: &UserId) -> RepositoryResult<Vec<Membership>> {
196 let rows = sqlx::query("select * from memberships where user_id = ?")
197 .bind(user_id.as_str())
198 .fetch_all(&self.pool)
199 .await
200 .map_err(backend)?;
201
202 rows.iter().map(Self::map).collect()
203 }
204
205 async fn save(&self, membership: &Membership) -> RepositoryResult<()> {
206 sqlx::query(
207 "insert into memberships (id, org_id, user_id, role)
208 values (?, ?, ?, ?)
209 on conflict (id) do update set role = excluded.role",
210 )
211 .bind(membership.id.as_str())
212 .bind(membership.org_id.as_str())
213 .bind(membership.user_id.as_str())
214 .bind(membership.role.as_str())
215 .execute(&self.pool)
216 .await
217 .map_err(backend)?;
218
219 Ok(())
220 }
221}
222
223#[cfg(test)]
224mod tests {
225 use super::*;
226 use crate::{
227 application::{OwnerSpec, claim_instance, port::PasswordHasher},
228 domain::SetupToken,
229 infrastructure::{database::test_support::test_pool, password::StubHasher},
230 };
231
232 struct Repos {
233 users: SqliteUserRepo,
234 orgs: SqliteOrgRepo,
235 memberships: SqliteMembershipRepo,
236 }
237
238 async fn repos() -> Repos {
239 let pool = test_pool().await;
240 Repos {
241 users: SqliteUserRepo::new(pool.clone()),
242 orgs: SqliteOrgRepo::new(pool.clone()),
243 memberships: SqliteMembershipRepo::new(pool),
244 }
245 }
246
247 async fn saved_org(repos: &Repos, name: &str) -> Organization {
248 let org = Organization::new(OrgId::generate(), name, None).expect("valid org");
249 repos.orgs.save(&org).await.expect("save org");
250 org
251 }
252
253 async fn saved_user(repos: &Repos, email: &str, org: &Organization) -> User {
254 let user = User::new(
255 UserId::generate(),
256 Email::new(email).expect("valid email"),
257 PasswordHash::from_trusted("$argon2id$test"),
258 org.id.clone(),
259 );
260 repos.users.save(&user).await.expect("save user");
261 user
262 }
263
264 #[tokio::test]
265 async fn a_saved_user_round_trips() {
266 let repos = repos().await;
267 let org = saved_org(&repos, "james").await;
268 let user = saved_user(&repos, "dev@example.com", &org).await;
269
270 let found = repos
271 .users
272 .find_by_id(&user.id)
273 .await
274 .expect("lookup")
275 .expect("user should exist");
276
277 assert_eq!(found, user);
278 }
279
280 #[tokio::test]
281 async fn users_are_found_by_email_case_insensitively() {
282 let repos = repos().await;
283 let org = saved_org(&repos, "james").await;
284 saved_user(&repos, "dev@example.com", &org).await;
285
286 // Email lowercases on construction, but a row written before that rule would
287 // still need finding.
288 let found = repos
289 .users
290 .find_by_email(&Email::from_trusted("DEV@EXAMPLE.COM"))
291 .await
292 .expect("lookup");
293
294 assert!(found.is_some(), "collate nocase should make this match");
295 }
296
297 #[tokio::test]
298 async fn an_org_round_trips_with_its_display_name() {
299 let repos = repos().await;
300 let org = Organization::new(OrgId::generate(), "acme", Some("Acme".to_owned()))
301 .expect("valid org");
302 repos.orgs.save(&org).await.expect("save");
303
304 let found = repos
305 .orgs
306 .find_by_name(&org.name)
307 .await
308 .expect("lookup")
309 .expect("org should exist");
310
311 assert_eq!(found, org);
312 assert_eq!(found.label(), "Acme");
313 }
314
315 #[tokio::test]
316 async fn a_missing_org_is_none_not_an_error() {
317 let repos = repos().await;
318
319 let found = repos
320 .orgs
321 .find_by_name(&OrgName::new("nobody").unwrap())
322 .await
323 .expect("lookup");
324
325 assert_eq!(found, None);
326 }
327
328 #[tokio::test]
329 async fn a_membership_round_trips_with_its_role() {
330 let repos = repos().await;
331 let org = saved_org(&repos, "james").await;
332 let user = saved_user(&repos, "dev@example.com", &org).await;
333 let membership = Membership::new(
334 MembershipId::generate(),
335 org.id.clone(),
336 user.id.clone(),
337 Role::Owner,
338 );
339 repos.memberships.save(&membership).await.expect("save");
340
341 let found = repos
342 .memberships
343 .find(&org.id, &user.id)
344 .await
345 .expect("lookup")
346 .expect("membership should exist");
347
348 assert_eq!(found, membership);
349 assert!(found.can_write());
350 }
351
352 #[tokio::test]
353 async fn an_unreadable_role_surfaces_rather_than_downgrading_access() {
354 let repos = repos().await;
355 let pool = test_pool().await;
356 let memberships = SqliteMembershipRepo::new(pool.clone());
357 let org = Organization::new(OrgId::generate(), "james", None).expect("valid org");
358 SqliteOrgRepo::new(pool.clone())
359 .save(&org)
360 .await
361 .expect("save org");
362 let user = User::new(
363 UserId::generate(),
364 Email::new("dev@example.com").expect("valid email"),
365 PasswordHash::from_trusted("$argon2id$test"),
366 org.id.clone(),
367 );
368 SqliteUserRepo::new(pool.clone())
369 .save(&user)
370 .await
371 .expect("save user");
372
373 sqlx::query("insert into memberships (id, org_id, user_id, role) values (?, ?, ?, 'wat')")
374 .bind(MembershipId::generate().as_str())
375 .bind(org.id.as_str())
376 .bind(user.id.as_str())
377 .execute(&pool)
378 .await
379 .expect("insert");
380
381 let result = memberships.find(&org.id, &user.id).await;
382
383 assert!(
384 result.is_err(),
385 "a role we can't parse must not read as no membership"
386 );
387 drop(repos);
388 }
389
390 #[tokio::test]
391 async fn saving_a_user_before_its_org_is_refused() {
392 let repos = repos().await;
393 let orphan = User::new(
394 UserId::generate(),
395 Email::new("dev@example.com").expect("valid email"),
396 PasswordHash::from_trusted("$argon2id$test"),
397 OrgId::generate(),
398 );
399
400 let result = repos.users.save(&orphan).await;
401
402 assert!(
403 result.is_err(),
404 "the foreign key should reject a user whose org doesn't exist"
405 );
406 }
407
408 #[tokio::test]
409 async fn a_duplicate_handle_is_refused() {
410 let repos = repos().await;
411 saved_org(&repos, "james").await;
412
413 let clash = Organization::new(OrgId::generate(), "james", None).expect("valid org");
414 let result = repos.orgs.save(&clash).await;
415
416 assert!(result.is_err(), "orgs.name is unique");
417 }
418
419 #[tokio::test]
420 async fn a_differently_cased_handle_is_also_refused() {
421 let repos = repos().await;
422 saved_org(&repos, "james").await;
423
424 // OrgName lowercases, so this can only arrive via from_trusted -- but the
425 // constraint is what we're testing, not the value object.
426 let clash = Organization::from_trusted(
427 OrgId::generate(),
428 OrgName::from_trusted("JAMES"),
429 None,
430 None,
431 );
432 let result = repos.orgs.save(&clash).await;
433
434 assert!(result.is_err(), "collate nocase should catch this");
435 }
436
437 /// The claim use case checks `is_claimed` and then writes, which is TOCTOU. This
438 /// asserts the database is what actually stops a second owner being created.
439 #[tokio::test]
440 async fn a_second_claim_is_stopped_by_the_database_not_the_check() {
441 let repos = repos().await;
442 let token = SetupToken::generate();
443 let hasher = StubHasher::new();
444 let spec = OwnerSpec {
445 handle: "james".to_owned(),
446 email: "dev@example.com".to_owned(),
447 password: "hunter2".to_owned(),
448 };
449
450 claim_instance(
451 token.reveal(),
452 &token,
453 &spec,
454 &repos.users,
455 &repos.orgs,
456 &repos.memberships,
457 &hasher,
458 )
459 .await
460 .expect("first claim");
461
462 // Simulate the race: the second claimant passed is_claimed before the first
463 // committed, so it proceeds straight to the writes.
464 let intruder_org = Organization::new(OrgId::generate(), "james", None).expect("valid org");
465 let result = repos.orgs.save(&intruder_org).await;
466
467 assert!(
468 result.is_err(),
469 "unique(orgs.name) is what actually serialises concurrent claims"
470 );
471
472 let intruder_user = User::new(
473 UserId::generate(),
474 Email::new("dev@example.com").expect("valid email"),
475 hasher.hash("letmein").expect("hash"),
476 OrgId::generate(),
477 );
478 assert!(
479 repos.users.save(&intruder_user).await.is_err(),
480 "unique(users.email) closes the other half"
481 );
482 }
483}
484
485#[derive(Debug, Clone)]
486pub struct SqliteSessionRepo {
487 pool: SqlitePool,
488}
489
490impl SqliteSessionRepo {
491 pub fn new(pool: SqlitePool) -> Self {
492 Self { pool }
493 }
494}
495
496/// Unix seconds. Times before the epoch cannot occur here — sessions always expire in
497/// the future — so saturating at 0 is safe rather than lossy.
498fn to_unix(time: SystemTime) -> i64 {
499 time.duration_since(SystemTime::UNIX_EPOCH)
500 .map(|d| d.as_secs() as i64)
501 .unwrap_or(0)
502}
503
504fn from_unix(seconds: i64) -> SystemTime {
505 SystemTime::UNIX_EPOCH + Duration::from_secs(seconds.max(0) as u64)
506}
507
508impl SessionRepository for SqliteSessionRepo {
509 async fn find(&self, token_hash: &SessionTokenHash) -> RepositoryResult<Option<Session>> {
510 let row = sqlx::query("select * from sessions where token_hash = ?")
511 .bind(token_hash.as_str())
512 .fetch_optional(&self.pool)
513 .await
514 .map_err(backend)?;
515
516 Ok(row.map(|row| {
517 Session::new(
518 SessionTokenHash::from_trusted(row.get::<String, _>("token_hash")),
519 UserId::from_trusted(row.get::<String, _>("user_id")),
520 from_unix(row.get::<i64, _>("expires_at")),
521 )
522 }))
523 }
524
525 async fn save(&self, session: &Session) -> RepositoryResult<()> {
526 sqlx::query(
527 "insert into sessions (token_hash, user_id, expires_at)
528 values (?, ?, ?)
529 on conflict (token_hash) do update set
530 user_id = excluded.user_id,
531 expires_at = excluded.expires_at",
532 )
533 .bind(session.token_hash.as_str())
534 .bind(session.user_id.as_str())
535 .bind(to_unix(session.expires_at))
536 .execute(&self.pool)
537 .await
538 .map_err(backend)?;
539
540 Ok(())
541 }
542
543 async fn delete(&self, token_hash: &SessionTokenHash) -> RepositoryResult<()> {
544 sqlx::query("delete from sessions where token_hash = ?")
545 .bind(token_hash.as_str())
546 .execute(&self.pool)
547 .await
548 .map_err(backend)?;
549
550 Ok(())
551 }
552
553 async fn delete_expired(&self, now: SystemTime) -> RepositoryResult<u64> {
554 let result = sqlx::query("delete from sessions where expires_at <= ?")
555 .bind(to_unix(now))
556 .execute(&self.pool)
557 .await
558 .map_err(backend)?;
559
560 Ok(result.rows_affected())
561 }
562}