Skip to main content

systemprompt_users/services/user/
provider.rs

1//! `UserProvider` implementation over `UserService`.
2//!
3//! Copyright (c) systemprompt.io — Business Source License 1.1.
4//! See <https://systemprompt.io> for licensing details.
5
6use std::str::FromStr;
7
8use async_trait::async_trait;
9use systemprompt_identifiers::UserId;
10use systemprompt_traits::auth::{
11    AuthProviderError, AuthResult, AuthUser, FederatedIdentityClaims, RoleProvider, UserProvider,
12};
13
14use super::UserService;
15use crate::models::{User, UserRole};
16
17impl From<User> for AuthUser {
18    fn from(user: User) -> Self {
19        let is_active = user.is_active();
20        Self {
21            id: user.id,
22            name: user.name,
23            email: user.email,
24            roles: user.roles,
25            is_active,
26        }
27    }
28}
29
30#[async_trait]
31impl UserProvider for UserService {
32    async fn find_by_id(&self, id: &UserId) -> AuthResult<Option<AuthUser>> {
33        self.find_by_id(id)
34            .await
35            .map(|opt| opt.map(AuthUser::from))
36            .map_err(|e| AuthProviderError::Internal(e.into()))
37    }
38
39    async fn find_by_email(&self, email: &str) -> AuthResult<Option<AuthUser>> {
40        Self::find_by_email(self, email)
41            .await
42            .map(|opt| opt.map(AuthUser::from))
43            .map_err(|e| AuthProviderError::Internal(e.into()))
44    }
45
46    async fn find_by_name(&self, name: &str) -> AuthResult<Option<AuthUser>> {
47        Self::find_by_name(self, name)
48            .await
49            .map(|opt| opt.map(AuthUser::from))
50            .map_err(|e| AuthProviderError::Internal(e.into()))
51    }
52
53    async fn create_user(
54        &self,
55        name: &str,
56        email: &str,
57        full_name: Option<&str>,
58    ) -> AuthResult<AuthUser> {
59        Self::create(self, name, email, full_name, full_name)
60            .await
61            .map(AuthUser::from)
62            .map_err(|e| AuthProviderError::Internal(e.into()))
63    }
64
65    async fn create_anonymous(&self, fingerprint: &str) -> AuthResult<AuthUser> {
66        Self::create_anonymous(self, fingerprint)
67            .await
68            .map(AuthUser::from)
69            .map_err(|e| AuthProviderError::Internal(e.into()))
70    }
71
72    async fn assign_roles(&self, user_id: &UserId, roles: &[String]) -> AuthResult<()> {
73        Self::assign_roles(self, user_id, roles)
74            .await
75            .map(|_| ())
76            .map_err(|e| AuthProviderError::Internal(e.into()))
77    }
78
79    async fn find_or_create_federated(
80        &self,
81        issuer: &str,
82        external_sub: &str,
83        claims: &FederatedIdentityClaims,
84    ) -> AuthResult<UserId> {
85        Self::find_or_create_federated(self, issuer, external_sub, claims)
86            .await
87            .map(|u| u.id)
88            .map_err(|e| AuthProviderError::Internal(e.into()))
89    }
90
91    async fn promote_anonymous(&self, source: &UserId, target: &UserId) -> AuthResult<u64> {
92        Self::promote_anonymous(self, source, target)
93            .await
94            .map(|result| result.total_rows)
95            .map_err(|e| AuthProviderError::Internal(e.into()))
96    }
97}
98
99impl RoleProvider for UserService {
100    async fn get_roles(&self, user_id: &UserId) -> AuthResult<Vec<String>> {
101        match Self::find_by_id(self, user_id).await {
102            Ok(Some(user)) => Ok(user.roles),
103            Ok(None) => Err(AuthProviderError::UserNotFound),
104            Err(e) => Err(AuthProviderError::Internal(e.into())),
105        }
106    }
107
108    async fn assign_role(&self, user_id: &UserId, role: &str) -> AuthResult<()> {
109        let user = match Self::find_by_id(self, user_id).await {
110            Ok(Some(u)) => u,
111            Ok(None) => return Err(AuthProviderError::UserNotFound),
112            Err(e) => return Err(AuthProviderError::Internal(e.into())),
113        };
114
115        let mut roles = user.roles;
116        let role_str = role.to_owned();
117        if !roles.contains(&role_str) {
118            roles.push(role_str);
119        }
120
121        Self::assign_roles(self, user_id, &roles)
122            .await
123            .map(|_| ())
124            .map_err(|e| AuthProviderError::Internal(e.into()))
125    }
126
127    async fn revoke_role(&self, user_id: &UserId, role: &str) -> AuthResult<()> {
128        let user = match Self::find_by_id(self, user_id).await {
129            Ok(Some(u)) => u,
130            Ok(None) => return Err(AuthProviderError::UserNotFound),
131            Err(e) => return Err(AuthProviderError::Internal(e.into())),
132        };
133
134        let roles: Vec<String> = user.roles.into_iter().filter(|r| r != role).collect();
135
136        Self::assign_roles(self, user_id, &roles)
137            .await
138            .map(|_| ())
139            .map_err(|e| AuthProviderError::Internal(e.into()))
140    }
141
142    async fn list_users_by_role(&self, role: &str) -> AuthResult<Vec<AuthUser>> {
143        let Ok(user_role) = UserRole::from_str(role) else {
144            return Ok(vec![]);
145        };
146
147        Self::list_by_role(self, user_role)
148            .await
149            .map(|users| users.into_iter().map(AuthUser::from).collect())
150            .map_err(|e| AuthProviderError::Internal(e.into()))
151    }
152}