luckychacha_sqlx_pg_tester/
lib.rs1use 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}