88 lines
2.8 KiB
Python
88 lines
2.8 KiB
Python
import os
|
|
from pathlib import Path
|
|
import socket
|
|
import sys
|
|
import threading
|
|
import unittest
|
|
|
|
|
|
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_server
|
|
from host_agent_protocol import (
|
|
HOST_AGENT_RUNTIME_UID,
|
|
MAX_RESPONSE_PAYLOAD_BYTES,
|
|
HostAgentAction,
|
|
HostAgentRequest,
|
|
HostAgentResponse,
|
|
HostAgentStatus,
|
|
decode_response_frame,
|
|
encode_request_frame,
|
|
receive_frame,
|
|
require_eof,
|
|
)
|
|
|
|
|
|
@unittest.skipUnless(sys.platform == 'linux', 'Linux peer credentials are required')
|
|
class HostAgentLinuxTests(unittest.TestCase):
|
|
def test_kernel_peer_credentials_admit_exact_runtime_uid(self):
|
|
if os.geteuid() != HOST_AGENT_RUNTIME_UID:
|
|
self.skipTest('test must run as the fixed runtime UID')
|
|
request = HostAgentRequest(
|
|
operation_id='12345678-1234-4234-8234-123456789abc',
|
|
action=HostAgentAction.RESTART,
|
|
active_config_sha256='a' * 64,
|
|
active_secrets_sha256='b' * 64,
|
|
candidate_config_sha256=None,
|
|
candidate_secrets_sha256=None,
|
|
)
|
|
server, client = socket.socketpair(socket.AF_UNIX, socket.SOCK_STREAM)
|
|
result = []
|
|
|
|
def run():
|
|
with server:
|
|
result.append(host_agent_server.serve_connection(server))
|
|
|
|
thread = threading.Thread(target=run)
|
|
thread.start()
|
|
with client:
|
|
client.sendall(encode_request_frame(request))
|
|
client.shutdown(socket.SHUT_WR)
|
|
deadline = __import__('time').monotonic() + 2
|
|
response = decode_response_frame(receive_frame(
|
|
client, maximum=MAX_RESPONSE_PAYLOAD_BYTES, deadline=deadline,
|
|
))
|
|
require_eof(client, deadline=deadline)
|
|
thread.join(2)
|
|
self.assertFalse(thread.is_alive())
|
|
self.assertEqual(result, [True])
|
|
self.assertEqual(
|
|
response,
|
|
HostAgentResponse(request.operation_id, HostAgentStatus.UNAVAILABLE),
|
|
)
|
|
|
|
def test_partial_disconnect_and_trailing_frame_do_not_dispatch(self):
|
|
if os.geteuid() != HOST_AGENT_RUNTIME_UID:
|
|
self.skipTest('test must run as the fixed runtime UID')
|
|
dispatched = []
|
|
server, client = socket.socketpair(socket.AF_UNIX, socket.SOCK_STREAM)
|
|
thread = threading.Thread(
|
|
target=lambda: host_agent_server.serve_connection(
|
|
server, handler=lambda request: dispatched.append(request),
|
|
)
|
|
)
|
|
thread.start()
|
|
client.sendall(b'\x00\x00\x00\x20{}')
|
|
client.close()
|
|
thread.join(2)
|
|
server.close()
|
|
self.assertFalse(thread.is_alive())
|
|
self.assertEqual(dispatched, [])
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main()
|