Files
ai-balance-monitor/app/monitor.py
T

217 lines
8.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""监控调度器:线程池并发轮询所有启用账号。
调度逻辑:
- 每个账号记录 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