import unittest from unittest.mock import MagicMock, patch import app.services.ops_agent_service as ops_agent_service import app.services.ops_job_service as ops_job_service import app.services.ops_release_service as ops_release_service from psycopg2 import errors class OpsSchemaInitTests(unittest.TestCase): @patch("app.services.ops_job_service.get_db") def test_ensure_ops_schema_uses_advisory_lock_and_skips_after_ready(self, mock_get_db) -> None: conn = MagicMock() cursor = MagicMock() cursor.fetchone.return_value = None db_ctx = MagicMock() cursor_ctx = MagicMock() db_ctx.__enter__.return_value = conn db_ctx.__exit__.return_value = False cursor_ctx.__enter__.return_value = cursor cursor_ctx.__exit__.return_value = False conn.cursor.return_value = cursor_ctx mock_get_db.return_value = db_ctx previous_ready = ops_job_service._OPS_SCHEMA_READY ops_job_service._OPS_SCHEMA_READY = False try: ops_job_service.ensure_ops_schema() ops_job_service.ensure_ops_schema() finally: ops_job_service._OPS_SCHEMA_READY = previous_ready self.assertEqual(1, mock_get_db.call_count) cursor.execute.assert_any_call( "SELECT pg_advisory_xact_lock(%s)", (ops_job_service._OPS_SCHEMA_ADVISORY_LOCK_KEY,), ) cursor.execute.assert_any_call(ops_job_service._OPS_SCHEMA_SQL) conn.commit.assert_called_once() @patch("app.services.ops_job_service.get_db") def test_ensure_ops_schema_skips_ddl_when_required_schema_already_exists(self, mock_get_db) -> None: conn = MagicMock() cursor = MagicMock() cursor.fetchone.side_effect = [(f"public.{name}",) for name in ops_job_service._OPS_REQUIRED_TABLES] cursor.fetchall.side_effect = [ [(column,) for column in ops_job_service._OPS_REQUIRED_COLUMNS["ops_jobs"]], [(column,) for column in ops_job_service._OPS_REQUIRED_COLUMNS["ops_job_steps"]], ] db_ctx = MagicMock() cursor_ctx = MagicMock() db_ctx.__enter__.return_value = conn db_ctx.__exit__.return_value = False cursor_ctx.__enter__.return_value = cursor cursor_ctx.__exit__.return_value = False conn.cursor.return_value = cursor_ctx mock_get_db.return_value = db_ctx previous_ready = ops_job_service._OPS_SCHEMA_READY ops_job_service._OPS_SCHEMA_READY = False try: ops_job_service.ensure_ops_schema() finally: ops_job_service._OPS_SCHEMA_READY = previous_ready self.assertFalse(any(call.args[0] == ops_job_service._OPS_SCHEMA_SQL for call in cursor.execute.call_args_list)) conn.commit.assert_not_called() @patch("app.services.ops_job_service.get_db") def test_ensure_ops_schema_accepts_deadlock_when_required_schema_already_exists(self, mock_get_db) -> None: class _Cursor: def __init__(self, *, raise_on_schema=False, fetchone_values=None, fetchall_values=None) -> None: self.raise_on_schema = raise_on_schema self.fetchone_values = list(fetchone_values or []) self.fetchall_values = list(fetchall_values or []) def execute(self, sql, params=None): if self.raise_on_schema and sql == ops_job_service._OPS_SCHEMA_SQL: raise errors.DeadlockDetected() def fetchone(self): if self.fetchone_values: return self.fetchone_values.pop(0) return None def fetchall(self): if self.fetchall_values: return self.fetchall_values.pop(0) return [] class _CursorContext: def __init__(self, cursor) -> None: self.cursor = cursor def __enter__(self): return self.cursor def __exit__(self, exc_type, exc, tb): return False conn = MagicMock() conn.cursor.side_effect = [ _CursorContext(_Cursor(fetchone_values=[None])), _CursorContext(_Cursor(raise_on_schema=True)), _CursorContext( _Cursor( fetchone_values=[(f"public.{name}",) for name in ops_job_service._OPS_REQUIRED_TABLES], fetchall_values=[ [(column,) for column in ops_job_service._OPS_REQUIRED_COLUMNS["ops_jobs"]], [(column,) for column in ops_job_service._OPS_REQUIRED_COLUMNS["ops_job_steps"]], ], ) ), ] db_ctx = MagicMock() db_ctx.__enter__.return_value = conn mock_get_db.return_value = db_ctx previous_ready = ops_job_service._OPS_SCHEMA_READY ops_job_service._OPS_SCHEMA_READY = False try: ops_job_service.ensure_ops_schema() finally: ops_job_service._OPS_SCHEMA_READY = previous_ready conn.rollback.assert_called_once() conn.commit.assert_not_called() @patch("app.services.ops_agent_service.ensure_ops_schema") @patch("app.services.ops_agent_service.get_db") def test_ensure_ops_agent_schema_uses_advisory_lock_and_skips_after_ready( self, mock_get_db, mock_ensure_ops_schema, ) -> None: conn = MagicMock() cursor = MagicMock() cursor.fetchone.return_value = None db_ctx = MagicMock() cursor_ctx = MagicMock() db_ctx.__enter__.return_value = conn db_ctx.__exit__.return_value = False cursor_ctx.__enter__.return_value = cursor cursor_ctx.__exit__.return_value = False conn.cursor.return_value = cursor_ctx mock_get_db.return_value = db_ctx previous_ready = ops_agent_service._OPS_AGENT_SCHEMA_READY ops_agent_service._OPS_AGENT_SCHEMA_READY = False try: ops_agent_service.ensure_ops_agent_schema() ops_agent_service.ensure_ops_agent_schema() finally: ops_agent_service._OPS_AGENT_SCHEMA_READY = previous_ready self.assertEqual(2, mock_ensure_ops_schema.call_count) self.assertEqual(1, mock_get_db.call_count) cursor.execute.assert_any_call( "SELECT pg_advisory_xact_lock(%s)", (ops_agent_service._OPS_AGENT_SCHEMA_ADVISORY_LOCK_KEY,), ) cursor.execute.assert_any_call(ops_agent_service._AGENT_SCHEMA_SQL) conn.commit.assert_called_once() @patch("app.services.ops_agent_service.ensure_ops_schema") @patch("app.services.ops_agent_service.get_db") def test_ensure_ops_agent_schema_skips_ddl_when_required_schema_already_exists( self, mock_get_db, mock_ensure_ops_schema, ) -> None: conn = MagicMock() cursor = MagicMock() cursor.fetchone.side_effect = [(f"public.{name}",) for name in ops_agent_service._OPS_AGENT_REQUIRED_TABLES] cursor.fetchall.side_effect = [ [(column,) for column in ops_agent_service._OPS_AGENT_REQUIRED_COLUMNS["ops_node_tokens"]], [(column,) for column in ops_agent_service._OPS_AGENT_REQUIRED_COLUMNS["ops_job_events"]], [(column,) for column in ops_agent_service._OPS_AGENT_REQUIRED_COLUMNS["ops_jobs"]], ] db_ctx = MagicMock() cursor_ctx = MagicMock() db_ctx.__enter__.return_value = conn db_ctx.__exit__.return_value = False cursor_ctx.__enter__.return_value = cursor cursor_ctx.__exit__.return_value = False conn.cursor.return_value = cursor_ctx mock_get_db.return_value = db_ctx previous_ready = ops_agent_service._OPS_AGENT_SCHEMA_READY ops_agent_service._OPS_AGENT_SCHEMA_READY = False try: ops_agent_service.ensure_ops_agent_schema() finally: ops_agent_service._OPS_AGENT_SCHEMA_READY = previous_ready self.assertFalse(any(call.args[0] == ops_agent_service._AGENT_SCHEMA_SQL for call in cursor.execute.call_args_list)) conn.commit.assert_not_called() self.assertEqual(1, mock_ensure_ops_schema.call_count) @patch("app.services.ops_agent_service.ensure_ops_schema") @patch("app.services.ops_agent_service.get_db") def test_ensure_ops_agent_schema_accepts_deadlock_when_required_schema_already_exists( self, mock_get_db, mock_ensure_ops_schema, ) -> None: class _Cursor: def __init__(self, *, raise_on_schema=False, fetchone_values=None, fetchall_values=None) -> None: self.raise_on_schema = raise_on_schema self.fetchone_values = list(fetchone_values or []) self.fetchall_values = list(fetchall_values or []) def execute(self, sql, params=None): if self.raise_on_schema and sql == ops_agent_service._AGENT_SCHEMA_SQL: raise errors.DeadlockDetected() def fetchone(self): if self.fetchone_values: return self.fetchone_values.pop(0) return None def fetchall(self): if self.fetchall_values: return self.fetchall_values.pop(0) return [] class _CursorContext: def __init__(self, cursor) -> None: self.cursor = cursor def __enter__(self): return self.cursor def __exit__(self, exc_type, exc, tb): return False conn = MagicMock() conn.cursor.side_effect = [ _CursorContext(_Cursor(fetchone_values=[None, None])), _CursorContext(_Cursor(raise_on_schema=True)), _CursorContext( _Cursor( fetchone_values=[(f"public.{name}",) for name in ops_agent_service._OPS_AGENT_REQUIRED_TABLES], fetchall_values=[ [(column,) for column in ops_agent_service._OPS_AGENT_REQUIRED_COLUMNS["ops_node_tokens"]], [(column,) for column in ops_agent_service._OPS_AGENT_REQUIRED_COLUMNS["ops_job_events"]], [(column,) for column in ops_agent_service._OPS_AGENT_REQUIRED_COLUMNS["ops_jobs"]], ], ) ), ] db_ctx = MagicMock() db_ctx.__enter__.return_value = conn mock_get_db.return_value = db_ctx previous_ready = ops_agent_service._OPS_AGENT_SCHEMA_READY ops_agent_service._OPS_AGENT_SCHEMA_READY = False try: ops_agent_service.ensure_ops_agent_schema() finally: ops_agent_service._OPS_AGENT_SCHEMA_READY = previous_ready conn.rollback.assert_called_once() conn.commit.assert_not_called() self.assertEqual(1, mock_ensure_ops_schema.call_count) @patch("app.services.ops_release_service.get_db") def test_ensure_ops_release_schema_uses_advisory_lock_and_skips_after_ready(self, mock_get_db) -> None: conn = MagicMock() cursor = MagicMock() cursor.fetchone.return_value = None db_ctx = MagicMock() cursor_ctx = MagicMock() db_ctx.__enter__.return_value = conn db_ctx.__exit__.return_value = False cursor_ctx.__enter__.return_value = cursor cursor_ctx.__exit__.return_value = False conn.cursor.return_value = cursor_ctx mock_get_db.return_value = db_ctx previous_ready = ops_release_service._RELEASE_SCHEMA_READY ops_release_service._RELEASE_SCHEMA_READY = False try: ops_release_service.ensure_ops_release_schema() ops_release_service.ensure_ops_release_schema() finally: ops_release_service._RELEASE_SCHEMA_READY = previous_ready self.assertEqual(1, mock_get_db.call_count) cursor.execute.assert_any_call( "SELECT pg_advisory_xact_lock(%s)", (ops_release_service._RELEASE_SCHEMA_ADVISORY_LOCK_KEY,), ) cursor.execute.assert_any_call(ops_release_service._RELEASE_SCHEMA_SQL) conn.commit.assert_called_once() @patch("app.services.ops_release_service.get_db") def test_ensure_ops_release_schema_skips_ddl_when_required_schema_already_exists(self, mock_get_db) -> None: conn = MagicMock() cursor = MagicMock() cursor.fetchone.side_effect = [(f"public.{name}",) for name in ops_release_service._RELEASE_REQUIRED_TABLES] cursor.fetchall.side_effect = [ [(column,) for column in ops_release_service._RELEASE_REQUIRED_COLUMNS["ops_release_rollouts"]], ] db_ctx = MagicMock() cursor_ctx = MagicMock() db_ctx.__enter__.return_value = conn db_ctx.__exit__.return_value = False cursor_ctx.__enter__.return_value = cursor cursor_ctx.__exit__.return_value = False conn.cursor.return_value = cursor_ctx mock_get_db.return_value = db_ctx previous_ready = ops_release_service._RELEASE_SCHEMA_READY ops_release_service._RELEASE_SCHEMA_READY = False try: ops_release_service.ensure_ops_release_schema() finally: ops_release_service._RELEASE_SCHEMA_READY = previous_ready self.assertFalse(any(call.args[0] == ops_release_service._RELEASE_SCHEMA_SQL for call in cursor.execute.call_args_list)) conn.commit.assert_not_called() @patch("app.services.ops_release_service.get_db") def test_ensure_ops_release_schema_accepts_deadlock_when_required_schema_already_exists(self, mock_get_db) -> None: class _Cursor: def __init__(self, *, raise_on_schema=False, fetchone_values=None, fetchall_values=None) -> None: self.raise_on_schema = raise_on_schema self.fetchone_values = list(fetchone_values or []) self.fetchall_values = list(fetchall_values or []) def execute(self, sql, params=None): if self.raise_on_schema and sql == ops_release_service._RELEASE_SCHEMA_SQL: raise errors.DeadlockDetected() def fetchone(self): if self.fetchone_values: return self.fetchone_values.pop(0) return None def fetchall(self): if self.fetchall_values: return self.fetchall_values.pop(0) return [] class _CursorContext: def __init__(self, cursor) -> None: self.cursor = cursor def __enter__(self): return self.cursor def __exit__(self, exc_type, exc, tb): return False conn = MagicMock() conn.cursor.side_effect = [ _CursorContext(_Cursor(fetchone_values=[None])), _CursorContext(_Cursor(raise_on_schema=True)), _CursorContext( _Cursor( fetchone_values=[(f"public.{name}",) for name in ops_release_service._RELEASE_REQUIRED_TABLES], fetchall_values=[ [(column,) for column in ops_release_service._RELEASE_REQUIRED_COLUMNS["ops_release_rollouts"]], ], ) ), ] db_ctx = MagicMock() db_ctx.__enter__.return_value = conn mock_get_db.return_value = db_ctx previous_ready = ops_release_service._RELEASE_SCHEMA_READY ops_release_service._RELEASE_SCHEMA_READY = False try: ops_release_service.ensure_ops_release_schema() finally: ops_release_service._RELEASE_SCHEMA_READY = previous_ready conn.rollback.assert_called_once() conn.commit.assert_not_called() if __name__ == "__main__": unittest.main()