Files
truf-server/tests/test_db_backend_safety.py
2026-09-30 20:30:56 +03:00

326 lines
14 KiB
Python

from pathlib import Path
import sys
import types
import unittest
from unittest import mock
ROOT = Path(__file__).resolve().parents[1]
APP_DIR = ROOT / 'app'
sys.path.insert(0, str(APP_DIR))
import db_backend
class PostgresTimeoutTests(unittest.TestCase):
@staticmethod
def successful_connection():
raw_connection = mock.Mock()
cursor = mock.Mock()
cursor.fetchone.return_value = {
'schema_name': 'public',
'search_path': 'public',
'public_create': False,
}
raw_connection.cursor.return_value = cursor
return raw_connection
def test_connect_and_session_timeouts_are_explicit_and_overridable(self):
raw_connection = mock.Mock()
cursor = mock.Mock()
cursor.fetchone.return_value = {
'schema_name': 'public',
'search_path': 'public',
'public_create': False,
}
raw_connection.cursor.return_value = cursor
psycopg = types.ModuleType('psycopg')
psycopg.connect = mock.Mock(return_value=raw_connection)
rows = types.ModuleType('psycopg.rows')
rows.dict_row = object()
with mock.patch.dict(sys.modules, {'psycopg': psycopg, 'psycopg.rows': rows}):
connection = db_backend.connect_postgres(
'postgresql://user:secret@127.0.0.1:5432/db',
connect_timeout_sec=3,
statement_timeout_ms=4000,
lock_timeout_ms=1200,
idle_in_transaction_timeout_ms=5000,
tcp_user_timeout_ms=6000,
)
self.assertIs(connection._conn, raw_connection)
kwargs = psycopg.connect.call_args.kwargs
self.assertEqual(kwargs['connect_timeout'], 3)
self.assertIn('statement_timeout=4000', kwargs['options'])
self.assertIn('lock_timeout=1200', kwargs['options'])
self.assertIn('idle_in_transaction_session_timeout=5000', kwargs['options'])
self.assertIn('search_path=public', kwargs['options'])
self.assertEqual(kwargs['tcp_user_timeout'], 6000)
self.assertEqual(kwargs['keepalives_count'], 2)
cursor.execute.assert_called_once()
raw_connection.rollback.assert_called_once()
def test_public_create_privilege_fails_closed(self):
raw_connection = mock.Mock()
cursor = mock.Mock()
cursor.fetchone.return_value = {
'schema_name': 'public',
'search_path': 'public',
'public_create': True,
}
raw_connection.cursor.return_value = cursor
psycopg = types.ModuleType('psycopg')
psycopg.connect = mock.Mock(return_value=raw_connection)
rows = types.ModuleType('psycopg.rows')
rows.dict_row = object()
with mock.patch.dict(sys.modules, {'psycopg': psycopg, 'psycopg.rows': rows}):
with self.assertRaisesRegex(RuntimeError, 'search_path'):
db_backend.connect_postgres('postgresql://user:secret@127.0.0.1:5432/db')
raw_connection.close.assert_called_once()
def test_transient_windows_child_start_failure_is_retried_boundedly(self):
raw_connection = self.successful_connection()
psycopg = types.ModuleType('psycopg')
psycopg.connect = mock.Mock(side_effect=[
OSError('connection failed: server closed the connection unexpectedly'),
raw_connection,
])
rows = types.ModuleType('psycopg.rows')
rows.dict_row = object()
with mock.patch.dict(sys.modules, {'psycopg': psycopg, 'psycopg.rows': rows}), \
mock.patch.object(db_backend.time, 'sleep') as sleep:
connection = db_backend.connect_postgres(
'postgresql://user:secret@127.0.0.1:5432/db'
)
self.assertIs(connection._conn, raw_connection)
self.assertEqual(psycopg.connect.call_count, 2)
sleep.assert_called_once_with(0.05)
def test_nontransient_connection_failure_is_not_retried(self):
psycopg = types.ModuleType('psycopg')
psycopg.connect = mock.Mock(side_effect=OSError('password authentication failed'))
rows = types.ModuleType('psycopg.rows')
rows.dict_row = object()
with mock.patch.dict(sys.modules, {'psycopg': psycopg, 'psycopg.rows': rows}), \
mock.patch.object(db_backend.time, 'sleep') as sleep:
with self.assertRaisesRegex(OSError, 'password authentication failed'):
db_backend.connect_postgres(
'postgresql://user:secret@127.0.0.1:5432/db'
)
psycopg.connect.assert_called_once()
sleep.assert_not_called()
def test_host_agent_connection_uses_only_fixed_root_peer_authority(self):
raw_connection = mock.Mock()
cursor = mock.Mock()
cursor.fetchone.return_value = {
'schema_name': 'public',
'search_path': 'public',
'current_user': 'truf',
'database_name': 'truf',
'unix_socket': True,
'port': 5432,
'public_create': False,
}
raw_connection.cursor.return_value = cursor
psycopg = types.ModuleType('psycopg')
psycopg.connect = mock.Mock(return_value=raw_connection)
rows = types.ModuleType('psycopg.rows')
rows.dict_row = object()
with mock.patch.dict(sys.modules, {'psycopg': psycopg, 'psycopg.rows': rows}), \
mock.patch.object(db_backend.os, 'name', 'posix'), \
mock.patch.object(db_backend.os, 'geteuid', return_value=0, create=True):
connection = db_backend.connect_host_agent_postgres()
self.assertIs(connection._conn, raw_connection)
self.assertEqual(psycopg.connect.call_args.args, ())
self.assertEqual(psycopg.connect.call_args.kwargs['host'], '/run/truf-postgres')
self.assertEqual(psycopg.connect.call_args.kwargs['dbname'], 'truf')
self.assertEqual(psycopg.connect.call_args.kwargs['user'], 'truf')
self.assertEqual(psycopg.connect.call_args.kwargs['port'], 5432)
self.assertEqual(psycopg.connect.call_args.kwargs['sslmode'], 'disable')
raw_connection.rollback.assert_called_once()
def test_host_agent_connection_rejects_nonroot_and_wrong_authority(self):
psycopg = types.ModuleType('psycopg')
psycopg.connect = mock.Mock()
rows = types.ModuleType('psycopg.rows')
rows.dict_row = object()
with mock.patch.dict(sys.modules, {'psycopg': psycopg, 'psycopg.rows': rows}), \
mock.patch.object(db_backend.os, 'name', 'posix'), \
mock.patch.object(db_backend.os, 'geteuid', return_value=10001, create=True):
with self.assertRaisesRegex(RuntimeError, 'requires root'):
db_backend.connect_host_agent_postgres()
psycopg.connect.assert_not_called()
raw_connection = mock.Mock()
cursor = mock.Mock()
cursor.fetchone.return_value = {
'schema_name': 'public', 'search_path': 'public',
'current_user': 'truf', 'database_name': 'truf',
'unix_socket': False, 'port': 5432, 'public_create': False,
}
raw_connection.cursor.return_value = cursor
psycopg.connect.return_value = raw_connection
with mock.patch.dict(sys.modules, {'psycopg': psycopg, 'psycopg.rows': rows}), \
mock.patch.object(db_backend.os, 'name', 'posix'), \
mock.patch.object(db_backend.os, 'geteuid', return_value=0, create=True):
with self.assertRaisesRegex(RuntimeError, 'authority validation'):
db_backend.connect_host_agent_postgres()
raw_connection.close.assert_called_once()
class CanonicalPostgresUrlTests(unittest.TestCase):
def test_query_parameters_cannot_override_managed_authority(self):
attacks = (
'?host=attacker',
'?hostaddr=203.0.113.10',
'?port=6543',
'?dbname=other',
'?user=other',
'?service=evil',
'?servicefile=C%3A%5Cevil.conf',
'?options=-c%20search_path%3Devil',
)
base = 'postgresql://truf:secret@127.0.0.1:5432/truf'
for query in attacks:
with self.subTest(query=query), self.assertRaises(db_backend.DatabaseUrlError):
db_backend.canonical_postgres_url(base + query, 'truf', 'truf', 5432)
def test_canonical_url_preserves_password_but_forces_identity(self):
value = db_backend.canonical_postgres_url(
'postgres://truf:p%40ss@127.0.0.1/truf',
'truf',
'truf',
5432,
)
self.assertEqual(value, 'postgresql://truf:p%40ss@127.0.0.1:5432/truf')
def test_redirecting_urls_never_reach_psycopg_connector(self):
attacks = (
'postgresql://truf:secret@127.0.0.1:5432/truf?host=attacker',
'postgresql://truf:secret@127.0.0.1:5432/truf?hostaddr=203.0.113.8',
'postgresql://truf:secret@127.0.0.1:5432/truf?service=evil',
'postgresql://truf:secret@127.0.0.1:5432/truf?servicefile=C%3A%5Cevil.conf',
'postgresql://truf:secret@127.0.0.1:5432,attacker:5432/truf',
'postgresql://truf:secret@%31%32%37.0.0.1:5432/truf',
'host=attacker dbname=truf user=truf',
)
psycopg = types.ModuleType('psycopg')
psycopg.connect = mock.Mock()
rows = types.ModuleType('psycopg.rows')
rows.dict_row = object()
with mock.patch.dict(sys.modules, {'psycopg': psycopg, 'psycopg.rows': rows}):
for attack in attacks:
with self.subTest(attack=attack), self.assertRaises(db_backend.DatabaseUrlError):
db_backend.connect_postgres(attack)
psycopg.connect.assert_not_called()
class CatalogCursor:
def __init__(self, raw):
self.raw = raw
self.rows = []
def execute(self, sql, params):
self.raw.calls.append((sql, params))
if 'FROM pg_catalog.pg_attribute' in sql:
self.rows = [{
'name': 'id', 'type': 'bigint', 'not_null': True,
'identity_generation': 'd', 'generated_kind': '',
'default_sql': None, 'has_default': False,
'sequence_name': 'public.runs_id_seq', 'primary_key': True,
}]
elif 'FROM pg_catalog.pg_index' in sql:
self.rows = [{
'name': 'runs_pkey', 'is_unique': True, 'is_primary': True,
'is_valid': False, 'is_ready': True, 'is_live': True,
'predicate': None, 'columns': ['id'], 'sql': 'CREATE UNIQUE INDEX runs_pkey ON public.runs USING btree (id)',
}]
elif "con.contype = 'c'" in sql:
self.rows = [{
'name': 'runs_state_check',
'expression': "state = 'ready'::text",
'definition': "CHECK (state = 'ready'::text)",
'is_valid': False,
}]
elif 'FROM pg_catalog.pg_constraint' in sql:
self.rows = [{
'name': 'source_cycles_run_id_fkey', 'columns': ['run_id'],
'referenced_schema': 'public', 'referenced_table': 'runs',
'referenced_columns': ['id'], 'update_action': 'a',
'delete_action': 'c', 'is_valid': True,
}]
def fetchall(self):
return self.rows
class CatalogRawConnection:
def __init__(self):
self.calls = []
def cursor(self):
return CatalogCursor(self)
class PostgresCatalogTests(unittest.TestCase):
def test_catalog_exposes_generation_index_state_foreign_keys_and_public_scope(self):
raw = CatalogRawConnection()
connection = db_backend.DatabaseConnection('postgres', raw, application_schema='public')
details = connection.table_column_details('runs')['id']
index = connection.table_indexes('runs')['runs_pkey']
foreign_key = connection.table_foreign_keys('source_cycles')['source_cycles_run_id_fkey']
check = connection.table_check_constraints('runs')['runs_state_check']
self.assertEqual(details['identity'], 'd')
self.assertEqual(details['sequence'], 'public.runs_id_seq')
self.assertFalse(details['has_default'])
self.assertFalse(index['valid'])
self.assertTrue(index['ready'])
self.assertEqual(foreign_key['columns'], ['run_id'])
self.assertEqual(foreign_key['referenced_schema'], 'public')
self.assertEqual(foreign_key['update_action'], 'NO ACTION')
self.assertEqual(foreign_key['delete_action'], 'CASCADE')
self.assertEqual(check['expression'], "state = 'ready'::text")
self.assertFalse(check['valid'])
self.assertTrue(all(params[0] == 'public' for _, params in raw.calls))
self.assertTrue(any('attidentity' in sql and 'pg_get_serial_sequence' in sql for sql, _ in raw.calls))
self.assertTrue(any('indisvalid' in sql and 'indisready' in sql for sql, _ in raw.calls))
def test_postgres_schema_conversion_defers_only_the_cyclic_queue_fk(self):
converted = db_backend._postgres_schema_sql('''
CREATE TABLE target_scans (
id INTEGER PRIMARY KEY AUTOINCREMENT,
queue_id INTEGER,
FOREIGN KEY(queue_id) REFERENCES target_queue(id)
);
''')
self.assertIn('BIGINT GENERATED BY DEFAULT AS IDENTITY PRIMARY KEY', converted)
self.assertNotIn('FOREIGN KEY(queue_id)', converted)
def test_sqlite_check_parser_preserves_named_nested_constraints(self):
constraints = db_backend._sqlite_check_constraints('''
CREATE TABLE fixture (
value TEXT,
enabled INTEGER CHECK (enabled = 0 OR enabled = 1),
CONSTRAINT fixture_value_check CHECK (
length(replace(value, 'a', '')) = 0
AND (value IS NULL OR lower(value) = value)
)
)
''')
self.assertIn('fixture_value_check', constraints)
self.assertTrue(constraints['fixture_value_check']['valid'])
self.assertIn("replace(value, 'a', '')", constraints['fixture_value_check']['expression'])
self.assertEqual(
len([name for name in constraints if name.startswith('__unnamed_check_')]),
1,
)
if __name__ == '__main__':
unittest.main()