Initial server source import
This commit is contained in:
@@ -0,0 +1,836 @@
|
||||
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,
|
||||
)
|
||||
Reference in New Issue
Block a user