| 1 | import copy |
| 2 | import json |
| 3 | import os |
| 4 | from pathlib import Path |
| 5 | import tempfile |
| 6 | from types import SimpleNamespace |
| 7 | import unittest |
| 8 | from unittest.mock import patch |
| 9 | |
| 10 | from native_disconnect import interrupt, verify_disconnect |
| 11 | |
| 12 | |
| 13 | class DisconnectOracle(unittest.TestCase): |
| 14 | def setUp(self): |
| 15 | temporary = tempfile.TemporaryDirectory() |
| 16 | self.addCleanup(temporary.cleanup) |
| 17 | self.root = Path(temporary.name) |
| 18 | (self.root / 'rust').mkdir() |
| 19 | self.config = {'stress_clients': 4, 'rust_writers': 4, 'rust_readers': 4, 'stress_operations': 80} |
| 20 | for name, stamp in [('start', 0), ('stop', 10**9)]: |
| 21 | path = self.root / 'rust' / name |
| 22 | path.touch() |
| 23 | os.utime(path, ns=(stamp, stamp)) |
| 24 | for i in range(4): |
| 25 | folder = self.root / f'n{i}' |
| 26 | folder.mkdir() |
| 27 | (folder / 'stress-events.jsonl').write_text('\n'.join(json.dumps({ |
| 28 | 'update_started_ticks': j * 10**7, 'updated_ticks': (j + 1) * 10**7}) for j in range(80))) |
| 29 | actors = [f'{role}{i}' for role in ('w', 'r') for i in range(4)] |
| 30 | self.samples = [{'cycle': i, 'before': dict.fromkeys(actors, 3 + 3 * i), |
| 31 | 'after': dict.fromkeys(actors, 6 + 3 * i), 'native_before': [5, 6, 7, 8], |
| 32 | 'before_errors': dict.fromkeys(actors, i), 'after_errors': dict.fromkeys(actors, i + 1)} |
| 33 | for i in range(2)] |
| 34 | for actor in actors: |
| 35 | (self.root / 'rust' / (actor + '.jsonl')).write_text('\n'.join(json.dumps({'event': event}) for event in |
| 36 | ['ready', 'transport_read_error', 'transport_connected', 'transport_read_error', 'transport_connected', 'done'])) |
| 37 | with (self.root / 'rust' / (actor + '.jsonl')).open('a') as stream: |
| 38 | stream.write('\n' + json.dumps({'event': 'commit' if actor.startswith('w') else 'read', 'finished_us': 1000000})) |
| 39 | for actor, published in [('w0', False), ('w1', True)]: |
| 40 | with (self.root / 'rust' / (actor + '.jsonl')).open('a') as stream: |
| 41 | stream.write('\n' + json.dumps({'event': 'transport_commit_error', 'state': 'Unknown', 'token': actor}) + '\n') |
| 42 | stream.write(json.dumps({'event': 'transport_reconciled', 'published': published, 'flush_confirmed': published, 'token': actor}) + '\n') |
| 43 | (self.root / 'smb-trace.jsonl').write_text('{"cut":{}}\n{"control":{"phase":"reconnected-0"}}\n{"control":{"phase":"disconnect-1"}}\n{"cut":{}}\n{"control":{"phase":"reconnected-1"}}\n') |
| 44 | |
| 45 | def check(self): |
| 46 | (self.root / 'run.json').write_text(json.dumps(self.config)) |
| 47 | (self.root / 'disconnect-progress.json').write_text(json.dumps(self.samples)) |
| 48 | with patch('native_disconnect.verify', return_value={'guarded': True}) as overlap: |
| 49 | result = verify_disconnect(self.root) |
| 50 | self.assertEqual([call.kwargs['phase'] for call in overlap.call_args_list], ['reconnected-0', 'reconnected-1']) |
| 51 | self.assertEqual([len(call.args[0]) for call in overlap.call_args_list], [2, 5]) |
| 52 | return result |
| 53 | |
| 54 | def test_both_interruptions_require_post_reconnect_overlap(self): |
| 55 | self.assertEqual(self.check()['interruptions'], 2) |
| 56 | |
| 57 | def test_stalled_and_backward_clients_are_rejected(self): |
| 58 | for actor in ['w3', 'r2']: |
| 59 | path = self.root / 'rust' / (actor + '.jsonl') |
| 60 | original = path.read_text() |
| 61 | for stamp in [120000001, -1]: |
| 62 | path.write_text(original.replace('"finished_us": 1000000', f'"finished_us": {stamp}')) |
| 63 | with self.assertRaisesRegex(AssertionError, 'progress stalled or went backwards'): self.check() |
| 64 | path.write_text(original) |
| 65 | path = self.root / 'n2/stress-events.jsonl' |
| 66 | rows = [json.loads(line) for line in path.read_text().splitlines()] |
| 67 | rows[30]['updated_ticks'] += 121 * 10**7 |
| 68 | path.write_text('\n'.join(map(json.dumps, rows))) |
| 69 | with self.assertRaisesRegex(AssertionError, 'n2: client progress'): self.check() |
| 70 | |
| 71 | def test_reader_progress_must_cover_the_stop_barrier(self): |
| 72 | stamp = 121000001000 |
| 73 | os.utime(self.root / 'rust/stop', ns=(stamp, stamp)) |
| 74 | with self.assertRaisesRegex(AssertionError, 'r0: client progress'): self.check() |
| 75 | |
| 76 | def test_one_idle_or_unaffected_client_is_rejected(self): |
| 77 | original = copy.deepcopy(self.samples) |
| 78 | for field, value in [('after', 3), ('after_errors', 0)]: |
| 79 | self.samples = copy.deepcopy(original) |
| 80 | self.samples[0][field]['r3'] = value |
| 81 | with self.assertRaises(AssertionError): self.check() |
| 82 | |
| 83 | def test_missing_actor_is_rejected(self): |
| 84 | for sample in self.samples: |
| 85 | for field in ('before', 'after', 'before_errors', 'after_errors'): del sample[field]['r3'] |
| 86 | with self.assertRaises(AssertionError): self.check() |
| 87 | |
| 88 | def test_completed_native_writer_is_rejected(self): |
| 89 | self.samples[0]['native_before'][2] = 80 |
| 90 | with self.assertRaises(AssertionError): self.check() |
| 91 | |
| 92 | def test_wrong_cut_count_is_rejected(self): |
| 93 | (self.root / 'smb-trace.jsonl').write_text('{"cut":{}}\n') |
| 94 | with self.assertRaises(AssertionError): self.check() |
| 95 | |
| 96 | def test_missing_reconnection_is_rejected(self): |
| 97 | (self.root / 'rust/r2.jsonl').write_text('{"event":"transport_read_error"}\n') |
| 98 | with self.assertRaises(AssertionError): self.check() |
| 99 | |
| 100 | def test_visibility_without_flush_confirmation_is_rejected(self): |
| 101 | path = self.root / 'rust/w1.jsonl' |
| 102 | path.write_text(path.read_text().replace('"flush_confirmed": true', '"flush_confirmed": false')) |
| 103 | with self.assertRaisesRegex(AssertionError, 'durable acknowledgement'): self.check() |
| 104 | |
| 105 | def test_one_uncertain_outcome_is_insufficient(self): |
| 106 | path = self.root / 'rust/w0.jsonl' |
| 107 | path.write_text(path.read_text().replace('"state": "Unknown"', '"state": "NotCommitted"')) |
| 108 | with self.assertRaisesRegex(AssertionError, 'both uncertain outcomes'): self.check() |
| 109 | |
| 110 | def test_failed_cut_observation_restores_the_connection(self): |
| 111 | (self.root / 'run.json').write_text(json.dumps({**self.config, 'server': 'owned-test'})) |
| 112 | (self.root / 'rust/w0.jsonl').write_text('{"event":"commit"}\n' * 3) |
| 113 | controls = [] |
| 114 | (self.root / 'harness').mkdir() |
| 115 | (self.root / 'harness/verify_smb_overlap.py').write_text('test oracle') |
| 116 | |
| 117 | def ssh(server, command, timeout): |
| 118 | self.assertEqual(server, 'owned-test') |
| 119 | if command.startswith('printf'): |
| 120 | controls.append(json.loads(command[command.index('{'):command.index('}') + 1])) |
| 121 | output = '' |
| 122 | elif (failed_phase_ack and controls[-1]['phase'] == 'disconnect-0') or command.startswith('grep -c'): |
| 123 | raise RuntimeError('Fault observation failed') |
| 124 | else: output = json.dumps({'control': controls[-1]}) |
| 125 | return SimpleNamespace(stdout=output, stderr='', returncode=0, check_returncode=lambda: None) |
| 126 | |
| 127 | def capture(remote, destination, name): |
| 128 | destination.write_text('{"operation":0}\n') |
| 129 | return {'error': None} |
| 130 | |
| 131 | for failed_phase_ack in [False, True]: |
| 132 | controls.clear() |
| 133 | with self.subTest(failed_phase_ack=failed_phase_ack), \ |
| 134 | patch('native_disconnect.linux_vm.run_ssh', side_effect=ssh), \ |
| 135 | patch('native_disconnect.linux_vm.ssh_argv', return_value=['ssh', 'owned-test']), \ |
| 136 | patch('native_disconnect.subprocess.run'), \ |
| 137 | patch('native_disconnect.windows.do_get', side_effect=capture): |
| 138 | with self.assertRaisesRegex(RuntimeError, 'Fault observation failed'): |
| 139 | interrupt(self.root, [{'name': 'native'}], [1], {'w0': SimpleNamespace(poll=lambda: None)}) |
| 140 | self.assertEqual([control['phase'] for control in controls], ['disconnect-0', 'reconnected-0']) |
| 141 | |
| 142 | |
| 143 | if __name__ == '__main__': unittest.main() |