837 lines
35 KiB
Python
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,
|
|
)
|