1use std::{collections::HashSet, fmt, time::Duration};
11
12use anyhow::anyhow;
13use clap::Args;
14use diesel::{
15 PgConnection, QueryableByName,
16 connection::BoxableConnection,
17 query_dsl::RunQueryDsl,
18 r2d2::{ConnectionManager, Pool, PooledConnection, R2D2Connection},
19};
20use serde::{Deserialize, Serialize};
21use strum::IntoEnumIterator;
22use tracing::{error, info};
23use url::Url;
24
25use crate::{errors::IndexerError, pruning::pruner::PrunableTable};
26
27pub type ConnectionPool = Pool<ConnectionManager<PgConnection>>;
28pub type PoolConnection = PooledConnection<ConnectionManager<PgConnection>>;
29
30#[derive(Args, Debug, Clone)]
31pub struct ConnectionPoolConfig {
32 #[arg(long, default_value_t = 100)]
33 #[arg(env = "DB_POOL_SIZE")]
34 pub pool_size: u32,
35 #[arg(long, value_parser = parse_duration, default_value = "30")]
36 #[arg(env = "DB_CONNECTION_TIMEOUT")]
37 pub connection_timeout: Duration,
38 #[arg(long, value_parser = parse_duration, default_value = "3600")]
39 #[arg(env = "DB_STATEMENT_TIMEOUT")]
40 pub statement_timeout: Duration,
41}
42
43fn parse_duration(arg: &str) -> Result<std::time::Duration, std::num::ParseIntError> {
44 let seconds = arg.parse()?;
45 Ok(std::time::Duration::from_secs(seconds))
46}
47
48impl ConnectionPoolConfig {
49 pub const DEFAULT_POOL_SIZE: u32 = 100;
50 pub const DEFAULT_CONNECTION_TIMEOUT: u64 = 30;
51 pub const DEFAULT_STATEMENT_TIMEOUT: u64 = 3600;
52
53 fn connection_config(&self) -> ConnectionConfig {
54 ConnectionConfig {
55 statement_timeout: self.statement_timeout,
56 read_only: false,
57 }
58 }
59
60 pub fn set_pool_size(&mut self, size: u32) {
61 self.pool_size = size;
62 }
63
64 pub fn set_connection_timeout(&mut self, timeout: Duration) {
65 self.connection_timeout = timeout;
66 }
67
68 pub fn set_statement_timeout(&mut self, timeout: Duration) {
69 self.statement_timeout = timeout;
70 }
71}
72
73impl Default for ConnectionPoolConfig {
74 fn default() -> Self {
75 Self {
76 pool_size: Self::DEFAULT_POOL_SIZE,
77 connection_timeout: Duration::from_secs(Self::DEFAULT_CONNECTION_TIMEOUT),
78 statement_timeout: Duration::from_secs(Self::DEFAULT_STATEMENT_TIMEOUT),
79 }
80 }
81}
82
83#[derive(Debug, Clone, Copy)]
84pub struct ConnectionConfig {
85 pub statement_timeout: Duration,
86 pub read_only: bool,
87}
88
89impl<T: R2D2Connection + 'static> diesel::r2d2::CustomizeConnection<T, diesel::r2d2::Error>
90 for ConnectionConfig
91{
92 fn on_acquire(&self, _conn: &mut T) -> std::result::Result<(), diesel::r2d2::Error> {
93 _conn
94 .as_any_mut()
95 .downcast_mut::<diesel::PgConnection>()
96 .map_or_else(
97 || {
98 Err(diesel::r2d2::Error::QueryError(
99 diesel::result::Error::DeserializationError(
100 "failed to downcast connection to PgConnection"
101 .to_string()
102 .into(),
103 ),
104 ))
105 },
106 |pg_conn| {
107 diesel::sql_query(format!(
108 "SET statement_timeout = {}",
109 self.statement_timeout.as_millis(),
110 ))
111 .execute(pg_conn)
112 .map_err(diesel::r2d2::Error::QueryError)?;
113
114 if self.read_only {
115 diesel::sql_query("SET default_transaction_read_only = 't'")
116 .execute(pg_conn)
117 .map_err(diesel::r2d2::Error::QueryError)?;
118 }
119 Ok(())
120 },
121 )?;
122 Ok(())
123 }
124}
125
126#[derive(Serialize, Deserialize, Clone, Eq, PartialEq)]
131#[serde(transparent)]
132pub struct DbUrl(String);
133
134impl DbUrl {
135 pub fn as_str(&self) -> &str {
136 &self.0
137 }
138
139 fn redacted(&self) -> String {
142 const HIDDEN: &str = "****";
143
144 let Ok(mut url) = Url::parse(&self.0) else {
145 return HIDDEN.to_string();
146 };
147
148 let password = url.password().map(|_| HIDDEN);
149 if url.set_password(password).is_err() {
150 return HIDDEN.to_string();
151 }
152
153 url.to_string()
154 }
155}
156
157impl From<String> for DbUrl {
158 fn from(url: String) -> Self {
159 Self(url)
160 }
161}
162
163impl From<&str> for DbUrl {
164 fn from(url: &str) -> Self {
165 Self(url.to_string())
166 }
167}
168
169impl fmt::Debug for DbUrl {
170 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
171 f.debug_tuple("DbUrl").field(&self.redacted()).finish()
172 }
173}
174
175pub fn new_connection_pool(
176 db_url: &DbUrl,
177 config: &ConnectionPoolConfig,
178) -> Result<ConnectionPool, IndexerError> {
179 let manager = ConnectionManager::<PgConnection>::new(db_url.as_str());
180
181 Pool::builder()
182 .max_size(config.pool_size)
183 .connection_timeout(config.connection_timeout)
184 .connection_customizer(Box::new(config.connection_config()))
185 .build(manager)
186 .map_err(|e| {
187 error!("failed to initialize connection pool: {e:?}");
188 IndexerError::PgConnectionPoolInit
189 })
190}
191
192pub fn get_pool_connection(pool: &ConnectionPool) -> Result<PoolConnection, IndexerError> {
193 pool.get().map_err(|e| {
194 error!("failed to get connection from PG connection pool: {e:?}");
195 IndexerError::PgPoolConnection
196 })
197}
198
199pub fn reset_database(conn: &mut PoolConnection) -> Result<(), anyhow::Error> {
200 {
201 conn.as_any_mut()
202 .downcast_mut::<PoolConnection>()
203 .map_or_else(
204 || Err(anyhow!("failed to downcast connection to PgConnection")),
205 |pg_conn| {
206 setup_postgres::reset_database(pg_conn)?;
207 Ok(())
208 },
209 )?;
210 }
211 Ok(())
212}
213
214pub async fn check_prunable_tables_valid(conn: &mut PoolConnection) -> Result<(), IndexerError> {
216 info!("Starting compatibility check");
217
218 use diesel::RunQueryDsl;
219
220 let select_parent_tables = r#"
221 SELECT c.relname AS table_name
222 FROM pg_class c
223 JOIN pg_namespace n ON n.oid = c.relnamespace
224 LEFT JOIN pg_partitioned_table pt ON pt.partrelid = c.oid
225 WHERE c.relkind IN ('r', 'p') -- 'r' for regular tables, 'p' for partitioned tables
226 AND n.nspname = 'public'
227 AND (
228 pt.partrelid IS NOT NULL -- This is a partitioned (parent) table
229 OR NOT EXISTS ( -- This is not a partition (child table)
230 SELECT 1
231 FROM pg_inherits i
232 WHERE i.inhrelid = c.oid
233 )
234 );
235 "#;
236
237 #[derive(QueryableByName)]
238 struct TableName {
239 #[diesel(sql_type = diesel::sql_types::Text)]
240 table_name: String,
241 }
242
243 let result: Vec<TableName> = diesel::sql_query(select_parent_tables)
244 .load(conn)
245 .map_err(|e| IndexerError::DbMigration(format!("failed to fetch tables: {e}")))?;
246
247 let parent_tables_from_db: HashSet<_> = result.into_iter().map(|t| t.table_name).collect();
248
249 for key in PrunableTable::iter() {
250 if !parent_tables_from_db.contains(key.as_ref()) {
251 return Err(IndexerError::Generic(format!(
252 "invalid retention policy override provided for table {key}: does not exist in the database",
253 )));
254 }
255 }
256
257 info!("Compatibility check passed");
258 Ok(())
259}
260
261pub mod setup_postgres {
262 use anyhow::anyhow;
263 use diesel::{
264 RunQueryDsl,
265 migration::{Migration, MigrationConnection, MigrationSource, MigrationVersion},
266 pg::Pg,
267 prelude::*,
268 };
269 use diesel_migrations::{EmbeddedMigrations, MigrationHarness, embed_migrations};
270 use tracing::info;
271
272 use crate::{IndexerError, db::PoolConnection};
273
274 table! {
275 __diesel_schema_migrations (version) {
276 version -> VarChar,
277 run_on -> Timestamp,
278 }
279 }
280
281 const MIGRATIONS: EmbeddedMigrations = embed_migrations!("migrations/pg");
282
283 pub fn reset_database(conn: &mut PoolConnection) -> Result<(), anyhow::Error> {
284 info!("Resetting PG database ...");
285
286 let drop_all_tables = "
287 DO $$ DECLARE
288 r RECORD;
289 BEGIN
290 FOR r IN (SELECT tablename FROM pg_tables WHERE schemaname = 'public')
291 LOOP
292 EXECUTE 'DROP TABLE IF EXISTS ' || quote_ident(r.tablename) || ' CASCADE';
293 END LOOP;
294 END $$;";
295 diesel::sql_query(drop_all_tables).execute(conn)?;
296 info!("Dropped all tables.");
297
298 let drop_all_procedures = "
299 DO $$ DECLARE
300 r RECORD;
301 BEGIN
302 FOR r IN (SELECT proname, oidvectortypes(proargtypes) as argtypes
303 FROM pg_proc INNER JOIN pg_namespace ns ON (pg_proc.pronamespace = ns.oid)
304 WHERE ns.nspname = 'public' AND prokind = 'p')
305 LOOP
306 EXECUTE 'DROP PROCEDURE IF EXISTS ' || quote_ident(r.proname) || '(' || r.argtypes || ') CASCADE';
307 END LOOP;
308 END $$;";
309 diesel::sql_query(drop_all_procedures).execute(conn)?;
310 info!("Dropped all procedures.");
311
312 let drop_all_functions = "
313 DO $$ DECLARE
314 r RECORD;
315 BEGIN
316 FOR r IN (SELECT proname, oidvectortypes(proargtypes) as argtypes
317 FROM pg_proc INNER JOIN pg_namespace ON (pg_proc.pronamespace = pg_namespace.oid)
318 WHERE pg_namespace.nspname = 'public' AND prokind = 'f')
319 LOOP
320 EXECUTE 'DROP FUNCTION IF EXISTS ' || quote_ident(r.proname) || '(' || r.argtypes || ') CASCADE';
321 END LOOP;
322 END $$;";
323 diesel::sql_query(drop_all_functions).execute(conn)?;
324 info!("Dropped all functions.");
325
326 conn.setup()?;
327 info!("Created __diesel_schema_migrations table.");
328
329 run_migrations(conn)?;
330 info!("Reset database complete.");
331 Ok(())
332 }
333
334 pub fn run_migrations(conn: &mut PoolConnection) -> Result<(), anyhow::Error> {
336 let pending_migrations = conn
337 .pending_migrations(MIGRATIONS)
338 .map_err(|e| anyhow!("failed to identify pending migrations {e}"))?;
339 for migration in pending_migrations {
340 info!("Applying migration {}", migration.name());
341 conn.run_migration(&migration)
342 .map_err(|e| anyhow!("failed to run migration {e}"))?;
343 }
344 Ok(())
345 }
346
347 pub fn check_db_migration_consistency(conn: &mut PoolConnection) -> Result<(), IndexerError> {
368 info!("Starting compatibility check");
369 let migrations: Vec<Box<dyn Migration<Pg>>> = MIGRATIONS.migrations().map_err(|err| {
370 IndexerError::DbMigration(format!(
371 "failed to fetch local migrations from schema: {err}"
372 ))
373 })?;
374
375 let local_migrations = migrations
376 .iter()
377 .map(|m| m.name().version())
378 .collect::<Vec<_>>();
379
380 check_db_migration_consistency_impl(conn, local_migrations)?;
381 info!("Compatibility check passed");
382 Ok(())
383 }
384
385 fn check_db_migration_consistency_impl(
386 conn: &mut PoolConnection,
387 local_migrations: Vec<MigrationVersion>,
388 ) -> Result<(), IndexerError> {
389 let applied_migrations: Vec<MigrationVersion> = __diesel_schema_migrations::table
394 .select(__diesel_schema_migrations::version)
395 .order(__diesel_schema_migrations::version.asc())
396 .load(conn)?;
397
398 if local_migrations.len() > applied_migrations.len() {
400 return Err(IndexerError::DbMigration(format!(
401 "the number of local migrations is greater than the number of applied migrations. Local migrations: {local_migrations:?}, Applied migrations: {applied_migrations:?}",
402 )));
403 }
404 for (local_migration, applied_migration) in local_migrations.iter().zip(&applied_migrations)
405 {
406 if local_migration != applied_migration {
407 return Err(IndexerError::DbMigration(format!(
408 "the next applied migration `{applied_migration:?}` diverges from the local migration `{local_migration:?}`",
409 )));
410 }
411 }
412 Ok(())
413 }
414
415 #[cfg(feature = "pg_integration")]
416 #[cfg(test)]
417 mod tests {
418 use diesel::{
419 migration::{Migration, MigrationSource},
420 pg::Pg,
421 };
422 use diesel_migrations::MigrationHarness;
423
424 use crate::{
425 db::setup_postgres::{self, MIGRATIONS},
426 test_utils::{TestDatabase, db_url},
427 };
428
429 #[test]
432 fn db_migration_consistency_smoke_test() {
433 let mut database = TestDatabase::new(db_url("db_migration_consistency_smoke_test"));
434 database.recreate();
435 database.reset_db();
436 {
437 let pool = database.to_connection_pool();
438 let mut conn = pool.get().unwrap();
439 setup_postgres::check_db_migration_consistency(&mut conn).unwrap();
440 }
441 database.drop_if_exists();
442 }
443
444 #[test]
445 fn db_migration_consistency_non_prefix_test() {
446 let mut database =
447 TestDatabase::new(db_url("db_migration_consistency_non_prefix_test"));
448 database.recreate();
449 database.reset_db();
450 {
451 let pool = database.to_connection_pool();
452 let mut conn = pool.get().unwrap();
453 conn.revert_migration(MIGRATIONS.migrations().unwrap().last().unwrap())
454 .unwrap();
455 assert!(setup_postgres::check_db_migration_consistency(&mut conn).is_err());
458
459 conn.run_pending_migrations(MIGRATIONS).unwrap();
460 setup_postgres::check_db_migration_consistency(&mut conn).unwrap();
462 }
463 database.drop_if_exists();
464 }
465
466 #[test]
467 fn db_migration_consistency_prefix_test() {
468 let mut database = TestDatabase::new(db_url("db_migration_consistency_prefix_test"));
469 database.recreate();
470 database.reset_db();
471 {
472 let pool = database.to_connection_pool();
473 let mut conn = pool.get().unwrap();
474
475 let migrations: Vec<Box<dyn Migration<Pg>>> = MIGRATIONS.migrations().unwrap();
476 let mut local_migrations: Vec<_> =
477 migrations.iter().map(|m| m.name().version()).collect();
478 local_migrations.pop();
479 setup_postgres::check_db_migration_consistency_impl(&mut conn, local_migrations)
482 .unwrap();
483 }
484 database.drop_if_exists();
485 }
486 }
487}
488
489#[cfg(test)]
490mod tests {
491 use super::*;
492
493 #[test]
494 fn test_db_url_debug_hides_password() {
495 let cases = [
496 (
497 "postgres://user:hunter2@localhost:5432/iota_indexer",
498 "postgres://user:****@localhost:5432/iota_indexer",
499 ),
500 (
501 "postgres://user@localhost:5432/iota_indexer",
502 "postgres://user@localhost:5432/iota_indexer",
503 ),
504 (
505 "postgres://localhost:5432/iota_indexer",
506 "postgres://localhost:5432/iota_indexer",
507 ),
508 ("user:hunter2@localhost", "****"),
510 ("host=localhost password=hunter2", "****"),
511 ];
512
513 for (url, expect) in cases {
514 let db_url = DbUrl::from(url);
515
516 assert_eq!(format!("{db_url:?}"), format!(r#"DbUrl("{expect}")"#));
517 assert_eq!(db_url.as_str(), url);
518 }
519 }
520}