mirror of
https://github.com/ChuckBuilds/LEDMatrix.git
synced 2026-10-04 14:25:08 +00:00
feat(ipc): display control socket, stage 1: on-demand with acks
The display process now serves a Unix socket, /run/ledmatrix/control.sock, carrying versioned newline-delimited JSON commands that are acknowledged. Stage 1 moves on-demand start/stop (plus status, hello and ping) onto it; the cache-file mailbox stays as the fallback for one release. - src/ipc/contract.py: typed request/response envelopes, command args, error codes, NDJSON framing with a 64 KiB limit, socket path rules. - src/ipc/server.py: threaded server owned by the display. Handlers only queue onto a bounded queue and ack with the request id; the render thread drains it where it reads the mailbox. Bounded clients, timeouts, garbage/oversize/disconnect handling; 0660 socket in the cache dir's group plus SO_PEERCRED checks; skips cleanly on Windows or when off. - src/ipc/client.py: one short-timeout request; any failure raises ControlError(reason). - api_v3/display.py: on-demand start/stop try the socket, fall back to the mailbox exactly as before, and report transport/socket_error. - display_controller.py: start/close the server; the mailbox handler body is extracted into _handle_on_demand_request and shared by both paths. - docs/IPC_CONTROL_SOCKET.md: protocol, security model, stage plan. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
This commit is contained in:
@@ -315,6 +315,19 @@ def _hermetic_display_watchdog(monkeypatch):
|
||||
display_watchdog.RenderWatchdog(environ={}, heartbeat_dir=None))
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _hermetic_control_socket(monkeypatch):
|
||||
"""Keep the control socket (src/ipc) off the host.
|
||||
|
||||
DisplayController.run() would serve /run/ledmatrix/control.sock -- or
|
||||
find the live display's already there, when the suite runs on a device
|
||||
-- and the web routes would send on-demand commands to that display.
|
||||
Off by default; the socket tests point it at a tmp_path of their own.
|
||||
"""
|
||||
from src.ipc.contract import SOCKET_PATH_ENV
|
||||
monkeypatch.setenv(SOCKET_PATH_ENV, 'off')
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_logging():
|
||||
"""Reset logging configuration before each test."""
|
||||
|
||||
@@ -0,0 +1,192 @@
|
||||
"""POST /display/on-demand/start and /stop: control socket first, mailbox fallback.
|
||||
|
||||
The routes hand the request to the display over the control socket
|
||||
(src/ipc) and get an acknowledgement. On any failure -- no socket (a stopped
|
||||
display, or one older than the socket), a timeout, a refusal, a bug in the
|
||||
client -- they write the file mailbox exactly as they did before the socket
|
||||
existed. These tests pin both paths, that exactly one of them is used, that
|
||||
the response says which, and that the request id is the same either way (the
|
||||
display deduplicates on it).
|
||||
|
||||
The socket client is patched at the route's module attribute; the last class
|
||||
runs a real server on a temp socket (Linux/macOS only).
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||
|
||||
from test._api_v3_test_helpers import api_v3_client, api_v3_module # noqa: F401,E402
|
||||
|
||||
from src.ipc import client as control_client # noqa: E402
|
||||
from src.ipc import contract as c # noqa: E402
|
||||
|
||||
START_URL = "/api/v3/display/on-demand/start"
|
||||
STOP_URL = "/api/v3/display/on-demand/stop"
|
||||
MAILBOX = "display_on_demand_request"
|
||||
CLIENT = "web_interface.blueprints.api_v3.display.control_client"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def service(api_v3_module):
|
||||
"""A running display service; records systemctl calls and mailbox writes."""
|
||||
api_v3_module.api_v3.plugin_catalog = None
|
||||
api_v3_module.api_v3.config_manager = None
|
||||
state = {"active": True}
|
||||
calls = []
|
||||
|
||||
def status():
|
||||
return {"active": state["active"]}
|
||||
|
||||
def systemctl(args):
|
||||
calls.append(("systemctl", args[-2]))
|
||||
if args[-2:] == ["start", "ledmatrix.service"]:
|
||||
state["active"] = True
|
||||
return {"returncode": 0, "stdout": "", "stderr": ""}
|
||||
|
||||
cache = api_v3_module.api_v3.cache_manager
|
||||
cache.set.side_effect = lambda key, value, *a, **kw: calls.append(("cache", key))
|
||||
with patch("web_interface.blueprints.api_v3._get_display_service_status",
|
||||
side_effect=status), \
|
||||
patch("web_interface.blueprints.api_v3.display._get_display_service_status",
|
||||
side_effect=status), \
|
||||
patch("web_interface.blueprints.api_v3._run_systemctl_command",
|
||||
side_effect=systemctl), \
|
||||
patch("web_interface.blueprints.api_v3.display._stop_display_service"):
|
||||
yield {"state": state, "cache": cache, "calls": calls}
|
||||
|
||||
|
||||
def _mailbox_writes(cache):
|
||||
return [call.args[1] for call in cache.set.call_args_list
|
||||
if call.args and call.args[0] == MAILBOX]
|
||||
|
||||
|
||||
def _ack(request_id, *a, **kw):
|
||||
return {"accepted": True, "request_id": request_id, "queued": 1}
|
||||
|
||||
|
||||
class TestSocketPath:
|
||||
def test_start_goes_over_the_socket_and_skips_the_mailbox(self, api_v3_client, service):
|
||||
with patch(f"{CLIENT}.on_demand_start", side_effect=_ack) as start:
|
||||
resp = api_v3_client.post(START_URL, json={
|
||||
"plugin_id": "weather", "mode": "weather_current",
|
||||
"duration": 60, "pinned": True})
|
||||
assert resp.status_code == 200, resp.get_json()
|
||||
data = resp.get_json()["data"]
|
||||
assert data["transport"] == "socket"
|
||||
assert "socket_error" not in data
|
||||
assert _mailbox_writes(service["cache"]) == []
|
||||
start.assert_called_once_with(data["request_id"], "weather", "weather_current", 60, True)
|
||||
|
||||
def test_a_callers_request_id_is_passed_through(self, api_v3_client, service):
|
||||
with patch(f"{CLIENT}.on_demand_start", side_effect=_ack) as start:
|
||||
data = api_v3_client.post(START_URL, json={
|
||||
"plugin_id": "weather", "request_id": "ha-123"}).get_json()["data"]
|
||||
assert data["request_id"] == "ha-123"
|
||||
assert start.call_args.args[0] == "ha-123"
|
||||
|
||||
def test_stop_goes_over_the_socket(self, api_v3_client, service):
|
||||
with patch(f"{CLIENT}.on_demand_stop", side_effect=_ack) as stop:
|
||||
data = api_v3_client.post(STOP_URL, json={}).get_json()["data"]
|
||||
assert data["transport"] == "socket"
|
||||
stop.assert_called_once_with(data["request_id"])
|
||||
assert _mailbox_writes(service["cache"]) == []
|
||||
|
||||
|
||||
class TestMailboxFallback:
|
||||
@pytest.mark.parametrize("reason", [
|
||||
"no_socket", "refused", "timeout", "closed", "bad_response", "invalid_request",
|
||||
"busy", "unknown_command", "unsupported_version", "disabled", "unsupported",
|
||||
])
|
||||
def test_any_socket_failure_writes_the_mailbox_as_before(
|
||||
self, api_v3_client, service, reason):
|
||||
with patch(f"{CLIENT}.on_demand_start",
|
||||
side_effect=control_client.ControlError(reason, "x")):
|
||||
resp = api_v3_client.post(START_URL, json={
|
||||
"plugin_id": "weather", "mode": "weather_current",
|
||||
"duration": 60, "pinned": True})
|
||||
assert resp.status_code == 200
|
||||
data = resp.get_json()["data"]
|
||||
assert data["transport"] == "mailbox"
|
||||
assert data["socket_error"] == reason
|
||||
[write] = _mailbox_writes(service["cache"])
|
||||
assert write["request_id"] == data["request_id"]
|
||||
assert write["action"] == "start"
|
||||
assert (write["plugin_id"], write["mode"], write["duration"], write["pinned"]) == \
|
||||
("weather", "weather_current", 60, True)
|
||||
|
||||
def test_a_client_bug_still_falls_back(self, api_v3_client, service):
|
||||
with patch(f"{CLIENT}.on_demand_start", side_effect=RuntimeError("boom")):
|
||||
resp = api_v3_client.post(START_URL, json={"plugin_id": "weather"})
|
||||
assert resp.status_code == 200
|
||||
assert resp.get_json()["data"]["socket_error"] == "internal"
|
||||
assert len(_mailbox_writes(service["cache"])) == 1
|
||||
|
||||
def test_stop_falls_back(self, api_v3_client, service):
|
||||
with patch(f"{CLIENT}.on_demand_stop",
|
||||
side_effect=control_client.ControlError("timeout")):
|
||||
data = api_v3_client.post(STOP_URL, json={}).get_json()["data"]
|
||||
assert data["transport"] == "mailbox"
|
||||
[write] = _mailbox_writes(service["cache"])
|
||||
assert write == {"request_id": data["request_id"], "action": "stop",
|
||||
"timestamp": write["timestamp"]}
|
||||
|
||||
def test_a_stopped_display_gets_the_mailbox_before_it_is_started(
|
||||
self, api_v3_client, service):
|
||||
service["state"]["active"] = False
|
||||
with patch(f"{CLIENT}.on_demand_start",
|
||||
side_effect=control_client.ControlError("no_socket")):
|
||||
resp = api_v3_client.post(START_URL, json={"plugin_id": "weather"})
|
||||
assert resp.status_code == 200
|
||||
assert service["calls"] == [("cache", MAILBOX), ("systemctl", "start")]
|
||||
|
||||
def test_the_socket_is_off_in_the_test_suite(self, api_v3_client, service):
|
||||
# conftest's _hermetic_control_socket: a suite run on a device must
|
||||
# not drive the live display.
|
||||
assert os.environ[c.SOCKET_PATH_ENV] == "off"
|
||||
data = api_v3_client.post(START_URL, json={"plugin_id": "weather"}).get_json()["data"]
|
||||
assert data["transport"] == "mailbox"
|
||||
assert data["socket_error"] in ("disabled", "unsupported") # Linux, Windows
|
||||
|
||||
|
||||
@pytest.mark.skipif(not c.socket_supported(), reason="AF_UNIX sockets are Linux/macOS only")
|
||||
class TestRealSocket:
|
||||
@pytest.fixture
|
||||
def live(self, monkeypatch):
|
||||
import shutil
|
||||
import tempfile
|
||||
from src.ipc.server import ControlServer
|
||||
d = tempfile.mkdtemp(prefix="lmipc-")
|
||||
path = os.path.join(d, "control.sock")
|
||||
server = ControlServer(path, status_provider=dict)
|
||||
assert server.start()
|
||||
monkeypatch.setenv(c.SOCKET_PATH_ENV, path)
|
||||
yield server
|
||||
server.close()
|
||||
shutil.rmtree(d, ignore_errors=True)
|
||||
|
||||
def test_start_is_acked_and_queued(self, api_v3_client, service, live):
|
||||
data = api_v3_client.post(START_URL, json={
|
||||
"plugin_id": "weather", "duration": "30"}).get_json()["data"]
|
||||
assert data["transport"] == "socket"
|
||||
[cmd] = live.drain()
|
||||
payload = cmd.as_on_demand_request()
|
||||
assert payload["request_id"] == data["request_id"]
|
||||
assert payload["plugin_id"] == "weather" and payload["duration"] == 30.0
|
||||
assert _mailbox_writes(service["cache"]) == []
|
||||
|
||||
def test_stop_is_acked_and_queued(self, api_v3_client, service, live):
|
||||
data = api_v3_client.post(STOP_URL, json={}).get_json()["data"]
|
||||
assert data["transport"] == "socket"
|
||||
assert [x.request_id for x in live.drain()] == [data["request_id"]]
|
||||
|
||||
def test_a_display_that_went_away_falls_back(self, api_v3_client, service, live):
|
||||
live.close()
|
||||
data = api_v3_client.post(START_URL, json={"plugin_id": "weather"}).get_json()["data"]
|
||||
assert data["transport"] == "mailbox" and data["socket_error"] == "no_socket"
|
||||
assert len(_mailbox_writes(service["cache"])) == 1
|
||||
@@ -0,0 +1,275 @@
|
||||
"""The control socket's contract (src/ipc/contract.py): messages and framing.
|
||||
|
||||
Pure data, so every test here runs on every platform. What they pin:
|
||||
|
||||
* a request and a response survive encode -> decode -> parse unchanged, and
|
||||
the on-demand arguments carry exactly what the file mailbox carries;
|
||||
* the envelope and the arguments refuse what the display could not act on
|
||||
(missing ids, wrong types, a non-finite duration) with a stable error code;
|
||||
* framing never holds more than one message's worth of bytes, however the
|
||||
bytes arrive;
|
||||
* where the socket is looked for, and how it is switched off.
|
||||
"""
|
||||
|
||||
import json
|
||||
import math
|
||||
|
||||
import pytest
|
||||
|
||||
from src.ipc import contract as c
|
||||
from src.ipc.contract import (
|
||||
Command, ErrorCode, FrameReader, OnDemandStartArgs, OnDemandStopArgs,
|
||||
ProtocolError, Request, Response,
|
||||
)
|
||||
|
||||
|
||||
def _wire(obj):
|
||||
"""Encode then decode, as one side's bytes reach the other."""
|
||||
data = c.encode_message(obj)
|
||||
assert data.endswith(b'\n') and data.count(b'\n') == 1
|
||||
return c.decode_message(data)
|
||||
|
||||
|
||||
class TestRoundTrip:
|
||||
def test_request(self):
|
||||
req = Request(id='abc-1', cmd=Command.ON_DEMAND_START,
|
||||
args={'plugin_id': 'clock', 'mode': None, 'duration': 30.0,
|
||||
'pinned': True})
|
||||
back = Request.from_dict(_wire(req.to_dict()))
|
||||
assert back == req
|
||||
assert back.v == c.PROTOCOL_VERSION
|
||||
|
||||
def test_success_response(self):
|
||||
resp = Response.success('abc-1', {'accepted': True, 'request_id': 'abc-1', 'queued': 1})
|
||||
back = Response.from_dict(_wire(resp.to_dict()))
|
||||
assert back == resp
|
||||
assert back.ok and back.error is None
|
||||
|
||||
def test_failure_response(self):
|
||||
resp = Response.failure('abc-1', ErrorCode.BUSY, 'queue full')
|
||||
wire = _wire(resp.to_dict())
|
||||
assert wire == {'v': 1, 'id': 'abc-1', 'ok': False,
|
||||
'error': {'code': 'busy', 'message': 'queue full'}}
|
||||
assert Response.from_dict(wire) == resp
|
||||
|
||||
def test_failure_without_an_id(self):
|
||||
wire = _wire(Response.failure(None, ErrorCode.BAD_JSON, 'nope').to_dict())
|
||||
assert wire['id'] is None
|
||||
assert Response.from_dict(wire).id is None
|
||||
|
||||
def test_start_args_round_trip(self):
|
||||
args = OnDemandStartArgs(plugin_id='clock', mode='clock_main', duration=45.0,
|
||||
pinned=True)
|
||||
assert OnDemandStartArgs.from_dict(_wire(args.to_dict())) == args
|
||||
|
||||
def test_encoded_messages_are_ascii_single_lines(self):
|
||||
data = c.encode_message({'v': 1, 'id': 'x', 'cmd': 'ping',
|
||||
'args': {'text': 'line1\nline2 café'}})
|
||||
assert data.count(b'\n') == 1
|
||||
data.decode('ascii')
|
||||
assert c.decode_message(data)['args']['text'] == 'line1\nline2 café'
|
||||
|
||||
|
||||
class TestEnvelopeValidation:
|
||||
@pytest.mark.parametrize('obj', [
|
||||
[], 'x', 1, None,
|
||||
])
|
||||
def test_not_an_object(self, obj):
|
||||
with pytest.raises(ProtocolError) as e:
|
||||
Request.from_dict(obj)
|
||||
assert e.value.code == ErrorCode.BAD_REQUEST
|
||||
|
||||
@pytest.mark.parametrize('bad_id', [None, '', 7, 'x' * (c.MAX_ID_LENGTH + 1), 'a\nb'])
|
||||
def test_bad_id(self, bad_id):
|
||||
with pytest.raises(ProtocolError) as e:
|
||||
Request.from_dict({'v': 1, 'id': bad_id, 'cmd': 'ping'})
|
||||
assert e.value.code == ErrorCode.BAD_REQUEST
|
||||
assert e.value.request_id is None
|
||||
|
||||
@pytest.mark.parametrize('v', [None, '1', 1.0, True])
|
||||
def test_bad_version_type_keeps_the_id(self, v):
|
||||
with pytest.raises(ProtocolError) as e:
|
||||
Request.from_dict({'v': v, 'id': 'r1', 'cmd': 'ping'})
|
||||
assert e.value.code == ErrorCode.BAD_REQUEST
|
||||
assert e.value.request_id == 'r1'
|
||||
|
||||
def test_missing_cmd(self):
|
||||
with pytest.raises(ProtocolError) as e:
|
||||
Request.from_dict({'v': 1, 'id': 'r1'})
|
||||
assert e.value.code == ErrorCode.BAD_REQUEST
|
||||
|
||||
def test_args_default_to_empty(self):
|
||||
assert Request.from_dict({'v': 1, 'id': 'r', 'cmd': 'ping'}).args == {}
|
||||
assert Request.from_dict({'v': 1, 'id': 'r', 'cmd': 'ping', 'args': None}).args == {}
|
||||
|
||||
def test_args_must_be_an_object(self):
|
||||
with pytest.raises(ProtocolError) as e:
|
||||
Request.from_dict({'v': 1, 'id': 'r', 'cmd': 'ping', 'args': [1]})
|
||||
assert e.value.code == ErrorCode.BAD_REQUEST
|
||||
|
||||
def test_an_unknown_version_parses(self):
|
||||
# The server, not the parser, decides about versions, so that hello
|
||||
# can negotiate.
|
||||
assert Request.from_dict({'v': 99, 'id': 'r', 'cmd': 'hello'}).v == 99
|
||||
|
||||
@pytest.mark.parametrize('obj', [
|
||||
{'v': 1, 'id': 'r', 'ok': 'yes'},
|
||||
{'v': 1, 'id': 'r', 'ok': False},
|
||||
{'v': 1, 'id': 'r', 'ok': False, 'error': {'message': 'x'}},
|
||||
{'v': 1, 'id': 5, 'ok': True},
|
||||
{'v': 1, 'id': 'r', 'ok': True, 'result': [1]},
|
||||
{'id': 'r', 'ok': True},
|
||||
])
|
||||
def test_malformed_responses(self, obj):
|
||||
with pytest.raises(ProtocolError):
|
||||
Response.from_dict(obj)
|
||||
|
||||
|
||||
class TestOnDemandArgs:
|
||||
def test_plugin_or_mode_is_required(self):
|
||||
with pytest.raises(ProtocolError) as e:
|
||||
OnDemandStartArgs.from_dict({'duration': 10})
|
||||
assert e.value.code == ErrorCode.INVALID_ARGS
|
||||
|
||||
def test_mode_alone_is_enough(self):
|
||||
assert OnDemandStartArgs.from_dict({'mode': 'nfl_live'}).mode == 'nfl_live'
|
||||
|
||||
@pytest.mark.parametrize('raw, seconds', [
|
||||
(None, None), ('', None), (0, None), (45, 45.0), (2.5, 2.5), ('30', 30.0),
|
||||
])
|
||||
def test_duration(self, raw, seconds):
|
||||
assert OnDemandStartArgs.from_dict({'plugin_id': 'p', 'duration': raw}).duration == seconds
|
||||
|
||||
@pytest.mark.parametrize('raw', [-1, 'soon', True, [5], math.inf, math.nan, 'inf'])
|
||||
def test_bad_duration(self, raw):
|
||||
with pytest.raises(ProtocolError) as e:
|
||||
OnDemandStartArgs.from_dict({'plugin_id': 'p', 'duration': raw})
|
||||
assert e.value.code == ErrorCode.INVALID_ARGS
|
||||
|
||||
@pytest.mark.parametrize('pinned', ['true', 1, 'false'])
|
||||
def test_pinned_must_be_a_real_boolean(self, pinned):
|
||||
# The web route coerces "false" to False before it gets here; the
|
||||
# contract does not guess (bool("false") is True).
|
||||
with pytest.raises(ProtocolError):
|
||||
OnDemandStartArgs.from_dict({'plugin_id': 'p', 'pinned': pinned})
|
||||
|
||||
@pytest.mark.parametrize('name', [5, 'x' * (c.MAX_NAME_LENGTH + 1), 'a\nb'])
|
||||
def test_bad_names(self, name):
|
||||
with pytest.raises(ProtocolError):
|
||||
OnDemandStartArgs.from_dict({'plugin_id': name})
|
||||
|
||||
def test_unknown_command(self):
|
||||
with pytest.raises(ProtocolError) as e:
|
||||
c.parse_args('reboot', {})
|
||||
assert e.value.code == ErrorCode.UNKNOWN_COMMAND
|
||||
|
||||
@pytest.mark.parametrize('cmd', c.COMMANDS)
|
||||
def test_every_command_has_an_argument_type(self, cmd):
|
||||
args = {'plugin_id': 'p'} if cmd == Command.ON_DEMAND_START else {}
|
||||
c.parse_args(cmd, args)
|
||||
|
||||
def test_hello_versions(self):
|
||||
assert c.HelloArgs.from_dict({'versions': [1, 2], 'client': 'web'}).versions == (1, 2)
|
||||
for bad in ([], ['1'], 'x', [True]):
|
||||
with pytest.raises(ProtocolError):
|
||||
c.HelloArgs.from_dict({'versions': bad})
|
||||
|
||||
def test_negotiation(self):
|
||||
assert c.negotiate_version((1,)) == 1
|
||||
assert c.negotiate_version((1, 7)) == 1
|
||||
assert c.negotiate_version((7,)) is None
|
||||
|
||||
|
||||
class TestMailboxShape:
|
||||
"""Socket commands are handed to the mailbox's own handler, so they must
|
||||
look exactly like what the web route writes to the mailbox."""
|
||||
|
||||
def test_start(self):
|
||||
args = OnDemandStartArgs(plugin_id='clock', mode='clock_main', duration=60.0,
|
||||
pinned=True)
|
||||
payload = c.on_demand_request('rid', args, 123.0)
|
||||
assert payload == {'request_id': 'rid', 'action': 'start', 'plugin_id': 'clock',
|
||||
'mode': 'clock_main', 'duration': 60.0, 'pinned': True,
|
||||
'timestamp': 123.0, 'source': 'socket'}
|
||||
|
||||
def test_stop(self):
|
||||
payload = c.on_demand_request('rid', OnDemandStopArgs(), 5.0)
|
||||
assert payload['action'] == 'stop' and payload['request_id'] == 'rid'
|
||||
|
||||
|
||||
class TestFraming:
|
||||
def test_one_message_in_pieces(self):
|
||||
data = c.encode_message({'v': 1, 'id': 'a', 'cmd': 'ping'})
|
||||
reader = FrameReader()
|
||||
out = []
|
||||
for i in range(len(data)):
|
||||
out += reader.feed(data[i:i + 1])
|
||||
assert [json.loads(x) for x in out] == [{'v': 1, 'id': 'a', 'cmd': 'ping'}]
|
||||
assert reader.pending == 0
|
||||
|
||||
def test_several_messages_in_one_chunk(self):
|
||||
data = b''.join(c.encode_message({'n': n}) for n in range(3))
|
||||
assert [json.loads(x)['n'] for x in FrameReader().feed(data)] == [0, 1, 2]
|
||||
|
||||
def test_blank_lines_are_skipped(self):
|
||||
assert FrameReader().feed(b'\n\r\n \n{"a":1}\n') == [b'{"a":1}']
|
||||
|
||||
def test_a_line_over_the_limit_is_refused(self):
|
||||
reader = FrameReader(max_bytes=32)
|
||||
with pytest.raises(ProtocolError) as e:
|
||||
reader.feed(b'x' * 40 + b'\n')
|
||||
assert e.value.code == ErrorCode.MESSAGE_TOO_LARGE
|
||||
|
||||
def test_a_line_that_never_ends_is_refused_at_the_limit(self):
|
||||
reader = FrameReader(max_bytes=32)
|
||||
reader.feed(b'x' * 31)
|
||||
with pytest.raises(ProtocolError) as e:
|
||||
reader.feed(b'x')
|
||||
assert e.value.code == ErrorCode.MESSAGE_TOO_LARGE
|
||||
|
||||
def test_exactly_the_limit_is_allowed(self):
|
||||
reader = FrameReader(max_bytes=8)
|
||||
assert reader.feed(b'1234567\n') == [b'1234567']
|
||||
|
||||
def test_encode_refuses_an_oversize_message(self):
|
||||
with pytest.raises(ProtocolError) as e:
|
||||
c.encode_message({'blob': 'x' * c.MAX_MESSAGE_BYTES})
|
||||
assert e.value.code == ErrorCode.MESSAGE_TOO_LARGE
|
||||
|
||||
def test_encode_refuses_non_json(self):
|
||||
with pytest.raises(ProtocolError):
|
||||
c.encode_message({'x': math.nan})
|
||||
with pytest.raises(ProtocolError):
|
||||
c.encode_message({'x': object()})
|
||||
|
||||
@pytest.mark.parametrize('line', [b'{', b'[1,2]', b'"x"', b'\xff\xfe', b'null'])
|
||||
def test_decode_garbage(self, line):
|
||||
with pytest.raises(ProtocolError) as e:
|
||||
c.decode_message(line)
|
||||
assert e.value.code == ErrorCode.BAD_JSON
|
||||
|
||||
|
||||
class TestSocketLocation:
|
||||
def test_default(self):
|
||||
paths = c.client_socket_paths({})
|
||||
assert paths[0] == '/run/ledmatrix/control.sock'
|
||||
assert len(paths) == 2 and paths[1].endswith('control.sock')
|
||||
|
||||
def test_configured_path_is_the_only_one_tried(self):
|
||||
assert c.client_socket_paths({c.SOCKET_PATH_ENV: '/x/y.sock'}) == ['/x/y.sock']
|
||||
|
||||
@pytest.mark.parametrize('value', ['off', 'OFF', '0', 'false', 'disabled', ' none '])
|
||||
def test_switched_off(self, value):
|
||||
env = {c.SOCKET_PATH_ENV: value}
|
||||
assert c.socket_disabled(env)
|
||||
assert c.client_socket_paths(env) == []
|
||||
assert c.configured_socket_path(env) is None
|
||||
|
||||
def test_dev_path_is_per_user(self):
|
||||
assert c.dev_socket_path(1000) != c.dev_socket_path(1001)
|
||||
assert 'ledmatrix-1000' in c.dev_socket_path(1000)
|
||||
|
||||
def test_the_default_dir_is_the_heartbeats(self):
|
||||
# One RuntimeDirectory= serves both (#687).
|
||||
from src import display_watchdog
|
||||
assert c.DEFAULT_SOCKET_DIR == display_watchdog.HEARTBEAT_DIR
|
||||
@@ -0,0 +1,229 @@
|
||||
"""DisplayController's side of the control socket.
|
||||
|
||||
The server's handlers only queue; the render thread drains the queue where
|
||||
it reads the file mailbox (_poll_on_demand_requests) and hands each command
|
||||
to the mailbox's own handler (_handle_on_demand_request). These tests pin
|
||||
that hook:
|
||||
|
||||
* a socket command is applied by the same code as a mailbox request, with
|
||||
its request id, and without waiting for the mailbox's 0.25 s read floor;
|
||||
* a request that arrives both ways (a client that timed out after the
|
||||
command was queued, then wrote the mailbox) is activated once;
|
||||
* a command that fails is contained, and the ones after it still run;
|
||||
* cleanup closes the socket; a disabled socket changes nothing.
|
||||
"""
|
||||
|
||||
import os
|
||||
import time
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from src.ipc import client
|
||||
from src.ipc import contract as c
|
||||
from src.ipc.contract import Command, OnDemandStartArgs, OnDemandStopArgs
|
||||
from src.ipc.server import QueuedCommand
|
||||
|
||||
|
||||
def _start(rid, plugin_id='clock', **kw):
|
||||
return QueuedCommand(request_id=rid, cmd=Command.ON_DEMAND_START,
|
||||
args=OnDemandStartArgs(plugin_id=plugin_id, **kw),
|
||||
received_at=time.time())
|
||||
|
||||
|
||||
def _stop(rid):
|
||||
return QueuedCommand(request_id=rid, cmd=Command.ON_DEMAND_STOP,
|
||||
args=OnDemandStopArgs(), received_at=time.time())
|
||||
|
||||
|
||||
class FakeServer:
|
||||
def __init__(self, *commands):
|
||||
self.commands = list(commands)
|
||||
self.closed = False
|
||||
|
||||
@property
|
||||
def has_pending(self):
|
||||
return bool(self.commands)
|
||||
|
||||
def drain(self):
|
||||
out, self.commands = self.commands, []
|
||||
return out
|
||||
|
||||
def close(self):
|
||||
self.closed = True
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def controller(test_display_controller):
|
||||
c_ = test_display_controller
|
||||
c_.on_demand_active = False
|
||||
c_.on_demand_request_id = None
|
||||
c_._last_on_demand_poll = None
|
||||
mailbox = {'value': None}
|
||||
|
||||
def fake_get(key, *a, **kw):
|
||||
if key == 'display_on_demand_request':
|
||||
return mailbox['value']
|
||||
return None
|
||||
|
||||
def fake_delete(key):
|
||||
if key == 'display_on_demand_request':
|
||||
mailbox['value'] = None
|
||||
|
||||
c_.cache_manager.get = MagicMock(side_effect=fake_get)
|
||||
c_.cache_manager.set = MagicMock()
|
||||
c_.cache_manager.delete = MagicMock(side_effect=fake_delete)
|
||||
c_._activate_on_demand = MagicMock()
|
||||
c_.mailbox = mailbox
|
||||
return c_
|
||||
|
||||
|
||||
class TestDrain:
|
||||
def test_a_socket_start_goes_through_the_mailbox_handler(self, controller):
|
||||
controller._control_server = FakeServer(_start('sock-1', duration=30.0, pinned=True))
|
||||
controller._poll_on_demand_requests()
|
||||
controller._activate_on_demand.assert_called_once()
|
||||
request = controller._activate_on_demand.call_args.args[0]
|
||||
assert request['request_id'] == 'sock-1'
|
||||
assert request['action'] == 'start'
|
||||
assert request['plugin_id'] == 'clock'
|
||||
assert request['duration'] == 30.0 and request['pinned'] is True
|
||||
assert controller.on_demand_request_id == 'sock-1'
|
||||
# The same restart-replay guard as a mailbox request.
|
||||
controller.cache_manager.set.assert_any_call(
|
||||
'display_on_demand_processed_id', 'sock-1', ttl=3600)
|
||||
|
||||
def test_socket_commands_skip_the_mailbox_floor(self, controller):
|
||||
server = FakeServer()
|
||||
controller._control_server = server
|
||||
controller._poll_on_demand_requests() # reads the mailbox, sets the floor
|
||||
reads = controller.cache_manager.get.call_count
|
||||
server.commands.append(_start('quick'))
|
||||
controller._poll_on_demand_requests() # within the floor
|
||||
controller._activate_on_demand.assert_called_once()
|
||||
mailbox_reads = [call for call in controller.cache_manager.get.call_args_list[reads:]
|
||||
if call.args[0] == 'display_on_demand_request']
|
||||
# Only _consume_on_demand_request's compare-before-delete re-read.
|
||||
assert len(mailbox_reads) <= 1
|
||||
|
||||
def test_a_request_that_came_both_ways_is_activated_once(self, controller):
|
||||
controller._control_server = FakeServer(_start('dup'))
|
||||
controller.mailbox['value'] = {'request_id': 'dup', 'action': 'start',
|
||||
'plugin_id': 'clock'}
|
||||
controller._poll_on_demand_requests()
|
||||
controller._last_on_demand_poll = None
|
||||
controller._poll_on_demand_requests()
|
||||
controller._activate_on_demand.assert_called_once()
|
||||
assert controller.mailbox['value'] is None, "the duplicate was left in the mailbox"
|
||||
|
||||
def test_a_fallback_write_landing_later_is_ignored(self, controller):
|
||||
controller._control_server = FakeServer(_start('late'))
|
||||
controller._poll_on_demand_requests()
|
||||
controller.mailbox['value'] = {'request_id': 'late', 'action': 'start',
|
||||
'plugin_id': 'clock'}
|
||||
controller._last_on_demand_poll = None
|
||||
controller._poll_on_demand_requests()
|
||||
controller._activate_on_demand.assert_called_once()
|
||||
|
||||
def test_the_mailbox_still_works_alongside(self, controller):
|
||||
controller._control_server = FakeServer()
|
||||
controller.mailbox['value'] = {'request_id': 'mb', 'action': 'start', 'plugin_id': 'p'}
|
||||
controller._poll_on_demand_requests()
|
||||
assert controller._activate_on_demand.call_args.args[0]['request_id'] == 'mb'
|
||||
|
||||
def test_a_socket_stop_ends_on_demand(self, controller):
|
||||
controller.on_demand_active = True
|
||||
controller._clear_on_demand = MagicMock()
|
||||
controller._control_server = FakeServer(_stop('halt'))
|
||||
controller._poll_on_demand_requests()
|
||||
controller._clear_on_demand.assert_called_once_with(reason='requested-stop')
|
||||
|
||||
def test_commands_run_in_arrival_order(self, controller):
|
||||
seen = []
|
||||
controller._activate_on_demand = MagicMock(
|
||||
side_effect=lambda r: seen.append(r['request_id']))
|
||||
controller._control_server = FakeServer(_start('a'), _start('b'), _start('c'))
|
||||
controller._poll_on_demand_requests()
|
||||
assert seen == ['a', 'b', 'c']
|
||||
|
||||
def test_a_failing_command_is_contained(self, controller):
|
||||
calls = []
|
||||
|
||||
def activate(request):
|
||||
calls.append(request['request_id'])
|
||||
if request['request_id'] == 'bad':
|
||||
raise RuntimeError('plugin exploded')
|
||||
|
||||
controller._activate_on_demand = MagicMock(side_effect=activate)
|
||||
controller._control_server = FakeServer(_start('bad'), _start('good'))
|
||||
controller._poll_on_demand_requests()
|
||||
assert calls == ['bad', 'good']
|
||||
|
||||
def test_no_server_means_mailbox_only(self, controller):
|
||||
controller._control_server = None
|
||||
controller._poll_on_demand_requests()
|
||||
controller._activate_on_demand.assert_not_called()
|
||||
|
||||
|
||||
class TestPendingChangesFloor:
|
||||
def test_a_queued_command_skips_the_floor(self, controller):
|
||||
server = FakeServer()
|
||||
controller._control_server = server
|
||||
controller._service_pending_changes()
|
||||
server.commands.append(_start('now'))
|
||||
controller._service_pending_changes() # well inside the 0.25 s floor
|
||||
controller._activate_on_demand.assert_called_once()
|
||||
|
||||
def test_nothing_queued_keeps_the_floor(self, controller):
|
||||
controller._control_server = FakeServer()
|
||||
controller._poll_on_demand_requests = MagicMock()
|
||||
controller._service_pending_changes()
|
||||
controller._service_pending_changes()
|
||||
assert controller._poll_on_demand_requests.call_count == 1
|
||||
|
||||
|
||||
class TestLifecycle:
|
||||
def test_status_snapshot(self, controller):
|
||||
controller.current_display_mode = 'clock_main'
|
||||
controller.on_demand_active = True
|
||||
controller.on_demand_plugin_id = 'clock'
|
||||
controller.on_demand_expires_at = None
|
||||
status = controller._control_status()
|
||||
assert status['current_mode'] == 'clock_main'
|
||||
assert status['on_demand']['active'] is True
|
||||
assert status['on_demand']['plugin_id'] == 'clock'
|
||||
c.encode_message(status) # it has to fit on the wire
|
||||
|
||||
def test_cleanup_closes_the_socket(self, controller):
|
||||
server = FakeServer()
|
||||
controller._control_server = server
|
||||
controller.cleanup()
|
||||
assert server.closed
|
||||
assert controller._control_server is None
|
||||
|
||||
def test_disabled_socket_starts_nothing(self, controller):
|
||||
# conftest sets LEDMATRIX_CONTROL_SOCKET=off for every test.
|
||||
controller._start_control_server()
|
||||
assert controller._control_server is None
|
||||
|
||||
@pytest.mark.skipif(not c.socket_supported(), reason='AF_UNIX sockets are Linux/macOS only')
|
||||
def test_end_to_end(self, controller, monkeypatch):
|
||||
import shutil
|
||||
import tempfile
|
||||
d = tempfile.mkdtemp(prefix='lmipc-')
|
||||
path = os.path.join(d, 'control.sock')
|
||||
monkeypatch.setenv(c.SOCKET_PATH_ENV, path)
|
||||
try:
|
||||
controller._start_control_server()
|
||||
assert controller._control_server is not None
|
||||
ack = client.on_demand_start('e2e', 'clock', None, 15, False, paths=[path])
|
||||
assert ack['accepted'] is True and ack['request_id'] == 'e2e'
|
||||
status = client.on_demand_status(paths=[path])
|
||||
assert 'on_demand' in status and 'current_mode' in status
|
||||
controller._service_pending_changes()
|
||||
request = controller._activate_on_demand.call_args.args[0]
|
||||
assert request['request_id'] == 'e2e' and request['duration'] == 15.0
|
||||
controller.cleanup()
|
||||
assert not os.path.exists(path)
|
||||
finally:
|
||||
shutil.rmtree(d, ignore_errors=True)
|
||||
@@ -0,0 +1,570 @@
|
||||
"""The display side of the control socket (src/ipc/server.py).
|
||||
|
||||
Two layers:
|
||||
|
||||
* ``handle_line`` and the permission model are plain functions of their
|
||||
input, tested on every platform: every request gets exactly one answer,
|
||||
garbage is answered rather than raised, queued commands are acked with
|
||||
their request id, and the render thread drains them in order.
|
||||
* The socket itself (``TestLiveSocket``, ``TestPermissions``) needs AF_UNIX,
|
||||
so those tests are skipped on Windows and run on Linux (CI, WSL, a Pi): a
|
||||
real server on a tmp_path socket, driven by the real client and by raw
|
||||
sockets that misbehave -- garbage, oversize lines, a client that hangs up
|
||||
mid-message, one that never finishes -- while the server keeps serving.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import socket
|
||||
import stat
|
||||
import threading
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from src.ipc import client
|
||||
from src.ipc import contract as c
|
||||
from src.ipc import server as srv
|
||||
from src.ipc.contract import Command, ErrorCode
|
||||
from src.ipc.server import ControlServer, PeerCredentials, peer_allowed
|
||||
|
||||
needs_unix_sockets = pytest.mark.skipif(not c.socket_supported(),
|
||||
reason='AF_UNIX sockets are Linux/macOS only')
|
||||
|
||||
|
||||
def _line(obj):
|
||||
return json.dumps(obj).encode()
|
||||
|
||||
|
||||
def _req(cmd, args=None, rid='r1', v=1):
|
||||
return _line({'v': v, 'id': rid, 'cmd': cmd, 'args': args or {}})
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def status():
|
||||
return {'on_demand': {'active': False, 'status': 'idle'}, 'current_mode': 'clock'}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def server(status, tmp_path):
|
||||
"""A server that is never started: handle_line and drain only."""
|
||||
return ControlServer(str(tmp_path / 'unused.sock'), status_provider=lambda: dict(status),
|
||||
queue_size=3)
|
||||
|
||||
|
||||
class TestHandleLine:
|
||||
def test_ping(self, server):
|
||||
resp = server.handle_line(_req(Command.PING))
|
||||
assert resp.ok and resp.id == 'r1' and resp.result == {'pong': True}
|
||||
|
||||
def test_hello_negotiates(self, server):
|
||||
resp = server.handle_line(_req(Command.HELLO, {'versions': [1, 5], 'client': 't'}, v=5))
|
||||
assert resp.ok
|
||||
assert resp.result['version'] == 1
|
||||
assert resp.result['commands'] == list(c.COMMANDS)
|
||||
assert resp.result['max_message_bytes'] == c.MAX_MESSAGE_BYTES
|
||||
assert resp.v == 1
|
||||
|
||||
def test_hello_with_nothing_in_common(self, server):
|
||||
resp = server.handle_line(_req(Command.HELLO, {'versions': [9]}, v=9))
|
||||
assert not resp.ok and resp.error.code == ErrorCode.UNSUPPORTED_VERSION
|
||||
|
||||
def test_other_commands_need_a_supported_version(self, server):
|
||||
resp = server.handle_line(_req(Command.PING, v=2))
|
||||
assert not resp.ok and resp.error.code == ErrorCode.UNSUPPORTED_VERSION
|
||||
assert resp.id == 'r1'
|
||||
|
||||
@pytest.mark.parametrize('line, code, rid', [
|
||||
(b'not json', ErrorCode.BAD_JSON, None),
|
||||
(b'[1]', ErrorCode.BAD_JSON, None),
|
||||
(b'\xff', ErrorCode.BAD_JSON, None),
|
||||
(_line({'v': 1, 'cmd': 'ping'}), ErrorCode.BAD_REQUEST, None),
|
||||
(_line({'v': 'one', 'id': 'q', 'cmd': 'ping'}), ErrorCode.BAD_REQUEST, 'q'),
|
||||
(_req('shutdown_the_pi'), ErrorCode.UNKNOWN_COMMAND, 'r1'),
|
||||
(_req(Command.ON_DEMAND_START, {}), ErrorCode.INVALID_ARGS, 'r1'),
|
||||
(_req(Command.ON_DEMAND_START, {'plugin_id': 'p', 'duration': 'x'}),
|
||||
ErrorCode.INVALID_ARGS, 'r1'),
|
||||
])
|
||||
def test_garbage_is_answered_not_raised(self, server, line, code, rid):
|
||||
resp = server.handle_line(line)
|
||||
assert not resp.ok
|
||||
assert resp.error.code == code
|
||||
assert resp.id == rid
|
||||
assert server.drain() == []
|
||||
|
||||
def test_start_is_queued_and_acked(self, server):
|
||||
resp = server.handle_line(_req(Command.ON_DEMAND_START,
|
||||
{'plugin_id': 'clock', 'duration': 30, 'pinned': True},
|
||||
rid='abc'))
|
||||
assert resp.ok
|
||||
assert resp.result == {'accepted': True, 'request_id': 'abc', 'queued': 1}
|
||||
assert server.has_pending
|
||||
[cmd] = server.drain()
|
||||
assert not server.has_pending
|
||||
payload = cmd.as_on_demand_request()
|
||||
assert payload['request_id'] == 'abc'
|
||||
assert payload['action'] == 'start'
|
||||
assert payload['plugin_id'] == 'clock'
|
||||
assert payload['duration'] == 30.0 and payload['pinned'] is True
|
||||
|
||||
def test_stop_is_queued_and_acked(self, server):
|
||||
resp = server.handle_line(_req(Command.ON_DEMAND_STOP, rid='s1'))
|
||||
assert resp.ok and resp.result['request_id'] == 's1'
|
||||
assert [x.as_on_demand_request()['action'] for x in server.drain()] == ['stop']
|
||||
|
||||
def test_drain_keeps_arrival_order(self, server):
|
||||
for rid in ('a', 'b', 'c'):
|
||||
server.handle_line(_req(Command.ON_DEMAND_START, {'plugin_id': 'p'}, rid=rid))
|
||||
assert [x.request_id for x in server.drain()] == ['a', 'b', 'c']
|
||||
assert server.drain() == []
|
||||
|
||||
def test_a_full_queue_says_busy_and_queues_nothing_more(self, server):
|
||||
for rid in ('a', 'b', 'c'):
|
||||
assert server.handle_line(_req(Command.ON_DEMAND_STOP, rid=rid)).ok
|
||||
resp = server.handle_line(_req(Command.ON_DEMAND_STOP, rid='d'))
|
||||
assert not resp.ok and resp.error.code == ErrorCode.BUSY and resp.id == 'd'
|
||||
assert [x.request_id for x in server.drain()] == ['a', 'b', 'c']
|
||||
|
||||
def test_status_answers_from_the_provider_without_queueing(self, server, status):
|
||||
status['on_demand']['active'] = True
|
||||
resp = server.handle_line(_req(Command.ON_DEMAND_STATUS))
|
||||
assert resp.ok and resp.result['on_demand']['active'] is True
|
||||
assert not server.has_pending
|
||||
|
||||
def test_a_failing_status_provider_is_an_internal_error(self, tmp_path):
|
||||
def boom():
|
||||
raise RuntimeError('render thread mid-update')
|
||||
s = ControlServer(str(tmp_path / 'x.sock'), status_provider=boom)
|
||||
resp = s.handle_line(_req(Command.ON_DEMAND_STATUS, rid='z'))
|
||||
assert not resp.ok and resp.error.code == ErrorCode.INTERNAL and resp.id == 'z'
|
||||
|
||||
|
||||
class TestPermissionModel:
|
||||
"""root, the display's own user, or the shared group -- nobody else."""
|
||||
OWN, GROUP = 0, 990
|
||||
|
||||
@pytest.mark.parametrize('cred, groups, allowed', [
|
||||
(PeerCredentials(1, 0, 0), None, True), # root
|
||||
(PeerCredentials(1, 1000, 1000), frozenset({990}), True), # web user, in group
|
||||
(PeerCredentials(1, 1000, 990), frozenset(), True), # primary group
|
||||
(PeerCredentials(1, 1001, 1001), frozenset({27, 44}), False),
|
||||
(PeerCredentials(1, 65534, 65534), frozenset(), False), # nobody
|
||||
])
|
||||
def test_model(self, cred, groups, allowed):
|
||||
assert peer_allowed(cred, self.OWN, self.GROUP, groups) is allowed
|
||||
|
||||
def test_own_user_without_a_group(self):
|
||||
assert peer_allowed(PeerCredentials(1, 1000, 1000), 1000, None, frozenset())
|
||||
assert not peer_allowed(PeerCredentials(1, 1001, 1001), 1000, None, frozenset({1}))
|
||||
|
||||
def test_group_database_decides_when_proc_is_unreadable(self):
|
||||
seen = []
|
||||
|
||||
def in_group(uid, gid):
|
||||
seen.append((uid, gid))
|
||||
return uid == 1000
|
||||
|
||||
cred = PeerCredentials(1, 1000, 1000)
|
||||
assert peer_allowed(cred, 0, 990, None, in_group=in_group)
|
||||
assert not peer_allowed(PeerCredentials(1, 1001, 1001), 0, 990, None, in_group=in_group)
|
||||
assert seen == [(1000, 990), (1001, 990)]
|
||||
|
||||
@pytest.mark.skipif(os.name != 'posix', reason='POSIX permission bits')
|
||||
def test_socket_group_follows_a_shared_cache_dir(self, tmp_path):
|
||||
shared = tmp_path / 'cache'
|
||||
shared.mkdir()
|
||||
os.chmod(shared, 0o2775)
|
||||
assert srv.resolve_socket_group(str(shared)) == shared.stat().st_gid
|
||||
|
||||
@pytest.mark.skipif(os.name != 'posix', reason='POSIX permission bits')
|
||||
def test_a_private_cache_dir_falls_back_to_the_project_group(self, tmp_path, monkeypatch):
|
||||
private = tmp_path / 'cache'
|
||||
private.mkdir()
|
||||
os.chmod(private, 0o755)
|
||||
from src.common import permission_utils
|
||||
monkeypatch.setattr(permission_utils, 'get_shared_group_gid', lambda: 4242)
|
||||
assert srv.resolve_socket_group(str(private)) == 4242
|
||||
|
||||
def test_mode_is_group_only_with_a_group(self, tmp_path):
|
||||
assert ControlServer(str(tmp_path / 'a'), group=990).socket_mode == 0o660
|
||||
assert ControlServer(str(tmp_path / 'b'), group=None).socket_mode == 0o600
|
||||
|
||||
|
||||
class TestWhereTheServerListens:
|
||||
def test_off_means_no_server(self):
|
||||
assert srv.server_socket_path({c.SOCKET_PATH_ENV: 'off'}) is None
|
||||
assert srv.start_control_server(environ={c.SOCKET_PATH_ENV: 'off'}) is None
|
||||
|
||||
@needs_unix_sockets
|
||||
def test_configured(self):
|
||||
assert srv.server_socket_path({c.SOCKET_PATH_ENV: '/tmp/x.sock'}) == '/tmp/x.sock'
|
||||
|
||||
@needs_unix_sockets
|
||||
def test_unprivileged_dev_run_uses_the_per_user_path(self, monkeypatch):
|
||||
monkeypatch.setattr(os, 'geteuid', lambda: 1000)
|
||||
monkeypatch.setattr(os, 'access', lambda p, m: False)
|
||||
assert srv.server_socket_path({}) == c.dev_socket_path()
|
||||
|
||||
@needs_unix_sockets
|
||||
def test_root_uses_run(self, monkeypatch):
|
||||
monkeypatch.setattr(os, 'geteuid', lambda: 0)
|
||||
assert srv.server_socket_path({}) == c.DEFAULT_SOCKET_PATH
|
||||
|
||||
@pytest.mark.skipif(c.socket_supported(), reason='Windows only')
|
||||
def test_windows_skips_cleanly(self):
|
||||
assert srv.server_socket_path({}) is None
|
||||
assert ControlServer('x.sock').start() is False
|
||||
with pytest.raises(client.ControlError) as e:
|
||||
client.ping(paths=['x.sock'])
|
||||
assert e.value.reason == 'unsupported'
|
||||
|
||||
|
||||
# -- the real socket --------------------------------------------------------------------
|
||||
|
||||
@pytest.fixture
|
||||
def sock_path(tmp_path_factory):
|
||||
# AF_UNIX paths are limited to ~107 bytes; pytest's tmp_path can be longer.
|
||||
import tempfile
|
||||
d = tempfile.mkdtemp(prefix='lmipc-')
|
||||
yield os.path.join(d, 'control.sock')
|
||||
import shutil
|
||||
shutil.rmtree(d, ignore_errors=True)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def live(sock_path, status):
|
||||
servers = []
|
||||
|
||||
def make(**kwargs):
|
||||
kwargs.setdefault('status_provider', lambda: dict(status))
|
||||
s = ControlServer(sock_path, **kwargs)
|
||||
assert s.start()
|
||||
servers.append(s)
|
||||
return s
|
||||
|
||||
yield make
|
||||
for s in servers:
|
||||
s.close()
|
||||
|
||||
|
||||
def _raw(path, timeout=2.0):
|
||||
s = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
|
||||
s.settimeout(timeout)
|
||||
s.connect(path)
|
||||
return s
|
||||
|
||||
|
||||
def _read_line(s):
|
||||
buf = b''
|
||||
while not buf.endswith(b'\n'):
|
||||
try:
|
||||
chunk = s.recv(1) # one byte at a time: never eat the next line
|
||||
except ConnectionResetError:
|
||||
break
|
||||
if not chunk:
|
||||
break
|
||||
buf += chunk
|
||||
return json.loads(buf) if buf.endswith(b'\n') else None
|
||||
|
||||
|
||||
@needs_unix_sockets
|
||||
class TestLiveSocket:
|
||||
def test_client_round_trip(self, live, sock_path):
|
||||
live()
|
||||
assert client.ping(paths=[sock_path]) == {'pong': True}
|
||||
assert client.hello(paths=[sock_path])['version'] == 1
|
||||
assert client.on_demand_status(paths=[sock_path])['current_mode'] == 'clock'
|
||||
|
||||
def test_ack_path(self, live, sock_path):
|
||||
server = live()
|
||||
ack = client.on_demand_start('req-1', 'clock', None, 20, True, paths=[sock_path])
|
||||
assert ack == {'accepted': True, 'request_id': 'req-1', 'queued': 1}
|
||||
ack = client.on_demand_stop('req-2', paths=[sock_path])
|
||||
assert ack['request_id'] == 'req-2'
|
||||
assert [(x.request_id, x.cmd) for x in server.drain()] == [
|
||||
('req-1', Command.ON_DEMAND_START), ('req-2', Command.ON_DEMAND_STOP)]
|
||||
|
||||
def test_socket_file_mode_and_cleanup(self, live, sock_path):
|
||||
server = live(group=os.getgid())
|
||||
st = os.lstat(sock_path)
|
||||
assert stat.S_ISSOCK(st.st_mode)
|
||||
assert stat.S_IMODE(st.st_mode) == 0o660
|
||||
assert st.st_gid == os.getgid()
|
||||
assert not [f for f in os.listdir(os.path.dirname(sock_path)) if f.endswith('.tmp')]
|
||||
server.close()
|
||||
assert not os.path.exists(sock_path)
|
||||
|
||||
def test_without_a_group_only_the_owner_may_connect(self, live, sock_path):
|
||||
live(group=None)
|
||||
assert stat.S_IMODE(os.lstat(sock_path).st_mode) == 0o600
|
||||
|
||||
def test_garbage_then_a_good_request_on_one_connection(self, live, sock_path):
|
||||
live()
|
||||
s = _raw(sock_path)
|
||||
try:
|
||||
s.sendall(b'this is not json\n')
|
||||
assert _read_line(s)['error']['code'] == ErrorCode.BAD_JSON
|
||||
s.sendall(_req(Command.PING, rid='after') + b'\n')
|
||||
resp = _read_line(s)
|
||||
assert resp['ok'] and resp['id'] == 'after'
|
||||
finally:
|
||||
s.close()
|
||||
|
||||
def test_two_requests_in_one_write(self, live, sock_path):
|
||||
live()
|
||||
s = _raw(sock_path)
|
||||
try:
|
||||
s.sendall(_req(Command.PING, rid='a') + b'\n' + _req(Command.PING, rid='b') + b'\n')
|
||||
assert _read_line(s)['id'] == 'a'
|
||||
assert _read_line(s)['id'] == 'b'
|
||||
finally:
|
||||
s.close()
|
||||
|
||||
def test_oversize_is_refused_and_the_server_lives_on(self, live, sock_path):
|
||||
live()
|
||||
s = _raw(sock_path)
|
||||
try:
|
||||
try:
|
||||
# Exactly the limit with no newline: the server has read it
|
||||
# all when it refuses, so its answer is not lost to a reset.
|
||||
s.sendall(b'{"pad":"' + b'x' * (c.MAX_MESSAGE_BYTES - 8))
|
||||
except OSError:
|
||||
pass # the server may hang up before we finish writing
|
||||
resp = _read_line(s)
|
||||
assert resp['error']['code'] == ErrorCode.MESSAGE_TOO_LARGE
|
||||
assert s.recv(10) == b'' # and hung up
|
||||
finally:
|
||||
s.close()
|
||||
assert client.ping(paths=[sock_path]) == {'pong': True}
|
||||
|
||||
def test_a_client_that_hangs_up_mid_message(self, live, sock_path):
|
||||
server = live()
|
||||
s = _raw(sock_path)
|
||||
s.sendall(b'{"v":1,"id":"half","cmd":"on_demand.st')
|
||||
s.close()
|
||||
time.sleep(0.2)
|
||||
assert client.ping(paths=[sock_path]) == {'pong': True}
|
||||
assert server.drain() == []
|
||||
|
||||
def test_a_slow_client_is_dropped_and_blocks_nobody(self, live, sock_path):
|
||||
live(io_timeout=0.2, message_timeout=0.5)
|
||||
slow = _raw(sock_path, timeout=3)
|
||||
try:
|
||||
slow.sendall(b'{"v":1,') # ...and never finishes
|
||||
t0 = time.monotonic()
|
||||
assert client.ping(paths=[sock_path]) == {'pong': True}
|
||||
assert time.monotonic() - t0 < 0.5, 'a slow client held up another'
|
||||
assert slow.recv(100) == b'' # hung up on, not answered
|
||||
finally:
|
||||
slow.close()
|
||||
|
||||
def test_an_idle_connection_is_closed(self, live, sock_path):
|
||||
live(io_timeout=0.1, idle_timeout=0.3)
|
||||
s = _raw(sock_path, timeout=3)
|
||||
try:
|
||||
assert s.recv(100) == b''
|
||||
finally:
|
||||
s.close()
|
||||
|
||||
def test_too_many_clients_are_told_busy(self, live, sock_path):
|
||||
live(max_clients=2, io_timeout=0.2, idle_timeout=5)
|
||||
held = [_raw(sock_path) for _ in range(2)]
|
||||
try:
|
||||
time.sleep(0.1)
|
||||
extra = _raw(sock_path)
|
||||
try:
|
||||
assert _read_line(extra)['error']['code'] == ErrorCode.BUSY
|
||||
finally:
|
||||
extra.close()
|
||||
finally:
|
||||
for s in held:
|
||||
s.close()
|
||||
time.sleep(0.3)
|
||||
assert client.ping(paths=[sock_path]) == {'pong': True}
|
||||
|
||||
def test_many_concurrent_clients(self, live, sock_path):
|
||||
server = live(queue_size=64)
|
||||
errors = []
|
||||
|
||||
def go(n):
|
||||
try:
|
||||
client.on_demand_start(f'r{n}', 'p', None, paths=[sock_path], timeout=3)
|
||||
except client.ControlError as e: # busy is allowed under load
|
||||
if e.reason != ErrorCode.BUSY:
|
||||
errors.append(e)
|
||||
|
||||
threads = [threading.Thread(target=go, args=(n,)) for n in range(20)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join()
|
||||
assert errors == []
|
||||
assert 0 < len(server.drain()) <= 20
|
||||
|
||||
def test_a_stale_socket_is_replaced(self, sock_path, status):
|
||||
dead = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
|
||||
dead.bind(sock_path)
|
||||
dead.close() # file left behind, nothing listening
|
||||
s = ControlServer(sock_path, status_provider=lambda: status)
|
||||
try:
|
||||
assert s.start()
|
||||
assert client.ping(paths=[sock_path]) == {'pong': True}
|
||||
finally:
|
||||
s.close()
|
||||
|
||||
def test_a_live_socket_is_not_stolen(self, live, sock_path, status):
|
||||
live()
|
||||
second = ControlServer(sock_path, status_provider=lambda: status)
|
||||
assert second.start() is False
|
||||
assert client.ping(paths=[sock_path]) == {'pong': True}
|
||||
|
||||
def test_a_regular_file_is_never_removed(self, sock_path):
|
||||
with open(sock_path, 'w') as f:
|
||||
f.write('precious')
|
||||
assert ControlServer(sock_path).start() is False
|
||||
with open(sock_path) as f:
|
||||
assert f.read() == 'precious'
|
||||
|
||||
def test_close_leaves_a_successor_s_socket_alone(self, sock_path, status):
|
||||
first = ControlServer(sock_path, status_provider=lambda: status)
|
||||
assert first.start()
|
||||
first._close_socket() # dead, but still owns the path
|
||||
os.unlink(sock_path)
|
||||
second = ControlServer(sock_path, status_provider=lambda: status)
|
||||
assert second.start()
|
||||
try:
|
||||
first.close() # must not unlink second's file
|
||||
assert client.ping(paths=[sock_path]) == {'pong': True}
|
||||
finally:
|
||||
second.close()
|
||||
|
||||
def test_client_reasons(self, sock_path, tmp_path):
|
||||
with pytest.raises(client.ControlError) as e:
|
||||
client.ping(paths=[sock_path])
|
||||
assert e.value.reason == 'no_socket'
|
||||
dead = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
|
||||
dead.bind(sock_path)
|
||||
try:
|
||||
with pytest.raises(client.ControlError) as e:
|
||||
client.ping(paths=[sock_path])
|
||||
assert e.value.reason == 'refused'
|
||||
finally:
|
||||
dead.close()
|
||||
with pytest.raises(client.ControlError) as e:
|
||||
client.ping(paths=[])
|
||||
assert e.value.reason == 'disabled'
|
||||
with pytest.raises(client.ControlError) as e:
|
||||
client.on_demand_start('x', None, None, paths=[sock_path])
|
||||
assert e.value.reason == 'invalid_request'
|
||||
|
||||
def test_a_display_that_never_answers_times_out(self, sock_path):
|
||||
mute = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
|
||||
mute.bind(sock_path)
|
||||
mute.listen(1) # accepts at the kernel, never replies
|
||||
try:
|
||||
t0 = time.monotonic()
|
||||
with pytest.raises(client.ControlError) as e:
|
||||
client.ping(paths=[sock_path], timeout=0.3)
|
||||
assert e.value.reason == 'timeout'
|
||||
assert time.monotonic() - t0 < 1.0
|
||||
finally:
|
||||
mute.close()
|
||||
|
||||
def test_the_dev_directory_must_be_private(self, monkeypatch, tmp_path):
|
||||
d = tmp_path / 'shared'
|
||||
d.mkdir()
|
||||
target = d / 'control.sock'
|
||||
monkeypatch.setattr(srv, 'dev_socket_path', lambda: str(target))
|
||||
monkeypatch.setattr(os, 'geteuid', lambda: os.getuid() + 1) # "someone else's"
|
||||
assert ControlServer(str(target)).start() is False
|
||||
|
||||
|
||||
@needs_unix_sockets
|
||||
@pytest.mark.skipif(not hasattr(socket, 'SO_PEERCRED'), reason='SO_PEERCRED is Linux-only')
|
||||
class TestPermissions:
|
||||
def test_peer_credentials_are_read(self, live, sock_path):
|
||||
server = live()
|
||||
s = _raw(sock_path)
|
||||
try:
|
||||
s.sendall(_req(Command.PING) + b'\n')
|
||||
assert _read_line(s)['ok']
|
||||
finally:
|
||||
s.close()
|
||||
a, b = socket.socketpair(socket.AF_UNIX)
|
||||
try:
|
||||
cred = srv.peer_credentials(a)
|
||||
assert cred.uid == os.geteuid() and cred.pid == os.getpid()
|
||||
finally:
|
||||
a.close()
|
||||
b.close()
|
||||
assert srv.process_groups(os.getpid()) == frozenset(os.getgroups())
|
||||
assert server.running
|
||||
|
||||
@pytest.mark.skipif(hasattr(os, 'geteuid') and os.geteuid() == 0,
|
||||
reason='root may always connect')
|
||||
def test_a_peer_outside_the_model_is_refused(self, live, sock_path):
|
||||
server = live(group=None)
|
||||
# Pretend the display runs as someone else: this process is then
|
||||
# neither root, the display's user, nor in its (absent) group.
|
||||
server._own_uid = os.geteuid() + 12345
|
||||
s = _raw(sock_path)
|
||||
try:
|
||||
resp = _read_line(s)
|
||||
assert resp['error']['code'] == ErrorCode.FORBIDDEN
|
||||
assert s.recv(10) == b''
|
||||
finally:
|
||||
s.close()
|
||||
with pytest.raises(client.ControlError) as e:
|
||||
client.on_demand_stop('nope', paths=[sock_path])
|
||||
assert e.value.reason == ErrorCode.FORBIDDEN
|
||||
assert server.drain() == []
|
||||
|
||||
@pytest.mark.skipif(not (hasattr(os, 'geteuid') and os.geteuid() == 0),
|
||||
reason='needs root to switch users (run under WSL as root, or on a Pi)')
|
||||
def test_the_kernel_enforces_the_group(self, live, sock_path):
|
||||
"""The real deployment shape: root serves, an unprivileged user connects.
|
||||
|
||||
nobody in the socket's group gets in; nobody outside it gets EACCES
|
||||
from connect() -- the kernel's check, before any byte is read.
|
||||
"""
|
||||
import pwd
|
||||
nobody = pwd.getpwnam('nobody')
|
||||
allowed_gid = nobody.pw_gid
|
||||
live(group=allowed_gid)
|
||||
os.chmod(os.path.dirname(sock_path), 0o755)
|
||||
|
||||
def try_as(gid):
|
||||
r, w = os.pipe()
|
||||
pid = os.fork()
|
||||
if pid == 0: # child: drop to nobody with only `gid`
|
||||
os.close(r)
|
||||
try:
|
||||
os.setgroups([])
|
||||
os.setgid(gid)
|
||||
os.setuid(nobody.pw_uid)
|
||||
result = json.dumps(client.ping(paths=[sock_path]))
|
||||
except client.ControlError as e:
|
||||
result = 'error:' + e.reason
|
||||
except Exception as e: # report anything else to the parent
|
||||
result = 'crash:' + repr(e)
|
||||
os.write(w, result.encode())
|
||||
os._exit(0)
|
||||
os.close(w)
|
||||
out = b''
|
||||
while True:
|
||||
chunk = os.read(r, 4096)
|
||||
if not chunk:
|
||||
break
|
||||
out += chunk
|
||||
os.close(r)
|
||||
os.waitpid(pid, 0)
|
||||
return out.decode()
|
||||
|
||||
assert try_as(allowed_gid) == '{"pong": true}'
|
||||
other_gid = allowed_gid - 1 if allowed_gid > 1 else allowed_gid + 1
|
||||
assert try_as(other_gid) == 'error:refused'
|
||||
# Someone loosens the mode by hand: the kernel lets the outsider
|
||||
# connect, and SO_PEERCRED still turns it away.
|
||||
os.chmod(sock_path, 0o666)
|
||||
assert try_as(other_gid) == 'error:forbidden'
|
||||
assert try_as(allowed_gid) == '{"pong": true}'
|
||||
Reference in New Issue
Block a user