From d03babfddc321d60fc26ef6e33d719b9a65df8fd Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 22:42:55 -0700 Subject: [PATCH] =?UTF-8?q?refactor(gateway):=20wave-2=20final=20compactio?= =?UTF-8?q?n=20=E2=80=94=20Platform.=5Fmissing=5F=20single=20exit,=20=5Fcl?= =?UTF-8?q?eanup=5Fexpired=20keep-set,=20expression-form=20display=20norma?= =?UTF-8?q?lisers=20and=20small=20predicates?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- gateway/authz_mixin.py | 6 ++---- gateway/channel_directory.py | 9 +++------ gateway/config.py | 35 +++++++++++++---------------------- gateway/config_env.py | 15 ++++----------- gateway/config_loader.py | 6 ++---- gateway/display_config.py | 12 +++--------- gateway/pairing.py | 30 +++++++++++++----------------- gateway/platform_registry.py | 4 +--- 8 files changed, 41 insertions(+), 76 deletions(-) diff --git a/gateway/authz_mixin.py b/gateway/authz_mixin.py index 0c5fea66a0..6f5efab2b3 100644 --- a/gateway/authz_mixin.py +++ b/gateway/authz_mixin.py @@ -325,9 +325,7 @@ class GatewayAuthorizationMixin: """Per-profile PairingStore for a source, else the global ``self.pairing_store``.""" per_profile = getattr(self, "pairing_stores", None) or {} profile = getattr(source, "profile", None) - if profile and profile in per_profile: - return per_profile[profile] - return getattr(self, "pairing_store", None) + return per_profile[profile] if profile and profile in per_profile else getattr(self, "pairing_store", None) def _adapter_extra_for_source(self, source) -> dict: return _adapter_config_extra(self._adapter_for_source(source)) @@ -396,7 +394,7 @@ class GatewayAuthorizationMixin: resolved_ids = resolver() if not isinstance(resolved_ids, (set, frozenset, list, tuple)): return set() - return {str(entry).strip() for entry in resolved_ids if isinstance(entry, (str, int)) and str(entry).strip()} + return {s for e in resolved_ids if isinstance(e, (str, int)) and (s := str(e).strip())} def _chat_scoped_grant(self, source, adapter_profile, is_group: bool, allow_adapter_delegation: bool) -> bool: """Grants that need no ``user_id`` (checked before the no-user-id guard).""" diff --git a/gateway/channel_directory.py b/gateway/channel_directory.py index c895b27a68..3ad56083a4 100644 --- a/gateway/channel_directory.py +++ b/gateway/channel_directory.py @@ -76,16 +76,13 @@ def _apply_channel_aliases(platforms: Dict[str, Any]) -> None: for chat_id, friendly in id_map.items(): if not isinstance(friendly, str) or not friendly.strip(): continue - chat_id = str(chat_id) - friendly = friendly.strip() + chat_id, friendly = str(chat_id), friendly.strip() matches = [e for e in entries if isinstance(e, dict) and e.get("id") == chat_id] for e in matches: e["name"] = friendly if not matches: - entries.append({ - "id": chat_id, "name": friendly, - "type": "group" if chat_id.endswith("@g.us") else "dm", "thread_id": None, - }) + entries.append({"id": chat_id, "name": friendly, "thread_id": None, + "type": "group" if chat_id.endswith("@g.us") else "dm"}) def _normalize_channel_query(value: str) -> str: diff --git a/gateway/config.py b/gateway/config.py index a0a20c0491..4b810ec934 100644 --- a/gateway/config.py +++ b/gateway/config.py @@ -166,11 +166,8 @@ def _coerce_dict(value: Any) -> Dict[str, Any]: def _normalize_choice(value: Any, choices: set, default: str) -> str: """Lower-cased *value* when it is one of *choices*, else *default*.""" - if isinstance(value, str): - normalized = value.strip().lower() - if normalized in choices: - return normalized - return default + normalized = value.strip().lower() if isinstance(value, str) else None + return normalized if normalized in choices else default def _dict_slot(container: dict, key: str) -> dict: @@ -192,8 +189,7 @@ def _getenv(name: str, default: Optional[str] = None) -> Optional[str]: def _getenv_str(name: str, default: str = "") -> str: - val = _getenv(name, default) - return val if val is not None else default + return val if (val := _getenv(name, default)) is not None else default _Platform__bundled_plugin_names: Optional[set] = None # cached outside the enum: never a member @@ -235,17 +231,15 @@ class Platform(Enum): value = value.strip().lower() if value in cls._value2member_map_: return cls._value2member_map_[value] - global _Platform__bundled_plugin_names if _Platform__bundled_plugin_names is None: _Platform__bundled_plugin_names = cls._scan_bundled_plugin_platforms() - if value in _Platform__bundled_plugin_names: - return cls._add_pseudo_member(value) - with contextlib.suppress(Exception): - from gateway.platform_registry import platform_registry - if platform_registry.is_registered(value): - return cls._add_pseudo_member(value) - return None + registered = value in _Platform__bundled_plugin_names + if not registered: + with contextlib.suppress(Exception): + from gateway.platform_registry import platform_registry + registered = platform_registry.is_registered(value) + return cls._add_pseudo_member(value) if registered else None @classmethod def _add_pseudo_member(cls, value: str) -> "Platform": @@ -306,9 +300,8 @@ class HomeChannel: scope_id: Optional[str] = None def to_dict(self) -> Dict[str, Any]: - result = {"platform": self.platform.value, "chat_id": self.chat_id, "name": self.name} - result.update({k: v for k in ("thread_id", "user_id", "scope_id") if (v := getattr(self, k))}) - return result + optional = {k: v for k in ("thread_id", "user_id", "scope_id") if (v := getattr(self, k))} + return {"platform": self.platform.value, "chat_id": self.chat_id, "name": self.name, **optional} @classmethod def from_dict(cls, data: Dict[str, Any]) -> "HomeChannel": @@ -319,7 +312,6 @@ class HomeChannel: def persist_home_channel(home: HomeChannel, *, enabled_if_new: bool = False) -> None: """Persist a logical home without falsely enabling a Relay-fronted adapter.""" from hermes_cli.config import load_config, save_config - config = load_config() platform_config = _dict_slot(_dict_slot(config, "platforms"), home.platform.value) if enabled_if_new: @@ -502,9 +494,9 @@ def _has_usable_api_server_key(key: object) -> bool: return False try: from hermes_cli.auth import has_usable_secret + return has_usable_secret(key, min_length=16) except ImportError: return len(str(key).strip()) >= 16 - return has_usable_secret(key, min_length=16) def _needs_extra(*keys: str) -> Callable[[PlatformConfig], bool]: @@ -624,8 +616,7 @@ class GatewayConfig: return False def get_home_channel(self, platform: Platform) -> Optional[HomeChannel]: - config = self.platforms.get(platform) - return config.home_channel if config else None + return self.platforms[platform].home_channel if self.platforms.get(platform) else None def get_reset_policy(self, platform: Optional[Platform] = None, session_type: Optional[str] = None) -> SessionResetPolicy: """Priority: platform override > type override > default.""" diff --git a/gateway/config_env.py b/gateway/config_env.py index 9b3aed0f0e..5cf58c7077 100644 --- a/gateway/config_env.py +++ b/gateway/config_env.py @@ -244,8 +244,7 @@ def _ReplyMode(platform: Platform, env: str): # --- platform-unique branches ------------------------------------------------ def _telegram_fallback_ips(config: GatewayConfig) -> None: - ips = getenv("TELEGRAM_FALLBACK_IPS") - if ips: + if ips := getenv("TELEGRAM_FALLBACK_IPS"): config.platforms.setdefault(Platform.TELEGRAM, PlatformConfig()).extra["fallback_ips"] = _csv_list(ips) @@ -265,8 +264,7 @@ def _whatsapp(config: GatewayConfig) -> None: def _slack_home(config: GatewayConfig) -> None: """SLACK_HOME_CHANNEL creates a disabled Slack entry if needed; user_id/scope_id provenance survives an unchanged chat_id.""" - slack_home = getenv("SLACK_HOME_CHANNEL") - if not slack_home: + if not (slack_home := getenv("SLACK_HOME_CHANNEL")): return slack_config = config.platforms.setdefault(Platform.SLACK, PlatformConfig(enabled=False)) existing_home = slack_config.home_channel @@ -281,10 +279,7 @@ def _slack_home(config: GatewayConfig) -> None: def _matrix_e2ee(config: GatewayConfig, matrix_config: PlatformConfig) -> None: mode = getenv("MATRIX_E2EE_MODE").strip().lower() - matrix_config.extra["encryption"] = ( - mode in ("required", "require", "optional", "prefer", "preferred") - or is_truthy_value(getenv("MATRIX_ENCRYPTION")) - ) + matrix_config.extra["encryption"] = mode in ("required", "require", "optional", "prefer", "preferred") or is_truthy_value(getenv("MATRIX_ENCRYPTION")) if mode: matrix_config.extra["e2ee_mode"] = mode _env_extras(matrix_config.extra, (("device_id", "MATRIX_DEVICE_ID"),)) @@ -348,8 +343,7 @@ def _qq_home(config: GatewayConfig, qq_config: PlatformConfig) -> None: def _session_settings(config: GatewayConfig) -> None: for env, attr in (("SESSION_IDLE_MINUTES", "idle_minutes"), ("SESSION_RESET_HOUR", "at_hour")): - raw = getenv(env) - if raw: + if raw := getenv(env): with contextlib.suppress(ValueError): setattr(config.default_reset_policy, attr, int(raw)) @@ -486,7 +480,6 @@ def _scrub_explicit_markers(config: GatewayConfig) -> None: for platform_config in config.platforms.values(): platform_config.extra.pop("_enabled_explicit", None) - # Order is significant: a home channel only attaches to a platform that already exists (Telegram's # reply mode may create the entry first; Discord reads home first). Relay disabling runs after the # plugin pass; the marker scrub must be last. diff --git a/gateway/config_loader.py b/gateway/config_loader.py index ec6dbad729..2c88960570 100644 --- a/gateway/config_loader.py +++ b/gateway/config_loader.py @@ -176,10 +176,8 @@ def platform_section(yaml_cfg: dict, name: str, gateway_platforms: Any) -> tuple section = yaml_cfg.get(name) toplevel = isinstance(section, dict) if not toplevel: - for src in (gateway_platforms, yaml_cfg.get("platforms")): - if isinstance(src, dict) and isinstance(src.get(name), dict): - section = src[name] - break + nested = (src[name] for src in (gateway_platforms, yaml_cfg.get("platforms")) if isinstance(src, dict) and isinstance(src.get(name), dict)) + section = next(nested, section) return section, toplevel diff --git a/gateway/display_config.py b/gateway/display_config.py index 0003ce7ea3..b95086aa61 100644 --- a/gateway/display_config.py +++ b/gateway/display_config.py @@ -121,21 +121,15 @@ def _norm_tristate(on: str, off: str, choices: set, extra_truthy: set = frozense def _norm_bool(value: Any) -> bool: - if isinstance(value, str): - return value.strip().lower() in _TRUTHY | {"raw", "verbose"} - return bool(value) + return value.strip().lower() in _TRUTHY | {"raw", "verbose"} if isinstance(value, str) else bool(value) def _norm_long_running(value: Any) -> Any: - if isinstance(value, str) and value.strip().lower() == "generic": - return "generic" - return _norm_bool(value) + return "generic" if isinstance(value, str) and value.strip().lower() == "generic" else _norm_bool(value) def _norm_cleanup_progress(value: Any) -> bool: - if isinstance(value, str): - return value.lower() in _TRUTHY - return bool(value) + return value.lower() in _TRUTHY if isinstance(value, str) else bool(value) def _norm_choice(choices: tuple[str, ...]) -> Any: diff --git a/gateway/pairing.py b/gateway/pairing.py index 278980316a..5754254329 100644 --- a/gateway/pairing.py +++ b/gateway/pairing.py @@ -87,9 +87,7 @@ def _platform_uses_whatsapp_identity(platform: str) -> bool: def _normalize_user_id(platform: str, user_id: str) -> str: """Normalize platform-specific user IDs before persisting / comparing them.""" raw_user_id = str(user_id or "").strip() - if _platform_uses_whatsapp_identity(platform): - return normalize_whatsapp_identifier(raw_user_id) or raw_user_id - return raw_user_id + return (normalize_whatsapp_identifier(raw_user_id) or raw_user_id) if _platform_uses_whatsapp_identity(platform) else raw_user_id def _user_id_aliases(platform: str, user_id: str) -> set[str]: @@ -351,6 +349,18 @@ class PairingStore: def _rate_limit_path(self) -> Path: return self._dir / "_rate_limits.json" + def _cleanup_expired(self, platform: str) -> None: + """Remove expired pending codes; malformed/legacy entries (no numeric ``created_at``) count as expired.""" + path = self._pending_path(platform) + pending = self._load_json(path) + now = time.time() + live = { + k: v for k, v in pending.items() + if (created := _entry_created_at(v)) is not None and (now - created) <= CODE_TTL_SECONDS + } + if len(live) != len(pending): + self._save_json(path, live) + _load_json = staticmethod(_load_json_file) _save_json = staticmethod(_save_json_file) @@ -568,20 +578,6 @@ class PairingStore: limits[fail_key] = 0 self._save_limits(limits) - # ----- Cleanup ----- - - def _cleanup_expired(self, platform: str) -> None: - """Remove expired pending codes; malformed/legacy entries (no numeric ``created_at``) count as expired.""" - path = self._pending_path(platform) - pending = self._load_json(path) - now = time.time() - expired = [ - entry_id for entry_id, info in pending.items() - if (created := _entry_created_at(info)) is None or (now - created) > CODE_TTL_SECONDS - ] - if expired: - self._save_json(path, {k: v for k, v in pending.items() if k not in expired}) - def _all_platforms(self, suffix: str) -> list: """Platforms that have a ``-.json`` data file (``_``-prefixed files are shared state).""" tail = f"-{suffix}.json" diff --git a/gateway/platform_registry.py b/gateway/platform_registry.py index 97d9ffd1a6..3feb57dde5 100644 --- a/gateway/platform_registry.py +++ b/gateway/platform_registry.py @@ -138,9 +138,7 @@ class PlatformRegistry: return entry, loader def _prune_scope(self, scope: Optional[str]) -> None: - if scope is None: - return - for maps in (self._scoped_entries, self._scoped_deferred): + for maps in (self._scoped_entries, self._scoped_deferred) if scope is not None else (): if not maps.get(scope): maps.pop(scope, None)