1import importlib.util
2import asyncio
3import json
4from pathlib import Path
5import socket
6import struct
7import sys
8import tempfile
9import unittest
10
11spec = importlib.util.spec_from_file_location('smb_proxy', Path(__file__).with_name('smb-proxy.py'))
12proxy = importlib.util.module_from_spec(spec)
13spec.loader.exec_module(proxy)
14
15
16class HeaderTrace(unittest.TestCase):
17 def test_fields_follow_complete_payload_ranges(self):
18 header = bytearray(1024)
19 header[96:100] = (257).to_bytes(4, 'little')
20 header[212:228] = bytes(range(16))
21 header[228:236] = (999).to_bytes(8, 'little')
22 header[236:252] = bytes(range(16, 32))
23 expected = {'transactions': 257, 'version': bytes(range(16)).hex(),
24 'generation': 999, 'deny_read': bytes(range(16, 32)).hex()}
25 self.assertEqual(proxy.header_fields(0, header), expected)
26 self.assertEqual(proxy.header_fields(212, header[212:252]), {k:v for k,v in expected.items() if k != 'transactions'})
27 self.assertEqual(proxy.header_fields(96, header[96:99]), {})
28 self.assertEqual(proxy.header_fields(100, header[100:212]), {})
29 self.assertEqual(proxy.header_fields(228, header[228:251]), {'generation':999})
30 self.assertEqual(proxy.header_fields(252, bytes(1024)), {})
31
32
33class WriteTrace(unittest.IsolatedAsyncioTestCase):
34 async def test_opt_in_payload_and_partial_write_response_are_traced(self):
35 frames = []
36
37 async def respond(reader, writer):
38 try:
39 for _ in range(3):
40 prefix = await reader.readexactly(4)
41 frame = await reader.readexactly(int.from_bytes(prefix[1:], 'big'))
42 frames.append(frame)
43 reply = bytearray(80)
44 reply[:64] = frame[:64]
45 struct.pack_into('<I', reply, 16, 1)
46 struct.pack_into('<I', reply, 68, 3)
47 writer.write(len(reply).to_bytes(4, 'big') + reply)
48 await writer.drain()
49 finally:
50 writer.close()
51 await writer.wait_closed()
52
53 server = await asyncio.start_server(respond, '127.0.0.1', 0)
54 server_port = server.sockets[0].getsockname()[1]
55 with socket.socket() as reservation:
56 reservation.bind(('127.0.0.1', 0))
57 proxy_port = reservation.getsockname()[1]
58 with tempfile.TemporaryDirectory() as directory:
59 control = Path(directory) / 'control.json'
60 control.write_text('{}')
61 process = await asyncio.create_subprocess_exec(
62 sys.executable, str(Path(proxy.__file__)), str(control),
63 '--port', str(proxy_port), '--server', '127.0.0.1', '--server-port', str(server_port),
64 stdout=asyncio.subprocess.PIPE, stderr=asyncio.subprocess.PIPE)
65 records = []
66 writer = None
67 try:
68 while not any('control' in row for row in records):
69 records.append(json.loads(await asyncio.wait_for(process.stdout.readline(), 5)))
70 reader, writer = await asyncio.open_connection('127.0.0.1', proxy_port)
71 sent = []
72 for message, command in enumerate((9, 9, 17)):
73 if message == 1:
74 control.write_text('{"record_writes":true}')
75 while not any(row.get('control', {}).get('record_writes') for row in records):
76 records.append(json.loads(await asyncio.wait_for(process.stdout.readline(), 5)))
77 if message == 2:
78 control.write_text('{}')
79 while True:
80 row = json.loads(await asyncio.wait_for(process.stdout.readline(), 5))
81 records.append(row)
82 if row.get('control') == {}:
83 break
84 data = b'abcde' if command == 9 else (4096).to_bytes(8, 'little')
85 offset = 112 if command == 9 else 96
86 frame = bytearray(offset)
87 frame[:4] = b'\xfeSMB'
88 struct.pack_into('<H', frame, 12, command)
89 struct.pack_into('<Q', frame, 24, message)
90 frame[80:96] = bytes(range(16))
91 if command == 9:
92 struct.pack_into('<HIQ', frame, 66, offset, len(data), 4096)
93 else:
94 struct.pack_into('<BBIH', frame, 66, 1, 20, len(data), offset)
95 frame += data
96 sent.append(bytes(frame))
97 writer.write(len(frame).to_bytes(4, 'big') + frame)
98 await writer.drain()
99 prefix = await asyncio.wait_for(reader.readexactly(4), 5)
100 await asyncio.wait_for(reader.readexactly(int.from_bytes(prefix[1:], 'big')), 5)
101 while not any(row.get('direction') == 'response' and row['message'] == message for row in records):
102 records.append(json.loads(await asyncio.wait_for(process.stdout.readline(), 5)))
103 self.assertEqual(frames, sent)
104 requests = [row for row in records if row.get('direction') == 'request']
105 self.assertNotIn('data', requests[0])
106 self.assertEqual(requests[1]['data'], b'abcde'.hex())
107 self.assertEqual(requests[1]['offset'], 4096)
108 self.assertEqual(requests[2]['data'], (4096).to_bytes(8, 'little').hex())
109 self.assertEqual(requests[2]['file_id'], bytes(range(16)).hex())
110 self.assertEqual([row['written'] for row in records if row.get('direction') == 'response' and row['command'] == 9], [3, 3])
111 finally:
112 if writer is not None:
113 writer.close()
114 await writer.wait_closed()
115 if process.returncode is None:
116 process.terminate()
117 await asyncio.wait_for(process.wait(), 5)
118 server.close()
119 await server.wait_closed()
120
121 async def test_connection_cut_keeps_existing_and_new_peers_usable(self):
122 async def respond(reader, writer):
123 try:
124 while True:
125 prefix = await reader.readexactly(4)
126 frame = await reader.readexactly(int.from_bytes(prefix[1:], 'big'))
127 reply = bytearray(80)
128 reply[:64] = frame[:64]
129 struct.pack_into('<I', reply, 16, 1)
130 struct.pack_into('<I', reply, 68, 4)
131 writer.write(len(reply).to_bytes(4, 'big') + reply)
132 await writer.drain()
133 except (OSError, asyncio.IncompleteReadError):
134 pass
135 finally:
136 writer.close()
137 await writer.wait_closed()
138
139 async def send(writer, message):
140 frame = bytearray(116)
141 frame[:4] = b'\xfeSMB'
142 struct.pack_into('<H', frame, 12, 9)
143 struct.pack_into('<Q', frame, 24, message)
144 struct.pack_into('<HIQ', frame, 66, 112, 4, 96)
145 writer.write(len(frame).to_bytes(4, 'big') + frame)
146 await writer.drain()
147
148 async def receive(reader):
149 prefix = await asyncio.wait_for(reader.readexactly(4), 5)
150 return await asyncio.wait_for(reader.readexactly(int.from_bytes(prefix[1:], 'big')), 5)
151
152 for scope in ('connection', 'all'):
153 with self.subTest(scope=scope), tempfile.TemporaryDirectory() as directory:
154 server = await asyncio.start_server(respond, '127.0.0.1', 0)
155 control = Path(directory) / 'control.json'
156 control.write_text('{}')
157 process = await asyncio.create_subprocess_exec(
158 sys.executable, str(Path(proxy.__file__)), str(control), '--port', '0',
159 '--server', '127.0.0.1', '--server-port', str(server.sockets[0].getsockname()[1]),
160 stdout=asyncio.subprocess.PIPE, stderr=asyncio.subprocess.PIPE)
161 clients, records = [], []
162 try:
163 while not any('control' in row for row in records):
164 records.append(json.loads(await asyncio.wait_for(process.stdout.readline(), 5)))
165 port, = [row['listening'] for row in records if 'listening' in row]
166 for index in range(2):
167 reader, writer = await asyncio.open_connection('127.0.0.1', port)
168 clients.append((reader, writer))
169 await send(writer, index)
170 self.assertEqual(struct.unpack_from('<I', await receive(reader), 68)[0], 4)
171 setting = dict(cut=9, offset=96, scope=scope)
172 control.write_text(json.dumps(setting))
173 while not any(row.get('control') == setting for row in records):
174 records.append(json.loads(await asyncio.wait_for(process.stdout.readline(), 5)))
175 await send(clients[0][1], 2)
176 with self.assertRaises(asyncio.IncompleteReadError): await receive(clients[0][0])
177 while not any('cut' in row for row in records):
178 records.append(json.loads(await asyncio.wait_for(process.stdout.readline(), 5)))
179 cut, = [row['cut'] for row in records if 'cut' in row]
180 self.assertEqual((cut['direction'], cut['status'], cut['written']), ('response', '0x0', 4))
181 reader, writer = await asyncio.open_connection('127.0.0.1', port)
182 clients.append((reader, writer))
183 for reader, writer in clients[1:]:
184 if scope == 'connection':
185 await send(writer, 3)
186 self.assertEqual(struct.unpack_from('<I', await receive(reader), 68)[0], 4)
187 else:
188 with self.assertRaises(asyncio.IncompleteReadError): await receive(reader)
189 finally:
190 for _, writer in clients:
191 writer.close()
192 await writer.wait_closed()
193 if process.returncode is None: process.terminate()
194 await asyncio.wait_for(process.wait(), 5)
195 server.close()
196 await server.wait_closed()