376 lines
15 KiB
Python
376 lines
15 KiB
Python
"""Gateway outbound agent: dial the platform, dispatch tasks, push events.
|
||
|
||
方向反转后网关只有一条自己发起的 WS 连接(/v1/agent,连接层 Bearer key 认证):
|
||
- 心跳线程周期上报 {type: heartbeat, version},平台据此维护 online/last_seen_at;
|
||
- 读循环接收 task(线程池执行,复用 HTTP 路由语义)与 event_ack/event_nack;
|
||
- result 帧回传路由结果;重复 task id 幂等拒绝,不重复执行有副作用的操作;
|
||
- 事件桥把 SubscriptionManager 的 pending deliveries 推为 event 帧,
|
||
平台事务保存成功后回 event_ack;未确认批次在 nack 或断线重连后重推。
|
||
协议与平台侧 internal/controlplane/api/gateway_channel.go 一一对应。
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import base64
|
||
import json
|
||
import logging
|
||
import threading
|
||
import time
|
||
from concurrent.futures import ThreadPoolExecutor
|
||
from dataclasses import dataclass
|
||
from typing import Any, Callable
|
||
|
||
import websocket
|
||
|
||
LOG = logging.getLogger("creatorhub.gateway")
|
||
|
||
GW_VERSION = "1.0.0-outbound"
|
||
HEARTBEAT_INTERVAL = 20.0
|
||
RECONNECT_BASE = 1.0
|
||
RECONNECT_MAX = 30.0
|
||
WORKER_COUNT = 4
|
||
IDEMPOTENT_STATUS = 409
|
||
|
||
|
||
@dataclass
|
||
class AgentConfig:
|
||
platform_url: str
|
||
access_key: str
|
||
version: str = GW_VERSION
|
||
heartbeat_interval: float = HEARTBEAT_INTERVAL
|
||
reconnect_base: float = RECONNECT_BASE
|
||
reconnect_max: float = RECONNECT_MAX
|
||
worker_count: int = WORKER_COUNT
|
||
|
||
|
||
class OutboundAgent:
|
||
"""单网关出站连接的完整生命周期;run() 阻塞运行,stop() 置位退出。"""
|
||
|
||
def __init__(self, gateway: Any, config: AgentConfig) -> None:
|
||
self.gateway = gateway
|
||
self.config = config
|
||
self.connected = threading.Event()
|
||
self.last_error = ""
|
||
self.reconnect_attempts = 0
|
||
self._stop_requested = threading.Event()
|
||
self._write_lock = threading.Lock()
|
||
self._conn: websocket.WebSocket | None = None
|
||
self._sent_events: dict[str, dict[str, float]] = {}
|
||
self._sent_lock = threading.Lock()
|
||
self._seen_tasks: set[int] = set()
|
||
self._seen_tasks_lock = threading.Lock()
|
||
self._push_wakeup = threading.Event()
|
||
|
||
# ---- 生命周期 -------------------------------------------------------
|
||
|
||
def run(self) -> None:
|
||
backoff = self.config.reconnect_base
|
||
while not self._stop_requested.is_set():
|
||
try:
|
||
self.reconnect_attempts += 1
|
||
conn = self._connect()
|
||
except Exception as exc:
|
||
self.last_error = str(exc) or exc.__class__.__name__
|
||
LOG.warning("platform connection failed: %s", self.last_error)
|
||
if self._wait(backoff):
|
||
break
|
||
backoff = min(backoff * 2, self.config.reconnect_max)
|
||
continue
|
||
self.reconnect_attempts = 0
|
||
self.last_error = ""
|
||
self._serve(conn)
|
||
backoff = self.config.reconnect_base
|
||
if not self._stop_requested.is_set():
|
||
LOG.info("platform connection lost; reconnecting")
|
||
|
||
def stop(self) -> None:
|
||
self._stop_requested.set()
|
||
with self._write_lock:
|
||
conn = self._conn
|
||
if conn is not None:
|
||
try:
|
||
conn.close()
|
||
except Exception:
|
||
pass
|
||
|
||
def _wait(self, seconds: float) -> bool:
|
||
return self._stop_requested.wait(seconds)
|
||
|
||
def _connect(self) -> websocket.WebSocket:
|
||
url = self.config.platform_url.rstrip("/") + "/v1/agent"
|
||
if url.startswith("https://"):
|
||
url = "wss://" + url[len("https://"):]
|
||
elif url.startswith("http://"):
|
||
url = "ws://" + url[len("http://"):]
|
||
conn = websocket.create_connection(
|
||
url,
|
||
header=[f"Authorization: Bearer {self.config.access_key}"],
|
||
timeout=10,
|
||
)
|
||
# timeout 仅用于握手;读循环必须阻塞,否则空闲连接会在 recv 超时后假性断线。
|
||
conn.settimeout(None)
|
||
with self._write_lock:
|
||
self._conn = conn
|
||
return conn
|
||
|
||
# ---- 会话线程 -------------------------------------------------------
|
||
|
||
def _serve(self, conn: websocket.WebSocket) -> None:
|
||
self.connected.set()
|
||
with self._sent_lock:
|
||
self._sent_events.clear()
|
||
executor = ThreadPoolExecutor(max_workers=self.config.worker_count, thread_name_prefix="gw-task")
|
||
heartbeat = threading.Thread(target=self._heartbeat_loop, args=(conn,), daemon=True)
|
||
sender = threading.Thread(target=self._event_sender_loop, args=(conn,), daemon=True)
|
||
heartbeat.start()
|
||
sender.start()
|
||
failure = ""
|
||
try:
|
||
while not self._stop_requested.is_set():
|
||
raw = conn.recv()
|
||
if not isinstance(raw, str):
|
||
continue
|
||
try:
|
||
frame = json.loads(raw)
|
||
except ValueError:
|
||
failure = "gateway channel message is not JSON"
|
||
break
|
||
if not isinstance(frame, dict):
|
||
failure = "gateway channel message invalid"
|
||
break
|
||
kind = frame.get("type")
|
||
if kind == "task":
|
||
future = executor.submit(self._handle_task, conn, frame)
|
||
future.add_done_callback(self._log_task_failure)
|
||
elif kind == "event_ack":
|
||
self._handle_event_ack(frame)
|
||
elif kind == "event_nack":
|
||
self._handle_event_nack(frame)
|
||
else:
|
||
failure = f"gateway channel message type invalid: {kind!r}"
|
||
break
|
||
except Exception as exc:
|
||
if not self._stop_requested.is_set():
|
||
failure = str(exc) or exc.__class__.__name__
|
||
LOG.exception("gateway channel receive loop failed: %s", failure)
|
||
finally:
|
||
self.connected.clear()
|
||
executor.shutdown(wait=False, cancel_futures=True)
|
||
self._stop_threads((heartbeat, sender))
|
||
with self._write_lock:
|
||
if self._conn is conn:
|
||
self._conn = None
|
||
try:
|
||
conn.close()
|
||
except Exception:
|
||
pass
|
||
if failure:
|
||
self.last_error = failure
|
||
|
||
@staticmethod
|
||
def _stop_threads(threads: tuple[threading.Thread, ...]) -> None:
|
||
for thread in threads:
|
||
thread.join(timeout=2)
|
||
|
||
def _heartbeat_loop(self, conn: websocket.WebSocket) -> None:
|
||
while not self._stop_requested.wait(self.config.heartbeat_interval):
|
||
try:
|
||
self._send(conn, {"type": "heartbeat", "version": self.config.version})
|
||
except Exception:
|
||
try:
|
||
conn.close()
|
||
except Exception:
|
||
pass
|
||
return
|
||
|
||
# ---- 帧读写 ---------------------------------------------------------
|
||
|
||
def _send(self, conn: websocket.WebSocket, value: dict) -> None:
|
||
payload = json.dumps(value, ensure_ascii=False, separators=(",", ":"))
|
||
with self._write_lock:
|
||
conn.send(payload)
|
||
|
||
# ---- 任务分发 -------------------------------------------------------
|
||
|
||
def _handle_task(self, conn: websocket.WebSocket, frame: dict) -> None:
|
||
task_id = frame.get("id")
|
||
method = frame.get("method")
|
||
path = frame.get("path")
|
||
if not isinstance(task_id, int) or not isinstance(method, str) or not isinstance(path, str):
|
||
self.last_error = "gateway task frame invalid"
|
||
try:
|
||
conn.close()
|
||
except Exception:
|
||
pass
|
||
return
|
||
with self._seen_tasks_lock:
|
||
duplicate = task_id in self._seen_tasks
|
||
self._seen_tasks.add(task_id)
|
||
if duplicate:
|
||
self._send(conn, {
|
||
"type": "result",
|
||
"id": task_id,
|
||
"status": IDEMPOTENT_STATUS,
|
||
"body": "",
|
||
"error": "duplicate task id; previous execution retained",
|
||
})
|
||
return
|
||
status, body, error = self._execute_task(method, path, frame.get("payload"))
|
||
self._send(conn, {
|
||
"type": "result",
|
||
"id": task_id,
|
||
"status": status,
|
||
"body": base64.b64encode(body).decode("ascii"),
|
||
"error": error,
|
||
})
|
||
|
||
def _execute_task(self, method: str, path: str, payload: Any) -> tuple[int, bytes, str]:
|
||
from .server.http import RequestError, route_gateway_request
|
||
|
||
if payload is None:
|
||
body: dict = {}
|
||
elif isinstance(payload, dict):
|
||
body = payload
|
||
else:
|
||
return 400, b"", "task payload must be a JSON object"
|
||
try:
|
||
result = route_gateway_request(self.gateway, method, path, {}, body)
|
||
except RequestError as exc:
|
||
return exc.status, b"", str(exc)
|
||
except FileNotFoundError as exc:
|
||
return 404, b"", str(exc)
|
||
except ValueError as exc:
|
||
return 400, b"", str(exc)
|
||
except Exception as exc: # 路由内部故障必须回传给平台,不能吞掉
|
||
LOG.exception("gateway task %s %s failed", method, path)
|
||
return 500, b"", "gateway operation failed"
|
||
if result is None:
|
||
return 204, b"", ""
|
||
if isinstance(result, tuple):
|
||
status, value = result
|
||
return int(status), json.dumps(value, ensure_ascii=False, separators=(",", ":")).encode("utf-8"), ""
|
||
return 200, json.dumps(result, ensure_ascii=False, separators=(",", ":")).encode("utf-8"), ""
|
||
|
||
@staticmethod
|
||
def _log_task_failure(future: Any) -> None:
|
||
exc = future.exception()
|
||
if exc is not None:
|
||
LOG.exception("gateway task dispatch failed", exc_info=exc)
|
||
|
||
# ---- 事件桥 ---------------------------------------------------------
|
||
|
||
def _event_sender_loop(self, conn: websocket.WebSocket) -> None:
|
||
manager = self.gateway.subscriptions
|
||
while not self._stop_requested.is_set():
|
||
try:
|
||
with manager.changed:
|
||
version = manager.version
|
||
self._push_wakeup.clear()
|
||
self._push_pending(conn, version)
|
||
with manager.changed:
|
||
manager.changed.wait_for(
|
||
lambda: self._stop_requested.is_set()
|
||
or manager.version != version
|
||
or self._push_wakeup.is_set(),
|
||
timeout=5,
|
||
)
|
||
except Exception:
|
||
if self._stop_requested.is_set():
|
||
return
|
||
LOG.exception("event sender loop failed; closing connection")
|
||
try:
|
||
conn.close()
|
||
except Exception:
|
||
pass
|
||
return
|
||
|
||
@staticmethod
|
||
def _subscription_items(manager: Any) -> list[tuple[str, Any]]:
|
||
# 从注册表快照枚举活动订阅;_items 缺失(测试替身)时视为空。
|
||
items = getattr(manager, "_items", {})
|
||
if not isinstance(items, dict):
|
||
return []
|
||
lock = getattr(manager, "_lock", None)
|
||
if not isinstance(lock, type(threading.RLock())):
|
||
return list(items.items())
|
||
with lock:
|
||
return list(items.items())
|
||
|
||
def _push_pending(self, conn: websocket.WebSocket, version: int) -> None:
|
||
# 每轮全量对账 pending 与已推记录的差集:新事件推送一次;nack/重连清空
|
||
# 已推记录后,同一批次会自然重推;无新差集时为空操作。
|
||
from .platform.douyin import DouyinError
|
||
|
||
manager = self.gateway.subscriptions
|
||
with self._sent_lock:
|
||
sent = self._sent_events
|
||
for alias, item in self._subscription_items(manager):
|
||
try:
|
||
current = manager._get(alias)
|
||
error = item.error or (
|
||
"event subscription replaced or stopped"
|
||
if current is not item or item.stopped
|
||
else ""
|
||
)
|
||
except (DouyinError, Exception) as exc:
|
||
error = str(exc) or exc.__class__.__name__
|
||
if error:
|
||
self._send(conn, {"type": "error", "subscription": "", "alias": alias, "error": error})
|
||
sent.pop(alias, None)
|
||
continue
|
||
try:
|
||
already = sent.setdefault(alias, {})
|
||
pending = [d for d in item.pending() if d["delivery_id"] not in already]
|
||
except Exception as exc:
|
||
self._send(conn, {"type": "error", "subscription": "", "alias": alias, "error": str(exc)})
|
||
sent.pop(alias, None)
|
||
continue
|
||
if not pending:
|
||
continue
|
||
subscription = f"{alias}:{item.session_id}"
|
||
self._send(conn, {
|
||
"type": "event",
|
||
"event": {"subscription": subscription, "alias": alias, "deliveries": pending},
|
||
})
|
||
stamp = time.monotonic()
|
||
already.update({d["delivery_id"]: stamp for d in pending})
|
||
LOG.info("event pushed alias=%s subscription=%s count=%s", alias, subscription, len(pending))
|
||
|
||
def _handle_event_ack(self, frame: dict) -> None:
|
||
alias = frame.get("alias")
|
||
ids = frame.get("delivery_ids")
|
||
if not isinstance(alias, str) or not isinstance(ids, list):
|
||
self.last_error = "event acknowledgement invalid"
|
||
return
|
||
manager = self.gateway.subscriptions
|
||
with self._sent_lock:
|
||
already = self._sent_events.get(alias, {})
|
||
unknown = [i for i in ids if i not in already]
|
||
if unknown:
|
||
# 平台确认了网关没有记录的投递(重连后旧 ack 等),仅清理已知项。
|
||
LOG.warning("event ack references unknown deliveries alias=%s count=%s", alias, len(unknown))
|
||
known = [i for i in ids if i in already]
|
||
if not known:
|
||
return
|
||
manager.ack(alias, known)
|
||
with self._sent_lock:
|
||
already = self._sent_events.get(alias, {})
|
||
for delivery_id in known:
|
||
already.pop(delivery_id, None)
|
||
|
||
def _handle_event_nack(self, frame: dict) -> None:
|
||
alias = frame.get("alias")
|
||
if not isinstance(alias, str):
|
||
return
|
||
with self._sent_lock:
|
||
self._sent_events.pop(alias, None)
|
||
self._wake_sender()
|
||
|
||
def _wake_sender(self) -> None:
|
||
# Condition.wait_for 只在 notify/超时时求值谓词;清空已推记录后必须唤醒。
|
||
self._push_wakeup.set()
|
||
manager = self.gateway.subscriptions
|
||
changed = getattr(manager, "changed", None)
|
||
if isinstance(changed, threading.Condition):
|
||
with changed:
|
||
changed.notify_all()
|