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()