1
use indoc::formatdoc;
2
use itertools::Itertools;
3
use rand::distr::{Alphanumeric, SampleString};
4
use sqlx::AssertSqlSafe;
5
use std::collections::BTreeMap;
6

            
7
use serde::{Deserialize, Serialize};
8

            
9
use sqlx::MySqlConnection;
10
use sqlx::prelude::*;
11

            
12
use crate::core::protocol::request_validation::GroupDenylist;
13
use crate::core::protocol::request_validation::validate_db_or_user_request;
14
use crate::core::types::DbOrUser;
15
use crate::{
16
    core::{
17
        common::UnixUser,
18
        database_privileges::DATABASE_PRIVILEGE_FIELDS,
19
        protocol::{
20
            CreateUserError, CreateUsersResponse, DropUserError, DropUsersResponse,
21
            ListAllUsersError, ListAllUsersResponse, ListUsersError, ListUsersResponse,
22
            LockUserError, LockUsersResponse, PasswordSource, SetPasswordError,
23
            SetUserPasswordResponse, UnlockUserError, UnlockUsersResponse,
24
        },
25
        types::MySQLUser,
26
    },
27
    server::{
28
        common::{create_user_group_matching_regex, try_get_with_binary_fallback},
29
        sql::quote_literal,
30
    },
31
};
32

            
33
const GENERATED_PASSWORD_LENGTH: usize = 24;
34
const MAX_SHOW_USER_RELATED_ITEMS: usize = 5;
35

            
36
// NOTE: this function is unsafe because it does no input validation.
37
pub(super) async fn unsafe_user_exists(
38
    db_user: &str,
39
    connection: &mut MySqlConnection,
40
) -> Result<bool, sqlx::Error> {
41
    let result = sqlx::query(
42
        r"
43
          SELECT EXISTS(
44
            SELECT 1
45
            FROM `mysql`.`user`
46
            WHERE `User` = ?
47
              AND `Host` = '%'
48
          )
49
        ",
50
    )
51
    .bind(db_user)
52
    .fetch_one(connection)
53
    .await
54
    .map(|row| row.get::<bool, _>(0));
55

            
56
    if let Err(err) = &result {
57
        tracing::error!("Failed to check if database user exists: {:?}", err);
58
    }
59

            
60
    result
61
}
62

            
63
pub async fn complete_user_name(
64
    user_prefix: &str,
65
    unix_user: &UnixUser,
66
    connection: &mut MySqlConnection,
67
    _db_is_mariadb: bool,
68
    group_denylist: &GroupDenylist,
69
) -> Vec<MySQLUser> {
70
    let result = sqlx::query(
71
        r"
72
          SELECT `User` AS `user`
73
          FROM `mysql`.`user`
74
          WHERE `User` REGEXP ?
75
            AND `User` LIKE ?
76
            AND `Host` = '%'
77
        ",
78
    )
79
    .bind(create_user_group_matching_regex(unix_user, group_denylist))
80
    .bind(format!("{user_prefix}%"))
81
    .fetch_all(connection)
82
    .await;
83

            
84
    match result {
85
        Ok(rows) => rows
86
            .into_iter()
87
            .filter_map(|row| {
88
                let user: String = try_get_with_binary_fallback(&row, "user").ok()?;
89
                Some(user.into())
90
            })
91
            .collect(),
92
        Err(err) => {
93
            tracing::error!(
94
                "Failed to complete user name for prefix '{}' and user '{}': {:?}",
95
                user_prefix,
96
                unix_user.username,
97
                err
98
            );
99
            vec![]
100
        }
101
    }
102
}
103

            
104
pub async fn create_database_users(
105
    db_users: &[MySQLUser],
106
    unix_user: &UnixUser,
107
    connection: &mut MySqlConnection,
108
    _db_is_mariadb: bool,
109
    group_denylist: &GroupDenylist,
110
) -> CreateUsersResponse {
111
    let mut results = BTreeMap::new();
112

            
113
    for db_user in db_users.iter().cloned() {
114
        if let Err(err) =
115
            validate_db_or_user_request(&DbOrUser::User(db_user.clone()), unix_user, group_denylist)
116
                .map_err(CreateUserError::ValidationError)
117
        {
118
            results.insert(db_user, Err(err));
119
            continue;
120
        }
121

            
122
        match unsafe_user_exists(&db_user, &mut *connection).await {
123
            Ok(true) => {
124
                results.insert(db_user, Err(CreateUserError::UserAlreadyExists));
125
                continue;
126
            }
127
            Err(err) => {
128
                results.insert(db_user, Err(CreateUserError::MySqlError(err.to_string())));
129
                continue;
130
            }
131
            _ => {}
132
        }
133

            
134
        let statement = AssertSqlSafe(format!("CREATE USER {}@'%'", quote_literal(&db_user),));
135
        let result = sqlx::query(statement)
136
            .execute(&mut *connection)
137
            .await
138
            .map(|_| ())
139
            .map_err(|err| CreateUserError::MySqlError(err.to_string()));
140

            
141
        if let Err(err) = &result {
142
            tracing::error!("Failed to create database user '{}': {:?}", &db_user, err);
143
        }
144

            
145
        results.insert(db_user, result);
146
    }
147

            
148
    results
149
}
150

            
151
pub async fn drop_database_users(
152
    db_users: &[MySQLUser],
153
    unix_user: &UnixUser,
154
    connection: &mut MySqlConnection,
155
    _db_is_mariadb: bool,
156
    group_denylist: &GroupDenylist,
157
) -> DropUsersResponse {
158
    let mut results = BTreeMap::new();
159

            
160
    for db_user in db_users.iter().cloned() {
161
        if let Err(err) =
162
            validate_db_or_user_request(&DbOrUser::User(db_user.clone()), unix_user, group_denylist)
163
                .map_err(DropUserError::ValidationError)
164
        {
165
            results.insert(db_user, Err(err));
166
            continue;
167
        }
168

            
169
        match unsafe_user_exists(&db_user, &mut *connection).await {
170
            Ok(false) => {
171
                results.insert(db_user, Err(DropUserError::UserDoesNotExist));
172
                continue;
173
            }
174
            Err(err) => {
175
                results.insert(db_user, Err(DropUserError::MySqlError(err.to_string())));
176
                continue;
177
            }
178
            _ => {}
179
        }
180

            
181
        let statement = AssertSqlSafe(format!("DROP USER {}@'%'", quote_literal(&db_user),));
182
        let result = sqlx::query(statement)
183
            .execute(&mut *connection)
184
            .await
185
            .map(|_| ())
186
            .map_err(|err| DropUserError::MySqlError(err.to_string()));
187

            
188
        if let Err(err) = &result {
189
            tracing::error!("Failed to drop database user '{}': {:?}", &db_user, err);
190
        }
191

            
192
        results.insert(db_user, result);
193
    }
194

            
195
    results
196
}
197

            
198
pub async fn set_password_for_database_user(
199
    db_user: &MySQLUser,
200
    password: &PasswordSource,
201
    unix_user: &UnixUser,
202
    connection: &mut MySqlConnection,
203
    _db_is_mariadb: bool,
204
    group_denylist: &GroupDenylist,
205
) -> SetUserPasswordResponse {
206
    validate_db_or_user_request(&DbOrUser::User(db_user.clone()), unix_user, group_denylist)
207
        .map_err(SetPasswordError::ValidationError)?;
208

            
209
    match unsafe_user_exists(db_user, &mut *connection).await {
210
        Ok(false) => return Err(SetPasswordError::UserDoesNotExist),
211
        Err(err) => return Err(SetPasswordError::MySqlError(err.to_string())),
212
        _ => {}
213
    }
214

            
215
    let generated_password = match password {
216
        PasswordSource::Explicit(_) | PasswordSource::Clear => None,
217
        PasswordSource::Generate => {
218
            Some(Alphanumeric.sample_string(&mut rand::rng(), GENERATED_PASSWORD_LENGTH))
219
        }
220
    };
221
    let password = match (&generated_password, password) {
222
        (Some(generated), _) => generated.as_str(),
223
        (None, PasswordSource::Explicit(password)) => password.as_str(),
224
        (None, PasswordSource::Clear) => "",
225
        (None, PasswordSource::Generate) => {
226
            unreachable!("generated_password is always Some for PasswordSource::Generate")
227
        }
228
    };
229

            
230
    let statement = AssertSqlSafe(format!(
231
        "ALTER USER {}@'%' IDENTIFIED BY {}",
232
        quote_literal(db_user),
233
        quote_literal(password).as_str(),
234
    ));
235
    let result = sqlx::query(statement)
236
        .execute(&mut *connection)
237
        .await
238
        .map(|_| generated_password)
239
        .map_err(|err| SetPasswordError::MySqlError(err.to_string()));
240

            
241
    if result.is_err() {
242
        tracing::error!(
243
            "Failed to set password for database user '{}': <REDACTED>",
244
            &db_user,
245
        );
246
    }
247

            
248
    result
249
}
250

            
251
const DATABASE_USER_LOCK_STATUS_QUERY_MARIADB: &str = r#"
252
    SELECT COALESCE(
253
        JSON_EXTRACT(`mysql`.`global_priv`.`priv`, "$.account_locked"),
254
        'false'
255
    ) != 'false'
256
    FROM `mysql`.`global_priv`
257
    WHERE `User` = ?
258
    AND `Host` = '%'
259
"#;
260

            
261
const DATABASE_USER_LOCK_STATUS_QUERY_MYSQL: &str = r"
262
    SELECT `mysql`.`user`.`account_locked` = 'Y'
263
    FROM `mysql`.`user`
264
    WHERE `User` = ?
265
    AND `Host` = '%'
266
";
267

            
268
// NOTE: this function is unsafe because it does no input validation.
269
async fn database_user_is_locked_unsafe(
270
    db_user: &str,
271
    connection: &mut MySqlConnection,
272
    db_is_mariadb: bool,
273
) -> Result<bool, sqlx::Error> {
274
    let result = sqlx::query(if db_is_mariadb {
275
        DATABASE_USER_LOCK_STATUS_QUERY_MARIADB
276
    } else {
277
        DATABASE_USER_LOCK_STATUS_QUERY_MYSQL
278
    })
279
    .bind(db_user)
280
    .fetch_one(connection)
281
    .await
282
    .map(|row| row.try_get(0))
283
    .and_then(|res| res);
284

            
285
    if let Err(err) = &result {
286
        tracing::error!(
287
            "Failed to check if database user is locked '{}': {:?}",
288
            &db_user,
289
            err
290
        );
291
    }
292

            
293
    result
294
}
295

            
296
pub async fn lock_database_users(
297
    db_users: &[MySQLUser],
298
    unix_user: &UnixUser,
299
    connection: &mut MySqlConnection,
300
    db_is_mariadb: bool,
301
    group_denylist: &GroupDenylist,
302
) -> LockUsersResponse {
303
    let mut results = BTreeMap::new();
304

            
305
    for db_user in db_users.iter().cloned() {
306
        if let Err(err) =
307
            validate_db_or_user_request(&DbOrUser::User(db_user.clone()), unix_user, group_denylist)
308
                .map_err(LockUserError::ValidationError)
309
        {
310
            results.insert(db_user, Err(err));
311
            continue;
312
        }
313

            
314
        match unsafe_user_exists(&db_user, &mut *connection).await {
315
            Ok(true) => {}
316
            Ok(false) => {
317
                results.insert(db_user, Err(LockUserError::UserDoesNotExist));
318
                continue;
319
            }
320
            Err(err) => {
321
                results.insert(db_user, Err(LockUserError::MySqlError(err.to_string())));
322
                continue;
323
            }
324
        }
325

            
326
        match database_user_is_locked_unsafe(&db_user, &mut *connection, db_is_mariadb).await {
327
            Ok(false) => {}
328
            Ok(true) => {
329
                results.insert(db_user, Err(LockUserError::UserIsAlreadyLocked));
330
                continue;
331
            }
332
            Err(err) => {
333
                results.insert(db_user, Err(LockUserError::MySqlError(err.to_string())));
334
                continue;
335
            }
336
        }
337

            
338
        let statement = AssertSqlSafe(format!(
339
            "ALTER USER {}@'%' ACCOUNT LOCK",
340
            quote_literal(&db_user),
341
        ));
342
        let result = sqlx::query(statement)
343
            .execute(&mut *connection)
344
            .await
345
            .map(|_| ())
346
            .map_err(|err| LockUserError::MySqlError(err.to_string()));
347

            
348
        if let Err(err) = &result {
349
            tracing::error!("Failed to lock database user '{}': {:?}", &db_user, err);
350
        }
351

            
352
        results.insert(db_user, result);
353
    }
354

            
355
    results
356
}
357

            
358
pub async fn unlock_database_users(
359
    db_users: &[MySQLUser],
360
    unix_user: &UnixUser,
361
    connection: &mut MySqlConnection,
362
    db_is_mariadb: bool,
363
    group_denylist: &GroupDenylist,
364
) -> UnlockUsersResponse {
365
    let mut results = BTreeMap::new();
366

            
367
    for db_user in db_users.iter().cloned() {
368
        if let Err(err) =
369
            validate_db_or_user_request(&DbOrUser::User(db_user.clone()), unix_user, group_denylist)
370
                .map_err(UnlockUserError::ValidationError)
371
        {
372
            results.insert(db_user, Err(err));
373
            continue;
374
        }
375

            
376
        match unsafe_user_exists(&db_user, &mut *connection).await {
377
            Ok(false) => {
378
                results.insert(db_user, Err(UnlockUserError::UserDoesNotExist));
379
                continue;
380
            }
381
            Err(err) => {
382
                results.insert(db_user, Err(UnlockUserError::MySqlError(err.to_string())));
383
                continue;
384
            }
385
            _ => {}
386
        }
387

            
388
        match database_user_is_locked_unsafe(&db_user, &mut *connection, db_is_mariadb).await {
389
            Ok(false) => {
390
                results.insert(db_user, Err(UnlockUserError::UserIsAlreadyUnlocked));
391
                continue;
392
            }
393
            Err(err) => {
394
                results.insert(db_user, Err(UnlockUserError::MySqlError(err.to_string())));
395
                continue;
396
            }
397
            _ => {}
398
        }
399

            
400
        let statement = AssertSqlSafe(format!(
401
            "ALTER USER {}@'%' ACCOUNT UNLOCK",
402
            quote_literal(&db_user),
403
        ));
404
        let result = sqlx::query(statement)
405
            .execute(&mut *connection)
406
            .await
407
            .map(|_| ())
408
            .map_err(|err| UnlockUserError::MySqlError(err.to_string()));
409

            
410
        if let Err(err) = &result {
411
            tracing::error!("Failed to unlock database user '{}': {:?}", &db_user, err);
412
        }
413

            
414
        results.insert(db_user, result);
415
    }
416

            
417
    results
418
}
419

            
420
/// This struct contains information about a database user.
421
/// This can be extended if we need more information in the future.
422
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
423
pub struct DatabaseUser {
424
    pub user: MySQLUser,
425
    #[serde(skip)]
426
    pub host: String,
427
    pub has_password: bool,
428
    pub is_locked: bool,
429
    pub databases: Vec<String>,
430
    pub database_count: u64,
431
}
432

            
433
impl FromRow<'_, sqlx::mysql::MySqlRow> for DatabaseUser {
434
    fn from_row(row: &sqlx::mysql::MySqlRow) -> Result<Self, sqlx::Error> {
435
        Ok(Self {
436
            user: try_get_with_binary_fallback(row, "User")?.into(),
437
            host: try_get_with_binary_fallback(row, "Host")?,
438
            has_password: row.try_get("has_password")?,
439
            is_locked: row.try_get("account_locked")?,
440
            databases: Vec::new(),
441
            database_count: 0,
442
        })
443
    }
444
}
445

            
446
const DB_USER_SELECT_STATEMENT_MARIADB: &str = r#"
447
SELECT
448
  `user`.`User`,
449
  `user`.`Host`,
450
  `user`.`Password` != '' OR `user`.`authentication_string` != '' AS `has_password`,
451
  COALESCE(
452
    JSON_EXTRACT(`global_priv`.`priv`, "$.account_locked"),
453
    'false'
454
  ) != 'false' AS `account_locked`
455
FROM `user`
456
JOIN `global_priv` ON
457
  `user`.`User` = `global_priv`.`User`
458
  AND `user`.`Host` = `global_priv`.`Host`
459
"#;
460

            
461
const DB_USER_SELECT_STATEMENT_MYSQL: &str = r"
462
SELECT
463
  `user`.`User`,
464
  `user`.`Host`,
465
  `user`.`authentication_string` != '' AS `has_password`,
466
  `user`.`account_locked` = 'Y' AS `account_locked`
467
FROM `user`
468
";
469

            
470
pub async fn list_database_users(
471
    db_users: &[MySQLUser],
472
    unix_user: &UnixUser,
473
    connection: &mut MySqlConnection,
474
    db_is_mariadb: bool,
475
    group_denylist: &GroupDenylist,
476
    include_all_databases: bool,
477
) -> ListUsersResponse {
478
    let mut results = BTreeMap::new();
479

            
480
    for db_user in db_users.iter().cloned() {
481
        if let Err(err) =
482
            validate_db_or_user_request(&DbOrUser::User(db_user.clone()), unix_user, group_denylist)
483
                .map_err(ListUsersError::ValidationError)
484
        {
485
            results.insert(db_user, Err(err));
486
            continue;
487
        }
488

            
489
        let statement = AssertSqlSafe(
490
            if db_is_mariadb {
491
                DB_USER_SELECT_STATEMENT_MARIADB.to_string()
492
            } else {
493
                DB_USER_SELECT_STATEMENT_MYSQL.to_string()
494
            } + "WHERE `mysql`.`user`.`User` = ? AND `mysql`.`user`.`Host` = '%'",
495
        );
496
        let mut result = sqlx::query_as::<_, DatabaseUser>(statement)
497
            .bind(db_user.as_str())
498
            .fetch_optional(&mut *connection)
499
            .await;
500

            
501
        if let Err(err) = &result {
502
            tracing::error!("Failed to list database user '{}': {:?}", &db_user, err);
503
        }
504

            
505
        if let Ok(Some(user)) = result.as_mut()
506
            && let Err(err) = set_databases_where_user_has_privileges(
507
                user,
508
                &mut *connection,
509
                include_all_databases,
510
            )
511
            .await
512
        {
513
            result = Err(err);
514
        }
515

            
516
        match result {
517
            Ok(Some(user)) => results.insert(db_user, Ok(user)),
518
            Ok(None) => results.insert(db_user, Err(ListUsersError::UserDoesNotExist)),
519
            Err(err) => results.insert(db_user, Err(ListUsersError::MySqlError(err.to_string()))),
520
        };
521
    }
522

            
523
    results
524
}
525

            
526
pub async fn list_all_database_users_for_unix_user(
527
    unix_user: &UnixUser,
528
    connection: &mut MySqlConnection,
529
    db_is_mariadb: bool,
530
    group_denylist: &GroupDenylist,
531
    include_all_databases: bool,
532
) -> ListAllUsersResponse {
533
    let statement = AssertSqlSafe(
534
        if db_is_mariadb {
535
            DB_USER_SELECT_STATEMENT_MARIADB.to_string()
536
        } else {
537
            DB_USER_SELECT_STATEMENT_MYSQL.to_string()
538
        } + "WHERE `user`.`User` REGEXP ? AND `user`.`Host` = '%'",
539
    );
540
    let mut result = sqlx::query_as::<_, DatabaseUser>(statement)
541
        .bind(create_user_group_matching_regex(unix_user, group_denylist))
542
        .fetch_all(&mut *connection)
543
        .await
544
        .map_err(|err| ListAllUsersError::MySqlError(err.to_string()));
545

            
546
    if let Err(err) = &result {
547
        tracing::error!("Failed to list all database users: {:?}", err);
548
    }
549

            
550
    if let Ok(users) = result.as_mut() {
551
        for user in users {
552
            if let Err(mysql_error) = set_databases_where_user_has_privileges(
553
                user,
554
                &mut *connection,
555
                include_all_databases,
556
            )
557
            .await
558
            {
559
                return Err(ListAllUsersError::MySqlError(mysql_error.to_string()));
560
            }
561
        }
562
    }
563

            
564
    result
565
}
566

            
567
/// This function sets the `databases` field of the given `DatabaseUser`
568
/// where the user has any privileges.
569
pub async fn set_databases_where_user_has_privileges(
570
    db_user: &mut DatabaseUser,
571
    connection: &mut MySqlConnection,
572
    include_all_databases: bool,
573
) -> Result<(), sqlx::Error> {
574
    let limit_clause = if include_all_databases {
575
        String::new()
576
    } else {
577
        format!(" LIMIT {MAX_SHOW_USER_RELATED_ITEMS}")
578
    };
579

            
580
    let statement = AssertSqlSafe(formatdoc!(
581
        r"
582
            SELECT
583
                `Db` AS `database`,
584
                CAST(COUNT(*) OVER() AS UNSIGNED) AS `database_count`
585
            FROM `db`
586
            WHERE `User` = ?  AND `Host` = '%' AND ({})
587
            ORDER BY `Db`{limit_clause}
588
        ",
589
        DATABASE_PRIVILEGE_FIELDS
590
            .iter()
591
            .map(|field| format!("`{field}` = 'Y'"))
592
            .join(" OR "),
593
    ));
594
    let database_list = sqlx::query(statement)
595
        .bind(db_user.user.as_str())
596
        .fetch_all(&mut *connection)
597
        .await;
598

            
599
    if let Err(err) = &database_list {
600
        tracing::error!(
601
            "Failed to list databases for user '{}': {:?}",
602
            &db_user.user,
603
            err
604
        );
605
    }
606

            
607
    let rows = database_list?;
608

            
609
    db_user.database_count = rows
610
        .first()
611
        .map(|row| row.try_get::<u64, _>("database_count"))
612
        .transpose()?
613
        .unwrap_or(0);
614

            
615
    db_user.databases = rows
616
        .into_iter()
617
        .map(|row| try_get_with_binary_fallback(&row, "database"))
618
        .collect::<Result<Vec<String>, sqlx::Error>>()?;
619

            
620
    Ok(())
621
}