Skip to main content

systemprompt_database/services/postgres/
connection.rs

1//! Initial-connect retry policy for `PostgresProvider`.
2//!
3//! Wraps the first `PgPool` connect in a bounded exponential backoff so
4//! transient startup races (Postgres still booting, SSL handshake racing
5//! the TCP listener) recover without surfacing as user-visible failures.
6//! The retry loop intentionally targets a narrow set of error shapes so
7//! permanent failures (auth, missing database, bad URL) fail fast. The
8//! backoff itself runs on [`crate::resilience::retry::retry_async`].
9//!
10//! Copyright (c) systemprompt.io — Business Source License 1.1.
11//! See <https://systemprompt.io> for licensing details.
12
13use std::future::Future;
14use std::str::FromStr;
15use std::sync::OnceLock;
16use std::sync::atomic::{AtomicU64, Ordering};
17use std::time::{Duration, Instant};
18
19use sqlx::postgres::{PgConnectOptions, PgPool, PgPoolOptions};
20
21use crate::error::DatabaseResult;
22use crate::resilience::classify::Outcome;
23use crate::resilience::config::RetryConfig;
24use crate::resilience::retry::retry_async;
25
26const RETRY_DELAYS_MS: &[u64] = &[100, 200, 400, 800, 1600];
27const MAX_ATTEMPTS: u32 = 5;
28pub const DEFAULT_STATEMENT_CACHE_CAPACITY: usize = 100;
29
30static POOL_CLOCK: OnceLock<Instant> = OnceLock::new();
31static SCHEMA_CHANGED_AT_NANOS: AtomicU64 = AtomicU64::new(0);
32
33fn pool_clock() -> Instant {
34    *POOL_CLOCK.get_or_init(Instant::now)
35}
36
37fn nanos_since_clock(at: Duration) -> u64 {
38    u64::try_from(at.as_nanos()).unwrap_or(u64::MAX).max(1)
39}
40
41pub fn mark_schema_changed() {
42    let now = nanos_since_clock(pool_clock().elapsed());
43    SCHEMA_CHANGED_AT_NANOS.store(now, Ordering::Release);
44}
45
46fn opened_after_schema_change(age: Duration) -> bool {
47    let changed = SCHEMA_CHANGED_AT_NANOS.load(Ordering::Acquire);
48    if changed == 0 {
49        return true;
50    }
51    let opened = pool_clock().elapsed().saturating_sub(age);
52    nanos_since_clock(opened) >= changed
53}
54
55/// Operator-tunable connection-pool sizing for a `PostgresProvider`.
56///
57/// Only the sizing/lifetime knobs an operator needs to fit the pool to their
58/// Postgres `max_connections` and replica count are exposed; the connect, SSL
59/// and retry behaviour is fixed.
60#[derive(Debug, Clone, Copy)]
61pub struct PoolConfig {
62    pub max_connections: u32,
63    pub min_connections: u32,
64    pub acquire_timeout: Duration,
65    pub idle_timeout: Duration,
66    pub max_lifetime: Duration,
67    pub statement_cache_capacity: usize,
68}
69
70impl Default for PoolConfig {
71    fn default() -> Self {
72        Self {
73            max_connections: 50,
74            min_connections: 0,
75            acquire_timeout: Duration::from_secs(30),
76            idle_timeout: Duration::from_mins(5),
77            max_lifetime: Duration::from_mins(30),
78            statement_cache_capacity: DEFAULT_STATEMENT_CACHE_CAPACITY,
79        }
80    }
81}
82
83#[must_use]
84pub fn build_pool_options(cfg: &PoolConfig) -> PgPoolOptions {
85    pool_clock();
86    PgPoolOptions::new()
87        .max_connections(cfg.max_connections)
88        .min_connections(cfg.min_connections)
89        .max_lifetime(cfg.max_lifetime)
90        .acquire_timeout(cfg.acquire_timeout)
91        .idle_timeout(cfg.idle_timeout)
92        // Why: a cached prepared statement fails with SQLSTATE 0A000 ("cached
93        // plan must not change result type") once DDL changes the table under
94        // it. In-process migrations call `mark_schema_changed`, and every
95        // connection opened before that is closed on its next acquire instead
96        // of being handed out with stale plans.
97        .before_acquire(|_conn, meta| {
98            Box::pin(async move { Ok(opened_after_schema_change(meta.age)) })
99        })
100}
101
102pub fn connect_options(database_url: &str) -> DatabaseResult<PgConnectOptions> {
103    let options = PgConnectOptions::from_str(database_url)?
104        .application_name("systemprompt")
105        // Why: sqlx 0.9 `PgConnection::get_or_prepare` Parses every persistent
106        // query as a NAMED statement and only sends Close when the cache is
107        // enabled and evicts, so capacity 0 never deallocates them: each
108        // backend grew without bound until Postgres OOMed. A bounded cache
109        // closes on eviction; stale plans after in-process DDL are handled by
110        // `mark_schema_changed` in `build_pool_options`.
111        .statement_cache_capacity(DEFAULT_STATEMENT_CACHE_CAPACITY)
112        .options([("client_min_messages", "warning")]);
113    Ok(options)
114}
115
116pub async fn connect_with_retry(
117    options: PgPoolOptions,
118    connect_options: PgConnectOptions,
119) -> DatabaseResult<PgPool> {
120    let connector = |opts: PgConnectOptions| {
121        let options = options.clone();
122        async move { options.connect_with(opts).await }
123    };
124    connect_with_retry_using(connect_options, MAX_ATTEMPTS, RETRY_DELAYS_MS, connector).await
125}
126
127pub async fn connect_with_retry_using<T, F, Fut>(
128    connect_options: PgConnectOptions,
129    max_attempts: u32,
130    delays_ms: &[u64],
131    connector: F,
132) -> DatabaseResult<T>
133where
134    T: Send,
135    F: Fn(PgConnectOptions) -> Fut + Send + Sync,
136    Fut: Future<Output = Result<T, sqlx::Error>> + Send,
137{
138    let cfg = RetryConfig {
139        max_attempts,
140        base_delay: Duration::from_millis(delays_ms.first().copied().unwrap_or(100)),
141        max_delay: Duration::from_millis(delays_ms.iter().copied().max().unwrap_or(1600)),
142        jitter: false,
143    };
144    let classify = |err: &sqlx::Error| {
145        if is_retryable(err) {
146            Outcome::Transient { retry_after: None }
147        } else {
148            Outcome::Permanent
149        }
150    };
151    retry_async(&cfg, "postgres-connect", classify, || {
152        connector(connect_options.clone())
153    })
154    .await
155    .map_err(Into::into)
156}
157
158fn is_retryable(err: &sqlx::Error) -> bool {
159    if let sqlx::Error::Io(io_err) = err
160        && io_err.kind() == std::io::ErrorKind::ConnectionRefused
161    {
162        return true;
163    }
164    let msg = err.to_string();
165    msg.contains("unexpected response from SSLRequest") || msg.contains("starting up")
166}