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

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