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
23impl Context {
26 pub fn get_client_id(&self) -> u64 {
27 unsafe { RedisModule_GetClientId.unwrap()(self.ctx) }
28 }
29
30 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 pub fn get_client_name(&self) -> ValkeyResult<ValkeyString> {
45 self.get_client_name_by_id(self.get_client_id())
46 }
47
48 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 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 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 pub fn get_client_username(&self) -> ValkeyResult<ValkeyString> {
75 self.get_client_username_by_id(self.get_client_id())
76 }
77
78 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 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 pub fn get_client_info(&self) -> ValkeyResult<RedisModuleClientInfo> {
107 self.get_client_info_by_id(self.get_client_id())
108 }
109
110 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 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}