import tempfile import unittest from pathlib import Path from rex.client.config import ClientAction, ClientConfig, ClientConfigManager from rex.server.config import ServerConfig, ServerConfigManager from rex.server.connection_manager import ConnectionManager class ConfigTests(unittest.TestCase): def test_server_config_is_toml_and_persists_scoped_keys(self) -> None: with tempfile.TemporaryDirectory() as directory: path = Path(directory) / "server.toml" manager = ServerConfigManager(path) key = manager.create_key("device", {"connect"}) self.assertIsNotNone(key) self.assertIn("[[keys]]", path.read_text()) self.assertIsNotNone(ServerConfigManager(path).get_key(key.secret, "connect")) self.assertIsNone(ServerConfigManager(path).get_key(key.secret, "execute")) def test_client_url_omits_optional_port(self) -> None: self.assertEqual(ClientConfig().websocket_url(), "wss://127.0.0.1/ws/default-device") self.assertEqual(ClientConfig(endpoint="rex.test", scheme="ws", port=8080).websocket_url(), "ws://rex.test:8080/ws/default-device") def test_client_actions_are_local_argv_only(self) -> None: action = ClientAction(name="lock", argv=["/bin/echo", "locked"]) self.assertEqual(action.argv[0], "/bin/echo") with self.assertRaises(ValueError): ClientAction(name="bad name", argv=["/bin/echo"]) def test_reload_preserves_last_valid_server_configuration(self) -> None: with tempfile.TemporaryDirectory() as directory: path = Path(directory) / "server.toml" manager = ServerConfigManager(path) manager.config = ServerConfig(devices=["desk"]) manager.save() manager.reload() path.write_text("rate_limit_per_minute = 0\n") self.assertFalse(manager.reload()) self.assertEqual(manager.config.devices, ["desk"]) path.write_text("rate_limit_per_minute = 100\ndevices = [\"laptop\"]\n") self.assertTrue(manager.reload()) self.assertEqual(manager.config.devices, ["laptop"]) def test_reload_preserves_last_valid_client_configuration(self) -> None: with tempfile.TemporaryDirectory() as directory: path = Path(directory) / "client.toml" manager = ClientConfigManager(path) manager.config = ClientConfig(device_name="desk") manager.save() manager.reload() path.write_text("scheme = \"https\"\n") self.assertFalse(manager.reload()) self.assertEqual(manager.config.device_name, "desk") class ConnectionManagerTests(unittest.IsolatedAsyncioTestCase): async def test_only_registered_actions_are_sent(self) -> None: class Socket: def __init__(self) -> None: self.messages: list[dict[str, str]] = [] async def accept(self) -> None: pass async def send_json(self, message: dict[str, str]) -> None: self.messages.append(message) socket = Socket() manager = ConnectionManager() await manager.connect("desk", socket, {"lock"}) # type: ignore[arg-type] self.assertFalse(await manager.send("desk", "shell")) self.assertTrue(await manager.send("desk", "lock")) self.assertEqual(socket.messages, [{"type": "action", "name": "lock"}]) async def test_disconnect_all_closes_registered_clients(self) -> None: class Socket: def __init__(self) -> None: self.closed = False async def close(self, **_: object) -> None: self.closed = True socket = Socket() manager = ConnectionManager() await manager.connect("desk", socket, {"lock"}) # type: ignore[arg-type] await manager.disconnect_all() self.assertTrue(socket.closed) self.assertEqual(manager.actions("desk"), [])