ForcePilot/backend/package/yuxi/channel/gateway/net_utils.py
Kris ecd3c90e80 feat(channel/gateway): 新增完整网关通道模块
新增设备身份管理、认证限流、并发通道、Webhook路由、RBAC权限控制、SSE/轮询降级等全套网关通道功能,包含:
1. 设备身份生成与签名验证
2. 设备令牌认证与速率限制
3. 内存+数据库双重设备注册表
4. 并发通道限流管理
5. Webhook安全处理与路由
6. RBAC权限校验系统
7. OpenAI API兼容适配层
8. Tailscale认证支持
9. HTTP轮询降级机制
2026-05-21 10:26:33 +08:00

374 lines
10 KiB
Python

import asyncio
import ipaddress
import logging
import os
import socket
from enum import StrEnum
from urllib.parse import urlparse
logger = logging.getLogger(__name__)
_LOOPBACK_HOSTS = frozenset(["localhost", "127.0.0.1", "::1"])
_PRIVATE_IP_NETS = [
ipaddress.ip_network("10.0.0.0/8"),
ipaddress.ip_network("172.16.0.0/12"),
ipaddress.ip_network("192.168.0.0/16"),
ipaddress.ip_network("127.0.0.0/8"),
ipaddress.ip_network("::1/128"),
ipaddress.ip_network("fc00::/7"),
ipaddress.ip_network("fe80::/10"),
]
_TAILNET_IPV4_NET = ipaddress.ip_network("100.64.0.0/10")
_ENV_ALLOW_INSECURE_PRIVATE_WS = "YUXI_ALLOW_INSECURE_PRIVATE_WS"
class GatewayBindMode(StrEnum):
LOOPBACK = "loopback"
LAN = "lan"
TAILNET = "tailnet"
AUTO = "auto"
CUSTOM = "custom"
def is_loopback_address(ip: str | None) -> bool:
if not ip:
return False
try:
return ipaddress.ip_address(ip.strip()).is_loopback
except ValueError:
return False
def is_trusted_proxy_address(ip: str | None, trusted_proxies: list[str] | None) -> bool:
if not ip or not trusted_proxies:
return False
try:
addr = ipaddress.ip_address(ip.strip())
except ValueError:
return False
for proxy in trusted_proxies:
candidate = proxy.strip()
if not candidate:
continue
try:
net = ipaddress.ip_network(candidate, strict=False)
except ValueError:
continue
if addr in net:
return True
return False
TAILSCALE_TRUSTED_PROXIES = ["127.0.0.1", "::1"]
def _parse_ip_literal(raw: str | None) -> str | None:
if not raw:
return None
trimmed = raw.strip()
if not trimmed:
return None
if trimmed.startswith("[") and "]" in trimmed:
trimmed = trimmed[1 : trimmed.index("]")]
if ":" in trimmed and "." in trimmed:
last_colon = trimmed.rfind(":")
candidate = trimmed[:last_colon]
try:
ipaddress.IPv4Address(candidate)
trimmed = candidate
except ValueError:
pass
try:
ipaddress.ip_address(trimmed)
return trimmed
except ValueError:
return None
def resolve_forwarded_client_ip(
forwarded_for: str | None,
trusted_proxies: list[str] | None,
) -> str | None:
if not trusted_proxies:
return None
chain: list[str] = []
for entry in (forwarded_for or "").split(","):
normalized = _parse_ip_literal(entry)
if normalized:
chain.append(normalized)
if not chain:
return None
for hop in reversed(chain):
if is_loopback_address(hop):
continue
if not is_trusted_proxy_address(hop, trusted_proxies):
return hop
return None
def resolve_client_ip(
remote_addr: str | None,
forwarded_for: str | None = None,
real_ip: str | None = None,
trusted_proxies: list[str] | None = None,
allow_real_ip_fallback: bool = False,
) -> str | None:
remote = _parse_ip_literal(remote_addr)
if not remote:
return None
if not is_trusted_proxy_address(remote, trusted_proxies):
return remote
forwarded = resolve_forwarded_client_ip(forwarded_for, trusted_proxies)
if forwarded:
return forwarded
if allow_real_ip_fallback:
return _parse_ip_literal(real_ip)
return None
def has_forwarded_request_headers(headers: dict) -> bool:
return bool(
headers.get("forwarded")
or headers.get("x-forwarded-for")
or headers.get("x-forwarded-proto")
or headers.get("x-real-ip")
or headers.get("x-forwarded-host")
)
def is_local_direct_request(
remote_addr: str | None,
headers: dict | None = None,
) -> bool:
if not remote_addr:
return False
if headers and has_forwarded_request_headers(headers):
return False
return is_loopback_address(remote_addr)
def has_tailscale_proxy_headers(headers: dict | None) -> bool:
if not headers:
return False
return bool(headers.get("x-forwarded-for") and headers.get("x-forwarded-proto") and headers.get("x-forwarded-host"))
def is_tailscale_proxy_request(
remote_addr: str | None,
headers: dict | None = None,
) -> bool:
if not remote_addr:
return False
return is_loopback_address(remote_addr) and has_tailscale_proxy_headers(headers)
def resolve_tailscale_client_ip(
remote_addr: str | None,
headers: dict | None = None,
) -> str | None:
return resolve_client_ip(
remote_addr=remote_addr,
forwarded_for=headers.get("x-forwarded-for") if headers else None,
trusted_proxies=list(TAILSCALE_TRUSTED_PROXIES),
)
def is_loopback_host(host: str) -> bool:
host = host.strip().lower().rstrip(".")
if not host:
return False
if host in _LOOPBACK_HOSTS:
return True
try:
addr = ipaddress.ip_address(host)
except ValueError:
return False
return addr.is_loopback
def is_private_host(host: str) -> bool:
host = host.strip().lower().rstrip(".")
if not host:
return False
try:
addr = ipaddress.ip_address(host)
except ValueError:
return False
return addr.is_private
def is_private_or_loopback_host(host: str) -> bool:
host = host.strip().lower().rstrip(".")
if not host:
return False
try:
addr = ipaddress.ip_address(host)
except ValueError:
return False
return addr.is_private or addr.is_loopback or addr.is_link_local
def is_localish_host(host: str | None) -> bool:
if not host:
return False
host = host.strip().lower().rstrip(".")
return is_loopback_host(host) or host.endswith(".ts.net")
def is_secure_ws_url(url: str, allow_private_ws: bool = False) -> bool:
try:
parsed = urlparse(url)
except ValueError:
return False
protocol = parsed.scheme.lower()
if protocol == "wss":
return True
if protocol not in ("ws", "http"):
return False
hostname = _extract_ws_hostname(parsed)
if not hostname:
return False
if is_loopback_host(hostname):
return True
if allow_private_ws and is_private_or_loopback_host(hostname):
return True
return False
def _extract_ws_hostname(parsed) -> str:
netloc = parsed.netloc or parsed.hostname or ""
if "@" in netloc:
netloc = netloc.rsplit("@", 1)[-1]
if "[" in netloc and "]" in netloc:
start = netloc.index("[") + 1
end = netloc.index("]")
return netloc[start:end]
if ":" in netloc:
return netloc.rsplit(":", 1)[0]
return netloc
async def can_bind_to_host(host: str) -> bool:
loop = asyncio.get_running_loop()
return await loop.run_in_executor(None, _sync_can_bind_to_host, host)
def _sync_can_bind_to_host(host: str) -> bool:
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
try:
sock.bind((host, 0))
return True
except OSError:
return False
finally:
sock.close()
_can_bind_cache: dict[str, bool] = {}
async def can_bind_to_host_cached(host: str) -> bool:
cached = _can_bind_cache.get(host)
if cached is not None:
return cached
result = await can_bind_to_host(host)
_can_bind_cache[host] = result
return result
def _is_container_environment() -> bool:
if os.path.exists("/.dockerenv"):
return True
try:
with open("/proc/1/cgroup") as f:
content = f.read()
if "docker" in content or "kubepods" in content:
return True
except OSError:
pass
return False
def _pick_primary_tailnet_ipv4() -> str | None:
try:
for info in socket.getaddrinfo(socket.gethostname(), None, socket.AF_INET):
ip = info[4][0]
try:
addr = ipaddress.ip_address(ip)
except ValueError:
continue
if addr in _TAILNET_IPV4_NET:
return ip
except OSError:
pass
return None
async def resolve_gateway_bind_host(
mode: GatewayBindMode | None = None,
custom_host: str | None = None,
) -> str:
mode = mode or GatewayBindMode.LOOPBACK
if mode == GatewayBindMode.LOOPBACK:
if await can_bind_to_host("127.0.0.1"):
return "127.0.0.1"
return "0.0.0.0"
if mode == GatewayBindMode.TAILNET:
tailnet_ip = _pick_primary_tailnet_ipv4()
if tailnet_ip and await can_bind_to_host(tailnet_ip):
return tailnet_ip
if await can_bind_to_host("127.0.0.1"):
return "127.0.0.1"
return "0.0.0.0"
if mode == GatewayBindMode.LAN:
return "0.0.0.0"
if mode == GatewayBindMode.CUSTOM:
host = (custom_host or "").strip()
if not host:
return "0.0.0.0"
try:
ipaddress.ip_address(host)
except ValueError:
logger.warning("gateway bind=custom: invalid IP '%s', falling back to 0.0.0.0", host)
return "0.0.0.0"
if await can_bind_to_host(host):
return host
logger.warning("gateway bind=custom: cannot bind '%s', falling back to 0.0.0.0", host)
return "0.0.0.0"
if mode == GatewayBindMode.AUTO:
if _is_container_environment():
return "0.0.0.0"
if await can_bind_to_host("127.0.0.1"):
return "127.0.0.1"
return "0.0.0.0"
return "0.0.0.0"
def build_ws_security_error(display_host: str) -> str:
allow_private = os.environ.get(_ENV_ALLOW_INSECURE_PRIVATE_WS) == "1"
msg = (
f'SECURITY ERROR: Cannot connect to "{display_host}" over plaintext ws://. '
"Both credentials and chat data would be exposed to network interception. "
"Use wss:// for remote URLs. Safe defaults: keep gateway.bind=loopback and "
"connect via SSH tunnel "
"(ssh -N -L 18789:127.0.0.1:18789 user@gateway-host), or use Tailscale Serve/Funnel."
)
if not allow_private:
msg += f" Break-glass (trusted private networks only): set {_ENV_ALLOW_INSECURE_PRIVATE_WS}=1."
return msg