Files
truf-server/app/db_backend.py
T
2026-09-30 20:30:56 +03:00

837 lines
35 KiB
Python

import os
import re
import sqlite3
import time
from urllib.parse import parse_qsl, quote, unquote, urlencode, urlsplit, urlunsplit
POSTGRES_SCHEMES = ('postgresql://', 'postgres://')
DEFAULT_POSTGRES_CONNECT_TIMEOUT_SEC = 10
DEFAULT_POSTGRES_STATEMENT_TIMEOUT_MS = 30000
DEFAULT_POSTGRES_LOCK_TIMEOUT_MS = 10000
DEFAULT_POSTGRES_IDLE_TRANSACTION_TIMEOUT_MS = 30000
DEFAULT_POSTGRES_TCP_USER_TIMEOUT_MS = 30000
POSTGRES_APPLICATION_SCHEMA = 'public'
POSTGRES_CHILD_START_RETRY_ATTEMPTS = 3
POSTGRES_CHILD_START_RETRY_MARKERS = (
'server closed the connection unexpectedly',
'connection reset by peer',
'connection was forcibly closed by the remote host',
)
HOST_AGENT_POSTGRES_SOCKET_DIRECTORY = '/run/truf-postgres'
HOST_AGENT_POSTGRES_DATABASE = 'truf'
HOST_AGENT_POSTGRES_USER = 'truf'
HOST_AGENT_POSTGRES_PORT = 5432
class DatabaseUrlError(ValueError):
pass
def is_postgres_url(value):
return str(value or '').strip().lower().startswith(POSTGRES_SCHEMES)
def database_url_from_env():
# Supervised processes receive TRUF_MANAGED_POSTGRES_DSN as the canonical
# connection authority. Use that same precedence everywhere that derives
# an endpoint identity or opens a managed connection.
return (
os.getenv('TRUF_MANAGED_POSTGRES_DSN')
or os.getenv('SCANNER_DB_URL')
or os.getenv('DATABASE_URL')
)
def parse_postgres_url(value):
"""Parse one URL-form libpq DSN without allowing alternate authorities."""
text = str(value or '').strip()
if not text or any(character in text for character in ('\x00', '\r', '\n')):
raise DatabaseUrlError('invalid PostgreSQL database URL')
try:
parsed = urlsplit(text)
port = parsed.port or 5432
host = parsed.hostname or ''
username = unquote(parsed.username or '')
password = unquote(parsed.password or '')
database = unquote((parsed.path or '')[1:]) if (parsed.path or '').startswith('/') else ''
except (TypeError, ValueError) as exc:
raise DatabaseUrlError('invalid PostgreSQL database URL') from exc
if parsed.scheme.lower() not in ('postgresql', 'postgres'):
raise DatabaseUrlError('database URL must use the PostgreSQL scheme')
if parsed.query:
raise DatabaseUrlError('PostgreSQL database URL query parameters are forbidden')
if parsed.fragment:
raise DatabaseUrlError('PostgreSQL database URL fragments are forbidden')
if not parsed.netloc or not host or not username or not database:
raise DatabaseUrlError('PostgreSQL database URL must include one host, user, and database')
authority = parsed.netloc.rsplit('@', 1)[-1]
decoded_authority = unquote(authority)
if (
',' in decoded_authority
or any(character.isspace() for character in decoded_authority)
or '%' in authority
or any(character in host for character in (',', '/', '\\', '\x00'))
):
raise DatabaseUrlError('PostgreSQL database URL must contain exactly one literal host authority')
if not 0 < int(port) <= 65535:
raise DatabaseUrlError('PostgreSQL database URL port is invalid')
if any(character in database for character in ('/', '\\', '?', '#', '\x00')):
raise DatabaseUrlError('PostgreSQL database URL database name contains encoded authority syntax')
if parsed.path.count('/') != 1:
raise DatabaseUrlError('PostgreSQL database URL must contain exactly one database path segment')
return {
'parsed': parsed,
'host': host.lower(),
'port': int(port),
'database': database,
'user': username,
'password': password,
}
def canonical_postgres_url(value, database, user, port, host='127.0.0.1'):
"""Return one libpq URL whose endpoint cannot be redirected by DSN options."""
values = parse_postgres_url(value)
expected_host = str(host).lower()
if values['host'] != expected_host or values['port'] != int(port):
raise DatabaseUrlError('PostgreSQL database URL endpoint does not match managed cluster authority')
if values['database'] != str(database) or values['user'] != str(user):
raise DatabaseUrlError('PostgreSQL database URL identity does not match managed cluster authority')
credentials = quote(str(user), safe='')
if values['parsed'].password is not None:
credentials += ':' + quote(values['password'], safe='')
netloc = f'{credentials}@{expected_host}:{int(port)}'
return urlunsplit(('postgresql', netloc, '/' + quote(str(database), safe=''), '', ''))
def redact_database_url(value):
text = str(value or '')
if not is_postgres_url(text):
if text.strip().lower().startswith(('postgres', 'postgre')):
return 'postgresql://***'
if '://' in text:
return text.split('://', 1)[0] + '://***'
return text
try:
parsed = urlsplit(text)
username = parsed.username or ''
host = parsed.hostname or ''
port = f':{parsed.port}' if parsed.port else ''
netloc = parsed.netloc
if parsed.password:
netloc = f'{username}:***@{host}{port}' if username else f'***@{host}{port}'
query = []
for key, value in parse_qsl(parsed.query, keep_blank_values=True):
key_lower = key.lower()
sensitive = any(part in key_lower for part in ('password', 'passwd', 'pwd', 'token', 'secret', 'credential'))
query.append((key, '***' if sensitive else value))
return urlunsplit((parsed.scheme, netloc, parsed.path, urlencode(query), parsed.fragment))
except Exception:
return 'postgresql://***'
def _split_sql_script(script):
statements = []
current = []
quote = None
escape = False
for char in str(script or ''):
current.append(char)
if escape:
escape = False
continue
if char == '\\':
escape = True
continue
if quote:
if char == quote:
quote = None
continue
if char in ("'", '"'):
quote = char
continue
if char == ';':
statement = ''.join(current).strip()
if statement:
statements.append(statement[:-1].strip())
current = []
tail = ''.join(current).strip()
if tail:
statements.append(tail)
return [statement for statement in statements if statement]
def _convert_qmark_to_psycopg(sql):
out = []
quote = None
escape = False
for char in str(sql or ''):
if escape:
out.append(char)
escape = False
continue
if char == '\\':
out.append(char)
escape = True
continue
if quote:
out.append('%%' if char == '%' else char)
if char == quote:
quote = None
continue
if char in ("'", '"'):
out.append(char)
quote = char
continue
if char == '?':
out.append('%s')
elif char == '%':
out.append('%%')
else:
out.append(char)
return ''.join(out)
def _postgres_schema_sql(script):
converted = str(script or '').replace(
'INTEGER PRIMARY KEY AUTOINCREMENT',
'BIGINT GENERATED BY DEFAULT AS IDENTITY PRIMARY KEY',
)
# PostgreSQL cannot create either side of the target_queue/target_scans
# cycle with both inline FKs. The offline migration adds this edge after
# both tables exist; SQLite can retain it in the base schema.
converted = converted.replace(
',\n FOREIGN KEY(queue_id) REFERENCES target_queue(id)',
'',
)
for future_foreign_key in (
',\n FOREIGN KEY(result_reservation_id) REFERENCES result_reservations(id)',
',\n FOREIGN KEY(reservation_id) REFERENCES result_reservations(id)',
',\n FOREIGN KEY(current_result_reservation_id) REFERENCES result_reservations(id)',
',\n FOREIGN KEY(candidate_id) REFERENCES keycheck_candidates(id)',
',\n FOREIGN KEY(credential_id) REFERENCES keycheck_credentials(id)',
',\n FOREIGN KEY(last_append_id) REFERENCES projection_appends(id)',
',\n FOREIGN KEY(experiment_id) REFERENCES docker_depth_experiments(id)',
):
converted = converted.replace(future_foreign_key, '')
for column in (
'run_id', 'cycle_id', 'target_scan_id', 'finding_id', 'keycheck_result_id',
'last_run_id', 'last_cycle_id', 'queue_id', 'byte_offset', 'line_number',
'current_result_reservation_id', 'result_reservation_id', 'reservation_id',
'projection_job_id', 'keycheck_candidate_id', 'candidate_id', 'credential_id',
'keycheck_result_id', 'target_scan_id', 'job_id', 'last_append_id',
'last_job_id', 'last_result_id', 'object_id', 'lease_reservation_id',
'covered_reservation_id', 'declared_bytes', 'verified_bytes',
'experiment_id', 'pass_id', 'page_id', 'source_cycle_id', 'retry_work_id',
'repository_queue_id', 'first_cycle_id', 'last_cycle_id', 'first_page_id',
'last_page_id', 'target_queue_id', 'manifest_id', 'manifest_layer_id',
'eligibility_page_id', 'experiment_repository_id', 'experiment_target_id',
'scan_binding_id', 'manifest_size_bytes', 'layer_size_bytes',
'fence_generation', 'resolver_generation', 'dispatch_order', 'total_count',
'user_id', 'remote_user_id', 'remote_device_id',
'expected_revision', 'resulting_revision', 'revision',
'before_bytes', 'after_bytes', 'previous_event_id',
):
converted = converted.replace(f'{column} INTEGER', f'{column} BIGINT')
converted = converted.replace(
'CREATE VIEW IF NOT EXISTS keycheck_latest_state AS',
'CREATE OR REPLACE VIEW keycheck_latest_state AS',
)
return converted
def _sqlite_check_constraints(sql):
text = str(sql or '')
constraints = {}
index = 0
position = 0
quote = None
while position < len(text):
char = text[position]
if quote:
if char == quote:
if position + 1 < len(text) and text[position + 1] == quote:
position += 2
continue
quote = None
position += 1
continue
if char in ("'", '"', '`'):
quote = char
position += 1
continue
if (
text[position:position + 5].lower() != 'check'
or (position and (text[position - 1].isalnum() or text[position - 1] == '_'))
or (
position + 5 < len(text)
and (text[position + 5].isalnum() or text[position + 5] == '_')
)
):
position += 1
continue
opening = position + 5
while opening < len(text) and text[opening].isspace():
opening += 1
if opening >= len(text) or text[opening] != '(':
position += 5
continue
depth = 1
closing = opening + 1
expression_quote = None
while closing < len(text) and depth:
current = text[closing]
if expression_quote:
if current == expression_quote:
if closing + 1 < len(text) and text[closing + 1] == expression_quote:
closing += 2
continue
expression_quote = None
elif current in ("'", '"', '`'):
expression_quote = current
elif current == '(':
depth += 1
elif current == ')':
depth -= 1
closing += 1
if depth:
break
prefix = text[:position]
named = re.search(
r'\bCONSTRAINT\s+(?:"([A-Za-z_][A-Za-z0-9_$]*)"|'
r'([A-Za-z_][A-Za-z0-9_$]*))\s*$',
prefix,
re.IGNORECASE,
)
name = (named.group(1) or named.group(2)) if named else f'__unnamed_check_{index}'
expression = text[opening + 1:closing - 1]
constraint = {
'expression': expression,
'definition': f'CHECK ({expression})',
'valid': True,
}
if name in constraints:
constraints[name]['valid'] = False
name = f'__duplicate_check_{index}_{name}'
constraint['valid'] = False
constraints[name] = constraint
index += 1
position = closing
return constraints
class DatabaseConnection:
def __init__(self, dialect, conn, application_schema=None):
self.dialect = dialect
self._conn = conn
self.application_schema = application_schema if dialect == 'postgres' else None
@property
def is_postgres(self):
return self.dialect == 'postgres'
@property
def is_sqlite(self):
return self.dialect == 'sqlite'
def execute(self, sql, params=None):
params = tuple(params or ())
if self.is_postgres:
params = tuple(value.replace('\x00', '') if isinstance(value, str) else value for value in params)
cur = self._conn.cursor()
cur.execute(_convert_qmark_to_psycopg(sql), params)
return cur
return self._conn.execute(sql, params)
def executescript(self, script):
if self.is_sqlite:
return self._conn.executescript(script)
for statement in _split_sql_script(_postgres_schema_sql(script)):
self.execute(statement)
return None
def commit(self):
return self._conn.commit()
def rollback(self):
return self._conn.rollback()
def close(self):
return self._conn.close()
def table_columns(self, table):
if self.is_postgres:
rows = self.execute(
'''SELECT a.attname AS name
FROM pg_catalog.pg_attribute a
JOIN pg_catalog.pg_class c ON c.oid = a.attrelid
JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace
WHERE n.nspname = ?
AND c.relname = ?
AND a.attnum > 0
AND NOT a.attisdropped''',
(self.application_schema, table),
).fetchall()
return {row['name'] for row in rows}
return {row['name'] for row in self.execute(f'PRAGMA table_info({table})').fetchall()}
def table_column_details(self, table):
if self.is_postgres:
rows = self.execute(
'''SELECT a.attname AS name,
pg_catalog.format_type(a.atttypid, a.atttypmod) AS type,
a.attnotnull AS not_null,
a.attidentity AS identity_generation,
a.attgenerated AS generated_kind,
pg_catalog.pg_get_expr(d.adbin, d.adrelid) AS default_sql,
(d.oid IS NOT NULL) AS has_default,
pg_catalog.pg_get_serial_sequence(
pg_catalog.quote_ident(n.nspname) || '.' || pg_catalog.quote_ident(c.relname),
a.attname
) AS sequence_name,
COALESCE(i.indisprimary, false) AS primary_key
FROM pg_catalog.pg_attribute a
JOIN pg_catalog.pg_class c ON c.oid = a.attrelid
JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace
LEFT JOIN pg_catalog.pg_attrdef d ON d.adrelid = a.attrelid AND d.adnum = a.attnum
LEFT JOIN pg_catalog.pg_index i ON i.indrelid = a.attrelid
AND i.indisprimary AND a.attnum = ANY(i.indkey)
WHERE n.nspname = ?
AND c.relname = ?
AND a.attnum > 0
AND NOT a.attisdropped
ORDER BY a.attnum''',
(self.application_schema, table),
).fetchall()
return {
row['name']: {
'type': str(row['type'] or '').lower(),
'not_null': bool(row['not_null']),
'default': str(row['default_sql'] or ''),
'has_default': bool(row['has_default']),
'primary_key': bool(row['primary_key']),
'identity': str(row['identity_generation'] or ''),
'generated': str(row['generated_kind'] or ''),
'sequence': str(row['sequence_name'] or ''),
}
for row in rows
}
rows = self.execute(f'PRAGMA table_info({table})').fetchall()
return {
row['name']: {
'type': str(row['type'] or '').lower(),
'not_null': bool(row['notnull']) or bool(row['pk']),
'default': str(row['dflt_value'] or ''),
'has_default': row['dflt_value'] is not None,
'primary_key': bool(row['pk']),
'identity': '',
'generated': '',
'sequence': '',
}
for row in rows
}
def table_indexes(self, table):
if self.is_postgres:
rows = self.execute(
'''SELECT idx.relname AS name,
i.indisunique AS is_unique,
i.indisprimary AS is_primary,
i.indisvalid AS is_valid,
i.indisready AS is_ready,
i.indislive AS is_live,
pg_catalog.pg_get_expr(i.indpred, i.indrelid) AS predicate,
ARRAY(
SELECT pg_catalog.pg_get_indexdef(i.indexrelid, position, true)
FROM pg_catalog.generate_series(1, i.indnkeyatts) AS position
ORDER BY position
) AS columns,
pg_catalog.pg_get_indexdef(i.indexrelid) AS sql
FROM pg_catalog.pg_index i
JOIN pg_catalog.pg_class tbl ON tbl.oid = i.indrelid
JOIN pg_catalog.pg_namespace n ON n.oid = tbl.relnamespace
JOIN pg_catalog.pg_class idx ON idx.oid = i.indexrelid
WHERE n.nspname = ? AND tbl.relname = ?''',
(self.application_schema, table),
).fetchall()
return {
row['name']: {
'unique': bool(row['is_unique']),
'primary': bool(row['is_primary']),
'valid': bool(row['is_valid']),
'ready': bool(row['is_ready']),
'live': bool(row['is_live']),
'predicate': str(row['predicate'] or ''),
'columns': [str(value).strip('"') for value in (row['columns'] or [])],
'sql': str(row['sql'] or ''),
}
for row in rows
}
output = {}
for row in self.execute(f'PRAGMA index_list({table})').fetchall():
name = row['name']
columns = [item['name'] for item in self.execute(f'PRAGMA index_info({name})').fetchall()]
sql_row = self.execute(
"SELECT sql FROM sqlite_master WHERE type = 'index' AND name = ?",
(name,),
).fetchone()
sql = str(sql_row['sql'] if sql_row else '')
predicate = sql.split(' WHERE ', 1)[1] if ' WHERE ' in sql.upper() else ''
if ' WHERE ' in sql.upper():
position = sql.upper().index(' WHERE ')
predicate = sql[position + 7:]
output[name] = {
'unique': bool(row['unique']),
'primary': str(row['origin'] or '') == 'pk',
'valid': True,
'ready': True,
'live': True,
'predicate': predicate,
'columns': columns,
'sql': sql,
}
return output
def table_foreign_keys(self, table):
if self.is_postgres:
rows = self.execute(
'''SELECT con.conname AS name,
ARRAY(
SELECT src.attname
FROM pg_catalog.unnest(con.conkey) WITH ORDINALITY AS keys(attnum, position)
JOIN pg_catalog.pg_attribute src
ON src.attrelid = con.conrelid AND src.attnum = keys.attnum
ORDER BY keys.position
) AS columns,
ref_n.nspname AS referenced_schema,
ref.relname AS referenced_table,
ARRAY(
SELECT dst.attname
FROM pg_catalog.unnest(con.confkey) WITH ORDINALITY AS keys(attnum, position)
JOIN pg_catalog.pg_attribute dst
ON dst.attrelid = con.confrelid AND dst.attnum = keys.attnum
ORDER BY keys.position
) AS referenced_columns,
con.confupdtype::text AS update_action,
con.confdeltype::text AS delete_action,
con.convalidated AS is_valid
FROM pg_catalog.pg_constraint con
JOIN pg_catalog.pg_class tbl ON tbl.oid = con.conrelid
JOIN pg_catalog.pg_namespace n ON n.oid = tbl.relnamespace
JOIN pg_catalog.pg_class ref ON ref.oid = con.confrelid
JOIN pg_catalog.pg_namespace ref_n ON ref_n.oid = ref.relnamespace
WHERE con.contype = 'f' AND n.nspname = ? AND tbl.relname = ?''',
(self.application_schema, table),
).fetchall()
action_names = {
'a': 'NO ACTION', 'r': 'RESTRICT', 'c': 'CASCADE',
'n': 'SET NULL', 'd': 'SET DEFAULT',
}
return {
row['name']: {
'columns': [str(value) for value in (row['columns'] or [])],
'referenced_schema': str(row['referenced_schema'] or ''),
'referenced_table': str(row['referenced_table'] or ''),
'referenced_columns': [str(value) for value in (row['referenced_columns'] or [])],
'update_action': action_names.get(str(row['update_action'] or ''), str(row['update_action'] or '')),
'delete_action': action_names.get(str(row['delete_action'] or ''), str(row['delete_action'] or '')),
'valid': bool(row['is_valid']),
}
for row in rows
}
output = {}
for row in self.execute(f'PRAGMA foreign_key_list({table})').fetchall():
name = f'fk_{row["id"]}'
current = output.setdefault(name, {
'columns': [],
'referenced_schema': 'main',
'referenced_table': str(row['table'] or ''),
'referenced_columns': [],
'update_action': str(row['on_update'] or '').upper(),
'delete_action': str(row['on_delete'] or '').upper(),
'valid': True,
})
current['columns'].append(str(row['from'] or ''))
current['referenced_columns'].append(str(row['to'] or ''))
return output
def table_check_constraints(self, table):
if self.is_postgres:
rows = self.execute(
'''SELECT con.conname AS name,
pg_catalog.pg_get_expr(con.conbin, con.conrelid, true) AS expression,
pg_catalog.pg_get_constraintdef(con.oid, true) AS definition,
con.convalidated AS is_valid
FROM pg_catalog.pg_constraint con
JOIN pg_catalog.pg_class tbl ON tbl.oid = con.conrelid
JOIN pg_catalog.pg_namespace n ON n.oid = tbl.relnamespace
WHERE con.contype = 'c' AND n.nspname = ? AND tbl.relname = ?''',
(self.application_schema, table),
).fetchall()
return {
str(row['name']): {
'expression': str(row['expression'] or ''),
'definition': str(row['definition'] or ''),
'valid': bool(row['is_valid']),
}
for row in rows
}
row = self.execute(
"SELECT sql FROM sqlite_master WHERE type = 'table' AND name = ?",
(table,),
).fetchone()
return _sqlite_check_constraints(row['sql'] if row else '')
def table_triggers(self, table):
if self.is_postgres:
rows = self.execute(
'''SELECT trg.tgname AS name,
trg.tgenabled <> 'D' AS enabled,
pg_catalog.pg_get_triggerdef(trg.oid, true) AS sql,
pg_catalog.pg_get_functiondef(trg.tgfoid) AS function_sql
FROM pg_catalog.pg_trigger trg
JOIN pg_catalog.pg_class tbl ON tbl.oid = trg.tgrelid
JOIN pg_catalog.pg_namespace n ON n.oid = tbl.relnamespace
WHERE n.nspname = ? AND tbl.relname = ?
AND NOT trg.tgisinternal''',
(self.application_schema, table),
).fetchall()
return {
str(row['name']): {
'enabled': bool(row['enabled']),
'sql': str(row['sql'] or ''),
'function_sql': str(row['function_sql'] or ''),
}
for row in rows
}
rows = self.execute(
"SELECT name, sql FROM sqlite_master WHERE type = 'trigger' AND tbl_name = ?",
(table,),
).fetchall()
return {
str(row['name']): {
'enabled': True,
'sql': str(row['sql'] or ''),
'function_sql': '',
}
for row in rows
}
def table_exists(self, table):
if self.is_postgres:
row = self.execute(
'''SELECT c.oid AS name FROM pg_catalog.pg_class c
JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace
WHERE n.nspname = ? AND c.relname = ? AND c.relkind IN ('r', 'p')''',
(self.application_schema, table),
).fetchone()
return bool(row and row['name'])
row = self.execute("SELECT name FROM sqlite_master WHERE type = 'table' AND name = ?", (table,)).fetchone()
return bool(row)
def insert_returning_id(self, sql, params=None):
if self.is_postgres:
cur = self.execute(f'{sql.rstrip()} RETURNING id', params)
row = cur.fetchone()
return row['id'] if row else None
cur = self.execute(sql, params)
return cur.lastrowid if int(getattr(cur, 'rowcount', 0) or 0) != 0 else None
def json_extract(self, column, path):
if self.is_sqlite:
return f"json_extract({column}, '{path}')"
parts = str(path or '').lstrip('$.').split('.')
pg_path = ','.join(part for part in parts if part)
return f"(NULLIF({column}, '')::jsonb #>> '{{{pg_path}}}')"
def connect_sqlite(path, timeout_sec=30, read_only=False, immutable=False, check_same_thread=True):
if read_only:
params = 'mode=ro&immutable=1' if immutable else 'mode=ro'
uri = 'file:' + str(path).replace('\\', '/') + '?' + params
conn = sqlite3.connect(uri, uri=True, timeout=max(1, int(timeout_sec or 30)), check_same_thread=check_same_thread)
else:
conn = sqlite3.connect(path, timeout=max(1, int(timeout_sec or 30)), check_same_thread=check_same_thread)
conn.row_factory = sqlite3.Row
return DatabaseConnection('sqlite', conn)
def _bounded_int(value, default, minimum=1):
try:
return max(minimum, int(value))
except (TypeError, ValueError):
return max(minimum, int(default))
def connect_postgres(
url,
connect_timeout_sec=None,
statement_timeout_ms=None,
lock_timeout_ms=None,
idle_in_transaction_timeout_ms=None,
tcp_user_timeout_ms=None,
):
parse_postgres_url(url)
try:
import psycopg
from psycopg.rows import dict_row
except ImportError as exc:
raise RuntimeError('PostgreSQL backend requires psycopg[binary]. Install app requirements first.') from exc
connect_timeout_sec = _bounded_int(
connect_timeout_sec if connect_timeout_sec is not None else os.getenv('TRUF_DB_CONNECT_TIMEOUT_SEC'),
DEFAULT_POSTGRES_CONNECT_TIMEOUT_SEC,
)
statement_timeout_ms = _bounded_int(
statement_timeout_ms if statement_timeout_ms is not None else os.getenv('TRUF_DB_STATEMENT_TIMEOUT_MS'),
DEFAULT_POSTGRES_STATEMENT_TIMEOUT_MS,
)
lock_timeout_ms = _bounded_int(
lock_timeout_ms if lock_timeout_ms is not None else os.getenv('TRUF_DB_LOCK_TIMEOUT_MS'),
DEFAULT_POSTGRES_LOCK_TIMEOUT_MS,
)
idle_in_transaction_timeout_ms = _bounded_int(
idle_in_transaction_timeout_ms if idle_in_transaction_timeout_ms is not None else os.getenv('TRUF_DB_IDLE_TRANSACTION_TIMEOUT_MS'),
DEFAULT_POSTGRES_IDLE_TRANSACTION_TIMEOUT_MS,
)
tcp_user_timeout_ms = _bounded_int(
tcp_user_timeout_ms if tcp_user_timeout_ms is not None else os.getenv('TRUF_DB_TCP_USER_TIMEOUT_MS'),
DEFAULT_POSTGRES_TCP_USER_TIMEOUT_MS,
minimum=1000,
)
options = ' '.join((
f'-c search_path={POSTGRES_APPLICATION_SCHEMA}',
f'-c statement_timeout={statement_timeout_ms}',
f'-c lock_timeout={lock_timeout_ms}',
f'-c idle_in_transaction_session_timeout={idle_in_transaction_timeout_ms}',
))
for attempt in range(POSTGRES_CHILD_START_RETRY_ATTEMPTS):
try:
conn = psycopg.connect(
url,
row_factory=dict_row,
connect_timeout=connect_timeout_sec,
options=options,
tcp_user_timeout=tcp_user_timeout_ms,
keepalives=1,
keepalives_idle=5,
keepalives_interval=5,
keepalives_count=2,
)
break
except Exception as exc:
transient_child_start = any(
marker in str(exc).lower() for marker in POSTGRES_CHILD_START_RETRY_MARKERS
)
if not transient_child_start or attempt + 1 >= POSTGRES_CHILD_START_RETRY_ATTEMPTS:
raise
time.sleep(0.05 * (attempt + 1))
try:
cursor = conn.cursor()
cursor.execute(
"""SELECT pg_catalog.current_schema() AS schema_name,
pg_catalog.current_setting('search_path') AS search_path,
EXISTS (
SELECT 1
FROM pg_catalog.pg_namespace n
CROSS JOIN LATERAL pg_catalog.aclexplode(
COALESCE(n.nspacl, pg_catalog.acldefault('n', n.nspowner))
) acl
WHERE n.nspname = 'public'
AND acl.grantee = 0
AND acl.privilege_type = 'CREATE'
) AS public_create"""
)
row = cursor.fetchone()
cursor.close()
schema_name = row.get('schema_name') if isinstance(row, dict) else row[0] if row else None
search_path = row.get('search_path') if isinstance(row, dict) else row[1] if row else None
public_create = row.get('public_create', False) if isinstance(row, dict) else row[2] if row and len(row) > 2 else False
normalized_path = re.sub(r'[\s\"]', '', str(search_path or '').lower())
if schema_name != POSTGRES_APPLICATION_SCHEMA or normalized_path != 'public' or bool(public_create):
raise RuntimeError('PostgreSQL application schema/search_path validation failed')
conn.rollback()
except Exception:
try:
conn.close()
except Exception:
pass
raise
return DatabaseConnection('postgres', conn, application_schema=POSTGRES_APPLICATION_SCHEMA)
def connect_host_agent_postgres():
"""Open the one fixed peer-authenticated host-agent authority."""
if os.name != 'posix' or not hasattr(os, 'geteuid') or os.geteuid() != 0:
raise RuntimeError('PostgreSQL host-agent authority requires root on POSIX')
try:
import psycopg
from psycopg.rows import dict_row
except ImportError as exc:
raise RuntimeError(
'PostgreSQL backend requires psycopg[binary]. Install app requirements first.'
) from exc
options = ' '.join((
f'-c search_path={POSTGRES_APPLICATION_SCHEMA}',
f'-c statement_timeout={DEFAULT_POSTGRES_STATEMENT_TIMEOUT_MS}',
f'-c lock_timeout={DEFAULT_POSTGRES_LOCK_TIMEOUT_MS}',
f'-c idle_in_transaction_session_timeout={DEFAULT_POSTGRES_IDLE_TRANSACTION_TIMEOUT_MS}',
))
conn = psycopg.connect(
dbname=HOST_AGENT_POSTGRES_DATABASE,
user=HOST_AGENT_POSTGRES_USER,
host=HOST_AGENT_POSTGRES_SOCKET_DIRECTORY,
port=HOST_AGENT_POSTGRES_PORT,
row_factory=dict_row,
connect_timeout=DEFAULT_POSTGRES_CONNECT_TIMEOUT_SEC,
options=options,
sslmode='disable',
)
try:
cursor = conn.cursor()
cursor.execute(
"""SELECT pg_catalog.current_schema() AS schema_name,
pg_catalog.current_setting('search_path') AS search_path,
CURRENT_USER AS current_user,
current_database() AS database_name,
pg_catalog.inet_server_addr() IS NULL AS unix_socket,
pg_catalog.current_setting('port')::integer AS port,
EXISTS (
SELECT 1
FROM pg_catalog.pg_namespace n
CROSS JOIN LATERAL pg_catalog.aclexplode(
COALESCE(n.nspacl, pg_catalog.acldefault('n', n.nspowner))
) acl
WHERE n.nspname = 'public'
AND acl.grantee = 0
AND acl.privilege_type = 'CREATE'
) AS public_create"""
)
row = cursor.fetchone()
cursor.close()
normalized_path = re.sub(
r'[\s\"]', '', str((row or {}).get('search_path') or '').lower()
)
if (
not isinstance(row, dict)
or row.get('schema_name') != POSTGRES_APPLICATION_SCHEMA
or normalized_path != POSTGRES_APPLICATION_SCHEMA
or row.get('current_user') != HOST_AGENT_POSTGRES_USER
or row.get('database_name') != HOST_AGENT_POSTGRES_DATABASE
or row.get('unix_socket') is not True
or row.get('port') != HOST_AGENT_POSTGRES_PORT
or bool(row.get('public_create'))
):
raise RuntimeError('PostgreSQL host-agent authority validation failed')
conn.rollback()
except Exception:
try:
conn.close()
except Exception:
pass
raise
return DatabaseConnection(
'postgres', conn, application_schema=POSTGRES_APPLICATION_SCHEMA,
)