Skip to main content

revolt_database/drivers/
mod.rs

1#[cfg(feature = "mongodb")]
2mod mongodb;
3mod reference;
4
5
6use rand::Rng;
7use revolt_config::config;
8
9#[cfg(feature = "mongodb")]
10pub use self::mongodb::*;
11pub use self::reference::*;
12
13/// Database information to use to create a client
14pub enum DatabaseInfo {
15    /// Auto-detect the database in use
16    Auto,
17    /// Auto-detect the database in use and create an empty testing database
18    Test(String),
19    /// Use the mock database
20    Reference,
21    /// Connect to MongoDB
22    #[cfg(feature = "mongodb")]
23    MongoDb { uri: String, database_name: String },
24    /// Use existing MongoDB connection
25    #[cfg(feature = "mongodb")]
26    MongoDbFromClient(::mongodb::Client, String),
27}
28
29/// Database
30#[derive(Clone, Debug)]
31pub enum Database {
32    /// Mock database
33    Reference(ReferenceDb),
34    /// MongoDB database
35    #[cfg(feature = "mongodb")]
36    MongoDb(MongoDb),
37}
38
39impl DatabaseInfo {
40    /// Create a database client from the given database information
41    #[async_recursion]
42    pub async fn connect(self) -> Result<Database, String> {
43        let config = config().await;
44
45        match self {
46            DatabaseInfo::Auto => {
47                if std::env::var("TEST_DB").is_ok() {
48                    DatabaseInfo::Test(format!(
49                        "revolt_test_{}",
50                        rand::thread_rng().gen_range(1_000_000..10_000_000)
51                    ))
52                    .connect()
53                    .await
54                } else if !config.database.mongodb.is_empty() {
55                    #[cfg(feature = "mongodb")]
56                    return DatabaseInfo::MongoDb {
57                        uri: config.database.mongodb,
58                        database_name: "revolt".to_string(),
59                    }
60                    .connect()
61                    .await;
62
63                    #[cfg(not(feature = "mongodb"))]
64                    return Err("MongoDB not enabled.".to_string());
65                } else {
66                    DatabaseInfo::Reference.connect().await
67                }
68            }
69            DatabaseInfo::Test(database_name) => {
70                match std::env::var("TEST_DB")
71                    .expect("`TEST_DB` environment variable should be set to REFERENCE or MONGODB")
72                    .as_str()
73                {
74                    "REFERENCE" => DatabaseInfo::Reference.connect().await,
75                    "MONGODB" => {
76                        #[cfg(feature = "mongodb")]
77                        return DatabaseInfo::MongoDb {
78                            uri: config.database.mongodb,
79                            database_name,
80                        }
81                        .connect()
82                        .await;
83
84                        #[cfg(not(feature = "mongodb"))]
85                        return Err("MongoDB not enabled.".to_string());
86                    }
87                    _ => unreachable!("must specify REFERENCE or MONGODB"),
88                }
89            }
90            DatabaseInfo::Reference => Ok(Database::Reference(Default::default())),
91            #[cfg(feature = "mongodb")]
92            DatabaseInfo::MongoDb { uri, database_name } => {
93                let client = ::mongodb::Client::with_uri_str(uri)
94                    .await
95                    .map_err(|_| "Failed to init db connection.".to_string())?;
96
97                Ok(Database::MongoDb(MongoDb(client, database_name)))
98            }
99            #[cfg(feature = "mongodb")]
100            DatabaseInfo::MongoDbFromClient(client, database_name) => {
101                Ok(Database::MongoDb(MongoDb(client, database_name)))
102            }
103        }
104    }
105}