190 lines
7.4 KiB
Python
190 lines
7.4 KiB
Python
import json
|
|
import os
|
|
from pathlib import Path
|
|
import sys
|
|
import tempfile
|
|
from types import SimpleNamespace
|
|
import unittest
|
|
from unittest import mock
|
|
from datetime import datetime, timezone
|
|
|
|
|
|
ROOT = Path(__file__).resolve().parents[1]
|
|
APP_DIR = ROOT / 'app'
|
|
sys.path.insert(0, str(APP_DIR))
|
|
|
|
import console_runner
|
|
import scanner_db
|
|
from scanner_db import ScannerDB
|
|
|
|
|
|
class PublicationLeaseDatabaseTests(unittest.TestCase):
|
|
def setUp(self):
|
|
self.environment = mock.patch.dict(os.environ, {
|
|
'SCANNER_DB_URL': '',
|
|
'DATABASE_URL': '',
|
|
'TRUF_MANAGED_POSTGRES_DSN': '',
|
|
})
|
|
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()
|
|
|
|
def insert_publication(self, status='pending', owner=None, expires_at=None):
|
|
now = '2026-07-19T00:00:00+00:00'
|
|
target_scan_id = self.db.conn.insert_returning_id(
|
|
'INSERT INTO target_scans (raw_result_json, created_at) VALUES (?, ?)',
|
|
(json.dumps({'scan_event_id': 'event-1'}), now),
|
|
)
|
|
outbox_id = self.db.conn.insert_returning_id(
|
|
'''INSERT INTO scan_publication_outbox (
|
|
target_scan_id, payload_json, status, attempts, lease_owner,
|
|
lease_expires_at, created_at, updated_at
|
|
) VALUES (?, '', ?, 1, ?, ?, ?, ?)''',
|
|
(target_scan_id, status, owner, expires_at, now, now),
|
|
)
|
|
self.db.conn.commit()
|
|
return outbox_id
|
|
|
|
def test_renewal_is_fenced_by_row_owner_and_delivering_status_and_commits(self):
|
|
outbox_id = self.insert_publication(
|
|
status='delivering',
|
|
owner='owner-a',
|
|
expires_at='2026-07-19T00:05:00+00:00',
|
|
)
|
|
clock = [datetime(2026, 7, 19, 0, 1, tzinfo=timezone.utc).timestamp()]
|
|
|
|
with mock.patch.object(scanner_db, 'utc_now_iso', side_effect=lambda: datetime.fromtimestamp(
|
|
clock[0], timezone.utc,
|
|
).isoformat(timespec='seconds')), mock.patch.object(
|
|
scanner_db.time, 'time', side_effect=lambda: clock[0],
|
|
):
|
|
self.assertFalse(self.db.renew_scan_publication(outbox_id, 'owner-b', 300))
|
|
self.assertTrue(self.db.renew_scan_publication(outbox_id, 'owner-a', 300))
|
|
|
|
observer = ScannerDB(db_path=self.path, initialize=False)
|
|
try:
|
|
row = observer.conn.execute(
|
|
'SELECT status, lease_owner, lease_expires_at FROM scan_publication_outbox WHERE id = ?',
|
|
(outbox_id,),
|
|
).fetchone()
|
|
self.assertEqual(row['lease_expires_at'], '2026-07-19T00:06:00+00:00')
|
|
observer.conn.execute(
|
|
"UPDATE scan_publication_outbox SET status = 'pending' WHERE id = ?",
|
|
(outbox_id,),
|
|
)
|
|
observer.conn.commit()
|
|
finally:
|
|
observer.close()
|
|
|
|
self.assertFalse(self.db.renew_scan_publication(outbox_id, 'owner-a', 300))
|
|
|
|
def test_renewal_prevents_reclaim_after_original_300_second_expiry(self):
|
|
self.insert_publication()
|
|
start = datetime(2026, 7, 19, 0, 0, tzinfo=timezone.utc).timestamp()
|
|
clock = [start]
|
|
|
|
def now_iso():
|
|
return datetime.fromtimestamp(clock[0], timezone.utc).isoformat(timespec='seconds')
|
|
|
|
with mock.patch.object(scanner_db, 'utc_now_iso', side_effect=now_iso), mock.patch.object(
|
|
scanner_db.time, 'time', side_effect=lambda: clock[0],
|
|
):
|
|
claimed = self.db.claim_scan_publications('owner-a', 1, lease_seconds=300)
|
|
self.assertEqual(len(claimed), 1)
|
|
outbox_id = claimed[0]['id']
|
|
|
|
clock[0] = start + 200
|
|
self.assertTrue(self.db.renew_scan_publication(outbox_id, 'owner-a', 300))
|
|
|
|
clock[0] = start + 301
|
|
self.assertEqual(
|
|
self.db.claim_scan_publications('owner-b', 1, lease_seconds=300),
|
|
[],
|
|
)
|
|
|
|
row = self.db.conn.execute(
|
|
'SELECT status, lease_owner, lease_expires_at FROM scan_publication_outbox WHERE id = ?',
|
|
(outbox_id,),
|
|
).fetchone()
|
|
self.assertEqual((row['status'], row['lease_owner']), ('delivering', 'owner-a'))
|
|
self.assertEqual(row['lease_expires_at'], '2026-07-19T00:08:20+00:00')
|
|
|
|
|
|
class PublicationLeaseDrainTests(unittest.TestCase):
|
|
@staticmethod
|
|
def parent_db(finish_result=True):
|
|
class DB:
|
|
conn = SimpleNamespace(is_postgres=True)
|
|
path = None
|
|
url = 'postgresql://truf:secret@127.0.0.1:5432/truf'
|
|
|
|
def __init__(self):
|
|
self.claim_count = 0
|
|
self.finishes = []
|
|
|
|
def claim_scan_publications(self, owner, limit):
|
|
self.claim_count += 1
|
|
return [{'id': 7, 'payload_json': '{}'}]
|
|
|
|
def finish_scan_publication(self, outbox_id, owner, delivered, error):
|
|
self.finishes.append((outbox_id, owner, delivered, error))
|
|
return finish_result
|
|
|
|
return DB()
|
|
|
|
@staticmethod
|
|
def heartbeat_db():
|
|
return SimpleNamespace(
|
|
enabled=True,
|
|
conn=SimpleNamespace(
|
|
is_postgres=True,
|
|
execute=mock.Mock(),
|
|
commit=mock.Mock(),
|
|
),
|
|
last_error='',
|
|
require_runtime_safety_schema=mock.Mock(return_value=True),
|
|
renew_scan_publication=mock.Mock(return_value=True),
|
|
close=mock.Mock(),
|
|
)
|
|
|
|
def test_thread_setup_failure_requeues_before_publication(self):
|
|
parent = self.parent_db()
|
|
heartbeat = self.heartbeat_db()
|
|
with mock.patch.object(console_runner, 'ScannerDB', return_value=heartbeat), mock.patch.object(
|
|
console_runner.threading.Thread, 'start', side_effect=RuntimeError('thread start failed'),
|
|
), mock.patch.object(console_runner, 'publish_scan_payload') as publish:
|
|
self.assertEqual(console_runner.drain_scan_publication_outbox(parent, 1), 0)
|
|
|
|
publish.assert_not_called()
|
|
heartbeat.require_runtime_safety_schema.assert_called_once_with()
|
|
heartbeat.renew_scan_publication.assert_called_once()
|
|
heartbeat.close.assert_called()
|
|
self.assertEqual(len(parent.finishes), 1)
|
|
self.assertFalse(parent.finishes[0][2])
|
|
self.assertIn('thread start failed', parent.finishes[0][3])
|
|
|
|
def test_final_fence_loss_is_a_nonfatal_handoff_and_stops_drain(self):
|
|
parent = self.parent_db(finish_result=False)
|
|
heartbeat = self.heartbeat_db()
|
|
with mock.patch.object(console_runner, 'ScannerDB', return_value=heartbeat), mock.patch.object(
|
|
console_runner, 'publish_scan_payload', return_value=(True, ''),
|
|
) as publish, self.assertLogs(console_runner.logger, level='WARNING') as logs:
|
|
self.assertEqual(console_runner.drain_scan_publication_outbox(parent, 5), 0)
|
|
|
|
publish.assert_called_once_with({})
|
|
self.assertEqual(parent.claim_count, 1)
|
|
self.assertEqual(len(parent.finishes), 1)
|
|
self.assertTrue(parent.finishes[0][2])
|
|
self.assertTrue(any('ownership handoff' in message.lower() for message in logs.output))
|
|
heartbeat.close.assert_called()
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main()
|