Skip to main content

valkey_module/context/
client.rs

1use crate::{
2    Context, RedisModuleClientInfo, RedisModule_DeauthenticateAndCloseClient,
3    RedisModule_GetClientCertificate, RedisModule_GetClientId, RedisModule_GetClientInfoById,
4    RedisModule_GetClientNameById, RedisModule_GetClientUserNameById,
5    RedisModule_SetClientNameById, Status, ValkeyError, ValkeyResult, ValkeyString, ValkeyValue,
6};
7use std::ffi::CStr;
8use std::os::raw::c_void;
9
10impl Default for RedisModuleClientInfo {
11    fn default() -> Self {
12        Self {
13            version: 1,
14            flags: 0,
15            id: 0,
16            addr: [0; 46],
17            port: 0,
18            db: 0,
19        }
20    }
21}
22
23/// GetClientNameById, GetClientUserNameById and GetClientCertificate use autoMemoryAdd on the ValkeyModuleString pointer
24/// after the callback (command, server event handler, ...) these ValkeyModuleString pointers will be freed automatically
25impl Context {
26    pub fn get_client_id(&self) -> u64 {
27        unsafe { RedisModule_GetClientId.unwrap()(self.ctx) }
28    }
29
30    /// wrapper for RedisModule_GetClientNameById
31    pub fn get_client_name_by_id(&self, client_id: u64) -> ValkeyResult<ValkeyString> {
32        let client_name = unsafe { RedisModule_GetClientNameById.unwrap()(self.ctx, client_id) };
33        if client_name.is_null() {
34            Err(ValkeyError::Str("Client/Client name is null"))
35        } else {
36            Ok(ValkeyString::from_redis_module_string(
37                self.ctx,
38                client_name,
39            ))
40        }
41    }
42
43    /// wrapper for RedisModule_GetClientNameById using current client ID
44    pub fn get_client_name(&self) -> ValkeyResult<ValkeyString> {
45        self.get_client_name_by_id(self.get_client_id())
46    }
47
48    /// wrapper for RedisModule_SetClientNameById
49    pub fn set_client_name_by_id(&self, client_id: u64, client_name: &ValkeyString) -> Status {
50        let resp = unsafe { RedisModule_SetClientNameById.unwrap()(client_id, client_name.inner) };
51        Status::from(resp)
52    }
53
54    /// wrapper for RedisModule_SetClientNameById using current client ID
55    pub fn set_client_name(&self, client_name: &ValkeyString) -> Status {
56        self.set_client_name_by_id(self.get_client_id(), client_name)
57    }
58
59    /// wrapper for RedisModule_GetClientUserNameById
60    pub fn get_client_username_by_id(&self, client_id: u64) -> ValkeyResult<ValkeyString> {
61        let client_username =
62            unsafe { RedisModule_GetClientUserNameById.unwrap()(self.ctx, client_id) };
63        if client_username.is_null() {
64            Err(ValkeyError::Str("Client/Username is null"))
65        } else {
66            Ok(ValkeyString::from_redis_module_string(
67                self.ctx,
68                client_username,
69            ))
70        }
71    }
72
73    /// wrapper for RedisModule_GetClientUserNameById using current client ID
74    pub fn get_client_username(&self) -> ValkeyResult<ValkeyString> {
75        self.get_client_username_by_id(self.get_client_id())
76    }
77
78    /// wrapper for RedisModule_GetClientCertificate
79    pub fn get_client_cert(&self) -> ValkeyResult<ValkeyString> {
80        let client_id = self.get_client_id();
81        let client_cert = unsafe { RedisModule_GetClientCertificate.unwrap()(self.ctx, client_id) };
82        if client_cert.is_null() {
83            Err(ValkeyError::Str("Client/Cert is null"))
84        } else {
85            Ok(ValkeyString::from_redis_module_string(
86                self.ctx,
87                client_cert,
88            ))
89        }
90    }
91
92    /// wrapper for RedisModule_GetClientInfoById
93    pub fn get_client_info_by_id(&self, client_id: u64) -> ValkeyResult<RedisModuleClientInfo> {
94        let mut mci = RedisModuleClientInfo::default();
95        let mci_ptr: *mut c_void = &mut mci as *mut _ as *mut c_void;
96        let status: Status =
97            unsafe { RedisModule_GetClientInfoById.unwrap()(mci_ptr, client_id).into() };
98        if status != Status::Ok {
99            Err(ValkeyError::Str("Client/Info not found"))
100        } else {
101            Ok(mci)
102        }
103    }
104
105    /// wrapper for RedisModule_GetClientInfoById using current client ID
106    pub fn get_client_info(&self) -> ValkeyResult<RedisModuleClientInfo> {
107        self.get_client_info_by_id(self.get_client_id())
108    }
109
110    /// wrapper to get the client IP address from RedisModuleClientInfo
111    pub fn get_client_ip_by_id(&self, client_id: u64) -> ValkeyResult<String> {
112        let client_info = self.get_client_info_by_id(client_id)?;
113        let c_str_addr = unsafe { CStr::from_ptr(client_info.addr.as_ptr()) };
114        let ip_addr_as_string = c_str_addr.to_string_lossy().into_owned();
115        Ok(ip_addr_as_string)
116    }
117
118    /// wrapper to get the client IP address from RedisModuleClientInfo using current client ID
119    pub fn get_client_ip(&self) -> ValkeyResult<String> {
120        self.get_client_ip_by_id(self.get_client_id())
121    }
122
123    pub fn deauthenticate_and_close_client_by_id(&self, client_id: u64) -> Status {
124        let resp =
125            unsafe { RedisModule_DeauthenticateAndCloseClient.unwrap()(self.ctx, client_id) };
126        Status::from(resp)
127    }
128
129    pub fn deauthenticate_and_close_client(&self) -> Status {
130        self.deauthenticate_and_close_client_by_id(self.get_client_id())
131    }
132
133    pub fn config_get(&self, config: String) -> ValkeyResult<ValkeyString> {
134        match self.call("CONFIG", &["GET", &config])? {
135            ValkeyValue::Array(array) if array.len() == 2 => match &array[1] {
136                ValkeyValue::SimpleString(val) => Ok(ValkeyString::create(None, val.clone())),
137                _ => Err(ValkeyError::Str("Config value is not a string")),
138            },
139            _ => Err(ValkeyError::Str("Unexpected CONFIG GET response")),
140        }
141    }
142}
143
144#[cfg(test)]
145mod tests {
146    use super::*;
147
148    const TEST_CLIENT_CERTIFICATE: &str = concat!(
149        "-----BEGIN CERTIFICATE-----\n",
150        "VGhpcyBpcyBhIHRlc3QgY2xpZW50IGNlcnRpZmljYXRlLg==\n",
151        "-----END CERTIFICATE-----\n"
152    );
153
154    #[test]
155    fn returns_current_client_id() {
156        let mut context = Context::test();
157        context.expect_get_client_id(42);
158
159        assert_eq!(context.get_client_id(), 42);
160    }
161
162    #[test]
163    fn gets_client_name_and_rejects_unknown_client_id() {
164        let mut context = Context::test();
165        context.expect_get_client_name_by_id(42, "alice");
166
167        assert_eq!(
168            context
169                .get_client_name()
170                .expect("configured client name should be returned")
171                .as_slice(),
172            b"alice"
173        );
174        assert!(matches!(
175            context.get_client_name_by_id(7),
176            Err(ValkeyError::Str("Client/Client name is null"))
177        ));
178    }
179
180    #[test]
181    fn gets_client_username_and_rejects_unknown_client_id() {
182        let mut context = Context::test();
183        context.expect_get_client_username_by_id(42, "alice");
184
185        assert_eq!(
186            context
187                .get_client_username()
188                .expect("configured client username should be returned")
189                .as_slice(),
190            b"alice"
191        );
192        assert!(matches!(
193            context.get_client_username_by_id(7),
194            Err(ValkeyError::Str("Client/Username is null"))
195        ));
196    }
197
198    #[test]
199    fn deauthenticates_configured_client_id_and_rejects_unknown_client_id() {
200        let mut context = Context::test();
201        context.expect_deauthenticate_and_close_client_by_id(42);
202
203        assert_eq!(context.deauthenticate_and_close_client(), Status::Ok);
204        assert_eq!(
205            context.deauthenticate_and_close_client_by_id(42),
206            Status::Ok
207        );
208        assert_eq!(
209            context.deauthenticate_and_close_client_by_id(7),
210            Status::Err
211        );
212    }
213
214    #[test]
215    fn gets_configured_client_info_and_rejects_unknown_client_id() {
216        let mut context = Context::test();
217        let client_info = RedisModuleClientInfo {
218            id: 42,
219            addr: [1; 46],
220            port: 6379,
221            db: 2,
222            ..RedisModuleClientInfo::default()
223        };
224        context.expect_get_client_info_by_id(client_info);
225
226        assert!(context.get_client_info().is_ok());
227        assert!(context.get_client_info_by_id(42).is_ok());
228        assert!(matches!(
229            context.get_client_info_by_id(7),
230            Err(ValkeyError::Str("Client/Info not found"))
231        ));
232    }
233
234    #[test]
235    fn client_info_default_uses_version_one_and_zeroed_fields() {
236        let client_info = RedisModuleClientInfo::default();
237
238        assert_eq!(client_info.version, 1);
239        assert_eq!(client_info.flags, 0);
240        assert_eq!(client_info.id, 0);
241        assert_eq!(client_info.addr, [0; 46]);
242        assert_eq!(client_info.port, 0);
243        assert_eq!(client_info.db, 0);
244    }
245
246    #[test]
247    fn gets_configured_client_ip_and_rejects_unknown_client_id() {
248        let mut context = Context::test();
249        context.expect_get_client_ip_by_id(42, "127.0.0.1");
250
251        assert_eq!(
252            context
253                .get_client_ip()
254                .expect("configured client IP should be returned"),
255            "127.0.0.1"
256        );
257        assert!(matches!(
258            context.get_client_ip_by_id(7),
259            Err(ValkeyError::Str("Client/Info not found"))
260        ));
261
262        context.expect_get_client_ip_by_id(42, "2001:db8::1");
263
264        assert_eq!(
265            context
266                .get_client_ip()
267                .expect("configured IPv6 client IP should be returned"),
268            "2001:db8::1"
269        );
270    }
271
272    #[test]
273    fn gets_configured_client_certificate_and_rejects_missing_certificate() {
274        let mut context = Context::test();
275        context.expect_get_client_cert(TEST_CLIENT_CERTIFICATE);
276
277        assert_eq!(
278            context
279                .get_client_cert()
280                .expect("configured client certificate should be returned")
281                .as_slice(),
282            TEST_CLIENT_CERTIFICATE.as_bytes()
283        );
284
285        let context = Context::test();
286        assert!(matches!(
287            context.get_client_cert(),
288            Err(ValkeyError::Str("Client/Cert is null"))
289        ));
290    }
291
292    #[test]
293    fn sets_current_client_name_and_rejects_unknown_client_id() {
294        let mut context = Context::test();
295        context.expect_set_client_name_by_id(42);
296        let client_name = context.create_string("bob");
297
298        assert_eq!(context.set_client_name(&client_name), Status::Ok);
299        assert_eq!(context.set_client_name_by_id(7, &client_name), Status::Err);
300    }
301
302    mod config_get {
303        use super::*;
304
305        #[test]
306        fn gets_configured_value() {
307            let mut context = Context::test();
308            context.expect_config_get("hz", "10");
309
310            assert_eq!(
311                context
312                    .config_get("hz".to_owned())
313                    .expect("configured config value should be returned")
314                    .as_slice(),
315                b"10"
316            );
317        }
318
319        #[test]
320        fn rejects_reply_with_unexpected_shape() {
321            let mut context = Context::test();
322            context.expect_call(
323                "CONFIG",
324                &["GET", "hz"],
325                ValkeyValue::Array(vec![ValkeyValue::SimpleString("hz".to_owned())]),
326            );
327
328            assert!(matches!(
329                context.config_get("hz".to_owned()),
330                Err(ValkeyError::Str("Unexpected CONFIG GET response"))
331            ));
332        }
333
334        #[test]
335        fn rejects_reply_with_non_string_value() {
336            let mut context = Context::test();
337            context.expect_call(
338                "CONFIG",
339                &["GET", "hz"],
340                ValkeyValue::Array(vec![
341                    ValkeyValue::SimpleString("hz".to_owned()),
342                    ValkeyValue::Integer(10),
343                ]),
344            );
345
346            assert!(matches!(
347                context.config_get("hz".to_owned()),
348                Err(ValkeyError::Str("Config value is not a string"))
349            ));
350        }
351
352        #[test]
353        fn propagates_call_error() {
354            let mut context = Context::test();
355            context.expect_call(
356                "CONFIG",
357                &["GET", "hz"],
358                ValkeyValue::StaticError("ERR permission denied"),
359            );
360
361            assert!(matches!(
362                context.config_get("hz".to_owned()),
363                Err(ValkeyError::String(message)) if message == "ERR permission denied"
364            ));
365        }
366
367        #[test]
368        fn rejects_unconfigured_test_call() {
369            let context = Context::test();
370
371            assert!(matches!(
372                context.config_get("hz".to_owned()),
373                Err(ValkeyError::String(message)) if message == "unexpected call: CONFIG GET hz"
374            ));
375        }
376    }
377}