1use async_trait::async_trait;
2use deadpool::managed;
3use ldap3::{drive, Ldap, LdapConnAsync, LdapError, Mod, SearchEntry};
4use rand::prelude::SliceRandom;
5use rand::SeedableRng;
6use std::collections::HashSet;
7use std::sync::Arc;
8use trust_dns_resolver::{
9 config::{ResolverConfig, ResolverOpts},
10 AsyncResolver,
11};
12
13use super::search::SearchAttrs;
14use super::user::{LdapUser, LdapUserChangeSet};
15
16type Pool = managed::Pool<LdapManager>;
17
18#[derive(Clone)]
19pub struct LdapClient {
20 ldap: Arc<Pool>,
21}
22
23#[derive(Clone)]
24struct LdapManager {
25 ldap_servers: Vec<String>,
26 bind_dn: String,
27 bind_pw: String,
28}
29
30impl LdapManager {
31 pub async fn new(bind_dn: &str, bind_pw: &str) -> Self {
32 let ldap_servers = get_ldap_servers().await;
33
34 LdapManager {
35 ldap_servers,
36 bind_dn: bind_dn.to_owned(),
37 bind_pw: bind_pw.to_owned(),
38 }
39 }
40}
41
42#[async_trait]
43impl managed::Manager for LdapManager {
44 type Type = Ldap;
45 type Error = LdapError;
46
47 async fn create(&self) -> Result<Self::Type, Self::Error> {
48 let (conn, mut ldap) = LdapConnAsync::new(
49 self.ldap_servers
50 .choose(&mut rand::rngs::StdRng::from_entropy())
51 .unwrap(),
52 )
53 .await
54 .unwrap();
55 drive!(conn);
56
57 ldap.simple_bind(&self.bind_dn, &self.bind_pw)
58 .await
59 .unwrap();
60
61 Ok(ldap)
62 }
63
64 async fn recycle(&self, ldap: &mut Self::Type) -> managed::RecycleResult<Self::Error> {
65 ldap.extended(ldap3::exop::WhoAmI).await?;
66 Ok(())
67 }
68}
69
70impl LdapClient {
71 pub async fn new(bind_dn: &str, bind_pw: &str) -> Self {
72 let ldap_manager = LdapManager::new(bind_dn, bind_pw).await;
73 let ldap_pool = Pool::builder(ldap_manager).max_size(5).build().unwrap();
74
75 LdapClient {
76 ldap: Arc::new(ldap_pool),
77 }
78 }
79
80 pub async fn search_users(&mut self, query: &str) -> Vec<LdapUser> {
81 let mut ldap = self.ldap.get().await.unwrap();
82 ldap.with_timeout(std::time::Duration::from_secs(5));
83 let (results, _result) = ldap
84 .search(
85 "cn=users,cn=accounts,dc=csh,dc=rit,dc=edu",
86 ldap3::Scope::Subtree,
87 &format!("(|(uid=*{query}*)(cn=*{query}*))"),
88 SearchAttrs::default().finalize(),
89 )
90 .await
91 .unwrap()
92 .success()
93 .unwrap();
94
95 results
96 .iter()
97 .map(|result| {
98 let user = SearchEntry::construct(result.to_owned());
99 LdapUser::from_entry(&user)
100 })
101 .collect()
102 }
103
104 pub async fn _do_not_use_get_all_users(&mut self) -> Vec<LdapUser> {
105 let mut ldap = self.ldap.get().await.unwrap();
106
107 let (results, _result) = ldap
108 .search(
109 "cn=users,cn=accounts,dc=csh,dc=rit,dc=edu",
110 ldap3::Scope::Subtree,
111 "(objectClass=cshMember)",
112 SearchAttrs::default().finalize(),
113 )
114 .await
115 .unwrap()
116 .success()
117 .unwrap();
118
119 results
120 .iter()
121 .map(|result| {
122 let user = SearchEntry::construct(result.clone());
123 LdapUser::from_entry(&user)
124 })
125 .collect()
126 }
127
128 pub async fn get_user(&mut self, uid: &str) -> Option<LdapUser> {
129 let mut ldap = self.ldap.get().await.unwrap();
130
131 ldap.with_timeout(std::time::Duration::from_secs(5));
132 let (results, _result) = ldap
133 .search(
134 "cn=users,cn=accounts,dc=csh,dc=rit,dc=edu",
135 ldap3::Scope::Subtree,
136 &format!("uid={uid}"),
137 SearchAttrs::default().finalize(),
138 )
139 .await
140 .unwrap()
141 .success()
142 .unwrap();
143
144 if results.len() == 1 {
145 let user = SearchEntry::construct(results.get(0).unwrap().to_owned());
146 Some(LdapUser::from_entry(&user))
147 } else {
148 None
149 }
150 }
151
152 pub async fn get_user_by_ibutton(&mut self, ibutton: &str) -> Option<LdapUser> {
153 let mut ldap = self.ldap.get().await.unwrap();
154
155 ldap.with_timeout(std::time::Duration::from_secs(5));
156 let (results, _result) = ldap
157 .search(
158 "cn=users,cn=accounts,dc=csh,dc=rit,dc=edu",
159 ldap3::Scope::Subtree,
160 &format!("ibutton={ibutton}"),
161 SearchAttrs::default().finalize(),
162 )
163 .await
164 .unwrap()
165 .success()
166 .unwrap();
167
168 if results.len() == 1 {
169 let user = SearchEntry::construct(results.get(0).unwrap().to_owned());
170 Some(LdapUser::from_entry(&user))
171 } else {
172 None
173 }
174 }
175
176 pub async fn get_user_by_phone(&mut self, phone: &str) -> Option<LdapUser> {
177 let mut ldap = self.ldap.get().await.unwrap();
178 ldap.with_timeout(std::time::Duration::from_secs(5));
179 let (results, _result) = ldap
180 .search(
181 "cn=users,cn=accounts,dc=csh,dc=rit,dc=edu",
182 ldap3::Scope::Subtree,
183 &format!("mobile={phone}"),
184 SearchAttrs::default().finalize(),
185 )
186 .await
187 .unwrap()
188 .success()
189 .unwrap();
190
191 if results.len() == 1 {
192 let user = SearchEntry::construct(results.get(0).unwrap().to_owned());
193 Some(LdapUser::from_entry(&user))
194 } else {
195 None
196 }
197 }
198
199 pub async fn update_user(&mut self, change_set: &LdapUserChangeSet) {
200 let mut ldap = self.ldap.get().await.unwrap();
201
202 let mut changes = Vec::new();
203 if change_set.drinkBalance.is_some() {
204 changes.push(Mod::Replace(
205 String::from("drinkBalance"),
206 HashSet::from([change_set.drinkBalance.unwrap().to_string()]),
207 ));
208 }
209 match ldap.modify(&change_set.dn, changes).await {
210 Ok(_) => {}
211 Err(e) => eprintln!("{:#?}", e),
212 }
213 }
214}
215
216async fn get_ldap_servers() -> Vec<String> {
217 let resolver =
218 AsyncResolver::tokio(ResolverConfig::default(), ResolverOpts::default()).unwrap();
219 let response = resolver.srv_lookup("_ldap._tcp.csh.rit.edu").await.unwrap();
220
221 response
223 .iter()
224 .map(|record| {
225 format!(
226 "ldaps://{}",
227 record.target().to_string().trim_end_matches('.')
228 )
229 })
230 .collect()
231}