"""Discover and select server egress channels for account-bound IM traffic. One public address may be reached through a private address on the host (for example, an ECS secondary private IP mapped to an EIP). A channel therefore keeps both values: ``source_ip`` is bound on the socket and ``public_ip`` is what the remote service observes. """ from __future__ import annotations import asyncio import ipaddress import json import logging import os import socket import subprocess import threading import time from dataclasses import dataclass from typing import Iterable import httpx import requests from requests.adapters import HTTPAdapter logger = logging.getLogger("rpa_engine.egress") _DISCOVERY_TTL_SECONDS = 300.0 _PROBE_TIMEOUT_SECONDS = 6.0 _MAX_CHANNEL_ATTEMPTS = 8 _PROBE_URLS = ( "https://www.cloudflare.com/cdn-cgi/trace", "https://api64.ipify.org?format=json", ) @dataclass(frozen=True) class LocalAddress: source_ip: str | None interface: str is_default: bool = False @dataclass(frozen=True) class EgressChannel: public_ip: str source_ip: str | None interface: str = "" is_default: bool = False @property def id(self) -> str: return self.public_ip @dataclass(frozen=True) class EgressSnapshot: channels: tuple[EgressChannel, ...] errors: tuple[str, ...] detected_at: float class EgressChannelUnavailable(RuntimeError): pass _cache_lock = threading.Lock() _cached_snapshot: EgressSnapshot | None = None def clamp_attempts(value: int | None) -> int: try: parsed = int(value or 1) except (TypeError, ValueError): parsed = 1 return max(1, min(_MAX_CHANNEL_ATTEMPTS, parsed)) def _usable_source_ip(value: str) -> bool: try: addr = ipaddress.ip_address(str(value or "").strip()) except ValueError: return False return bool( addr.version == 4 and not addr.is_loopback and not addr.is_link_local and not addr.is_multicast and not addr.is_unspecified ) def _linux_local_addresses() -> list[LocalAddress]: if os.name != "posix": return [] try: proc = subprocess.run( ["ip", "-j", "-4", "addr", "show", "scope", "global"], capture_output=True, text=True, timeout=3, check=False, ) payload = json.loads(proc.stdout or "[]") if proc.returncode == 0 else [] except (OSError, subprocess.SubprocessError, json.JSONDecodeError): return [] found: list[LocalAddress] = [] for item in payload if isinstance(payload, list) else []: interface = str(item.get("ifname") or "") for info in item.get("addr_info") or []: source_ip = str(info.get("local") or "").strip() if _usable_source_ip(source_ip): found.append(LocalAddress(source_ip, interface)) return found def _socket_local_addresses() -> list[LocalAddress]: found: list[LocalAddress] = [] names = {socket.gethostname(), socket.getfqdn()} for name in names: try: records = socket.getaddrinfo(name, None, socket.AF_INET, socket.SOCK_STREAM) except OSError: continue for record in records: source_ip = str(record[4][0] or "").strip() if _usable_source_ip(source_ip): found.append(LocalAddress(source_ip, name)) return found def local_address_candidates() -> list[LocalAddress]: """Return the default route plus each bindable global/private IPv4.""" candidates = [LocalAddress(None, "default", True)] seen: set[str] = set() for item in [*_linux_local_addresses(), *_socket_local_addresses()]: source_ip = str(item.source_ip or "") if not source_ip or source_ip in seen: continue seen.add(source_ip) candidates.append(item) return candidates def _extract_public_ip(response: httpx.Response) -> str: text = response.text.strip() content_type = response.headers.get("content-type", "").lower() candidate = "" if "json" in content_type or text.startswith("{"): try: candidate = str(response.json().get("ip") or "").strip() except (ValueError, AttributeError): candidate = "" if not candidate: for line in text.splitlines(): if line.startswith("ip="): candidate = line.partition("=")[2].strip() break if not candidate and "\n" not in text and len(text) <= 64: candidate = text try: addr = ipaddress.ip_address(candidate) except ValueError: return "" return str(addr) if addr.version == 4 else "" async def _probe_local_address(candidate: LocalAddress) -> tuple[EgressChannel | None, str]: transport = httpx.AsyncHTTPTransport( local_address=candidate.source_ip, retries=0, ) last_error = "" try: async with httpx.AsyncClient( transport=transport, timeout=httpx.Timeout(_PROBE_TIMEOUT_SECONDS), follow_redirects=True, trust_env=False, ) as client: for url in _PROBE_URLS: try: response = await client.get(url, headers={"Accept": "text/plain, application/json"}) response.raise_for_status() public_ip = _extract_public_ip(response) if public_ip: return ( EgressChannel( public_ip=public_ip, source_ip=candidate.source_ip, interface=candidate.interface, is_default=candidate.is_default, ), "", ) last_error = "探测响应中没有 IPv4" except Exception as exc: # one endpoint may be unavailable last_error = str(exc) or type(exc).__name__ finally: await transport.aclose() label = candidate.source_ip or "默认路由" return None, f"{label}: {last_error or '无法访问公网探测服务'}" def _dedupe_channels(channels: Iterable[EgressChannel]) -> tuple[EgressChannel, ...]: by_public_ip: dict[str, EgressChannel] = {} order: list[str] = [] for channel in channels: existing = by_public_ip.get(channel.public_ip) if existing is None: by_public_ip[channel.public_ip] = channel order.append(channel.public_ip) continue # Keep an explicit bindable source when possible, while preserving the # fact that this is also the host's default public route. if existing.source_ip is None and channel.source_ip: by_public_ip[channel.public_ip] = EgressChannel( public_ip=channel.public_ip, source_ip=channel.source_ip, interface=channel.interface, is_default=existing.is_default or channel.is_default, ) elif channel.is_default and not existing.is_default: by_public_ip[channel.public_ip] = EgressChannel( public_ip=existing.public_ip, source_ip=existing.source_ip, interface=existing.interface, is_default=True, ) return tuple(by_public_ip[key] for key in order) async def discover_egress_channels(*, force: bool = False) -> EgressSnapshot: global _cached_snapshot now = time.time() with _cache_lock: cached = _cached_snapshot if not force and cached and now - cached.detected_at < _DISCOVERY_TTL_SECONDS: return cached candidates = await asyncio.to_thread(local_address_candidates) results = await asyncio.gather(*(_probe_local_address(item) for item in candidates)) channels = _dedupe_channels(item[0] for item in results if item[0] is not None) errors = tuple(item[1] for item in results if item[1]) snapshot = EgressSnapshot(channels=channels, errors=errors, detected_at=time.time()) with _cache_lock: _cached_snapshot = snapshot return snapshot async def resolve_fixed_channel(public_ip: str) -> EgressChannel: selected = str(public_ip or "").strip() if not selected: return EgressChannel(public_ip="", source_ip=None, interface="default", is_default=True) snapshot = await discover_egress_channels() for channel in snapshot.channels: if channel.public_ip == selected: return channel raise EgressChannelUnavailable( f"指定公网通道 {selected} 当前不可用;请在账号编辑中重新检测并选择可用通道" ) async def resolve_send_channels( preferred_public_ip: str = "", max_attempts: int = 1, ) -> list[EgressChannel]: """Order channels for one serial send operation. The ordinary one-channel automatic mode deliberately avoids discovery so a temporary outage of the probe service never blocks existing sends. """ preferred = str(preferred_public_ip or "").strip() attempts = clamp_attempts(max_attempts) if not preferred and attempts == 1: return [EgressChannel(public_ip="", source_ip=None, interface="default", is_default=True)] snapshot = await discover_egress_channels() channels = list(snapshot.channels) if not channels: if preferred: raise EgressChannelUnavailable( f"指定公网通道 {preferred} 无法探测;请检查服务器网卡、路由或公网访问" ) return [EgressChannel(public_ip="", source_ip=None, interface="default", is_default=True)] ordered: list[EgressChannel] = [] if preferred: selected = next((item for item in channels if item.public_ip == preferred), None) if selected is None: raise EgressChannelUnavailable( f"指定公网通道 {preferred} 当前不可用;请在账号编辑中重新检测" ) ordered.append(selected) else: default = next((item for item in channels if item.is_default), None) if default is not None: ordered.append(default) ordered.extend(item for item in channels if item not in ordered) return ordered[:attempts] class _SourceAddressAdapter(HTTPAdapter): """Requests adapter that binds outgoing sockets to one local IPv4.""" def __init__(self, source_ip: str, *args, **kwargs): self._source_address = (source_ip, 0) super().__init__(*args, **kwargs) def init_poolmanager(self, connections, maxsize, block=False, **pool_kwargs): pool_kwargs["source_address"] = self._source_address return super().init_poolmanager(connections, maxsize, block=block, **pool_kwargs) def proxy_manager_for(self, proxy, **proxy_kwargs): proxy_kwargs["source_address"] = self._source_address return super().proxy_manager_for(proxy, **proxy_kwargs) def source_bound_requests_session(source_ip: str | None = None) -> requests.Session: client = requests.Session() source = str(source_ip or "").strip() if source: client.trust_env = False adapter = _SourceAddressAdapter(source) client.mount("http://", adapter) client.mount("https://", adapter) return client def reset_egress_cache_for_tests() -> None: global _cached_snapshot with _cache_lock: _cached_snapshot = None