Files
dy/backend/rpa_engine/batch_start.py
T
2026-07-30 10:06:53 +08:00

400 lines
16 KiB
Python

"""Bounded, deduplicated background queue for bulk account starts."""
from __future__ import annotations
import asyncio
import logging
import os
import time
import uuid
from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import Any, Awaitable, Callable
logger = logging.getLogger("rpa.batch_start")
StartHandler = Callable[[int], Awaitable[dict[str, Any]]]
JobToken = tuple[str, int]
class _StartPreparationTimeout(Exception):
"""Internal marker for the queue's own per-account deadline."""
DEFAULT_CONCURRENCY = 6
MAX_CONCURRENCY = 32
def _configured_concurrency() -> int:
"""Admission width for account preparation.
Every account spends most of its startup waiting: for the shared network
lane, for a WebSocket handshake, for signing work in a thread. Admitting
only two at a time therefore left the network lane idle and made a fleet of
several hundred accounts take tens of minutes. Actual outbound traffic is
still capped by the traffic controller, so a wider admission window fills
the existing lane instead of adding load.
"""
try:
return max(
1,
min(
MAX_CONCURRENCY,
int(os.getenv("KEFU_BATCH_START_CONCURRENCY", str(DEFAULT_CONCURRENCY))),
),
)
except (TypeError, ValueError):
return DEFAULT_CONCURRENCY
def _configured_timeout_seconds() -> float:
try:
value = float(os.getenv("KEFU_BATCH_START_TIMEOUT_SECONDS", "90"))
except (TypeError, ValueError):
return 90.0
if value <= 0:
return 0.0
return max(5.0, min(600.0, value))
def _utc_now() -> str:
return datetime.now(timezone.utc).isoformat()
@dataclass
class _BatchRecord:
batch_id: str
owner_id: int | None = None
metadata: dict[str, Any] = field(default_factory=dict)
items: dict[int, dict[str, Any]] = field(default_factory=dict)
created_at: str = field(default_factory=_utc_now)
updated_at: str = field(default_factory=_utc_now)
class BatchStartQueue:
"""Run account preparation with a small process-wide concurrency cap.
The HTTP endpoint only enqueues account ids. Long credential checks then
run in these workers, so a batch of many accounts cannot block the request
that submitted it. ``_pending_accounts`` makes overlapping clicks and
overlapping batches idempotent for each account. Pending and active
ownership is tied to a concrete job token so cleanup from a cancelled old
job cannot release a newer submission for the same account.
"""
def __init__(
self,
handler: StartHandler,
concurrency: int | None = None,
max_batches: int = 100,
timeout_seconds: float | None = None,
) -> None:
self._handler = handler
self.concurrency = max(1, int(concurrency or _configured_concurrency()))
self.max_batches = max(10, int(max_batches or 100))
self.timeout_seconds = (
_configured_timeout_seconds()
if timeout_seconds is None
else max(0.0, float(timeout_seconds or 0.0))
)
self._queue: asyncio.Queue[JobToken] = asyncio.Queue()
self._pending_jobs: dict[int, JobToken] = {}
self._active_tasks: dict[int, tuple[JobToken, asyncio.Task]] = {}
self._batches: dict[str, _BatchRecord] = {}
self._workers: list[asyncio.Task] = []
self._worker_sequence = 0
self._lock = asyncio.Lock()
self._stopping = False
async def _ensure_workers(self) -> None:
async with self._lock:
self._workers = [task for task in self._workers if not task.done()]
if self._stopping:
return
# Top up to the configured width instead of only starting from
# zero. A worker that died on an unexpected error used to shrink
# the queue permanently, so later batches crawled through a single
# remaining worker with no way to recover short of a restart.
while len(self._workers) < self.concurrency:
self._worker_sequence += 1
self._workers.append(
asyncio.create_task(
self._worker(self._worker_sequence),
name=f"account-batch-start-{self._worker_sequence}",
)
)
async def submit(
self,
account_ids: list[int],
owner_id: int | None = None,
metadata: dict[str, Any] | None = None,
) -> dict[str, Any]:
await self._ensure_workers()
unique_ids = list(dict.fromkeys(int(value) for value in account_ids if int(value) > 0))
batch_id = uuid.uuid4().hex
record = _BatchRecord(
batch_id=batch_id,
owner_id=owner_id,
metadata=dict(metadata or {}),
)
async with self._lock:
if self._stopping:
raise RuntimeError("账号启动队列正在停止")
self._prune_locked()
self._batches[batch_id] = record
for account_id in unique_ids:
if account_id in self._pending_jobs:
record.items[account_id] = {
"account_id": account_id,
"status": "already_queued",
"message": "账号已在启动队列中",
}
continue
job_token = (batch_id, account_id)
self._pending_jobs[account_id] = job_token
record.items[account_id] = {
"account_id": account_id,
"status": "queued",
"message": "等待启动",
}
self._queue.put_nowait(job_token)
record.updated_at = _utc_now()
return self._snapshot_locked(record)
async def get_batch(
self,
batch_id: str,
owner_id: int | None = None,
include_items: bool = False,
) -> dict[str, Any] | None:
async with self._lock:
record = self._batches.get(str(batch_id or ""))
if record and owner_id is not None and record.owner_id != int(owner_id):
return None
return self._snapshot_locked(record, include_items=include_items) if record else None
async def _worker(self, worker_number: int) -> None:
while True:
batch_id, account_id = await self._queue.get()
job_token = (batch_id, account_id)
started_at = time.monotonic()
handler_task: asyncio.Task | None = None
try:
async with self._lock:
record = self._batches.get(batch_id)
if not record:
if self._pending_jobs.get(account_id) == job_token:
self._pending_jobs.pop(account_id, None)
continue
item = record.items[account_id]
if item.get("status") == "cancelled":
if self._pending_jobs.get(account_id) == job_token:
self._pending_jobs.pop(account_id, None)
continue
item.update(status="processing", message="正在校验并启动")
record.updated_at = _utc_now()
handler_task = asyncio.create_task(
self._handler(account_id),
name=f"account-start-{account_id}",
)
self._active_tasks[account_id] = (job_token, handler_task)
if self.timeout_seconds > 0:
try:
result = await asyncio.wait_for(
handler_task,
timeout=self.timeout_seconds,
)
except asyncio.TimeoutError as exc:
# wait_for cancels its task only when this queue's
# deadline expires. Preserve a TimeoutError raised by
# the handler itself as its real account failure.
if handler_task.cancelled():
raise _StartPreparationTimeout from exc
raise
else:
result = await handler_task
async with self._lock:
record = self._batches.get(batch_id)
if record:
item = record.items[account_id]
if item.get("status") != "cancelled":
item.update(
status="submitted",
message=str(result.get("message") or "已提交启动"),
login_mode=result.get("login_mode"),
skip_browser=bool(result.get("skip_browser", False)),
elapsed_seconds=round(time.monotonic() - started_at, 3),
)
record.updated_at = _utc_now()
except _StartPreparationTimeout:
elapsed = round(time.monotonic() - started_at, 3)
logger.warning(
"Batch start timed out account=%s worker=%s after %.1fs",
account_id,
worker_number,
self.timeout_seconds,
)
async with self._lock:
record = self._batches.get(batch_id)
if record:
item = record.items[account_id]
if item.get("status") != "cancelled":
item.update(
status="failed",
message=(
f"启动准备超过 {self.timeout_seconds:g} 秒,"
"已跳过并继续处理后续账号"
),
elapsed_seconds=elapsed,
)
record.updated_at = _utc_now()
except asyncio.CancelledError:
async with self._lock:
record = self._batches.get(batch_id)
if record:
record.items[account_id].update(
status="cancelled",
message="服务停止,启动任务已取消",
)
record.updated_at = _utc_now()
# Cancelling one account must not kill a long-lived queue
# worker. Re-raise only when stop() cancelled the worker.
if asyncio.current_task().cancelling():
raise
except Exception as exc:
logger.exception(
"Batch start failed account=%s worker=%s: %s",
account_id,
worker_number,
exc,
)
async with self._lock:
record = self._batches.get(batch_id)
if record:
detail = getattr(exc, "detail", None) or str(exc) or "启动失败"
item = record.items[account_id]
if item.get("status") != "cancelled":
item.update(
status="failed",
message=str(detail),
elapsed_seconds=round(time.monotonic() - started_at, 3),
)
record.updated_at = _utc_now()
finally:
async with self._lock:
active_entry = self._active_tasks.get(account_id)
if active_entry == (job_token, handler_task):
self._active_tasks.pop(account_id, None)
if self._pending_jobs.get(account_id) == job_token:
self._pending_jobs.pop(account_id, None)
self._queue.task_done()
async def cancel_account(self, account_id: int) -> int:
"""Cancel queued/processing work so stop/delete cannot restart it."""
account_key = int(account_id)
cancelled = 0
active_task: asyncio.Task | None = None
async with self._lock:
for record in self._batches.values():
item = record.items.get(account_key)
if not item or item.get("status") not in ("queued", "processing"):
continue
item.update(status="cancelled", message="启动任务已取消")
record.updated_at = _utc_now()
cancelled += 1
self._pending_jobs.pop(account_key, None)
active_entry = self._active_tasks.get(account_key)
active_task = active_entry[1] if active_entry else None
if active_task and not active_task.done():
active_task.cancel()
if active_task and not active_task.done():
await asyncio.gather(active_task, return_exceptions=True)
return cancelled
def _snapshot_locked(
self,
record: _BatchRecord,
*,
include_items: bool = True,
) -> dict[str, Any]:
counts = {
"queued": 0,
"processing": 0,
"submitted": 0,
"failed": 0,
"skipped": 0,
"cancelled": 0,
}
browser_required = 0
for item in record.items.values():
status = item.get("status")
if status == "submitted" and item.get("skip_browser") is False:
browser_required += 1
if status == "already_queued":
counts["skipped"] += 1
elif status in counts:
counts[status] += 1
active = counts["queued"] + counts["processing"]
snapshot = {
"batch_id": record.batch_id,
"total_count": len(record.items),
"accepted_count": len(record.items) - counts["skipped"],
"queued_count": counts["queued"],
"processing_count": counts["processing"],
"submitted_count": counts["submitted"],
"failed_count": counts["failed"],
"skipped_count": counts["skipped"],
"cancelled_count": counts["cancelled"],
"browser_required_count": browser_required,
"complete": active == 0,
"created_at": record.created_at,
"updated_at": record.updated_at,
}
snapshot.update(record.metadata)
if include_items:
snapshot["items"] = [dict(item) for item in record.items.values()]
return snapshot
def _prune_locked(self) -> None:
if len(self._batches) < self.max_batches:
return
removable = [
batch_id
for batch_id, record in self._batches.items()
if self._snapshot_locked(record, include_items=False)["complete"]
]
for batch_id in removable[: max(1, len(self._batches) - self.max_batches + 1)]:
self._batches.pop(batch_id, None)
async def stop(self) -> None:
async with self._lock:
self._stopping = True
workers = list(self._workers)
self._workers.clear()
for task in workers:
task.cancel()
if workers:
await asyncio.gather(*workers, return_exceptions=True)
async with self._lock:
self._active_tasks.clear()
while True:
try:
batch_id, account_id = self._queue.get_nowait()
except asyncio.QueueEmpty:
break
record = self._batches.get(batch_id)
if record:
record.items[account_id].update(
status="cancelled",
message="服务停止,启动任务已取消",
)
record.updated_at = _utc_now()
job_token = (batch_id, account_id)
if self._pending_jobs.get(account_id) == job_token:
self._pending_jobs.pop(account_id, None)
self._queue.task_done()