Skip to main content

luckychacha_sqlx_pg_tester/
lib.rs

1use sqlx::{migrate::Migrator, Connection, Executor, PgPool};
2use std::{path::Path, thread};
3
4use sqlx::PgConnection;
5use tokio::runtime::Runtime;
6use uuid::Uuid;
7
8pub struct TestDb {
9    pub dbname: String,
10    pub user: String,
11    pub password: String,
12    pub host: String,
13    pub port: u16,
14}
15
16impl TestDb {
17    pub fn new(
18        user: impl Into<String>,
19        password: impl Into<String>,
20        host: impl Into<String>,
21        port: u16,
22        migration_path: impl Into<String>,
23    ) -> Self {
24        let uuid = Uuid::new_v4();
25        let dbname = format!("test_{uuid}");
26        let dbname_clone = dbname.clone();
27
28        let user = user.into();
29        let password = password.into();
30        let host = host.into();
31        let migration_path = migration_path.into();
32
33        let tdb = Self {
34            dbname,
35            user,
36            password,
37            host,
38            port,
39        };
40
41        let server_url = tdb.server_url();
42        let db_url = tdb.url();
43
44        thread::spawn(move || {
45            let rt = Runtime::new().unwrap();
46            rt.block_on(async move {
47                let mut conn = PgConnection::connect(&server_url).await.unwrap();
48                conn.execute(format!(r#"CREATE DATABASE "{dbname_clone}""#).as_str())
49                    .await
50                    .expect("Failed when create database {dbname_clone}.");
51
52                let mut conn = PgConnection::connect(&db_url).await.unwrap();
53
54                let m = Migrator::new(Path::new(&migration_path)).await.unwrap();
55                m.run(&mut conn).await.unwrap();
56            })
57        })
58        .join()
59        .expect("Create database failed.");
60
61        tdb
62    }
63
64    pub fn server_url(&self) -> String {
65        if self.password.is_empty() {
66            format!("postgres://{}@{}:{}", self.user, self.host, self.port)
67        } else {
68            format!(
69                "postgres://{}:{}@{}:{}",
70                self.user, self.password, self.host, self.port
71            )
72        }
73    }
74
75    pub fn url(&self) -> String {
76        format!("{}/{}", self.server_url(), self.dbname)
77    }
78
79    pub async fn get_pool(&self) -> PgPool {
80        sqlx::postgres::PgPoolOptions::new()
81            .max_connections(5)
82            .connect(&self.url())
83            .await
84            .unwrap()
85    }
86}
87
88impl Drop for TestDb {
89    fn drop(&mut self) {
90        let server_url = self.server_url();
91        let database_name = self.dbname.clone();
92        thread::spawn(move || {
93            let rt = Runtime::new().unwrap();
94            rt.block_on(async move {
95                let mut conn = PgConnection::connect(&server_url).await.unwrap();
96
97                #[allow(clippy::expect_used)]
98                sqlx::query(&format!(r#"SELECT pg_terminate_backend(pid) FROM pg_stat_activity WHERE pid <> pg_backend_pid() AND datname = '{database_name}'"#))
99                    .execute(&mut conn)
100                    .await
101                    .expect("Terminate all other connections");
102                #[allow(clippy::expect_used)]
103                sqlx::query(&format!(r#"DROP DATABASE "{database_name}""#))
104                    .execute(&mut conn)
105                    .await
106                    .expect("Deleting the database");
107            })
108        });
109    }
110}
111
112#[cfg(test)]
113mod tests {
114    use super::*;
115
116    #[tokio::test]
117    async fn test_db_should_create_and_drop() {
118        let tdb = TestDb::new("postgres", "postgres", "localhost", 5432, "./migrations");
119        let pool = tdb.get_pool().await;
120        sqlx::query("INSERT INTO todos(title) VALUES ('test')")
121            .execute(&pool)
122            .await
123            .unwrap();
124        let (id, title) = sqlx::query_as::<_, (i32, String)>("SELECT id, title FROM todos")
125            .fetch_one(&pool)
126            .await
127            .unwrap();
128
129        assert_eq!(id, 1);
130        assert_eq!(title, "test");
131    }
132}