mirror of
https://github.com/ChuckBuilds/LEDMatrix.git
synced 2026-10-04 06:15:09 +00:00
feat(ipc): display control socket, stage 1 - on-demand with acks (#706)
The display serves a control socket (/run/ledmatrix/control.sock) carrying versioned JSON commands, one per line, each answered. Stage 1 covers on-demand start, stop and status; commands are queued on the socket thread and applied on the render thread through the mailbox's own handler, and the web interface falls back to the file mailbox when the socket is unavailable. Protocol and security model: docs/IPC_CONTROL_SOCKET.md. 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,206 @@
|
||||
"""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_an_unknown_reason_is_reported_as_other(self, api_v3_client, service):
|
||||
# Only known codes are echoed back; anything else stays server-side.
|
||||
with patch(f"{CLIENT}.on_demand_start",
|
||||
side_effect=control_client.ControlError("/run/secret/path", "x")):
|
||||
data = api_v3_client.post(START_URL, json={"plugin_id": "weather"}).get_json()["data"]
|
||||
assert data["transport"] == "mailbox"
|
||||
assert data["socket_error"] == "other"
|
||||
assert len(_mailbox_writes(service["cache"])) == 1
|
||||
|
||||
def test_every_display_error_code_is_reportable(self):
|
||||
from web_interface.blueprints.api_v3 import display
|
||||
codes = {v for k, v in vars(c.ErrorCode).items() if not k.startswith("_")}
|
||||
assert codes <= set(display._REPORTABLE_SOCKET_REASONS)
|
||||
|
||||
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