refactor(gateway): collapse authz signatures/ladders, pairing purge/expire, slack resolve branches

This commit is contained in:
Teknium
2026-09-02 19:09:23 -07:00
parent a6bde733fb
commit dc40ec89e9
3 changed files with 60 additions and 157 deletions

View File

@@ -179,11 +179,7 @@ class GatewayAuthorizationMixin:
def _primary_adapters(self) -> dict:
return getattr(self, "adapters", None) or {}
def _authorization_adapter(
self,
platform: Optional[Platform],
profile: Optional[str] = None,
):
def _authorization_adapter(self, platform: Optional[Platform], profile: Optional[str] = None):
"""Live adapter whose intake policy gates authorization.
Secondary-profile adapters live in ``_profile_adapters[profile]``; the primary
@@ -228,10 +224,7 @@ class GatewayAuthorizationMixin:
if getattr(source, "delivered_via_upstream_relay", False) is True:
return self._primary_adapters().get(Platform.RELAY)
# ``getattr``: test fixtures build bare SimpleNamespace sources without ``profile``.
return self._authorization_adapter(
getattr(source, "platform", None),
getattr(source, "profile", None),
)
return self._authorization_adapter(getattr(source, "platform", None), getattr(source, "profile", None))
def _owning_profile(self, adapter, platform):
"""Return (registered, profile) for a live adapter: profile is None for primary."""
@@ -273,12 +266,7 @@ class GatewayAuthorizationMixin:
adapter = self._authorization_adapter(platform, profile)
return adapter is not None and bool(getattr(adapter, name, False))
def _adapter_authorization_is_upstream(
self,
platform: Optional[Platform],
*,
profile: Optional[str] = None,
) -> bool:
def _adapter_authorization_is_upstream(self, platform: Optional[Platform], *, profile: Optional[str] = None) -> bool:
"""Whether the adapter delegates authz to a trusted authenticated upstream (relay).
Unlike ``_adapter_enforces_own_access_policy`` (a LOCAL policy only trusted when
@@ -286,12 +274,7 @@ class GatewayAuthorizationMixin:
"""
return self._adapter_flag(platform, "authorization_is_upstream", profile)
def _adapter_enforces_own_access_policy(
self,
platform: Optional[Platform],
*,
profile: Optional[str] = None,
) -> bool:
def _adapter_enforces_own_access_policy(self, platform: Optional[Platform], *, profile: Optional[str] = None) -> bool:
"""Whether the adapter gates access at intake itself (WeCom, Weixin, Yuanbao, QQBot, WhatsApp).
The flag alone is NOT "already authorized": these adapters default to ``open``,
@@ -301,13 +284,8 @@ class GatewayAuthorizationMixin:
def _config_extra(self, platform) -> dict:
"""``config.platforms[platform].extra`` as a dict ({} when absent)."""
config = getattr(self, "config", None)
platform_cfg = (
config.platforms.get(platform)
if config is not None and hasattr(config, "platforms")
else None
)
extra = getattr(platform_cfg, "extra", None) if platform_cfg else None
platforms = getattr(getattr(self, "config", None), "platforms", None)
extra = getattr(platforms.get(platform), "extra", None) if platforms is not None else None
return extra if isinstance(extra, dict) else {}
def _adapter_setting(self, platform, attr: str, extra_key: str, profile):
@@ -324,30 +302,16 @@ class GatewayAuthorizationMixin:
return ""
return str(self._adapter_setting(platform, attr, extra_key, profile) or "").strip().lower()
def _adapter_dm_policy(
self,
platform: Optional[Platform],
*,
profile: Optional[str] = None,
) -> str:
def _adapter_dm_policy(self, platform: Optional[Platform], *, profile: Optional[str] = None) -> str:
"""Lowercased effective ``dm_policy`` (open/allowlist/disabled/pairing), ``""`` if unknown."""
return self._adapter_policy(platform, "_dm_policy", "dm_policy", profile)
def _adapter_group_policy(
self,
platform: Optional[Platform],
*,
profile: Optional[str] = None,
) -> str:
def _adapter_group_policy(self, platform: Optional[Platform], *, profile: Optional[str] = None) -> str:
"""Lowercased effective ``group_policy`` (open/allowlist/disabled), ``""`` if unknown."""
return self._adapter_policy(platform, "_group_policy", "group_policy", profile)
def _adapter_group_has_sender_allowlist(
self,
platform: Optional[Platform],
chat_id: Optional[str],
*,
profile: Optional[str] = None,
self, platform: Optional[Platform], chat_id: Optional[str], *, profile: Optional[str] = None
) -> bool:
"""Whether a per-group sender allowlist (WeCom ``groups.<id>.allow_from``) gated this message.
@@ -439,12 +403,7 @@ class GatewayAuthorizationMixin:
allowed = {normalize(entry) or entry for entry in allowed}
return _allows(allowed, user_id)
def _is_user_authorized(
self,
source: SessionSource,
*,
allow_adapter_delegation: bool = True,
) -> bool:
def _is_user_authorized(self, source: SessionSource, *, allow_adapter_delegation: bool = True) -> bool:
"""Whether a user may use the bot.
Order: trusted-upstream delegation, chat-scoped group allowlists,
@@ -479,9 +438,7 @@ class GatewayAuthorizationMixin:
# admins, sender_chat posts, channel broadcasts): run before the no-user-id guard.
if is_group and source.chat_id:
chat_allowlist_env = _GROUP_CHAT_ENV.get(source.platform, "")
if chat_allowlist_env and _allows(
_coerce_allow_set(_platform_gate_env(chat_allowlist_env)), source.chat_id
):
if chat_allowlist_env and _allows(_coerce_allow_set(_platform_gate_env(chat_allowlist_env)), source.chat_id):
return True
# config.yaml fallback (``extra.group_allowed_chats``): Telegram observe-
# unmentioned mode strips user_id, so the env-only check above misses it.
@@ -557,9 +514,7 @@ class GatewayAuthorizationMixin:
# Backward-compat: TELEGRAM_GROUP_ALLOWED_USERS was once (mis)used as a chat-ID
# allowlist; "-"-prefixed values are chat IDs, honor them and warn once.
if source.platform == Platform.TELEGRAM and group_user_allowlist and is_group_or_forum and source.chat_id:
legacy_chat_ids = {
v.strip() for v in group_user_allowlist.split(",") if v.strip().startswith("-")
}
legacy_chat_ids = {v.strip() for v in group_user_allowlist.split(",") if v.strip().startswith("-")}
if legacy_chat_ids:
if not getattr(self, "_warned_telegram_group_users_legacy", False):
logger.warning(
@@ -599,9 +554,7 @@ class GatewayAuthorizationMixin:
resolved_ids = None
if isinstance(resolved_ids, (set, frozenset, list, tuple)):
allowed_ids.update(
str(entry).strip()
for entry in resolved_ids
if isinstance(entry, (str, int)) and str(entry).strip()
str(entry).strip() for entry in resolved_ids if isinstance(entry, (str, int)) and str(entry).strip()
)
if "*" in allowed_ids:
@@ -613,11 +566,7 @@ class GatewayAuthorizationMixin:
# WhatsApp (Baileys + Cloud): phone<->LID / JID aliases match the same principal.
if source.platform in {Platform.WHATSAPP, Platform.WHATSAPP_CLOUD}:
normalized_allowed_ids = set()
for allowed_id in allowed_ids:
normalized_allowed_ids.update(_expand_whatsapp_auth_aliases(allowed_id))
if normalized_allowed_ids:
allowed_ids = normalized_allowed_ids
allowed_ids = set().union(*(_expand_whatsapp_auth_aliases(a) for a in allowed_ids)) or allowed_ids
check_ids.update(_expand_whatsapp_auth_aliases(user_id))
normalized_user_id = _normalize_whatsapp_identifier(user_id)
@@ -633,19 +582,13 @@ class GatewayAuthorizationMixin:
# Buzz: allowlist may hold npub or hex; inbound pubkeys are hex.
if platform_value == "buzz":
allowed_ids = _normalize_nostr_allow_entries(allowed_ids)
if user_id.startswith("npub"):
hex_user = _npub_to_hex(user_id)
if hex_user:
check_ids.add(hex_user)
hex_user = _npub_to_hex(user_id) if user_id.startswith("npub") else None
if hex_user:
check_ids.add(hex_user)
return bool(check_ids & allowed_ids)
def _get_unauthorized_dm_behavior(
self,
platform: Optional[Platform],
*,
profile: Optional[str] = None,
) -> str:
def _get_unauthorized_dm_behavior(self, platform: Optional[Platform], *, profile: Optional[str] = None) -> str:
"""How unauthorized DMs are handled ("pair" / "ignore") for a platform.
Order: explicit per-platform config; Email → "ignore" (inboxes hold arbitrary
@@ -667,6 +610,7 @@ class GatewayAuthorizationMixin:
if config and hasattr(config, "unauthorized_dm_behavior") and config.unauthorized_dm_behavior != "pair":
return config.unauthorized_dm_behavior
allowlist_keys = ["GATEWAY_ALLOWED_USERS"]
if platform:
dm_policy = self._adapter_dm_policy(platform, profile=profile)
if not dm_policy:
@@ -675,15 +619,9 @@ class GatewayAuthorizationMixin:
return "pair"
if dm_policy in {"allowlist", "disabled"}:
return "ignore"
# Historical: Yuanbao is absent from this allowlist-aware default.
env_key = "" if platform == Platform.YUANBAO else _ALLOWED_USERS_ENV.get(platform, "")
group_keys = (_GROUP_USER_ENV.get(platform), _GROUP_CHAT_ENV.get(platform))
for key in (env_key, *group_keys):
if key and _platform_gate_env(key).strip():
return "ignore"
if _platform_gate_env("GATEWAY_ALLOWED_USERS").strip():
allowlist_keys = [env_key, _GROUP_USER_ENV.get(platform), _GROUP_CHAT_ENV.get(platform), *allowlist_keys]
if any(key and _platform_gate_env(key).strip() for key in allowlist_keys):
return "ignore"
return "pair"

View File

@@ -5,6 +5,7 @@ send_message reads it for action="list" and to resolve friendly channel names to
"""
import asyncio
import contextlib
import json
import logging
import time
@@ -75,12 +76,10 @@ def _apply_channel_aliases(platforms: Dict[str, Any]) -> None:
continue
chat_id = str(chat_id)
friendly = friendly.strip()
matched = False
for e in entries:
if isinstance(e, dict) and e.get("id") == chat_id:
e["name"] = friendly
matched = True
if not matched:
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,
@@ -218,11 +217,8 @@ def _build_discord(adapter) -> List[Dict[str, str]]:
def _slack_api_error_code(error: Exception) -> Optional[str]:
"""Slack Web API error code from SlackApiError-like exceptions."""
response = getattr(error, "response", None)
if response is None:
return None
try:
value = response.get("error")
value = getattr(error, "response").get("error")
except Exception:
return None
return str(value) if value else None
@@ -232,9 +228,7 @@ def _normalize_adapter_channels(raw_channels: Any) -> List[Dict[str, Any]]:
"""Validate and dedupe entries returned by an adapter's ``list_channels()`` hook."""
channels: List[Dict[str, Any]] = []
seen_ids = set()
if not isinstance(raw_channels, list):
return channels
for raw in raw_channels:
for raw in raw_channels if isinstance(raw_channels, list) else ():
if not isinstance(raw, dict):
continue
channel_id = str(raw.get("id") or "").strip()
@@ -276,8 +270,7 @@ async def _slack_team_channels(team_id: str, client, seen_ids: set) -> List[Dict
_report_slack_failure(team_id, error_code, f"users.conversations not ok: {error_code}")
break
for ch in response.get("channels", []):
cid = ch.get("id")
name = ch.get("name")
cid, name = ch.get("id"), ch.get("name")
if not cid or not name or cid in seen_ids:
continue
seen_ids.add(cid)
@@ -306,23 +299,19 @@ async def _slack_resolve_raw_names(client, channels: List[Dict[str, Any]]) -> No
if not resp.get("ok"):
return
ch_info = resp.get("channel", {})
resolved_name = None
resolved_type = None
if ch_info.get("is_im"):
peer_user = ch_info.get("user", "")
if peer_user:
user_resp = await client.users_info(user=peer_user)
if user_resp.get("ok"):
u = user_resp["user"]
resolved_name = u.get("profile", {}).get("display_name") or u.get("real_name") or u.get("name")
resolved_type = "dm"
else:
resolved_name = resolved_type = None
if not ch_info.get("is_im"):
resolved_name = ch_info.get("name") or ch_info.get("name_normalized")
if resolved_name:
for entry in entries:
entry["name"] = resolved_name
if resolved_type:
entry["type"] = resolved_type
elif ch_info.get("user", ""):
user_resp = await client.users_info(user=ch_info["user"])
if user_resp.get("ok"):
u = user_resp["user"]
resolved_name = u.get("profile", {}).get("display_name") or u.get("real_name") or u.get("name")
resolved_type = "dm"
for entry in entries if resolved_name else ():
entry["name"] = resolved_name
if resolved_type:
entry["type"] = resolved_type
except Exception as e:
logger.debug("Channel directory: failed to resolve %s: %s", base_id, e)
@@ -348,10 +337,8 @@ async def _build_slack(adapter) -> List[Dict[str, Any]]:
eid = entry.get("id")
if not isinstance(eid, str) or eid in seen_ids:
continue
if _slack_has_raw_name(entry):
base_id = _slack_base_id(eid)
if base_id in api_name_lookup:
entry["name"] = api_name_lookup[base_id]
if _slack_has_raw_name(entry) and _slack_base_id(eid) in api_name_lookup:
entry["name"] = api_name_lookup[_slack_base_id(eid)]
channels.append(entry)
seen_ids.add(eid)
@@ -404,14 +391,9 @@ def _build_from_sessions_db(platform_name: str) -> List[Dict[str, str]]:
release_or_close(db)
for row in rows:
origin = None
if row.get("origin_json"):
try:
parsed = json.loads(row["origin_json"])
if isinstance(parsed, dict) and parsed:
origin = parsed
except (TypeError, ValueError):
pass
if origin is None:
with contextlib.suppress(TypeError, ValueError):
origin = json.loads(row["origin_json"]) if row.get("origin_json") else None
if not isinstance(origin, dict) or not origin:
origin = {"chat_id": row.get("chat_id"), "thread_id": row.get("thread_id"), "chat_name": row.get("display_name")}
yield origin, row.get("chat_type") or "dm"
@@ -444,14 +426,12 @@ def load_directory() -> Dict[str, Any]:
"""Load the cached directory from disk, with aliases re-applied on read."""
directory_path = _directory_path()
if directory_path.exists():
try:
with contextlib.suppress(Exception):
with open(directory_path, encoding="utf-8") as f:
data = json.load(f)
# Aliases apply on read too, so new names take effect between timed rebuilds.
_apply_channel_aliases(data.setdefault("platforms", {}))
return data
except Exception:
pass
base = {"updated_at": None, "platforms": {}}
_apply_channel_aliases(base["platforms"])
return base

View File

@@ -176,10 +176,7 @@ def _write_allowlist_env(env_var: str, ids: list) -> None:
with contextlib.suppress(Exception):
from hermes_cli.config import save_env_value, remove_env_value
if ids:
save_env_value(env_var, ",".join(ids))
else:
remove_env_value(env_var)
save_env_value(env_var, ",".join(ids)) if ids else remove_env_value(env_var)
def _sync_allowlist_add(platform: str, user_id: str) -> None:
@@ -200,7 +197,7 @@ def _iter_live_gateway_adapters():
runner = _gateway_runner_ref()
except Exception:
return
runner = None
if runner is None:
return
mappings = [getattr(runner, "adapters", None) or {}]
@@ -221,15 +218,14 @@ def _adapter_platform_name(adapter) -> str:
def _purge_allowlist_entries(entries, platform: str, user_id: str):
"""Drop alias-equivalent allowlist entries while preserving ``*``."""
def keep(entry) -> bool:
entry = str(entry)
return entry.strip() == "*" or not _user_ids_match(platform, entry, str(user_id))
return str(entry).strip() == "*" or not _user_ids_match(platform, str(entry), str(user_id))
if isinstance(entries, str):
return ",".join(part for part in _split_allowlist(entries) if keep(part))
return ",".join(filter(keep, _split_allowlist(entries)))
if isinstance(entries, (set, frozenset)):
return {entry for entry in entries if keep(entry)}
return set(filter(keep, entries))
if isinstance(entries, (list, tuple)):
return [entry for entry in entries if keep(entry)]
return list(filter(keep, entries))
return entries
@@ -457,9 +453,7 @@ class PairingStore:
def _hash_code(code: str, salt: bytes) -> str:
return hashlib.sha256(salt + code.encode("utf-8")).hexdigest()
def _finish_approval(
self, platform: str, pending: dict, matched_key: str, matched_entry: dict
) -> dict:
def _finish_approval(self, platform: str, pending: dict, matched_key: str, matched_entry: dict) -> dict:
"""Remove a pending request and approve its user. Must hold self._lock."""
del pending[matched_key]
self._save_json(self._pending_path(platform), pending)
@@ -468,18 +462,11 @@ class PairingStore:
# must not carry over (isolated typos would accumulate into a spurious lockout).
self._reset_failed_attempts(platform)
self._approve_user(
platform, matched_entry["user_id"], matched_entry.get("user_name", "")
)
result = {"user_id": matched_entry["user_id"], "user_name": matched_entry.get("user_name", "")}
self._approve_user(platform, result["user_id"], result["user_name"])
return result
return {
"user_id": matched_entry["user_id"],
"user_name": matched_entry.get("user_name", ""),
}
def generate_code(
self, platform: str, user_id: str, user_name: str = ""
) -> Optional[str]:
def generate_code(self, platform: str, user_id: str, user_name: str = "") -> Optional[str]:
"""Generate a pairing code for a new user.
Returns None if the user is rate-limited, the platform hit
@@ -628,7 +615,7 @@ class PairingStore:
limits[fail_key] = fails
if fails >= MAX_FAILED_ATTEMPTS:
limits[f"_lockout:{platform}"] = time.time() + LOCKOUT_SECONDS
limits[fail_key] = 0 # Reset counter
limits[fail_key] = 0
print(f"[pairing] Platform {platform} locked out for {LOCKOUT_SECONDS}s "
f"after {MAX_FAILED_ATTEMPTS} failed attempts", flush=True)
self._save_json(self._rate_limit_path(), limits)
@@ -655,9 +642,7 @@ class PairingStore:
or (now - info["created_at"]) > CODE_TTL_SECONDS
]
if expired:
for entry_id in expired:
del pending[entry_id]
self._save_json(path, pending)
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 ``-<suffix>.json`` data file (``_``-prefixed files are shared state)."""