1#!/usr/bin/env python3
2"""Trace a dedicated test SMB session and cut selected requests or responses."""
3import argparse
4import asyncio
5import json
6from pathlib import Path
7import struct
8import time
9
10
11def lease_state(frame, at, contexts):
12 """The lease state of a CREATE's `RqLs` context, if any (MS-SMB2 2.2.13.2.8, 2.2.14.2.10)."""
13 offset, length = struct.unpack_from('<II', frame, at + contexts)
14 cursor = at + offset if length else None
15 while cursor is not None:
16 following, name_offset, name_length, _, data_offset = struct.unpack_from('<IHHHH', frame, cursor)
17 if frame[cursor + name_offset:cursor + name_offset + name_length] == b'RqLs':
18 return struct.unpack_from('<I', frame, cursor + data_offset + 16)[0]
19 cursor = cursor + following if following else None
20 return None
21
22
23def header_fields(offset, data):
24 result = {}
25 for name, start, length in [('transactions', 96, 4), ('version', 212, 16),
26 ('generation', 228, 8), ('deny_read', 236, 16)]:
27 if offset <= start and start + length <= offset + len(data):
28 value = data[start - offset:start - offset + length]
29 result[name] = value.hex() if length == 16 else int.from_bytes(value, 'little')
30 return result
31
32
33async def main():
34 parser = argparse.ArgumentParser(description=__doc__)
35 parser.add_argument('control', type=Path)
36 parser.add_argument('--port', type=int, default=11445)
37 parser.add_argument('--server', default='10.0.0.1')
38 parser.add_argument('--server-port', type=int, default=445)
39 parser.add_argument('--bind', default='127.0.0.1')
40 args = parser.parse_args()
41 writers = set()
42 state = {'mode': 'up'}
43 previous = None
44 blocked = False
45 matched = 0
46 connections = 0
47 record_writes = False
48
49 def record(**fields):
50 print(json.dumps({'time': time.time(), **fields}), flush=True)
51
52 async def controls():
53 nonlocal state, previous, blocked, matched, record_writes
54 while True:
55 try:
56 raw = args.control.read_bytes()
57 if raw != previous:
58 state = json.loads(raw)
59 if 'record_writes' in state:
60 record_writes = bool(state['record_writes'])
61 previous, matched = raw, 0
62 blocked = state.get('mode') == 'down'
63 record(control=state)
64 if blocked:
65 for writer in list(writers):
66 writer.transport.abort()
67 except (FileNotFoundError, json.JSONDecodeError):
68 pass
69 await asyncio.sleep(0.05)
70
71 async def connection(client, client_writer):
72 nonlocal blocked, matched, connections
73 connections += 1
74 connection_id = connections
75 record(connection=connection_id, peer=client_writer.get_extra_info("peername"), opened=True)
76 if blocked:
77 client_writer.close()
78 return
79 server_writer = None
80 try:
81 server, server_writer = await asyncio.open_connection(args.server, args.server_port)
82 writers.update((client_writer, server_writer))
83
84 async def forward(reader, writer, direction):
85 nonlocal blocked, matched
86 while True:
87 prefix = await reader.readexactly(4)
88 frame = await reader.readexactly(int.from_bytes(prefix[1:], 'big'))
89 at = 0
90 while frame[at:at+4] == b'\xfeSMB':
91 command = struct.unpack_from('<H', frame, at+12)[0]
92 entry = {'connection': connection_id, 'direction': direction, 'command': command, 'frame': frames[direction],
93 'message': struct.unpack_from('<Q', frame, at+24)[0],
94 'credit_charge': struct.unpack_from('<H', frame, at+6)[0],
95 'credits': struct.unpack_from('<H', frame, at+14)[0]}
96 if direction == 'response':
97 entry['status'] = hex(struct.unpack_from('<I', frame, at+8)[0])
98 elif command in (8, 9):
99 entry['length'], entry['offset'] = struct.unpack_from('<IQ', frame, at+68)
100 if command == 9:
101 data_offset = struct.unpack_from('<H', frame, at+66)[0]
102 data = frame[at+data_offset:at+data_offset+entry['length']]
103 entry.update(header_fields(entry['offset'], data))
104 if record_writes:
105 entry['data'] = data.hex()
106 elif command == 10:
107 count = struct.unpack_from('<H', frame, at+66)[0]
108 entry['locks'] = [struct.unpack_from('<QQI', frame, at+88+24*i) for i in range(count)]
109 if command == 18 and (direction == 'request' or entry['status'] == '0x0'):
110 size = struct.unpack_from('<H', frame, at+64)[0]
111 entry['oplock_body'] = frame[at+64:at+64+size].hex()
112 if direction == 'request':
113 if command in (6, 7, 8, 9, 10):
114 offset = 80 if command in (8, 9) else 72
115 entry['file_id'] = frame[at+offset:at+offset+16].hex()
116 elif command == 17:
117 entry['info_type'], entry['info_class'] = struct.unpack_from('<BB', frame, at+66)
118 if record_writes:
119 entry['file_id'] = frame[at+80:at+96].hex()
120 length, offset = struct.unpack_from('<IH', frame, at+68)
121 entry['data'] = frame[at+offset:at+offset+length].hex()
122 elif command == 5:
123 entry['oplock'] = frame[at+67]
124 entry['access'] = struct.unpack_from('<I', frame, at+88)[0]
125 entry['share'], entry['disposition'], entry['options'] = struct.unpack_from('<III', frame, at+96)
126 offset, length = struct.unpack_from('<HH', frame, at+108)
127 entry['path'] = frame[at+offset:at+offset+length].decode('utf-16-le')
128 if entry['oplock'] == 0xff:
129 entry['lease'] = lease_state(frame, at, 112)
130 elif command == 5 and entry['status'] == '0x0':
131 entry['oplock'] = frame[at+66]
132 entry['file_id'] = frame[at+128:at+144].hex()
133 if entry['oplock'] == 0xff:
134 entry['lease'] = lease_state(frame, at, 144)
135 elif command == 9 and entry['status'] == '0x0':
136 entry['written'] = struct.unpack_from('<I', frame, at+68)[0]
137 if direction == 'request':
138 requests[entry['message']] = entry
139 request = entry
140 else:
141 request = requests.get(entry['message'], {})
142 if entry['status'] != '0x103': requests.pop(entry['message'], None)
143 if command == 8 and entry['status'] == '0x0' and request.get('offset', 252) < 252:
144 data_offset = frame[at+66]
145 length = struct.unpack_from('<I', frame, at+68)[0]
146 entry.update(header_fields(request['offset'], frame[at+data_offset:at+data_offset+length]))
147 record(**entry)
148 if (state.get('cut') == command and direction == state.get('direction', 'response')
149 and (direction == 'request' or entry['status'] == state.get('status', '0x0'))
150 and ('peer' not in state or state['peer'] == client_writer.get_extra_info('peername')[0])
151 and ('offset' not in state or state['offset'] == request.get('offset'))):
152 matched += 1
153 if matched == state.get('occurrence', 1):
154 local = state.get('scope') == 'connection'
155 blocked = not local
156 record(cut=entry)
157 for stream in (client_writer, server_writer) if local else list(writers):
158 stream.transport.abort()
159 return
160 next_command = struct.unpack_from('<I', frame, at+20)[0]
161 if not next_command:
162 break
163 at += next_command
164 frames[direction] += 1
165 if frame[:4] == b'\xfdSMB':
166 record(encrypted=True, direction=direction)
167 if state.get('delay_ms'):
168 # Half the round trip each way, as on a wide-area link: every frame
169 # arrives that much later, in order, without queueing behind others.
170 loop = asyncio.get_running_loop()
171 loop.call_at(loop.time() + state['delay_ms'] / 2000, writer.write, prefix + frame)
172 else:
173 writer.write(prefix + frame)
174 await writer.drain()
175
176 requests = {}
177 frames = {'request': 0, 'response': 0}
178 tasks = [asyncio.create_task(forward(client, server_writer, 'request')),
179 asyncio.create_task(forward(server, client_writer, 'response'))]
180 done, pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
181 for task in pending:
182 task.cancel()
183 results = await asyncio.gather(*tasks, return_exceptions=True)
184 for result in results:
185 if isinstance(result, Exception) and not isinstance(result, (OSError, asyncio.IncompleteReadError)):
186 record(connection=connection_id, trace_error=repr(result))
187 except (OSError, asyncio.IncompleteReadError) as error:
188 record(error=str(error))
189 finally:
190 record(connection=connection_id, closed=True)
191 for writer in (client_writer, server_writer):
192 if writer:
193 writers.discard(writer)
194 writer.close()
195
196 server = await asyncio.start_server(connection, args.bind, args.port)
197 async with server:
198 record(listening=server.sockets[0].getsockname()[1])
199 await asyncio.gather(server.serve_forever(), controls())
200
201
202if __name__ == '__main__':
203 try:
204 asyncio.run(main())
205 except KeyboardInterrupt:
206 pass