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 use crate::random::random_token;
74
75 #[tokio::test]
76 async fn connect_creates_file_and_runs_migrations() {
77 let file =
80 std::env::temp_dir().join(format!("acme-proxy-test-{}.db", uuid::Uuid::new_v4()));
81 let url = format!("sqlite://{}", file.display());
82
83 let database = Database::connect(&url).await.unwrap();
84
85 let count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM nonces;")
87 .fetch_one(&database.pool)
88 .await
89 .unwrap();
90 assert_eq!(count, 0);
91
92 let journal: String = sqlx::query_scalar("PRAGMA journal_mode;")
95 .fetch_one(&database.pool)
96 .await
97 .unwrap();
98 assert_eq!(journal.to_lowercase(), "wal");
99 let foreign_keys: i64 = sqlx::query_scalar("PRAGMA foreign_keys;")
100 .fetch_one(&database.pool)
101 .await
102 .unwrap();
103 assert_eq!(foreign_keys, 1);
104
105 database.pool.close().await;
106 for suffix in ["", "-wal", "-shm"] {
108 let _ = std::fs::remove_file(format!("{}{suffix}", file.display()));
109 }
110 }
111
112 #[tokio::test]
115 async fn foreign_keys_and_the_nonce_sweep_are_indexed() {
116 let database = Database::connect_in_memory().await.unwrap();
117 let names: Vec<String> =
118 sqlx::query_scalar("SELECT name FROM sqlite_master WHERE type = 'index';")
119 .fetch_all(&database.pool)
120 .await
121 .unwrap();
122
123 for expected in [
124 "idx_orders_account_id",
125 "idx_authorizations_order",
126 "idx_challenges_authz",
127 "idx_nonces_created_at",
128 "idx_orders_cert_serial",
129 "idx_orders_replaces_claim",
130 ] {
131 assert!(
132 names.iter().any(|name| name == expected),
133 "missing index {expected}; have {names:?}"
134 );
135 }
136 }
137
138 #[tokio::test]
148 async fn declared_token_widths_match_random_token() {
149 let database = Database::connect_in_memory().await.unwrap();
150 let expected = format!("VARCHAR({})", random_token().len());
151
152 for (table, column) in [("nonces", "value"), ("challenges", "token")] {
153 let columns: Vec<(String, String)> =
154 sqlx::query_as("SELECT name, type FROM pragma_table_info(?);")
155 .bind(table)
156 .fetch_all(&database.pool)
157 .await
158 .unwrap();
159
160 let declared = columns
161 .iter()
162 .find(|(name, _)| name == column)
163 .map(|(_, declared)| declared.as_str())
164 .unwrap_or_else(|| panic!("no column {table}.{column}"));
165
166 assert_eq!(
167 declared, expected,
168 "{table}.{column} declares a width the value no longer has"
169 );
170 }
171 }
172
173 #[tokio::test]
181 async fn one_predecessor_can_only_be_claimed_by_one_live_order() {
182 let database = Database::connect_in_memory().await.unwrap();
183 let cert_id = "aYhba4dGQEHhs3uEe6CuLN4ByNQ.AIdlQyE";
184
185 sqlx::query(
186 "INSERT INTO accounts (id, profile, pubkey, contact, status, created_at) \
187 VALUES ('acct', 'default', X'00', '[]', 'valid', 0);",
188 )
189 .execute(&database.pool)
190 .await
191 .unwrap();
192
193 let insert = |id: &'static str, status: &'static str| {
194 let pool = database.pool.clone();
195 async move {
196 sqlx::query(
197 "INSERT INTO orders (id, profile, account_id, status, identifiers, expires, \
198 replaces, created_at) VALUES (?, 'default', 'acct', ?, '[]', 0, ?, 0);",
199 )
200 .bind(id)
201 .bind(status)
202 .bind(cert_id)
203 .execute(&pool)
204 .await
205 }
206 };
207
208 insert("first", "pending").await.unwrap();
209
210 let error = insert("second", "pending").await.unwrap_err();
217 match &error {
218 sqlx::Error::Database(db) => {
219 assert!(db.is_unique_violation(), "got {error}");
220 assert!(
221 db.message().contains("orders.replaces"),
222 "the violation must name the column, got {:?}",
223 db.message()
224 );
225 }
226 other => panic!("expected a database error, got {other}"),
227 }
228
229 insert("third", "invalid").await.unwrap();
232
233 sqlx::query("UPDATE orders SET status = 'invalid' WHERE id = 'first';")
235 .execute(&database.pool)
236 .await
237 .unwrap();
238 insert("fourth", "pending").await.unwrap();
239 }
240
241 #[tokio::test]
245 async fn status_columns_reject_values_outside_the_state_machine() {
246 let database = Database::connect_in_memory().await.unwrap();
247
248 let result = sqlx::query(
249 "INSERT INTO accounts (id, profile, pubkey, contact, status, created_at) \
250 VALUES ('a', 'default', X'00', '[]', 'definitely-not-a-status', 0);",
251 )
252 .execute(&database.pool)
253 .await;
254 assert!(result.is_err(), "an unknown account status must be refused");
255
256 sqlx::query(
258 "INSERT INTO accounts (id, profile, pubkey, contact, status, created_at) \
259 VALUES ('a', 'default', X'00', '[]', 'valid', 0);",
260 )
261 .execute(&database.pool)
262 .await
263 .unwrap();
264 }
265
266 #[tokio::test]
270 async fn deleting_an_account_cascades_to_its_orders() {
271 let database = Database::connect_in_memory().await.unwrap();
272
273 sqlx::query(
274 "INSERT INTO accounts (id, profile, pubkey, contact, status, created_at) \
275 VALUES ('acct', 'default', X'00', '[]', 'valid', 0);",
276 )
277 .execute(&database.pool)
278 .await
279 .unwrap();
280 sqlx::query(
281 "INSERT INTO orders (id, profile, account_id, status, identifiers, expires, created_at) \
282 VALUES ('ord', 'default', 'acct', 'pending', '[]', 0, 0);",
283 )
284 .execute(&database.pool)
285 .await
286 .unwrap();
287 sqlx::query(
288 "INSERT INTO authorizations (id, order_id, identifier, status, expires, created_at) \
289 VALUES ('az', 'ord', '{}', 'pending', 0, 0);",
290 )
291 .execute(&database.pool)
292 .await
293 .unwrap();
294 sqlx::query(
295 "INSERT INTO challenges (id, authz_id, type, token, status, created_at) \
296 VALUES ('ch', 'az', 'http-01', 't', 'pending', 0);",
297 )
298 .execute(&database.pool)
299 .await
300 .unwrap();
301
302 sqlx::query("DELETE FROM accounts WHERE id = 'acct';")
303 .execute(&database.pool)
304 .await
305 .unwrap();
306
307 for (table, query) in [
308 ("orders", "SELECT COUNT(*) FROM orders;"),
309 ("authorizations", "SELECT COUNT(*) FROM authorizations;"),
310 ("challenges", "SELECT COUNT(*) FROM challenges;"),
311 ] {
312 let count: i64 = sqlx::query_scalar(query)
313 .fetch_one(&database.pool)
314 .await
315 .unwrap();
316 assert_eq!(count, 0, "{table} should have been cascaded away");
317 }
318 }
319
320 #[tokio::test]
322 async fn an_order_cannot_have_duplicate_authorizations_for_one_identifier() {
323 let database = Database::connect_in_memory().await.unwrap();
324 sqlx::query(
325 "INSERT INTO accounts (id, profile, pubkey, contact, status, created_at) \
326 VALUES ('acct', 'default', X'00', '[]', 'valid', 0);",
327 )
328 .execute(&database.pool)
329 .await
330 .unwrap();
331 sqlx::query(
332 "INSERT INTO orders (id, profile, account_id, status, identifiers, expires, created_at) \
333 VALUES ('ord', 'default', 'acct', 'pending', '[]', 0, 0);",
334 )
335 .execute(&database.pool)
336 .await
337 .unwrap();
338
339 let insert = |id: &'static str| {
340 sqlx::query(
341 "INSERT INTO authorizations (id, order_id, identifier, status, expires, created_at) \
342 VALUES (?, 'ord', '{\"type\":\"dns\",\"value\":\"example.com\"}', 'pending', 0, 0);",
343 )
344 .bind(id)
345 .execute(&database.pool)
346 };
347
348 insert("az1").await.unwrap();
349 assert!(
350 insert("az2").await.is_err(),
351 "a second authorization for the same identifier must be refused"
352 );
353 }
354}