530 lines
22 KiB
Python
530 lines
22 KiB
Python
import json
|
|
from pathlib import Path
|
|
import socket
|
|
import struct
|
|
import sys
|
|
import threading
|
|
import time
|
|
import unittest
|
|
import uuid
|
|
from unittest import mock
|
|
|
|
|
|
ROOT = Path(__file__).resolve().parents[1]
|
|
APP = ROOT / 'app'
|
|
if str(APP) not in sys.path:
|
|
sys.path.insert(0, str(APP))
|
|
|
|
import host_agent_client
|
|
import host_agent_protocol
|
|
import host_agent_server
|
|
from host_agent_client import (
|
|
HostAgentClient,
|
|
HostAgentRejectedError,
|
|
HostAgentUnavailableError,
|
|
)
|
|
from host_agent_protocol import (
|
|
MAX_REQUEST_PAYLOAD_BYTES,
|
|
HostAgentAction,
|
|
HostAgentProtocolError,
|
|
HostAgentRequest,
|
|
HostAgentResponse,
|
|
HostAgentStatus,
|
|
decode_request_frame,
|
|
decode_request_payload,
|
|
decode_response_frame,
|
|
encode_request_frame,
|
|
encode_request_payload,
|
|
encode_response_frame,
|
|
)
|
|
|
|
|
|
def request_for(action='apply-config'):
|
|
candidates = {
|
|
'apply-config': ('c' * 64, None),
|
|
'apply-secrets': (None, 'd' * 64),
|
|
'apply-both': ('c' * 64, 'd' * 64),
|
|
'restart': (None, None),
|
|
}[action]
|
|
return HostAgentRequest(
|
|
operation_id=str(uuid.uuid4()),
|
|
action=HostAgentAction(action),
|
|
active_config_sha256='a' * 64,
|
|
active_secrets_sha256='b' * 64,
|
|
candidate_config_sha256=candidates[0],
|
|
candidate_secrets_sha256=candidates[1],
|
|
)
|
|
|
|
|
|
class HostAgentProtocolTests(unittest.TestCase):
|
|
def test_all_actions_round_trip_as_canonical_exact_schema(self):
|
|
for action in ('apply-config', 'apply-secrets', 'apply-both', 'restart'):
|
|
with self.subTest(action=action):
|
|
request = request_for(action)
|
|
payload = encode_request_payload(request)
|
|
self.assertEqual(decode_request_payload(payload), request)
|
|
self.assertEqual(decode_request_frame(encode_request_frame(request)), request)
|
|
parsed = json.loads(payload)
|
|
self.assertEqual(set(parsed), {
|
|
'operation_id', 'action', 'active_config_sha256',
|
|
'active_secrets_sha256', 'candidate_config_sha256',
|
|
'candidate_secrets_sha256',
|
|
})
|
|
self.assertEqual(
|
|
payload,
|
|
json.dumps(
|
|
parsed, ensure_ascii=True, sort_keys=True,
|
|
separators=(',', ':'), allow_nan=False,
|
|
).encode('ascii'),
|
|
)
|
|
|
|
def test_unknown_duplicate_missing_and_noncanonical_fields_are_rejected(self):
|
|
request = request_for()
|
|
canonical = encode_request_payload(request)
|
|
parsed = json.loads(canonical)
|
|
malformed = []
|
|
for field in tuple(parsed):
|
|
changed = dict(parsed)
|
|
changed.pop(field)
|
|
malformed.append(json.dumps(changed, sort_keys=True, separators=(',', ':')).encode())
|
|
for field in ('path', 'service', 'command', 'argv', 'environment', 'docker', 'options'):
|
|
changed = dict(parsed, **{field: 'forbidden'})
|
|
malformed.append(json.dumps(changed, sort_keys=True, separators=(',', ':')).encode())
|
|
malformed.extend((
|
|
canonical.replace(b'"action":', b' "action":', 1),
|
|
canonical + b'\n',
|
|
b'{' + canonical[1:canonical.find(b',') + 1] + canonical[1:],
|
|
b'[]',
|
|
b'{"action":NaN}',
|
|
b'\xff',
|
|
))
|
|
for payload in malformed:
|
|
with self.subTest(payload=payload[:60]):
|
|
with self.assertRaises(HostAgentProtocolError):
|
|
decode_request_payload(payload)
|
|
|
|
def test_identity_types_hashes_uuid_and_action_matrix_are_exact(self):
|
|
parsed = json.loads(encode_request_payload(request_for('apply-both')))
|
|
cases = (
|
|
('operation_id', str(uuid.UUID(int=0))),
|
|
('operation_id', parsed['operation_id'].upper()),
|
|
('operation_id', 1),
|
|
('action', 'shell'),
|
|
('active_config_sha256', 'A' * 64),
|
|
('active_secrets_sha256', 'b' * 63),
|
|
('candidate_config_sha256', None),
|
|
('candidate_secrets_sha256', True),
|
|
)
|
|
for field, value in cases:
|
|
with self.subTest(field=field, value=value):
|
|
changed = dict(parsed)
|
|
changed[field] = value
|
|
payload = json.dumps(changed, sort_keys=True, separators=(',', ':')).encode()
|
|
with self.assertRaises(HostAgentProtocolError):
|
|
decode_request_payload(payload)
|
|
|
|
def test_frames_reject_zero_oversize_truncation_and_trailing_data(self):
|
|
frame = encode_request_frame(request_for())
|
|
for changed in (
|
|
b'', b'\x00\x00\x00', struct.pack('!I', 0),
|
|
struct.pack('!I', MAX_REQUEST_PAYLOAD_BYTES + 1),
|
|
frame[:-1], frame + b'x',
|
|
):
|
|
with self.subTest(frame=changed[:20]):
|
|
with self.assertRaises(HostAgentProtocolError):
|
|
decode_request_frame(changed)
|
|
|
|
def test_response_schema_is_closed_and_identity_bound(self):
|
|
operation_id = str(uuid.uuid4())
|
|
for status in (
|
|
HostAgentStatus.ACCEPTED, HostAgentStatus.UNAVAILABLE,
|
|
HostAgentStatus.REJECTED,
|
|
):
|
|
response = HostAgentResponse(operation_id, status)
|
|
self.assertEqual(decode_response_frame(encode_response_frame(response)), response)
|
|
invalid = HostAgentResponse(None, HostAgentStatus.INVALID)
|
|
self.assertEqual(decode_response_frame(encode_response_frame(invalid)), invalid)
|
|
for response in (
|
|
HostAgentResponse(None, HostAgentStatus.ACCEPTED),
|
|
HostAgentResponse(operation_id, HostAgentStatus.INVALID),
|
|
):
|
|
with self.assertRaises(HostAgentProtocolError):
|
|
encode_response_frame(response)
|
|
|
|
def test_absolute_deadline_is_not_extended_by_trickled_bytes(self):
|
|
class TrickleSocket:
|
|
def __init__(self):
|
|
self.timeouts = []
|
|
|
|
def settimeout(self, value):
|
|
self.timeouts.append(value)
|
|
|
|
def recv(self, _length):
|
|
return b'\x00'
|
|
|
|
connection = TrickleSocket()
|
|
with mock.patch.object(
|
|
host_agent_protocol.time, 'monotonic',
|
|
side_effect=(1.0, 4.0, 7.0, 10.1),
|
|
):
|
|
with self.assertRaisesRegex(HostAgentProtocolError, 'invalid') as raised:
|
|
host_agent_protocol.receive_frame(
|
|
connection, maximum=MAX_REQUEST_PAYLOAD_BYTES, deadline=10.0,
|
|
)
|
|
self.assertEqual(raised.exception.category, 'timeout')
|
|
self.assertEqual(len(connection.timeouts), 3)
|
|
|
|
|
|
class _ClientSocket:
|
|
def __init__(self, response, *, peer_uid=0):
|
|
self.response = bytearray(response)
|
|
self.peer_uid = peer_uid
|
|
self.sent = bytearray()
|
|
self.closed = False
|
|
self.shutdown_how = None
|
|
self.connected = None
|
|
|
|
def settimeout(self, _value):
|
|
pass
|
|
|
|
def connect(self, path):
|
|
self.connected = path
|
|
|
|
def getsockopt(self, _level, _option, _length):
|
|
return struct.pack('3i', 123, self.peer_uid, 0)
|
|
|
|
def send(self, value):
|
|
self.sent.extend(value)
|
|
return len(value)
|
|
|
|
def recv(self, length):
|
|
if not self.response:
|
|
return b''
|
|
result = bytes(self.response[:length])
|
|
del self.response[:length]
|
|
return result
|
|
|
|
def shutdown(self, how):
|
|
self.shutdown_how = how
|
|
|
|
def close(self):
|
|
self.closed = True
|
|
|
|
|
|
class HostAgentClientTests(unittest.TestCase):
|
|
def test_client_sends_exact_request_and_accepts_only_root_response(self):
|
|
request = request_for('apply-both')
|
|
response = encode_response_frame(HostAgentResponse(
|
|
request.operation_id, HostAgentStatus.ACCEPTED,
|
|
))
|
|
connection = _ClientSocket(response)
|
|
with mock.patch.object(host_agent_client, '_fixed_socket_is_safe', return_value=True), \
|
|
mock.patch.object(
|
|
host_agent_client, '_peer_credentials',
|
|
side_effect=lambda value: (0, value.peer_uid, 0),
|
|
), \
|
|
mock.patch.object(host_agent_client.socket, 'AF_UNIX', 1, create=True), \
|
|
mock.patch.object(host_agent_client.socket, 'socket', return_value=connection):
|
|
result = HostAgentClient().dispatch(
|
|
operation_id=request.operation_id, action=request.action.value,
|
|
active_config_sha256=request.active_config_sha256,
|
|
active_secrets_sha256=request.active_secrets_sha256,
|
|
candidate_config_sha256=request.candidate_config_sha256,
|
|
candidate_secrets_sha256=request.candidate_secrets_sha256,
|
|
)
|
|
self.assertEqual(result.status, HostAgentStatus.ACCEPTED)
|
|
self.assertEqual(decode_request_frame(bytes(connection.sent)), request)
|
|
self.assertEqual(connection.shutdown_how, socket.SHUT_WR)
|
|
self.assertTrue(connection.closed)
|
|
|
|
def test_client_rejects_nonroot_mismatch_rejection_and_unavailable_socket(self):
|
|
request = request_for()
|
|
cases = (
|
|
(_ClientSocket(encode_response_frame(HostAgentResponse(
|
|
request.operation_id, HostAgentStatus.ACCEPTED,
|
|
)), peer_uid=1000), HostAgentUnavailableError),
|
|
(_ClientSocket(encode_response_frame(HostAgentResponse(
|
|
str(uuid.uuid4()), HostAgentStatus.ACCEPTED,
|
|
))), HostAgentUnavailableError),
|
|
(_ClientSocket(encode_response_frame(HostAgentResponse(
|
|
request.operation_id, HostAgentStatus.REJECTED,
|
|
))), HostAgentRejectedError),
|
|
(_ClientSocket(encode_response_frame(HostAgentResponse(
|
|
request.operation_id, HostAgentStatus.UNAVAILABLE,
|
|
))), HostAgentUnavailableError),
|
|
)
|
|
for connection, error in cases:
|
|
with self.subTest(error=error.__name__), \
|
|
mock.patch.object(host_agent_client, '_fixed_socket_is_safe', return_value=True), \
|
|
mock.patch.object(
|
|
host_agent_client, '_peer_credentials',
|
|
side_effect=lambda value: (123, value.peer_uid, 0),
|
|
), \
|
|
mock.patch.object(host_agent_client.socket, 'AF_UNIX', 1, create=True), \
|
|
mock.patch.object(host_agent_client.socket, 'socket', return_value=connection):
|
|
with self.assertRaises(error):
|
|
HostAgentClient().dispatch(
|
|
operation_id=request.operation_id, action=request.action.value,
|
|
active_config_sha256=request.active_config_sha256,
|
|
active_secrets_sha256=request.active_secrets_sha256,
|
|
candidate_config_sha256=request.candidate_config_sha256,
|
|
candidate_secrets_sha256=request.candidate_secrets_sha256,
|
|
)
|
|
self.assertTrue(connection.closed)
|
|
with mock.patch.object(host_agent_client, '_fixed_socket_is_safe', return_value=False), \
|
|
mock.patch.object(host_agent_client.socket, 'socket') as socket_factory:
|
|
with self.assertRaises(HostAgentUnavailableError):
|
|
HostAgentClient().dispatch(
|
|
operation_id=request.operation_id, action=request.action.value,
|
|
active_config_sha256=request.active_config_sha256,
|
|
active_secrets_sha256=request.active_secrets_sha256,
|
|
candidate_config_sha256=request.candidate_config_sha256,
|
|
candidate_secrets_sha256=request.candidate_secrets_sha256,
|
|
)
|
|
socket_factory.assert_not_called()
|
|
|
|
def test_client_rejects_response_pipelining_and_closes_on_baseexception(self):
|
|
request = request_for()
|
|
response = encode_response_frame(HostAgentResponse(
|
|
request.operation_id, HostAgentStatus.ACCEPTED,
|
|
))
|
|
pipelined = _ClientSocket(response + b'x')
|
|
with mock.patch.object(host_agent_client, '_fixed_socket_is_safe', return_value=True), \
|
|
mock.patch.object(host_agent_client, '_peer_credentials', return_value=(123, 0, 0)), \
|
|
mock.patch.object(host_agent_client.socket, 'AF_UNIX', 1, create=True), \
|
|
mock.patch.object(host_agent_client.socket, 'socket', return_value=pipelined):
|
|
with self.assertRaises(HostAgentUnavailableError):
|
|
HostAgentClient().dispatch(
|
|
operation_id=request.operation_id, action=request.action.value,
|
|
active_config_sha256=request.active_config_sha256,
|
|
active_secrets_sha256=request.active_secrets_sha256,
|
|
candidate_config_sha256=request.candidate_config_sha256,
|
|
candidate_secrets_sha256=request.candidate_secrets_sha256,
|
|
)
|
|
self.assertTrue(pipelined.closed)
|
|
|
|
interrupted = _ClientSocket(response)
|
|
interrupted.send = mock.Mock(side_effect=KeyboardInterrupt())
|
|
with mock.patch.object(host_agent_client, '_fixed_socket_is_safe', return_value=True), \
|
|
mock.patch.object(host_agent_client, '_peer_credentials', return_value=(123, 0, 0)), \
|
|
mock.patch.object(host_agent_client.socket, 'AF_UNIX', 1, create=True), \
|
|
mock.patch.object(host_agent_client.socket, 'socket', return_value=interrupted):
|
|
with self.assertRaises(KeyboardInterrupt):
|
|
HostAgentClient().dispatch(
|
|
operation_id=request.operation_id, action=request.action.value,
|
|
active_config_sha256=request.active_config_sha256,
|
|
active_secrets_sha256=request.active_secrets_sha256,
|
|
candidate_config_sha256=request.candidate_config_sha256,
|
|
candidate_secrets_sha256=request.candidate_secrets_sha256,
|
|
)
|
|
self.assertTrue(interrupted.closed)
|
|
|
|
|
|
class HostAgentServerTests(unittest.TestCase):
|
|
def _exchange(self, request_bytes, *, credentials=(123, 10001, 10001), handler=None):
|
|
server, client = socket.socketpair()
|
|
result = []
|
|
|
|
def run():
|
|
with server, mock.patch.object(
|
|
host_agent_server, '_peer_credentials', return_value=credentials,
|
|
):
|
|
kwargs = {} if handler is None else {'handler': handler}
|
|
result.append(host_agent_server.serve_connection(server, **kwargs))
|
|
|
|
thread = threading.Thread(target=run)
|
|
thread.start()
|
|
try:
|
|
client.sendall(request_bytes)
|
|
client.shutdown(socket.SHUT_WR)
|
|
except (BrokenPipeError, ConnectionAbortedError, ConnectionResetError):
|
|
pass
|
|
response = bytearray()
|
|
while True:
|
|
try:
|
|
chunk = client.recv(4096)
|
|
except (ConnectionAbortedError, ConnectionResetError):
|
|
chunk = b''
|
|
if not chunk:
|
|
break
|
|
response.extend(chunk)
|
|
client.close()
|
|
thread.join(2)
|
|
self.assertFalse(thread.is_alive())
|
|
return result, bytes(response)
|
|
|
|
def test_authorized_valid_request_is_validated_but_unavailable_until_task_9_2(self):
|
|
request = request_for()
|
|
result, response = self._exchange(encode_request_frame(request))
|
|
self.assertEqual(result, [True])
|
|
self.assertEqual(
|
|
decode_response_frame(response),
|
|
HostAgentResponse(request.operation_id, HostAgentStatus.UNAVAILABLE),
|
|
)
|
|
|
|
def test_unauthorized_peer_is_rejected_before_body_read_and_gets_no_response(self):
|
|
request = request_for()
|
|
result, response = self._exchange(
|
|
encode_request_frame(request), credentials=(123, 1000, 10001),
|
|
)
|
|
self.assertEqual(result, [False])
|
|
self.assertEqual(response, b'')
|
|
|
|
server, client = socket.socketpair()
|
|
try:
|
|
with mock.patch.object(
|
|
host_agent_server, '_peer_credentials', return_value=(123, 1000, 10001),
|
|
), mock.patch.object(
|
|
host_agent_server, 'receive_frame',
|
|
side_effect=AssertionError('body must not be read'),
|
|
):
|
|
self.assertFalse(host_agent_server.serve_connection(server))
|
|
finally:
|
|
server.close()
|
|
client.close()
|
|
|
|
def test_invalid_and_trailing_requests_receive_only_bounded_invalid_response(self):
|
|
for frame in (
|
|
struct.pack('!I', 2) + b'{}',
|
|
encode_request_frame(request_for()) + b'x',
|
|
):
|
|
with self.subTest(frame=frame[:30]):
|
|
result, response = self._exchange(frame)
|
|
self.assertEqual(result, [True])
|
|
self.assertEqual(
|
|
decode_response_frame(response),
|
|
HostAgentResponse(None, HostAgentStatus.INVALID),
|
|
)
|
|
|
|
def test_deep_unknown_request_field_is_rejected_before_handler(self):
|
|
parsed = json.loads(encode_request_payload(request_for()))
|
|
nested = 'forbidden'
|
|
for _ in range(128):
|
|
nested = [nested]
|
|
parsed['options'] = nested
|
|
payload = json.dumps(
|
|
parsed, ensure_ascii=True, sort_keys=True,
|
|
separators=(',', ':'), allow_nan=False,
|
|
).encode('ascii')
|
|
self.assertLess(len(payload), MAX_REQUEST_PAYLOAD_BYTES)
|
|
handler = mock.Mock()
|
|
|
|
result, response = self._exchange(
|
|
struct.pack('!I', len(payload)) + payload,
|
|
handler=handler,
|
|
)
|
|
|
|
self.assertEqual(result, [True])
|
|
self.assertEqual(
|
|
decode_response_frame(response),
|
|
HostAgentResponse(None, HostAgentStatus.INVALID),
|
|
)
|
|
handler.assert_not_called()
|
|
|
|
def test_handler_cannot_change_operation_identity(self):
|
|
request = request_for()
|
|
result, response = self._exchange(
|
|
encode_request_frame(request),
|
|
handler=lambda _request: HostAgentResponse(
|
|
str(uuid.uuid4()), HostAgentStatus.ACCEPTED,
|
|
),
|
|
)
|
|
self.assertEqual(result, [True])
|
|
self.assertEqual(
|
|
decode_response_frame(response),
|
|
HostAgentResponse(request.operation_id, HostAgentStatus.UNAVAILABLE),
|
|
)
|
|
|
|
def test_listener_requires_exact_stream_type_and_accepting_state(self):
|
|
class FakeListener:
|
|
def __init__(self, *, socket_type=socket.SOCK_STREAM, accepting=1, family=None):
|
|
self.family = getattr(socket, 'AF_UNIX', 1) if family is None else family
|
|
self.socket_type = socket_type
|
|
self.accepting = accepting
|
|
|
|
def getsockopt(self, _level, option):
|
|
if option == socket.SO_TYPE:
|
|
return self.socket_type
|
|
if option == socket.SO_ACCEPTCONN:
|
|
return self.accepting
|
|
raise OSError('unexpected option')
|
|
|
|
def getsockname(self):
|
|
return host_agent_server.HOST_AGENT_SOCKET_PATH
|
|
|
|
details = mock.Mock(st_mode=__import__('stat').S_IFSOCK, st_uid=0)
|
|
cases = (
|
|
FakeListener(socket_type=getattr(socket, 'SOCK_SEQPACKET', 5)),
|
|
FakeListener(accepting=0),
|
|
FakeListener(family=socket.AF_INET),
|
|
)
|
|
for listener in cases:
|
|
with self.subTest(listener=listener.__dict__), \
|
|
mock.patch.object(host_agent_server.socket, 'socket', FakeListener), \
|
|
mock.patch.object(host_agent_server.socket, 'AF_UNIX', 1, create=True), \
|
|
mock.patch.object(host_agent_server.os, 'lstat', return_value=details):
|
|
with self.assertRaises(host_agent_server.HostAgentServerError):
|
|
host_agent_server._validate_listener(listener)
|
|
|
|
listener = FakeListener()
|
|
with mock.patch.object(host_agent_server.socket, 'socket', FakeListener), \
|
|
mock.patch.object(host_agent_server.socket, 'AF_UNIX', 1, create=True), \
|
|
mock.patch.object(host_agent_server.os, 'lstat', return_value=details):
|
|
host_agent_server._validate_listener(listener)
|
|
|
|
def test_invalid_inherited_listener_is_closed(self):
|
|
created = []
|
|
|
|
class ActivatedListener:
|
|
def __init__(self, *, fileno):
|
|
self.fileno = fileno
|
|
self.family = socket.AF_INET
|
|
self.closed = False
|
|
created.append(self)
|
|
|
|
def getsockopt(self, _level, option):
|
|
if option == socket.SO_TYPE:
|
|
return socket.SOCK_STREAM
|
|
if option == socket.SO_ACCEPTCONN:
|
|
return 1
|
|
raise OSError('unexpected option')
|
|
|
|
def close(self):
|
|
self.closed = True
|
|
|
|
environment = {
|
|
'LISTEN_PID': str(__import__('os').getpid()),
|
|
'LISTEN_FDS': '1',
|
|
}
|
|
with mock.patch.object(host_agent_server.os, 'geteuid', return_value=0, create=True), \
|
|
mock.patch.dict(host_agent_server.os.environ, environment, clear=True), \
|
|
mock.patch.object(host_agent_server.socket, 'socket', ActivatedListener), \
|
|
mock.patch.object(host_agent_server.socket, 'AF_UNIX', 1, create=True):
|
|
with self.assertRaises(host_agent_server.HostAgentServerError):
|
|
host_agent_server.inherited_systemd_listener()
|
|
self.assertEqual(len(created), 1)
|
|
self.assertEqual(created[0].fileno, host_agent_server.SYSTEMD_LISTEN_FD)
|
|
self.assertTrue(created[0].closed)
|
|
|
|
def test_server_stop_event_bounds_accept_loop(self):
|
|
stop_event = threading.Event()
|
|
|
|
class PollListener:
|
|
timeout = None
|
|
|
|
def settimeout(self, value):
|
|
self.timeout = value
|
|
|
|
def accept(self):
|
|
stop_event.set()
|
|
raise socket.timeout()
|
|
|
|
listener = PollListener()
|
|
with mock.patch.object(host_agent_server, '_validate_listener'):
|
|
host_agent_server.serve_forever(listener, stop_event=stop_event)
|
|
self.assertEqual(listener.timeout, host_agent_server.ACCEPT_POLL_SECONDS)
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main()
|