151 lines
6.7 KiB
Python
151 lines
6.7 KiB
Python
from __future__ import annotations
|
|
|
|
from typing import Any
|
|
|
|
_SCHEMA: dict[str, dict[str, Any]] = {
|
|
"homeserver": {"type": "str", "required": False, "default": "https://matrix.org"},
|
|
"user_id": {"type": "str", "required": True, "pattern": r"^@[\w._=\-/]+:[\w.\-]+(?:\.[\w.\-]+)*$"},
|
|
"access_token": {"type": "str", "required": False},
|
|
"password": {"type": "str", "required": False},
|
|
"device_id": {"type": "str", "required": False, "default": "yuxi-bot"},
|
|
"deviceName": {"type": "str", "required": False},
|
|
"crypto_store_dir": {"type": "str", "required": False},
|
|
"proxy": {"type": "str", "required": False},
|
|
"display_name": {"type": "str", "required": False},
|
|
"avatarUrl": {"type": "str", "required": False},
|
|
"sync_timeout_ms": {"type": "int", "required": False, "default": 30000, "min": 1000},
|
|
"loop_sleep_ms": {"type": "int", "required": False, "default": 100},
|
|
"probe_timeout_s": {"type": "float", "required": False, "min": 1},
|
|
"responsePrefix": {"type": "str", "required": False},
|
|
"chunkMode": {"type": "str", "required": False, "enum": ["length", "newline"], "default": "length"},
|
|
"mediaMaxMb": {"type": "int", "required": False, "default": 100, "min": 1},
|
|
"requireMention": {"type": "bool", "required": False, "default": True},
|
|
"dangerouslyAllowPrivateNetwork": {"type": "bool", "required": False, "default": False},
|
|
"autoJoin": {"type": "str", "required": False, "enum": ["off", "always", "allowlist", "invites"], "default": "off"},
|
|
"enabled": {"type": "bool", "required": False, "default": True},
|
|
"name": {"type": "str", "required": False},
|
|
"defaultAccount": {"type": "str", "required": False},
|
|
"dm.policy": {
|
|
"type": "str",
|
|
"required": False,
|
|
"enum": ["open", "disabled", "allowlist", "pairing"],
|
|
"default": "pairing",
|
|
},
|
|
"groupPolicy": {
|
|
"type": "str",
|
|
"required": False,
|
|
"enum": ["open", "disabled", "allowlist"],
|
|
"default": "allowlist",
|
|
},
|
|
"allowlistOnly": {"type": "bool", "required": False, "default": False},
|
|
"contextVisibility": {"type": "str", "required": False, "enum": ["all", "none", "dm", "group"], "default": "all"},
|
|
"replyToMode": {"type": "str", "required": False, "enum": ["off", "first", "all", "batched"], "default": "all"},
|
|
"ackReaction": {"type": "str", "required": False},
|
|
"ackReactionScope": {
|
|
"type": "str",
|
|
"required": False,
|
|
"enum": ["always", "dm-only", "group-mentions", "never"],
|
|
"default": "group-mentions",
|
|
},
|
|
"reactionNotifications": {"type": "str", "required": False, "enum": ["all", "own", "none"], "default": "own"},
|
|
"encryption": {"type": "bool", "required": False, "default": False},
|
|
"startupVerification": {
|
|
"type": "str",
|
|
"required": False,
|
|
"enum": ["if-unverified", "always", "never"],
|
|
"default": "if-unverified",
|
|
},
|
|
"startupVerificationCooldownHours": {"type": "int", "required": False, "default": 24},
|
|
"blockStreaming": {"type": "bool", "required": False, "default": False},
|
|
"textChunkLimit": {"type": "int", "required": False, "min": 100, "max": 65536},
|
|
"historyLimit": {"type": "int", "required": False, "default": 0},
|
|
"allowBots": {"type": "bool", "required": False},
|
|
"execApprovals.enabled": {"type": "str", "required": False, "enum": ["auto", "true", "false"], "default": "auto"},
|
|
}
|
|
|
|
_SCHEMA_FLAT: dict[str, dict[str, Any]] = _SCHEMA
|
|
|
|
_KNOWN_SECTIONS = {"dm", "rooms", "accounts", "threadBindings", "streaming", "actions", "groups", "execApprovals"}
|
|
|
|
|
|
def validate_config(config: dict[str, Any]) -> list[dict[str, str]]:
|
|
errors: list[dict[str, str]] = []
|
|
seen_auth = False
|
|
|
|
for key, rules in _SCHEMA_FLAT.items():
|
|
if "." in key:
|
|
parts = key.split(".")
|
|
section = parts[0]
|
|
sub_key = parts[1] if len(parts) == 2 else ""
|
|
value = config.get(section, {}).get(sub_key)
|
|
else:
|
|
value = config.get(key)
|
|
|
|
if rules.get("required") and value is None:
|
|
errors.append({"field": key, "error": "required field is missing"})
|
|
|
|
if value is not None:
|
|
errors.extend(_validate_field(key, value, rules))
|
|
|
|
if not config.get("access_token") and not config.get("password"):
|
|
errors.append({"field": "access_token|password", "error": "at least one auth method is required"})
|
|
|
|
if config.get("encryption") and not config.get("crypto_store_dir"):
|
|
errors.append({"field": "crypto_store_dir", "error": "required when encryption is enabled"})
|
|
|
|
return errors
|
|
|
|
|
|
def _validate_field(key: str, value: Any, rules: dict[str, Any]) -> list[dict[str, str]]:
|
|
import re
|
|
|
|
errors: list[dict[str, str]] = []
|
|
expected_type = rules.get("type")
|
|
|
|
if expected_type == "str" and not isinstance(value, str):
|
|
errors.append({"field": key, "error": f"expected str, got {type(value).__name__}"})
|
|
elif expected_type == "int" and not isinstance(value, int):
|
|
errors.append({"field": key, "error": f"expected int, got {type(value).__name__}"})
|
|
elif expected_type == "float" and not isinstance(value, (int, float)):
|
|
errors.append({"field": key, "error": f"expected float, got {type(value).__name__}"})
|
|
elif expected_type == "bool" and not isinstance(value, bool):
|
|
errors.append({"field": key, "error": f"expected bool, got {type(value).__name__}"})
|
|
|
|
if "enum" in rules and isinstance(value, str):
|
|
if value not in rules["enum"]:
|
|
errors.append({"field": key, "error": f"value '{value}' not in {rules['enum']}"})
|
|
|
|
if "min" in rules and isinstance(value, (int, float)):
|
|
if value < rules["min"]:
|
|
errors.append({"field": key, "error": f"value {value} below minimum {rules['min']}"})
|
|
|
|
if "max" in rules and isinstance(value, (int, float)):
|
|
if value > rules["max"]:
|
|
errors.append({"field": key, "error": f"value {value} above maximum {rules['max']}"})
|
|
|
|
if "pattern" in rules and isinstance(value, str):
|
|
pattern = rules["pattern"]
|
|
if not re.match(pattern, value):
|
|
errors.append({"field": key, "error": f"value does not match pattern {pattern}"})
|
|
|
|
return errors
|
|
|
|
|
|
def get_config_schema() -> dict[str, dict[str, Any]]:
|
|
return dict(_SCHEMA_FLAT)
|
|
|
|
|
|
def get_default_config() -> dict[str, Any]:
|
|
defaults: dict[str, Any] = {}
|
|
for key, rules in _SCHEMA_FLAT.items():
|
|
if "default" in rules:
|
|
if "." in key:
|
|
parts = key.split(".")
|
|
section = parts[0]
|
|
sub_key = parts[1] if len(parts) == 2 else ""
|
|
defaults.setdefault(section, {})
|
|
defaults[section][sub_key] = rules["default"]
|
|
else:
|
|
defaults[key] = rules["default"]
|
|
return defaults
|