| 1 | import unittest |
| 2 | |
| 3 | from verify_smb_overlap import verify |
| 4 | |
| 5 | |
| 6 | class Overlap(unittest.TestCase): |
| 7 | def history(self, serialized=False, failed_write=False, writer_reads=False): |
| 8 | events = [{'connection': 1, 'opened': True, 'peer': ['10.0.2.2', 1]}, |
| 9 | {'connection': 2, 'opened': True, 'peer': ['192.168.77.2', 2]}] |
| 10 | def exchange(connection, command, status='0x0', **fields): |
| 11 | message = len(events) |
| 12 | request = {'connection': connection, 'command': command, 'message': message, |
| 13 | 'direction': 'request', **fields} |
| 14 | response = {'connection': connection, 'command': command, 'message': message, |
| 15 | 'direction': 'response', 'status': status, 'file_id': str(connection)} |
| 16 | events.extend((request, response)) |
| 17 | for connection in (1, 2): |
| 18 | exchange(connection, 5, path='synthetic.one', access=0xc0000000 if connection == 2 or writer_reads else 0x80000000) |
| 19 | exchange(connection, 10, file_id=str(connection), locks=[(0xfffffffb, 1, 0x11)]) |
| 20 | if serialized: |
| 21 | exchange(1, 8, file_id='1', length=32) |
| 22 | exchange(1, 6, file_id='1') |
| 23 | exchange(2, 10, file_id='2', locks=[(0xfffffffd, 1, 0x12)]) |
| 24 | if not serialized: |
| 25 | exchange(1, 8, file_id='1', length=32) |
| 26 | exchange(2, 9, status='0xc0000054' if failed_write else '0x0', file_id='2', length=32) |
| 27 | exchange(2, 6, file_id='2') |
| 28 | events.append({'connection': 3, 'opened': True, 'peer': ['10.0.2.2', 3]}) |
| 29 | exchange(3, 5, path='synthetic.one', access=0xc0000000) |
| 30 | exchange(3, 10, file_id='3', locks=[(0xfffffffb, 1, 0x11), (0xfffffffd, 1, 0x12)]) |
| 31 | if not serialized: |
| 32 | exchange(1, 8, file_id='1', length=32) |
| 33 | exchange(3, 9, file_id='3', length=32) |
| 34 | return events |
| 35 | |
| 36 | def test_active_read_and_native_write(self): |
| 37 | result = verify(self.history()) |
| 38 | self.assertEqual(result['active_native_writer_pairs'], 1) |
| 39 | self.assertEqual(result['active_rust_writer_pairs'], 1) |
| 40 | self.assertEqual(result['overlapping_reads'], 2) |
| 41 | self.assertEqual(result['overlapping_writes'], 2) |
| 42 | |
| 43 | def test_progress_before_resume_does_not_satisfy_after_gate(self): |
| 44 | events = self.history() |
| 45 | events.append({'control': {'phase': 'resumed'}}) |
| 46 | with self.assertRaises(AssertionError): |
| 47 | verify(events, phase='resumed') |
| 48 | with self.assertRaises(AssertionError): |
| 49 | verify(self.history(), phase='resumed') |
| 50 | self.assertEqual(verify([{'control': {'phase': 'resumed'}}, *self.history()], phase='resumed')['active_rust_writer_pairs'], 1) |
| 51 | |
| 52 | def test_native_only_traffic_does_not_count(self): |
| 53 | events = self.history() |
| 54 | for event in events: |
| 55 | if event.get('opened'): event['peer'][0] = '192.168.77.' + str(event['connection'] + 1) |
| 56 | with self.assertRaises(AssertionError): |
| 57 | verify(events) |
| 58 | |
| 59 | def test_delayed_io_responses_do_not_count_as_resumed_requests(self): |
| 60 | events = self.history() |
| 61 | end = next(i for i, event in enumerate(events) if event.get('direction') == 'request' and event['command'] == 6) |
| 62 | responses = [event for event in events[:end] if event.get('direction') == 'response' and event['command'] in (8, 9)] |
| 63 | events = [event for event in events if event not in responses] |
| 64 | events[end - len(responses):end - len(responses)] = [{'control': {'phase': 'resumed'}}, *responses] |
| 65 | self.assertEqual(verify(events)['active_native_writer_pairs'], 1) |
| 66 | with self.assertRaises(AssertionError): |
| 67 | verify(events, phase='resumed') |
| 68 | |
| 69 | def test_serialized_calls_do_not_count(self): |
| 70 | with self.assertRaises(AssertionError): |
| 71 | verify(self.history(serialized=True)) |
| 72 | |
| 73 | def test_failed_write_does_not_count(self): |
| 74 | with self.assertRaises(AssertionError): |
| 75 | verify(self.history(failed_write=True)) |
| 76 | |
| 77 | def test_writers_validation_reads_do_not_count(self): |
| 78 | with self.assertRaises(AssertionError): |
| 79 | verify(self.history(writer_reads=True)) |
| 80 | |
| 81 | def test_lock_failure_does_not_count(self): |
| 82 | events = self.history() |
| 83 | for event in events: |
| 84 | if event.get('direction') == 'response' and event['command'] == 10: |
| 85 | event['status'] = '0xc0000055' |
| 86 | with self.assertRaises(AssertionError): |
| 87 | verify(events) |
| 88 | |
| 89 | def test_separate_lock_lifetimes_do_not_combine(self): |
| 90 | events = self.history() |
| 91 | inserted = [] |
| 92 | for message, flags in [(1000, 4), (1001, 0x12)]: |
| 93 | inserted.extend([ |
| 94 | {'connection': 3, 'direction': 'request', 'command': 10, 'message': message, |
| 95 | 'file_id': '3', 'locks': [(0xfffffffd, 1, flags)]}, |
| 96 | {'connection': 3, 'direction': 'response', 'command': 10, 'message': message, |
| 97 | 'status': '0x0'}, |
| 98 | ]) |
| 99 | events[-2:-2] = inserted |
| 100 | with self.assertRaises(AssertionError): |
| 101 | verify(events) |
| 102 | |
| 103 | def test_renamed_path_does_not_combine_file_versions(self): |
| 104 | events = self.history() |
| 105 | events[-2:-2] = [ |
| 106 | {'connection': 3, 'direction': 'request', 'command': 17, 'message': 1000, |
| 107 | 'info_type': 1, 'info_class': 10}, |
| 108 | {'connection': 3, 'direction': 'response', 'command': 17, 'message': 1000, 'status': '0x0'}, |
| 109 | ] |
| 110 | with self.assertRaises(AssertionError): |
| 111 | verify(events) |
| 112 | |
| 113 | def test_open_response_crossing_rename_is_not_attributed(self): |
| 114 | events = self.history() |
| 115 | at = next(i for i, event in enumerate(events) if event.get('connection') == 3 |
| 116 | and event.get('command') == 5 and event['direction'] == 'response') |
| 117 | inserted = [ |
| 118 | {'connection': 4, 'opened': True, 'peer': ['192.168.77.4', 4]}, |
| 119 | {'connection': 4, 'direction': 'request', 'command': 17, 'message': 1000, |
| 120 | 'info_type': 1, 'info_class': 10}, |
| 121 | {'connection': 4, 'direction': 'response', 'command': 17, 'message': 1000, 'status': '0x0'}, |
| 122 | ] |
| 123 | for message, command, fields in [ |
| 124 | (1001, 5, {'path': 'synthetic.one', 'access': 0x80000000}), |
| 125 | (1002, 10, {'file_id': '1', 'locks': [(0xfffffffb, 1, 0x11)]}), |
| 126 | ]: |
| 127 | inserted.extend([ |
| 128 | {'connection': 1, 'direction': 'request', 'command': command, 'message': message, **fields}, |
| 129 | {'connection': 1, 'direction': 'response', 'command': command, 'message': message, |
| 130 | 'status': '0x0', 'file_id': '1'}, |
| 131 | ]) |
| 132 | events[at:at] = inserted |
| 133 | with self.assertRaises(AssertionError): |
| 134 | verify(events) |
| 135 | |
| 136 | def test_connection_close_ends_guard(self): |
| 137 | events = self.history() |
| 138 | events.insert(-2, {'connection': 1, 'closed': True}) |
| 139 | with self.assertRaises(AssertionError): |
| 140 | verify(events) |
| 141 | |
| 142 | |
| 143 | if __name__ == '__main__': |
| 144 | unittest.main() |