217 lines
8.3 KiB
Python
217 lines
8.3 KiB
Python
"""监控调度器:线程池并发轮询所有启用账号。
|
||
|
||
调度逻辑:
|
||
- 每个账号记录 next_check_at(epoch 秒),间隔取平台覆盖值或全局默认
|
||
- 调度线程每秒扫描一次,到期的账号且无进行中任务则提交线程池
|
||
- 任务完成后把 next_check_at 更新为 now + interval,避免重复执行
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import base64
|
||
import logging
|
||
import threading
|
||
import time
|
||
from concurrent.futures import Future, ThreadPoolExecutor
|
||
|
||
from app import db
|
||
from app.alert import evaluate, notify_disabled
|
||
from app.config import Config
|
||
from app.fetcher import fetch_balance
|
||
from app.providers import get_provider
|
||
|
||
logger = logging.getLogger("monitor.scheduler")
|
||
|
||
# JOIN 查询后拆分为账号/平台两个 dict 用到的字段
|
||
_ACCOUNT_FIELDS = [
|
||
"id", "platform_id", "name", "api_key", "threshold", "enabled",
|
||
"alert_armed", "last_balance", "last_status", "last_error",
|
||
"last_check_at", "note",
|
||
]
|
||
_PLATFORM_FIELDS = [
|
||
"platform_id", "provider_id", "platform_name", "currency", "icon",
|
||
"interval_seconds", "retry_count", "timeout_seconds", "platform_enabled", "note",
|
||
]
|
||
|
||
_ACCOUNT_SQL = """
|
||
SELECT a.*,
|
||
p.provider_id, p.name AS platform_name, p.currency, p.icon,
|
||
p.interval_seconds, p.retry_count, p.timeout_seconds,
|
||
p.enabled AS platform_enabled, p.note AS platform_note
|
||
FROM accounts a JOIN platforms p ON p.id = a.platform_id
|
||
"""
|
||
|
||
|
||
def _unb64(v: str) -> str:
|
||
"""api_key 以 base64 存储,内部使用时必须解码回明文。"""
|
||
try:
|
||
return base64.b64decode(v.encode("ascii")).decode("utf-8")
|
||
except Exception:
|
||
return v
|
||
|
||
|
||
def _split(row: dict) -> tuple[dict, dict]:
|
||
"""把 JOIN 行拆成 (account, platform) 两个 dict。"""
|
||
account = {k: row[k] for k in _ACCOUNT_FIELDS if k in row}
|
||
account["api_key"] = _unb64(account["api_key"])
|
||
platform = {k: row[k] for k in _PLATFORM_FIELDS if k in row}
|
||
platform["name"] = platform.pop("platform_name")
|
||
platform["id"] = platform.pop("platform_id")
|
||
platform["enabled"] = platform.pop("platform_enabled")
|
||
return account, platform
|
||
|
||
|
||
def _account_interval(platform: dict, cfg: Config) -> float:
|
||
return float(platform.get("interval_seconds") or cfg.global_interval_seconds)
|
||
|
||
|
||
class Monitor:
|
||
def __init__(self, cfg: Config) -> None:
|
||
self.cfg = cfg
|
||
self._stop = threading.Event()
|
||
self._pool: ThreadPoolExecutor | None = None
|
||
self._thread: threading.Thread | None = None
|
||
self._next_check: dict[int, float] = {}
|
||
self._inflight: set[int] = set()
|
||
self._lock = threading.Lock()
|
||
|
||
# ---------- 生命周期 ----------
|
||
|
||
def start(self) -> None:
|
||
if self._thread is not None:
|
||
return
|
||
self._pool = ThreadPoolExecutor(max_workers=self.cfg.max_workers, thread_name_prefix="check")
|
||
self._thread = threading.Thread(target=self._loop, name="monitor-loop", daemon=True)
|
||
self._thread.start()
|
||
logger.info(
|
||
"监控调度已启动,全局间隔 %ss,并发 %s",
|
||
self.cfg.global_interval_seconds,
|
||
self.cfg.max_workers,
|
||
)
|
||
|
||
def stop(self) -> None:
|
||
self._stop.set()
|
||
if self._thread is not None:
|
||
self._thread.join(timeout=5)
|
||
if self._pool is not None:
|
||
self._pool.shutdown(wait=False, cancel_futures=True)
|
||
logger.info("监控调度已停止")
|
||
|
||
# ---------- 调度循环 ----------
|
||
|
||
def _loop(self) -> None:
|
||
while not self._stop.is_set():
|
||
try:
|
||
self._dispatch()
|
||
except Exception:
|
||
logger.exception("调度循环异常")
|
||
self._stop.wait(1)
|
||
|
||
def _dispatch(self) -> None:
|
||
with db.get_conn() as conn:
|
||
rows = conn.execute(
|
||
_ACCOUNT_SQL + " WHERE a.enabled = 1 AND p.enabled = 1"
|
||
).fetchall()
|
||
now = time.time()
|
||
for row in rows:
|
||
data = dict(row)
|
||
aid = data["id"]
|
||
with self._lock:
|
||
if aid in self._inflight:
|
||
continue
|
||
if now < self._next_check.get(aid, 0):
|
||
continue
|
||
self._inflight.add(aid)
|
||
assert self._pool is not None
|
||
future = self._pool.submit(self._check_one, data)
|
||
future.add_done_callback(lambda f, a=aid: self._on_done(a, f))
|
||
|
||
def _on_done(self, account_id: int, future: Future) -> None:
|
||
with self._lock:
|
||
self._inflight.discard(account_id)
|
||
if future.exception():
|
||
logger.error("账号 %s 检查任务异常: %s", account_id, future.exception())
|
||
|
||
# ---------- 单次检查 ----------
|
||
|
||
def _check_one(self, row: dict) -> None:
|
||
cfg = self.cfg
|
||
account, platform = _split(row)
|
||
provider = get_provider(platform["provider_id"])
|
||
if provider is None:
|
||
logger.warning("平台 %s 的 provider 不存在: %r", platform["name"], platform["provider_id"])
|
||
with db.get_conn() as conn:
|
||
conn.execute(
|
||
"UPDATE accounts SET last_status='error', last_error=?, "
|
||
"last_check_at=datetime('now','localtime') WHERE id=?",
|
||
(f"provider 不存在: {platform['provider_id']}", account["id"]),
|
||
)
|
||
with self._lock:
|
||
self._next_check[account["id"]] = time.time() + _account_interval(platform, cfg)
|
||
return
|
||
retry = platform.get("retry_count")
|
||
if retry is None:
|
||
retry = cfg.retry_count
|
||
timeout = platform.get("timeout_seconds")
|
||
if timeout is None:
|
||
timeout = cfg.timeout_seconds
|
||
result = fetch_balance(provider, account["api_key"], int(retry), int(timeout))
|
||
|
||
with db.get_conn() as conn:
|
||
if result.ok:
|
||
conn.execute(
|
||
"UPDATE accounts SET last_balance=?, last_status='ok', last_error='', "
|
||
"last_check_at=datetime('now','localtime') WHERE id=?",
|
||
(result.balance, account["id"]),
|
||
)
|
||
conn.execute(
|
||
"INSERT INTO balance_history (account_id, balance) VALUES (?, ?)",
|
||
(account["id"], result.balance),
|
||
)
|
||
elif result.auth_error:
|
||
conn.execute(
|
||
"UPDATE accounts SET enabled=0, last_status='disabled', last_error=?, "
|
||
"last_check_at=datetime('now','localtime') WHERE id=?",
|
||
(result.error, account["id"]),
|
||
)
|
||
else:
|
||
conn.execute(
|
||
"UPDATE accounts SET last_status='error', last_error=?, "
|
||
"last_check_at=datetime('now','localtime') WHERE id=?",
|
||
(result.error, account["id"]),
|
||
)
|
||
|
||
if result.auth_error:
|
||
logger.warning("认证失败已禁用: %s / %s(%s)", platform["name"], account["name"], result.error)
|
||
notify_disabled(account, platform, cfg)
|
||
elif result.ok:
|
||
logger.info("检查成功: %s / %s 余额 %.4f %s", platform["name"], account["name"], result.balance, platform["currency"])
|
||
evaluate(account, platform, result.balance, cfg)
|
||
else:
|
||
logger.warning("检查失败: %s / %s -> %s", platform["name"], account["name"], result.error)
|
||
|
||
# 调度:无论成功失败,按间隔排下一次
|
||
with self._lock:
|
||
self._next_check[account["id"]] = time.time() + _account_interval(platform, cfg)
|
||
|
||
# ---------- 手动触发 ----------
|
||
|
||
def check_now(self, account_id: int) -> bool:
|
||
"""API 手动立即检查;账号/平台未启用或已在检查中则返回 False。"""
|
||
with db.get_conn() as conn:
|
||
row = conn.execute(
|
||
_ACCOUNT_SQL + " WHERE a.id = ? AND a.enabled = 1 AND p.enabled = 1",
|
||
(account_id,),
|
||
).fetchone()
|
||
if row is None:
|
||
return False
|
||
data = dict(row)
|
||
with self._lock:
|
||
if account_id in self._inflight:
|
||
return False
|
||
self._inflight.add(account_id)
|
||
assert self._pool is not None
|
||
future = self._pool.submit(self._check_one, data)
|
||
future.add_done_callback(lambda f, a=account_id: self._on_done(a, f))
|
||
return True
|