Initial server source import

This commit is contained in:
sashatrask
2026-09-30 20:30:56 +03:00
commit 170dd941b9
498 changed files with 261563 additions and 0 deletions
+189
View File
@@ -0,0 +1,189 @@
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()