from __future__ import annotations import os from typing import TYPE_CHECKING, Any if TYPE_CHECKING: from nio import AsyncClient from .config_schema import validate_config from .encryption import check_crypto_store_ready from .sync_store import SyncTokenStore, SSRFGuard class MatrixDoctor: def __init__(self, config: dict[str, Any], client: AsyncClient | None = None): self._config = config self._client = client self._issues: list[dict[str, Any]] = [] async def diagnose(self) -> list[dict[str, Any]]: self._issues = [] self._check_config_validity() self._check_auth_method() self._check_homeserver_reachability() self._check_crypto_store() self._check_sync_token() self._check_private_network() self._check_encryption_config() return self._issues def _check_config_validity(self) -> None: errors = validate_config(self._config) for err in errors: self._issues.append( { "severity": "error", "category": "config", "field": err["field"], "message": err["error"], } ) def _check_auth_method(self) -> None: has_token = bool(self._config.get("access_token")) has_password = bool(self._config.get("password")) if not has_token and not has_password: self._issues.append( { "severity": "error", "category": "auth", "message": "No authentication method configured (access_token or password required)", } ) def _check_homeserver_reachability(self) -> None: homeserver = self._config.get("homeserver", "") if not homeserver: self._issues.append( { "severity": "warning", "category": "connection", "message": "Homeserver URL not configured", } ) def _check_crypto_store(self) -> None: crypto_dir = self._config.get("crypto_store_dir", "") if not crypto_dir and self._config.get("encryption"): self._issues.append( { "severity": "error", "category": "encryption", "message": "Encryption enabled but crypto_store_dir not configured", } ) return if crypto_dir: result = check_crypto_store_ready(crypto_dir) if result["status"] == "unavailable": self._issues.append( { "severity": "error", "category": "encryption", "message": f"Crypto store unavailable: {', '.join(result.get('details', []))}", } ) elif result["status"] == "empty" and self._config.get("encryption"): self._issues.append( { "severity": "info", "category": "encryption", "message": "Crypto store empty — keys will be created on first encrypted message", } ) def _check_sync_token(self) -> None: crypto_dir = self._config.get("crypto_store_dir", "./matrix_crypto_store") user_id = self._config.get("user_id", "unknown") store = SyncTokenStore(crypto_dir, user_id) token = store.load() if token: self._issues.append( { "severity": "info", "category": "sync", "message": "Sync token found — incremental sync will be used", } ) def _check_private_network(self) -> None: homeserver = self._config.get("homeserver", "") if homeserver: from urllib.parse import urlparse parsed = urlparse(homeserver) hostname = parsed.hostname or "" if hostname and SSRFGuard.is_private_host(hostname): allowed = self._config.get("dangerouslyAllowPrivateNetwork", False) if not allowed: self._issues.append( { "severity": "warning", "category": "security", "message": f"Homeserver {homeserver} appears to be on a private network. Set dangerouslyAllowPrivateNetwork=true to allow.", } ) def _check_encryption_config(self) -> None: encryption_enabled = self._config.get("encryption", False) if encryption_enabled: if not self._config.get("crypto_store_dir"): self._issues.append( { "severity": "error", "category": "encryption", "message": "Encryption requires crypto_store_dir to be configured", } ) async def repair(self) -> list[dict[str, Any]]: repairs: list[dict[str, Any]] = [] crypto_dir = self._config.get("crypto_store_dir", "") if crypto_dir and not os.path.exists(crypto_dir): try: os.makedirs(crypto_dir, exist_ok=True) repairs.append( { "status": "repaired", "category": "encryption", "message": f"Created crypto store directory: {crypto_dir}", } ) except OSError as e: repairs.append( { "status": "failed", "category": "encryption", "message": f"Failed to create crypto store directory: {e}", } ) return repairs def _check_homeserver_ssrf(self) -> None: from urllib.parse import urlparse homeserver = self._config.get("homeserver", "") if not homeserver: return parsed = urlparse(homeserver) hostname = parsed.hostname or "" if hostname: is_private = SSRFGuard.is_private_host(hostname) allowed = self._config.get("dangerouslyAllowPrivateNetwork", False) if is_private and not allowed: self._issues.append( { "severity": "warning", "category": "security", "message": f"Homeserver {homeserver} is on a private network. Set dangerouslyAllowPrivateNetwork=true to connect.", } ) def get_summary(self) -> dict[str, int]: summary = {"error": 0, "warning": 0, "info": 0} for issue in self._issues: severity = issue.get("severity", "info") if severity in summary: summary[severity] += 1 return summary @property def has_critical_issues(self) -> bool: return any(issue.get("severity") == "error" for issue in self._issues)