miryad_core/users/
membership.rs1use sea_orm::entity::prelude::*;
2use sea_orm::{ConnectionTrait, Set};
3
4use crate::users::group::ensure_group;
5
6#[derive(Clone, Debug, PartialEq, DeriveEntityModel)]
7#[sea_orm(table_name = "miryad_group_memberships")]
8pub struct Model {
9 #[sea_orm(primary_key)]
10 pub id: i32,
11 pub user_id: i32,
12 pub group_id: i32,
13}
14
15#[derive(Copy, Clone, Debug, EnumIter, DeriveRelation)]
16pub enum Relation {}
17
18impl ActiveModelBehavior for ActiveModel {}
19
20pub type GroupMembership = Entity;
21
22pub async fn sync_group_memberships<C: ConnectionTrait>(
27 db: &C,
28 user_id: i32,
29 groups: &[String],
30) -> Result<(), DbErr> {
31 let mut wanted_group_ids = Vec::with_capacity(groups.len());
32 for name in groups {
33 wanted_group_ids.push(ensure_group(db, name).await?);
34 }
35
36 let current = Entity::find().filter(Column::UserId.eq(user_id)).all(db).await?;
37
38 for membership in ¤t {
39 if !wanted_group_ids.contains(&membership.group_id) {
40 Entity::delete_by_id(membership.id).exec(db).await?;
41 }
42 }
43
44 let current_group_ids: Vec<i32> = current.iter().map(|m| m.group_id).collect();
45 for group_id in wanted_group_ids {
46 if !current_group_ids.contains(&group_id) {
47 let active = ActiveModel {
48 user_id: Set(user_id),
49 group_id: Set(group_id),
50 ..Default::default()
51 };
52 Entity::insert(active)
59 .on_conflict_do_nothing_on([Column::UserId, Column::GroupId])
60 .exec(db)
61 .await?;
62 }
63 }
64
65 Ok(())
66}
67
68#[cfg(test)]
69mod tests {
70 use super::*;
71 use crate::migration::Migrator;
72 use crate::users::user::resolve_user;
73 use sea_orm_migration::MigratorTrait;
74
75 async fn test_db() -> DatabaseConnection {
76 let db = sea_orm::Database::connect("sqlite::memory:")
77 .await
78 .expect("in-memory sqlite connects");
79 Migrator::up(&db, None).await.expect("migrations apply cleanly");
80 db
81 }
82
83 #[tokio::test]
84 async fn sync_adds_missing_groups_including_unknown_ones() {
85 let db = test_db().await;
86 let user = resolve_user(&db, "sub-1", None).await.expect("resolve succeeds");
87
88 sync_group_memberships(&db, user.id, &["admin".to_string(), "editors".to_string()])
89 .await
90 .expect("sync succeeds");
91
92 let memberships = Entity::find()
93 .filter(Column::UserId.eq(user.id))
94 .all(&db)
95 .await
96 .expect("query succeeds");
97 assert_eq!(memberships.len(), 2);
98 }
99
100 #[tokio::test]
101 async fn second_sync_with_fewer_groups_removes_stale_memberships() {
102 let db = test_db().await;
103 let user = resolve_user(&db, "sub-1", None).await.expect("resolve succeeds");
104
105 sync_group_memberships(&db, user.id, &["admin".to_string(), "editors".to_string()])
106 .await
107 .expect("first sync succeeds");
108 sync_group_memberships(&db, user.id, &["editors".to_string()])
109 .await
110 .expect("second sync succeeds");
111
112 let memberships = Entity::find()
113 .filter(Column::UserId.eq(user.id))
114 .all(&db)
115 .await
116 .expect("query succeeds");
117 assert_eq!(memberships.len(), 1);
118 }
119
120 #[tokio::test]
121 async fn sync_is_idempotent() {
122 let db = test_db().await;
123 let user = resolve_user(&db, "sub-1", None).await.expect("resolve succeeds");
124
125 sync_group_memberships(&db, user.id, &["editors".to_string()])
126 .await
127 .expect("first sync succeeds");
128 sync_group_memberships(&db, user.id, &["editors".to_string()])
129 .await
130 .expect("second sync succeeds");
131
132 let memberships = Entity::find()
133 .filter(Column::UserId.eq(user.id))
134 .all(&db)
135 .await
136 .expect("query succeeds");
137 assert_eq!(memberships.len(), 1);
138 }
139
140 #[tokio::test]
147 async fn concurrent_syncs_for_the_same_user_do_not_violate_the_unique_constraint() {
148 let db = test_db().await;
149 let user = resolve_user(&db, "sub-1", None).await.expect("resolve succeeds");
150 let groups = vec!["admin".to_string(), "editors".to_string()];
151
152 let (first, second) = tokio::join!(
153 sync_group_memberships(&db, user.id, &groups),
154 sync_group_memberships(&db, user.id, &groups),
155 );
156 first.expect("first concurrent sync succeeds");
157 second.expect("second concurrent sync succeeds");
158
159 let memberships = Entity::find()
160 .filter(Column::UserId.eq(user.id))
161 .all(&db)
162 .await
163 .expect("query succeeds");
164 assert_eq!(memberships.len(), 2);
165 }
166}