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:
Chuck
2026-09-30 21:28:31 -04:00
co-authored by Claude Opus 5.5
parent f4bda50710
commit af8dc3940a
18 changed files with 3129 additions and 42 deletions
+13
View File
@@ -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."""
+192
View File
@@ -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
+275
View File
@@ -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
+229
View File
@@ -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)
+570
View File
@@ -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}'