Files
truf-server/tests/test_host_agent_linux.py
T
2026-09-30 20:30:56 +03:00

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()