from __future__ import annotations import asyncio import json import os import signal import socket import struct from pathlib import Path from .config import socket_path, pid_path, setting from .config import db_path from .store import Store from .registry import RegistryError, validate_targets from .control import ControlModeClient PROTOCOL_VERSION = "0.1" class Service: def __init__(self, path: Path | None = None, store: Store | None = None, poll_interval: float | None = None): self.path = path or socket_path() self.store = store or Store(db_path()) self.pidfile = pid_path() self.server: asyncio.AbstractServer | None = None configured_interval = setting("delivery_poll_interval", "0.5") try: self.poll_interval = poll_interval if poll_interval is not None else max(0.05, float(configured_interval)) except ValueError: self.poll_interval = poll_interval if poll_interval is not None else 0.5 async def run(self) -> None: self.path.parent.mkdir(parents=True, exist_ok=True) self.pidfile.write_text(str(os.getpid()), encoding="utf-8") os.chmod(self.pidfile, 0o600) if self.path.exists(): self.path.unlink() self.server = await asyncio.start_unix_server(self.handle, path=str(self.path)) os.chmod(self.path, 0o600) stopped = asyncio.Event() loop = asyncio.get_running_loop() for sig in (signal.SIGTERM, signal.SIGINT): try: loop.add_signal_handler(sig, stopped.set) except (NotImplementedError, RuntimeError): pass async with self.server: delivery = asyncio.create_task(self._delivery_loop(stopped)) try: await stopped.wait() finally: delivery.cancel() await asyncio.gather(delivery, return_exceptions=True) self.server.close() await self.server.wait_closed() self.store.disconnect_all() self.path.unlink(missing_ok=True) self.pidfile.unlink(missing_ok=True) self.store.close() async def _delivery_loop(self, stopped: asyncio.Event) -> None: """Continuously inject pending messages into registered live endpoints.""" while not stopped.is_set(): try: self._deliver_once() except Exception: # A disappearing tmux session must not take down the queue. pass try: await asyncio.wait_for(stopped.wait(), timeout=self.poll_interval) except asyncio.TimeoutError: continue def _deliver_once(self) -> None: for endpoint in self.store.endpoints(): repos = json.loads(endpoint["repos"]) control = ControlModeClient(endpoint["session"]) if not control.session_exists(expected_pid=int(endpoint["pid"])): self.store.disconnect_endpoint(endpoint["endpoint_id"]) continue pending = [row for row in self.store.list(state="pending") if row["endpoint_id"] in (None, endpoint["endpoint_id"]) and row["target_repo"] in repos] if not pending: continue try: control.start() for row in pending: lease = self.store.claim(row["message_id"], endpoint["endpoint_id"]) if lease is None: continue try: control.inject(f'{endpoint["session"]}:{row["target_repo"]}', f'#{row["sender_repo"]}: {row["body"]}') except Exception: continue self.store.release(row["message_id"], lease, "injected") finally: control.close() async def handle(self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: if not self._peer_allowed(writer): writer.write(b'{"ok":false,"error":"unauthorized peer"}\n') await writer.drain(); writer.close(); await writer.wait_closed(); return line = await reader.readline() try: request = json.loads(line or b"{}") op = request.get("op") protocol = request.get("protocol", PROTOCOL_VERSION) if protocol.split(".")[0] != PROTOCOL_VERSION.split(".")[0]: response = {"ok": False, "error": f"incompatible protocol: {protocol}", "protocol": PROTOCOL_VERSION} elif op == "ping": response = {"ok": True, "protocol": PROTOCOL_VERSION, "capabilities": ["register", "send", "history"]} elif op == "register": required = ("endpoint_id", "pid", "session", "repos") if any(key not in request for key in required): response = {"ok": False, "error": "register requires endpoint_id, pid, session, repos"} else: try: validate_targets(list(request["repos"])) endpoint_id = request.get("instance_id", request["endpoint_id"]) self.store.register_endpoint(endpoint_id, int(request["pid"]), request["session"], list(request["repos"])) response = {"ok": True, "endpoint_id": request["endpoint_id"], "instance_id": endpoint_id, "protocol": PROTOCOL_VERSION} except (RegistryError, ValueError) as exc: response = {"ok": False, "error": str(exc)} elif op == "send": required = ("sender_repo", "target_repo", "body") if any(key not in request for key in required): response = {"ok": False, "error": "send requires sender_repo, target_repo, body"} else: try: validate_targets([request["target_repo"]]) endpoint_id = request.get("endpoint_id") if endpoint_id: endpoint = self.store.endpoint(endpoint_id) if endpoint is None: raise RegistryError("endpoint is not registered") import json as _json if request["target_repo"] not in _json.loads(endpoint["repos"]): raise RegistryError("target is not attached to endpoint") endpoint_id = endpoint["endpoint_id"] message_id = self.store.add(request["sender_repo"], request["target_repo"], request["body"], endpoint=endpoint_id) response = {"ok": True, "message_id": message_id, "state": "pending"} except (RegistryError, ValueError) as exc: response = {"ok": False, "error": str(exc)} elif op == "history": response = {"ok": True, "messages": [dict(row) for row in self.store.list(request.get("target_repo"), request.get("state"))]} elif op == "ack": if not request.get("message_id"): response = {"ok": False, "error": "ack requires message_id"} else: ok = self.store.acknowledge(request["message_id"]) response = {"ok": ok, "message_id": request["message_id"], "state": "acknowledged" if ok else "missing"} elif op == "endpoints": response = {"ok": True, "endpoints": [dict(row) for row in self.store.endpoints()]} elif op == "disconnect": endpoint_id = request.get("endpoint_id") if not endpoint_id: response = {"ok": False, "error": "disconnect requires endpoint_id"} else: self.store.disconnect_endpoint(endpoint_id) response = {"ok": True, "endpoint_id": endpoint_id, "state": "disconnected"} else: response = {"ok": False, "error": "unsupported operation"} except (json.JSONDecodeError, UnicodeDecodeError): response = {"ok": False, "error": "invalid JSON"} writer.write((json.dumps(response) + "\n").encode()) await writer.drain() writer.close() await writer.wait_closed() @staticmethod def _peer_allowed(writer: asyncio.StreamWriter) -> bool: sock = writer.get_extra_info("socket") if sock is None or not hasattr(socket, "SO_PEERCRED"): return True try: _, uid, _ = struct.unpack("3i", sock.getsockopt(socket.SOL_SOCKET, socket.SO_PEERCRED, struct.calcsize("3i"))) return uid == os.getuid() except OSError: return False def close(self) -> None: self.store.close() async def ping(path: Path | None = None) -> bool: try: reader, writer = await asyncio.open_unix_connection(str(path or socket_path())) writer.write(b'{"op":"ping"}\n'); await writer.drain() result = json.loads(await reader.readline()) writer.close(); await writer.wait_closed() return bool(result.get("ok")) except (OSError, json.JSONDecodeError): return False async def request(payload: dict, path: Path | None = None) -> dict: reader, writer = await asyncio.open_unix_connection(str(path or socket_path())) writer.write((json.dumps(payload) + "\n").encode()) await writer.drain() response = json.loads(await reader.readline()) writer.close() await writer.wait_closed() return response