1
use std::sync::{
2
    Arc,
3
    atomic::{AtomicU64, Ordering},
4
};
5

            
6
use futures_util::{SinkExt, StreamExt};
7
use indoc::concatdoc;
8
use itertools::Itertools;
9
use sqlx::{MySqlConnection, MySqlPool};
10
use tokio::net::UnixStream;
11
use tracing::Instrument;
12

            
13
use crate::{
14
    core::{
15
        common::UnixUser,
16
        protocol::{
17
            PasswordSource, Request, Response, ServerToClientMessageStream, SetPasswordError,
18
            create_server_to_client_message_stream, request_validation::GroupDenylist,
19
        },
20
    },
21
    server::{
22
        authorization::check_authorization,
23
        common::get_user_filtered_groups,
24
        sql::{
25
            database_operations::{
26
                complete_database_name, create_databases, drop_databases,
27
                list_all_databases_for_user, list_databases,
28
            },
29
            database_privilege_operations::{
30
                apply_privilege_diffs, get_all_database_privileges, get_databases_privilege_data,
31
            },
32
            user_operations::{
33
                complete_user_name, create_database_users, drop_database_users,
34
                list_all_database_users_for_unix_user, list_database_users, lock_database_users,
35
                set_password_for_database_user, unlock_database_users,
36
            },
37
        },
38
    },
39
};
40

            
41
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
42
pub struct SessionId(u64);
43

            
44
impl SessionId {
45
    pub fn new(id: u64) -> Self {
46
        SessionId(id)
47
    }
48

            
49
    pub fn inner(&self) -> u64 {
50
        self.0
51
    }
52
}
53

            
54
// TODO: don't use database connection unless necessary.
55

            
56
pub async fn session_handler(
57
    socket: UnixStream,
58
    session_id: SessionId,
59
    db_pool: Arc<MySqlPool>,
60
    db_is_mariadb: bool,
61
    group_denylist: &GroupDenylist,
62
    request_counter: Arc<AtomicU64>,
63
) -> anyhow::Result<()> {
64
    let uid = match socket.peer_cred() {
65
        Ok(cred) => cred.uid(),
66
        Err(e) => {
67
            tracing::error!("Failed to get peer credentials from socket: {}", e);
68
            let mut message_stream = create_server_to_client_message_stream(socket);
69
            message_stream
70
                .send(Response::Error(
71
                    (concatdoc! {
72
                        "Server failed to get peer credentials from socket\n",
73
                        "Please check the server logs or contact the system administrators"
74
                    })
75
                    .to_string(),
76
                ))
77
                .await
78
                .ok();
79
            anyhow::bail!("Failed to get peer credentials from socket");
80
        }
81
    };
82

            
83
    tracing::trace!("Validated peer UID: {}", uid);
84

            
85
    let unix_user = match UnixUser::from_uid(uid) {
86
        Ok(user) => user,
87
        Err(e) => {
88
            tracing::error!("Failed to get username from uid: {}", e);
89
            let mut message_stream = create_server_to_client_message_stream(socket);
90
            message_stream
91
                .send(Response::Error(
92
                    (concatdoc! {
93
                        "Server failed to get user data from the system\n",
94
                        "Please check the server logs or contact the system administrators"
95
                    })
96
                    .to_string(),
97
                ))
98
                .await
99
                .ok();
100
            anyhow::bail!("Failed to get username from uid: {e}");
101
        }
102
    };
103

            
104
    let span = tracing::info_span!(
105
        "user_session",
106
        session_id = session_id.inner(),
107
        user = %unix_user,
108
    );
109

            
110
    (async move {
111
        tracing::debug!("Accepted connection from user: {}", unix_user);
112

            
113
        let result = session_handler_with_unix_user(
114
            socket,
115
            session_id,
116
            &unix_user,
117
            db_pool,
118
            db_is_mariadb,
119
            group_denylist,
120
            request_counter,
121
        )
122
        .await;
123

            
124
        tracing::debug!(
125
            "Finished handling requests for connection from user: {}",
126
            unix_user,
127
        );
128

            
129
        result
130
    })
131
    .instrument(span)
132
    .await
133
}
134

            
135
pub async fn session_handler_with_unix_user(
136
    socket: UnixStream,
137
    session_id: SessionId,
138
    unix_user: &UnixUser,
139
    db_pool: Arc<MySqlPool>,
140
    db_is_mariadb: bool,
141
    group_denylist: &GroupDenylist,
142
    request_counter: Arc<AtomicU64>,
143
) -> anyhow::Result<()> {
144
    let mut message_stream = create_server_to_client_message_stream(socket);
145

            
146
    tracing::trace!("Requesting database connection from pool");
147
    let mut db_connection = match db_pool.acquire().await {
148
        Ok(connection) => connection,
149
        Err(err) => {
150
            message_stream
151
                .send(Response::Error(
152
                    (concatdoc! {
153
                        "Server failed to connect to database\n",
154
                        "Please check the server logs or contact the system administrators"
155
                    })
156
                    .to_string(),
157
                ))
158
                .await?;
159
            message_stream.flush().await?;
160
            return Err(err.into());
161
        }
162
    };
163
    tracing::trace!("Successfully acquired database connection from pool");
164

            
165
    let result = session_handler_with_db_connection(
166
        message_stream,
167
        session_id,
168
        unix_user,
169
        &mut db_connection,
170
        db_is_mariadb,
171
        group_denylist,
172
        request_counter,
173
    )
174
    .await;
175

            
176
    tracing::trace!("Releasing database connection back to pool");
177

            
178
    result
179
}
180

            
181
// TODO: ensure proper db_connection hygiene for functions that invoke
182
//       this function
183

            
184
async fn session_handler_with_db_connection(
185
    mut stream: ServerToClientMessageStream,
186
    session_id: SessionId,
187
    unix_user: &UnixUser,
188
    db_connection: &mut MySqlConnection,
189
    db_is_mariadb: bool,
190
    group_denylist: &GroupDenylist,
191
    request_counter: Arc<AtomicU64>,
192
) -> anyhow::Result<()> {
193
    stream.send(Response::Ready).await?;
194
    loop {
195
        // TODO: better error handling
196
        // TODO: timeout for receiving requests
197
        // TODO: cancel on request by supervisor
198
        let request = match stream.next().await {
199
            Some(Ok(request)) => request,
200
            Some(Err(e)) => return Err(e.into()),
201
            None => {
202
                tracing::warn!("Client disconnected without sending an exit message");
203
                break;
204
            }
205
        };
206

            
207
        request_counter.fetch_add(1, Ordering::Relaxed);
208

            
209
        let request_span = tracing::info_span!("request", command = request.command_name());
210

            
211
        if !handle_request(
212
            request,
213
            session_id,
214
            unix_user,
215
            db_connection,
216
            db_is_mariadb,
217
            group_denylist,
218
            &mut stream,
219
        )
220
        .instrument(request_span)
221
        .await?
222
        {
223
            break;
224
        }
225
    }
226

            
227
    Ok(())
228
}
229

            
230
/// Handle a single request from a client.
231
///
232
/// If the function returns `true`, the session should continue.
233
async fn handle_request(
234
    request: Request,
235
    session_id: SessionId,
236
    unix_user: &UnixUser,
237
    db_connection: &mut MySqlConnection,
238
    db_is_mariadb: bool,
239
    group_denylist: &GroupDenylist,
240
    stream: &mut ServerToClientMessageStream,
241
) -> anyhow::Result<bool> {
242
    match &request {
243
        Request::Exit => tracing::debug!("Request: exit"),
244
        Request::PasswdUser((db_user, password)) => tracing::debug!(
245
            "Request:\n{}",
246
            serde_json::to_string_pretty(&Request::PasswdUser((
247
                db_user.to_owned(),
248
                match password {
249
                    PasswordSource::Explicit(_) =>
250
                        PasswordSource::Explicit("<REDACTED>".to_string()),
251
                    PasswordSource::Generate => PasswordSource::Generate,
252
                    PasswordSource::Clear => PasswordSource::Clear,
253
                }
254
            )))?
255
        ),
256
        request => tracing::debug!("Request:\n{}", serde_json::to_string_pretty(request)?),
257
    }
258

            
259
    let affected_dbs = request.affected_databases();
260
    if !affected_dbs.is_empty() {
261
        tracing::trace!(
262
            "Affected databases: {}",
263
            affected_dbs.into_iter().map(|db| db.to_string()).join(", ")
264
        );
265
    }
266

            
267
    let affected_users = request.affected_users();
268
    if !affected_users.is_empty() {
269
        tracing::trace!(
270
            "Affected users: {}",
271
            affected_users.into_iter().map(|u| u.to_string()).join(", "),
272
        );
273
    }
274

            
275
    let response = match request {
276
        Request::CheckAuthorization(ref dbs_or_users) => {
277
            let result = check_authorization(dbs_or_users, unix_user, group_denylist).await;
278
            Response::CheckAuthorization(result)
279
        }
280
        Request::ListValidNamePrefixes => {
281
            let mut result = Vec::with_capacity(unix_user.groups.len() + 1);
282
            result.push(unix_user.username.clone());
283

            
284
            for group in get_user_filtered_groups(unix_user, group_denylist) {
285
                result.push(group.clone());
286
            }
287

            
288
            Response::ListValidNamePrefixes(result)
289
        }
290
        Request::CompleteDatabaseName(ref partial_database_name) => {
291
            // TODO: more correct validation here
292
            if partial_database_name
293
                .chars()
294
                .all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-')
295
            {
296
                let result = complete_database_name(
297
                    partial_database_name,
298
                    unix_user,
299
                    db_connection,
300
                    db_is_mariadb,
301
                    group_denylist,
302
                )
303
                .await;
304
                Response::CompleteDatabaseName(result)
305
            } else {
306
                Response::CompleteDatabaseName(vec![])
307
            }
308
        }
309
        Request::CompleteUserName(ref partial_user_name) => {
310
            // TODO: more correct validation here
311
            if partial_user_name
312
                .chars()
313
                .all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-')
314
            {
315
                let result = complete_user_name(
316
                    partial_user_name,
317
                    unix_user,
318
                    db_connection,
319
                    db_is_mariadb,
320
                    group_denylist,
321
                )
322
                .await;
323
                Response::CompleteUserName(result)
324
            } else {
325
                Response::CompleteUserName(vec![])
326
            }
327
        }
328
        Request::CreateDatabases(ref databases_names) => {
329
            let result = create_databases(
330
                databases_names,
331
                unix_user,
332
                db_connection,
333
                db_is_mariadb,
334
                group_denylist,
335
            )
336
            .await;
337
            Response::CreateDatabases(result)
338
        }
339
        Request::DropDatabases(ref databases_names) => {
340
            let result = drop_databases(
341
                databases_names,
342
                unix_user,
343
                db_connection,
344
                db_is_mariadb,
345
                group_denylist,
346
            )
347
            .await;
348
            Response::DropDatabases(result)
349
        }
350
        Request::ListDatabases(ref request) => {
351
            if let Some(database_names) = &request.names {
352
                let result = list_databases(
353
                    database_names,
354
                    unix_user,
355
                    db_connection,
356
                    db_is_mariadb,
357
                    group_denylist,
358
                    request.include_all_tables_and_users,
359
                )
360
                .await;
361
                Response::ListDatabases(result)
362
            } else {
363
                let result = list_all_databases_for_user(
364
                    unix_user,
365
                    db_connection,
366
                    db_is_mariadb,
367
                    group_denylist,
368
                    request.include_all_tables_and_users,
369
                )
370
                .await;
371
                Response::ListAllDatabases(result)
372
            }
373
        }
374
        Request::ListPrivileges(ref database_names) => {
375
            if let Some(database_names) = database_names {
376
                let privilege_data = get_databases_privilege_data(
377
                    database_names,
378
                    unix_user,
379
                    db_connection,
380
                    db_is_mariadb,
381
                    group_denylist,
382
                )
383
                .await;
384
                Response::ListPrivileges(privilege_data)
385
            } else {
386
                let privilege_data = get_all_database_privileges(
387
                    unix_user,
388
                    db_connection,
389
                    db_is_mariadb,
390
                    group_denylist,
391
                )
392
                .await;
393
                Response::ListAllPrivileges(privilege_data)
394
            }
395
        }
396
        Request::ModifyPrivileges(ref database_privilege_diffs) => {
397
            let result = apply_privilege_diffs(
398
                database_privilege_diffs,
399
                unix_user,
400
                db_connection,
401
                db_is_mariadb,
402
                group_denylist,
403
            )
404
            .await;
405
            Response::ModifyPrivileges(result)
406
        }
407
        Request::CreateUsers(ref db_users) => {
408
            let result = create_database_users(
409
                db_users,
410
                unix_user,
411
                db_connection,
412
                db_is_mariadb,
413
                group_denylist,
414
            )
415
            .await;
416
            Response::CreateUsers(result)
417
        }
418
        Request::DropUsers(ref db_users) => {
419
            let result = drop_database_users(
420
                db_users,
421
                unix_user,
422
                db_connection,
423
                db_is_mariadb,
424
                group_denylist,
425
            )
426
            .await;
427
            Response::DropUsers(result)
428
        }
429
        Request::PasswdUser((ref db_user, ref password)) => {
430
            let result = set_password_for_database_user(
431
                db_user,
432
                password,
433
                unix_user,
434
                db_connection,
435
                db_is_mariadb,
436
                group_denylist,
437
            )
438
            .await;
439
            Response::SetUserPassword(result)
440
        }
441
        Request::ListUsers(ref request) => {
442
            if let Some(db_users) = &request.names {
443
                let result = list_database_users(
444
                    db_users,
445
                    unix_user,
446
                    db_connection,
447
                    db_is_mariadb,
448
                    group_denylist,
449
                    request.include_all_databases,
450
                )
451
                .await;
452
                Response::ListUsers(result)
453
            } else {
454
                let result = list_all_database_users_for_unix_user(
455
                    unix_user,
456
                    db_connection,
457
                    db_is_mariadb,
458
                    group_denylist,
459
                    request.include_all_databases,
460
                )
461
                .await;
462
                Response::ListAllUsers(result)
463
            }
464
        }
465
        Request::LockUsers(ref db_users) => {
466
            let result = lock_database_users(
467
                db_users,
468
                unix_user,
469
                db_connection,
470
                db_is_mariadb,
471
                group_denylist,
472
            )
473
            .await;
474
            Response::LockUsers(result)
475
        }
476
        Request::UnlockUsers(ref db_users) => {
477
            let result = unlock_database_users(
478
                db_users,
479
                unix_user,
480
                db_connection,
481
                db_is_mariadb,
482
                group_denylist,
483
            )
484
            .await;
485
            Response::UnlockUsers(result)
486
        }
487
        Request::Exit => {
488
            return Ok(false);
489
        }
490
    };
491

            
492
    let response_to_display = match &response {
493
        Response::SetUserPassword(Err(SetPasswordError::MySqlError(_))) => {
494
            &Response::SetUserPassword(Err(SetPasswordError::MySqlError("<REDACTED>".to_string())))
495
        }
496
        Response::SetUserPassword(Ok(Some(_))) => {
497
            &Response::SetUserPassword(Ok(Some("<REDACTED>".to_string())))
498
        }
499
        response => response,
500
    };
501
    tracing::debug!(
502
        "Response:\n{}",
503
        serde_json::to_string_pretty(&response_to_display)?
504
    );
505

            
506
    log_request(session_id, unix_user, &request, &response);
507

            
508
    stream.send(response).await?;
509
    stream.flush().await?;
510
    tracing::trace!("Successfully processed request");
511

            
512
    Ok(true)
513
}
514

            
515
/// Log a summary of the request and its result.
516
fn log_request(
517
    session_id: SessionId,
518
    unix_user: &UnixUser,
519
    request: &Request,
520
    response: &Response,
521
) {
522
    tracing::info!(
523
        "[{}|session:{}|user:{unix_user}] {}",
524
        response.ok_status(),
525
        session_id.inner(),
526
        request.log_summary(),
527
    );
528
}