use derive_builder::Builder;
use regex::Regex;
use serde::{Deserialize, Serialize};
use sqlx::{
AssertSqlSafe, ConnectOptions, PgPool,
postgres::{PgConnectOptions, PgConnection},
};
fn validate_sql_identifier(value: &str, label: &str) -> anyhow::Result<()> {
if value.is_empty() {
anyhow::bail!("{label} must not be empty");
}
if !value.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') {
anyhow::bail!(
"{label} contains invalid characters (only alphanumeric and underscore allowed): {value}"
);
}
Ok(())
}
fn escape_sql_string(value: &str) -> String {
value.replace('\'', "''")
}
#[derive(Debug, Clone, Serialize, Deserialize, Builder)]
#[serde(deny_unknown_fields)]
#[builder(default)]
#[cfg_attr(
feature = "python",
pyo3::pyclass(
module = "nautilus_trader.core.nautilus_pyo3.infrastructure",
from_py_object
)
)]
#[cfg_attr(
feature = "python",
pyo3_stub_gen::derive::gen_stub_pyclass(module = "nautilus_trader.infrastructure")
)]
#[allow(
clippy::unsafe_derive_deserialize,
reason = "config type deserializes plain field values; unsafe PyO3 methods are unrelated"
)]
pub struct PostgresConnectOptions {
pub host: String,
pub port: u16,
pub username: String,
pub password: String,
pub database: String,
}
impl PostgresConnectOptions {
#[must_use]
pub const fn new(
host: String,
port: u16,
username: String,
password: String,
database: String,
) -> Self {
Self {
host,
port,
username,
password,
database,
}
}
#[must_use]
pub fn connection_string(&self) -> String {
format!(
"postgres://{username}:{password}@{host}:{port}/{database}",
username = self.username,
password = self.password,
host = self.host,
port = self.port,
database = self.database
)
}
#[must_use]
pub fn connection_string_masked(&self) -> String {
format!(
"postgres://{username}:***@{host}:{port}/{database}",
username = self.username,
host = self.host,
port = self.port,
database = self.database
)
}
#[must_use]
pub fn default_administrator() -> Self {
Self::new(
String::from("localhost"),
5432,
String::from("nautilus"),
String::from("pass"),
String::from("nautilus"),
)
}
}
impl Default for PostgresConnectOptions {
fn default() -> Self {
Self::new(
String::from("localhost"),
5432,
String::from("nautilus"),
String::from("pass"),
String::from("nautilus"),
)
}
}
impl From<PostgresConnectOptions> for PgConnectOptions {
fn from(opt: PostgresConnectOptions) -> Self {
Self::new()
.host(opt.host.as_str())
.port(opt.port)
.username(opt.username.as_str())
.password(opt.password.as_str())
.database(opt.database.as_str())
.disable_statement_logging()
}
}
#[must_use]
pub fn get_postgres_connect_options(
host: Option<String>,
port: Option<u16>,
username: Option<String>,
password: Option<String>,
database: Option<String>,
) -> PostgresConnectOptions {
let defaults = PostgresConnectOptions::default_administrator();
let host = host
.or_else(|| std::env::var("POSTGRES_HOST").ok())
.unwrap_or(defaults.host);
let port = port
.or_else(|| {
std::env::var("POSTGRES_PORT")
.map(|port| port.parse::<u16>().unwrap())
.ok()
})
.unwrap_or(defaults.port);
let username = username
.or_else(|| std::env::var("POSTGRES_USERNAME").ok())
.unwrap_or(defaults.username);
let database = database
.or_else(|| std::env::var("POSTGRES_DATABASE").ok())
.unwrap_or(defaults.database);
let password = password
.or_else(|| std::env::var("POSTGRES_PASSWORD").ok())
.unwrap_or(defaults.password);
PostgresConnectOptions::new(host, port, username, password, database)
}
pub async fn connect_pg(options: PgConnectOptions) -> anyhow::Result<PgPool> {
Ok(PgPool::connect_with(options).await?)
}
fn get_schema_dir() -> anyhow::Result<String> {
std::env::var("SCHEMA_DIR").or_else(|_| {
let nautilus_git_repo_name = "nautilus_trader";
let binding = std::env::current_dir().unwrap();
let current_dir = binding.to_str().unwrap();
match current_dir.find(nautilus_git_repo_name){
Some(index) => {
let schema_path = current_dir[0..index + nautilus_git_repo_name.len()].to_string() + "/schema/sql";
Ok(schema_path)
}
None => anyhow::bail!("Could not calculate schema dir from current directory path or SCHEMA_DIR env variable")
}
})
}
pub async fn init_postgres(
pg: &PgPool,
database: String,
password: String,
schema_dir: Option<String>,
) -> anyhow::Result<()> {
log::info!("Initializing Postgres database with target permissions and schema");
validate_sql_identifier(&database, "database")?;
let mut connection = pg.acquire().await?;
match sqlx::query("CREATE SCHEMA IF NOT EXISTS public;")
.execute(&mut *connection)
.await
{
Ok(_) => log::info!("Schema public created successfully"),
Err(e) => log::error!("Error creating schema public: {e:?}"),
}
let escaped_password = escape_sql_string(&password);
match sqlx::query(AssertSqlSafe(format!(
"CREATE ROLE {database} PASSWORD '{escaped_password}' LOGIN;"
)))
.execute(&mut *connection)
.await
{
Ok(_) => log::info!("Role {database} created successfully"),
Err(e) => {
if e.to_string().contains("already exists") {
log::info!("Role {database} already exists");
} else {
log::error!("Error creating role {database}: {e:?}");
}
}
}
let schema_dir = schema_dir.unwrap_or_else(|| get_schema_dir().unwrap());
assign_schema_ownership(&mut connection, &database).await?;
sqlx::query(AssertSqlSafe(format!(
"ALTER DATABASE {database} OWNER TO {database};"
)))
.execute(&mut *connection)
.await?;
sqlx::query(AssertSqlSafe(format!(
"ALTER SCHEMA public OWNER TO {database};"
)))
.execute(&mut *connection)
.await?;
execute_schema_as_role(&mut connection, &database, &schema_dir).await?;
match sqlx::query(AssertSqlSafe(format!(
"GRANT CONNECT ON DATABASE {database} TO {database};"
)))
.execute(&mut *connection)
.await
{
Ok(_) => log::info!("Connect privileges granted to role {database}"),
Err(e) => log::error!("Error granting connect privileges to role {database}: {e:?}"),
}
match sqlx::query(AssertSqlSafe(format!(
"GRANT ALL PRIVILEGES ON SCHEMA public TO {database};"
)))
.execute(&mut *connection)
.await
{
Ok(_) => log::info!("All schema privileges granted to role {database}"),
Err(e) => log::error!("Error granting all privileges to role {database}: {e:?}"),
}
match sqlx::query(AssertSqlSafe(format!(
"GRANT ALL PRIVILEGES ON ALL TABLES IN SCHEMA public TO {database};"
)))
.execute(&mut *connection)
.await
{
Ok(_) => log::info!("All tables privileges granted to role {database}"),
Err(e) => log::error!("Error granting all privileges to role {database}: {e:?}"),
}
match sqlx::query(AssertSqlSafe(format!(
"GRANT ALL PRIVILEGES ON ALL SEQUENCES IN SCHEMA public TO {database};"
)))
.execute(&mut *connection)
.await
{
Ok(_) => log::info!("All sequences privileges granted to role {database}"),
Err(e) => log::error!("Error granting all privileges to role {database}: {e:?}"),
}
match sqlx::query(AssertSqlSafe(format!(
"GRANT EXECUTE ON ALL FUNCTIONS IN SCHEMA public TO {database};"
)))
.execute(&mut *connection)
.await
{
Ok(_) => log::info!("All functions privileges granted to role {database}"),
Err(e) => log::error!("Error granting all privileges to role {database}: {e:?}"),
}
Ok(())
}
#[expect(
clippy::too_many_lines,
reason = "The catalog query stays intact as one ownership migration boundary"
)]
async fn assign_schema_ownership(
connection: &mut PgConnection,
database: &str,
) -> anyhow::Result<()> {
let statements: Vec<String> = sqlx::query_scalar(
"
SELECT statement
FROM (
SELECT
1 AS object_order,
CASE c.relkind
WHEN 'S' THEN format(
'ALTER SEQUENCE %I.%I OWNER TO %I',
n.nspname,
c.relname,
$1
)
WHEN 'v' THEN format(
'ALTER VIEW %I.%I OWNER TO %I',
n.nspname,
c.relname,
$1
)
WHEN 'm' THEN format(
'ALTER MATERIALIZED VIEW %I.%I OWNER TO %I',
n.nspname,
c.relname,
$1
)
WHEN 'f' THEN format(
'ALTER FOREIGN TABLE %I.%I OWNER TO %I',
n.nspname,
c.relname,
$1
)
ELSE format(
'ALTER TABLE %I.%I OWNER TO %I',
n.nspname,
c.relname,
$1
)
END AS statement
FROM pg_class c
JOIN pg_namespace n ON n.oid = c.relnamespace
WHERE n.nspname = 'public'
AND c.relkind IN ('r', 'p', 'v', 'm', 'S', 'f')
AND (
c.relkind <> 'S'
OR NOT EXISTS (
SELECT 1
FROM pg_depend d
WHERE d.classid = 'pg_class'::regclass
AND d.objid = c.oid
AND d.refclassid = 'pg_class'::regclass
AND d.deptype IN ('a', 'i')
)
)
UNION ALL
SELECT
2 AS object_order,
format(
'ALTER %s %I.%I(%s) OWNER TO %I',
CASE p.prokind WHEN 'p' THEN 'PROCEDURE' ELSE 'FUNCTION' END,
n.nspname,
p.proname,
pg_get_function_identity_arguments(p.oid),
$1
) AS statement
FROM pg_proc p
JOIN pg_namespace n ON n.oid = p.pronamespace
WHERE n.nspname = 'public'
AND p.prokind IN ('f', 'p', 'w')
UNION ALL
SELECT
3 AS object_order,
CASE t.typtype
WHEN 'd' THEN format(
'ALTER DOMAIN %I.%I OWNER TO %I',
n.nspname,
t.typname,
$1
)
ELSE format(
'ALTER TYPE %I.%I OWNER TO %I',
n.nspname,
t.typname,
$1
)
END AS statement
FROM pg_type t
JOIN pg_namespace n ON n.oid = t.typnamespace
WHERE n.nspname = 'public'
AND t.typtype IN ('d', 'e')
) objects
ORDER BY object_order, statement
",
)
.bind(database)
.fetch_all(&mut *connection)
.await?;
for statement in statements {
sqlx::query(AssertSqlSafe(statement))
.execute(&mut *connection)
.await?;
}
Ok(())
}
async fn execute_schema_as_role(
connection: &mut PgConnection,
database: &str,
schema_dir: &str,
) -> anyhow::Result<()> {
sqlx::query(AssertSqlSafe(format!("SET ROLE {database};")))
.execute(&mut *connection)
.await?;
let result = async {
let sql_files = ["types.sql", "functions.sql", "partitions.sql", "tables.sql"];
let plpgsql_regex =
Regex::new(r"\$\$ LANGUAGE plpgsql(?:[ \t\r\n]+SECURITY[ \t\r\n]+DEFINER)?;")?;
for file_name in sql_files {
log::info!("Executing schema file: {file_name:?}");
let file_path = format!("{schema_dir}/{file_name}");
let sql_content = std::fs::read_to_string(&file_path)?;
let sql_statements = match file_name {
"functions.sql" | "partitions.sql" => {
let mut statements = Vec::new();
let mut last_end = 0;
for mat in plpgsql_regex.find_iter(&sql_content) {
let statement = sql_content[last_end..mat.end()].to_string();
if !statement.trim().is_empty() {
statements.push(statement);
}
last_end = mat.end();
}
statements
}
_ => split_sql_statements(&sql_content),
};
for sql_statement in sql_statements {
if let Err(e) = sqlx::query(AssertSqlSafe(sql_statement.as_str()))
.execute(&mut *connection)
.await
{
if e.to_string().contains("already exists") {
log::info!("Already exists error on statement, skipping");
} else {
anyhow::bail!(
"Error executing statement {sql_statement} with error: {e:?}"
);
}
}
}
}
Ok(())
}
.await;
let reset_result = sqlx::query("RESET ROLE;").execute(connection).await;
match (result, reset_result) {
(Err(e), Err(reset_error)) => {
log::error!("Error resetting Postgres role after schema failure: {reset_error:?}");
Err(e)
}
(Err(e), Ok(_)) => Err(e),
(Ok(()), Err(e)) => Err(e.into()),
(Ok(()), Ok(_)) => Ok(()),
}
}
fn split_sql_statements(sql: &str) -> Vec<String> {
let mut statements = Vec::new();
let mut current = String::new();
let mut chars = sql.chars().peekable();
let mut in_string = false;
let mut in_dollar_quote = false;
while let Some(c) = chars.next() {
match c {
'\'' if !in_dollar_quote => {
in_string = !in_string;
current.push(c);
}
'$' if !in_string && chars.peek() == Some(&'$') => {
chars.next();
in_dollar_quote = !in_dollar_quote;
current.push_str("$$");
}
'-' if !in_string && !in_dollar_quote && chars.peek() == Some(&'-') => {
for next in chars.by_ref() {
if next == '\n' {
current.push('\n');
break;
}
}
}
';' if !in_string && !in_dollar_quote => {
let trimmed = current.trim();
if !trimmed.is_empty() {
statements.push(format!("{trimmed};"));
}
current.clear();
}
_ => current.push(c),
}
}
let trimmed = current.trim();
if !trimmed.is_empty() {
statements.push(format!("{trimmed};"));
}
statements
}
pub async fn drop_postgres(pg: &PgPool, database: String) -> anyhow::Result<()> {
validate_sql_identifier(&database, "database")?;
sqlx::query(AssertSqlSafe(format!(
"ALTER DATABASE {database} OWNER TO SESSION_USER"
)))
.execute(pg)
.await?;
match sqlx::query(AssertSqlSafe(format!("DROP OWNED BY {database}")))
.execute(pg)
.await
{
Ok(_) => log::info!("Dropped owned objects by role {database}"),
Err(e) => {
let err_msg = e.to_string();
if err_msg.contains("2BP01") || err_msg.contains("required by the database system") {
log::warn!("Skipping system-required objects for role {database}");
} else {
log::error!("Error dropping owned by role {database}: {e:?}");
}
}
}
match sqlx::query(AssertSqlSafe(format!(
"REVOKE CONNECT ON DATABASE {database} FROM {database};"
)))
.execute(pg)
.await
{
Ok(_) => log::info!("Revoked connect privileges from role {database}"),
Err(e) => log::error!("Error revoking connect privileges from role {database}: {e:?}"),
}
match sqlx::query(AssertSqlSafe(format!(
"REVOKE ALL PRIVILEGES ON DATABASE {database} FROM {database};"
)))
.execute(pg)
.await
{
Ok(_) => log::info!("Revoked all privileges from role {database}"),
Err(e) => log::error!("Error revoking all privileges from role {database}: {e:?}"),
}
match sqlx::query("DROP SCHEMA IF EXISTS public CASCADE")
.execute(pg)
.await
{
Ok(_) => log::info!("Dropped schema public"),
Err(e) => log::error!("Error dropping schema public: {e:?}"),
}
match sqlx::query(AssertSqlSafe(format!("DROP ROLE IF EXISTS {database};")))
.execute(pg)
.await
{
Ok(_) => log::info!("Dropped role {database}"),
Err(e) => {
let err_msg = e.to_string();
if err_msg.contains("55006") || err_msg.contains("current user cannot be dropped") {
log::warn!("Cannot drop currently connected role {database}");
} else {
anyhow::bail!("Error dropping role {database}: {e:?}");
}
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use super::*;
#[rstest]
fn test_postgres_connect_options_toml_round_trip() {
let config: PostgresConnectOptions = toml::from_str(
r#"
host = "localhost"
port = 5432
username = "nautilus"
password = "secret"
database = "nautilus"
"#,
)
.unwrap();
assert_eq!(config.host, "localhost");
assert_eq!(config.port, 5432);
assert_eq!(config.username, "nautilus");
assert_eq!(config.database, "nautilus");
}
#[rstest]
fn test_split_sql_statements_basic() {
let sql = "CREATE TABLE a (id INT); CREATE TABLE b (id INT);";
assert_eq!(
split_sql_statements(sql),
vec!["CREATE TABLE a (id INT);", "CREATE TABLE b (id INT);"]
);
}
#[rstest]
fn test_split_sql_statements_ignores_semicolon_in_line_comment() {
let sql = "\
-- start points; a later run re-validates them.
ALTER TABLE pool_snapshot ADD COLUMN IF NOT EXISTS validation_state TEXT;";
assert_eq!(
split_sql_statements(sql),
vec!["ALTER TABLE pool_snapshot ADD COLUMN IF NOT EXISTS validation_state TEXT;"]
);
}
#[rstest]
fn test_split_sql_statements_keeps_code_before_trailing_comment() {
let sql = "CREATE TABLE a (\n id INT, -- REFERENCES x;\n name TEXT\n);";
assert_eq!(
split_sql_statements(sql),
vec!["CREATE TABLE a (\n id INT, \n name TEXT\n);"]
);
}
#[rstest]
fn test_split_sql_statements_keeps_dollar_quoted_body_intact() {
let sql = "\
DO $$
BEGIN
IF EXISTS (SELECT 1 FROM information_schema.columns WHERE column_name = 'avg_px') THEN
ALTER TABLE \"order\" ALTER COLUMN avg_px TYPE NUMERIC;
END IF;
END $$;
SELECT 1;";
let statements = split_sql_statements(sql);
assert_eq!(statements.len(), 2);
assert!(statements[0].starts_with("DO $$"));
assert!(statements[0].ends_with("END $$;"));
assert!(statements[0].contains("ALTER COLUMN avg_px TYPE NUMERIC;"));
assert_eq!(statements[1], "SELECT 1;");
}
#[rstest]
fn test_split_sql_statements_ignores_semicolon_in_string_literal() {
let sql = "INSERT INTO t VALUES ('a;b'); SELECT 1;";
assert_eq!(
split_sql_statements(sql),
vec!["INSERT INTO t VALUES ('a;b');", "SELECT 1;"]
);
}
#[rstest]
fn test_split_sql_statements_drops_comment_only_lines() {
let sql =
"------------------- ENUMS -------------------\nCREATE TYPE x AS ENUM ('A', 'B');";
assert_eq!(
split_sql_statements(sql),
vec!["CREATE TYPE x AS ENUM ('A', 'B');"]
);
}
}