steid

@jamesgill /

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