Skip to main content

iota_indexer/
db.rs

1// Copyright (c) Mysten Labs, Inc.
2// Modifications Copyright (c) 2024 IOTA Stiftung
3// SPDX-License-Identifier: Apache-2.0
4
5//! Types and logic to setup and maintain the database.
6//!
7//! Creating connections, applying or validating migrations are examples of
8//! operations included in this scope.
9
10use 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/// A database connection URL, which can contain a password.
127///
128/// Its `Debug` impl hides the password, so that printing a config that
129/// contains it does not write password to the logs.
130#[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    /// The URL with its password replaced by `****`. If it's impossible to
140    /// parse the url to hide only the password then the whole url is hidden.
141    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
214/// Check that prunable tables exist in the database.
215pub 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    /// Execute all unapplied migrations.
335    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    /// Checks that the local migration scripts are a prefix of the records in
348    /// the database. This allows to run migration scripts against a DB at
349    /// any time, without worrying about existing readers failing over.
350    ///
351    /// # Deployment Requirement
352    /// Whenever deploying a new version of either the reader or writer,
353    /// migration scripts **must** be run first. This ensures that there are
354    /// never more local migration scripts than those recorded in the database.
355    ///
356    /// # Backward Compatibility
357    /// All new migrations must be **backward compatible** with the previous
358    /// data model. Do **not** remove or rename columns, tables, or change types
359    /// in a way that would break older versions of the reader or writer.
360    ///
361    /// Only after all services are running the new code and no old versions
362    /// are in use, can you safely remove deprecated fields or make breaking
363    /// changes.
364    ///
365    /// This approach supports rolling upgrades and prevents unnecessary
366    /// failures during deployment.
367    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        // Unfortunately we cannot call applied_migrations() directly on the connection,
390        // since it implicitly creates the __diesel_schema_migrations table if it
391        // doesn't exist, which is a write operation that we don't want to do in
392        // this function.
393        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        // We check that the local migrations is a prefix of the applied migrations.
399        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        // Check that the migration records in the database created from the local
430        // schema pass the consistency check.
431        #[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                // Local migrations is one record more than the applied migrations.
456                // This will fail the consistency check since it's not a prefix.
457                assert!(setup_postgres::check_db_migration_consistency(&mut conn).is_err());
458
459                conn.run_pending_migrations(MIGRATIONS).unwrap();
460                // After running pending migrations they should be consistent.
461                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                // Local migrations is one record less than the applied migrations.
480                // This should pass the consistency check since it's still a prefix.
481                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            // url fails to parse, password can be anywhere
509            ("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}