Initial implementation

This commit is contained in:
2026-08-26 15:13:41 +01:00
commit 7e28074eff
18 changed files with 2306 additions and 0 deletions
+10
View File
@@ -0,0 +1,10 @@
# Python-generated files
__pycache__/
*.py[oc]
build/
dist/
wheels/
*.egg-info
# Virtual environments
.venv
+1
View File
@@ -0,0 +1 @@
3.13
+19
View File
@@ -0,0 +1,19 @@
## Rex
Run the public and management servers:
```console
uv run rex server
```
Open the terminal management interface in another shell:
```console
uv run rex manage
```
The manager connects to `http://127.0.0.1:8001` by default. Use
`uv run rex manage --url http://host:port` to connect elsewhere. The footer
shows all controls: `H/L` switch tabs, `h/j/k/l` move, `y` copies the selected
value, `n` creates a key or device, `a` adds an action, `d` deletes the
selection, `r` refreshes, and `q` quits.
+24
View File
@@ -0,0 +1,24 @@
[project]
name = "rex"
version = "0.1.0"
description = "Add your description here"
readme = "README.md"
authors = [
{ name = "Hector van der Aa", email = "hector@h3cx.dev" }
]
requires-python = ">=3.13"
dependencies = [
"fastapi[standard]>=0.141.1",
"httpx>=0.28.1",
"pydantic>=2.13.4",
"textual>=6.1.0",
"uvicorn>=0.52.4",
"websockets>=17.0.1",
]
[project.scripts]
rex = "rex:main"
[build-system]
requires = ["uv_build>=0.12.5,<0.13.0"]
build-backend = "uv_build"
+29
View File
@@ -0,0 +1,29 @@
import argparse
def main() -> None:
parser = argparse.ArgumentParser(prog="rex")
subparsers = parser.add_subparsers(dest="command", required=True)
subparsers.add_parser("server", help="run the Rex server")
subparsers.add_parser("client", help="run the Rex client")
manage_parser = subparsers.add_parser("manage", help="manage a Rex server")
manage_parser.add_argument(
"--url",
default="http://127.0.0.1:8011",
help="management API URL (default: %(default)s)",
)
args = parser.parse_args()
if args.command == "server":
from rex.server import main as server_main
server_main()
elif args.command == "client":
from rex.client import main as client_main
client_main()
elif args.command == "manage":
from rex.management import main as management_main
management_main(args.url)
+109
View File
@@ -0,0 +1,109 @@
import asyncio
import os
import signal
import websockets
from rex.client.action_manager import ClientActionManager
from rex.client.config_manager import ClientConfigManager
from rex.server.device_manager import Action
client_config = ClientConfigManager("~/.config/rex/client.json")
action_manager = ClientActionManager("~/.config/rex/actions.json")
async def run_command(command: str) -> None:
process = await asyncio.create_subprocess_shell(
command,
start_new_session=os.name == "posix",
)
try:
await process.wait()
except asyncio.CancelledError:
if process.returncode is None:
if os.name == "posix":
try:
os.killpg(process.pid, signal.SIGTERM)
except ProcessLookupError:
pass
else:
process.terminate()
try:
await asyncio.wait_for(process.wait(), timeout=5)
except TimeoutError:
if os.name == "posix":
try:
os.killpg(process.pid, signal.SIGKILL)
except ProcessLookupError:
pass
else:
process.kill()
await process.wait()
raise
async def receive_actions(websocket: websockets.ClientConnection) -> None:
command_tasks: set[asyncio.Task[None]] = set()
try:
async for message in websocket:
action = Action.model_validate_json(message)
client_action = action_manager.get_action(action.name)
if client_action is not None:
task = asyncio.create_task(run_command(client_action.command))
command_tasks.add(task)
task.add_done_callback(command_tasks.discard)
except websockets.ConnectionClosed:
pass
finally:
for task in command_tasks:
task.cancel()
await asyncio.gather(*command_tasks, return_exceptions=True)
async def listen() -> None:
uri = f"ws://{client_config.get_endpoint()}:{client_config.get_port()}/ws/{client_config.get_name()}"
print(uri)
stop = asyncio.Event()
loop = asyncio.get_running_loop()
handles_sigterm = False
try:
loop.add_signal_handler(signal.SIGTERM, stop.set)
handles_sigterm = True
except (NotImplementedError, RuntimeError, ValueError):
pass
try:
async with websockets.connect(
uri, additional_headers={"X-API-Key": client_config.get_key()}
) as websocket:
print("Connected")
receiver = asyncio.create_task(receive_actions(websocket))
shutdown = asyncio.create_task(stop.wait())
try:
done, _ = await asyncio.wait(
(receiver, shutdown), return_when=asyncio.FIRST_COMPLETED
)
if receiver in done:
await receiver
finally:
receiver.cancel()
shutdown.cancel()
await asyncio.gather(receiver, shutdown, return_exceptions=True)
finally:
if handles_sigterm:
loop.remove_signal_handler(signal.SIGTERM)
def main() -> None:
print(f"Rex client name: {client_config.get_name()}")
print(f"API Key: {client_config.get_key()}")
asyncio.run(listen())
+56
View File
@@ -0,0 +1,56 @@
from pathlib import Path
from pydantic import BaseModel
class ClientAction(BaseModel):
name: str
command: str
class ClientActionStore(BaseModel):
actions: list[ClientAction]
class ClientActionManager:
def __init__(self, file_path: str | Path) -> None:
self.file_path = Path(file_path).expanduser()
if not self.file_path.exists():
self.action_store = ClientActionStore(actions=[])
self.file_path.parent.mkdir(parents=True, exist_ok=True)
self.file_path.write_text(self.action_store.model_dump_json(indent=4))
self.action_store = ClientActionStore.model_validate_json(
self.file_path.read_text()
)
def create_action(self, action_name: str, command: str) -> ClientAction | None:
if self.get_action(action_name) is not None:
return None
new_action = ClientAction(name=action_name, command=command)
self.action_store.actions.append(new_action)
self._save()
return new_action
def _save(self) -> None:
self.file_path.parent.mkdir(parents=True, exist_ok=True)
self.file_path.write_text(self.action_store.model_dump_json(indent=4))
def get_action(self, action_name: str) -> ClientAction | None:
for action in self.action_store.actions:
if action_name == action.name:
return action
return None
def get_all_actions(self) -> list[ClientAction]:
return self.action_store.actions
def delete_action(self, action_name: str) -> bool:
action = self.get_action(action_name)
if action is None:
return False
self.action_store.actions.remove(action)
self._save()
return True
+38
View File
@@ -0,0 +1,38 @@
from pathlib import Path
from pydantic import BaseModel
class ClientConfig(BaseModel):
device_name: str
api_key: str
endpoint: str
port: int
class ClientConfigManager:
def __init__(self, file_path: str | Path) -> None:
self.file_path = Path(file_path).expanduser()
if not self.file_path.exists():
self.client_config = ClientConfig(device_name="default_device")
self.file_path.parent.mkdir(parents=True, exist_ok=True)
self.file_path.write_text(self.client_config.model_dump_json(indent=4))
self.client_config = ClientConfig.model_validate_json(
self.file_path.read_text()
)
def _save(self) -> None:
self.file_path.parent.mkdir(parents=True, exist_ok=True)
self.file_path.write_text(self.client_config.model_dump_json(indent=4))
def get_name(self) -> str:
return self.client_config.device_name
def get_key(self) -> str:
return self.client_config.api_key
def get_endpoint(self) -> str:
return self.client_config.endpoint
def get_port(self) -> int:
return self.client_config.port
+5
View File
@@ -0,0 +1,5 @@
from rex.management.app import RexManager
def main(url: str) -> None:
RexManager(url).run()
+264
View File
@@ -0,0 +1,264 @@
from collections.abc import Callable, Coroutine
from typing import Any, ClassVar
from textual.app import App, ComposeResult
from textual.containers import Container
from textual.screen import ModalScreen
from textual.widgets import (
DataTable,
Footer,
Header,
Input,
Static,
TabbedContent,
TabPane,
Tree,
)
from rex.management.client import ManagementClient, ManagementError
class NamePrompt(ModalScreen[str | None]):
BINDINGS: ClassVar = [("escape", "cancel", "Cancel")]
CSS = """
NamePrompt { align: center middle; background: $background 60%; }
#prompt-dialog {
width: 60;
max-width: 90%;
height: auto;
padding: 1 2;
border: round $primary;
background: $surface;
}
#prompt-label { margin-bottom: 1; }
"""
def __init__(self, prompt: str) -> None:
super().__init__()
self.prompt = prompt
def compose(self) -> ComposeResult:
with Container(id="prompt-dialog"):
yield Static(self.prompt, id="prompt-label")
yield Input(placeholder="Type a name and press Enter", id="prompt-input")
def on_mount(self) -> None:
self.query_one(Input).focus()
def on_input_submitted(self, event: Input.Submitted) -> None:
name = event.value.strip()
if name:
self.dismiss(name)
def action_cancel(self) -> None:
self.dismiss(None)
class RexManager(App[None]):
TITLE = "Rex Manager"
SUB_TITLE = "Management API"
CSS = """
Screen { background: $surface; }
TabbedContent { margin: 1 2; }
TabPane { padding: 1 2; }
DataTable, Tree { height: 1fr; border: round $primary; }
#status { height: auto; padding: 0 3 1 3; color: $text-muted; }
"""
BINDINGS: ClassVar = [
("shift+h", "previous_tab", "Previous tab"),
("shift+l", "next_tab", "Next tab"),
("h", "move_left", "Move left"),
("l", "move_right", "Move right"),
("j", "move_down", "Move down"),
("k", "move_up", "Move up"),
("y", "yank", "Yank"),
("n", "new", "New"),
("a", "add_action", "Add action"),
("d", "delete", "Delete"),
("r", "refresh", "Refresh"),
("q", "quit", "Quit"),
]
def __init__(self, url: str) -> None:
super().__init__()
self.url = url.rstrip("/")
self.client = ManagementClient(self.url)
def compose(self) -> ComposeResult:
yield Header()
with TabbedContent():
with TabPane("API keys", id="keys-tab"):
yield DataTable(id="keys")
with TabPane("Devices", id="devices-tab"):
yield Tree("Devices", id="devices")
yield Static(f"Connecting to {self.url}", id="status")
yield Footer()
async def on_mount(self) -> None:
keys = self.query_one("#keys", DataTable)
keys.add_columns("Name", "Secret")
await self.action_refresh()
keys.focus()
async def on_unmount(self) -> None:
await self.client.close()
async def action_refresh(self) -> None:
try:
keys, devices = await self.client.keys(), await self.client.devices()
except ManagementError as error:
self._set_status(str(error), error=True)
return
key_table = self.query_one("#keys", DataTable)
key_table.clear()
for key in keys:
key_table.add_row(key["name"], key["secret"], key=key["name"])
tree = self.query_one("#devices", Tree)
tree.clear()
for device in devices:
device_name = device["name"]
device_node = tree.root.add(device_name, data=("device", device_name))
for action in device["actions"]:
device_node.add_leaf(
action["name"], data=("action", device_name, action["name"])
)
device_node.expand()
tree.root.expand()
self._set_status(f"Connected to {self.url}")
def action_previous_tab(self) -> None:
self._select_tab("keys-tab")
def action_next_tab(self) -> None:
self._select_tab("devices-tab")
def action_move_down(self) -> None:
self._active_widget().action_cursor_down()
def action_move_up(self) -> None:
self._active_widget().action_cursor_up()
def action_move_left(self) -> None:
if self.query_one(TabbedContent).active == "keys-tab":
self.query_one("#keys", DataTable).action_cursor_left()
def action_move_right(self) -> None:
if self.query_one(TabbedContent).active == "keys-tab":
self.query_one("#keys", DataTable).action_cursor_right()
def action_yank(self) -> None:
if self.query_one(TabbedContent).active == "keys-tab":
table = self.query_one("#keys", DataTable)
if table.row_count == 0:
self._set_status("Select a key first", error=True)
return
value = str(table.get_cell_at(table.cursor_coordinate))
else:
selection = self._tree_selection()
if selection is None:
self._set_status("Select a device or action first", error=True)
return
value = selection[-1]
self.copy_to_clipboard(value)
self._set_status("Copied selected value")
def action_new(self) -> None:
tabs = self.query_one(TabbedContent)
if tabs.active == "keys-tab":
self._prompt("New API key", self._create_key)
else:
self._prompt("New device", self._create_device)
def action_add_action(self) -> None:
if self.query_one(TabbedContent).active != "devices-tab":
self._set_status("Switch to Devices to add an action", error=True)
return
selection = self._tree_selection()
if selection is None:
self._set_status("Select a device first", error=True)
return
device = selection[1]
self._prompt(
f"New action for {device}",
lambda name: self._create_action(device, name),
)
def action_delete(self) -> None:
if self.query_one(TabbedContent).active == "keys-tab":
request = self._delete_key()
else:
request = self._delete_node()
if request is not None:
self.run_worker(request)
def _create_key(self, name: str) -> None:
self.run_worker(self._change(self.client.create_key(name), "Key created"))
def _create_device(self, name: str) -> None:
self.run_worker(self._change(self.client.create_device(name), "Device created"))
def _create_action(self, device: str, name: str) -> None:
self.run_worker(
self._change(self.client.create_action(device, name), "Action created")
)
def _delete_key(self) -> Coroutine[Any, Any, None] | None:
table = self.query_one("#keys", DataTable)
if table.row_count == 0:
self._set_status("Select a key first", error=True)
return None
name = str(table.get_row_at(table.cursor_row)[0])
return self._change(self.client.delete_key(name), "Key deleted")
def _delete_node(self) -> Coroutine[Any, Any, None] | None:
selection = self._tree_selection()
if selection is None:
self._set_status("Select a device or action first", error=True)
return None
if selection[0] == "device":
request = self.client.delete_device(selection[1])
message = "Device deleted"
else:
request = self.client.delete_action(selection[1], selection[2])
message = "Action deleted"
return self._change(request, message)
async def _change(self, request: Any, success: str) -> None:
try:
await request
except ManagementError as error:
self._set_status(str(error), error=True)
return
await self.action_refresh()
self._set_status(success)
def _prompt(self, label: str, callback: Callable[[str], None]) -> None:
def submitted(name: str | None) -> None:
if name is not None:
callback(name)
self.push_screen(NamePrompt(label), submitted)
def _select_tab(self, tab: str) -> None:
tabs = self.query_one(TabbedContent)
tabs.active = tab
self._active_widget().focus()
def _active_widget(self) -> DataTable[Any] | Tree[Any]:
if self.query_one(TabbedContent).active == "keys-tab":
return self.query_one("#keys", DataTable)
return self.query_one("#devices", Tree)
def _tree_selection(self) -> tuple[str, ...] | None:
return self.query_one("#devices", Tree).cursor_node.data
def _set_status(self, message: str, *, error: bool = False) -> None:
status = self.query_one("#status", Static)
status.update(message)
status.styles.color = "red" if error else "green"
+65
View File
@@ -0,0 +1,65 @@
from typing import Any
from urllib.parse import quote
import httpx
class ManagementError(Exception):
pass
class ManagementClient:
def __init__(
self, base_url: str, *, transport: httpx.AsyncBaseTransport | None = None
) -> None:
self._client = httpx.AsyncClient(
base_url=base_url.rstrip("/"), timeout=5, transport=transport
)
async def close(self) -> None:
await self._client.aclose()
async def keys(self) -> list[dict[str, str]]:
return (await self._request("GET", "/keys"))["keys"]
async def create_key(self, name: str) -> None:
await self._request("POST", f"/key/{quote(name, safe='')}")
async def delete_key(self, name: str) -> None:
await self._request("DELETE", f"/key/{quote(name, safe='')}")
async def devices(self) -> list[dict[str, Any]]:
return (await self._request("GET", "/devices"))["devices"]
async def create_device(self, name: str) -> None:
await self._request("POST", f"/device/{quote(name, safe='')}")
async def delete_device(self, name: str) -> None:
await self._request("DELETE", f"/device/{quote(name, safe='')}")
async def create_action(self, device: str, name: str) -> None:
await self._request(
"POST", f"/device/{quote(device, safe='')}/action/{quote(name, safe='')}"
)
async def delete_action(self, device: str, name: str) -> None:
await self._request(
"DELETE", f"/device/{quote(device, safe='')}/action/{quote(name, safe='')}"
)
async def _request(self, method: str, path: str) -> Any:
try:
response = await self._client.request(method, path)
response.raise_for_status()
except httpx.HTTPStatusError as error:
try:
message = error.response.json().get("detail", str(error))
except ValueError:
message = str(error)
raise ManagementError(message) from error
except httpx.HTTPError as error:
raise ManagementError(f"Could not reach the server: {error}") from error
if not response.content:
return None
return response.json()
+214
View File
@@ -0,0 +1,214 @@
import asyncio
import uvicorn
from fastapi import (
Depends,
FastAPI,
HTTPException,
Header,
WebSocket,
WebSocketDisconnect,
WebSocketException,
status,
)
from rex.server.connection_manager import ConnectionManager
from rex.server.device_manager import DeviceManager
from rex.server.key_manager import KeyManager
public_app = FastAPI()
management_app = FastAPI()
key_manager = KeyManager("~/.config/rex/keys.json")
device_manager = DeviceManager("~/.config/rex/devices.json")
connection_manager = ConnectionManager()
async def verify_api_key(x_api_key: str = Header()):
if not key_manager.compare_key(x_api_key):
raise HTTPException(
status_code=401,
detail="Invalid API key",
)
async def verify_api_key_ws(websocket: WebSocket) -> None:
api_key = websocket.headers.get("x-api-key")
if api_key is None or not key_manager.compare_key(api_key):
raise WebSocketException(
code=status.WS_1008_POLICY_VIOLATION,
reason="Invalid API key",
)
@public_app.websocket("/ws/{device_name}")
async def websocket_endpoint(
websocket: WebSocket, device_name: str, _: None = Depends(verify_api_key_ws)
):
if device_manager.get_device(device_name) is None:
raise WebSocketException(
code=status.WS_1008_POLICY_VIOLATION,
reason="Device not found",
)
await connection_manager.connect(device_name, websocket)
try:
while True:
await websocket.receive_json()
except WebSocketDisconnect:
connection_manager.disconnect(device_name)
@public_app.get("/device/{device_name}/action/{action_name}")
async def send_action(
device_name: str, action_name: str, _: None = Depends(verify_api_key)
):
action = device_manager.get_action(device_name, action_name)
if action is None:
raise HTTPException(
status_code=404,
detail="Action not found for this device",
)
await connection_manager.send(device_name, action)
@management_app.get("/tryauth")
async def tryauth(_: None = Depends(verify_api_key)):
return
@management_app.get("/keys")
async def get_keys():
return {"keys": key_manager.get_all_keys()}
@management_app.get("/key/{key_name}")
async def get_key_by_name(key_name: str):
key = key_manager.get_key(key_name)
if key is None:
raise HTTPException(
status_code=404,
detail="Key not found",
)
return {"key": key}
@management_app.post("/key/{key_name}")
async def create_key(key_name: str):
key = key_manager.create_key(key_name)
if key is None:
raise HTTPException(status_code=409, detail="Key already exists with this name")
return {"key": key}
@management_app.delete("/key/{key_name}")
async def delete_key(key_name: str):
val = key_manager.delete_key(key_name)
if not val:
raise HTTPException(status_code=404, detail="No key found with this name")
return
@management_app.get("/devices")
async def get_devices():
return {"devices": device_manager.get_all_devices()}
@management_app.get("/device/{device_name}")
async def get_device_by_name(device_name: str):
device = device_manager.get_device(device_name)
if device is None:
raise HTTPException(
status_code=404,
detail="Device not found",
)
return {"device": device}
@management_app.post("/device/{device_name}")
async def create_device(device_name: str):
device = device_manager.create_device(device_name)
if device is None:
raise HTTPException(
status_code=409, detail="Device already exists with this name"
)
return {"device": device}
@management_app.delete("/device/{device_name}")
async def delete_device(device_name: str):
val = device_manager.delete_device(device_name)
if not val:
raise HTTPException(status_code=404, detail="No device found with this name")
return
@management_app.get("/device/{device_name}/actions")
async def get_all_actions(device_name: str):
return {"devices": device_manager.get_all_actions(device_name)}
@management_app.get("/device/{device_name}/action/{action_name}")
async def get_action_by_name(device_name: str, action_name: str):
action = device_manager.get_action(device_name, action_name)
if action is None:
raise HTTPException(
status_code=404,
detail="Action not found on this device",
)
return {"device": action}
@management_app.post("/device/{device_name}/action/{action_name}")
async def create_action(device_name: str, action_name: str):
device = device_manager.create_action(device_name, action_name)
if device is None:
raise HTTPException(
status_code=409,
detail="Action already exists with this name on this device",
)
return {"device": device}
@management_app.delete("/device/{device_name}/action/{action_name}")
async def delete_action(device_name: str, action_name: str):
val = device_manager.delete_action(device_name, action_name)
if not val:
raise HTTPException(
status_code=404, detail="No action found with this name on this device"
)
return
async def run_servers() -> None:
public_server = uvicorn.Server(
uvicorn.Config(
public_app,
host="127.0.0.1",
port=8010,
)
)
management_server = uvicorn.Server(
uvicorn.Config(
management_app,
host="127.0.0.1",
port=8011,
)
)
await asyncio.gather(
public_server.serve(),
management_server.serve(),
)
def main() -> None:
try:
asyncio.run(run_servers())
except KeyboardInterrupt:
pass
+27
View File
@@ -0,0 +1,27 @@
from fastapi import WebSocket
from pydantic import BaseModel
class ConnectionManager:
def __init__(self) -> None:
self.connections: dict[str, WebSocket] = {}
async def connect(self, client_name: str, websocket: WebSocket) -> None:
await websocket.accept()
self.connections[client_name] = websocket
def disconnect(self, client_name: str) -> None:
self.connections.pop(client_name, None)
async def send(self, client_name: str, message: BaseModel) -> bool:
websocket = self.connections.get(client_name)
if websocket is None:
return False
try:
await websocket.send_json(message.model_dump())
return True
except Exception:
self.disconnect(client_name)
return False
+118
View File
@@ -0,0 +1,118 @@
from pathlib import Path
from pydantic import BaseModel
class Action(BaseModel):
name: str
class Device(BaseModel):
name: str
actions: list[Action]
def add_action(self, action_name: str) -> Action | None:
for action in self.actions:
if action.name == action_name:
return None
new_action = Action(name=action_name)
self.actions.append(new_action)
return new_action
def get_action(self, action_name: str) -> Action | None:
for action in self.actions:
if action.name == action_name:
return action
return None
def get_all_actions(self) -> list[Action]:
return self.actions
def delete_action(self, action_name: str) -> bool:
for action in self.actions:
if action.name == action_name:
self.actions.remove(action)
return True
return False
class DeviceStore(BaseModel):
devices: list[Device]
class DeviceManager:
def __init__(self, file_path: str | Path) -> None:
self.file_path = Path(file_path).expanduser()
if not self.file_path.exists():
self.device_store = DeviceStore(devices=[])
self.file_path.parent.mkdir(parents=True, exist_ok=True)
self.file_path.write_text(self.device_store.model_dump_json(indent=4))
self.device_store = DeviceStore.model_validate_json(self.file_path.read_text())
def create_device(self, device_name: str) -> Device | None:
if self.get_device(device_name) is not None:
return None
new_device = Device(name=device_name, actions=[])
self.device_store.devices.append(new_device)
self._save()
return new_device
def _save(self) -> None:
self.file_path.parent.mkdir(parents=True, exist_ok=True)
self.file_path.write_text(self.device_store.model_dump_json(indent=4))
def get_device(self, device_name: str) -> Device | None:
for device in self.device_store.devices:
if device_name == device.name:
return device
return None
def get_all_devices(self) -> list[Device]:
return self.device_store.devices
def delete_device(self, device_name: str) -> bool:
device = self.get_device(device_name)
if device is None:
return False
self.device_store.devices.remove(device)
self._save()
return True
def create_action(self, device_name: str, action_name: str) -> Action | None:
for device in self.device_store.devices:
if device.name == device_name:
action = device.add_action(action_name)
if action is not None:
self._save()
return action
return None
def get_action(self, device_name: str, action_name: str) -> Action | None:
for device in self.device_store.devices:
if device.name == device_name:
return device.get_action(action_name)
return None
def get_all_actions(self, device_name: str) -> list[Action]:
for device in self.device_store.devices:
if device.name == device_name:
return device.get_all_actions()
return []
def delete_action(self, device_name: str, action_name: str) -> bool:
for device in self.device_store.devices:
if device.name == device_name:
deleted = device.delete_action(action_name)
if deleted:
self._save()
return deleted
return False
+61
View File
@@ -0,0 +1,61 @@
from pathlib import Path
import secrets
from pydantic import BaseModel
class Key(BaseModel):
secret: str
name: str
class KeyStore(BaseModel):
keys: list[Key]
class KeyManager:
def __init__(self, file_path: str | Path) -> None:
self.file_path = Path(file_path).expanduser()
if not self.file_path.exists():
self.key_store = KeyStore(keys=[])
self.file_path.parent.mkdir(parents=True, exist_ok=True)
self.file_path.write_text(self.key_store.model_dump_json(indent=4))
self.key_store = KeyStore.model_validate_json(self.file_path.read_text())
def create_key(self, key_name: str) -> Key | None:
if self.get_key(key_name) is not None:
return None
new_key = Key(secret=secrets.token_urlsafe(32), name=key_name)
self.key_store.keys.append(new_key)
self._save()
return new_key
def _save(self) -> None:
self.file_path.parent.mkdir(parents=True, exist_ok=True)
self.file_path.write_text(self.key_store.model_dump_json(indent=4))
def get_key(self, key_name: str) -> Key | None:
for key in self.key_store.keys:
if key_name == key.name:
return key
return None
def get_all_keys(self) -> list[Key]:
return self.key_store.keys
def delete_key(self, key_name: str) -> bool:
key = self.get_key(key_name)
if key is None:
return False
self.key_store.keys.remove(key)
self._save()
return True
def compare_key(self, key: str) -> bool:
for k in self.key_store.keys:
if secrets.compare_digest(key, k.secret):
return True
return False
+89
View File
@@ -0,0 +1,89 @@
import asyncio
import signal
import unittest
from unittest.mock import AsyncMock, patch
from rex.client import action_manager, receive_actions, run_command
from rex.client.action_manager import ClientAction
class FakeWebSocket:
def __init__(self, messages: list[str]) -> None:
self.messages = iter(messages)
def __aiter__(self):
return self
async def __anext__(self) -> str:
await asyncio.sleep(0)
try:
return next(self.messages)
except StopIteration:
raise StopAsyncIteration
class ClientTests(unittest.IsolatedAsyncioTestCase):
async def test_run_command_waits_for_process(self) -> None:
process = AsyncMock()
process.wait.return_value = 0
with patch("rex.client.asyncio.create_subprocess_shell", return_value=process):
await run_command("test command")
process.wait.assert_awaited_once()
async def test_cancelling_command_terminates_process_group(self) -> None:
process = AsyncMock()
process.pid = 123
process.returncode = None
wait_started = asyncio.Event()
wait_calls = 0
async def wait() -> int:
nonlocal wait_calls
wait_calls += 1
if wait_calls == 1:
wait_started.set()
await asyncio.Event().wait()
return 0
process.wait.side_effect = wait
with (
patch("rex.client.asyncio.create_subprocess_shell", return_value=process),
patch("rex.client.os.name", "posix"),
patch("rex.client.os.killpg") as killpg,
):
task = asyncio.create_task(run_command("test command"))
await wait_started.wait()
task.cancel()
await asyncio.gather(task, return_exceptions=True)
killpg.assert_called_once_with(123, signal.SIGTERM)
self.assertEqual(process.wait.await_count, 2)
async def test_disconnect_cancels_running_commands(self) -> None:
action = ClientAction(name="wake", command="test command")
command_started = asyncio.Event()
command_cancelled = asyncio.Event()
async def running_command(_: str) -> None:
command_started.set()
try:
await asyncio.Event().wait()
finally:
command_cancelled.set()
websocket = FakeWebSocket(['{"name":"wake"}'])
with (
patch.object(action_manager, "get_action", return_value=action),
patch("rex.client.run_command", side_effect=running_command),
):
await receive_actions(websocket) # type: ignore[arg-type]
self.assertTrue(command_started.is_set())
self.assertTrue(command_cancelled.is_set())
if __name__ == "__main__":
unittest.main()
+125
View File
@@ -0,0 +1,125 @@
import tempfile
import unittest
from pathlib import Path
from unittest.mock import patch
import httpx
from textual.widgets import DataTable, TabbedContent, Tree
from rex.management.app import RexManager
from rex.management.client import ManagementClient, ManagementError
from rex.server.device_manager import DeviceManager
class ManagementClientTests(unittest.IsolatedAsyncioTestCase):
async def test_reads_keys_and_encodes_names(self) -> None:
requests: list[httpx.Request] = []
def handler(request: httpx.Request) -> httpx.Response:
requests.append(request)
if request.method == "GET":
return httpx.Response(
200, json={"keys": [{"name": "home", "secret": "abc"}]}
)
return httpx.Response(200, json={"key": {}})
client = ManagementClient(
"http://rex.test", transport=httpx.MockTransport(handler)
)
self.addAsyncCleanup(client.close)
self.assertEqual((await client.keys())[0]["name"], "home")
await client.create_key("living room")
self.assertEqual(requests[1].url.raw_path, b"/key/living%20room")
async def test_reports_api_error_detail(self) -> None:
def handler(_: httpx.Request) -> httpx.Response:
return httpx.Response(409, json={"detail": "already exists"})
client = ManagementClient(
"http://rex.test", transport=httpx.MockTransport(handler)
)
self.addAsyncCleanup(client.close)
with self.assertRaisesRegex(ManagementError, "already exists"):
await client.create_device("desk")
class DeviceManagerTests(unittest.TestCase):
def test_action_changes_are_persisted(self) -> None:
with tempfile.TemporaryDirectory() as directory:
path = Path(directory) / "devices.json"
manager = DeviceManager(path)
manager.create_device("desk")
manager.create_action("desk", "wake")
reloaded = DeviceManager(path)
self.assertIsNotNone(reloaded.get_action("desk", "wake"))
reloaded.delete_action("desk", "wake")
self.assertIsNone(DeviceManager(path).get_action("desk", "wake"))
class FakeManagementClient:
def __init__(self) -> None:
self.created_keys: list[str] = []
async def close(self) -> None:
pass
async def keys(self) -> list[dict[str, str]]:
return [{"name": "home", "secret": "abc"}]
async def devices(self) -> list[dict[str, object]]:
return [{"name": "desk", "actions": [{"name": "wake"}]}]
async def create_key(self, name: str) -> None:
self.created_keys.append(name)
class ManagerAppTests(unittest.IsolatedAsyncioTestCase):
async def test_refresh_populates_keys_and_devices(self) -> None:
app = RexManager("http://rex.test")
app.client = FakeManagementClient() # type: ignore[assignment]
async with app.run_test() as pilot:
await pilot.pause()
keys = app.query_one("#keys", DataTable)
devices = app.query_one("#devices", Tree)
self.assertEqual(keys.get_row_at(0), ["home", "abc"])
self.assertEqual(devices.root.children[0].label.plain, "desk")
self.assertEqual(devices.root.children[0].children[0].label.plain, "wake")
async def test_vim_navigation_and_create_prompt(self) -> None:
app = RexManager("http://rex.test")
client = FakeManagementClient()
app.client = client # type: ignore[assignment]
async with app.run_test() as pilot:
await pilot.press("n")
await pilot.press(*"office", "enter")
await pilot.pause()
self.assertEqual(client.created_keys, ["office"])
await pilot.press("shift+l", "j")
self.assertEqual(app.query_one(TabbedContent).active, "devices-tab")
self.assertEqual(
app.query_one("#devices", Tree).cursor_node.data,
("device", "desk"),
)
async def test_moves_between_key_columns_and_yanks_value(self) -> None:
app = RexManager("http://rex.test")
app.client = FakeManagementClient() # type: ignore[assignment]
async with app.run_test() as pilot:
with patch.object(app, "copy_to_clipboard") as copy:
await pilot.press("l", "y")
self.assertEqual(app.query_one("#keys", DataTable).cursor_column, 1)
copy.assert_called_once_with("abc")
if __name__ == "__main__":
unittest.main()
Generated
+1052
View File
File diff suppressed because it is too large Load Diff