Skip to main content

csh_ldap/
client.rs

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    // TODO: Make sure servers are working
222    response
223        .iter()
224        .map(|record| {
225            format!(
226                "ldaps://{}",
227                record.target().to_string().trim_end_matches('.')
228            )
229        })
230        .collect()
231}