1import unittest
2
3from verify_smb_overlap import verify
4
5
6class 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
143if __name__ == '__main__':
144 unittest.main()