import json import os from pathlib import Path import sqlite3 import sys import tempfile import unittest import uuid from unittest import mock ROOT = Path(__file__).resolve().parents[1] APP_DIR = ROOT / 'app' sys.path.insert(0, str(APP_DIR)) import scanner_db from scanner_db import ScannerDB class OperationsControlTests(unittest.TestCase): def setUp(self): self.environment = mock.patch.dict( os.environ, {'SCANNER_DB_URL': '', 'DATABASE_URL': ''}, ) self.environment.start() self.temp = tempfile.TemporaryDirectory() self.path = os.path.join(self.temp.name, 'scanner.db') self.db = ScannerDB(db_path=self.path) def tearDown(self): self.db.close() self.temp.cleanup() self.environment.stop() @staticmethod def operation_id(): return str(uuid.uuid4()) def counts(self): return ( int(self.db.conn.execute( 'SELECT COUNT(*) AS count FROM runtime_operations' ).fetchone()['count']), int(self.db.conn.execute( 'SELECT COUNT(*) AS count FROM runtime_audit_events' ).fetchone()['count']), ) def add_reservation(self, *, assignment_kind='remote', resolved=False): suffix = uuid.uuid4().hex target = f'https://example.invalid/{suffix}' self.db.enqueue_targets('github', 'github', 'drain-fixture', [target]) queue_id = int(self.db.conn.execute( 'SELECT id FROM target_queue WHERE normalized_target = ?', (target,), ).fetchone()['id']) now = scanner_db.utc_now_iso() cursor = self.db.conn.execute( '''INSERT INTO result_reservations( reservation_token, bundle_id, scan_event_id, queue_id, source, platform, query, target, normalized_target, claim_lease_owner, claim_lease_token, claim_batch, producer_instance_id, producer_pid, producer_creation_time, producer_executable, assignment_kind, remote_resolved_at, declared_bundle_bytes, reserved_bundle_bytes, reserved_projection_bytes, reserved_candidate_items, reserved_candidate_bytes, ready_relative_path, state, producer_lease_expires_at, created_at, updated_at ) VALUES (?, ?, ?, ?, 'github', 'github', 'drain-fixture', ?, ?, 'fixture-owner', ?, 'fixture-batch', 'fixture-instance', 1, 'fixture-creation', 'fixture-executable', ?, ?, 1024, 1024, 1024, 1, 1024, ?, 'scanning', ?, ?, ?)''', ( suffix, 'b' + suffix, 'e' + suffix, queue_id, target, target, 'lease-' + suffix, assignment_kind, now if resolved else None, f'ready/{suffix}.trb', '2999-01-01T00:00:00+00:00', now, now, ), ) self.db.conn.commit() return int(cursor.lastrowid) def add_bundle(self, state='ready'): reservation_id = self.add_reservation(assignment_kind='local', resolved=False) suffix = uuid.uuid4().hex now = scanner_db.utc_now_iso() self.db.conn.execute( '''INSERT INTO result_bundles( reservation_id, bundle_id, scan_event_id, scan_event_hash, format_version, relative_path, actual_bytes, frame_count, finding_count, error_count, candidate_count, state, ready_at, updated_at ) VALUES (?, ?, ?, ?, 2, ?, 128, 1, 0, 0, 0, ?, ?, ?)''', ( reservation_id, 'bundle-' + suffix, 'event-' + suffix, 'a' * 64, f'ready/bundle-{suffix}.trb', state, now, now, ), ) self.db.conn.commit() return reservation_id def test_typed_controls_preserve_explicit_pauses_and_chain_audit_events(self): initial = self.db.runtime_control_state() self.assertEqual(initial['revision'], 0) self.assertFalse(initial['discovery_paused']) self.assertFalse(initial['dispatch_paused']) self.assertFalse(initial['effective_discovery_paused']) self.assertFalse(initial['effective_dispatch_paused']) self.assertEqual(initial['drain_state'], 'normal') operations = [] operation_id = self.operation_id() operations.append((operation_id, self.db.set_runtime_discovery_paused( True, expected_revision=0, actor='operator:alice', operation_id=operation_id, ))) operation_id = self.operation_id() operations.append((operation_id, self.db.set_runtime_dispatch_paused( True, expected_revision=1, actor='operator:alice', operation_id=operation_id, ))) operation_id = self.operation_id() operations.append((operation_id, self.db.start_runtime_drain( expected_revision=2, actor='operator:alice', operation_id=operation_id, ))) draining = self.db.runtime_control_state() self.assertEqual(draining['drain_state'], 'draining') self.assertTrue(draining['effective_discovery_paused']) self.assertTrue(draining['effective_dispatch_paused']) operation_id = self.operation_id() operations.append((operation_id, self.db.cancel_runtime_drain( expected_revision=3, actor='operator:alice', operation_id=operation_id, ))) canceled = self.db.runtime_control_state() self.assertEqual(canceled['drain_state'], 'normal') self.assertTrue(canceled['discovery_paused']) self.assertTrue(canceled['dispatch_paused']) self.assertTrue(canceled['effective_discovery_paused']) self.assertTrue(canceled['effective_dispatch_paused']) operation_id = self.operation_id() operations.append((operation_id, self.db.set_runtime_discovery_paused( False, expected_revision=4, actor='operator:alice', operation_id=operation_id, ))) operation_id = self.operation_id() operations.append((operation_id, self.db.set_runtime_dispatch_paused( False, expected_revision=5, actor='operator:alice', operation_id=operation_id, ))) final = self.db.runtime_control_state() self.assertEqual(final['revision'], 6) self.assertFalse(final['effective_discovery_paused']) self.assertFalse(final['effective_dispatch_paused']) self.assertEqual(self.counts(), (6, 6)) events = self.db.conn.execute( '''SELECT id, previous_event_id, previous_event_sha256, event_sha256 FROM runtime_audit_events ORDER BY id''' ).fetchall() for index, event in enumerate(events): if index == 0: self.assertIsNone(event['previous_event_id']) self.assertIsNone(event['previous_event_sha256']) else: self.assertEqual(int(event['previous_event_id']), int(events[index - 1]['id'])) self.assertEqual(event['previous_event_sha256'], events[index - 1]['event_sha256']) for operation_id, original in operations: replay = self.db._runtime_control_mutation( action=original['action'], target_ref=original['action'].split('.')[1], expected_revision=original['before']['revision'], actor='operator:alice', operation_id=operation_id, ) self.assertTrue(replay['replayed']) self.assertEqual(replay['audit_event_sha256'], original['audit_event_sha256']) self.assertEqual(self.counts(), (6, 6)) self.assertEqual(self.db.runtime_control_state()['revision'], 6) def test_stale_revision_has_no_durable_side_effects(self): with self.assertRaises(scanner_db.RuntimeControlRevisionConflictError) as raised: self.db.set_runtime_discovery_paused( True, expected_revision=1, actor='operator:alice', operation_id=self.operation_id(), ) self.assertEqual(raised.exception.expected_revision, 1) self.assertEqual(raised.exception.current_state['revision'], 0) self.assertEqual(self.counts(), (0, 0)) self.assertEqual(self.db.runtime_control_state()['revision'], 0) def test_exact_replay_is_idempotent_and_uuid_reuse_conflicts(self): operation_id = self.operation_id() first = self.db.set_runtime_discovery_paused( True, expected_revision=0, actor='operator:alice', operation_id=operation_id, ) replay = self.db.set_runtime_discovery_paused( True, expected_revision=0, actor='operator:alice', operation_id=operation_id, ) self.assertFalse(first['replayed']) self.assertTrue(replay['replayed']) self.assertEqual(first['after'], replay['after']) self.assertEqual(first['audit_event_id'], replay['audit_event_id']) self.assertEqual(self.counts(), (1, 1)) with self.assertRaises(scanner_db.RuntimeOperationIdentityConflictError): self.db.set_runtime_discovery_paused( True, expected_revision=0, actor='operator:bob', operation_id=operation_id, ) with self.assertRaises(scanner_db.RuntimeOperationIdentityConflictError): self.db.set_runtime_discovery_paused( True, expected_revision=1, actor='operator:alice', operation_id=operation_id, ) self.assertEqual(self.counts(), (1, 1)) def test_redundant_and_invalid_transitions_have_no_side_effects(self): with self.assertRaises(scanner_db.RuntimeControlTransitionError): self.db.set_runtime_discovery_paused( False, expected_revision=0, actor='operator:alice', operation_id=self.operation_id(), ) with self.assertRaises(scanner_db.RuntimeControlTransitionError): self.db.cancel_runtime_drain( expected_revision=0, actor='operator:alice', operation_id=self.operation_id(), ) self.assertEqual(self.counts(), (0, 0)) self.assertEqual(self.db.runtime_control_state()['revision'], 0) def test_inputs_are_validated_before_a_transaction(self): valid_id = self.operation_id() invalid_calls = ( lambda: self.db.set_runtime_discovery_paused( 1, expected_revision=0, actor='operator:alice', operation_id=valid_id, ), lambda: self.db.set_runtime_discovery_paused( True, expected_revision=True, actor='operator:alice', operation_id=valid_id, ), lambda: self.db.set_runtime_discovery_paused( True, expected_revision=0, actor='', operation_id=valid_id, ), lambda: self.db.set_runtime_discovery_paused( True, expected_revision=0, actor='operator\nalice', operation_id=valid_id, ), lambda: self.db.set_runtime_discovery_paused( True, expected_revision=0, actor='operator:alice', operation_id='00000000-0000-0000-0000-000000000000', ), lambda: self.db.set_runtime_discovery_paused( True, expected_revision=0, actor='operator:alice', operation_id=valid_id.upper(), ), ) for call in invalid_calls: with self.subTest(call=call): with self.assertRaises(ValueError): call() self.assertEqual(self.counts(), (0, 0)) def test_audit_failure_rolls_back_operation_and_control(self): with mock.patch.object( self.db.conn, 'insert_returning_id', side_effect=RuntimeError('injected audit failure'), ): with self.assertRaisesRegex(RuntimeError, 'injected audit failure'): self.db.set_runtime_dispatch_paused( True, expected_revision=0, actor='operator:alice', operation_id=self.operation_id(), ) self.assertEqual(self.counts(), (0, 0)) state = self.db.runtime_control_state() self.assertEqual(state['revision'], 0) self.assertFalse(state['dispatch_paused']) def test_corrupt_replay_evidence_fails_closed(self): operation_id = self.operation_id() self.db.start_runtime_drain( expected_revision=0, actor='operator:alice', operation_id=operation_id, ) stored = self.db.conn.execute( '''SELECT resulting_identity_json FROM runtime_operations WHERE operation_id = ?''', (operation_id,), ).fetchone()['resulting_identity_json'] self.db.conn.execute( '''UPDATE runtime_operations SET resulting_identity_json = ? WHERE operation_id = ?''', (json.dumps(json.loads(stored), indent=2), operation_id), ) self.db.conn.commit() with self.assertRaisesRegex( scanner_db.RuntimeSafetySchemaError, 'not canonical', ): self.db.start_runtime_drain( expected_revision=0, actor='operator:alice', operation_id=operation_id, ) self.assertEqual(self.db.runtime_control_state()['revision'], 1) self.assertEqual(self.counts(), (1, 1)) def test_control_state_and_replay_survive_restart(self): operation_id = self.operation_id() original = self.db.set_runtime_dispatch_paused( True, expected_revision=0, actor='operator:alice', operation_id=operation_id, ) self.db.close() self.db = ScannerDB(db_path=self.path) state = self.db.runtime_control_state() self.assertEqual(state['revision'], 1) self.assertTrue(state['dispatch_paused']) replay = self.db.set_runtime_dispatch_paused( True, expected_revision=0, actor='operator:alice', operation_id=operation_id, ) self.assertTrue(replay['replayed']) self.assertEqual(replay['audit_event_sha256'], original['audit_event_sha256']) self.assertEqual(self.counts(), (1, 1)) def test_drain_progress_counts_only_live_remote_and_precommit_bundles(self): remote_id = self.add_reservation(assignment_kind='remote', resolved=False) self.add_reservation(assignment_kind='remote', resolved=True) self.add_reservation(assignment_kind='local', resolved=False) bundle_reservation_id = self.add_bundle('ready') progress = self.db.runtime_drain_progress() self.assertEqual(progress['live_remote_assignments'], 1) self.assertEqual(progress['precommit_result_bundles'], 1) self.assertEqual(progress['blocker_count'], 2) self.db.conn.execute( 'UPDATE result_reservations SET remote_resolved_at = ? WHERE id = ?', (scanner_db.utc_now_iso(), remote_id), ) for state, expected in ( ('ingesting', 1), ('db_committed', 0), ('acknowledged', 0), ('quarantined', 0), ): self.db.conn.execute( 'UPDATE result_bundles SET state = ? WHERE reservation_id = ?', (state, bundle_reservation_id), ) self.db.conn.commit() progress = self.db.runtime_drain_progress() self.assertEqual(progress['live_remote_assignments'], 0) self.assertEqual(progress['precommit_result_bundles'], expected) self.assertEqual(progress['blocker_count'], expected) def test_drain_reconciliation_completes_once_and_preserves_explicit_pauses(self): operation_id = self.operation_id() self.db.set_runtime_discovery_paused( True, expected_revision=0, actor='operator:alice', operation_id=self.operation_id(), ) self.db.set_runtime_dispatch_paused( True, expected_revision=1, actor='operator:alice', operation_id=self.operation_id(), ) self.db.start_runtime_drain( expected_revision=2, actor='operator:alice', operation_id=operation_id, ) completed = self.db.reconcile_runtime_drain() self.assertEqual(completed['action'], 'control.drain.complete') self.assertFalse(completed['replayed']) self.assertEqual(completed['before']['drain_state'], 'draining') self.assertEqual(completed['after']['drain_state'], 'drained') self.assertTrue(completed['after']['discovery_paused']) self.assertTrue(completed['after']['dispatch_paused']) self.assertEqual(completed['after']['revision'], 4) stored = self.db.conn.execute( 'SELECT actor, action FROM runtime_operations WHERE operation_id = ?', (completed['operation_id'],), ).fetchone() self.assertEqual(stored['actor'], 'system:drain-reconciler') self.assertEqual(stored['action'], 'control.drain.complete') self.assertEqual(self.counts(), (4, 4)) self.assertIsNone(self.db.reconcile_runtime_drain()) self.assertEqual(self.counts(), (4, 4)) replay = self.db._runtime_control_mutation( action='control.drain.complete', target_ref='drain', expected_revision=3, actor='system:drain-reconciler', operation_id=completed['operation_id'], ) self.assertTrue(replay['replayed']) self.db.cancel_runtime_drain( expected_revision=4, actor='operator:alice', operation_id=self.operation_id(), ) state = self.db.runtime_control_state() self.assertEqual(state['drain_state'], 'normal') self.assertTrue(state['discovery_paused']) self.assertTrue(state['dispatch_paused']) def test_blockers_delay_completion_and_audit_failure_rolls_back(self): self.db.start_runtime_drain( expected_revision=0, actor='operator:alice', operation_id=self.operation_id(), ) remote_id = self.add_reservation(assignment_kind='remote', resolved=False) self.assertIsNone(self.db.reconcile_runtime_drain()) self.assertEqual(self.db.runtime_control_state()['drain_state'], 'draining') self.assertEqual(self.counts(), (1, 1)) self.db.conn.execute( 'UPDATE result_reservations SET remote_resolved_at = ? WHERE id = ?', (scanner_db.utc_now_iso(), remote_id), ) self.db.conn.commit() with mock.patch.object( self.db.conn, 'insert_returning_id', side_effect=RuntimeError('audit failed'), ): with self.assertRaisesRegex(RuntimeError, 'audit failed'): self.db.reconcile_runtime_drain() state = self.db.runtime_control_state() self.assertEqual(state['revision'], 1) self.assertEqual(state['drain_state'], 'draining') self.assertEqual(self.counts(), (1, 1)) def test_cancelled_drain_is_not_reconciled(self): self.db.start_runtime_drain( expected_revision=0, actor='operator:alice', operation_id=self.operation_id(), ) self.db.cancel_runtime_drain( expected_revision=1, actor='operator:alice', operation_id=self.operation_id(), ) self.assertIsNone(self.db.reconcile_runtime_drain()) self.assertEqual(self.db.runtime_control_state()['revision'], 2) self.assertEqual(self.counts(), (2, 2)) def test_offline_migration_allows_previous_image_after_protocol2_drain(self): for table in ( 'runtime_audit_events', 'runtime_operations_control', 'runtime_operations', 'remote_worker_devices', 'remote_worker_users', ): self.db.conn.execute(f'DROP TABLE {table}') self.db.conn.execute( 'DELETE FROM runtime_schema_migrations WHERE version IN (?, ?)', ( scanner_db.REMOTE_WORKER_MIGRATION, scanner_db.OPERATIONS_CONTROL_MIGRATION, ), ) self.db.conn.commit() self.assertTrue(scanner_db.migrate_runtime_safety_schema(self.db)) remote_id = self.add_reservation(assignment_kind='remote', resolved=False) bundle_reservation_id = self.add_bundle('ready') self.db.start_runtime_drain( expected_revision=0, actor='operator:rollback', operation_id=self.operation_id(), ) self.assertIsNone(self.db.reconcile_runtime_drain()) now = scanner_db.utc_now_iso() self.db.conn.execute( 'UPDATE result_reservations SET remote_resolved_at = ? WHERE id = ?', (now, remote_id), ) self.db.conn.execute( "UPDATE result_bundles SET state = 'db_committed' WHERE reservation_id = ?", (bundle_reservation_id,), ) self.db.conn.commit() completed = self.db.reconcile_runtime_drain() self.assertEqual(completed['after']['drain_state'], 'drained') self.assertEqual(self.db.runtime_drain_progress()['blocker_count'], 0) operation_counts = self.counts() self.db.close() legacy_target = 'https://example.invalid/previous-image' previous_image = sqlite3.connect(self.path) try: previous_image.execute( '''INSERT INTO target_queue( source, platform, query, target, normalized_target, status, created_at, updated_at ) VALUES ('github', 'github', 'rollback-fixture', ?, ?, 'pending', ?, ?)''', (legacy_target, legacy_target, now, now), ) previous_image.commit() finally: previous_image.close() self.db = ScannerDB(db_path=self.path, initialize=False) self.assertEqual(self.db.runtime_control_state()['drain_state'], 'drained') self.assertEqual(self.counts(), operation_counts) legacy = self.db.conn.execute( 'SELECT status FROM target_queue WHERE normalized_target = ?', (legacy_target,), ).fetchone() self.assertEqual(legacy['status'], 'pending') for table in ( 'remote_worker_users', 'remote_worker_devices', 'runtime_operations', 'runtime_operations_control', 'runtime_audit_events', ): with self.subTest(table=table): self.assertTrue(self.db.conn.table_exists(table)) if __name__ == '__main__': unittest.main()