1
mod check_authorization;
2
mod complete_database_name;
3
mod complete_user_name;
4
mod create_databases;
5
mod create_users;
6
mod drop_databases;
7
mod drop_users;
8
mod list_all_databases;
9
mod list_all_privileges;
10
mod list_all_users;
11
mod list_databases;
12
mod list_privileges;
13
mod list_users;
14
mod list_valid_name_prefixes;
15
mod lock_users;
16
mod modify_privileges;
17
mod passwd_user;
18
mod unlock_users;
19

            
20
pub use check_authorization::*;
21
pub use complete_database_name::*;
22
pub use complete_user_name::*;
23
pub use create_databases::*;
24
pub use create_users::*;
25
pub use drop_databases::*;
26
pub use drop_users::*;
27
pub use list_all_databases::*;
28
pub use list_all_privileges::*;
29
pub use list_all_users::*;
30
pub use list_databases::*;
31
pub use list_privileges::*;
32
pub use list_users::*;
33
pub use list_valid_name_prefixes::*;
34
pub use lock_users::*;
35
pub use modify_privileges::*;
36
pub use passwd_user::*;
37
pub use unlock_users::*;
38

            
39
use std::collections::BTreeSet;
40
use std::fmt;
41

            
42
use serde::{Deserialize, Serialize};
43
use tokio::net::UnixStream;
44
use tokio_serde::{Framed as SerdeFramed, formats::Json};
45
use tokio_util::codec::{Framed, LengthDelimitedCodec};
46

            
47
use crate::core::types::{MySQLDatabase, MySQLUser};
48

            
49
pub type ServerToClientMessageStream = SerdeFramed<
50
    Framed<UnixStream, LengthDelimitedCodec>,
51
    Request,
52
    Response,
53
    Json<Request, Response>,
54
>;
55

            
56
pub type ClientToServerMessageStream = SerdeFramed<
57
    Framed<UnixStream, LengthDelimitedCodec>,
58
    Response,
59
    Request,
60
    Json<Response, Request>,
61
>;
62

            
63
const MAX_REQUEST_FRAME_LENGTH: usize = 100 * 1024; // 100 KB
64
const MAX_RESPONSE_FRAME_LENGTH: usize = 1024 * 1024; // 1 MB
65

            
66
pub fn create_client_to_server_message_stream(socket: UnixStream) -> ClientToServerMessageStream {
67
    let codec = {
68
        let mut codec = LengthDelimitedCodec::new();
69
        codec.set_max_frame_length(MAX_REQUEST_FRAME_LENGTH);
70
        codec
71
    };
72
    let length_delimited = Framed::new(socket, codec);
73
    tokio_serde::Framed::new(length_delimited, Json::default())
74
}
75

            
76
pub fn create_server_to_client_message_stream(socket: UnixStream) -> ServerToClientMessageStream {
77
    let codec = {
78
        let mut codec = LengthDelimitedCodec::new();
79
        codec.set_max_frame_length(MAX_RESPONSE_FRAME_LENGTH);
80
        codec
81
    };
82
    let length_delimited = Framed::new(socket, codec);
83
    tokio_serde::Framed::new(length_delimited, Json::default())
84
}
85

            
86
#[non_exhaustive]
87
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
88
pub enum Request {
89
    CheckAuthorization(CheckAuthorizationRequest),
90

            
91
    ListValidNamePrefixes,
92
    CompleteDatabaseName(CompleteDatabaseNameRequest),
93
    CompleteUserName(CompleteUserNameRequest),
94

            
95
    CreateDatabases(CreateDatabasesRequest),
96
    DropDatabases(DropDatabasesRequest),
97
    ListDatabases(ListDatabasesRequest),
98
    ListPrivileges(ListPrivilegesRequest),
99
    ModifyPrivileges(ModifyPrivilegesRequest),
100

            
101
    CreateUsers(CreateUsersRequest),
102
    DropUsers(DropUsersRequest),
103
    PasswdUser(SetUserPasswordRequest),
104
    ListUsers(ListUsersRequest),
105
    LockUsers(LockUsersRequest),
106
    UnlockUsers(UnlockUsersRequest),
107

            
108
    // Commit,
109
    Exit,
110
}
111

            
112
impl Request {
113
    /// Get the command name associated with this request.
114
    pub fn command_name(&self) -> &str {
115
        match self {
116
            Request::CheckAuthorization(_) => "check-authorization",
117
            Request::ListValidNamePrefixes => "list-valid-name-prefixes",
118
            Request::CompleteDatabaseName(_) => "complete-database-name",
119
            Request::CompleteUserName(_) => "complete-user-name",
120
            Request::CreateDatabases(_) => "create-databases",
121
            Request::DropDatabases(_) => "drop-databases",
122
            Request::ListDatabases(_) => "list-databases",
123
            Request::ListPrivileges(_) => "list-privileges",
124
            Request::ModifyPrivileges(_) => "modify-privileges",
125
            Request::CreateUsers(_) => "create-users",
126
            Request::DropUsers(_) => "drop-users",
127
            Request::PasswdUser(_) => "passwd-user",
128
            Request::ListUsers(_) => "list-users",
129
            Request::LockUsers(_) => "lock-users",
130
            Request::UnlockUsers(_) => "unlock-users",
131
            Request::Exit => "exit",
132
        }
133
    }
134

            
135
    /// Generate a short summary string representing this request for logging purposes.
136
    pub fn log_summary(&self) -> String {
137
        match self {
138
            Request::CheckAuthorization(req) => format!("{}({})", self.command_name(), req.len()),
139

            
140
            Request::CreateDatabases(req) => format!("{}({})", self.command_name(), req.len()),
141
            Request::DropDatabases(req) => format!("{}({})", self.command_name(), req.len()),
142
            Request::ListDatabases(req) => format!(
143
                "{}{}",
144
                self.command_name(),
145
                req.names
146
                    .as_ref()
147
                    .map_or("".to_string(), |r| format!("({})", r.len()))
148
            ),
149
            Request::ListPrivileges(req) => format!(
150
                "{}{}",
151
                self.command_name(),
152
                req.as_ref()
153
                    .map_or("".to_string(), |r| format!("({})", r.len()))
154
            ),
155
            Request::ModifyPrivileges(req) => format!("{}({})", self.command_name(), req.len()),
156

            
157
            Request::CreateUsers(req) => format!("{}({})", self.command_name(), req.len()),
158
            Request::DropUsers(req) => format!("{}({})", self.command_name(), req.len()),
159
            Request::ListUsers(req) => format!(
160
                "{}{}",
161
                self.command_name(),
162
                req.names
163
                    .as_ref()
164
                    .map_or("".to_string(), |r| format!("({})", r.len()))
165
            ),
166
            Request::LockUsers(req) => format!("{}({})", self.command_name(), req.len()),
167
            Request::UnlockUsers(req) => format!("{}({})", self.command_name(), req.len()),
168

            
169
            _ => self.command_name().to_string(),
170
        }
171
    }
172

            
173
    /// Get the set of users affected by this request.
174
    pub fn affected_users(&self) -> BTreeSet<MySQLUser> {
175
        match self {
176
            Request::CheckAuthorization(_) => Default::default(),
177
            Request::ListValidNamePrefixes => Default::default(),
178
            Request::CompleteDatabaseName(_) => Default::default(),
179
            Request::CompleteUserName(_) => Default::default(),
180
            Request::CreateDatabases(_) => Default::default(),
181
            Request::DropDatabases(_) => Default::default(),
182
            Request::ListDatabases(_) => Default::default(),
183
            Request::ListPrivileges(_) => Default::default(),
184
            Request::ModifyPrivileges(priv_diffs) => priv_diffs
185
                .iter()
186
                .map(|priv_diff| priv_diff.get_user_name().clone())
187
                .collect(),
188
            Request::CreateUsers(users) => users.iter().cloned().collect(),
189
            Request::DropUsers(users) => users.iter().cloned().collect(),
190
            Request::PasswdUser(user_passwd_req) => {
191
                let mut result = BTreeSet::new();
192
                result.insert(user_passwd_req.0.clone());
193
                result
194
            }
195
            Request::ListUsers(request) => request
196
                .names
197
                .clone()
198
                .unwrap_or_default()
199
                .into_iter()
200
                .collect(),
201
            Request::LockUsers(users) => users.iter().cloned().collect(),
202
            Request::UnlockUsers(users) => users.iter().cloned().collect(),
203
            Request::Exit => Default::default(),
204
        }
205
    }
206

            
207
    /// Get the set of databases affected by this request.
208
    pub fn affected_databases(&self) -> BTreeSet<MySQLDatabase> {
209
        match self {
210
            Request::CheckAuthorization(_) => Default::default(),
211
            Request::ListValidNamePrefixes => Default::default(),
212
            Request::CompleteDatabaseName(_) => Default::default(),
213
            Request::CompleteUserName(_) => Default::default(),
214
            Request::CreateDatabases(databases) => databases.iter().cloned().collect(),
215
            Request::DropDatabases(databases) => databases.iter().cloned().collect(),
216
            Request::ListDatabases(request) => request
217
                .names
218
                .clone()
219
                .unwrap_or_default()
220
                .into_iter()
221
                .collect(),
222
            Request::ListPrivileges(databases) => {
223
                databases.clone().unwrap_or_default().into_iter().collect()
224
            }
225
            Request::ModifyPrivileges(priv_diffs) => priv_diffs
226
                .iter()
227
                .map(|priv_diff| priv_diff.get_database_name().clone())
228
                .collect(),
229
            Request::CreateUsers(_) => Default::default(),
230
            Request::DropUsers(_) => Default::default(),
231
            Request::PasswdUser(_) => Default::default(),
232
            Request::ListUsers(_) => Default::default(),
233
            Request::LockUsers(_) => Default::default(),
234
            Request::UnlockUsers(_) => Default::default(),
235
            Request::Exit => Default::default(),
236
        }
237
    }
238
}
239

            
240
// TODO: include a generic "message" that will display a message to the user?
241

            
242
#[non_exhaustive]
243
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
244
pub enum Response {
245
    CheckAuthorization(CheckAuthorizationResponse),
246

            
247
    ListValidNamePrefixes(ListValidNamePrefixesResponse),
248
    CompleteDatabaseName(CompleteDatabaseNameResponse),
249
    CompleteUserName(CompleteUserNameResponse),
250

            
251
    // Specific data for specific commands
252
    CreateDatabases(CreateDatabasesResponse),
253
    DropDatabases(DropDatabasesResponse),
254
    ListDatabases(ListDatabasesResponse),
255
    ListAllDatabases(ListAllDatabasesResponse),
256
    ListPrivileges(ListPrivilegesResponse),
257
    ListAllPrivileges(ListAllPrivilegesResponse),
258
    ModifyPrivileges(ModifyPrivilegesResponse),
259

            
260
    CreateUsers(CreateUsersResponse),
261
    DropUsers(DropUsersResponse),
262
    SetUserPassword(SetUserPasswordResponse),
263
    ListUsers(ListUsersResponse),
264
    ListAllUsers(ListAllUsersResponse),
265
    LockUsers(LockUsersResponse),
266
    UnlockUsers(UnlockUsersResponse),
267

            
268
    // Generic responses
269
    Ready,
270
    Error(String),
271
}
272

            
273
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
274
pub enum ResponseOkStatus {
275
    Success,
276
    PartialSuccess(usize, usize), // succeeded, total
277
    Error,
278
}
279

            
280
impl ResponseOkStatus {
281
    pub fn from_counts(total: usize, succeeded: usize) -> Self {
282
        if succeeded == total {
283
            ResponseOkStatus::Success
284
        } else if succeeded == 0 {
285
            ResponseOkStatus::Error
286
        } else {
287
            ResponseOkStatus::PartialSuccess(succeeded, total)
288
        }
289
    }
290

            
291
    pub fn from_bool(is_ok: bool) -> Self {
292
        if is_ok {
293
            ResponseOkStatus::Success
294
        } else {
295
            ResponseOkStatus::Error
296
        }
297
    }
298
}
299

            
300
impl fmt::Display for ResponseOkStatus {
301
    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
302
        match self {
303
            ResponseOkStatus::Success => write!(f, "OK"),
304
            ResponseOkStatus::PartialSuccess(succeeded, total) => {
305
                write!(f, "PARTIAL_OK({}/{})", succeeded, total)
306
            }
307
            ResponseOkStatus::Error => write!(f, "ERR"),
308
        }
309
    }
310
}
311

            
312
impl Response {
313
    pub fn ok_status(&self) -> ResponseOkStatus {
314
        match self {
315
            Response::CheckAuthorization(res) => {
316
                ResponseOkStatus::from_counts(res.len(), res.values().filter(|v| v.is_ok()).count())
317
            }
318

            
319
            Response::ListValidNamePrefixes(_) => ResponseOkStatus::Success,
320
            Response::CompleteDatabaseName(_) => ResponseOkStatus::Success,
321
            Response::CompleteUserName(_) => ResponseOkStatus::Success,
322

            
323
            Response::CreateDatabases(res) => {
324
                ResponseOkStatus::from_counts(res.len(), res.values().filter(|v| v.is_ok()).count())
325
            }
326
            Response::DropDatabases(res) => {
327
                ResponseOkStatus::from_counts(res.len(), res.values().filter(|v| v.is_ok()).count())
328
            }
329
            Response::ListDatabases(res) => {
330
                ResponseOkStatus::from_counts(res.len(), res.values().filter(|v| v.is_ok()).count())
331
            }
332
            Response::ListAllDatabases(res) => ResponseOkStatus::from_bool(res.is_ok()),
333
            Response::ListPrivileges(res) => {
334
                ResponseOkStatus::from_counts(res.len(), res.values().filter(|v| v.is_ok()).count())
335
            }
336
            Response::ListAllPrivileges(res) => ResponseOkStatus::from_bool(res.is_ok()),
337
            Response::ModifyPrivileges(res) => ResponseOkStatus::from_counts(
338
                res.len(),
339
                res.values()
340
                    .map(|user_map| user_map.values().filter(|v| v.is_ok()).count())
341
                    .sum(),
342
            ),
343

            
344
            Response::CreateUsers(res) => {
345
                ResponseOkStatus::from_counts(res.len(), res.values().filter(|v| v.is_ok()).count())
346
            }
347
            Response::DropUsers(res) => {
348
                ResponseOkStatus::from_counts(res.len(), res.values().filter(|v| v.is_ok()).count())
349
            }
350
            Response::SetUserPassword(res) => ResponseOkStatus::from_bool(res.is_ok()),
351
            Response::ListUsers(res) => {
352
                ResponseOkStatus::from_counts(res.len(), res.values().filter(|v| v.is_ok()).count())
353
            }
354
            Response::ListAllUsers(res) => ResponseOkStatus::from_bool(res.is_ok()),
355
            Response::LockUsers(res) => {
356
                ResponseOkStatus::from_counts(res.len(), res.values().filter(|v| v.is_ok()).count())
357
            }
358
            Response::UnlockUsers(res) => {
359
                ResponseOkStatus::from_counts(res.len(), res.values().filter(|v| v.is_ok()).count())
360
            }
361

            
362
            Response::Ready => ResponseOkStatus::Success,
363
            Response::Error(_) => ResponseOkStatus::Error,
364
        }
365
    }
366
}
367

            
368
// Utility function to format a list of items, used in list commands.
369
4
pub(crate) fn format_possibly_truncated_list<'a>(
370
4
    items: impl Iterator<Item = &'a str>,
371
4
    total_count: u64,
372
4
) -> String {
373
4
    let items: Vec<&str> = items.collect();
374
4
    let mut joined = items.join("\n");
375

            
376
4
    let remaining = total_count.saturating_sub(items.len() as u64);
377
4
    if remaining > 0 {
378
2
        if !joined.is_empty() {
379
1
            joined.push('\n');
380
1
        }
381
2
        joined.push_str(&format!("... ({remaining} more)"));
382
2
    }
383

            
384
4
    joined
385
4
}
386

            
387
#[cfg(test)]
388
mod tests {
389
    use super::*;
390

            
391
    #[test]
392
1
    fn test_format_possibly_truncated_list() {
393
1
        assert_eq!(format_possibly_truncated_list([].into_iter(), 0), "");
394

            
395
1
        assert_eq!(
396
1
            format_possibly_truncated_list(["a", "b"].into_iter(), 2),
397
            "a\nb"
398
        );
399

            
400
1
        assert_eq!(
401
1
            format_possibly_truncated_list(["a", "b"].into_iter(), 5),
402
            "a\nb\n... (3 more)"
403
        );
404

            
405
1
        assert_eq!(
406
1
            format_possibly_truncated_list([].into_iter(), 5),
407
            "... (5 more)"
408
        );
409
1
    }
410
}