1use std::{sync::Arc, time::Duration};
4
5use sqlx::AnyConnection;
6use sqlx::any::AnyPoolOptions;
7
8use super::backend::{DatabaseBackend, configure_postgres_timeouts};
9use crate::config::{PostgresConfig, SqliteConfig};
10
11const SQLITE_MEMORY_MAX_CONNECTIONS: u32 = 1;
12const DEFAULT_MAX_CONNECTIONS: u32 = 10;
13const SQLITE_BUSY_TIMEOUT_MS: u64 = 5_000;
14const SQLITE_WAL_MAX_ATTEMPTS: u32 = 5;
15const SQLITE_WAL_RETRY_DELAY: Duration = Duration::from_secs(1);
16const SQLITE_BUSY: i32 = 5;
17const SQLITE_LOCKED: i32 = 6;
18
19pub type DbPool = sqlx::Pool<sqlx::Any>;
21
22pub type DbTransaction<'a> = sqlx::Transaction<'a, sqlx::Any>;
24
25pub type DbResult<T> = Result<T, sqlx::Error>;
29
30struct PreparedDatabaseUrl {
31 value: String,
32 backend: DatabaseBackend,
33}
34
35fn prepare_db_url(url: Option<&str>) -> DbResult<PreparedDatabaseUrl> {
43 let url = url.unwrap_or("sqlite://./agentic_api.db");
44 let backend = DatabaseBackend::from_url(url).map_err(|error| sqlx::Error::Configuration(Box::new(error)))?;
45 let value = if backend != DatabaseBackend::Sqlite || has_query_param(url, "mode") {
46 url.to_owned()
47 } else {
48 append_query_param(url, "mode=rwc")
49 };
50 Ok(PreparedDatabaseUrl { value, backend })
51}
52
53fn has_query_param(url: &str, param: &str) -> bool {
54 query_param_value(url, param).is_some()
55}
56
57fn query_param_value<'a>(url: &'a str, param: &str) -> Option<&'a str> {
58 let (_, query) = url.split_once('?')?;
59 let query = query.split_once('#').map_or(query, |(query, _)| query);
60 query.split('&').find_map(|part| {
61 let (key, value) = part.split_once('=').map_or((part, ""), |(key, value)| (key, value));
62 (key == param).then_some(value)
63 })
64}
65
66fn append_query_param(url: &str, param: &str) -> String {
67 let (base, fragment) = url
68 .split_once('#')
69 .map_or((url, None), |(base, fragment)| (base, Some(fragment)));
70 let separator = if base.contains('?') { '&' } else { '?' };
71 let mut prepared = format!("{base}{separator}{param}");
72 if let Some(fragment) = fragment {
73 prepared.push('#');
74 prepared.push_str(fragment);
75 }
76 prepared
77}
78
79fn sqlite_is_memory_url(url: &str) -> bool {
80 query_param_value(url, "mode").is_some_and(|mode| mode.eq_ignore_ascii_case("memory")) || url.contains(":memory:")
81}
82
83fn sqlite_should_enable_wal(url: &str, backend: DatabaseBackend) -> bool {
84 backend == DatabaseBackend::Sqlite
85 && !sqlite_is_memory_url(url)
86 && !query_param_value(url, "mode").is_some_and(|mode| mode.eq_ignore_ascii_case("ro"))
87}
88
89fn sqlite_max_connections(url: &str, config: SqliteConfig) -> u32 {
90 if sqlite_is_memory_url(url) {
91 SQLITE_MEMORY_MAX_CONNECTIONS
92 } else {
93 config.max_connections
94 }
95}
96
97fn pool_options(
98 url: &str,
99 backend: DatabaseBackend,
100 sqlite_config: SqliteConfig,
101 postgres_config: PostgresConfig,
102) -> AnyPoolOptions {
103 match backend {
104 DatabaseBackend::Sqlite => AnyPoolOptions::new()
105 .max_connections(sqlite_max_connections(url, sqlite_config))
106 .after_connect(move |conn, _meta| Box::pin(configure_sqlite_connection(conn, sqlite_config))),
107 DatabaseBackend::Postgres => AnyPoolOptions::new()
108 .max_connections(postgres_config.max_connections)
109 .acquire_timeout(postgres_config.acquire_timeout)
110 .idle_timeout(postgres_config.idle_timeout)
111 .max_lifetime(postgres_config.max_lifetime)
112 .after_connect(move |conn, _meta| Box::pin(configure_postgres_connection(conn, postgres_config))),
113 DatabaseBackend::Other => AnyPoolOptions::new().max_connections(DEFAULT_MAX_CONNECTIONS),
114 }
115}
116
117async fn configure_sqlite_connection(conn: &mut AnyConnection, config: SqliteConfig) -> DbResult<()> {
118 sqlx::query(&format!("PRAGMA busy_timeout = {SQLITE_BUSY_TIMEOUT_MS}"))
119 .execute(&mut *conn)
120 .await?;
121 sqlx::query(&format!(
122 "PRAGMA journal_size_limit = {}",
123 config.journal_size_limit_bytes
124 ))
125 .execute(&mut *conn)
126 .await?;
127 sqlx::query(&format!("PRAGMA temp_store = {}", config.temp_store.as_pragma_value()))
128 .execute(&mut *conn)
129 .await?;
130 sqlx::query(&format!("PRAGMA mmap_size = {}", config.mmap_size_bytes))
131 .execute(&mut *conn)
132 .await?;
133 sqlx::query("PRAGMA foreign_keys = ON").execute(&mut *conn).await?;
134 sqlx::query("PRAGMA synchronous = NORMAL").execute(&mut *conn).await?;
135
136 Ok(())
137}
138
139async fn configure_postgres_connection(conn: &mut AnyConnection, config: PostgresConfig) -> DbResult<()> {
140 configure_postgres_timeouts(conn, config.lock_timeout, config.statement_timeout).await?;
141 super::schema::pin_postgres_persistence_schema(conn).await
142}
143
144async fn enable_sqlite_wal(pool: &DbPool) -> DbResult<()> {
145 enable_sqlite_wal_with_retry(pool, SQLITE_WAL_MAX_ATTEMPTS, SQLITE_WAL_RETRY_DELAY).await
146}
147
148async fn enable_sqlite_wal_with_retry(pool: &DbPool, max_attempts: u32, retry_delay: Duration) -> DbResult<()> {
149 debug_assert!(max_attempts > 0);
150
151 for attempt in 1..=max_attempts {
152 match sqlx::query("PRAGMA journal_mode = WAL").execute(pool).await {
153 Ok(_) => return Ok(()),
154 Err(error) if attempt < max_attempts && sqlite_is_busy_or_locked(&error) => {
155 tokio::time::sleep(retry_delay).await;
156 }
157 Err(error) => return Err(error),
158 }
159 }
160
161 unreachable!("WAL retry loop always returns on success or final failure")
162}
163
164fn sqlite_is_busy_or_locked(error: &sqlx::Error) -> bool {
165 let sqlx::Error::Database(database_error) = error else {
166 return false;
167 };
168 let Some(code) = database_error.code().and_then(|code| code.parse::<i32>().ok()) else {
169 return false;
170 };
171
172 matches!(code & 0xff, SQLITE_BUSY | SQLITE_LOCKED)
173}
174
175pub async fn create_pool(db_url: Option<&str>) -> DbResult<Arc<DbPool>> {
198 create_pool_with_sqlite_config(db_url, SqliteConfig::default()).await
199}
200
201pub async fn create_pool_with_sqlite_config(
207 db_url: Option<&str>,
208 sqlite_config: SqliteConfig,
209) -> DbResult<Arc<DbPool>> {
210 create_pool_with_configs(db_url, sqlite_config, PostgresConfig::default()).await
211}
212
213pub async fn create_pool_with_configs(
219 db_url: Option<&str>,
220 sqlite_config: SqliteConfig,
221 postgres_config: PostgresConfig,
222) -> DbResult<Arc<DbPool>> {
223 sqlx::any::install_default_drivers();
225
226 let prepared = prepare_db_url(db_url)?;
228
229 let options = pool_options(&prepared.value, prepared.backend, sqlite_config, postgres_config);
230 let pool = options.connect(&prepared.value).await?;
231 if sqlite_should_enable_wal(&prepared.value, prepared.backend) {
232 enable_sqlite_wal(&pool).await?;
233 }
234
235 Ok(Arc::new(pool))
237}
238
239pub async fn create_pool_with_schema(db_url: Option<&str>) -> DbResult<Arc<DbPool>> {
251 create_pool_with_schema_and_sqlite_config(db_url, SqliteConfig::default()).await
252}
253
254pub async fn create_pool_with_schema_and_sqlite_config(
260 db_url: Option<&str>,
261 sqlite_config: SqliteConfig,
262) -> DbResult<Arc<DbPool>> {
263 create_pool_with_schema_and_configs(db_url, sqlite_config, PostgresConfig::default()).await
264}
265
266pub async fn create_pool_with_schema_and_configs(
272 db_url: Option<&str>,
273 sqlite_config: SqliteConfig,
274 postgres_config: PostgresConfig,
275) -> DbResult<Arc<DbPool>> {
276 use crate::storage::PoolWithSchema;
277
278 let pool = create_pool_with_configs(db_url, sqlite_config, postgres_config).await?;
279 let pool_with_schema = PoolWithSchema::with_postgres_migration_timeout(pool, postgres_config.migration_timeout);
280 pool_with_schema.ensure_schema_ready().await?;
281
282 Ok(pool_with_schema.pool().clone())
283}
284
285#[cfg(test)]
286mod tests {
287 use super::*;
288 use crate::config::{
289 DEFAULT_SQLITE_JOURNAL_SIZE_LIMIT_BYTES, DEFAULT_SQLITE_MAX_CONNECTIONS, DEFAULT_SQLITE_MMAP_SIZE_BYTES,
290 PostgresConfig, SqliteTempStore,
291 };
292 use sqlx::Connection;
293
294 fn prepared_url(url: Option<&str>) -> PreparedDatabaseUrl {
295 prepare_db_url(url).expect("valid database URL")
296 }
297
298 #[test]
299 fn test_prepare_sqlite_url_without_params() {
300 let url = "sqlite://test.db";
301 let prepared = prepared_url(Some(url));
302 assert_eq!(prepared.value, "sqlite://test.db?mode=rwc");
303 assert_eq!(prepared.backend, DatabaseBackend::Sqlite);
304 }
305
306 #[test]
307 fn test_prepare_uppercase_sqlite_url_without_params() {
308 let url = "SQLITE://test.db";
309 let prepared = prepared_url(Some(url));
310 assert_eq!(prepared.value, "SQLITE://test.db?mode=rwc");
311 assert_eq!(prepared.backend, DatabaseBackend::Sqlite);
312 }
313
314 #[test]
315 fn test_prepare_sqlite_url_with_params() {
316 let url = "sqlite://test.db?cache=shared";
317 let prepared = prepared_url(Some(url));
318 assert_eq!(prepared.value, "sqlite://test.db?cache=shared&mode=rwc");
319 }
320
321 #[test]
322 fn test_prepare_sqlite_url_with_fragment() {
323 let url = "sqlite://test.db?cache=shared#frag";
324 let prepared = prepared_url(Some(url));
325 assert_eq!(prepared.value, "sqlite://test.db?cache=shared&mode=rwc#frag");
326 }
327
328 #[test]
329 fn test_prepare_sqlite_url_with_existing_mode() {
330 let url = "sqlite://test.db?mode=ro";
331 let prepared = prepared_url(Some(url));
332 assert_eq!(prepared.value, "sqlite://test.db?mode=ro");
333 }
334
335 #[test]
336 fn test_prepare_sqlite_memory_url_keeps_memory_mode() {
337 let url = "sqlite://?mode=memory";
338 let prepared = prepared_url(Some(url));
339 assert_eq!(prepared.value, "sqlite://?mode=memory");
340 }
341
342 #[test]
343 fn test_prepare_sqlite_memory_shorthand_classifies_backend() {
344 let url = "sqlite::memory:";
345 let prepared = prepared_url(Some(url));
346 assert_eq!(prepared.value, "sqlite::memory:?mode=rwc");
347 assert_eq!(prepared.backend, DatabaseBackend::Sqlite);
348 }
349
350 #[test]
351 fn test_sqlite_wal_enabled_for_writable_file_urls_only() {
352 assert!(sqlite_should_enable_wal(
353 "sqlite://test.db?mode=rwc",
354 DatabaseBackend::Sqlite
355 ));
356 assert!(sqlite_should_enable_wal(
357 "sqlite://test.db?mode=rw",
358 DatabaseBackend::Sqlite
359 ));
360 assert!(!sqlite_should_enable_wal(
361 "sqlite://test.db?mode=ro",
362 DatabaseBackend::Sqlite
363 ));
364 assert!(!sqlite_should_enable_wal(
365 "sqlite://?mode=memory",
366 DatabaseBackend::Sqlite
367 ));
368 assert!(!sqlite_should_enable_wal("sqlite::memory:", DatabaseBackend::Sqlite));
369 assert!(!sqlite_should_enable_wal(
370 "postgresql://localhost/db",
371 DatabaseBackend::Postgres
372 ));
373 }
374
375 #[test]
376 fn test_sqlite_max_connections_for_file_and_memory_urls() {
377 let config = SqliteConfig {
378 max_connections: 6,
379 ..SqliteConfig::default()
380 };
381
382 assert_eq!(sqlite_max_connections("sqlite://test.db?mode=rwc", config), 6);
383 assert_eq!(
384 sqlite_max_connections("sqlite://?mode=memory", config),
385 SQLITE_MEMORY_MAX_CONNECTIONS
386 );
387 assert_eq!(
388 sqlite_max_connections("sqlite::memory:", config),
389 SQLITE_MEMORY_MAX_CONNECTIONS
390 );
391 }
392
393 #[test]
394 fn test_prepare_postgres_url() {
395 let url = "postgresql://user:pass@localhost/db";
396 let prepared = prepared_url(Some(url));
397 assert_eq!(prepared.value, "postgresql://user:pass@localhost/db");
398 assert_eq!(prepared.backend, DatabaseBackend::Postgres);
399 }
400
401 #[test]
402 fn test_postgres_pool_options_use_explicit_config() {
403 let postgres_config = PostgresConfig {
404 max_connections: 7,
405 acquire_timeout: Duration::from_secs(11),
406 lock_timeout: Duration::from_secs(13),
407 migration_timeout: Duration::from_secs(17),
408 statement_timeout: Duration::from_secs(23),
409 idle_timeout: None,
410 max_lifetime: Some(Duration::from_secs(19)),
411 };
412
413 let options = pool_options(
414 "postgresql://user:pass@localhost/db",
415 DatabaseBackend::Postgres,
416 SqliteConfig::default(),
417 postgres_config,
418 );
419
420 assert_eq!(options.get_max_connections(), 7);
421 assert_eq!(options.get_acquire_timeout(), Duration::from_secs(11));
422 assert_eq!(options.get_idle_timeout(), None);
423 assert_eq!(options.get_max_lifetime(), Some(Duration::from_secs(19)));
424 }
425
426 #[test]
427 fn test_uppercase_postgres_pool_options_use_explicit_config() {
428 let postgres_config = PostgresConfig {
429 max_connections: 7,
430 acquire_timeout: Duration::from_secs(11),
431 lock_timeout: Duration::from_secs(13),
432 migration_timeout: Duration::from_secs(17),
433 statement_timeout: Duration::from_secs(23),
434 idle_timeout: None,
435 max_lifetime: Some(Duration::from_secs(19)),
436 };
437
438 let options = pool_options(
439 "POSTGRESQL://user:pass@localhost/db",
440 DatabaseBackend::from_url("POSTGRESQL://user:pass@localhost/db").expect("valid uppercase PostgreSQL URL"),
441 SqliteConfig::default(),
442 postgres_config,
443 );
444
445 assert_eq!(options.get_max_connections(), 7);
446 assert_eq!(options.get_acquire_timeout(), Duration::from_secs(11));
447 }
448
449 #[test]
450 fn test_prepare_mysql_url() {
451 let url = "mysql://user:pass@localhost/db";
452 let prepared = prepared_url(Some(url));
453 assert_eq!(prepared.value, "mysql://user:pass@localhost/db");
454 assert_eq!(prepared.backend, DatabaseBackend::Other);
455 }
456
457 #[test]
458 fn test_prepare_default_sqlite_url() {
459 let prepared = prepared_url(None);
460 assert_eq!(prepared.value, "sqlite://./agentic_api.db?mode=rwc");
461 assert_eq!(prepared.backend, DatabaseBackend::Sqlite);
462 }
463
464 #[tokio::test]
465 async fn test_sqlite_connection_pragmas_are_configured() {
466 let db_path = std::env::temp_dir().join(format!("pragma_{}.db", uuid::Uuid::now_v7()));
467 let db_url = format!("sqlite://{}", db_path.display());
468 let pool = create_pool(Some(&db_url)).await.expect("failed to create pool");
469
470 let journal_mode: String = sqlx::query_scalar("PRAGMA journal_mode")
471 .fetch_one(pool.as_ref())
472 .await
473 .expect("journal_mode query failed");
474 let busy_timeout: i64 = sqlx::query_scalar("PRAGMA busy_timeout")
475 .fetch_one(pool.as_ref())
476 .await
477 .expect("busy_timeout query failed");
478 let foreign_keys: i64 = sqlx::query_scalar("PRAGMA foreign_keys")
479 .fetch_one(pool.as_ref())
480 .await
481 .expect("foreign_keys query failed");
482 let synchronous: i64 = sqlx::query_scalar("PRAGMA synchronous")
483 .fetch_one(pool.as_ref())
484 .await
485 .expect("synchronous query failed");
486 let journal_size_limit: i64 = sqlx::query_scalar("PRAGMA journal_size_limit")
487 .fetch_one(pool.as_ref())
488 .await
489 .expect("journal_size_limit query failed");
490 let temp_store: i64 = sqlx::query_scalar("PRAGMA temp_store")
491 .fetch_one(pool.as_ref())
492 .await
493 .expect("temp_store query failed");
494 let mmap_size: i64 = sqlx::query_scalar("PRAGMA mmap_size")
495 .fetch_one(pool.as_ref())
496 .await
497 .expect("mmap_size query failed");
498
499 assert_eq!(journal_mode.to_lowercase(), "wal");
500 assert_eq!(pool.options().get_max_connections(), DEFAULT_SQLITE_MAX_CONNECTIONS);
501 assert_eq!(
502 busy_timeout,
503 i64::try_from(SQLITE_BUSY_TIMEOUT_MS).expect("default fits in i64")
504 );
505 assert_eq!(foreign_keys, 1);
506 assert_eq!(synchronous, 1);
507 assert_eq!(
508 journal_size_limit,
509 i64::try_from(DEFAULT_SQLITE_JOURNAL_SIZE_LIMIT_BYTES).expect("default fits in i64")
510 );
511 assert_eq!(temp_store, 2);
512 assert_eq!(
513 mmap_size,
514 i64::try_from(DEFAULT_SQLITE_MMAP_SIZE_BYTES).expect("default fits in i64")
515 );
516 }
517
518 #[tokio::test]
519 async fn test_sqlite_connection_pragmas_use_explicit_config() {
520 let db_path = std::env::temp_dir().join(format!("pragma_custom_{}.db", uuid::Uuid::now_v7()));
521 let db_url = format!("sqlite://{}", db_path.display());
522 let config = SqliteConfig {
523 max_connections: 3,
524 journal_size_limit_bytes: 131_072,
525 temp_store: SqliteTempStore::File,
526 mmap_size_bytes: 1_048_576,
527 };
528 let pool = create_pool_with_sqlite_config(Some(&db_url), config)
529 .await
530 .expect("failed to create pool");
531
532 let journal_size_limit: i64 = sqlx::query_scalar("PRAGMA journal_size_limit")
533 .fetch_one(pool.as_ref())
534 .await
535 .expect("journal_size_limit query failed");
536 let temp_store: i64 = sqlx::query_scalar("PRAGMA temp_store")
537 .fetch_one(pool.as_ref())
538 .await
539 .expect("temp_store query failed");
540 let mmap_size: i64 = sqlx::query_scalar("PRAGMA mmap_size")
541 .fetch_one(pool.as_ref())
542 .await
543 .expect("mmap_size query failed");
544
545 assert_eq!(pool.options().get_max_connections(), 3);
546 assert_eq!(journal_size_limit, 131_072);
547 assert_eq!(temp_store, 1);
548 assert_eq!(mmap_size, 1_048_576);
549 }
550
551 #[tokio::test]
552 async fn test_sqlite_wal_preflight_retries_locked_database() {
553 sqlx::any::install_default_drivers();
554
555 let db_path = std::env::temp_dir().join(format!("wal_retry_{}.db", uuid::Uuid::now_v7()));
556 let db_url = format!("sqlite://{}?mode=rwc", db_path.display());
557 let mut lock_conn = AnyConnection::connect(&db_url)
558 .await
559 .expect("failed to open lock connection");
560 sqlx::query("CREATE TABLE locked_write (id INTEGER PRIMARY KEY)")
561 .execute(&mut lock_conn)
562 .await
563 .expect("failed to create test table");
564 sqlx::query("BEGIN EXCLUSIVE")
565 .execute(&mut lock_conn)
566 .await
567 .expect("failed to acquire exclusive lock");
568
569 let pool = AnyPoolOptions::new()
570 .max_connections(1)
571 .after_connect(|conn, _meta| {
572 Box::pin(async move {
573 sqlx::query("PRAGMA busy_timeout = 20").execute(&mut *conn).await?;
574 Ok(())
575 })
576 })
577 .connect(&db_url)
578 .await
579 .expect("failed to create test pool");
580
581 let wal_pool = pool.clone();
582 let wal_task = tokio::spawn(async move {
583 enable_sqlite_wal_with_retry(&wal_pool, SQLITE_WAL_MAX_ATTEMPTS, Duration::from_millis(50)).await
584 });
585
586 tokio::time::sleep(Duration::from_millis(30)).await;
587 sqlx::query("ROLLBACK")
588 .execute(&mut lock_conn)
589 .await
590 .expect("failed to release exclusive lock");
591
592 wal_task
593 .await
594 .expect("WAL preflight task panicked")
595 .expect("WAL preflight should retry after SQLITE_BUSY");
596
597 let journal_mode: String = sqlx::query_scalar("PRAGMA journal_mode")
598 .fetch_one(&pool)
599 .await
600 .expect("journal_mode query failed");
601 assert_eq!(journal_mode.to_lowercase(), "wal");
602 }
603
604 #[tokio::test]
605 async fn test_sqlite_memory_pool_stays_single_connection() {
606 let config = SqliteConfig {
607 max_connections: 3,
608 ..SqliteConfig::default()
609 };
610 let pool = create_pool_with_sqlite_config(Some("sqlite://?mode=memory"), config)
611 .await
612 .expect("failed to create pool");
613
614 assert_eq!(pool.options().get_max_connections(), SQLITE_MEMORY_MAX_CONNECTIONS);
615 }
616}