1
use std::collections::BTreeMap;
2

            
3
use sqlx::AssertSqlSafe;
4
use sqlx::MySqlConnection;
5
use sqlx::prelude::*;
6

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

            
9
use crate::core::protocol::CompleteDatabaseNameResponse;
10
use crate::core::protocol::request_validation::GroupDenylist;
11
use crate::core::protocol::request_validation::validate_db_or_user_request;
12
use crate::core::types::DbOrUser;
13
use crate::core::types::MySQLDatabase;
14
use crate::core::types::MySQLUser;
15
use crate::{
16
    core::{
17
        common::UnixUser,
18
        protocol::{
19
            CreateDatabaseError, CreateDatabasesResponse, DropDatabaseError, DropDatabasesResponse,
20
            ListAllDatabasesError, ListAllDatabasesResponse, ListDatabasesError,
21
            ListDatabasesResponse,
22
        },
23
    },
24
    server::{common::create_user_group_matching_regex, sql::quote_identifier},
25
};
26

            
27
const MAX_SHOW_DB_RELATED_ITEMS: usize = 5;
28

            
29
// NOTE: this function is unsafe because it does no input validation.
30
pub(super) async fn unsafe_database_exists(
31
    database_name: &str,
32
    connection: &mut MySqlConnection,
33
) -> Result<bool, sqlx::Error> {
34
    let result =
35
        sqlx::query("SELECT SCHEMA_NAME FROM information_schema.SCHEMATA WHERE SCHEMA_NAME = ?")
36
            .bind(database_name)
37
            .fetch_optional(connection)
38
            .await;
39

            
40
    if let Err(err) = &result {
41
        tracing::error!(
42
            "Failed to check if database '{}' exists: {:?}",
43
            &database_name,
44
            err
45
        );
46
    }
47

            
48
    Ok(result?.is_some())
49
}
50

            
51
pub async fn complete_database_name(
52
    database_prefix: &str,
53
    unix_user: &UnixUser,
54
    connection: &mut MySqlConnection,
55
    _db_is_mariadb: bool,
56
    group_denylist: &GroupDenylist,
57
) -> CompleteDatabaseNameResponse {
58
    let result = sqlx::query(
59
        r"
60
          SELECT CAST(`SCHEMA_NAME` AS CHAR(64)) AS `database`
61
          FROM `information_schema`.`SCHEMATA`
62
          WHERE `SCHEMA_NAME` NOT IN ('information_schema', 'performance_schema', 'mysql', 'sys')
63
            AND `SCHEMA_NAME` REGEXP ?
64
            AND `SCHEMA_NAME` LIKE ?
65
        ",
66
    )
67
    .bind(create_user_group_matching_regex(unix_user, group_denylist))
68
    .bind(format!("{database_prefix}%"))
69
    .fetch_all(connection)
70
    .await;
71

            
72
    match result {
73
        Ok(rows) => rows
74
            .into_iter()
75
            .filter_map(|row| {
76
                let database: String = row.try_get("database").ok()?;
77
                Some(database.into())
78
            })
79
            .collect(),
80
        Err(err) => {
81
            tracing::error!(
82
                "Failed to complete database name for prefix '{}' and user '{}': {:?}",
83
                database_prefix,
84
                unix_user.username,
85
                err
86
            );
87
            vec![]
88
        }
89
    }
90
}
91

            
92
pub async fn create_databases(
93
    database_names: &[MySQLDatabase],
94
    unix_user: &UnixUser,
95
    connection: &mut MySqlConnection,
96
    _db_is_mariadb: bool,
97
    group_denylist: &GroupDenylist,
98
) -> CreateDatabasesResponse {
99
    let mut results = BTreeMap::new();
100

            
101
    for database_name in database_names.iter().cloned() {
102
        if let Err(err) = validate_db_or_user_request(
103
            &DbOrUser::Database(database_name.clone()),
104
            unix_user,
105
            group_denylist,
106
        )
107
        .map_err(CreateDatabaseError::ValidationError)
108
        {
109
            results.insert(database_name.clone(), Err(err));
110
            continue;
111
        }
112

            
113
        match unsafe_database_exists(&database_name, &mut *connection).await {
114
            Ok(true) => {
115
                results.insert(
116
                    database_name.clone(),
117
                    Err(CreateDatabaseError::DatabaseAlreadyExists),
118
                );
119
                continue;
120
            }
121
            Err(err) => {
122
                results.insert(
123
                    database_name.clone(),
124
                    Err(CreateDatabaseError::MySqlError(err.to_string())),
125
                );
126
                continue;
127
            }
128
            _ => {}
129
        }
130

            
131
        let statement = AssertSqlSafe(format!(
132
            "CREATE DATABASE {}",
133
            quote_identifier(&database_name)
134
        ));
135
        let result = sqlx::query(statement)
136
            .execute(&mut *connection)
137
            .await
138
            .map(|_| ())
139
            .map_err(|err| CreateDatabaseError::MySqlError(err.to_string()));
140

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

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

            
148
    results
149
}
150

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

            
160
    for database_name in database_names.iter().cloned() {
161
        if let Err(err) = validate_db_or_user_request(
162
            &DbOrUser::Database(database_name.clone()),
163
            unix_user,
164
            group_denylist,
165
        )
166
        .map_err(DropDatabaseError::ValidationError)
167
        {
168
            results.insert(database_name.clone(), Err(err));
169
            continue;
170
        }
171

            
172
        match unsafe_database_exists(&database_name, &mut *connection).await {
173
            Ok(false) => {
174
                results.insert(
175
                    database_name.clone(),
176
                    Err(DropDatabaseError::DatabaseDoesNotExist),
177
                );
178
                continue;
179
            }
180
            Err(err) => {
181
                results.insert(
182
                    database_name.clone(),
183
                    Err(DropDatabaseError::MySqlError(err.to_string())),
184
                );
185
                continue;
186
            }
187
            _ => {}
188
        }
189

            
190
        let statement = AssertSqlSafe(format!(
191
            "DROP DATABASE {}",
192
            quote_identifier(&database_name)
193
        ));
194
        let result = sqlx::query(statement)
195
            .execute(&mut *connection)
196
            .await
197
            .map(|_| ())
198
            .map_err(|err| DropDatabaseError::MySqlError(err.to_string()));
199

            
200
        if let Err(err) = &result {
201
            tracing::error!("Failed to drop database '{}': {:?}", &database_name, err);
202
        }
203

            
204
        results.insert(database_name, result);
205
    }
206

            
207
    results
208
}
209

            
210
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
211
pub struct DatabaseRow {
212
    pub database: MySQLDatabase,
213
    pub tables: Vec<String>,
214
    pub table_count: u64,
215
    pub users: Vec<MySQLUser>,
216
    pub user_count: u64,
217
    pub collation: Option<String>,
218
    pub character_set: Option<String>,
219
    pub size_bytes: u64,
220
}
221

            
222
impl FromRow<'_, sqlx::mysql::MySqlRow> for DatabaseRow {
223
    fn from_row(row: &sqlx::mysql::MySqlRow) -> Result<Self, sqlx::Error> {
224
        Ok(DatabaseRow {
225
            database: row.try_get::<String, _>("database")?.into(),
226
            tables: {
227
                let s: Option<String> = row.try_get("tables")?;
228
                s.and_then(|s| {
229
                    if s.is_empty() {
230
                        None
231
                    } else {
232
                        Some(s.split(',').map(std::borrow::ToOwned::to_owned).collect())
233
                    }
234
                })
235
                .unwrap_or_default()
236
            },
237
            table_count: row.try_get::<u64, _>("table_count")?,
238
            users: {
239
                let s: Option<String> = row.try_get("users")?;
240
                s.and_then(|s| {
241
                    if s.is_empty() {
242
                        None
243
                    } else {
244
                        Some(s.split(',').map(|s| s.to_owned().into()).collect())
245
                    }
246
                })
247
                .unwrap_or_default()
248
            },
249
            user_count: row.try_get::<u64, _>("user_count")?,
250
            collation: row.try_get::<Option<String>, _>("collation")?,
251
            character_set: row.try_get::<Option<String>, _>("character_set")?,
252
            size_bytes: row.try_get::<u64, _>("size_bytes")?,
253
        })
254
    }
255
}
256

            
257
fn list_database_query(include_all_tables_and_users: bool) -> AssertSqlSafe<String> {
258
    let limit_clause = if include_all_tables_and_users {
259
        "".to_string()
260
    } else {
261
        format!(" LIMIT {}", MAX_SHOW_DB_RELATED_ITEMS)
262
    };
263

            
264
    AssertSqlSafe(format!(
265
        r"
266
            SELECT
267
                CAST(s.SCHEMA_NAME AS CHAR(64)) AS `database`,
268
                t.tables,
269
                CAST(COALESCE(sz.table_count, 0) AS UNSIGNED) AS table_count,
270
                u.users,
271
                CAST(COALESCE(uc.user_count, 0) AS UNSIGNED) AS user_count,
272
                s.DEFAULT_COLLATION_NAME AS `collation`,
273
                s.DEFAULT_CHARACTER_SET_NAME AS `character_set`,
274
                CAST(COALESCE(sz.size_bytes, 0) AS UNSIGNED) AS size_bytes
275
            FROM information_schema.SCHEMATA s
276

            
277
            LEFT JOIN (
278
                SELECT
279
                    x.TABLE_SCHEMA,
280
                    GROUP_CONCAT(x.TABLE_NAME ORDER BY x.TABLE_NAME SEPARATOR ',') AS tables
281
                FROM (
282
                    SELECT
283
                        TABLE_SCHEMA,
284
                        TABLE_NAME
285
                    FROM information_schema.TABLES
286
                    WHERE TABLE_SCHEMA = ?
287
                    ORDER BY TABLE_NAME{limit_clause}
288
                ) x
289
                GROUP BY x.TABLE_SCHEMA
290
            ) t
291
                ON t.TABLE_SCHEMA = s.SCHEMA_NAME
292

            
293
            LEFT JOIN (
294
                SELECT
295
                    x.DB,
296
                    GROUP_CONCAT(x.User ORDER BY x.User SEPARATOR ',') AS users
297
                FROM (
298
                    SELECT DISTINCT
299
                        DB,
300
                        User
301
                    FROM mysql.db
302
                    WHERE DB = ?
303
                    ORDER BY User{limit_clause}
304
                ) x
305
                GROUP BY x.DB
306
            ) u
307
                ON u.DB = s.SCHEMA_NAME
308

            
309
            LEFT JOIN (
310
                SELECT
311
                    TABLE_SCHEMA,
312
                    SUM(DATA_LENGTH + INDEX_LENGTH) AS size_bytes,
313
                    COUNT(*) AS table_count
314
                FROM information_schema.TABLES
315
                WHERE TABLE_SCHEMA = ?
316
                GROUP BY TABLE_SCHEMA
317
            ) sz
318
                ON sz.TABLE_SCHEMA = s.SCHEMA_NAME
319

            
320
            LEFT JOIN (
321
                SELECT
322
                    DB,
323
                    COUNT(DISTINCT User) AS user_count
324
                FROM mysql.db
325
                WHERE DB = ?
326
                GROUP BY DB
327
            ) uc
328
                ON uc.DB = s.SCHEMA_NAME
329

            
330
            WHERE s.SCHEMA_NAME REGEXP ?
331
            AND s.SCHEMA_NAME NOT IN (
332
                'information_schema',
333
                'performance_schema',
334
                'mysql',
335
                'sys'
336
            )
337
        "
338
    ))
339
}
340

            
341
pub async fn list_databases(
342
    database_names: &[MySQLDatabase],
343
    unix_user: &UnixUser,
344
    connection: &mut MySqlConnection,
345
    _db_is_mariadb: bool,
346
    group_denylist: &GroupDenylist,
347
    include_all_tables_and_users: bool,
348
) -> ListDatabasesResponse {
349
    let mut results = BTreeMap::new();
350

            
351
    for database_name in database_names.iter().cloned() {
352
        if let Err(err) = validate_db_or_user_request(
353
            &DbOrUser::Database(database_name.clone()),
354
            unix_user,
355
            group_denylist,
356
        )
357
        .map_err(ListDatabasesError::ValidationError)
358
        {
359
            results.insert(database_name.clone(), Err(err));
360
            continue;
361
        }
362

            
363
        let query = list_database_query(include_all_tables_and_users);
364

            
365
        let result = sqlx::query_as::<_, DatabaseRow>(query)
366
            .bind(database_name.to_string())
367
            .bind(database_name.to_string())
368
            .bind(database_name.to_string())
369
            .bind(database_name.to_string())
370
            .bind(database_name.to_string())
371
            .fetch_optional(&mut *connection)
372
            .await
373
            .map_err(|err| ListDatabasesError::MySqlError(err.to_string()))
374
            .and_then(|database| {
375
                database.map_or_else(|| Err(ListDatabasesError::DatabaseDoesNotExist), Ok)
376
            });
377

            
378
        if let Err(err) = &result {
379
            tracing::error!("Failed to list database '{}': {:?}", &database_name, err);
380
        }
381

            
382
        // TODO: should we assert that the users are also owned by the unix_user from the request?
383

            
384
        results.insert(database_name, result);
385
    }
386

            
387
    results
388
}
389

            
390
fn list_all_databases_for_user_query(include_all_tables_and_users: bool) -> AssertSqlSafe<String> {
391
    let row_limit_clause = if include_all_tables_and_users {
392
        String::new()
393
    } else {
394
        format!("WHERE row_num <= {MAX_SHOW_DB_RELATED_ITEMS}")
395
    };
396

            
397
    AssertSqlSafe(format!(
398
        r"
399
            SELECT
400
                CAST(s.SCHEMA_NAME AS CHAR(64)) AS `database`,
401
                t.tables,
402
                CAST(COALESCE(sz.table_count, 0) AS UNSIGNED) AS table_count,
403
                u.users,
404
                CAST(COALESCE(uc.user_count, 0) AS UNSIGNED) AS user_count,
405
                s.DEFAULT_COLLATION_NAME AS `collation`,
406
                s.DEFAULT_CHARACTER_SET_NAME AS `character_set`,
407
                CAST(COALESCE(sz.size_bytes, 0) AS UNSIGNED) AS size_bytes
408
            FROM information_schema.SCHEMATA s
409

            
410
            LEFT JOIN (
411
                SELECT
412
                    x.TABLE_SCHEMA,
413
                    GROUP_CONCAT(x.TABLE_NAME ORDER BY x.TABLE_NAME SEPARATOR ',') AS tables
414
                FROM (
415
                    SELECT
416
                        TABLE_SCHEMA,
417
                        TABLE_NAME,
418
                        ROW_NUMBER() OVER (PARTITION BY TABLE_SCHEMA ORDER BY TABLE_NAME) AS row_num
419
                    FROM information_schema.TABLES
420
                    WHERE TABLE_SCHEMA REGEXP ?
421
                ) x
422
                {row_limit_clause}
423
                GROUP BY x.TABLE_SCHEMA
424
            ) t
425
                ON t.TABLE_SCHEMA = s.SCHEMA_NAME
426

            
427
            LEFT JOIN (
428
                SELECT
429
                    x.DB,
430
                    GROUP_CONCAT(x.User ORDER BY x.User SEPARATOR ',') AS users
431
                FROM (
432
                    SELECT
433
                        DB,
434
                        User,
435
                        ROW_NUMBER() OVER (PARTITION BY DB ORDER BY User) AS row_num
436
                    FROM (
437
                        SELECT DISTINCT DB, User
438
                        FROM mysql.db
439
                        WHERE DB REGEXP ?
440
                    ) d
441
                ) x
442
                {row_limit_clause}
443
                GROUP BY x.DB
444
            ) u
445
                ON u.DB = s.SCHEMA_NAME
446

            
447
            LEFT JOIN (
448
                SELECT
449
                    TABLE_SCHEMA,
450
                    SUM(DATA_LENGTH + INDEX_LENGTH) AS size_bytes,
451
                    COUNT(*) AS table_count
452
                FROM information_schema.TABLES
453
                WHERE TABLE_SCHEMA REGEXP ?
454
                GROUP BY TABLE_SCHEMA
455
            ) sz
456
                ON sz.TABLE_SCHEMA = s.SCHEMA_NAME
457

            
458
            LEFT JOIN (
459
                SELECT
460
                    DB,
461
                    COUNT(DISTINCT User) AS user_count
462
                FROM mysql.db
463
                WHERE DB REGEXP ?
464
                GROUP BY DB
465
            ) uc
466
                ON uc.DB = s.SCHEMA_NAME
467

            
468
            WHERE s.SCHEMA_NAME REGEXP ?
469
            AND s.SCHEMA_NAME NOT IN (
470
                'information_schema',
471
                'performance_schema',
472
                'mysql',
473
                'sys'
474
            )
475

            
476
            ORDER BY s.SCHEMA_NAME
477
        "
478
    ))
479
}
480

            
481
pub async fn list_all_databases_for_user(
482
    unix_user: &UnixUser,
483
    connection: &mut MySqlConnection,
484
    _db_is_mariadb: bool,
485
    group_denylist: &GroupDenylist,
486
    include_all_tables_and_users: bool,
487
) -> ListAllDatabasesResponse {
488
    let query = list_all_databases_for_user_query(include_all_tables_and_users);
489
    let user_group_regex = create_user_group_matching_regex(unix_user, group_denylist);
490

            
491
    let result = sqlx::query_as::<_, DatabaseRow>(query)
492
        .bind(&user_group_regex)
493
        .bind(&user_group_regex)
494
        .bind(&user_group_regex)
495
        .bind(&user_group_regex)
496
        .bind(&user_group_regex)
497
        .fetch_all(connection)
498
        .await
499
        .map_err(|err| ListAllDatabasesError::MySqlError(err.to_string()));
500

            
501
    // TODO: should we assert that the users are also owned by the unix_user from the request?
502

            
503
    if let Err(err) = &result {
504
        tracing::error!(
505
            "Failed to list databases for user '{}': {:?}",
506
            unix_user.username,
507
            err
508
        );
509
    }
510

            
511
    result
512
}