268 lines
11 KiB
Python
268 lines
11 KiB
Python
from concurrent.futures import Future
|
|
import os
|
|
from pathlib import Path
|
|
import sys
|
|
import tempfile
|
|
import unittest
|
|
|
|
|
|
ROOT = Path(__file__).resolve().parents[1]
|
|
APP_DIR = ROOT / 'app'
|
|
sys.path.insert(0, str(APP_DIR))
|
|
|
|
import postgres_runtime
|
|
from runtime_security import ensure_private_directory, private_file_ready
|
|
|
|
|
|
class ImmediateExecutor:
|
|
def submit(self, function):
|
|
future = Future()
|
|
try:
|
|
future.set_result(function())
|
|
except BaseException as exc:
|
|
future.set_exception(exc)
|
|
return future
|
|
|
|
|
|
class ControlledStartExecutor:
|
|
def __init__(self):
|
|
self.submissions = []
|
|
|
|
def submit(self, function):
|
|
future = Future()
|
|
self.submissions.append((function.__name__, future))
|
|
if function.__name__ == 'stop':
|
|
try:
|
|
future.set_result(function())
|
|
except BaseException as exc:
|
|
future.set_exception(exc)
|
|
return future
|
|
|
|
|
|
class ScriptedBackend:
|
|
def __init__(self, probes, starts=()):
|
|
self.probes = list(probes)
|
|
self.starts = list(starts)
|
|
self.start_calls = 0
|
|
self.stop_calls = 0
|
|
self.closed = False
|
|
|
|
def probe(self):
|
|
return self.probes.pop(0)
|
|
|
|
def start(self):
|
|
self.start_calls += 1
|
|
return self.starts.pop(0)
|
|
|
|
def stop(self):
|
|
self.stop_calls += 1
|
|
return postgres_runtime.StopResult(True, True, 'stopped')
|
|
|
|
def close(self):
|
|
self.closed = True
|
|
|
|
|
|
@unittest.skipUnless(os.name == 'nt', 'Windows ACL inheritance semantics required')
|
|
class PostgresCollectorAclTests(unittest.TestCase):
|
|
def make_logging_backend(self, postgres_dir):
|
|
backend = postgres_runtime.PostgresBackend.__new__(postgres_runtime.PostgresBackend)
|
|
backend.paths = {
|
|
'log_dir': os.path.join(postgres_dir, 'logs'),
|
|
'log_path': os.path.join(postgres_dir, 'postgres.log'),
|
|
}
|
|
backend.log_max_mb = 1
|
|
backend.log_keep = 2
|
|
return backend
|
|
|
|
def test_prepare_hardens_recreated_startup_log_with_inherited_acl(self):
|
|
with tempfile.TemporaryDirectory() as temp_dir:
|
|
postgres_dir = os.path.join(temp_dir, 'postgres')
|
|
ensure_private_directory(postgres_dir, reject_reparse=True)
|
|
backend = self.make_logging_backend(postgres_dir)
|
|
backend._prepare_logging()
|
|
path = backend.paths['log_path']
|
|
os.remove(path)
|
|
Path(path).write_text('pg_ctl startup output', encoding='ascii')
|
|
|
|
self.assertFalse(private_file_ready(path))
|
|
backend._prepare_logging()
|
|
self.assertTrue(private_file_ready(path))
|
|
|
|
def test_prepare_rejects_startup_log_hard_link_without_hardening_target(self):
|
|
with tempfile.TemporaryDirectory() as temp_dir:
|
|
postgres_dir = os.path.join(temp_dir, 'postgres')
|
|
ensure_private_directory(postgres_dir, reject_reparse=True)
|
|
backend = self.make_logging_backend(postgres_dir)
|
|
target = os.path.join(temp_dir, 'external.log')
|
|
Path(target).write_text('external', encoding='ascii')
|
|
os.link(target, backend.paths['log_path'])
|
|
|
|
with self.assertRaisesRegex(postgres_runtime.ClusterIdentityError, 'unsafe'):
|
|
backend._prepare_logging()
|
|
|
|
self.assertFalse(private_file_ready(target))
|
|
self.assertEqual(Path(target).read_text(encoding='ascii'), 'external')
|
|
|
|
def test_prune_hardens_normally_created_collector_file(self):
|
|
with tempfile.TemporaryDirectory() as temp_dir:
|
|
log_dir = os.path.join(temp_dir, 'collector')
|
|
ensure_private_directory(log_dir, reject_reparse=True)
|
|
path = os.path.join(log_dir, 'postgresql-20260719-123456.log')
|
|
Path(path).write_text('active log', encoding='ascii')
|
|
|
|
self.assertFalse(private_file_ready(path))
|
|
self.assertEqual(postgres_runtime.prune_postgres_collector_logs(log_dir, 2), 0)
|
|
self.assertTrue(private_file_ready(path))
|
|
|
|
def test_prune_rejects_collector_hard_link_without_hardening_target(self):
|
|
with tempfile.TemporaryDirectory() as temp_dir:
|
|
log_dir = os.path.join(temp_dir, 'collector')
|
|
ensure_private_directory(log_dir, reject_reparse=True)
|
|
target = os.path.join(temp_dir, 'external.log')
|
|
Path(target).write_text('external', encoding='ascii')
|
|
link = os.path.join(log_dir, 'postgresql-20260719-123456.log')
|
|
os.link(target, link)
|
|
|
|
with self.assertRaisesRegex(postgres_runtime.ClusterIdentityError, 'unsafe'):
|
|
postgres_runtime.prune_postgres_collector_logs(log_dir, 2)
|
|
|
|
self.assertFalse(private_file_ready(target))
|
|
self.assertEqual(Path(target).read_text(encoding='ascii'), 'external')
|
|
|
|
|
|
class PostgresOwnedLifecycleTests(unittest.TestCase):
|
|
def make_controller(self, backend):
|
|
return postgres_runtime.PostgresController(
|
|
backend,
|
|
executor=ImmediateExecutor(),
|
|
health_interval_sec=1,
|
|
stable_ready_interval_sec=0,
|
|
)
|
|
|
|
def pending_start(self):
|
|
backend = ScriptedBackend([])
|
|
executor = ControlledStartExecutor()
|
|
controller = postgres_runtime.PostgresController(backend, executor=executor)
|
|
controller.tick(0)
|
|
executor.submissions[0][1].set_result(
|
|
postgres_runtime.ProbeResult(postgres_runtime.ProbeKind.STOPPED, 'initially offline')
|
|
)
|
|
controller.tick(0)
|
|
self.assertEqual([name for name, _ in executor.submissions], ['probe', 'start'])
|
|
controller.request_stop()
|
|
return controller, backend, executor
|
|
|
|
def test_shutdown_after_rejected_stale_start_closes_without_stop(self):
|
|
results = (
|
|
postgres_runtime.StartResult(False, 'start refused while cluster state is READY'),
|
|
postgres_runtime.StartResult(
|
|
False,
|
|
'PostgreSQL logging setup failed while cluster remained offline',
|
|
foreign_or_config_error=True,
|
|
),
|
|
)
|
|
for result in results:
|
|
with self.subTest(detail=result.detail):
|
|
controller, backend, executor = self.pending_start()
|
|
executor.submissions[1][1].set_result(result)
|
|
|
|
controller.tick(0)
|
|
|
|
self.assertEqual([name for name, _ in executor.submissions], ['probe', 'start'])
|
|
self.assertEqual(backend.stop_calls, 0)
|
|
self.assertTrue(controller.terminal)
|
|
self.assertTrue(controller.authority_release_safe)
|
|
self.assertFalse(controller.lifecycle_action_required)
|
|
self.assertIn('no controller-owned side effect', controller.detail)
|
|
self.assertTrue(controller.close(timeout_sec=0.1))
|
|
self.assertTrue(backend.closed)
|
|
|
|
def test_shutdown_after_accepted_or_uncertain_stale_start_compensates_once(self):
|
|
results = (
|
|
postgres_runtime.StartResult(True, 'accepted'),
|
|
postgres_runtime.StartResult(True, 'outcome uncertain', uncertain=True),
|
|
)
|
|
for result in results:
|
|
with self.subTest(detail=result.detail):
|
|
controller, backend, executor = self.pending_start()
|
|
executor.submissions[1][1].set_result(result)
|
|
|
|
controller.tick(0)
|
|
|
|
self.assertEqual([name for name, _ in executor.submissions], ['probe', 'start', 'stop'])
|
|
self.assertEqual(backend.stop_calls, 1)
|
|
self.assertFalse(controller.authority_release_safe)
|
|
controller.tick(0)
|
|
self.assertEqual(controller.state, postgres_runtime.PostgresState.STOPPED)
|
|
self.assertTrue(controller.authority_release_safe)
|
|
|
|
def test_shutdown_after_stale_start_exception_fails_closed_with_stop(self):
|
|
controller, backend, executor = self.pending_start()
|
|
executor.submissions[1][1].set_exception(RuntimeError('uncertain worker failure'))
|
|
|
|
controller.tick(0)
|
|
|
|
self.assertEqual([name for name, _ in executor.submissions], ['probe', 'start', 'stop'])
|
|
self.assertEqual(backend.stop_calls, 1)
|
|
self.assertFalse(controller.authority_release_safe)
|
|
controller.tick(0)
|
|
self.assertTrue(controller.authority_release_safe)
|
|
|
|
def test_owned_ready_recovery_death_backs_off_and_starts_again(self):
|
|
backend = ScriptedBackend(
|
|
probes=[
|
|
postgres_runtime.ProbeResult(postgres_runtime.ProbeKind.STOPPED, 'initially stopped'),
|
|
postgres_runtime.ProbeResult(postgres_runtime.ProbeKind.READY, 'owned ready', 'epoch-1'),
|
|
postgres_runtime.ProbeResult(postgres_runtime.ProbeKind.RECOVERING, 'owned recovery', 'epoch-1'),
|
|
postgres_runtime.ProbeResult(postgres_runtime.ProbeKind.STOPPED, 'owned cluster stopped'),
|
|
postgres_runtime.ProbeResult(postgres_runtime.ProbeKind.STOPPED, 'backoff verification'),
|
|
postgres_runtime.ProbeResult(postgres_runtime.ProbeKind.READY, 'restarted', 'epoch-2'),
|
|
],
|
|
starts=[
|
|
postgres_runtime.StartResult(True, 'first start accepted'),
|
|
postgres_runtime.StartResult(True, 'second start accepted'),
|
|
],
|
|
)
|
|
controller = self.make_controller(backend)
|
|
|
|
for now in (0, 0, 0, 0):
|
|
controller.tick(now)
|
|
self.assertEqual(controller.state, postgres_runtime.PostgresState.READY)
|
|
self.assertTrue(controller._owned_start)
|
|
|
|
for now in (1, 1, 2, 2):
|
|
controller.tick(now)
|
|
self.assertEqual(controller.state, postgres_runtime.PostgresState.BACKOFF)
|
|
self.assertFalse(controller._owned_start)
|
|
|
|
for now in (32, 32, 32, 32):
|
|
controller.tick(now)
|
|
self.assertEqual(controller.state, postgres_runtime.PostgresState.READY)
|
|
self.assertEqual(backend.start_calls, 2)
|
|
self.assertEqual(backend.stop_calls, 0)
|
|
|
|
def test_preexisting_cluster_remains_inert_after_ready_and_recovery(self):
|
|
backend = ScriptedBackend(probes=[
|
|
postgres_runtime.ProbeResult(postgres_runtime.ProbeKind.READY, 'preexisting ready', 'epoch-1'),
|
|
postgres_runtime.ProbeResult(postgres_runtime.ProbeKind.RECOVERING, 'preexisting recovery', 'epoch-1'),
|
|
postgres_runtime.ProbeResult(postgres_runtime.ProbeKind.STOPPED, 'preexisting cluster stopped'),
|
|
postgres_runtime.ProbeResult(postgres_runtime.ProbeKind.STOPPED, 'still stopped'),
|
|
])
|
|
controller = self.make_controller(backend)
|
|
|
|
controller.tick(0)
|
|
controller.tick(0)
|
|
self.assertEqual(controller.state, postgres_runtime.PostgresState.READY)
|
|
self.assertTrue(controller.snapshot()['lifecycle_inert'])
|
|
self.assertFalse(controller._owned_start)
|
|
|
|
for now in (1, 1, 2, 2, 3, 3):
|
|
controller.tick(now)
|
|
self.assertEqual(controller.state, postgres_runtime.PostgresState.FOREIGN_OR_CONFIG_ERROR)
|
|
self.assertEqual(backend.start_calls, 0)
|
|
self.assertEqual(backend.stop_calls, 0)
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main()
|