1
// TODO: fix comment
2
//! Database privilege operations
3
//!
4
//! This module contains functions for querying, modifying,
5
//! displaying and comparing database privileges.
6
//!
7
//! A lot of the complexity comes from two core components:
8
//!
9
//! - The privilege editor that needs to be able to print
10
//!   an editable table of privileges and reparse the content
11
//!   after the user has made manual changes.
12
//!
13
//! - The comparison functionality that tells the user what
14
//!   changes will be made when applying a set of changes
15
//!   to the list of database privileges.
16

            
17
use std::collections::{BTreeMap, BTreeSet};
18

            
19
use indoc::indoc;
20
use itertools::Itertools;
21
use sqlx::{AssertSqlSafe, MySqlConnection, mysql::MySqlRow, prelude::*};
22

            
23
use crate::{
24
    core::{
25
        common::{UnixUser, rev_yn, yn},
26
        database_privileges::{
27
            DATABASE_PRIVILEGE_FIELDS, DatabasePrivilegeChange, DatabasePrivilegeRow,
28
            DatabasePrivilegesDiff,
29
        },
30
        protocol::{
31
            DiffDoesNotApplyError, ListAllPrivilegesError, ListAllPrivilegesResponse,
32
            ListPrivilegesError, ListPrivilegesResponse, ModifyDatabasePrivilegesError,
33
            ModifyPrivilegesResponse,
34
            request_validation::{GroupDenylist, validate_db_or_user_request},
35
        },
36
        types::{DbOrUser, MySQLDatabase, MySQLUser},
37
    },
38
    server::{
39
        common::{create_user_group_matching_regex, try_get_with_binary_fallback},
40
        sql::{
41
            database_operations::unsafe_database_exists, quote_identifier,
42
            user_operations::unsafe_user_exists,
43
        },
44
    },
45
};
46

            
47
// TODO: get by name instead of row tuple position
48

            
49
#[inline]
50
fn get_mysql_row_priv_field(row: &MySqlRow, position: usize) -> Result<bool, sqlx::Error> {
51
    let field = DATABASE_PRIVILEGE_FIELDS[position];
52
    let value = row.try_get(position)?;
53
    if let Some(val) = rev_yn(value) {
54
        Ok(val)
55
    } else {
56
        tracing::warn!(r#"Invalid value for privilege "{}": '{}'"#, field, value);
57
        Ok(false)
58
    }
59
}
60

            
61
impl FromRow<'_, MySqlRow> for DatabasePrivilegeRow {
62
    fn from_row(row: &MySqlRow) -> Result<Self, sqlx::Error> {
63
        Ok(Self {
64
            db: try_get_with_binary_fallback(row, "Db")?.into(),
65
            user: try_get_with_binary_fallback(row, "User")?.into(),
66
            select_priv: get_mysql_row_priv_field(row, 2)?,
67
            insert_priv: get_mysql_row_priv_field(row, 3)?,
68
            update_priv: get_mysql_row_priv_field(row, 4)?,
69
            delete_priv: get_mysql_row_priv_field(row, 5)?,
70
            create_priv: get_mysql_row_priv_field(row, 6)?,
71
            drop_priv: get_mysql_row_priv_field(row, 7)?,
72
            alter_priv: get_mysql_row_priv_field(row, 8)?,
73
            index_priv: get_mysql_row_priv_field(row, 9)?,
74
            create_tmp_table_priv: get_mysql_row_priv_field(row, 10)?,
75
            lock_tables_priv: get_mysql_row_priv_field(row, 11)?,
76
            references_priv: get_mysql_row_priv_field(row, 12)?,
77
            create_view_priv: get_mysql_row_priv_field(row, 13)?,
78
            show_view_priv: get_mysql_row_priv_field(row, 14)?,
79
            trigger_priv: get_mysql_row_priv_field(row, 15)?,
80
        })
81
    }
82
}
83

            
84
// NOTE: this function is unsafe because it does no input validation.
85
/// Get all users + privileges for a single database.
86
async fn unsafe_get_database_privileges(
87
    database_name: &str,
88
    connection: &mut MySqlConnection,
89
) -> Result<Vec<DatabasePrivilegeRow>, sqlx::Error> {
90
    let statement = AssertSqlSafe(format!(
91
        "SELECT {} FROM `db` WHERE `Db` = ?",
92
        DATABASE_PRIVILEGE_FIELDS
93
            .iter()
94
            .map(|field| quote_identifier(field))
95
            .join(","),
96
    ));
97
    let result = sqlx::query_as::<_, DatabasePrivilegeRow>(statement)
98
        .bind(database_name)
99
        .fetch_all(connection)
100
        .await;
101

            
102
    if let Err(e) = &result {
103
        tracing::error!(
104
            "Failed to get database privileges for '{}': {}",
105
            &database_name,
106
            e
107
        );
108
    }
109

            
110
    result
111
}
112

            
113
// NOTE: this function is unsafe because it does no input validation.
114
/// Get all users + privileges for a single database-user pair.
115
pub async fn unsafe_get_database_privileges_for_db_user_pair(
116
    database_name: &MySQLDatabase,
117
    user_name: &MySQLUser,
118
    connection: &mut MySqlConnection,
119
) -> Result<Option<DatabasePrivilegeRow>, sqlx::Error> {
120
    let statement = AssertSqlSafe(format!(
121
        "SELECT {} FROM `db` WHERE `Db` = ? AND `User` = ? AND `Host` = '%'",
122
        DATABASE_PRIVILEGE_FIELDS
123
            .iter()
124
            .map(|field| quote_identifier(field))
125
            .join(","),
126
    ));
127
    let result = sqlx::query_as::<_, DatabasePrivilegeRow>(statement)
128
        .bind(database_name.as_str())
129
        .bind(user_name.as_str())
130
        .fetch_optional(connection)
131
        .await;
132

            
133
    if let Err(e) = &result {
134
        tracing::error!(
135
            "Failed to get database privileges for '{}.{}': {}",
136
            &database_name,
137
            &user_name,
138
            e
139
        );
140
    }
141

            
142
    result
143
}
144

            
145
pub async fn get_databases_privilege_data(
146
    database_names: &[MySQLDatabase],
147
    unix_user: &UnixUser,
148
    connection: &mut MySqlConnection,
149
    _db_is_mariadb: bool,
150
    group_denylist: &GroupDenylist,
151
) -> ListPrivilegesResponse {
152
    let mut results = BTreeMap::new();
153

            
154
    for database_name in database_names.iter().cloned() {
155
        if let Err(err) = validate_db_or_user_request(
156
            &DbOrUser::Database(database_name.to_owned()),
157
            unix_user,
158
            group_denylist,
159
        )
160
        .map_err(ListPrivilegesError::ValidationError)
161
        {
162
            results.insert(database_name, Err(err));
163
            continue;
164
        }
165

            
166
        match unsafe_database_exists(&database_name, connection).await {
167
            Ok(false) => {
168
                results.insert(
169
                    database_name.to_owned(),
170
                    Err(ListPrivilegesError::DatabaseDoesNotExist),
171
                );
172
                continue;
173
            }
174
            Err(e) => {
175
                results.insert(
176
                    database_name.to_owned(),
177
                    Err(ListPrivilegesError::MySqlError(e.to_string())),
178
                );
179
                continue;
180
            }
181
            Ok(true) => {}
182
        }
183

            
184
        let result = unsafe_get_database_privileges(&database_name, connection)
185
            .await
186
            .map_err(|e| ListPrivilegesError::MySqlError(e.to_string()));
187

            
188
        results.insert(database_name.to_owned(), result);
189
    }
190

            
191
    debug_assert!(database_names.len() == results.len());
192

            
193
    results
194
}
195

            
196
/// TODO: make this constant
197
fn get_all_db_privs_query() -> AssertSqlSafe<String> {
198
    AssertSqlSafe(format!(
199
        indoc! {r"
200
            SELECT {} FROM `db` WHERE `db` IN
201
            (SELECT DISTINCT CAST(`SCHEMA_NAME` AS CHAR(64)) AS `database`
202
              FROM `information_schema`.`SCHEMATA`
203
              WHERE `SCHEMA_NAME` NOT IN ('information_schema', 'performance_schema', 'mysql', 'sys')
204
                AND `SCHEMA_NAME` REGEXP ?)
205
        "},
206
        DATABASE_PRIVILEGE_FIELDS
207
            .iter()
208
            .map(|field| quote_identifier(field))
209
            .join(","),
210
    ))
211
}
212

            
213
/// Get all database + user + privileges pairs that are owned by the current user.
214
pub async fn get_all_database_privileges(
215
    unix_user: &UnixUser,
216
    connection: &mut MySqlConnection,
217
    _db_is_mariadb: bool,
218
    group_denylist: &GroupDenylist,
219
) -> ListAllPrivilegesResponse {
220
    let result = sqlx::query_as::<_, DatabasePrivilegeRow>(get_all_db_privs_query())
221
        .bind(create_user_group_matching_regex(unix_user, group_denylist))
222
        .fetch_all(connection)
223
        .await
224
        .map_err(|e| ListAllPrivilegesError::MySqlError(e.to_string()));
225

            
226
    if let Err(e) = &result {
227
        tracing::error!("Failed to get all database privileges: {:?}", e);
228
    }
229

            
230
    result
231
}
232

            
233
// TODO: make these queries constant strings.
234
async fn unsafe_apply_privilege_diff(
235
    database_privilege_diff: &DatabasePrivilegesDiff,
236
    connection: &mut MySqlConnection,
237
) -> Result<(), sqlx::Error> {
238
    let result = match database_privilege_diff {
239
        DatabasePrivilegesDiff::New(p) => {
240
            let tables = DATABASE_PRIVILEGE_FIELDS
241
                .iter()
242
                .chain(&["Host"])
243
                .map(|field| quote_identifier(field))
244
                .join(",");
245

            
246
            let question_marks =
247
                std::iter::repeat_n("?", DATABASE_PRIVILEGE_FIELDS.len() + 1).join(",");
248

            
249
            let statement = AssertSqlSafe(format!(
250
                "INSERT INTO `db` ({tables}) VALUES ({question_marks})"
251
            ));
252
            sqlx::query(statement)
253
                .bind(p.db.to_string())
254
                .bind(p.user.to_string())
255
                .bind(yn(p.select_priv))
256
                .bind(yn(p.insert_priv))
257
                .bind(yn(p.update_priv))
258
                .bind(yn(p.delete_priv))
259
                .bind(yn(p.create_priv))
260
                .bind(yn(p.drop_priv))
261
                .bind(yn(p.alter_priv))
262
                .bind(yn(p.index_priv))
263
                .bind(yn(p.create_tmp_table_priv))
264
                .bind(yn(p.lock_tables_priv))
265
                .bind(yn(p.references_priv))
266
                .bind(yn(p.create_view_priv))
267
                .bind(yn(p.show_view_priv))
268
                .bind(yn(p.trigger_priv))
269
                .bind("%")
270
                .execute(connection)
271
                .await
272
                .map(|_| ())
273
        }
274
        DatabasePrivilegesDiff::Modified(p) => {
275
            let changes = DATABASE_PRIVILEGE_FIELDS
276
                .iter()
277
                .skip(2) // Skip Db and User fields
278
                .map(|field| {
279
                    format!(
280
                        "{} = COALESCE(?, {})",
281
                        quote_identifier(field),
282
                        quote_identifier(field)
283
                    )
284
                })
285
                .join(",");
286

            
287
            fn change_to_yn(change: DatabasePrivilegeChange) -> &'static str {
288
                match change {
289
                    DatabasePrivilegeChange::YesToNo => "N",
290
                    DatabasePrivilegeChange::NoToYes => "Y",
291
                }
292
            }
293

            
294
            let statement = AssertSqlSafe(format!(
295
                "UPDATE `db` SET {changes} WHERE `Db` = ? AND `User` = ? AND `Host` = ?"
296
            ));
297
            sqlx::query(statement)
298
                .bind(p.select_priv.map(change_to_yn))
299
                .bind(p.insert_priv.map(change_to_yn))
300
                .bind(p.update_priv.map(change_to_yn))
301
                .bind(p.delete_priv.map(change_to_yn))
302
                .bind(p.create_priv.map(change_to_yn))
303
                .bind(p.drop_priv.map(change_to_yn))
304
                .bind(p.alter_priv.map(change_to_yn))
305
                .bind(p.index_priv.map(change_to_yn))
306
                .bind(p.create_tmp_table_priv.map(change_to_yn))
307
                .bind(p.lock_tables_priv.map(change_to_yn))
308
                .bind(p.references_priv.map(change_to_yn))
309
                .bind(p.create_view_priv.map(change_to_yn))
310
                .bind(p.show_view_priv.map(change_to_yn))
311
                .bind(p.trigger_priv.map(change_to_yn))
312
                .bind(p.db.to_string())
313
                .bind(p.user.to_string())
314
                .bind("%")
315
                .execute(connection)
316
                .await
317
                .map(|_| ())
318
        }
319
        DatabasePrivilegesDiff::Deleted(p) => {
320
            sqlx::query("DELETE FROM `db` WHERE `Db` = ? AND `User` = ? AND `Host` = ?")
321
                .bind(p.db.to_string())
322
                .bind(p.user.to_string())
323
                .bind("%")
324
                .execute(connection)
325
                .await
326
                .map(|_| ())
327
        }
328
        DatabasePrivilegesDiff::Noop { .. } => Ok(()),
329
    };
330

            
331
    if let Err(e) = &result {
332
        tracing::error!("Failed to apply database privilege diff: {}", e);
333
    }
334

            
335
    result
336
}
337

            
338
async fn validate_diff(
339
    diff: &DatabasePrivilegesDiff,
340
    connection: &mut MySqlConnection,
341
) -> Result<(), ModifyDatabasePrivilegesError> {
342
    let privilege_row = unsafe_get_database_privileges_for_db_user_pair(
343
        diff.get_database_name(),
344
        diff.get_user_name(),
345
        connection,
346
    )
347
    .await;
348

            
349
    let privilege_row = match privilege_row {
350
        Ok(privilege_row) => privilege_row,
351
        Err(e) => return Err(ModifyDatabasePrivilegesError::MySqlError(e.to_string())),
352
    };
353

            
354
    match diff {
355
        DatabasePrivilegesDiff::New(_) => {
356
            if privilege_row.is_some() {
357
                Err(ModifyDatabasePrivilegesError::DiffDoesNotApply(Box::new(
358
                    DiffDoesNotApplyError::RowAlreadyExists(
359
                        diff.get_database_name().to_owned(),
360
                        diff.get_user_name().to_owned(),
361
                    ),
362
                )))
363
            } else {
364
                Ok(())
365
            }
366
        }
367
        DatabasePrivilegesDiff::Modified(_) if privilege_row.is_none() => {
368
            Err(ModifyDatabasePrivilegesError::DiffDoesNotApply(Box::new(
369
                DiffDoesNotApplyError::RowDoesNotExist(
370
                    diff.get_database_name().to_owned(),
371
                    diff.get_user_name().to_owned(),
372
                ),
373
            )))
374
        }
375
        DatabasePrivilegesDiff::Modified(row_diff) => {
376
            let row = privilege_row.unwrap();
377

            
378
            let error_exists = DATABASE_PRIVILEGE_FIELDS
379
                .iter()
380
                .skip(2) // Skip Db and User fields
381
                .any(
382
                    |field| match row_diff.get_privilege_change_by_name(field).unwrap() {
383
                        Some(DatabasePrivilegeChange::YesToNo) => {
384
                            !row.get_privilege_by_name(field).unwrap()
385
                        }
386
                        Some(DatabasePrivilegeChange::NoToYes) => {
387
                            row.get_privilege_by_name(field).unwrap()
388
                        }
389
                        None => false,
390
                    },
391
                );
392

            
393
            if error_exists {
394
                Err(ModifyDatabasePrivilegesError::DiffDoesNotApply(Box::new(
395
                    DiffDoesNotApplyError::RowPrivilegeChangeDoesNotApply(row_diff.to_owned(), row),
396
                )))
397
            } else {
398
                Ok(())
399
            }
400
        }
401
        DatabasePrivilegesDiff::Deleted(_) => {
402
            if privilege_row.is_none() {
403
                Err(ModifyDatabasePrivilegesError::DiffDoesNotApply(Box::new(
404
                    DiffDoesNotApplyError::RowDoesNotExist(
405
                        diff.get_database_name().to_owned(),
406
                        diff.get_user_name().to_owned(),
407
                    ),
408
                )))
409
            } else {
410
                Ok(())
411
            }
412
        }
413
        DatabasePrivilegesDiff::Noop { .. } => {
414
            tracing::warn!(
415
                "Server got sent a noop database privilege diff to validate, is the client buggy?"
416
            );
417
            Ok(())
418
        }
419
    }
420
}
421

            
422
/// Uses the result of [`diff_privileges`] to modify privileges in the database.
423
pub async fn apply_privilege_diffs(
424
    database_privilege_diffs: &BTreeSet<DatabasePrivilegesDiff>,
425
    unix_user: &UnixUser,
426
    connection: &mut MySqlConnection,
427
    _db_is_mariadb: bool,
428
    group_denylist: &GroupDenylist,
429
) -> ModifyPrivilegesResponse {
430
    let mut results: BTreeMap<(MySQLDatabase, MySQLUser), _> = BTreeMap::new();
431

            
432
    for diff in database_privilege_diffs {
433
        let key = (
434
            diff.get_database_name().to_owned(),
435
            diff.get_user_name().to_owned(),
436
        );
437
        if let Err(err) = validate_db_or_user_request(
438
            &DbOrUser::Database(diff.get_database_name().to_owned()),
439
            unix_user,
440
            group_denylist,
441
        )
442
        .map_err(ModifyDatabasePrivilegesError::UserValidationError)
443
        {
444
            results.insert(key, Err(err));
445
            continue;
446
        }
447

            
448
        if let Err(err) = validate_db_or_user_request(
449
            &DbOrUser::User(diff.get_user_name().to_owned()),
450
            unix_user,
451
            group_denylist,
452
        )
453
        .map_err(ModifyDatabasePrivilegesError::UserValidationError)
454
        {
455
            results.insert(key, Err(err));
456
            continue;
457
        }
458

            
459
        match unsafe_database_exists(diff.get_database_name(), connection).await {
460
            Ok(false) => {
461
                results.insert(
462
                    key,
463
                    Err(ModifyDatabasePrivilegesError::DatabaseDoesNotExist),
464
                );
465
                continue;
466
            }
467
            Err(e) => {
468
                results.insert(
469
                    key,
470
                    Err(ModifyDatabasePrivilegesError::MySqlError(e.to_string())),
471
                );
472
                continue;
473
            }
474
            Ok(true) => {}
475
        }
476

            
477
        match unsafe_user_exists(diff.get_user_name(), connection).await {
478
            Ok(false) => {
479
                results.insert(key, Err(ModifyDatabasePrivilegesError::UserDoesNotExist));
480
                continue;
481
            }
482
            Err(e) => {
483
                results.insert(
484
                    key,
485
                    Err(ModifyDatabasePrivilegesError::MySqlError(e.to_string())),
486
                );
487
                continue;
488
            }
489
            Ok(true) => {}
490
        }
491

            
492
        if let Err(err) = validate_diff(diff, connection).await {
493
            results.insert(key, Err(err));
494
            continue;
495
        }
496

            
497
        let result = unsafe_apply_privilege_diff(diff, connection)
498
            .await
499
            .map_err(|e| ModifyDatabasePrivilegesError::MySqlError(e.to_string()));
500

            
501
        results.insert(key, result);
502
    }
503

            
504
    if let Err(err) = connection.execute("FLUSH PRIVILEGES").await {
505
        tracing::error!("Failed to flush privileges: {}", err);
506
    }
507

            
508
    results
509
        .into_iter()
510
        .map(|((k1, k2), v)| (k1, (k2, v)))
511
        .into_group_map()
512
        .into_iter()
513
        .map(|(k1, pairs)| {
514
            let inner = pairs.into_iter().collect::<BTreeMap<_, _>>();
515
            (k1, inner)
516
        })
517
        .collect()
518
}