Initial server source import
This commit is contained in:
@@ -0,0 +1,87 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user