207 lines
7.2 KiB
Python
207 lines
7.2 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
from collections.abc import Callable
|
|
from typing import Any
|
|
|
|
from yuxi.utils.logging_config import logger
|
|
|
|
|
|
class AccountInfo:
|
|
def __init__(self, name: str, ship: str, url: str, code: str, metadata: dict[str, Any] | None = None):
|
|
self.name = name
|
|
self.ship = ship.lstrip("~")
|
|
self.url = url.rstrip("/")
|
|
self.code = code
|
|
self.metadata = metadata or {}
|
|
self.status: str = "disconnected"
|
|
self.last_probe_at: float = 0.0
|
|
|
|
@property
|
|
def ship_with_tilde(self) -> str:
|
|
return f"~{self.ship}"
|
|
|
|
def to_dict(self) -> dict[str, Any]:
|
|
return {
|
|
"name": self.name,
|
|
"ship": self.ship_with_tilde,
|
|
"url": self.url,
|
|
"status": self.status,
|
|
"metadata": self.metadata,
|
|
}
|
|
|
|
|
|
class AccountManager:
|
|
def __init__(self, accounts_config: list[dict[str, Any]] | None = None):
|
|
self._accounts: dict[str, AccountInfo] = {}
|
|
self._on_status_change: list[Callable[[str, str], Any]] = []
|
|
if accounts_config:
|
|
self.load_from_config(accounts_config)
|
|
|
|
def load_from_config(self, accounts_config: list[dict[str, Any]]) -> None:
|
|
self._accounts.clear()
|
|
for acct in accounts_config:
|
|
name = acct.get("name", acct.get("ship", "default"))
|
|
ship = acct.get("ship", "").lstrip("~")
|
|
url = acct.get("url", "")
|
|
code = acct.get("code", "")
|
|
if not ship or not url or not code:
|
|
logger.warning(f"[Urbit] Skipping invalid account: {name}")
|
|
continue
|
|
self._accounts[name] = AccountInfo(
|
|
name=name,
|
|
ship=ship,
|
|
url=url,
|
|
code=code,
|
|
metadata=acct.get("metadata", {}),
|
|
)
|
|
logger.info(f"[Urbit] Loaded {len(self._accounts)} accounts")
|
|
|
|
def add_account(self, name: str, ship: str, url: str, code: str, **meta) -> AccountInfo:
|
|
info = AccountInfo(name=name, ship=ship, url=url, code=code, metadata=meta)
|
|
self._accounts[name] = info
|
|
logger.info(f"[Urbit] Account '{name}' added (~{info.ship})")
|
|
return info
|
|
|
|
def remove_account(self, name: str) -> bool:
|
|
if name in self._accounts:
|
|
del self._accounts[name]
|
|
logger.info(f"[Urbit] Account '{name}' removed")
|
|
return True
|
|
return False
|
|
|
|
def get(self, name: str = "default") -> AccountInfo | None:
|
|
if not self._accounts and name == "default":
|
|
return None
|
|
return self._accounts.get(name)
|
|
|
|
def get_all(self) -> dict[str, AccountInfo]:
|
|
return dict(self._accounts)
|
|
|
|
def list_ships(self) -> list[str]:
|
|
return [f"~{a.ship}" for a in self._accounts.values()]
|
|
|
|
def set_status(self, name: str, status: str) -> None:
|
|
acct = self._accounts.get(name)
|
|
if acct:
|
|
old = acct.status
|
|
acct.status = status
|
|
if old != status:
|
|
for handler in self._on_status_change:
|
|
try:
|
|
handler(name, status)
|
|
except Exception as e:
|
|
logger.warning(f"[Urbit] Status change handler error: {e}")
|
|
|
|
def on_status_change(self, handler: Callable[[str, str], Any]) -> None:
|
|
self._on_status_change.append(handler)
|
|
|
|
@property
|
|
def is_empty(self) -> bool:
|
|
return len(self._accounts) == 0
|
|
|
|
@property
|
|
def count(self) -> int:
|
|
return len(self._accounts)
|
|
|
|
def has_multi_accounts(self) -> bool:
|
|
return len(self._accounts) > 1
|
|
|
|
|
|
class MultiAccountConnectionManager:
|
|
def __init__(
|
|
self,
|
|
account_manager: AccountManager,
|
|
connection_factory: Callable[[AccountInfo], Any],
|
|
):
|
|
self._account_manager = account_manager
|
|
self._connection_factory = connection_factory
|
|
self._connections: dict[str, Any] = {}
|
|
self._tasks: dict[str, asyncio.Task | None] = {}
|
|
self._running = False
|
|
|
|
@property
|
|
def connections(self) -> dict[str, Any]:
|
|
return dict(self._connections)
|
|
|
|
@property
|
|
def is_running(self) -> bool:
|
|
return self._running
|
|
|
|
async def start_all(self) -> None:
|
|
if self._running:
|
|
return
|
|
|
|
self._running = True
|
|
accounts = self._account_manager.get_all()
|
|
|
|
for name, account in accounts.items():
|
|
await self._start_account(name, account)
|
|
|
|
logger.info(f"[Urbit] MultiAccountConnectionManager started with {len(self._connections)} connections")
|
|
|
|
async def _start_account(self, name: str, account: AccountInfo) -> None:
|
|
try:
|
|
connection = await self._connection_factory(account)
|
|
self._connections[name] = connection
|
|
self._account_manager.set_status(name, "connected")
|
|
logger.info(f"[Urbit] Account '{name}' (~{account.ship}) connected")
|
|
except Exception as e:
|
|
self._account_manager.set_status(name, "error")
|
|
logger.error(f"[Urbit] Failed to connect account '{name}' (~{account.ship}): {e}")
|
|
|
|
async def start_account(self, name: str, account: AccountInfo) -> None:
|
|
if name in self._connections:
|
|
logger.warning(f"[Urbit] Account '{name}' already connected")
|
|
return
|
|
await self._start_account(name, account)
|
|
|
|
async def stop_account(self, name: str) -> None:
|
|
connection = self._connections.pop(name, None)
|
|
if connection is None:
|
|
return
|
|
|
|
if hasattr(connection, "disconnect"):
|
|
try:
|
|
await connection.disconnect()
|
|
except Exception as e:
|
|
logger.warning(f"[Urbit] Error disconnecting account '{name}': {e}")
|
|
|
|
self._account_manager.set_status(name, "disconnected")
|
|
logger.info(f"[Urbit] Account '{name}' disconnected")
|
|
|
|
async def stop_all(self) -> None:
|
|
if not self._running:
|
|
return
|
|
|
|
self._running = False
|
|
names = list(self._connections.keys())
|
|
for name in names:
|
|
await self.stop_account(name)
|
|
|
|
self._connections.clear()
|
|
logger.info("[Urbit] MultiAccountConnectionManager stopped all connections")
|
|
|
|
def get_connection(self, name: str = "default") -> Any | None:
|
|
return self._connections.get(name)
|
|
|
|
async def send_to_account(self, name: str, *args, **kwargs) -> Any:
|
|
connection = self._connections.get(name)
|
|
if connection is None:
|
|
raise RuntimeError(f"Account '{name}' not connected")
|
|
if hasattr(connection, "send"):
|
|
return await connection.send(*args, **kwargs)
|
|
raise RuntimeError(f"Connection for account '{name}' does not support send()")
|
|
|
|
async def health_check_all(self) -> dict[str, Any]:
|
|
result: dict[str, Any] = {}
|
|
for name, connection in self._connections.items():
|
|
if hasattr(connection, "health_check"):
|
|
try:
|
|
result[name] = await connection.health_check()
|
|
except Exception as e:
|
|
result[name] = {"status": "unhealthy", "error": str(e)}
|
|
else:
|
|
result[name] = {"status": "unknown"}
|
|
return result
|