1import copy
2import json
3import os
4from pathlib import Path
5import tempfile
6from types import SimpleNamespace
7import unittest
8from unittest.mock import patch
9
10from native_disconnect import interrupt, verify_disconnect
11
12
13class 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
143if __name__ == '__main__': unittest.main()