514 lines
22 KiB
Python
514 lines
22 KiB
Python
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()
|