1use std::str::FromStr;
2use std::time::Duration;
3
4use sqlx::migrate::Migrator;
5use sqlx::sqlite::{SqliteConnectOptions, SqliteJournalMode, SqlitePoolOptions};
6use sqlx::{Error, Pool, Sqlite, SqlitePool, migrate::MigrateDatabase};
7use tracing::{error, info};
8
9static MIGRATOR: Migrator = sqlx::migrate!(); pub struct Database {
12 pub pool: Pool<Sqlite>,
13}
14
15impl Database {
16 pub async fn connect(url: &str) -> Result<Database, Error> {
19 if !Sqlite::database_exists(url).await.unwrap_or(false) {
20 info!(event = "db_creation_started", outcome = "progress", database_url = %url);
21 Sqlite::create_database(url).await?;
22 info!(event = "db_creation_completed", outcome = "success", database_url = %url);
23 }
24
25 let options = SqliteConnectOptions::from_str(url)?
26 .foreign_keys(true)
30 .journal_mode(SqliteJournalMode::Wal)
34 .busy_timeout(Duration::from_secs(5));
35
36 let pool = SqlitePool::connect_with(options).await?;
37
38 run_migrations(&pool).await?;
39
40 Ok(Database { pool })
41 }
42
43 pub async fn connect_in_memory() -> Result<Database, Error> {
47 let pool = SqlitePoolOptions::new()
48 .max_connections(1)
49 .connect_with(SqliteConnectOptions::from_str("sqlite::memory:")?.foreign_keys(true))
50 .await?;
51
52 run_migrations(&pool).await?;
53
54 Ok(Database { pool })
55 }
56}
57
58async fn run_migrations(pool: &Pool<Sqlite>) -> Result<(), Error> {
59 MIGRATOR.run(pool).await.map_err(|error| {
60 error!(event = "db_migration_failed", outcome = "failure", error = %error);
63 Error::Migrate(Box::new(error))
64 })?;
65 info!(event = "db_migration_completed", outcome = "success");
66 Ok(())
67}
68
69#[cfg(test)]
70mod tests {
71 use super::*;
72
73 #[tokio::test]
74 async fn connect_creates_file_and_runs_migrations() {
75 let file =
78 std::env::temp_dir().join(format!("acme-proxy-test-{}.db", uuid::Uuid::new_v4()));
79 let url = format!("sqlite://{}", file.display());
80
81 let database = Database::connect(&url).await.unwrap();
82
83 let count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM nonces;")
85 .fetch_one(&database.pool)
86 .await
87 .unwrap();
88 assert_eq!(count, 0);
89
90 let journal: String = sqlx::query_scalar("PRAGMA journal_mode;")
93 .fetch_one(&database.pool)
94 .await
95 .unwrap();
96 assert_eq!(journal.to_lowercase(), "wal");
97 let foreign_keys: i64 = sqlx::query_scalar("PRAGMA foreign_keys;")
98 .fetch_one(&database.pool)
99 .await
100 .unwrap();
101 assert_eq!(foreign_keys, 1);
102
103 database.pool.close().await;
104 for suffix in ["", "-wal", "-shm"] {
106 let _ = std::fs::remove_file(format!("{}{suffix}", file.display()));
107 }
108 }
109
110 #[tokio::test]
113 async fn foreign_keys_and_the_nonce_sweep_are_indexed() {
114 let database = Database::connect_in_memory().await.unwrap();
115 let names: Vec<String> =
116 sqlx::query_scalar("SELECT name FROM sqlite_master WHERE type = 'index';")
117 .fetch_all(&database.pool)
118 .await
119 .unwrap();
120
121 for expected in [
122 "idx_orders_account_id",
123 "idx_authorizations_order",
124 "idx_challenges_authz",
125 "idx_nonces_created_at",
126 "idx_orders_cert_serial",
127 "idx_orders_replaces_claim",
128 ] {
129 assert!(
130 names.iter().any(|name| name == expected),
131 "missing index {expected}; have {names:?}"
132 );
133 }
134 }
135
136 #[tokio::test]
144 async fn one_predecessor_can_only_be_claimed_by_one_live_order() {
145 let database = Database::connect_in_memory().await.unwrap();
146 let cert_id = "aYhba4dGQEHhs3uEe6CuLN4ByNQ.AIdlQyE";
147
148 sqlx::query(
149 "INSERT INTO accounts (id, profile, pubkey, contact, status, created_at) \
150 VALUES ('acct', 'default', X'00', '[]', 'valid', 0);",
151 )
152 .execute(&database.pool)
153 .await
154 .unwrap();
155
156 let insert = |id: &'static str, status: &'static str| {
157 let pool = database.pool.clone();
158 async move {
159 sqlx::query(
160 "INSERT INTO orders (id, profile, account_id, status, identifiers, expires, \
161 replaces, created_at) VALUES (?, 'default', 'acct', ?, '[]', 0, ?, 0);",
162 )
163 .bind(id)
164 .bind(status)
165 .bind(cert_id)
166 .execute(&pool)
167 .await
168 }
169 };
170
171 insert("first", "pending").await.unwrap();
172
173 let error = insert("second", "pending").await.unwrap_err();
180 match &error {
181 sqlx::Error::Database(db) => {
182 assert!(db.is_unique_violation(), "got {error}");
183 assert!(
184 db.message().contains("orders.replaces"),
185 "the violation must name the column, got {:?}",
186 db.message()
187 );
188 }
189 other => panic!("expected a database error, got {other}"),
190 }
191
192 insert("third", "invalid").await.unwrap();
195
196 sqlx::query("UPDATE orders SET status = 'invalid' WHERE id = 'first';")
198 .execute(&database.pool)
199 .await
200 .unwrap();
201 insert("fourth", "pending").await.unwrap();
202 }
203
204 #[tokio::test]
208 async fn status_columns_reject_values_outside_the_state_machine() {
209 let database = Database::connect_in_memory().await.unwrap();
210
211 let result = sqlx::query(
212 "INSERT INTO accounts (id, profile, pubkey, contact, status, created_at) \
213 VALUES ('a', 'default', X'00', '[]', 'definitely-not-a-status', 0);",
214 )
215 .execute(&database.pool)
216 .await;
217 assert!(result.is_err(), "an unknown account status must be refused");
218
219 sqlx::query(
221 "INSERT INTO accounts (id, profile, pubkey, contact, status, created_at) \
222 VALUES ('a', 'default', X'00', '[]', 'valid', 0);",
223 )
224 .execute(&database.pool)
225 .await
226 .unwrap();
227 }
228
229 #[tokio::test]
233 async fn deleting_an_account_cascades_to_its_orders() {
234 let database = Database::connect_in_memory().await.unwrap();
235
236 sqlx::query(
237 "INSERT INTO accounts (id, profile, pubkey, contact, status, created_at) \
238 VALUES ('acct', 'default', X'00', '[]', 'valid', 0);",
239 )
240 .execute(&database.pool)
241 .await
242 .unwrap();
243 sqlx::query(
244 "INSERT INTO orders (id, profile, account_id, status, identifiers, expires, created_at) \
245 VALUES ('ord', 'default', 'acct', 'pending', '[]', 0, 0);",
246 )
247 .execute(&database.pool)
248 .await
249 .unwrap();
250 sqlx::query(
251 "INSERT INTO authorizations (id, order_id, identifier, status, expires, created_at) \
252 VALUES ('az', 'ord', '{}', 'pending', 0, 0);",
253 )
254 .execute(&database.pool)
255 .await
256 .unwrap();
257 sqlx::query(
258 "INSERT INTO challenges (id, authz_id, type, token, status, created_at) \
259 VALUES ('ch', 'az', 'http-01', 't', 'pending', 0);",
260 )
261 .execute(&database.pool)
262 .await
263 .unwrap();
264
265 sqlx::query("DELETE FROM accounts WHERE id = 'acct';")
266 .execute(&database.pool)
267 .await
268 .unwrap();
269
270 for (table, query) in [
271 ("orders", "SELECT COUNT(*) FROM orders;"),
272 ("authorizations", "SELECT COUNT(*) FROM authorizations;"),
273 ("challenges", "SELECT COUNT(*) FROM challenges;"),
274 ] {
275 let count: i64 = sqlx::query_scalar(query)
276 .fetch_one(&database.pool)
277 .await
278 .unwrap();
279 assert_eq!(count, 0, "{table} should have been cascaded away");
280 }
281 }
282
283 #[tokio::test]
285 async fn an_order_cannot_have_duplicate_authorizations_for_one_identifier() {
286 let database = Database::connect_in_memory().await.unwrap();
287 sqlx::query(
288 "INSERT INTO accounts (id, profile, pubkey, contact, status, created_at) \
289 VALUES ('acct', 'default', X'00', '[]', 'valid', 0);",
290 )
291 .execute(&database.pool)
292 .await
293 .unwrap();
294 sqlx::query(
295 "INSERT INTO orders (id, profile, account_id, status, identifiers, expires, created_at) \
296 VALUES ('ord', 'default', 'acct', 'pending', '[]', 0, 0);",
297 )
298 .execute(&database.pool)
299 .await
300 .unwrap();
301
302 let insert = |id: &'static str| {
303 sqlx::query(
304 "INSERT INTO authorizations (id, order_id, identifier, status, expires, created_at) \
305 VALUES (?, 'ord', '{\"type\":\"dns\",\"value\":\"example.com\"}', 'pending', 0, 0);",
306 )
307 .bind(id)
308 .execute(&database.pool)
309 };
310
311 insert("az1").await.unwrap();
312 assert!(
313 insert("az2").await.is_err(),
314 "a second authorization for the same identifier must be refused"
315 );
316 }
317}