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

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