diff --git a/UPSTREAM.md b/UPSTREAM.md new file mode 100644 index 0000000..b6b74ae --- /dev/null +++ b/UPSTREAM.md @@ -0,0 +1,129 @@ +# 与上游的差异管理 + +本仓库在 [NanmiCoder/MediaCrawler](https://github.com/NanmiCoder/MediaCrawler) 之上加了一层 +监控/鉴权/多平台面板。这份文档记录**改了上游哪些文件、为什么**,以及**上游更新时怎么合并**。 + +--- + +## 一、改动分三类 + +冲突风险从低到高: + +### 1. 纯新增文件(零冲突) + +上游怎么改都不会碰到它们: + +``` +api/auth.py WebUI 登录鉴权 +api/monitor/* 监控层整体(含 platforms.py 能力矩阵) +api/monitor/db.py MySQL 连接层(可回退 SQLite 供测试用) +api/monitor/migrate_from_sqlite.py SQLite → MySQL 一次性迁移脚本 +api/routers/{auth,monitor,settings}.py +api/schemas/{auth,monitor,settings}.py +api/services/interpreter.py 解释器探测(uv / .venv / 当前解释器) +webui/src/components/{monitor,settings,auth}/ 新视图 +webui/src/components/layout/{PlatformSwitcher,UnwiredPlatformNotice}.tsx +webui/src/{hooks/useMonitor.ts,hooks/usePlatform.ts,store/platformStore.ts,lib/monitorFormat.ts,types/monitor.ts} +docs/监控功能使用说明.md +tests/test_{auth,settings,platforms,monitor_*}.py +``` + +### 2. 加法改动(低冲突) + +只在既有文件里**新增**内容,不改动原有行: + +| 文件 | 加了什么 | +|---|---| +| `cmd_arg/arg.py` | typer 选项:`--enable_cdp_mode`、`--inject_all_cookies`、`--save_login_state`、`--cookies_file`、`--crawler_max_sleep_sec`,以及对应的 `config.*` 回写 | +| `api/schemas/crawler.py` | `CrawlerStartRequest` 的若干**可选**字段(默认 `None`,不传则不加对应 CLI 参数) | +| `config/base_config.py` | `INJECT_ALL_COOKIES = False` | +| `api/routers/__init__.py` | 导出新增的 router | +| `requirements.txt` | 补上 `websockets`(上游 `pyproject.toml` 里有、`requirements.txt` 里漏了) | +| `tests/conftest.py` | 新增 `_bypass_auth_for_non_auth_suites` fixture | + +### 3. 接线改动(中冲突,需要人看) + +| 文件 | 改了什么 | 上游若在此处变动 | +|---|---|---| +| `api/main.py` | 注册 4 个 router 并加 `Depends(require_auth)`;`lifespan` 里初始化监控库、启动调度器、跑设置键迁移;`load_dotenv`;CORS 可配;`docs/redoc/openapi` 关闭;监听地址改 env | **最需要人工合并的文件**。留意 router 注册块、lifespan、`__main__` | +| `api/routers/websocket.py` | 两个 WS 路由加 `dependencies=[Depends(require_ws_auth)]` | 上游若新增 WS 路由,**必须同样加上**,否则那条流是裸奔的 | +| `api/services/crawler_manager.py` | 解释器探测替换硬编码 `uv run`;`_build_command` 转发新增参数;新增 `is_busy()` / `run_and_wait()` 与完成事件 | 留意 `_build_command` 的参数拼装 | +| `media_platform/xhs/login.py` | `login_by_cookies` 在 `INJECT_ALL_COOKIES` 打开时注入**全部** cookie(默认关闭,行为不变) | 小改动,好合并 | + +### 4. 上游 bug 修复(建议回馈上游) + +| 文件 | 修的问题 | +|---|---| +| `media_platform/xhs/core.py` | 见下节 | +| `media_platform/xhs/login.py` | 同上(cookie 加固) | + +--- + +## 二、应该给上游提 PR 的两个修复 + +这两处是**上游自身的缺陷**,提上去以后就不用自己背着: + +### 1. 博主主页抓取失败会跳掉整个博主(`xhs/core.py`) + +`get_creator_info()` 抓主页 HTML 解析 `window.__INITIAL_STATE__`,解析失败抛 `JSONDecodeError`—— +它是 `ValueError` 的子类,被 `except ValueError` 误捕获,日志报成 +"Failed to parse creator URL"(**误导**,URL 根本没解析错),然后 `continue` **跳过整个博主**。 + +而那份资料只喂给 `save_creator()`,它在教学版里是**空函数**。也就是说: +一个喂给空函数的抓取失败,让真正要抓的作品一条都没抓到,表现为"0 篇作品", +和"登录失效"长得一模一样。 + +修复:把资料抓取改成**尽力而为**,失败只警告、继续抓作品。 + +### 2. cookie 登录只注入 `web_session`(`xhs/login.py`) + +`a1` / `webId` 等签名所需 cookie 只能靠持久化 profile 补,冷启动时签名会失败。 +默认行为保持不变,用 `INJECT_ALL_COOKIES` 开关控制。 + +--- + +## 三、上游更新时怎么操作 + +### 日常流程 + +```bash +git stash # 或先 commit 到自己的分支(推荐) +git fetch origin main +git rebase origin/main # 冲突只会出现在上表第 3、4 类文件里 +./.venv/Scripts/python.exe -m pytest tests/ -q # 486 个测试就是回归网 +``` + +### 强烈建议:先把改动提交掉 + +当前状态是**未提交**的(25 个上游文件被改 + 31 个新文件)。在 `main` 分支上裸着工作区, +一次 `git checkout .` 就全没了,而且没法 rebase。 + +```bash +git checkout -b local/monitor-panel +git add -A && git commit -m "监控面板 / 鉴权 / 多平台" +``` + +### 如果改动持续增长:fork + +把本仓库 fork 到自己名下,上游设为 remote: + +```bash +git remote rename origin upstream +git remote add origin <你的 fork> +git push -u origin local/monitor-panel +``` + +之后同步上游用 `git fetch upstream && git rebase upstream/main`。 + +--- + +## 四、合并时最容易忘的三件事 + +1. **新增的 `/api` 路由必须带鉴权**。跑一下 `tests/test_auth.py`—— + 里面有个测试会遍历 `app.routes`,断言除豁免集外每个 `/api` 路由无凭据都返回 401。 + 上游新增接口忘了加鉴权,这个测试会直接失败。 +2. **新增的 WebSocket 路由必须加 `require_ws_auth`**。 + `BaseHTTPMiddleware` 对 WS 完全不生效(`scope["type"] != "http"` 直接放行), + 只靠中间件会漏。同样有测试守着。 +3. **上游若改动 `AsyncFileWriter` 的输出路径规则**,`api/monitor/ingest.py::find_run_files` + 会跟着失效——它靠 glob `{out_dir}/{platform}/jsonl/*_contents_*.jsonl` 定位每轮的产物。 diff --git a/api/auth.py b/api/auth.py new file mode 100644 index 0000000..a50bc6e --- /dev/null +++ b/api/auth.py @@ -0,0 +1,369 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/auth.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Authentication for the WebUI. + +Design constraints that drove this, all verified against the codebase: + +* **Cookies, not bearer headers, are the primary transport.** Browser + WebSockets cannot set custom headers on the handshake, and the data-export + downloads use ``window.open`` (a navigation, also header-less). Only a cookie + is carried on both. The same opaque token is *also* accepted from an + ``Authorization: Bearer`` header so ``curl`` and scripts remain usable. +* **Enforcement is a ``Depends``, not middleware.** ``BaseHTTPMiddleware`` + returns early for any non-``http`` scope, so it never sees a WebSocket -- + a middleware-only gate would leave the live log stream wide open. It is also + overridable per-test via ``app.dependency_overrides``. +* **Sessions are server-side** so logout and password-change revoke immediately. + +Only the environment variable ``MC_PASSWORD`` can bypass the stored hash. That is +the documented way back in if the password is forgotten, which is why it is +never persisted. +""" + +import asyncio +import base64 +import binascii +import hashlib +import hmac +import os +import secrets +import time +from typing import Optional + +from anyio import to_thread +from fastapi import HTTPException, Request, WebSocket, WebSocketException, status +from sqlalchemy import delete +from sqlalchemy.ext.asyncio import AsyncSession + +from tools.time_util import get_current_timestamp + +from .monitor.db import get_session +from .monitor.models import ( + SETTING_AUTH_PASSWORD_HASH, + SETTING_AUTH_PASSWORD_UPDATED_AT, + AuthSession, +) +from .monitor.settings import get_setting, set_setting + +SESSION_COOKIE_NAME = "mc_session" + +# OWASP's current PBKDF2-HMAC-SHA256 guidance. Deliberately slow -- see +# verify_password() for why that cost must not land on the event loop. +PBKDF2_ITERATIONS = 600_000 +PBKDF2_ALGO = "pbkdf2_sha256" + +# A single generic message for every failure mode, so the response never +# reveals whether a password is set, wrong, or empty. +INVALID_CREDENTIALS = "用户名或密码错误" + +# Brute-force throttle. In-process is sufficient: this is a single-user tool and +# uvicorn runs one worker. Documented as reset-on-restart. +THROTTLE_THRESHOLD = 5 +THROTTLE_WINDOW_SECONDS = 900 +THROTTLE_MAX_LOCKOUT_SECONDS = 900 + +_failures: dict[str, list[float]] = {} +_throttle_lock = asyncio.Lock() + + +def _now() -> float: + """Monotonic clock, indirected so tests can drive it without sleeping.""" + return time.monotonic() + + +# --------------------------------------------------------------------------- +# Environment configuration (read at call time so tests can set it per-case) +# --------------------------------------------------------------------------- + +def env_password() -> str: + return os.getenv("MC_PASSWORD", "").strip() + + +def cookie_secure() -> bool: + return os.getenv("MC_COOKIE_SECURE", "").strip().lower() in ("1", "true", "yes", "y") + + +def session_ttl_ms() -> int: + try: + hours = int(os.getenv("MC_SESSION_TTL_HOURS", "336")) + except ValueError: + hours = 336 + return max(hours, 1) * 3_600_000 + + +# --------------------------------------------------------------------------- +# Password hashing +# --------------------------------------------------------------------------- + +def _b64(raw: bytes) -> str: + return base64.b64encode(raw).decode("ascii") + + +def hash_password(password: str, *, iterations: Optional[int] = None) -> str: + """Return a self-describing hash so the iteration count can be raised later + without a migration: ``pbkdf2_sha256$$$``. + + ``iterations`` is resolved at call time (not bound as a default) so tests can + lower it; the production value stays the module constant. + """ + iterations = iterations or PBKDF2_ITERATIONS + salt = secrets.token_bytes(16) + digest = hashlib.pbkdf2_hmac("sha256", password.encode("utf-8"), salt, iterations) + return f"{PBKDF2_ALGO}${iterations}${_b64(salt)}${_b64(digest)}" + + +def _verify_password_sync(password: str, stored: str) -> bool: + try: + algo, iterations_raw, salt_raw, digest_raw = stored.split("$") + if algo != PBKDF2_ALGO: + return False + salt = base64.b64decode(salt_raw) + expected = base64.b64decode(digest_raw) + actual = hashlib.pbkdf2_hmac("sha256", password.encode("utf-8"), salt, int(iterations_raw)) + except (ValueError, TypeError, binascii.Error): + return False + return hmac.compare_digest(actual, expected) + + +async def verify_password(password: str, stored: str) -> bool: + """Verify off the event loop. + + At 600k iterations this takes a few hundred milliseconds. Running it inline + in an async handler would block the loop entirely -- stalling the monitor + scheduler and every websocket ping -- and present as "the whole UI freezes + when I click login". + """ + return await to_thread.run_sync(_verify_password_sync, password, stored) + + +async def current_password_hash(session: AsyncSession) -> str: + return (await get_setting(session, SETTING_AUTH_PASSWORD_HASH)) or "" + + +async def set_password(session: AsyncSession, password: str) -> None: + await set_setting(session, SETTING_AUTH_PASSWORD_HASH, hash_password(password)) + await set_setting( + session, SETTING_AUTH_PASSWORD_UPDATED_AT, str(get_current_timestamp()) + ) + + +async def check_password(session: AsyncSession, password: str) -> bool: + """The environment override wins over the stored hash, always. + + That is the escape hatch: forgetting the password is recoverable by setting + MC_PASSWORD and restarting, without touching the database. + """ + override = env_password() + if override: + return hmac.compare_digest(password, override) + + stored = await current_password_hash(session) + if not stored: + return False + return await verify_password(password, stored) + + +async def ensure_initial_credential() -> Optional[str]: + """Seed a password on first run; returns it once so main() can print it. + + Deliberately NOT an unauthenticated "set your password" endpoint: on a + LAN-exposed bind that is a claim-the-instance race where whoever reaches the + page first becomes the administrator. Generating and printing a random + password avoids the race and also avoids locking the operator out. + """ + if env_password(): + return None + + async with get_session() as session: + if await current_password_hash(session): + return None + generated = secrets.token_urlsafe(12) + await set_password(session, generated) + return generated + + +# --------------------------------------------------------------------------- +# Sessions +# --------------------------------------------------------------------------- + +def _hash_token(token: str) -> str: + return hashlib.sha256(token.encode("utf-8")).hexdigest() + + +async def create_session(session: AsyncSession) -> tuple[str, int]: + """Issue a session. Returns (token, expires_at_ms). + + The caller receives the raw token; only its hash is stored. + """ + token = secrets.token_urlsafe(32) + now = get_current_timestamp() + expires_at = now + session_ttl_ms() + session.add( + AuthSession( + token_hash=_hash_token(token), + created_at=now, + expires_at=expires_at, + last_seen_at=now, + ) + ) + return token, expires_at + + +async def resolve_session(session: AsyncSession, token: str) -> Optional[AuthSession]: + if not token: + return None + + row = await session.get(AuthSession, _hash_token(token)) + if row is None: + return None + + now = get_current_timestamp() + if row.expires_at <= now: + await session.delete(row) + return None + + row.last_seen_at = now + return row + + +async def revoke_session(session: AsyncSession, token: str) -> None: + row = await session.get(AuthSession, _hash_token(token)) + if row is not None: + await session.delete(row) + + +async def revoke_all_sessions(session: AsyncSession) -> int: + """Used on password change, which is what makes "all devices logged out" + take effect immediately rather than at token expiry.""" + result = await session.execute(delete(AuthSession)) + return result.rowcount or 0 + + +async def purge_expired_sessions(session: AsyncSession) -> None: + await session.execute(delete(AuthSession).where(AuthSession.expires_at <= get_current_timestamp())) + + +# --------------------------------------------------------------------------- +# Credential extraction and enforcement +# --------------------------------------------------------------------------- + +def token_from_request(request: Request) -> str: + """Cookie first (browsers, websockets, navigations), then Bearer (scripts).""" + token = request.cookies.get(SESSION_COOKIE_NAME, "") + if token: + return token + header = request.headers.get("authorization", "") + if header.lower().startswith("bearer "): + return header[7:].strip() + return "" + + +def _unauthorized() -> HTTPException: + return HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=INVALID_CREDENTIALS, + headers={"WWW-Authenticate": "Bearer"}, + ) + + +async def require_auth(request: Request) -> None: + """FastAPI dependency guarding the protected routers. + + Applied per-router via ``include_router(..., dependencies=[Depends(...)])`` + rather than as app-wide middleware, so it appears in the OpenAPI schema, + returns a correct 401, and can be overridden in tests. + """ + token = token_from_request(request) + if not token: + raise _unauthorized() + + async with get_session() as session: + if await resolve_session(session, token) is None: + raise _unauthorized() + + +async def require_ws_auth(websocket: WebSocket) -> None: + """Guard for WebSocket routes. + + These need their own dependency: ``BaseHTTPMiddleware`` passes any non-http + scope straight through, and router-level HTTP dependencies do not apply to + websocket routes. Raising ``WebSocketException`` closes the handshake with + the given code; ``HTTPException`` would be meaningless here. + """ + token = websocket.cookies.get(SESSION_COOKIE_NAME, "") + if not token: + raise WebSocketException(code=status.WS_1008_POLICY_VIOLATION) + + async with get_session() as session: + if await resolve_session(session, token) is None: + raise WebSocketException(code=status.WS_1008_POLICY_VIOLATION) + + +# --------------------------------------------------------------------------- +# Brute-force throttle +# --------------------------------------------------------------------------- + +def client_key(request: Request) -> str: + """Identify the caller for throttling. + + ``X-Forwarded-For`` is only consulted when the operator explicitly opts in, + because otherwise any client could spoof the header and throttle someone + else (or evade its own throttle). + """ + if os.getenv("MC_TRUST_PROXY", "").strip() == "1": + forwarded = request.headers.get("x-forwarded-for", "") + if forwarded: + return forwarded.split(",")[0].strip() + return request.client.host if request.client else "unknown" + + +def _recent_failures(key: str) -> list[float]: + cutoff = _now() - THROTTLE_WINDOW_SECONDS + return [ts for ts in _failures.get(key, []) if ts >= cutoff] + + +async def retry_after_seconds(key: str) -> int: + """0 when not throttled, otherwise how long the caller must wait.""" + async with _throttle_lock: + recent = _recent_failures(key) + _failures[key] = recent + if len(recent) < THROTTLE_THRESHOLD: + return 0 + + # Lockout doubles per failure past the threshold, capped. + extra = len(recent) - THROTTLE_THRESHOLD + lockout = min(2 ** extra, THROTTLE_MAX_LOCKOUT_SECONDS) + elapsed = _now() - recent[-1] + remaining = int(lockout - elapsed) + return max(remaining, 1) + + +async def record_failure(key: str) -> None: + async with _throttle_lock: + _failures.setdefault(key, []).append(_now()) + + +async def clear_failures(key: str) -> None: + async with _throttle_lock: + _failures.pop(key, None) + + +def reset_throttle_state() -> None: + """Test hook: drop all throttle state.""" + _failures.clear() diff --git a/api/main.py b/api/main.py index af49b15..87c4111 100644 --- a/api/main.py +++ b/api/main.py @@ -17,7 +17,7 @@ # 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 """ -MediaCrawler WebUI API Server +综合采集平台 API Server Start command: uvicorn api.main:app --port 8080 --reload Or: python -m api.main """ @@ -25,28 +25,95 @@ import asyncio import os import sys import subprocess +from contextlib import asynccontextmanager from pathlib import Path -import uvicorn -from fastapi import FastAPI -from fastapi.middleware.cors import CORSMiddleware -from fastapi.staticfiles import StaticFiles -from fastapi.responses import FileResponse - -from .routers import crawler_router, data_router, websocket_router # Project root directory (used for running subprocesses like uv run main.py) PROJECT_ROOT = Path(__file__).parent.parent +# Load .env before importing anything that reads os.getenv at module import time +# (config/db_config.py does). python-dotenv was already a declared dependency but +# nothing ever called it, so the shipped .env.example had no effect. +from dotenv import load_dotenv + +load_dotenv(PROJECT_ROOT / ".env") + +import uvicorn +from fastapi import Depends, FastAPI +from fastapi.middleware.cors import CORSMiddleware +from fastapi.staticfiles import StaticFiles +from fastapi.responses import FileResponse + +from .auth import ensure_initial_credential, require_auth +from .routers import ( + auth_router, + crawler_router, + data_router, + monitor_router, + settings_router, + websocket_router, +) +from .services.interpreter import describe_interpreter, resolve_python_cmd + + +@asynccontextmanager +async def lifespan(_app: FastAPI): + """Start the monitor scheduler with the server, and shut it down cleanly. + + The scheduler is a background asyncio task, so it must not be tied to a + browser session the way the log broadcaster is -- a scheduled run has to + happen whether or not anyone has the UI open. + """ + from .monitor.db import dispose_engine, init_db + from .monitor.scheduler import monitor_scheduler + + await init_db() + + generated = await ensure_initial_credential() + if generated: + # Printed once, on the run that creates it. There is no unauthenticated + # "set your password" endpoint on purpose: on a LAN bind that would be a + # claim-the-instance race. + rule = "=" * 68 + print( + f"\n{rule}\n" + " WebUI 首次启动,已生成登录密码:\n" + f"\n {generated}\n" + "\n 请立即登录并修改。忘记密码时可设置环境变量 MC_PASSWORD 后重启。\n" + f"{rule}\n", + flush=True, + ) + + await monitor_scheduler.start() + try: + yield + finally: + await monitor_scheduler.stop() + await dispose_engine() + + +# Docs are disabled deliberately: /docs, /redoc and /openapi.json are +# unauthenticated by default, which would hand out a complete map of the API +# (and a "Try it out" console that 401s anyway). app = FastAPI( - title="MediaCrawler WebUI API", - description="API for controlling MediaCrawler from WebUI", - version="1.0.0" + title="综合采集平台 API", + description="API for controlling 综合采集平台 from WebUI", + version="1.0.0", + lifespan=lifespan, + docs_url=None, + redoc_url=None, + openapi_url=None, ) # Get webui static files directory WEBUI_DIR = os.path.join(os.path.dirname(__file__), "webui") -# CORS configuration - allow frontend dev server access +# CORS only matters for a split-origin setup. In production this app serves the +# SPA itself, and in development Vite proxies /api here (see webui/vite.config.ts), +# so the browser always sees a single origin and CORS never actually triggers. +# Kept as an explicit allowlist -- never "*", which is invalid next to +# allow_credentials -- and extensible via env for a dev server reached over LAN. +_extra_origins = [o.strip() for o in os.getenv("MC_CORS_ORIGINS", "").split(",") if o.strip()] app.add_middleware( CORSMiddleware, allow_origins=[ @@ -54,15 +121,24 @@ app.add_middleware( "http://localhost:3000", # Backup port "http://127.0.0.1:5173", "http://127.0.0.1:3000", + *_extra_origins, ], + allow_origin_regex=os.getenv("MC_CORS_ORIGIN_REGEX") or None, allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) -# Register routers -app.include_router(crawler_router, prefix="/api") -app.include_router(data_router, prefix="/api") +# Register routers. +# The auth router stays open -- it is the way in. Everything else under /api +# requires a session. Enforcement is a Depends applied per router rather than +# app-wide middleware, because middleware needs a hand-rolled path allowlist and, +# more importantly, never sees WebSocket scopes at all. +app.include_router(auth_router, prefix="/api") +app.include_router(crawler_router, prefix="/api", dependencies=[Depends(require_auth)]) +app.include_router(data_router, prefix="/api", dependencies=[Depends(require_auth)]) +app.include_router(monitor_router, prefix="/api", dependencies=[Depends(require_auth)]) +app.include_router(settings_router, prefix="/api", dependencies=[Depends(require_auth)]) app.include_router(websocket_router, prefix="/api") @@ -73,9 +149,8 @@ async def serve_frontend(): if os.path.exists(index_path): return FileResponse(index_path) return { - "message": "MediaCrawler WebUI API", + "message": "综合采集平台 API", "version": "1.0.0", - "docs": "/docs", "note": "WebUI not found, please build it first: cd webui && npm run build" } @@ -85,18 +160,21 @@ async def health_check(): return {"status": "ok"} -@app.get("/api/env/check") +@app.get("/api/env/check", dependencies=[Depends(require_auth)]) async def check_environment(): - """Check if MediaCrawler environment is configured correctly""" + """Check whether the crawler environment is configured correctly""" try: - # Run uv run main.py --help command to check environment - # Use PROJECT_ROOT so it works regardless of where uvicorn was started + # Run `main.py --help` to check the environment. + # Resolve the interpreter the same way the crawler manager does, so this + # check can never disagree with how main.py is actually executed. + # Use PROJECT_ROOT so it works regardless of where uvicorn was started. + python_cmd = resolve_python_cmd() if sys.platform == "win32": loop = asyncio.get_running_loop() process = await loop.run_in_executor( None, lambda: subprocess.run( - ["uv", "run", "main.py", "--help"], + [*python_cmd, "main.py", "--help"], capture_output=True, timeout=30.0, cwd=str(PROJECT_ROOT) @@ -105,7 +183,7 @@ async def check_environment(): stdout, stderr = process.stdout, process.stderr # bytes else: process = await asyncio.create_subprocess_exec( - "uv", "run", "main.py", "--help", + *python_cmd, "main.py", "--help", stdout=subprocess.PIPE, stderr=subprocess.PIPE, cwd=str(PROJECT_ROOT) # Project root directory @@ -117,7 +195,8 @@ async def check_environment(): if process.returncode == 0: return { "success": True, - "message": "MediaCrawler environment configured correctly", + "message": "环境配置正确", + "interpreter": describe_interpreter(), "output": stdout.decode("utf-8", errors="ignore")[:500] # Truncate to first 500 characters } else: @@ -136,8 +215,11 @@ async def check_environment(): except FileNotFoundError: return { "success": False, - "message": "uv command not found", - "error": "Please ensure uv is installed and configured in system PATH" + "message": "Python interpreter not found", + "error": ( + "Neither uv nor a usable interpreter was found. Install uv, or create a " + "project virtualenv (.venv) with the requirements installed." + ) } except Exception as e: return { @@ -147,29 +229,29 @@ async def check_environment(): } -@app.get("/api/config/platforms") +@app.get("/api/config/platforms", dependencies=[Depends(require_auth)]) async def get_platforms(): - """Get list of supported platforms""" - return { - "platforms": [ - {"value": "xhs", "label": "Xiaohongshu", "icon": "book-open"}, - {"value": "dy", "label": "Douyin", "icon": "music"}, - {"value": "ks", "label": "Kuaishou", "icon": "video"}, - {"value": "bili", "label": "Bilibili", "icon": "tv"}, - {"value": "wb", "label": "Weibo", "icon": "message-circle"}, - {"value": "tieba", "label": "Baidu Tieba", "icon": "messages-square"}, - {"value": "zhihu", "label": "Zhihu", "icon": "help-circle"}, - ] - } + """Platform capability matrix. + + Returns what each platform's crawler supports (modes, metrics, comment + levels, media) *and* whether the monitoring layer has been wired up for it. + The UI renders its platform switcher and metric columns from this, so the + two are never allowed to drift apart. + """ + from .monitor.platforms import describe_all + + return {"platforms": describe_all()} -@app.get("/api/config/options") +@app.get("/api/config/options", dependencies=[Depends(require_auth)]) async def get_config_options(): """Get all configuration options""" return { "login_types": [ - {"value": "qrcode", "label": "QR Code Login"}, - {"value": "cookie", "label": "Cookie Login"}, + {"value": "qrcode", "label": "扫码登录"}, + # Named for what it now does: the value itself is no longer typed + # here, it is reused from Settings. + {"value": "cookie", "label": "复用已保存的 Cookie"}, ], "crawler_types": [ {"value": "search", "label": "Search Mode"}, @@ -202,4 +284,18 @@ if os.path.exists(WEBUI_DIR): if __name__ == "__main__": - uvicorn.run(app, host="0.0.0.0", port=8080) + # Loopback by default: the safe choice for anyone who has not thought about + # exposure. Set MC_HOST=0.0.0.0 (e.g. in .env) for LAN access. Before this, + # `python -m api.main` bound 0.0.0.0 while the documented `uvicorn api.main:app` + # bound loopback -- two launch paths with different exposure. + host = os.getenv("MC_HOST", "127.0.0.1") + port = int(os.getenv("MC_PORT", "8080")) + + if host not in ("127.0.0.1", "localhost", "::1"): + print( + f"[综合采集平台] 监听 {host}:{port},局域网内其他机器可访问。\n" + f"[综合采集平台] 已启用密码鉴权;如需暴露到可信网络之外,请走 HTTPS 反向代理。", + flush=True, + ) + + uvicorn.run(app, host=host, port=port) diff --git a/api/monitor/__init__.py b/api/monitor/__init__.py new file mode 100644 index 0000000..be1c291 --- /dev/null +++ b/api/monitor/__init__.py @@ -0,0 +1,19 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/monitor/__init__.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Scheduled monitoring layer: repeated crawls with change detection.""" diff --git a/api/monitor/app_settings.py b/api/monitor/app_settings.py new file mode 100644 index 0000000..54ab30e --- /dev/null +++ b/api/monitor/app_settings.py @@ -0,0 +1,395 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/monitor/app_settings.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Application settings, declared once and rendered from that declaration. + +Every setting carries a **scope**, which is the whole reason this is not a flat +list: + +* ``platform`` -- each platform keeps its own copy. A cookie obviously differs, + but so do crawl pacing and proxies: what is safe on one platform is a rate + limit on another. Stored as ``platform.

.``. +* ``system`` -- one value for the whole instance. The notification webhook is + a single group chat, and the scheduler has a single active-hours window, so + scoping those per platform would be a fiction. + +The registry is the single source of truth: the API returns it and the Settings +page builds its form from it, so adding a setting does not mean editing a +matching list on the frontend. + +Two rules carry over from how the cookie and webhook were already handled: + +* **Secrets are never returned.** A sensitive key comes back as + ``{present, length, updated_at}``, never as a value. +* **Update is partial.** Only keys present in the request are written, so a form + that does not resubmit a secret cannot silently wipe it. +""" + +from dataclasses import dataclass +from typing import Any, Dict, List, Optional + +from sqlalchemy.ext.asyncio import AsyncSession + +from .platforms import PLATFORM_XHS +from .settings import ( + delete_setting, + get_setting, + platform_key, + set_setting, + system_key, +) + +SCOPE_PLATFORM = "platform" +SCOPE_SYSTEM = "system" + +TYPE_BOOL = "bool" +TYPE_INT = "int" +TYPE_STR = "str" +TYPE_SECRET = "secret" + +# Mirrors config/base_config.py. Nothing is written until the operator changes +# something; an unset value simply means "pass no CLI flag, so the config file's +# value applies". +_DEFAULT_SLEEP_SEC = 2 + + +@dataclass +class SettingSpec: + name: str + scope: str + type: str + label: str + help: str = "" + default: Any = None + minimum: Optional[int] = None + maximum: Optional[int] = None + choices: Optional[List[str]] = None + affects_new_runs: bool = True + + def key(self, platform: str = PLATFORM_XHS) -> str: + if self.scope == SCOPE_SYSTEM: + return system_key(self.name) + return platform_key(platform, self.name) + + +SETTING_SPECS: List[SettingSpec] = [ + # --- 平台设置 ----------------------------------------------------------- + SettingSpec( + name="cookie", + scope=SCOPE_PLATFORM, + type=TYPE_SECRET, + label="登录 Cookie", + help="定时监控必须持久化登录态。建议先手动登录一次再粘贴 Cookie。", + ), + SettingSpec( + name="default_interval_minutes", + scope=SCOPE_PLATFORM, + type=TYPE_INT, + label="新任务默认采集间隔(分钟)", + help="仅影响新建任务时的默认值,不会改动已有任务。", + default=360, + minimum=30, + maximum=10080, + ), + SettingSpec( + name="default_max_notes", + scope=SCOPE_PLATFORM, + type=TYPE_INT, + label="默认单轮作品上限", + default=20, + minimum=1, + maximum=500, + ), + SettingSpec( + name="default_max_comments", + scope=SCOPE_PLATFORM, + type=TYPE_INT, + label="默认每篇评论抓取条数", + help="接口无时间排序,只取平台默认排序的前 N 条;N 越大越容易发现新评论。", + default=50, + minimum=1, + maximum=500, + ), + SettingSpec( + name="enable_sub_comments", + scope=SCOPE_PLATFORM, + type=TYPE_BOOL, + label="抓取二级评论", + help="请求量显著增加,风控风险更高。", + default=False, + ), + SettingSpec( + name="crawl_sleep_sec", + scope=SCOPE_PLATFORM, + type=TYPE_INT, + label="请求间隔(秒)", + help="调大更慢但更不容易触发平台限流。各平台风控容忍度不同,故分开配置。", + default=_DEFAULT_SLEEP_SEC, + minimum=0, + maximum=600, + ), + SettingSpec( + name="enable_ip_proxy", + scope=SCOPE_PLATFORM, + type=TYPE_BOOL, + label="启用 IP 代理", + default=False, + ), + SettingSpec( + name="proxy_provider", + scope=SCOPE_PLATFORM, + type=TYPE_STR, + label="代理提供方", + default="kuaidaili", + choices=["kuaidaili", "wandouhttp", "static"], + ), + SettingSpec( + name="proxy_pool_count", + scope=SCOPE_PLATFORM, + type=TYPE_INT, + label="代理 IP 池大小", + default=2, + minimum=1, + maximum=100, + ), + SettingSpec( + name="static_proxy_url", + scope=SCOPE_PLATFORM, + type=TYPE_STR, + label="静态代理地址", + help="仅当提供方选择 static 时使用,格式 http://host:port", + default="", + ), + # --- 系统设置 ----------------------------------------------------------- + SettingSpec( + name="wecom_webhook", + scope=SCOPE_SYSTEM, + type=TYPE_SECRET, + label="企业微信 Webhook", + help="企业微信群机器人地址。所有平台共用同一个群,只有开了推送开关的任务才会发消息。", + ), + SettingSpec( + name="active_hours_start", + scope=SCOPE_SYSTEM, + type=TYPE_INT, + label="活跃时段开始(小时)", + help="只在此时段内触发定时采集。默认 0–23 即全天;支持跨午夜,如 22–6。", + default=0, + minimum=0, + maximum=23, + affects_new_runs=False, + ), + SettingSpec( + name="active_hours_end", + scope=SCOPE_SYSTEM, + type=TYPE_INT, + label="活跃时段结束(小时)", + default=23, + minimum=0, + maximum=23, + affects_new_runs=False, + ), +] + +SPECS_BY_NAME = {spec.name: spec for spec in SETTING_SPECS} + +# Managed by their own endpoints; never writable through the settings API. +# Suffix-matched rather than enumerated, because the cookie bookkeeping keys +# exist once per platform. +_HIDDEN_KEY_SUFFIXES = (".cookie_updated_at", ".cookie_last_ok_at") +_HIDDEN_KEYS = {"auth_password_hash", "auth_password_updated_at"} + + +def _is_hidden(key: str) -> bool: + return key in _HIDDEN_KEYS or key.endswith(_HIDDEN_KEY_SUFFIXES) + + +class SettingValidationError(ValueError): + """Raised for a value the registry will not accept.""" + + +def _coerce(spec: SettingSpec, raw: Any) -> Any: + if spec.type == TYPE_SECRET: + return str(raw) if raw is not None else "" + + if spec.type == TYPE_BOOL: + if isinstance(raw, bool): + return raw + text = str(raw).strip().lower() + if text in ("1", "true", "yes", "y", "on"): + return True + if text in ("0", "false", "no", "n", "off", ""): + return False + raise SettingValidationError(f"{spec.label}: 需要是/否") + + if spec.type == TYPE_INT: + try: + value = int(raw) + except (TypeError, ValueError): + raise SettingValidationError(f"{spec.label}: 需要整数") + if spec.minimum is not None and value < spec.minimum: + raise SettingValidationError(f"{spec.label}: 不能小于 {spec.minimum}") + if spec.maximum is not None and value > spec.maximum: + raise SettingValidationError(f"{spec.label}: 不能大于 {spec.maximum}") + return value + + value = str(raw) if raw is not None else "" + if spec.choices and value not in spec.choices: + raise SettingValidationError(f"{spec.label}: 只能是 {'/'.join(spec.choices)}") + return value + + +def _decode(spec: SettingSpec, raw: Optional[str]) -> Any: + if raw is None: + return spec.default + if spec.type == TYPE_BOOL: + return raw.strip().lower() in ("1", "true", "yes", "y", "on") + if spec.type == TYPE_INT: + try: + return int(raw) + except ValueError: + return spec.default + return raw + + +def _encode(spec: SettingSpec, value: Any) -> str: + if spec.type == TYPE_BOOL: + return "true" if value else "false" + return str(value) + + +def _describe(spec: SettingSpec, platform: str) -> Dict[str, Any]: + return { + "key": spec.key(platform), + "name": spec.name, + "scope": spec.scope, + "type": spec.type, + "label": spec.label, + "help": spec.help, + "default": spec.default, + "minimum": spec.minimum, + "maximum": spec.maximum, + "choices": spec.choices, + "affects_new_runs": spec.affects_new_runs, + } + + +async def get_all(session: AsyncSession, platform: str = PLATFORM_XHS) -> Dict[str, Any]: + """Every editable setting for one platform, plus the system-wide ones. + + Secrets come back masked, never in the clear. + """ + values: Dict[str, Any] = {} + secrets: Dict[str, Any] = {} + + for spec in SETTING_SPECS: + key = spec.key(platform) + raw = await get_setting(session, key) + + if spec.type == TYPE_SECRET: + secrets[key] = {"present": bool(raw), "length": len(raw or "")} + else: + values[key] = _decode(spec, raw) + + return { + "platform": platform, + "values": values, + "secrets": secrets, + "specs": [_describe(spec, platform) for spec in SETTING_SPECS], + } + + +def _spec_for_key(key: str, platform: str) -> Optional[SettingSpec]: + """Resolve a full key back to its spec, rejecting keys for another platform.""" + for spec in SETTING_SPECS: + if spec.key(platform) == key: + return spec + return None + + +async def update( + session: AsyncSession, payload: Dict[str, Any], platform: str = PLATFORM_XHS +) -> List[str]: + """Apply a partial update. Returns the keys that changed. + + Only keys present in ``payload`` are touched: a form that omits a secret must + not blank it. Keys belonging to a different platform are rejected rather than + silently written somewhere unexpected. + """ + changed: List[str] = [] + + for key, raw in payload.items(): + if _is_hidden(key): + continue + + spec = _spec_for_key(key, platform) + if spec is None: + raise SettingValidationError(f"未知的设置项:{key}") + + # An explicit empty string clears a secret -- that is how the UI removes + # one. For everything else it is just a value. + if spec.type == TYPE_SECRET and raw == "": + await delete_setting(session, key) + changed.append(key) + continue + + value = _coerce(spec, raw) + await set_setting(session, key, _encode(spec, value)) + changed.append(key) + + return changed + + +async def get_value( + session: AsyncSession, + name: str, + platform: str = PLATFORM_XHS, + fallback: Any = None, +) -> Any: + """Read one typed setting for internal callers (the runner, the scheduler).""" + spec = SPECS_BY_NAME.get(name) + if spec is None: + return fallback + raw = await get_setting(session, spec.key(platform)) + if raw is None: + return spec.default if fallback is None else fallback + return _decode(spec, raw) + + +async def defaults(session: AsyncSession, platform: str = PLATFORM_XHS) -> Dict[str, Any]: + """Defaults applied when creating a task on this platform. + + This is what makes the Settings page govern new tasks: the create endpoint + falls back to these for anything the caller omits. + """ + return { + "interval_minutes": int( + await get_value(session, "default_interval_minutes", platform, 360) + ), + "max_notes_count": int(await get_value(session, "default_max_notes", platform, 20)), + "max_comments_count": int( + await get_value(session, "default_max_comments", platform, 50) + ), + } + + +async def active_hours(session: AsyncSession) -> tuple[int, int]: + """The (start, end) hour window for scheduled runs. System-wide.""" + start = await get_value(session, "active_hours_start", fallback=0) + end = await get_value(session, "active_hours_end", fallback=23) + return int(start), int(end) diff --git a/api/monitor/db.py b/api/monitor/db.py new file mode 100644 index 0000000..df0ac81 --- /dev/null +++ b/api/monitor/db.py @@ -0,0 +1,293 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/monitor/db.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Database engine for the monitoring layer. + +**MySQL** by default (see ``config/db_config.py`` and ``.env``), with SQLite kept +as an option so the test suite can run without a reachable server. + +Three things here exist because of specific MySQL 5.7 behaviour: + +* **utf8mb4 is forced per table.** This instance's server *and* the target schema + default to ``latin1``; relying on either would mangle or reject Chinese text. + The charset is set on every table rather than on the database, so it holds no + matter what the schema default is. +* **Connections are recycled.** The monitor runs for weeks, and MySQL drops idle + connections after ``wait_timeout`` (8h by default). Without ``pool_recycle`` and + ``pool_pre_ping`` the first query after a quiet night fails with "server has + gone away". +* **The connected schema is asserted at startup.** A misconfigured database name + is caught immediately instead of silently writing to the wrong schema. + +Only the configured schema is ever touched: no ``CREATE DATABASE``, no ``USE``, +no cross-schema query. +""" + +import os +import sys +from contextlib import asynccontextmanager +from pathlib import Path +from typing import AsyncIterator, Optional + +from sqlalchemy import event, text +from sqlalchemy.ext.asyncio import ( + AsyncEngine, + AsyncSession, + async_sessionmaker, + create_async_engine, +) + +from .models import MonitorBase + +PROJECT_ROOT = Path(__file__).parent.parent.parent +DATA_DIR = PROJECT_ROOT / "data" +DEFAULT_SQLITE_PATH = DATA_DIR / "monitor.db" + +# Load .env here as well as in api/main.py: this module is imported directly by +# scripts and tests, and a configuration that only applies when the server is the +# entry point is a trap. load_dotenv does not override real environment variables. +from dotenv import load_dotenv + +load_dotenv(PROJECT_ROOT / ".env") + +# Kept identical to config/db_config.py's defaults so one .env drives both the +# monitor database and the crawler's own DB output. +MYSQL_HOST = lambda: os.getenv("MYSQL_DB_HOST", "localhost") # noqa: E731 +MYSQL_PORT = lambda: int(os.getenv("MYSQL_DB_PORT", "3306")) # noqa: E731 +MYSQL_USER = lambda: os.getenv("MYSQL_DB_USER", "root") # noqa: E731 +MYSQL_PWD = lambda: os.getenv("MYSQL_DB_PWD", "") # noqa: E731 +MYSQL_DB_NAME = lambda: os.getenv("MYSQL_DB_NAME", "mediacrawler") # noqa: E731 + +_engine: Optional[AsyncEngine] = None +_session_factory: Optional[async_sessionmaker[AsyncSession]] = None +# None means "resolve from the environment" (MySQL). Tests set a SQLite URL. +_db_url: Optional[str] = None +_expected_schema: Optional[str] = None + + +def resolve_db_url() -> str: + """Build the connection URL. MySQL unless overridden.""" + if _db_url is not None: + return _db_url + + from urllib.parse import quote_plus + + user = quote_plus(MYSQL_USER()) + password = quote_plus(MYSQL_PWD()) + host = MYSQL_HOST() + port = MYSQL_PORT() + name = MYSQL_DB_NAME() + return f"mysql+aiomysql://{user}:{password}@{host}:{port}/{name}?charset=utf8mb4" + + +def is_mysql() -> bool: + return resolve_db_url().startswith("mysql") + + +def set_sqlite_path(path: Path) -> None: + """Point the layer at SQLite. Used by the test suite only.""" + global _db_url, _engine, _session_factory, _expected_schema + _db_url = f"sqlite+aiosqlite:///{Path(path)}" + _engine = None + _session_factory = None + _expected_schema = None + + +def set_db_url(url: str, expected_schema: Optional[str] = None) -> None: + """Point the layer at an explicit URL. ``expected_schema`` enables the guard.""" + global _db_url, _engine, _session_factory, _expected_schema + _db_url = url + _engine = None + _session_factory = None + _expected_schema = expected_schema + + +def expected_schema() -> Optional[str]: + """The schema the connection must be using, if the guard applies.""" + if _expected_schema is not None: + return _expected_schema + return MYSQL_DB_NAME() if is_mysql() else None + + +def get_engine() -> AsyncEngine: + global _engine + if _engine is None: + url = resolve_db_url() + kwargs: dict = {"future": True} + + if url.startswith("mysql"): + # Recycle well inside MySQL's default 8h wait_timeout, and verify a + # pooled connection before handing it out. + kwargs.update(pool_recycle=3600, pool_pre_ping=True, pool_size=5, max_overflow=5) + kwargs["connect_args"] = {"charset": "utf8mb4"} + else: + Path(url.split("///", 1)[-1]).parent.mkdir(parents=True, exist_ok=True) + + _engine = create_async_engine(url, **kwargs) + + if url.startswith("sqlite"): + + @event.listens_for(_engine.sync_engine, "connect") + def _set_sqlite_pragmas(dbapi_connection, _connection_record): # pragma: no cover + cursor = dbapi_connection.cursor() + cursor.execute("PRAGMA journal_mode=WAL") + cursor.execute("PRAGMA foreign_keys=ON") + cursor.close() + + return _engine + + +def get_session_factory() -> async_sessionmaker[AsyncSession]: + global _session_factory + if _session_factory is None: + _session_factory = async_sessionmaker( + bind=get_engine(), + class_=AsyncSession, + expire_on_commit=False, + ) + return _session_factory + + +@asynccontextmanager +async def get_session() -> AsyncIterator[AsyncSession]: + """Transactional session. Commits on success, rolls back on error.""" + factory = get_session_factory() + async with factory() as session: + try: + yield session + await session.commit() + except Exception: + await session.rollback() + raise + + +async def _assert_correct_schema(conn) -> None: + """Refuse to run against anything but the configured schema. + + A guard, not the guarantee: the real protection is a MySQL account scoped to + this one schema (see UPSTREAM.md). This catches the ordinary mistake of a + wrong database name in configuration, before a single row is written. + """ + if not is_mysql(): + return + + expected = expected_schema() + if not expected: + return + + current = (await conn.execute(text("SELECT DATABASE()"))).scalar() + if current is None: + raise RuntimeError( + f"数据库连接未选定 schema,期望 {expected!r}。请检查 MYSQL_DB_NAME。" + ) + # lower_case_table_names=1 makes names case-insensitive server-side. + if current.lower() != expected.lower(): + raise RuntimeError( + f"连接的库是 {current!r},但配置要求 {expected!r}。" + f"为避免误写其它库,已拒绝启动。" + ) + print(f"[monitor.db] 已连接 MySQL schema: {current}", flush=True) + + +async def init_db() -> None: + """Create missing tables, then run the small in-place migrations.""" + engine = get_engine() + async with engine.begin() as conn: + await _assert_correct_schema(conn) + await conn.run_sync(MonitorBase.metadata.create_all) + await _ensure_columns(conn) + await _migrate_setting_keys(conn) + + +# Columns added to a table after it may already exist. ``create_all`` only +# creates missing *tables*, so new columns need an explicit ALTER TABLE. +_ADDED_COLUMNS: dict[str, list[tuple[str, str]]] = { + "monitor_task": [ + ("notify_enabled", "BOOLEAN NOT NULL DEFAULT 0"), + ("last_notified_at", "BIGINT NULL"), + ], +} + + +async def _existing_columns(conn, table: str) -> set[str]: + if is_mysql(): + rows = await conn.execute( + text( + "SELECT COLUMN_NAME FROM information_schema.COLUMNS " + "WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = :t" + ), + {"t": table}, + ) + return {row[0] for row in rows} + + rows = await conn.execute(text(f"PRAGMA table_info({table})")) + return {row[1] for row in rows} + + +async def _ensure_columns(conn) -> None: + for table, columns in _ADDED_COLUMNS.items(): + existing = await _existing_columns(conn, table) + if not existing: + # Table did not exist before this run; create_all built it complete. + continue + for name, ddl in columns: + if name not in existing: + await conn.execute(text(f"ALTER TABLE {table} ADD COLUMN {name} {ddl}")) + + +async def _migrate_setting_keys(conn) -> None: + """Move pre-namespacing setting keys to their scoped names. + + Idempotent: the legacy row is only renamed when the new key is absent, so an + operator's later value is never overwritten. + """ + from .models import LEGACY_SETTING_KEY_RENAMES + + for legacy, scoped in LEGACY_SETTING_KEY_RENAMES.items(): + exists = ( + await conn.execute( + text("SELECT 1 FROM monitor_setting WHERE `key` = :k"), {"k": legacy} + ) + ).first() + if not exists: + continue + + already = ( + await conn.execute( + text("SELECT 1 FROM monitor_setting WHERE `key` = :k"), {"k": scoped} + ) + ).first() + if already: + # Both present: the scoped one is authoritative; drop the stale row. + await conn.execute( + text("DELETE FROM monitor_setting WHERE `key` = :k"), {"k": legacy} + ) + continue + + await conn.execute( + text("UPDATE monitor_setting SET `key` = :new WHERE `key` = :old"), + {"new": scoped, "old": legacy}, + ) + + +async def dispose_engine() -> None: + global _engine, _session_factory + if _engine is not None: + await _engine.dispose() + _engine = None + _session_factory = None diff --git a/api/monitor/ingest.py b/api/monitor/ingest.py new file mode 100644 index 0000000..a49d797 --- /dev/null +++ b/api/monitor/ingest.py @@ -0,0 +1,589 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/monitor/ingest.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Turn one run's crawled jsonl into snapshots and change events. + +Pure-ish and offline testable: give it a directory of jsonl files, a run row and +a session, and it does the diffing. No network, no browser. + +Correctness notes that drive the code below: + +* Counts arrive as strings and may be abbreviated ("1.2万", "3亿"). A value that + cannot be parsed is stored as NULL, never 0 -- 0 would forge a large negative + delta on the next comparison. +* The comment endpoint has no time-sort, so only the platform's top-N window is + ever visible. A comment we have not seen before is therefore split into + "posted since last run" vs "seen for the first time", rather than claiming the + former always. +* A bad cookie does not make the crawler exit non-zero; it exits 0 having + fetched nothing. That is detected here as a suspected auth failure. +""" + +import json +import re +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, Dict, List, Optional + +from sqlalchemy import func, select +from sqlalchemy.ext.asyncio import AsyncSession + +from tools.time_util import get_current_timestamp + +from .platforms import PLATFORM_XHS +from .models import ( + EVENT_AUTH_FAILURE, + EVENT_METRIC_DELTA, + EVENT_NEW_COMMENT_POSTED, + EVENT_NEW_COMMENT_SEEN, + EVENT_NEW_NOTE, + EVENT_NO_DATA, + EVENT_RUN_FAILED, + MonitorComment, + MonitorEvent, + MonitorNote, + MonitorNoteMetric, + MonitorRun, + MonitorTask, + RUN_FAILED, + RUN_PARTIAL, + RUN_SUCCESS, +) + +_COUNT_UNITS = { + "": 1, + "万": 10_000, + "w": 10_000, + "W": 10_000, + "k": 1_000, + "K": 1_000, + "亿": 100_000_000, +} +_COUNT_RE = re.compile(r"^([\d.]+)\s*([万wWkK亿]?)$") + +# Metric fields shared by the snapshot table and the delta comparison. +_METRIC_FIELDS = ("liked_count", "comment_count", "collected_count", "share_count") + + +def parse_count(value: Any) -> Optional[int]: + """Parse an XHS interaction count into an int, or None if unintelligible. + + Handles plain numbers, thousands separators, and the Chinese abbreviations + the platform actually returns ("1.2万" -> 12000, "3亿" -> 300000000). + """ + if value is None or isinstance(value, bool): + return None + if isinstance(value, int): + return value + if isinstance(value, float): + return int(value) + + text = str(value).strip().replace(",", "").replace(" ", "") + if not text: + return None + + match = _COUNT_RE.match(text) + if not match: + return None + + try: + number = float(match.group(1)) + except ValueError: + return None + + return int(number * _COUNT_UNITS.get(match.group(2), 1)) + + +# Windows reports hard process failures as NTSTATUS values, which surface in the +# UI as meaningless large integers (e.g. 3221225794 = 0xC0000142). Translating +# the ones we actually see saves the reader a hex-decoding detour. +_WINDOWS_EXIT_REASONS = { + 0xC0000005: "进程访问冲突 (ACCESS_VIOLATION)", + 0xC00000FD: "栈溢出 (STACK_OVERFLOW)", + 0xC000013A: "进程被中断(控制台关闭或 Ctrl+C)", + 0xC0000142: "进程初始化失败 (STATUS_DLL_INIT_FAILED),属启动环境异常,重启服务后重试", + 0xC0000409: "栈缓冲区溢出 (STACK_BUFFER_OVERRUN)", +} + + +def describe_exit_code(code: int) -> str: + """Render an exit code so a human can act on it.""" + unsigned = code & 0xFFFFFFFF if code < 0 else code + reason = _WINDOWS_EXIT_REASONS.get(unsigned) + if reason: + return f"Crawler exited with code {code} (0x{unsigned:08X}): {reason}" + return f"Crawler exited with code {code}" + + +@dataclass +class IngestResult: + status: str + notes_fetched: int = 0 + comments_fetched: int = 0 + new_notes: int = 0 + new_comments: int = 0 + is_baseline: bool = False + error: Optional[str] = None + events: List[str] = field(default_factory=list) + + +def _read_jsonl(path: Path) -> List[Dict[str, Any]]: + """Read a jsonl file, skipping blank or malformed lines.""" + records: List[Dict[str, Any]] = [] + if not path.exists(): + return records + + with path.open("r", encoding="utf-8") as handle: + for line in handle: + line = line.strip() + if not line: + continue + try: + item = json.loads(line) + except json.JSONDecodeError: + continue + if isinstance(item, dict): + records.append(item) + return records + + +def find_run_files( + out_dir: Path, platform: str = PLATFORM_XHS +) -> tuple[List[Path], List[Path]]: + """Locate the contents/comments jsonl files a run produced. + + The crawler writes ``{save_data_path}/{platform}/jsonl/{type}_{item}_{date}.jsonl``. + Glob rather than reconstructing the name: both the crawler type and the date + are runtime-dependent. Returns lists because a crawl crossing midnight + produces one file per day. + """ + jsonl_dir = out_dir / platform / "jsonl" + if not jsonl_dir.is_dir(): + return [], [] + + return ( + sorted(jsonl_dir.glob("*_contents_*.jsonl")), + sorted(jsonl_dir.glob("*_comments_*.jsonl")), + ) + + +async def _emit( + session: AsyncSession, + run: MonitorRun, + event_type: str, + title: str, + *, + severity: str = "info", + target_kind: str = "", + target_id: str = "", + payload: Optional[Dict[str, Any]] = None, +) -> None: + session.add( + MonitorEvent( + task_id=run.task_id, + run_id=run.id, + type=event_type, + severity=severity, + target_kind=target_kind, + target_id=target_id, + title=title, + payload_json=json.dumps(payload or {}, ensure_ascii=False), + created_at=get_current_timestamp(), + ) + ) + + +async def _previous_run_started_at( + session: AsyncSession, task_id: int, run_id: int +) -> Optional[int]: + """Started-at of the most recent earlier successful run, in ms.""" + return await session.scalar( + select(MonitorRun.started_at) + .where( + MonitorRun.task_id == task_id, + MonitorRun.id != run_id, + MonitorRun.status.in_((RUN_SUCCESS, RUN_PARTIAL)), + MonitorRun.started_at.is_not(None), + # Same reasoning as _count_prior_successes: an empty run is a useless + # reference point for "was this comment posted since last time?". + MonitorRun.notes_fetched > 0, + ) + .order_by(MonitorRun.id.desc()) + .limit(1) + ) + + +# How far back to look for proof that the stored login still works. +_AUTH_PROOF_WINDOW_MS = 6 * 60 * 60 * 1000 + + +async def _another_task_succeeded_recently(session: AsyncSession, task_id: int) -> bool: + """Whether a different task fetched data recently, proving the login is valid.""" + since = get_current_timestamp() - _AUTH_PROOF_WINDOW_MS + count = await session.scalar( + select(func.count()) + .select_from(MonitorRun) + .where( + MonitorRun.task_id != task_id, + MonitorRun.status == RUN_SUCCESS, + MonitorRun.started_at.is_not(None), + MonitorRun.started_at >= since, + ) + ) + return bool(count) + + +async def _count_prior_successes(session: AsyncSession, task_id: int, run_id: int) -> int: + return ( + await session.scalar( + select(func.count()) + .select_from(MonitorRun) + .where( + MonitorRun.task_id == task_id, + MonitorRun.id != run_id, + MonitorRun.status.in_((RUN_SUCCESS, RUN_PARTIAL)), + # A run that fetched nothing established no baseline. Without this + # check the first run that actually works after a failed one looks + # like a flood of newly discovered works. + MonitorRun.notes_fetched > 0, + ) + ) + ) or 0 + + +async def _ingest_notes( + session: AsyncSession, + run: MonitorRun, + records: List[Dict[str, Any]], + is_baseline: bool, +) -> int: + """Upsert notes, write metric snapshots, and emit new-note/delta events.""" + now = get_current_timestamp() + new_count = 0 + + for record in records: + note_id = record.get("note_id") + if not note_id: + continue + + note = await session.scalar( + select(MonitorNote).where( + MonitorNote.task_id == run.task_id, + MonitorNote.note_id == note_id, + ) + ) + + title = (record.get("title") or "")[:500] + raw_images = record.get("image_list") or "" + cover = raw_images.split(",")[0] if raw_images else "" + + if note is None: + note = MonitorNote( + task_id=run.task_id, + note_id=note_id, + title=title, + note_url=record.get("note_url") or "", + cover=cover, + creator_hash=record.get("creator_hash") or "", + source_kind=record.get("type") or "", + published_at=_as_int(record.get("time")), + first_seen_run_id=run.id, + first_seen_at=now, + last_seen_run_id=run.id, + last_seen_at=now, + ) + session.add(note) + new_count += 1 + if not is_baseline: + await _emit( + session, + run, + EVENT_NEW_NOTE, + f"新作品:{title or note_id}", + target_kind="note", + target_id=note_id, + payload={"note_id": note_id, "title": title}, + ) + else: + # Only refresh descriptive fields; seen-tracking is updated below. + if title: + note.title = title + note.last_seen_run_id = run.id + note.last_seen_at = now + + await _snapshot_metrics(session, run, note_id, record, now, is_baseline) + + return new_count + + +def _as_int(value: Any) -> Optional[int]: + try: + return int(value) + except (TypeError, ValueError): + return None + + +async def _snapshot_metrics( + session: AsyncSession, + run: MonitorRun, + note_id: str, + record: Dict[str, Any], + now: int, + is_baseline: bool, +) -> None: + """Write this run's metric snapshot and report any change vs the previous one.""" + previous = await session.scalar( + select(MonitorNoteMetric) + .where( + MonitorNoteMetric.task_id == run.task_id, + MonitorNoteMetric.note_id == note_id, + MonitorNoteMetric.run_id != run.id, + ) + .order_by(MonitorNoteMetric.run_id.desc()) + .limit(1) + ) + + parsed = {name: parse_count(record.get(name)) for name in _METRIC_FIELDS} + + session.add( + MonitorNoteMetric( + task_id=run.task_id, + note_id=note_id, + run_id=run.id, + captured_at=now, + liked_count=parsed["liked_count"], + comment_count=parsed["comment_count"], + collected_count=parsed["collected_count"], + share_count=parsed["share_count"], + raw_liked_count=str(record.get("liked_count") or ""), + raw_comment_count=str(record.get("comment_count") or ""), + raw_collected_count=str(record.get("collected_count") or ""), + raw_share_count=str(record.get("share_count") or ""), + ) + ) + + if previous is None or is_baseline: + return + + deltas = {} + for name in _METRIC_FIELDS: + old, new = getattr(previous, name), parsed[name] + # A None on either side means the value was unparseable; skip rather + # than report a bogus change. + if old is None or new is None or old == new: + continue + deltas[name] = {"from": old, "to": new, "delta": new - old} + + if deltas: + summary = "、".join( + f"{_metric_label(name)} {info['from']}→{info['to']}" + for name, info in deltas.items() + ) + await _emit( + session, + run, + EVENT_METRIC_DELTA, + f"互动数据变化:{summary}", + target_kind="note", + target_id=note_id, + payload={"note_id": note_id, "deltas": deltas}, + ) + + +def _metric_label(name: str) -> str: + return { + "liked_count": "点赞", + "comment_count": "评论", + "collected_count": "收藏", + "share_count": "分享", + }.get(name, name) + + +async def _ingest_comments( + session: AsyncSession, + run: MonitorRun, + records: List[Dict[str, Any]], + is_baseline: bool, + previous_run_started_at: Optional[int], +) -> int: + """Upsert comments and emit events for ones never seen before.""" + now = get_current_timestamp() + new_count = 0 + + for record in records: + comment_id = record.get("comment_id") + note_id = record.get("note_id") + if not comment_id or not note_id: + continue + + exists = await session.scalar( + select(MonitorComment.id).where( + MonitorComment.task_id == run.task_id, + MonitorComment.note_id == note_id, + MonitorComment.comment_id == comment_id, + ) + ) + if exists is not None: + continue + + create_time = _as_int(record.get("create_time")) + session.add( + MonitorComment( + task_id=run.task_id, + note_id=note_id, + comment_id=comment_id, + content=(record.get("content") or "")[:2000], + nickname=record.get("nickname") or "", + creator_hash=record.get("creator_hash") or "", + create_time=create_time, + like_count=parse_count(record.get("like_count")), + sub_comment_count=_as_int(record.get("sub_comment_count")) or 0, + parent_comment_id=record.get("parent_comment_id") or "", + first_seen_run_id=run.id, + first_seen_at=now, + ) + ) + new_count += 1 + + if is_baseline: + continue + + # Without a time-sorted comment API we can only observe the top-N window, + # so distinguish a genuinely new comment from one that just surfaced. + posted = ( + create_time is not None + and previous_run_started_at is not None + and create_time > previous_run_started_at + ) + await _emit( + session, + run, + EVENT_NEW_COMMENT_POSTED if posted else EVENT_NEW_COMMENT_SEEN, + f"{'新评论' if posted else '新出现评论'}:{(record.get('content') or '')[:60]}", + target_kind="note", + target_id=note_id, + payload={ + "note_id": note_id, + "comment_id": comment_id, + "create_time": create_time, + "nickname": record.get("nickname") or "", + }, + ) + + return new_count + + +async def ingest_run( + session: AsyncSession, + run: MonitorRun, + task: MonitorTask, + out_dir: Path, +) -> IngestResult: + """Ingest one finished run and return what changed. + + Sets ``run.status``, ``run.is_baseline`` and the counters on the run row. + On a failed or untrustworthy run nothing is diffed -- the "seen" sets only + ever grow, so a partial run must never be allowed to look like deletions. + """ + # A non-zero exit is a genuine crash: trust nothing this run produced. + if run.exit_code not in (0, None): + run.status = RUN_FAILED + run.error_message = describe_exit_code(run.exit_code) + await _emit( + session, + run, + EVENT_RUN_FAILED, + f"采集进程异常退出(code={run.exit_code})", + severity="error", + payload={"exit_code": run.exit_code, "detail": run.error_message}, + ) + return IngestResult(status=RUN_FAILED, error=run.error_message) + + contents_paths, comment_paths = find_run_files(out_dir, task.platform) + contents = [record for path in contents_paths for record in _read_jsonl(path)] + comments = [record for path in comment_paths for record in _read_jsonl(path)] + + run.notes_fetched = len(contents) + run.comments_fetched = len(comments) + + # A bad cookie does NOT fail the process: XHS cookie login is never validated, + # so an unauthenticated session just returns zero notes with exit 0 -- and + # usually does not even create an output file. Treating that as "the creator + # posted nothing" would silently hide login outages, which is exactly what + # monitoring exists to catch. + if not contents: + run.status = RUN_PARTIAL + + # Blaming the cookie is only honest if nothing else is authenticating. + # A sibling task that just succeeded proves the login works, so the + # fault is with this target (bad/expired per-creator token, an empty + # account, or a page-structure change). + if await _another_task_succeeded_recently(session, run.task_id): + run.error_message = ( + "Crawler produced no notes for this target, but other tasks " + "succeeded recently, so the login is probably fine" + ) + await _emit( + session, + run, + EVENT_NO_DATA, + "本次未抓到任何作品:其他任务近期采集正常,登录态应该没问题,请检查该目标是否有效", + severity="warning", + payload={"out_dir": str(out_dir)}, + ) + else: + run.error_message = "Crawler produced no notes; the login cookie may have expired" + await _emit( + session, + run, + EVENT_AUTH_FAILURE, + "疑似登录态失效:本次未抓到任何作品,请检查 Cookie", + severity="error", + payload={"out_dir": str(out_dir)}, + ) + + return IngestResult( + status=RUN_PARTIAL, + error=run.error_message, + comments_fetched=len(comments), + ) + + is_baseline = await _count_prior_successes(session, run.task_id, run.id) == 0 + run.is_baseline = is_baseline + run.status = RUN_SUCCESS + run.error_message = None + + previous_started_at = ( + None if is_baseline else await _previous_run_started_at(session, run.task_id, run.id) + ) + + result = IngestResult( + status=RUN_SUCCESS, + notes_fetched=len(contents), + comments_fetched=len(comments), + is_baseline=is_baseline, + ) + result.new_notes = await _ingest_notes(session, run, contents, is_baseline) + if task.enable_comments: + result.new_comments = await _ingest_comments( + session, run, comments, is_baseline, previous_started_at + ) + + run.new_notes = result.new_notes + run.new_comments = result.new_comments + return result diff --git a/api/monitor/migrate_from_sqlite.py b/api/monitor/migrate_from_sqlite.py new file mode 100644 index 0000000..544cc64 --- /dev/null +++ b/api/monitor/migrate_from_sqlite.py @@ -0,0 +1,177 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/monitor/migrate_from_sqlite.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""One-off: copy the monitoring database from SQLite into MySQL. + + python -m api.monitor.migrate_from_sqlite [--source data/monitor.db] [--dry-run] + +Primary keys are preserved rather than reassigned, because rows in +``monitor_note`` / ``monitor_comment`` / ``monitor_run`` reference ``task_id``; +letting MySQL auto-assign new ids would silently break those links. + +Refuses to run against a target that already holds data unless ``--force`` is +given, so a second accidental run cannot double everything up. +""" + +import argparse +import sqlite3 +import sys +from pathlib import Path +from typing import Any, Dict, List + +PROJECT_ROOT = Path(__file__).parent.parent.parent + +# Insert order matters: monitor_target and monitor_run carry real foreign keys to +# monitor_task, so the parent rows have to land first. +TABLES_IN_ORDER = [ + "monitor_task", + "monitor_target", + "monitor_run", + "monitor_note", + "monitor_note_metric", + "monitor_comment", + "monitor_event", + "monitor_setting", + "auth_session", +] + + +def read_sqlite(path: Path) -> Dict[str, List[Dict[str, Any]]]: + if not path.exists(): + raise SystemExit(f"找不到源库:{path}") + + connection = sqlite3.connect(path) + connection.row_factory = sqlite3.Row + try: + existing = { + row[0] + for row in connection.execute( + "SELECT name FROM sqlite_master WHERE type='table'" + ) + } + data: Dict[str, List[Dict[str, Any]]] = {} + for table in TABLES_IN_ORDER: + if table not in existing: + continue + rows = [dict(row) for row in connection.execute(f"SELECT * FROM {table}")] + if rows: + data[table] = rows + return data + finally: + connection.close() + + +def migrate(source: Path, dry_run: bool, force: bool) -> None: + import pymysql + + from . import db as monitor_db + + data = read_sqlite(source) + if not data: + print("源库里没有可迁移的数据。") + return + + print("源库内容:") + for table, rows in data.items(): + print(f" {table:22} {len(rows)} 行") + + url = monitor_db.resolve_db_url() + if not url.startswith("mysql"): + raise SystemExit(f"目标不是 MySQL:{url}") + + connection = pymysql.connect( + host=monitor_db.MYSQL_HOST(), + port=monitor_db.MYSQL_PORT(), + user=monitor_db.MYSQL_USER(), + password=monitor_db.MYSQL_PWD(), + database=monitor_db.MYSQL_DB_NAME(), + charset="utf8mb4", + autocommit=False, + ) + + try: + with connection.cursor() as cursor: + # Never write outside the configured schema. + cursor.execute("SELECT DATABASE()") + current = cursor.fetchone()[0] + expected = monitor_db.MYSQL_DB_NAME() + if current.lower() != expected.lower(): + raise SystemExit( + f"当前连接的是 {current!r},配置要求 {expected!r};已中止。" + ) + + occupied = [] + for table in data: + cursor.execute(f"SELECT COUNT(*) FROM `{table}`") + if cursor.fetchone()[0]: + occupied.append(table) + + if occupied and not force: + raise SystemExit( + "目标库已有数据:" + ", ".join(occupied) + "\n" + "加 --force 才会继续(会与现有数据并存,造成重复)。" + ) + + if dry_run: + print("\n[试运行] 未写入任何数据。") + return + + total = 0 + for table, rows in data.items(): + columns = list(rows[0].keys()) + column_sql = ", ".join(f"`{c}`" for c in columns) + placeholders = ", ".join(["%s"] * len(columns)) + statement = ( + f"INSERT INTO `{table}` ({column_sql}) VALUES ({placeholders})" + ) + cursor.executemany( + statement, [[row[c] for c in columns] for row in rows] + ) + total += len(rows) + print(f" 已写入 {table:22} {len(rows)} 行") + + connection.commit() + print(f"\n完成,共迁移 {total} 行。") + print("提示:源 SQLite 文件仍在原处,确认无误后自行删除。") + + except Exception: + connection.rollback() + raise + finally: + connection.close() + + +def main(argv: List[str] | None = None) -> int: + parser = argparse.ArgumentParser(description="把监控库从 SQLite 迁到 MySQL") + parser.add_argument( + "--source", + default=str(PROJECT_ROOT / "data" / "monitor.db"), + help="SQLite 源文件路径", + ) + parser.add_argument("--dry-run", action="store_true", help="只检查,不写入") + parser.add_argument( + "--force", action="store_true", help="目标库已有数据时也继续" + ) + args = parser.parse_args(argv) + + migrate(Path(args.source), args.dry_run, args.force) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/api/monitor/models.py b/api/monitor/models.py new file mode 100644 index 0000000..f3e439d --- /dev/null +++ b/api/monitor/models.py @@ -0,0 +1,367 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/monitor/models.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Monitoring layer data model. + +Lives in its own SQLite database (``data/monitor.db``) with its own declarative +Base, deliberately separate from the crawler's ``database/models.py``. The +crawler's DB store overwrites ``liked_count`` and friends in place on every +re-crawl, so it cannot answer "how did this note change?". These tables keep the +history the crawler throws away. + +All timestamps are epoch **milliseconds** (BigInteger), matching the project's +own ``tools.time_util.get_current_timestamp()`` convention. Using ints +throughout avoids naive/aware datetime mixing bugs. +""" + +from typing import Optional + +from sqlalchemy import ( + BigInteger, + Boolean, + ForeignKey, + Integer, + String, + Text, + UniqueConstraint, +) +from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column, relationship + + +class MonitorBase(DeclarativeBase): + """Declarative base for the monitoring database.""" + + +# Run statuses +RUN_PENDING = "pending" +RUN_RUNNING = "running" +RUN_SUCCESS = "success" +RUN_PARTIAL = "partial" +RUN_FAILED = "failed" +RUN_TIMEOUT = "timeout" +RUN_INTERRUPTED = "interrupted" + +# Event types +EVENT_NEW_NOTE = "new_note" +EVENT_NEW_COMMENT_POSTED = "new_comment_posted" +EVENT_NEW_COMMENT_SEEN = "new_comment_seen" +EVENT_METRIC_DELTA = "metric_delta" +EVENT_RUN_FAILED = "run_failed" +EVENT_AUTH_FAILURE = "suspected_auth_failure" +# A run that completed cleanly yet fetched nothing, where the login is provably +# fine because another task just succeeded with it. The target, not the cookie, +# is what needs looking at. +EVENT_NO_DATA = "no_data_found" + +# Task modes. One subprocess handles exactly one crawler type, so a task is +# either creator-driven or note-driven -- never both. +MODE_CREATOR = "creator" +MODE_NOTE = "note" + + +class MonitorTask(MonitorBase): + """One monitored schedule: a set of targets plus an interval.""" + + __tablename__ = "monitor_task" + + id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) + name: Mapped[str] = mapped_column(String(200), nullable=False) + platform: Mapped[str] = mapped_column(String(32), nullable=False, default="xhs") + mode: Mapped[str] = mapped_column(String(16), nullable=False) + enabled: Mapped[bool] = mapped_column(Boolean, nullable=False, default=True) + interval_minutes: Mapped[int] = mapped_column(Integer, nullable=False, default=360) + + # Crawl window knobs, mirrored onto each run's CLI flags. + max_notes_count: Mapped[int] = mapped_column(Integer, nullable=False, default=20) + enable_comments: Mapped[bool] = mapped_column(Boolean, nullable=False, default=True) + max_comments_count: Mapped[int] = mapped_column(Integer, nullable=False, default=50) + run_timeout_seconds: Mapped[int] = mapped_column(Integer, nullable=False, default=3600) + + # Push notifications are opt-in per task. A task list that all pushes to one + # webhook turns noisy fast, so silence is the default. + notify_enabled: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False) + + # Scheduler state. Persisted so the schedule survives an API restart. + next_run_at: Mapped[Optional[int]] = mapped_column(BigInteger, index=True) + last_run_at: Mapped[Optional[int]] = mapped_column(BigInteger) + last_status: Mapped[str] = mapped_column(String(32), nullable=False, default="idle") + last_error: Mapped[Optional[str]] = mapped_column(Text) + # Lets the UI answer "why did I not get a push for this run?". + last_notified_at: Mapped[Optional[int]] = mapped_column(BigInteger) + + created_at: Mapped[int] = mapped_column(BigInteger, nullable=False) + updated_at: Mapped[int] = mapped_column(BigInteger, nullable=False) + + targets: Mapped[list["MonitorTarget"]] = relationship( + back_populates="task", + cascade="all, delete-orphan", + lazy="selectin", + ) + + +class MonitorTarget(MonitorBase): + """One watched creator or note belonging to a task. + + ``external_id`` is the stable identity (XHS user_id / note_id). It is kept + separate from ``xsec_token`` on purpose: tokens expire within weeks, so + treating a tokenised URL as the primary key would make every long-running + task fail eventually. + """ + + __tablename__ = "monitor_target" + __table_args__ = ( + UniqueConstraint("task_id", "kind", "external_id", name="uq_monitor_target"), + ) + + id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) + task_id: Mapped[int] = mapped_column( + ForeignKey("monitor_task.id", ondelete="CASCADE"), nullable=False, index=True + ) + kind: Mapped[str] = mapped_column(String(16), nullable=False) + external_id: Mapped[str] = mapped_column(String(128), nullable=False) + xsec_token: Mapped[str] = mapped_column(String(512), nullable=False, default="") + xsec_source: Mapped[str] = mapped_column(String(64), nullable=False, default="") + raw_value: Mapped[str] = mapped_column(Text, nullable=False, default="") + label: Mapped[str] = mapped_column(String(200), nullable=False, default="") + enabled: Mapped[bool] = mapped_column(Boolean, nullable=False, default=True) + created_at: Mapped[int] = mapped_column(BigInteger, nullable=False) + + task: Mapped["MonitorTask"] = relationship(back_populates="targets") + + +class MonitorRun(MonitorBase): + """One subprocess execution. The run history in the UI is this table.""" + + __tablename__ = "monitor_run" + + id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) + task_id: Mapped[int] = mapped_column( + ForeignKey("monitor_task.id", ondelete="CASCADE"), nullable=False, index=True + ) + trigger: Mapped[str] = mapped_column(String(16), nullable=False, default="scheduled") + status: Mapped[str] = mapped_column(String(16), nullable=False, default=RUN_PENDING, index=True) + phase: Mapped[str] = mapped_column(String(16), nullable=False) + + # Where this run's jsonl landed. Each run gets its own directory because the + # crawler's file writer names output by date only. + save_data_path: Mapped[str] = mapped_column(Text, nullable=False, default="") + + queued_at: Mapped[int] = mapped_column(BigInteger, nullable=False) + not_before: Mapped[int] = mapped_column(BigInteger, nullable=False, default=0) + started_at: Mapped[Optional[int]] = mapped_column(BigInteger) + finished_at: Mapped[Optional[int]] = mapped_column(BigInteger) + # BigInteger, not Integer: Windows reports failures as unsigned 32-bit + # NTSTATUS values (0xC0000142 = 3221225794), which overflow MySQL's signed + # INT. SQLite's dynamic typing hid this until the data was migrated. + exit_code: Mapped[Optional[int]] = mapped_column(BigInteger) + + notes_fetched: Mapped[int] = mapped_column(Integer, nullable=False, default=0) + comments_fetched: Mapped[int] = mapped_column(Integer, nullable=False, default=0) + new_notes: Mapped[int] = mapped_column(Integer, nullable=False, default=0) + new_comments: Mapped[int] = mapped_column(Integer, nullable=False, default=0) + + # The very first successful run of a task establishes the baseline: every + # note is "new" at that point, so emitting events would be pure noise. + is_baseline: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False) + + # Window actually used, so the UI can be honest that comments are the top N + # in the platform's own ordering rather than a complete set. + max_comments_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0) + + error_message: Mapped[Optional[str]] = mapped_column(Text) + + +class MonitorNote(MonitorBase): + """A note ever seen by a task, plus when it was first/last seen. + + Grain is (task, note) so the same note tracked by two tasks stays independent. + """ + + __tablename__ = "monitor_note" + __table_args__ = ( + UniqueConstraint("task_id", "note_id", name="uq_monitor_note"), + ) + + id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) + task_id: Mapped[int] = mapped_column( + ForeignKey("monitor_task.id", ondelete="CASCADE"), nullable=False, index=True + ) + note_id: Mapped[str] = mapped_column(String(128), nullable=False, index=True) + title: Mapped[str] = mapped_column(Text, nullable=False, default="") + note_url: Mapped[str] = mapped_column(Text, nullable=False, default="") + cover: Mapped[str] = mapped_column(Text, nullable=False, default="") + creator_hash: Mapped[str] = mapped_column(String(64), nullable=False, default="") + source_kind: Mapped[str] = mapped_column(String(16), nullable=False, default="") + published_at: Mapped[Optional[int]] = mapped_column(BigInteger) + + first_seen_run_id: Mapped[Optional[int]] = mapped_column(Integer) + first_seen_at: Mapped[int] = mapped_column(BigInteger, nullable=False) + last_seen_run_id: Mapped[Optional[int]] = mapped_column(Integer) + last_seen_at: Mapped[int] = mapped_column(BigInteger, nullable=False) + + +class MonitorNoteMetric(MonitorBase): + """One metric snapshot per (task, note, run) -- the time series. + + Raw strings are kept alongside the parsed integers so a mis-parsed "1.2万" + can always be audited after the fact. + """ + + __tablename__ = "monitor_note_metric" + __table_args__ = ( + UniqueConstraint("task_id", "note_id", "run_id", name="uq_note_metric"), + ) + + id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) + task_id: Mapped[int] = mapped_column(Integer, nullable=False, index=True) + note_id: Mapped[str] = mapped_column(String(128), nullable=False, index=True) + run_id: Mapped[int] = mapped_column(Integer, nullable=False, index=True) + captured_at: Mapped[int] = mapped_column(BigInteger, nullable=False) + + # NULL (not 0) when the platform value could not be parsed: storing 0 would + # forge a large negative delta on the next comparison. + liked_count: Mapped[Optional[int]] = mapped_column(Integer) + comment_count: Mapped[Optional[int]] = mapped_column(Integer) + collected_count: Mapped[Optional[int]] = mapped_column(Integer) + share_count: Mapped[Optional[int]] = mapped_column(Integer) + + raw_liked_count: Mapped[str] = mapped_column(String(64), nullable=False, default="") + raw_comment_count: Mapped[str] = mapped_column(String(64), nullable=False, default="") + raw_collected_count: Mapped[str] = mapped_column(String(64), nullable=False, default="") + raw_share_count: Mapped[str] = mapped_column(String(64), nullable=False, default="") + + +class MonitorComment(MonitorBase): + """A comment ever seen by a task. + + The (task, note, comment) uniqueness gives idempotent dedup across runs for + free -- re-running the same crawl cannot double-count. + """ + + __tablename__ = "monitor_comment" + __table_args__ = ( + UniqueConstraint("task_id", "note_id", "comment_id", name="uq_monitor_comment"), + ) + + id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) + task_id: Mapped[int] = mapped_column(Integer, nullable=False, index=True) + note_id: Mapped[str] = mapped_column(String(128), nullable=False, index=True) + comment_id: Mapped[str] = mapped_column(String(128), nullable=False) + content: Mapped[str] = mapped_column(Text, nullable=False, default="") + nickname: Mapped[str] = mapped_column(String(200), nullable=False, default="") + creator_hash: Mapped[str] = mapped_column(String(64), nullable=False, default="") + # Platform-stated publish time. Used to distinguish a genuinely new comment + # from one that merely entered the visible top-N window this run. + create_time: Mapped[Optional[int]] = mapped_column(BigInteger) + like_count: Mapped[Optional[int]] = mapped_column(Integer) + sub_comment_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0) + parent_comment_id: Mapped[str] = mapped_column(String(128), nullable=False, default="") + + first_seen_run_id: Mapped[Optional[int]] = mapped_column(Integer) + first_seen_at: Mapped[int] = mapped_column(BigInteger, nullable=False) + + +class MonitorEvent(MonitorBase): + """Append-only change feed. This is what the dashboard reads.""" + + __tablename__ = "monitor_event" + + id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) + task_id: Mapped[int] = mapped_column(Integer, nullable=False, index=True) + run_id: Mapped[Optional[int]] = mapped_column(Integer, index=True) + type: Mapped[str] = mapped_column(String(32), nullable=False, index=True) + severity: Mapped[str] = mapped_column(String(16), nullable=False, default="info") + target_kind: Mapped[str] = mapped_column(String(16), nullable=False, default="") + target_id: Mapped[str] = mapped_column(String(128), nullable=False, default="") + title: Mapped[str] = mapped_column(Text, nullable=False, default="") + payload_json: Mapped[str] = mapped_column(Text, nullable=False, default="{}") + created_at: Mapped[int] = mapped_column(BigInteger, nullable=False, index=True) + is_read: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False) + + +class MonitorSetting(MonitorBase): + """Key/value store. Holds the XHS cookie for unattended runs.""" + + __tablename__ = "monitor_setting" + + key: Mapped[str] = mapped_column(String(64), primary_key=True) + value: Mapped[str] = mapped_column(Text, nullable=False, default="") + updated_at: Mapped[int] = mapped_column(BigInteger, nullable=False) + + +class AuthSession(MonitorBase): + """A WebUI login session. + + Only the SHA-256 of the token is stored, never the token itself -- a leaked + database therefore does not hand over live sessions. This mirrors the + existing posture of never returning the XHS cookie or webhook value. + + A stateful table (rather than a signed stateless token) is what makes "log + out" and "password changed" take effect immediately. + """ + + __tablename__ = "auth_session" + + token_hash: Mapped[str] = mapped_column(String(64), primary_key=True) + created_at: Mapped[int] = mapped_column(BigInteger, nullable=False) + expires_at: Mapped[int] = mapped_column(BigInteger, nullable=False, index=True) + last_seen_at: Mapped[int] = mapped_column(BigInteger, nullable=False) + + +SETTING_AUTH_PASSWORD_HASH = "auth_password_hash" +SETTING_AUTH_PASSWORD_UPDATED_AT = "auth_password_updated_at" + +# Settings are namespaced by scope: `platform.

.` for values each +# platform keeps its own copy of, `system.` for values shared across all of +# them. Key builders live in settings.py. +SETTING_WECOM_WEBHOOK = "system.wecom_webhook" + +# Pre-namespacing keys, kept only so the startup migration can find and move +# them. Nothing should read these directly. +LEGACY_SETTING_KEY_RENAMES = { + # Pre-batch-2 flat keys. + "xhs_cookie": "platform.xhs.cookie", + "xhs_cookie_updated_at": "platform.xhs.cookie_updated_at", + "xhs_cookie_last_ok_at": "platform.xhs.cookie_last_ok_at", + "wecom_webhook": "system.wecom_webhook", + # Batch-2 keys, before settings gained a scope. Those values belonged to + # Xiaohongshu because it was the only platform, so they migrate to its scope; + # the two scheduling keys were always instance-wide. + "collect.default_interval_minutes": "platform.xhs.default_interval_minutes", + "collect.default_max_notes": "platform.xhs.default_max_notes", + "collect.default_max_comments": "platform.xhs.default_max_comments", + "collect.enable_sub_comments": "platform.xhs.enable_sub_comments", + "collect.crawl_sleep_sec": "platform.xhs.crawl_sleep_sec", + "collect.active_hours_start": "system.active_hours_start", + "collect.active_hours_end": "system.active_hours_end", + "proxy.enable_ip_proxy": "platform.xhs.enable_ip_proxy", + "proxy.provider": "platform.xhs.proxy_provider", + "proxy.pool_count": "platform.xhs.proxy_pool_count", + "proxy.static_proxy_url": "platform.xhs.static_proxy_url", +} + + +# utf8mb4 is forced on every table rather than left to the schema default: this +# deployment's MySQL server *and* the target database both default to latin1, +# which would mangle or reject Chinese text. Setting it per table means it holds +# regardless of what the schema default happens to be. +# +# Must run after every model is declared, hence the end of the module. +for _table in MonitorBase.metadata.tables.values(): + _table.kwargs["mysql_charset"] = "utf8mb4" + _table.kwargs["mysql_collate"] = "utf8mb4_unicode_ci" diff --git a/api/monitor/notify.py b/api/monitor/notify.py new file mode 100644 index 0000000..67d462a --- /dev/null +++ b/api/monitor/notify.py @@ -0,0 +1,200 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/monitor/notify.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Push notifications via a WeCom (企业微信) group robot webhook. + +Two rules shape this module: + +* **One message per run, not per event.** A run that finds twenty new notes must + produce one summary, not twenty pushes. +* **A failed push never fails the crawl.** Notification is best-effort: the run's + data is already committed by the time we get here, so every error is logged + and swallowed. +""" + +import json +from typing import Optional + +import httpx +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from tools.time_util import get_current_timestamp + +from .models import ( + EVENT_AUTH_FAILURE, + EVENT_NEW_NOTE, + EVENT_NO_DATA, + EVENT_RUN_FAILED, + MonitorEvent, + MonitorRun, + MonitorTask, +) +from .settings import get_setting + +# Short on purpose: the scheduler awaits the run, so a hanging webhook would +# stall every other task behind it. +WEBHOOK_TIMEOUT_SECONDS = 10.0 + +# Only these event types are worth interrupting someone for. NO_DATA is included +# because a run that fetched nothing at all is always anomalous -- a creator +# always has *some* notes -- even when the login is not the culprit. +NOTIFIABLE_EVENT_TYPES = ( + EVENT_AUTH_FAILURE, + EVENT_RUN_FAILED, + EVENT_NO_DATA, + EVENT_NEW_NOTE, +) + +# WeCom markdown is a limited subset; coloured text is the one bit of flair it +# supports and it makes failures stand out in a busy group chat. +_COLOR_WARNING = "warning" +_COLOR_INFO = "info" + + +async def send_wecom(webhook_url: str, content: str) -> tuple[bool, str]: + """Post a markdown message to a WeCom group robot. + + Returns (ok, detail) rather than raising, so callers can surface the reason + in the UI when the user clicks "send test". + """ + if not webhook_url: + return False, "Webhook 未配置" + + payload = {"msgtype": "markdown", "markdown": {"content": content}} + + try: + async with httpx.AsyncClient(timeout=WEBHOOK_TIMEOUT_SECONDS) as client: + response = await client.post(webhook_url, json=payload) + response.raise_for_status() + body = response.json() + except httpx.HTTPError as exc: + return False, f"请求失败:{exc}" + except json.JSONDecodeError: + return False, "返回内容不是合法 JSON,请检查 Webhook 地址" + + # WeCom answers 200 with a non-zero errcode on failure. + errcode = body.get("errcode") + if errcode != 0: + return False, f"企业微信返回 errcode={errcode} {body.get('errmsg', '')}" + + return True, "发送成功" + + +async def get_webhook_url(session: AsyncSession) -> str: + from .models import SETTING_WECOM_WEBHOOK + + return (await get_setting(session, SETTING_WECOM_WEBHOOK)) or "" + + +async def build_run_message( + session: AsyncSession, + task: MonitorTask, + run: MonitorRun, +) -> Optional[str]: + """Compose one markdown summary for a finished run, or None if nothing to say.""" + events = list( + ( + await session.scalars( + select(MonitorEvent) + .where( + MonitorEvent.run_id == run.id, + MonitorEvent.type.in_(NOTIFIABLE_EVENT_TYPES), + ) + .order_by(MonitorEvent.id) + ) + ).all() + ) + if not events: + return None + + failures = [ + e for e in events if e.type in (EVENT_AUTH_FAILURE, EVENT_RUN_FAILED, EVENT_NO_DATA) + ] + new_notes = [e for e in events if e.type == EVENT_NEW_NOTE] + + lines: list[str] = [] + + if failures: + # Word the header from what actually happened, not from whether new notes + # accompanied it: a login outage usually brings no new notes either. + unavailable = any(e.type == EVENT_NO_DATA for e in failures) and not any( + e.type in (EVENT_AUTH_FAILURE, EVENT_RUN_FAILED) for e in failures + ) + header = "监控任务未抓到数据" if unavailable else "监控任务异常" + lines.append(f"**⚠️ {header}:{task.name}**") + for event in failures: + lines.append(f'> {event.title}') + else: + lines.append(f"**📢 监控任务有新作品:{task.name}**") + + if new_notes: + lines.append(f"> 新增作品 **{len(new_notes)}** 篇") + # Cap the listing: a first-ever run or a long gap can produce a lot, and + # a wall of text is worse than a count. + for event in new_notes[:10]: + payload = _load_payload(event.payload_json) + title = payload.get("title") or event.target_id + note_id = payload.get("note_id") or event.target_id + url = f"https://www.xiaohongshu.com/explore/{note_id}" + lines.append(f"> [{title}]({url})") + if len(new_notes) > 10: + lines.append(f"> …等共 {len(new_notes)} 篇") + + if run.is_baseline: + lines.append("> (首轮基线,未计入新增统计)") + + return "\n".join(lines) + + +def _load_payload(raw: str) -> dict: + try: + payload = json.loads(raw or "{}") + except json.JSONDecodeError: + return {} + return payload if isinstance(payload, dict) else {} + + +async def notify_run(session: AsyncSession, task: MonitorTask, run: MonitorRun) -> Optional[str]: + """Push a summary for a finished run if the task opted in. + + Returns the message that was sent, or None. Never raises. + """ + try: + if not task.notify_enabled: + return None + + webhook_url = await get_webhook_url(session) + if not webhook_url: + return None + + message = await build_run_message(session, task, run) + if not message: + return None + + ok, detail = await send_wecom(webhook_url, message) + if not ok: + print(f"[monitor.notify] task {task.id} push failed: {detail}") + return None + + task.last_notified_at = get_current_timestamp() + return message + + except Exception as exc: # pragma: no cover - notification must never break a run + print(f"[monitor.notify] unexpected error: {exc}") + return None diff --git a/api/monitor/platforms.py b/api/monitor/platforms.py new file mode 100644 index 0000000..e605a50 --- /dev/null +++ b/api/monitor/platforms.py @@ -0,0 +1,197 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/monitor/platforms.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Platform capability matrix. + +The single source of truth for what each platform can do. The UI renders its +platform switcher and metric columns from this, and the API validates against +it. + +Two distinct things are recorded here, and conflating them would be misleading: + +* ``crawler_modes`` / ``metrics`` / ``comment_levels`` / ``media`` describe what + the upstream crawler module actually supports. These were read out of the + platform modules, not assumed -- all seven implement search/detail/creator; + the real differences are in which interaction metrics they capture. +* ``monitor_wired`` says whether the *monitoring layer* has been hooked up. It + currently covers only Xiaohongshu: ``runner.py`` pins the platform, + ``ingest.py`` reads a fixed ``xhs/jsonl`` directory, and ``service.py`` only + parses Xiaohongshu target URLs. + +A platform can therefore be fully crawlable by upstream and still not usable for +monitoring, which is exactly the state of the other six today. +""" + +from typing import Any, Dict, List, Optional + +PLATFORM_XHS = "xhs" + +PLATFORM_LABELS = { + "xhs": "小红书", + "dy": "抖音", + "ks": "快手", + "bili": "B站", + "wb": "微博", + "tieba": "贴吧", + "zhihu": "知乎", +} + +# Interaction metrics each platform's store actually persists. Xiaohongshu has no +# play count or danmaku; Bilibili has both and the widest set; Kuaishou carries +# no comment/share/collect at all; Tieba stores only reply counts. +PLATFORM_CAPABILITIES: Dict[str, Dict[str, Any]] = { + "xhs": { + "crawler_modes": ["search", "detail", "creator"], + "metrics": ["liked_count", "comment_count", "collected_count", "share_count"], + "comment_levels": 2, + "media": True, + "monitor_wired": True, + }, + "dy": { + "crawler_modes": ["search", "detail", "creator"], + "metrics": ["liked_count", "comment_count", "collected_count", "share_count"], + "comment_levels": 2, + "media": True, + "monitor_wired": False, + }, + "ks": { + "crawler_modes": ["search", "detail", "creator"], + # No comment/share/collect in the Kuaishou store; sub-comments are stored + # flat with no parent link and carry no like count. + "metrics": ["liked_count", "view_count"], + "comment_levels": 1, + "media": True, + "monitor_wired": False, + }, + "bili": { + "crawler_modes": ["search", "detail", "creator"], + "metrics": [ + "liked_count", + "video_play_count", + "video_danmaku", + "comment_count", + "video_favorite_count", + "video_coin_count", + "video_share_count", + ], + "comment_levels": 2, + "media": True, + "monitor_wired": False, + }, + "wb": { + "crawler_modes": ["search", "detail", "creator"], + # Weibo has no collect count, and its comment count field is named + # differently in the model. + "metrics": ["liked_count", "comments_count", "shared_count"], + "comment_levels": 2, + "media": True, + "monitor_wired": False, + }, + "tieba": { + "crawler_modes": ["search", "detail", "creator"], + "metrics": ["total_replay_num", "total_replay_page"], + "comment_levels": 2, + "media": False, + "monitor_wired": False, + }, + "zhihu": { + "crawler_modes": ["search", "detail", "creator"], + "metrics": ["voteup_count", "comment_count"], + "comment_levels": 2, + "media": False, + "monitor_wired": False, + }, +} + +METRIC_LABELS = { + "liked_count": "点赞", + "comment_count": "评论", + "collected_count": "收藏", + "share_count": "分享", + "view_count": "播放", + "video_play_count": "播放", + "video_danmaku": "弹幕", + "video_favorite_count": "收藏", + "video_coin_count": "投币", + "video_share_count": "分享", + "comments_count": "评论", + "shared_count": "转发", + "total_replay_num": "回复数", + "total_replay_page": "回复页数", + "voteup_count": "赞同", +} + +# Monitoring modes, mapped to the CLI crawler types upstream understands. +MONITOR_MODE_CREATOR = "creator" +MONITOR_MODE_NOTE = "note" +CLI_TYPE_FOR_MODE = { + MONITOR_MODE_CREATOR: "creator", + MONITOR_MODE_NOTE: "detail", +} + + +class UnsupportedPlatformError(ValueError): + """Raised for an unknown platform, or one the monitor layer cannot run.""" + + +def all_platforms() -> List[str]: + return list(PLATFORM_CAPABILITIES) + + +def is_known(platform: str) -> bool: + return platform in PLATFORM_CAPABILITIES + + +def is_monitor_wired(platform: str) -> bool: + return bool(PLATFORM_CAPABILITIES.get(platform, {}).get("monitor_wired")) + + +def describe(platform: str) -> Optional[Dict[str, Any]]: + capability = PLATFORM_CAPABILITIES.get(platform) + if capability is None: + return None + return { + "value": platform, + "label": PLATFORM_LABELS.get(platform, platform), + **capability, + "metric_labels": { + metric: METRIC_LABELS.get(metric, metric) for metric in capability["metrics"] + }, + } + + +def describe_all() -> List[Dict[str, Any]]: + return [describe(platform) for platform in all_platforms()] + + +def ensure_runnable(platform: str) -> None: + """Validate a platform for a monitoring task. + + An unwired platform is rejected outright rather than accepted and left to + silently produce nothing -- the same silent-failure shape that made a valid + creator look like an expired login earlier. + """ + if not is_known(platform): + raise UnsupportedPlatformError( + f"未知平台:{platform}(支持:{', '.join(all_platforms())})" + ) + if not is_monitor_wired(platform): + label = PLATFORM_LABELS.get(platform, platform) + raise UnsupportedPlatformError( + f"{label}的爬虫已支持,但监控层尚未接通,暂时无法创建监控任务。" + ) diff --git a/api/monitor/report.py b/api/monitor/report.py new file mode 100644 index 0000000..7467a2b --- /dev/null +++ b/api/monitor/report.py @@ -0,0 +1,203 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/monitor/report.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Cross-task reporting: what grew, and what is new, over a date range. + +Two families of numbers that answer different questions and are therefore kept +as separate columns: + +* **互动增量** — Σ(current − previous) across the selected notes. "How many likes + did this set of notes gain?" +* **新增内容** — count of newly discovered notes and comments. "How much new + material showed up?" + +The per-day interaction delta is defined as *last value on the day* minus *last +value before the day* (0 when the note was first seen on that day). That keeps +growth from a note's first observation counted once, rather than smeared across +every later day. + +Aggregation runs in Python over the snapshots rather than as one large SQL +query: the per-note-per-day baseline lookup is a windowed operation that SQLite +expresses awkwardly, and the row counts here are small enough that clarity is +worth more than the query planner. +""" + +from bisect import bisect_right +from datetime import date, datetime, time, timedelta +from typing import Any, Dict, Iterable, List, Optional, Sequence + +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from .models import MonitorComment, MonitorNote, MonitorNoteMetric + +METRIC_FIELDS = ("liked_count", "comment_count", "collected_count", "share_count") + +METRIC_LABELS = { + "liked_count": "点赞", + "comment_count": "评论", + "collected_count": "收藏", + "share_count": "分享", +} + + +def day_bounds(day: date) -> tuple[int, int]: + """Inclusive epoch-millisecond bounds for a local calendar day.""" + start = datetime.combine(day, time.min) + end = datetime.combine(day, time.max) + return int(start.timestamp() * 1000), int(end.timestamp() * 1000) + + +def iter_days(start: date, end: date) -> List[date]: + days = [] + cursor = start + while cursor <= end: + days.append(cursor) + cursor += timedelta(days=1) + return days + + +def compute_daily_rows( + series_by_note: Dict[str, List[tuple[int, Dict[str, Optional[int]]]]], + notes_per_day: Dict[date, int], + comments_per_day: Dict[date, int], + days: Sequence[date], +) -> List[Dict[str, Any]]: + """Pure aggregation. ``series_by_note`` must be sorted by timestamp ascending.""" + prepared = {note_id: ([ts for ts, _ in points], points) for note_id, points in series_by_note.items()} + + rows: List[Dict[str, Any]] = [] + for day in days: + day_start, day_end = day_bounds(day) + totals = {field: 0 for field in METRIC_FIELDS} + # Records *which* metric could not be compared, not just that something + # could not. A blanket flag loses all value the moment one permanently + # unparseable field makes every row "incomplete". + partial_metrics: set[str] = set() + + for times, points in prepared.values(): + end_index = bisect_right(times, day_end) - 1 + if end_index < 0: + # Not yet tracked on this day. + continue + + end_values = points[end_index][1] + start_index = bisect_right(times, day_start - 1) - 1 + # No earlier snapshot means the note first appeared in this window, + # so it starts from zero -- all of its count is genuinely new. + start_values = ( + points[start_index][1] if start_index >= 0 else {f: 0 for f in METRIC_FIELDS} + ) + + for field in METRIC_FIELDS: + end_value, start_value = end_values.get(field), start_values.get(field) + if end_value is None or start_value is None: + # An unparseable count on either side makes the delta unknown; + # skipping beats reporting a fabricated number. + partial_metrics.add(field) + continue + totals[field] += end_value - start_value + + row: Dict[str, Any] = { + "date": day.isoformat(), + "new_notes": notes_per_day.get(day, 0), + "new_comments": comments_per_day.get(day, 0), + "partial_metrics": sorted(partial_metrics), + } + row.update({f"{field}_delta": value for field, value in totals.items()}) + rows.append(row) + + return rows + + +async def build_report( + session: AsyncSession, + task_ids: Optional[Iterable[int]], + start_day: date, + end_day: date, +) -> Dict[str, Any]: + """Daily rows plus totals for the selected tasks over the given date range.""" + start_ms, _ = day_bounds(start_day) + _, end_ms = day_bounds(end_day) + + scope = list(task_ids) if task_ids else None + days = iter_days(start_day, end_day) + + # Fetch every snapshot up to the range end: the delta on the first day needs + # the last value from *before* the range, so a lower bound would be wrong. + metric_stmt = select(MonitorNoteMetric).where(MonitorNoteMetric.captured_at <= end_ms) + if scope is not None: + metric_stmt = metric_stmt.where(MonitorNoteMetric.task_id.in_(scope)) + metric_stmt = metric_stmt.order_by(MonitorNoteMetric.note_id, MonitorNoteMetric.run_id) + + series_by_note: Dict[str, List[tuple[int, Dict[str, Optional[int]]]]] = {} + included_note_ids: set[str] = set() + for snapshot in (await session.scalars(metric_stmt)).all(): + included_note_ids.add(snapshot.note_id) + series_by_note.setdefault(snapshot.note_id, []).append( + ( + snapshot.captured_at, + {field: getattr(snapshot, field) for field in METRIC_FIELDS}, + ) + ) + + note_stmt = select(MonitorNote.first_seen_at).where( + MonitorNote.first_seen_at >= start_ms, MonitorNote.first_seen_at <= end_ms + ) + if scope is not None: + note_stmt = note_stmt.where(MonitorNote.task_id.in_(scope)) + + comment_stmt = select(MonitorComment.first_seen_at).where( + MonitorComment.first_seen_at >= start_ms, MonitorComment.first_seen_at <= end_ms + ) + if scope is not None: + comment_stmt = comment_stmt.where(MonitorComment.task_id.in_(scope)) + + notes_per_day = _count_by_day((await session.scalars(note_stmt)).all()) + comments_per_day = _count_by_day((await session.scalars(comment_stmt)).all()) + + rows = compute_daily_rows(series_by_note, notes_per_day, comments_per_day, days) + + totals = { + "new_notes": sum(row["new_notes"] for row in rows), + "new_comments": sum(row["new_comments"] for row in rows), + } + for field in METRIC_FIELDS: + totals[f"{field}_delta"] = sum(row[f"{field}_delta"] for row in rows) + + return { + "start_date": start_day.isoformat(), + "end_date": end_day.isoformat(), + "task_ids": scope, + "rows": rows, + "totals": totals, + "note_count": len(included_note_ids), + "has_partial_data": any(row["partial_metrics"] for row in rows), + "partial_metrics": sorted({field for row in rows for field in row["partial_metrics"]}), + "metric_labels": METRIC_LABELS, + } + + +def _count_by_day(timestamps: Iterable[Optional[int]]) -> Dict[date, int]: + counts: Dict[date, int] = {} + for ts in timestamps: + if ts is None: + continue + day = datetime.fromtimestamp(ts / 1000).date() + counts[day] = counts.get(day, 0) + 1 + return counts diff --git a/api/monitor/runner.py b/api/monitor/runner.py new file mode 100644 index 0000000..d0895aa --- /dev/null +++ b/api/monitor/runner.py @@ -0,0 +1,277 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/monitor/runner.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Execute a single monitoring run: build the command, wait, then ingest. + +Runs reuse ``CrawlerManager`` so that monitor crawls share the existing +single-subprocess guarantee and their logs stream to the existing Terminal +component over the existing log WebSocket. +""" + +import asyncio +import os +from pathlib import Path +from typing import Iterable, List, Optional + +from tools.time_util import get_current_timestamp + +from ..schemas import ( + CrawlerStartRequest, + CrawlerTypeEnum, + LoginTypeEnum, + PlatformEnum, + SaveDataOptionEnum, +) +from ..services import crawler_manager +from . import app_settings, notify +from .db import get_session +from .ingest import IngestResult, ingest_run +from .models import ( + MODE_CREATOR, + RUN_FAILED, + RUN_PENDING, + RUN_RUNNING, + RUN_TIMEOUT, + MonitorRun, + MonitorTarget, + MonitorTask, +) +from .settings import get_cookie, mark_cookie_ok + +PROJECT_ROOT = Path(__file__).parent.parent.parent +MONITOR_RUNS_DIR = PROJECT_ROOT / "data" / "monitor_runs" + +# Monitor platform ids align with PlatformEnum's values, but mapping explicitly +# beats relying on that coincidence. +_PLATFORM_ENUM = { + "xhs": PlatformEnum.XHS, + "dy": PlatformEnum.DOUYIN, + "ks": PlatformEnum.KUAISHOU, + "bili": PlatformEnum.BILIBILI, + "wb": PlatformEnum.WEIBO, + "tieba": PlatformEnum.TIEBA, + "zhihu": PlatformEnum.ZHIHU, +} + +_XHS_WEB_BASE = "https://www.xiaohongshu.com" +_CREATOR_PATH = "/user/profile" +_NOTE_PATH = "/explore" + +# Timeout used when the caller does not care; tasks carry their own. +DEFAULT_RUN_TIMEOUT_SECONDS = 3600 + + +def build_target_url(value: str, kind: str) -> str: + """Turn a stored target into a URL the crawler's parser accepts. + + Always emits a full URL rather than a bare id: the XHS parser accepts a bare + 24-hex id only, so the URL form is the safer universal input. The + ``xsec_token`` is appended when present but is deliberately optional -- it + expires, and the id alone is what keeps a long-running task alive. + """ + path = _CREATOR_PATH if kind == MODE_CREATOR else _NOTE_PATH + return f"{_XHS_WEB_BASE}{path}/{value}" + + +def build_target_urls(mode: str, targets: Iterable[MonitorTarget]) -> List[str]: + urls = [] + for target in targets: + url = build_target_url(target.external_id, target.kind) + if target.xsec_token: + url = f"{url}?xsec_token={target.xsec_token}" + if target.xsec_source: + url = f"{url}&xsec_source={target.xsec_source}" + urls.append(url) + return urls + + +async def _strategy_settings(session, platform: str) -> dict: + """Crawl-strategy and proxy settings for one platform. + + Per-platform because the values genuinely differ: what is a safe request + interval on one site is a rate limit on another. Read per run rather than + cached, so a change takes effect on the next scheduled run. + """ + return { + "enable_sub_comments": bool( + await app_settings.get_value(session, "enable_sub_comments", platform, False) + ), + "crawl_sleep_sec": int( + await app_settings.get_value(session, "crawl_sleep_sec", platform, 2) + ), + "enable_ip_proxy": bool( + await app_settings.get_value(session, "enable_ip_proxy", platform, False) + ), + "proxy_provider": await app_settings.get_value( + session, "proxy_provider", platform, "kuaidaili" + ), + "proxy_pool_count": int( + await app_settings.get_value(session, "proxy_pool_count", platform, 2) + ), + "static_proxy_url": await app_settings.get_value( + session, "static_proxy_url", platform, "" + ), + } + + +def _write_cookie_file(path: Path, cookie: str) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(cookie, encoding="utf-8") + + +def _remove_cookie_file(path: Path) -> None: + """Best-effort removal; the cookie is a credential, do not leave it around.""" + try: + os.remove(path) + except OSError: + pass + + +async def execute_task(task_id: int, trigger: str = "manual") -> IngestResult: + """Run one monitoring cycle for ``task_id`` and ingest its output. + + Split into three phases with separate short-lived DB sessions so no + transaction is held open across the multi-minute subprocess run. + """ + # --- Phase 1: book the run and work out where its output goes ------------- + async with get_session() as session: + task = await session.get(MonitorTask, task_id) + if task is None: + raise ValueError(f"Monitor task {task_id} not found") + + targets = [target for target in task.targets if target.enabled] + if not targets: + raise ValueError(f"Monitor task {task_id} has no enabled targets") + + platform = task.platform + urls = build_target_urls(task.mode, targets) + cookie = await get_cookie(session, platform) + strategy = await _strategy_settings(session, platform) + + run = MonitorRun( + task_id=task.id, + trigger=trigger, + status=RUN_PENDING, + phase=task.mode, + save_data_path="", + queued_at=get_current_timestamp(), + not_before=0, + max_comments_count=task.max_comments_count if task.enable_comments else 0, + ) + session.add(run) + await session.flush() + + run_id = run.id + out_dir = MONITOR_RUNS_DIR / str(task.id) / str(run_id) + run.save_data_path = str(out_dir) + + # Snapshot the values the subprocess needs; `task` is detached after commit. + mode = task.mode + enable_comments = task.enable_comments + max_notes_count = task.max_notes_count + max_comments_count = task.max_comments_count + timeout_seconds = task.run_timeout_seconds + + # --- Phase 2: run the crawler outside any transaction --------------------- + cookie_file = out_dir / ".cookies" + _write_cookie_file(cookie_file, cookie) + + request = CrawlerStartRequest( + platform=_PLATFORM_ENUM[platform], + login_type=LoginTypeEnum.COOKIE, + crawler_type=CrawlerTypeEnum.CREATOR if mode == MODE_CREATOR else CrawlerTypeEnum.DETAIL, + creator_ids=",".join(urls) if mode == MODE_CREATOR else "", + specified_ids=",".join(urls) if mode != MODE_CREATOR else "", + start_page=1, + enable_comments=enable_comments, + enable_sub_comments=strategy["enable_sub_comments"], + enable_media=False, + save_option=SaveDataOptionEnum.JSONL, + cookies="", + headless=True, + max_notes_count=max_notes_count, + max_comments_count=max_comments_count, + # Isolate this run's output: the crawler names files by date only, so + # otherwise same-day runs would append into one shared file. + save_data_path=str(out_dir), + # Unattended runs must not try to attach to the user's desktop Chrome. + enable_cdp_mode=False, + # Only injecting web_session is not enough to sign requests from a cold + # browser profile. + inject_all_cookies=True, + save_login_state=True, + cookies_file=str(cookie_file), + max_concurrency_num=1, + # Strategy + proxy, surfaced on the Settings page. + crawler_max_sleep_sec=strategy["crawl_sleep_sec"], + enable_ip_proxy=strategy["enable_ip_proxy"], + ip_proxy_pool_count=strategy["proxy_pool_count"], + ip_proxy_provider_name=strategy["proxy_provider"], + static_proxy_url=strategy["static_proxy_url"] or None, + ) + + async with get_session() as session: + run = await session.get(MonitorRun, run_id) + if run is not None: + run.status = RUN_RUNNING + run.started_at = get_current_timestamp() + + try: + exit_code = await crawler_manager.run_and_wait(request, timeout=timeout_seconds) + finally: + _remove_cookie_file(cookie_file) + + # --- Phase 3: ingest ------------------------------------------------------ + async with get_session() as session: + run = await session.get(MonitorRun, run_id) + task = await session.get(MonitorTask, task_id) + if run is None or task is None: + raise ValueError(f"Run {run_id} or task {task_id} vanished during execution") + + if exit_code == -1 and not (out_dir / "xhs").exists(): + # run_and_wait returns -1 when the process could not start or timed out. + run.status = RUN_TIMEOUT + run.finished_at = get_current_timestamp() + run.exit_code = exit_code + run.error_message = "Run was killed by timeout or failed to start" + result = IngestResult(status=RUN_TIMEOUT, error=run.error_message) + else: + run.exit_code = exit_code + run.finished_at = get_current_timestamp() + result = await ingest_run(session, run, task, out_dir) + + # A run that authenticated fine is the only useful signal that the + # stored cookie still works. + if result.notes_fetched > 0: + await mark_cookie_ok(session, task.platform) + + task.last_run_at = run.finished_at + task.last_status = result.status + task.last_error = result.error + + # --- Phase 4: notify ------------------------------------------------------ + # Runs after the ingest transaction has committed, in its own session. A push + # failure must never roll back collected data, and notify_run() swallows its + # own errors for the same reason. + async with get_session() as session: + task = await session.get(MonitorTask, task_id) + run = await session.get(MonitorRun, run_id) + if task is not None and run is not None: + await notify.notify_run(session, task, run) + + return result diff --git a/api/monitor/scheduler.py b/api/monitor/scheduler.py new file mode 100644 index 0000000..0792f1a --- /dev/null +++ b/api/monitor/scheduler.py @@ -0,0 +1,185 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/monitor/scheduler.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Background scheduler for monitor tasks. + +One asyncio loop polls for due tasks and hands them to the runner. A plain loop +is enough here: there is exactly one process, one global crawler subprocess, and +therefore no concurrency to coordinate -- a cron-style library would add a +dependency without adding a capability. + +Scheduling is **fixed-delay**, not fixed-rate: ``next_run_at`` is set from the +moment a run starts, so a slow run cannot make its task fire back-to-back. +""" + +import asyncio +import random +from datetime import datetime +from typing import Optional + +from sqlalchemy import select + +from tools.time_util import get_current_timestamp + +from ..services import crawler_manager +from . import app_settings +from .db import get_session +from .models import MonitorRun, MonitorTask, RUN_INTERRUPTED, RUN_RUNNING +from .runner import execute_task +from .settings import get_cookie + +POLL_INTERVAL_SECONDS = 20 +# Spread tasks sharing an interval so they do not all come due on the same tick. +JITTER_SECONDS = 60 + +_MS_PER_MINUTE = 60_000 + + +class MonitorScheduler: + """Polls the task table and runs whatever is due.""" + + def __init__(self) -> None: + self._loop_task: Optional[asyncio.Task] = None + self._stopping = asyncio.Event() + # Avoids logging "no cookie" on every single tick. + self._warned_no_cookie = False + + async def start(self) -> None: + if self._loop_task is not None and not self._loop_task.done(): + return + self._stopping.clear() + self._loop_task = asyncio.create_task(self._run_loop()) + + async def stop(self) -> None: + self._stopping.set() + if self._loop_task is not None: + self._loop_task.cancel() + try: + await self._loop_task + except asyncio.CancelledError: + pass + self._loop_task = None + + async def _run_loop(self) -> None: + try: + await self.recover() + except Exception as exc: # pragma: no cover - defensive + print(f"[monitor.scheduler] recovery failed: {exc}") + + while not self._stopping.is_set(): + try: + await self.tick() + except Exception as exc: # pragma: no cover - keep the loop alive + print(f"[monitor.scheduler] tick failed: {exc}") + await asyncio.sleep(POLL_INTERVAL_SECONDS) + + async def recover(self) -> None: + """Clean up state left behind by a server restart. + + A run still marked ``running`` cannot be running -- its subprocess died + with the previous process. Marking it interrupted stops it from blocking + the UI as a phantom in-flight run. + """ + async with get_session() as session: + stale = ( + await session.scalars( + select(MonitorRun).where(MonitorRun.status == RUN_RUNNING) + ) + ).all() + for run in stale: + run.status = RUN_INTERRUPTED + run.finished_at = get_current_timestamp() + if stale: + print( + f"[monitor.scheduler] marked {len(stale)} interrupted run(s) " + f"left over from a previous process" + ) + + async def tick(self) -> None: + """Run one due task, if the crawler is free and we are in the active window.""" + # The crawler subprocess is a global singleton, so a manual crawl and a + # monitor run cannot overlap. Returning without advancing next_run_at + # leaves the task due, and it is picked up on a later tick. + if crawler_manager.is_busy(): + return + + async with get_session() as session: + if not await self._within_active_hours(session): + # Deliberately does not advance next_run_at: the task simply runs + # when the window next opens, rather than being skipped for a day. + return + + await self._run_due_task() + + async def _within_active_hours(self, session) -> bool: + """Whether scheduled runs are allowed right now (local time).""" + start, end = await app_settings.active_hours(session) + hour = datetime.now().hour + if start <= end: + return start <= hour <= end + # Window wraps past midnight, e.g. 22 -> 6. + return hour >= start or hour <= end + + async def _run_due_task(self) -> None: + async with get_session() as session: + task = await session.scalar( + select(MonitorTask) + .where( + MonitorTask.enabled.is_(True), + MonitorTask.next_run_at.is_not(None), + MonitorTask.next_run_at <= get_current_timestamp(), + ) + .order_by(MonitorTask.next_run_at) + .limit(1) + ) + + if task is None: + return + + # No cookie means every run would report an auth failure. Leave the + # task due rather than advancing: it starts working the moment the + # user pastes one. + cookie = await get_cookie(session) + if not cookie: + if not self._warned_no_cookie: + print( + "[monitor.scheduler] no XHS cookie configured; " + "scheduled tasks will not run until one is set" + ) + self._warned_no_cookie = True + return + self._warned_no_cookie = False + + # Advance before running so a crash mid-run cannot cause an immediate + # re-fire, and so a long outage coalesces into a single run instead + # of one run per missed interval. + task.next_run_at = ( + get_current_timestamp() + + task.interval_minutes * _MS_PER_MINUTE + + random.randint(0, JITTER_SECONDS) * 1000 + ) + task_id = task.id + + try: + await execute_task(task_id, trigger="scheduled") + except Exception as exc: + print(f"[monitor.scheduler] task {task_id} failed: {exc}") + + +# Global singleton, mirroring the crawler_manager pattern. +monitor_scheduler = MonitorScheduler() diff --git a/api/monitor/service.py b/api/monitor/service.py new file mode 100644 index 0000000..20c03bd --- /dev/null +++ b/api/monitor/service.py @@ -0,0 +1,720 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/monitor/service.py +# GitHub: https://github.com/NanmiCoder +# Non-commercial learning license 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Task CRUD and dashboard queries for the monitoring layer.""" + +import asyncio +import re +from typing import Any, Dict, List, Optional +from urllib.parse import parse_qs, urlparse + +from sqlalchemy import delete, func, select +from sqlalchemy.ext.asyncio import AsyncSession + +from tools.time_util import get_current_timestamp + +from . import app_settings, platforms +from .db import get_session +from .platforms import PLATFORM_XHS +from .models import ( + MODE_CREATOR, + MODE_NOTE, + MonitorComment, + MonitorEvent, + MonitorNote, + MonitorNoteMetric, + MonitorRun, + MonitorTarget, + MonitorTask, + RUN_SUCCESS, + RUN_PARTIAL, +) +from .runner import execute_task + +# Keep strong references to in-flight manual runs; asyncio only holds weak ones, +# so without this a run can be garbage collected mid-flight. +_background_runs: set[asyncio.Task] = set() + +MIN_INTERVAL_MINUTES = 30 +MAX_INTERVAL_MINUTES = 7 * 24 * 60 + +_CREATOR_URL_RE = re.compile(r"xiaohongshu\.com/user/profile/([A-Za-z0-9_-]+)") +_NOTE_URL_RE = re.compile(r"xiaohongshu\.com/(?:explore|discovery/item)/([A-Za-z0-9_-]+)") +# XHS user ids and note ids are 24-char hex; allow a slightly wider range so a +# format change degrades into "still accepted" rather than "rejected". +_BARE_ID_RE = re.compile(r"^[A-Za-z0-9_-]{8,64}$") + + +class TargetParseError(ValueError): + """Raised when a pasted monitoring target cannot be understood.""" + + +def parse_target_input( + value: str, mode: str, platform: str = PLATFORM_XHS +) -> Dict[str, str]: + """Parse a pasted creator/note value into a stable id plus a refreshable token. + + Accepts either a full URL (with or without ``xsec_token``) or a bare id. + Storing the id separately from the token is what keeps a long-running task + alive: tokens expire, ids do not. + + URL shapes are platform-specific. Only Xiaohongshu is wired, so anything else + is rejected here as well as at task creation -- parsing a Douyin link as if it + were a Xiaohongshu one would be worse than refusing it. + """ + if platform != PLATFORM_XHS: + raise TargetParseError(f"暂不支持解析该平台({platform})的目标链接") + + raw = (value or "").strip() + if not raw: + raise TargetParseError("Empty target") + + external_id = "" + if raw.startswith("http") or "/" in raw: + # xhslink.com and other short links are not resolvable without a network + # round-trip, so only the direct profile/explore forms are supported. + match = _CREATOR_URL_RE.search(raw) if mode == MODE_CREATOR else _NOTE_URL_RE.search(raw) + if not match: + expected = "博主主页" if mode == MODE_CREATOR else "笔记" + raise TargetParseError(f"无法从链接中解析出{expected} ID:{raw}") + external_id = match.group(1) + elif _BARE_ID_RE.match(raw): + external_id = raw + else: + raise TargetParseError(f"无法识别的目标:{raw}") + + params = parse_qs(urlparse(raw).query) if raw.startswith("http") else {} + return { + "external_id": external_id, + "xsec_token": (params.get("xsec_token") or [""])[0], + "xsec_source": (params.get("xsec_source") or [""])[0], + "raw_value": raw, + } + + +async def platform_task_ids(session: AsyncSession, platform: str) -> List[int]: + """Ids of the tasks belonging to a platform. + + Note/comment/event tables carry no platform column -- they hang off a task -- + so scoping a query to a platform means scoping it to that task set. + """ + return list( + await session.scalars(select(MonitorTask.id).where(MonitorTask.platform == platform)) + ) + + +# --------------------------------------------------------------------------- +# Task CRUD +# --------------------------------------------------------------------------- + +async def create_task(session: AsyncSession, payload: Dict[str, Any]) -> MonitorTask: + mode = payload["mode"] + if mode not in (MODE_CREATOR, MODE_NOTE): + raise ValueError(f"Unsupported mode: {mode}") + + # Rejects unknown platforms and, more importantly, platforms whose crawler + # exists upstream but whose monitoring is not wired up -- accepting those + # would create a task that can never produce data. + platform = payload.get("platform") or PLATFORM_XHS + platforms.ensure_runnable(platform) + + now = get_current_timestamp() + + # Fall back to the configured defaults for anything the caller left out, so + # the Settings page actually governs new tasks. + defaults = await app_settings.defaults(session, platform) + interval_minutes = payload.get("interval_minutes") or defaults["interval_minutes"] + interval_ms = int(interval_minutes) * 60_000 + + task = MonitorTask( + name=payload["name"], + platform=platform, + mode=mode, + enabled=payload.get("enabled", True), + interval_minutes=interval_minutes, + max_notes_count=payload.get("max_notes_count") or defaults["max_notes_count"], + enable_comments=payload.get("enable_comments", True), + max_comments_count=payload.get("max_comments_count") or defaults["max_comments_count"], + run_timeout_seconds=payload.get("run_timeout_seconds", 3600), + notify_enabled=payload.get("notify_enabled", False), + next_run_at=now + interval_ms, + last_status="idle", + created_at=now, + updated_at=now, + ) + session.add(task) + await session.flush() + + seen: set[str] = set() + for value in payload.get("targets", []): + parsed = parse_target_input(value, mode, platform) + if parsed["external_id"] in seen: + continue + seen.add(parsed["external_id"]) + session.add( + MonitorTarget( + task_id=task.id, + kind=mode, + external_id=parsed["external_id"], + xsec_token=parsed["xsec_token"], + xsec_source=parsed["xsec_source"], + raw_value=parsed["raw_value"], + label=parsed["external_id"], + enabled=True, + created_at=now, + ) + ) + + await session.flush() + return task + + +async def update_task(session: AsyncSession, task_id: int, payload: Dict[str, Any]) -> MonitorTask: + task = await session.get(MonitorTask, task_id) + if task is None: + raise ValueError(f"Task {task_id} not found") + + for field in ( + "name", + "enabled", + "interval_minutes", + "max_notes_count", + "enable_comments", + "max_comments_count", + "run_timeout_seconds", + "notify_enabled", + ): + if field in payload and payload[field] is not None: + setattr(task, field, payload[field]) + + # Replacing targets resets the baseline implicitly: a note set that now + # includes new ids will simply report them as new on the next run. + if payload.get("targets") is not None: + await session.execute(delete(MonitorTarget).where(MonitorTarget.task_id == task_id)) + now = get_current_timestamp() + seen: set[str] = set() + for value in payload["targets"]: + parsed = parse_target_input(value, task.mode) + if parsed["external_id"] in seen: + continue + seen.add(parsed["external_id"]) + session.add( + MonitorTarget( + task_id=task_id, + kind=task.mode, + external_id=parsed["external_id"], + xsec_token=parsed["xsec_token"], + xsec_source=parsed["xsec_source"], + raw_value=parsed["raw_value"], + label=parsed["external_id"], + enabled=True, + created_at=now, + ) + ) + + if "interval_minutes" in payload and payload["interval_minutes"]: + task.next_run_at = get_current_timestamp() + payload["interval_minutes"] * 60_000 + + task.updated_at = get_current_timestamp() + await session.flush() + return task + + +async def delete_task(session: AsyncSession, task_id: int) -> None: + task = await session.get(MonitorTask, task_id) + if task is None: + raise ValueError(f"Task {task_id} not found") + await session.delete(task) + + +def trigger_manual_run(task_id: int) -> None: + """Fire a run in the background and return immediately. + + A crawl takes minutes, so the HTTP request must not wait for it. The UI + follows progress through the logs WebSocket and the run history. + """ + task = asyncio.create_task(execute_task(task_id, trigger="manual")) + _background_runs.add(task) + task.add_done_callback(_background_runs.discard) + + +# --------------------------------------------------------------------------- +# Dashboard queries +# --------------------------------------------------------------------------- + +async def _latest_successful_run_id(session: AsyncSession, task_id: int) -> Optional[int]: + return await session.scalar( + select(MonitorRun.id) + .where( + MonitorRun.task_id == task_id, + MonitorRun.status.in_((RUN_SUCCESS, RUN_PARTIAL)), + ) + .order_by(MonitorRun.id.desc()) + .limit(1) + ) + + +def _delta(current: Optional[int], previous: Optional[int]) -> Optional[int]: + if current is None or previous is None: + return None + return current - previous + + +async def list_notes( + session: AsyncSession, + task_id: Optional[int] = None, + only_new: bool = False, + limit: int = 200, + platform: Optional[str] = None, +) -> List[Dict[str, Any]]: + """Tracked notes with their latest metrics and change vs the previous run.""" + query = select(MonitorNote).order_by(MonitorNote.last_seen_at.desc()).limit(limit) + if task_id is not None: + query = query.where(MonitorNote.task_id == task_id) + if platform is not None: + scoped = await platform_task_ids(session, platform) + if not scoped: + return [] + query = query.where(MonitorNote.task_id.in_(scoped)) + + notes = list((await session.scalars(query)).all()) + if not notes: + return [] + + note_ids = [note.note_id for note in notes] + + # Fetch every snapshot for these notes in one go and gather the two most + # recent per note, rather than issuing two queries per note. + snapshots = list( + ( + await session.scalars( + select(MonitorNoteMetric) + .where(MonitorNoteMetric.note_id.in_(note_ids)) + .order_by(MonitorNoteMetric.note_id, MonitorNoteMetric.run_id.desc()) + ) + ).all() + ) + by_note: Dict[str, List[MonitorNoteMetric]] = {} + for snapshot in snapshots: + by_note.setdefault(snapshot.note_id, []).append(snapshot) + + latest_run_ids: Dict[int, Optional[int]] = {} + result: List[Dict[str, Any]] = [] + + for note in notes: + series = by_note.get(note.note_id, []) + current = series[0] if series else None + previous = series[1] if len(series) > 1 else None + + if only_new: + if note.task_id not in latest_run_ids: + latest_run_ids[note.task_id] = await _latest_successful_run_id(session, note.task_id) + if note.first_seen_run_id != latest_run_ids[note.task_id]: + continue + + result.append( + { + "task_id": note.task_id, + "note_id": note.note_id, + "title": note.title, + "note_url": note.note_url, + "cover": note.cover, + "first_seen_at": note.first_seen_at, + "last_seen_at": note.last_seen_at, + "is_new": note.first_seen_run_id == latest_run_ids.get(note.task_id), + "metrics": { + "liked_count": current.liked_count if current else None, + "comment_count": current.comment_count if current else None, + "collected_count": current.collected_count if current else None, + "share_count": current.share_count if current else None, + }, + "deltas": { + "liked_count": _delta( + current.liked_count if current else None, + previous.liked_count if previous else None, + ), + "comment_count": _delta( + current.comment_count if current else None, + previous.comment_count if previous else None, + ), + "collected_count": _delta( + current.collected_count if current else None, + previous.collected_count if previous else None, + ), + "share_count": _delta( + current.share_count if current else None, + previous.share_count if previous else None, + ), + }, + "snapshot_count": len(series), + } + ) + + return result + + +async def note_series(session: AsyncSession, note_id: str, task_id: Optional[int] = None) -> List[Dict[str, Any]]: + """Metric time series for one note.""" + query = ( + select(MonitorNoteMetric) + .where(MonitorNoteMetric.note_id == note_id) + .order_by(MonitorNoteMetric.run_id) + ) + if task_id is not None: + query = query.where(MonitorNoteMetric.task_id == task_id) + + return [ + { + "run_id": row.run_id, + "captured_at": row.captured_at, + "liked_count": row.liked_count, + "comment_count": row.comment_count, + "collected_count": row.collected_count, + "share_count": row.share_count, + } + for row in (await session.scalars(query)).all() + ] + + +async def _note_meta_map( + session: AsyncSession, note_ids: List[str] +) -> Dict[str, Dict[str, Any]]: + """Look up note title/cover/url for a set of note ids. + + Fetched as one query and joined in Python rather than as a SQL join: the + comment table has no foreign key to the note table (both are keyed by the + platform's note id, per task), and a single IN() is easier to follow here. + """ + if not note_ids: + return {} + + rows = ( + await session.scalars(select(MonitorNote).where(MonitorNote.note_id.in_(set(note_ids)))) + ).all() + return { + row.note_id: { + "note_title": row.title, + "note_cover": row.cover, + "note_url": row.note_url, + "task_id": row.task_id, + } + for row in rows + } + + +async def list_comments( + session: AsyncSession, + task_id: Optional[int] = None, + note_id: Optional[str] = None, + limit: int = 200, + platform: Optional[str] = None, +) -> List[Dict[str, Any]]: + """Comments, each carrying the note it belongs to. + + The note association is the point: without it a comment stream is unreadable, + since a bare note_id tells the operator nothing. + """ + query = select(MonitorComment).order_by(MonitorComment.first_seen_at.desc()).limit(limit) + if task_id is not None: + query = query.where(MonitorComment.task_id == task_id) + if note_id is not None: + query = query.where(MonitorComment.note_id == note_id) + if platform is not None: + scoped = await platform_task_ids(session, platform) + if not scoped: + return [] + query = query.where(MonitorComment.task_id.in_(scoped)) + + comments = list((await session.scalars(query)).all()) + meta = await _note_meta_map(session, [row.note_id for row in comments]) + + return [ + { + "task_id": row.task_id, + "note_id": row.note_id, + "comment_id": row.comment_id, + "content": row.content, + "nickname": row.nickname, + "create_time": row.create_time, + "like_count": row.like_count, + "sub_comment_count": row.sub_comment_count, + "first_seen_at": row.first_seen_at, + "note_title": meta.get(row.note_id, {}).get("note_title", ""), + "note_cover": meta.get(row.note_id, {}).get("note_cover", ""), + "note_url": meta.get(row.note_id, {}).get("note_url", ""), + } + for row in comments + ] + + +async def comment_note_groups( + session: AsyncSession, task_id: Optional[int] = None, platform: Optional[str] = None +) -> List[Dict[str, Any]]: + """Notes that have comments, newest first, with their comment counts. + + Feeds the comment filter dropdown: the operator picks a work by title, so + the counts need to be visible before choosing. + """ + scoped_ids: Optional[List[int]] = None + if platform is not None: + scoped_ids = await platform_task_ids(session, platform) + if not scoped_ids: + return [] + + count_query = select( + MonitorComment.note_id, func.count().label("comment_count") + ).group_by(MonitorComment.note_id) + if task_id is not None: + count_query = count_query.where(MonitorComment.task_id == task_id) + if scoped_ids is not None: + count_query = count_query.where(MonitorComment.task_id.in_(scoped_ids)) + + counts = {row.note_id: row.comment_count for row in (await session.execute(count_query)).all()} + if not counts: + return [] + + latest_query = ( + select(MonitorComment.note_id, func.max(MonitorComment.first_seen_at).label("latest")) + .where(MonitorComment.note_id.in_(set(counts))) + .group_by(MonitorComment.note_id) + ) + if task_id is not None: + latest_query = latest_query.where(MonitorComment.task_id == task_id) + if scoped_ids is not None: + latest_query = latest_query.where(MonitorComment.task_id.in_(scoped_ids)) + latest = {row.note_id: row.latest for row in (await session.execute(latest_query)).all()} + + meta = await _note_meta_map(session, list(counts)) + + groups = [ + { + "note_id": note_id, + "note_title": meta.get(note_id, {}).get("note_title", ""), + "note_cover": meta.get(note_id, {}).get("note_cover", ""), + "note_url": meta.get(note_id, {}).get("note_url", ""), + "comment_count": count, + "latest_at": latest.get(note_id, 0), + } + for note_id, count in counts.items() + ] + groups.sort(key=lambda group: group["latest_at"], reverse=True) + return groups + + +async def list_events( + session: AsyncSession, + task_id: Optional[int] = None, + event_type: Optional[str] = None, + since_id: Optional[int] = None, + limit: int = 200, + platform: Optional[str] = None, +) -> List[Dict[str, Any]]: + query = select(MonitorEvent).order_by(MonitorEvent.id.desc()).limit(limit) + if task_id is not None: + query = query.where(MonitorEvent.task_id == task_id) + if event_type is not None: + query = query.where(MonitorEvent.type == event_type) + if since_id is not None: + query = query.where(MonitorEvent.id > since_id) + if platform is not None: + scoped = await platform_task_ids(session, platform) + if not scoped: + return [] + query = query.where(MonitorEvent.task_id.in_(scoped)) + + return [ + { + "id": row.id, + "task_id": row.task_id, + "run_id": row.run_id, + "type": row.type, + "severity": row.severity, + "target_kind": row.target_kind, + "target_id": row.target_id, + "title": row.title, + "created_at": row.created_at, + "is_read": row.is_read, + } + for row in (await session.scalars(query)).all() + ] + + +async def list_runs(session: AsyncSession, task_id: int, limit: int = 50) -> List[Dict[str, Any]]: + rows = ( + await session.scalars( + select(MonitorRun) + .where(MonitorRun.task_id == task_id) + .order_by(MonitorRun.id.desc()) + .limit(limit) + ) + ).all() + + return [ + { + "id": row.id, + "task_id": row.task_id, + "status": row.status, + "trigger": row.trigger, + "queued_at": row.queued_at, + "started_at": row.started_at, + "finished_at": row.finished_at, + "exit_code": row.exit_code, + "notes_fetched": row.notes_fetched, + "comments_fetched": row.comments_fetched, + "new_notes": row.new_notes, + "new_comments": row.new_comments, + "is_baseline": row.is_baseline, + "max_comments_count": row.max_comments_count, + "error_message": row.error_message, + } + for row in rows + ] + + +async def list_tasks( + session: AsyncSession, platform: Optional[str] = None +) -> List[Dict[str, Any]]: + query = select(MonitorTask).order_by(MonitorTask.id) + if platform is not None: + query = query.where(MonitorTask.platform == platform) + + tasks = list((await session.scalars(query)).all()) + if not tasks: + return [] + + counts = dict( + ( + await session.execute( + select(MonitorTarget.task_id, func.count()) + .group_by(MonitorTarget.task_id) + ) + ).all() + ) + unread = dict( + ( + await session.execute( + select(MonitorEvent.task_id, func.count()) + .where(MonitorEvent.is_read.is_(False)) + .group_by(MonitorEvent.task_id) + ) + ).all() + ) + + return [ + { + "id": task.id, + "name": task.name, + "platform": task.platform, + "mode": task.mode, + "enabled": task.enabled, + "interval_minutes": task.interval_minutes, + "max_notes_count": task.max_notes_count, + "enable_comments": task.enable_comments, + "max_comments_count": task.max_comments_count, + "run_timeout_seconds": task.run_timeout_seconds, + "notify_enabled": task.notify_enabled, + "next_run_at": task.next_run_at, + "last_run_at": task.last_run_at, + "last_status": task.last_status, + "last_error": task.last_error, + "last_notified_at": task.last_notified_at, + "target_count": counts.get(task.id, 0), + "targets": [ + {"id": t.id, "external_id": t.external_id, "raw_value": t.raw_value, "enabled": t.enabled} + for t in task.targets + ], + "unread_events": unread.get(task.id, 0), + } + for task in tasks + ] + + +async def overview(session: AsyncSession, platform: Optional[str] = None) -> Dict[str, Any]: + """Headline numbers for the dashboard tiles, scoped to one platform.""" + now = get_current_timestamp() + day_ago = now - 24 * 60 * 60 * 1000 + + # Nothing but the task table carries a platform column, so the other counts + # are scoped through the platform's task ids. + scoped: Optional[List[int]] = None + if platform is not None: + scoped = await platform_task_ids(session, platform) + + def by_task(stmt, column): + return stmt if scoped is None else stmt.where(column.in_(scoped)) + + task_count = select(func.count()).select_from(MonitorTask) + if platform is not None: + task_count = task_count.where(MonitorTask.platform == platform) + + enabled_count = select(func.count()).select_from(MonitorTask).where( + MonitorTask.enabled.is_(True) + ) + if platform is not None: + enabled_count = enabled_count.where(MonitorTask.platform == platform) + + return { + "platform": platform, + "tasks": await session.scalar(task_count) or 0, + "enabled_tasks": await session.scalar(enabled_count) or 0, + "notes": await session.scalar( + by_task(select(func.count()).select_from(MonitorNote), MonitorNote.task_id) + ) + or 0, + "comments": await session.scalar( + by_task(select(func.count()).select_from(MonitorComment), MonitorComment.task_id) + ) + or 0, + "events_24h": await session.scalar( + by_task( + select(func.count()) + .select_from(MonitorEvent) + .where(MonitorEvent.created_at >= day_ago), + MonitorEvent.task_id, + ) + ) + or 0, + "unread_events": await session.scalar( + by_task( + select(func.count()) + .select_from(MonitorEvent) + .where(MonitorEvent.is_read.is_(False)), + MonitorEvent.task_id, + ) + ) + or 0, + "running_runs": await session.scalar( + by_task( + select(func.count()) + .select_from(MonitorRun) + .where(MonitorRun.status == "running"), + MonitorRun.task_id, + ) + ) + or 0, + } + + +async def mark_events_read(session: AsyncSession, task_id: Optional[int] = None) -> int: + query = select(MonitorEvent).where(MonitorEvent.is_read.is_(False)) + if task_id is not None: + query = query.where(MonitorEvent.task_id == task_id) + rows = list((await session.scalars(query)).all()) + for row in rows: + row.is_read = True + return len(rows) diff --git a/api/monitor/settings.py b/api/monitor/settings.py new file mode 100644 index 0000000..8ac1f4b --- /dev/null +++ b/api/monitor/settings.py @@ -0,0 +1,108 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/monitor/settings.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Key/value settings for the monitoring layer, plus cookie health helpers. + +The XHS cookie is what makes scheduled runs unattended. It expires every few +weeks, so alongside the value we track when it was last seen working -- that is +what lets the UI warn before a task silently stops collecting. +""" + +from typing import Optional + +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from tools.time_util import get_current_timestamp + +from .models import MonitorSetting +from .platforms import PLATFORM_XHS + + +def platform_key(platform: str, name: str) -> str: + """Key for a setting that each platform keeps its own copy of.""" + return f"platform.{platform}.{name}" + + +def system_key(name: str) -> str: + """Key for a setting shared across every platform.""" + return f"system.{name}" + + +def cookie_key(platform: str) -> str: + return platform_key(platform, "cookie") + + +def cookie_updated_key(platform: str) -> str: + return platform_key(platform, "cookie_updated_at") + + +def cookie_last_ok_key(platform: str) -> str: + return platform_key(platform, "cookie_last_ok_at") + + +async def get_setting(session: AsyncSession, key: str) -> Optional[str]: + return await session.scalar(select(MonitorSetting.value).where(MonitorSetting.key == key)) + + +async def set_setting(session: AsyncSession, key: str, value: str) -> None: + row = await session.get(MonitorSetting, key) + now = get_current_timestamp() + if row is None: + session.add(MonitorSetting(key=key, value=value, updated_at=now)) + else: + row.value = value + row.updated_at = now + + +async def delete_setting(session: AsyncSession, key: str) -> None: + row = await session.get(MonitorSetting, key) + if row is not None: + await session.delete(row) + + +async def get_cookie(session: AsyncSession, platform: str = PLATFORM_XHS) -> str: + return (await get_setting(session, cookie_key(platform))) or "" + + +async def set_cookie( + session: AsyncSession, cookie: str, platform: str = PLATFORM_XHS +) -> None: + await set_setting(session, cookie_key(platform), cookie) + await set_setting(session, cookie_updated_key(platform), str(get_current_timestamp())) + + +async def mark_cookie_ok(session: AsyncSession, platform: str = PLATFORM_XHS) -> None: + """Record that a run authenticated successfully.""" + await set_setting(session, cookie_last_ok_key(platform), str(get_current_timestamp())) + + +async def get_cookie_status(session: AsyncSession, platform: str = PLATFORM_XHS) -> dict: + """Cookie health for the UI. Never returns the cookie value itself.""" + cookie = await get_cookie(session, platform) + updated_at = await get_setting(session, cookie_updated_key(platform)) + last_ok_at = await get_setting(session, cookie_last_ok_key(platform)) + + return { + "platform": platform, + "present": bool(cookie), + # Enough to eyeball whether the pasted value looks right, not enough to leak it. + "length": len(cookie), + "updated_at": int(updated_at) if updated_at else None, + "last_ok_at": int(last_ok_at) if last_ok_at else None, + } diff --git a/api/routers/__init__.py b/api/routers/__init__.py index 123cbc0..70c6dc6 100644 --- a/api/routers/__init__.py +++ b/api/routers/__init__.py @@ -16,8 +16,18 @@ # 详细许可条款请参阅项目根目录下的LICENSE文件。 # 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 +from .auth import router as auth_router from .crawler import router as crawler_router from .data import router as data_router +from .monitor import router as monitor_router +from .settings import router as settings_router from .websocket import router as websocket_router -__all__ = ["crawler_router", "data_router", "websocket_router"] +__all__ = [ + "auth_router", + "crawler_router", + "data_router", + "monitor_router", + "settings_router", + "websocket_router", +] diff --git a/api/routers/auth.py b/api/routers/auth.py new file mode 100644 index 0000000..852b9ed --- /dev/null +++ b/api/routers/auth.py @@ -0,0 +1,173 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/routers/auth.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Login / logout endpoints. + +Deliberately exempt from ``require_auth``: +* ``/login`` -- it is the way in. +* ``/logout`` -- exempt so an already-expired session still gets a clean 200 + and a cleared cookie instead of a confusing 401, which + would leave the browser holding a stale cookie. +""" + +from fastapi import APIRouter, Depends, HTTPException, Request, Response, status + +from ..auth import ( + INVALID_CREDENTIALS, + SESSION_COOKIE_NAME, + check_password, + require_auth, + clear_failures, + client_key, + cookie_secure, + create_session, + purge_expired_sessions, + record_failure, + resolve_session, + retry_after_seconds, + revoke_all_sessions, + revoke_session, + set_password, + token_from_request, +) +from ..monitor.db import get_session +from ..schemas.auth import ChangePasswordPayload, LoginPayload +from tools.time_util import get_current_timestamp + +router = APIRouter(prefix="/auth", tags=["auth"]) + + +def _apply_session_cookie(response: Response, token: str, expires_at: int) -> None: + """Attach the session cookie. + + ``secure`` is off by default because the panel is served over plain HTTP on + a LAN; setting it there means the browser silently discards the cookie and + the login page just loops with no error. ``SameSite=lax`` is also what + blocks cross-site POSTs, i.e. the CSRF defence for the write endpoints. + """ + max_age = max((expires_at - get_current_timestamp()) // 1000, 60) + response.set_cookie( + key=SESSION_COOKIE_NAME, + value=token, + max_age=max_age, + httponly=True, + secure=cookie_secure(), + samesite="lax", + path="/", + ) + + +@router.post("/login") +async def login(payload: LoginPayload, request: Request, response: Response): + key = client_key(request) + + wait = await retry_after_seconds(key) + if wait: + raise HTTPException( + status_code=status.HTTP_429_TOO_MANY_REQUESTS, + detail=f"尝试过于频繁,请 {wait} 秒后再试", + headers={"Retry-After": str(wait)}, + ) + + async with get_session() as session: + valid = await check_password(session, payload.password) + token = "" + expires_at = 0 + if valid: + await purge_expired_sessions(session) + token, expires_at = await create_session(session) + + if not valid: + await record_failure(key) + # One generic message regardless of whether the password was wrong, + # empty, or simply not set yet -- no oracle. + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, detail=INVALID_CREDENTIALS + ) + + await clear_failures(key) + _apply_session_cookie(response, token, expires_at) + return {"expires_at": expires_at} + + +@router.post("/logout") +async def logout(request: Request, response: Response): + token = token_from_request(request) + if token: + async with get_session() as session: + await revoke_session(session, token) + + response.delete_cookie(SESSION_COOKIE_NAME, path="/") + return {"message": "已退出登录"} + + +@router.get("/me") +async def me(request: Request): + """Identity probe. The SPA treats a 401 here as "show the login page". + + Does its own resolution rather than using ``require_auth`` so it can also + report the expiry. + """ + token = token_from_request(request) + if not token: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, detail=INVALID_CREDENTIALS + ) + + async with get_session() as session: + row = await resolve_session(session, token) + if row is None: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, detail=INVALID_CREDENTIALS + ) + return {"authenticated": True, "expires_at": row.expires_at} + + +# Auth as a route dependency, not only inside the handler: FastAPI validates the +# request body before the endpoint body runs, so an unauthenticated caller would +# otherwise get a 422 that confirms the endpoint and its schema exist. +@router.post("/password", dependencies=[Depends(require_auth)]) +async def change_password( + payload: ChangePasswordPayload, request: Request, response: Response +): + """Change the password and log every device out. + + Revoking all sessions is the point: a password change is usually a response + to suspicion, and leaving other sessions alive would defeat it. + """ + token = token_from_request(request) + + async with get_session() as session: + current = await resolve_session(session, token) + if current is None: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, detail=INVALID_CREDENTIALS + ) + + if not await check_password(session, payload.current): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, detail="当前密码不正确" + ) + + await set_password(session, payload.new) + await revoke_all_sessions(session) + # Issue a fresh session so the caller is not bounced mid-use. + new_token, expires_at = await create_session(session) + + _apply_session_cookie(response, new_token, expires_at) + return {"message": "密码已更新,其他设备的登录已全部失效", "expires_at": expires_at} diff --git a/api/routers/crawler.py b/api/routers/crawler.py index eead9e1..7547721 100644 --- a/api/routers/crawler.py +++ b/api/routers/crawler.py @@ -18,7 +18,9 @@ from fastapi import APIRouter, HTTPException -from ..schemas import CrawlerStartRequest, CrawlerStatusResponse +from ..monitor.db import get_session +from ..monitor.settings import get_cookie +from ..schemas import CrawlerStartRequest, CrawlerStatusResponse, LoginTypeEnum from ..services import crawler_manager router = APIRouter(prefix="/crawler", tags=["crawler"]) @@ -26,7 +28,28 @@ router = APIRouter(prefix="/crawler", tags=["crawler"]) @router.post("/start") async def start_crawler(request: CrawlerStartRequest): - """Start crawler task""" + """Start crawler task. + + A cookie login with no cookie supplied falls back to the one stored for the + selected platform. The manual crawl and the monitor therefore share a single + credential; keeping a second paste field on the crawl page meant it went + stale and could silently disagree with the monitor's. + """ + if ( + request.login_type == LoginTypeEnum.COOKIE + and not request.cookies + and not request.cookies_file + ): + async with get_session() as session: + stored = await get_cookie(session, request.platform.value) + + if not stored: + raise HTTPException( + status_code=400, + detail="该平台尚未保存 Cookie,请到「设置 → 登录态」配置,或改用扫码登录", + ) + request.cookies = stored + success = await crawler_manager.start(request) if not success: # Handle concurrent/duplicate requests: if process is already running, return 400 instead of 500 diff --git a/api/routers/monitor.py b/api/routers/monitor.py new file mode 100644 index 0000000..1957b7f --- /dev/null +++ b/api/routers/monitor.py @@ -0,0 +1,502 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/routers/monitor.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""HTTP API for scheduled monitoring tasks.""" + +from datetime import date, timedelta +from typing import Any, Dict, List, Optional + +from fastapi import APIRouter, HTTPException, Query, Response + +from ..monitor import notify, report, service +from ..monitor.db import get_session +from ..monitor.platforms import PLATFORM_XHS +from ..monitor.settings import ( + cookie_key, + delete_setting, + get_cookie_status, + get_setting, + set_cookie, + set_setting, +) +from ..monitor.models import SETTING_WECOM_WEBHOOK, MonitorTask +from ..schemas.monitor import ( + CookiePayload, + MonitorTaskCreate, + MonitorTaskUpdate, + WebhookPayload, + WebhookTestPayload, +) + +router = APIRouter(prefix="/monitor", tags=["monitor"]) + + +@router.get("/overview") +async def get_overview(platform: Optional[str] = None): + """Headline numbers for the dashboard tiles, scoped to one platform.""" + async with get_session() as session: + return await service.overview(session, platform) + + +# --------------------------------------------------------------------------- +# Tasks +# --------------------------------------------------------------------------- + +@router.get("/tasks") +async def list_tasks(platform: Optional[str] = None): + async with get_session() as session: + return {"tasks": await service.list_tasks(session, platform)} + + +@router.post("/tasks", status_code=201) +async def create_task(payload: MonitorTaskCreate): + async with get_session() as session: + try: + task = await service.create_task(session, payload.model_dump()) + except service.TargetParseError as exc: + raise HTTPException(status_code=400, detail=str(exc)) + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) + return {"id": task.id, "message": "Monitoring task created"} + + +@router.patch("/tasks/{task_id}") +async def update_task(task_id: int, payload: MonitorTaskUpdate): + async with get_session() as session: + try: + await service.update_task(session, task_id, payload.model_dump(exclude_unset=True)) + except service.TargetParseError as exc: + raise HTTPException(status_code=400, detail=str(exc)) + except ValueError as exc: + raise HTTPException(status_code=404, detail=str(exc)) + return {"message": "Monitoring task updated"} + + +@router.delete("/tasks/{task_id}") +async def delete_task(task_id: int): + async with get_session() as session: + try: + await service.delete_task(session, task_id) + except ValueError as exc: + raise HTTPException(status_code=404, detail=str(exc)) + return {"message": "Monitoring task deleted"} + + +@router.post("/tasks/{task_id}/run") +async def run_task_now(task_id: int): + """Queue a run immediately and return; the crawl itself takes minutes.""" + async with get_session() as session: + task = await session.get(MonitorTask, task_id) + if task is None: + raise HTTPException(status_code=404, detail=f"Task {task_id} not found") + if not any(target.enabled for target in task.targets): + raise HTTPException(status_code=400, detail="Task has no enabled targets") + + service.trigger_manual_run(task_id) + return {"message": "Run queued"} + + +@router.get("/tasks/{task_id}/runs") +async def list_runs(task_id: int, limit: int = Query(default=50, ge=1, le=500)): + async with get_session() as session: + return {"runs": await service.list_runs(session, task_id, limit=limit)} + + +# --------------------------------------------------------------------------- +# Collected data +# --------------------------------------------------------------------------- + +@router.get("/notes") +async def list_notes( + task_id: Optional[int] = None, + only_new: bool = False, + limit: int = Query(default=200, ge=1, le=2000), + platform: Optional[str] = None, +): + async with get_session() as session: + return { + "notes": await service.list_notes(session, task_id, only_new, limit, platform) + } + + +@router.get("/notes/{note_id}/series") +async def note_series(note_id: str, task_id: Optional[int] = None): + """Metric time series for a single note.""" + async with get_session() as session: + return {"series": await service.note_series(session, note_id, task_id)} + + +@router.get("/comments") +async def list_comments( + task_id: Optional[int] = None, + note_id: Optional[str] = None, + group_by: Optional[str] = Query( + default=None, description="传 note 则按作品分组返回,便于阅读" + ), + limit: int = Query(default=200, ge=1, le=2000), + platform: Optional[str] = None, +): + """Comments, each carrying the work it belongs to. + + ``note_id`` filters to one work; ``group_by=note`` returns them bucketed per + work instead of as a flat stream. + """ + async with get_session() as session: + comments = await service.list_comments(session, task_id, note_id, limit, platform) + + if group_by != "note": + return {"comments": comments, "total": len(comments)} + + buckets: Dict[str, Dict[str, Any]] = {} + for comment in comments: + bucket = buckets.setdefault( + comment["note_id"], + { + "note_id": comment["note_id"], + "note_title": comment["note_title"], + "note_cover": comment["note_cover"], + "note_url": comment["note_url"], + "comments": [], + }, + ) + bucket["comments"].append(comment) + + ordered = sorted( + buckets.values(), + key=lambda group: group["comments"][0]["first_seen_at"], + reverse=True, + ) + return {"groups": ordered, "total": len(comments)} + + +@router.get("/comment-notes") +async def list_comment_notes(task_id: Optional[int] = None, platform: Optional[str] = None): + """Works that have comments, newest first, with counts. + + Feeds the comment filter dropdown so the operator can pick by title. + """ + async with get_session() as session: + return {"notes": await service.comment_note_groups(session, task_id, platform)} + + +@router.get("/events") +async def list_events( + task_id: Optional[int] = None, + type: Optional[str] = None, + since_id: Optional[int] = None, + limit: int = Query(default=200, ge=1, le=2000), + platform: Optional[str] = None, +): + async with get_session() as session: + events = await service.list_events(session, task_id, type, since_id, limit, platform) + return {"events": events, "latest_id": events[0]["id"] if events else since_id} + + +@router.post("/events/read") +async def mark_events_read(task_id: Optional[int] = None): + async with get_session() as session: + count = await service.mark_events_read(session, task_id) + return {"marked": count} + + +# --------------------------------------------------------------------------- +# Cookie / login health +# --------------------------------------------------------------------------- + +@router.get("/cookie") +async def get_cookie_endpoint(platform: str = Query(default=PLATFORM_XHS)): + """Cookie health only -- deliberately never returns the cookie value. + + ``platform`` defaults to Xiaohongshu so existing callers keep working; the + key it reads is the namespaced one. + """ + async with get_session() as session: + return await get_cookie_status(session, platform) + + +@router.post("/cookie") +async def set_cookie_endpoint(payload: CookiePayload, platform: str = Query(default=PLATFORM_XHS)): + async with get_session() as session: + await set_cookie(session, payload.cookie.strip(), platform) + return {"message": "Cookie saved"} + + +@router.delete("/cookie") +async def clear_cookie_endpoint(platform: str = Query(default=PLATFORM_XHS)): + async with get_session() as session: + await delete_setting(session, cookie_key(platform)) + return {"message": "Cookie cleared"} + + +# --------------------------------------------------------------------------- +# Report +# --------------------------------------------------------------------------- + +async def _resolve_scope( + session, task_ids: Optional[List[int]], platform: Optional[str] +) -> Optional[List[int]]: + """Combine an explicit task selection with an optional platform filter. + + ``None`` means "no restriction"; an explicit list is intersected with the + platform's tasks so a stale selection cannot leak another platform's data + into a scoped report. + """ + if platform is None: + return task_ids + + platform_ids = set(await service.platform_task_ids(session, platform)) + if task_ids is None: + return list(platform_ids) + return [task for task in task_ids if task in platform_ids] + + +@router.get("/export") +async def export_data( + kind: str = Query(..., description="notes | comments | report"), + task_id: Optional[List[int]] = Query(default=None), + note_id: Optional[str] = None, + start_date: Optional[str] = Query(default=None, description="YYYY-MM-DD,report 用"), + end_date: Optional[str] = Query(default=None, description="YYYY-MM-DD,report 用"), + days: int = Query(default=7, ge=1, le=365), + file_format: str = Query(default="csv", alias="format", description="csv | xlsx"), + platform: Optional[str] = None, +): + """Download collected data as CSV or Excel. + + Reached by the browser as a navigation (``window.open``), which cannot carry + an Authorization header -- this is one of the reasons the session lives in a + cookie. + """ + if kind not in ("notes", "comments", "report"): + raise HTTPException(status_code=400, detail="kind 必须是 notes / comments / report") + if file_format not in ("csv", "xlsx"): + raise HTTPException(status_code=400, detail="format 必须是 csv 或 xlsx") + + async with get_session() as session: + scoped = await _resolve_scope(session, task_id, platform) + + if kind == "notes": + single_task = scoped[0] if scoped and len(scoped) == 1 else None + rows = await service.list_notes(session, single_task, False, 5000, platform) + elif kind == "comments": + single_task = scoped[0] if scoped and len(scoped) == 1 else None + rows = await service.list_comments(session, single_task, note_id, 5000, platform) + else: + try: + end_day = date.fromisoformat(end_date) if end_date else date.today() + start_day = ( + date.fromisoformat(start_date) + if start_date + else end_day - timedelta(days=days - 1) + ) + except ValueError: + raise HTTPException(status_code=400, detail="日期格式应为 YYYY-MM-DD") + # Not `report = ...`: that would make `report` a local name for the + # whole function and shadow the module import on this very line. + report_data = await report.build_report(session, scoped, start_day, end_day) + rows = report_data["rows"] + + if not rows: + raise HTTPException(status_code=404, detail="该范围内没有数据可导出") + + columns = _export_columns(kind) + stamp = date.today().isoformat() + # ASCII on purpose: a non-ASCII filename needs RFC 5987 encoding in + # Content-Disposition, and the plain `filename="..."` form used below would + # mangle it. + filename = f"export_{kind}_{stamp}.{file_format}" + + if file_format == "xlsx": + payload = _to_xlsx(rows, columns, kind) + media_type = "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet" + else: + payload = _to_csv(rows, columns) + media_type = "text/csv; charset=utf-8" + + return Response( + content=payload, + media_type=media_type, + headers={"Content-Disposition": f'attachment; filename="{filename}"'}, + ) + + +def _export_columns(kind: str) -> List[tuple[str, str]]: + """(key, header) pairs per export kind.""" + if kind == "notes": + return [ + ("note_id", "作品ID"), + ("title", "标题"), + ("note_url", "链接"), + ("liked_count", "点赞"), + ("comment_count", "评论"), + ("collected_count", "收藏"), + ("share_count", "分享"), + ("liked_count_delta", "点赞增量"), + ("comment_count_delta", "评论增量"), + ("first_seen_at", "首次发现"), + ("last_seen_at", "最近采集"), + ] + if kind == "comments": + return [ + ("note_title", "所属作品"), + ("note_id", "作品ID"), + ("comment_id", "评论ID"), + ("content", "内容"), + ("nickname", "昵称"), + ("like_count", "点赞"), + ("sub_comment_count", "子评论数"), + ("create_time", "发布时间"), + ("first_seen_at", "首次发现"), + ] + return [ + ("date", "日期"), + ("new_notes", "新增作品"), + ("new_comments", "新增评论"), + ("liked_count_delta", "点赞增量"), + ("comment_count_delta", "评论增量"), + ("collected_count_delta", "收藏增量"), + ("share_count_delta", "分享增量"), + ] + + +def _cell(value: Any) -> Any: + if value is None: + return "" + if isinstance(value, (list, dict)): + return ", ".join(str(v) for v in value) if isinstance(value, list) else str(value) + return value + + +def _to_csv(rows: List[Dict[str, Any]], columns: List[tuple[str, str]]) -> bytes: + import csv + import io + + buffer = io.StringIO() + writer = csv.writer(buffer) + writer.writerow([header for _, header in columns]) + for row in rows: + writer.writerow([_cell(row.get(key)) for key, _ in columns]) + + # utf-8-sig: without the BOM Excel opens Chinese CSV as mojibake, which is + # the single most common complaint about CSV exports here. + return buffer.getvalue().encode("utf-8-sig") + + +def _to_xlsx(rows: List[Dict[str, Any]], columns: List[tuple[str, str]], sheet: str) -> bytes: + import io + + from openpyxl import Workbook + + workbook = Workbook() + worksheet = workbook.active + worksheet.title = {"notes": "作品", "comments": "评论"}.get(sheet, "报表") + worksheet.append([header for _, header in columns]) + for row in rows: + worksheet.append([_cell(row.get(key)) for key, _ in columns]) + + output = io.BytesIO() + workbook.save(output) + return output.getvalue() + + +@router.get("/report") +async def get_report( + task_id: Optional[List[int]] = Query( + default=None, description="Repeat to include several tasks; omit for all" + ), + start_date: Optional[str] = Query(default=None, description="YYYY-MM-DD"), + end_date: Optional[str] = Query(default=None, description="YYYY-MM-DD"), + days: int = Query(default=7, ge=1, le=365, description="Window used when dates are omitted"), + platform: Optional[str] = None, +): + """Daily new-content counts and interaction deltas for the selected tasks.""" + try: + end_day = date.fromisoformat(end_date) if end_date else date.today() + start_day = date.fromisoformat(start_date) if start_date else end_day - timedelta(days=days - 1) + except ValueError: + raise HTTPException(status_code=400, detail="日期格式应为 YYYY-MM-DD") + + if start_day > end_day: + raise HTTPException(status_code=400, detail="开始日期不能晚于结束日期") + + async with get_session() as session: + scoped = await _resolve_scope(session, task_id, platform) + return await report.build_report(session, scoped, start_day, end_day) + + +# --------------------------------------------------------------------------- +# WeCom webhook +# --------------------------------------------------------------------------- + +def _mask_webhook(url: str) -> str: + """Show enough of the URL to recognise it, without exposing the robot key.""" + if not url: + return "" + key_marker = "key=" + index = url.find(key_marker) + if index == -1: + return url[:12] + "..." if len(url) > 12 else url + prefix = url[: index + len(key_marker)] + key = url[index + len(key_marker) :] + if len(key) <= 8: + return prefix + "*" * len(key) + return f"{prefix}{key[:4]}...{key[-4:]}" + + +@router.get("/webhook") +async def get_webhook(): + async with get_session() as session: + url = (await get_setting(session, SETTING_WECOM_WEBHOOK)) or "" + return {"configured": bool(url), "masked": _mask_webhook(url)} + + +@router.post("/webhook") +async def set_webhook(payload: WebhookPayload): + url = payload.url.strip() + if url and "qyapi.weixin.qq.com" not in url: + # Catches the common mistake of pasting a group-chat invite or the app + # URL instead of the robot webhook. + raise HTTPException( + status_code=400, + detail="这不像企业微信机器人 Webhook 地址(应包含 qyapi.weixin.qq.com)", + ) + + async with get_session() as session: + await set_setting(session, SETTING_WECOM_WEBHOOK, url) + return {"message": "Webhook 已保存" if url else "Webhook 已清空", "configured": bool(url)} + + +@router.delete("/webhook") +async def clear_webhook(): + async with get_session() as session: + await delete_setting(session, SETTING_WECOM_WEBHOOK) + return {"message": "Webhook 已删除"} + + +@router.post("/webhook/test") +async def test_webhook(payload: WebhookTestPayload): + """Send a test message so the user can verify the robot works before relying on it.""" + async with get_session() as session: + url = payload.url.strip() if payload.url else await notify.get_webhook_url(session) + + ok, detail = await notify.send_wecom( + url, "**综合采集平台 通知测试**\n> 如果你看到这条消息,说明 Webhook 配置成功。" + ) + if not ok: + raise HTTPException(status_code=400, detail=detail) + return {"message": detail} diff --git a/api/routers/settings.py b/api/routers/settings.py new file mode 100644 index 0000000..8b4e0d1 --- /dev/null +++ b/api/routers/settings.py @@ -0,0 +1,70 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/routers/settings.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Unified settings endpoint. + +Consolidates what used to be scattered across the monitor router. The older +``/api/monitor/cookie`` and ``/api/monitor/webhook`` endpoints are deliberately +left in place -- they still work and removing them would be a breaking change for +no gain. +""" + +from typing import Any, Dict + +from fastapi import APIRouter, HTTPException, Query + +from ..monitor import app_settings +from ..monitor.db import get_session +from ..monitor.platforms import PLATFORM_XHS +from ..schemas.settings import SettingsUpdatePayload + +router = APIRouter(prefix="/settings", tags=["settings"]) + + +@router.get("") +async def read_settings(platform: str = Query(default=PLATFORM_XHS)): + """Settings for one platform, plus the system-wide ones. + + Platform-scoped values are returned for the requested platform; system-scoped + values are the same regardless. Every spec carries its resolved ``key`` so the + UI can PUT changes straight back. + + Sensitive values are returned as ``{present, length}`` only. + """ + async with get_session() as session: + return await app_settings.get_all(session, platform) + + +@router.put("") +async def write_settings( + payload: SettingsUpdatePayload, platform: str = Query(default=PLATFORM_XHS) +): + """Partial update: only the keys present in the body are written. + + A key belonging to a different platform is rejected rather than written + somewhere unexpected. + """ + values: Dict[str, Any] = payload.values() + + async with get_session() as session: + try: + changed = await app_settings.update(session, values, platform) + except app_settings.SettingValidationError as exc: + raise HTTPException(status_code=400, detail=str(exc)) + + return {"message": f"已保存 {len(changed)} 项设置", "changed": changed} diff --git a/api/routers/websocket.py b/api/routers/websocket.py index 215d4ee..a6c92f2 100644 --- a/api/routers/websocket.py +++ b/api/routers/websocket.py @@ -19,8 +19,9 @@ import asyncio from typing import Set, Optional -from fastapi import APIRouter, WebSocket, WebSocketDisconnect +from fastapi import APIRouter, Depends, WebSocket, WebSocketDisconnect +from ..auth import require_ws_auth from ..services import crawler_manager router = APIRouter(tags=["websocket"]) @@ -86,7 +87,10 @@ def start_broadcaster(): _broadcaster_task = asyncio.create_task(log_broadcaster()) -@router.websocket("/ws/logs") +# Websocket routes need their own auth dependency: BaseHTTPMiddleware returns +# early for any non-http scope, and HTTP router-level dependencies do not reach +# websocket routes. Without this the live crawl log stream would be wide open. +@router.websocket("/ws/logs", dependencies=[Depends(require_ws_auth)]) async def websocket_logs(websocket: WebSocket): """WebSocket log stream""" print("[WS] New connection attempt") @@ -134,7 +138,7 @@ async def websocket_logs(websocket: WebSocket): print(f"[WS] Cleanup done, active connections: {len(manager.active_connections)}") -@router.websocket("/ws/status") +@router.websocket("/ws/status", dependencies=[Depends(require_ws_auth)]) async def websocket_status(websocket: WebSocket): """WebSocket status stream""" await websocket.accept() diff --git a/api/schemas/auth.py b/api/schemas/auth.py new file mode 100644 index 0000000..8f76c58 --- /dev/null +++ b/api/schemas/auth.py @@ -0,0 +1,34 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/schemas/auth.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Request models for authentication endpoints.""" + +from pydantic import BaseModel, Field + +# A floor, not a policy: this is a single-operator internal panel, so the goal is +# only to reject obviously weak input. +MIN_PASSWORD_LENGTH = 8 + + +class LoginPayload(BaseModel): + password: str = Field(min_length=1) + + +class ChangePasswordPayload(BaseModel): + current: str = Field(min_length=1) + new: str = Field(min_length=MIN_PASSWORD_LENGTH, max_length=256) diff --git a/api/schemas/crawler.py b/api/schemas/crawler.py index cfc995e..c9bcc8b 100644 --- a/api/schemas/crawler.py +++ b/api/schemas/crawler.py @@ -78,6 +78,34 @@ class CrawlerStartRequest(BaseModel): max_notes_count: Optional[int] = Field(default=None, ge=1, le=MAX_API_LIMIT_COUNT) max_comments_count: Optional[int] = Field(default=None, ge=1, le=MAX_API_LIMIT_COUNT) + # --- Options only used by scheduled monitor runs. Each defaults to None so + # the corresponding CLI flag is omitted entirely for manual Crawl-tab runs, + # which keeps their behaviour byte-identical to before. + + # Isolate this run's output in its own directory. The crawler's own file + # writer names files by date only, so same-day runs would otherwise append + # into one shared file and could not be told apart. + save_data_path: Optional[str] = None + # Unattended runs must not try to attach to the user's desktop Chrome. + enable_cdp_mode: Optional[bool] = None + # XHS cookie login only injects `web_session` by default, which is not enough + # to sign API requests from a cold browser profile. + inject_all_cookies: Optional[bool] = None + save_login_state: Optional[bool] = None + # Preferred over `cookies`: a value on the command line is visible in the + # process list. + cookies_file: Optional[str] = None + max_concurrency_num: Optional[int] = Field(default=None, ge=1, le=MAX_API_LIMIT_COUNT) + + # Crawl-strategy and proxy knobs surfaced on the Settings page. Like the + # fields above, each stays None unless the caller sets it, so the CLI flag is + # omitted entirely and the config-file default applies. + crawler_max_sleep_sec: Optional[int] = Field(default=None, ge=0, le=600) + enable_ip_proxy: Optional[bool] = None + ip_proxy_pool_count: Optional[int] = Field(default=None, ge=1, le=100) + ip_proxy_provider_name: Optional[str] = None + static_proxy_url: Optional[str] = None + class CrawlerStatusResponse(BaseModel): """Crawler status response""" diff --git a/api/schemas/monitor.py b/api/schemas/monitor.py new file mode 100644 index 0000000..f396503 --- /dev/null +++ b/api/schemas/monitor.py @@ -0,0 +1,81 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/schemas/monitor.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Request models for the monitoring API.""" + +from typing import List, Literal, Optional + +from pydantic import BaseModel, Field + +# A floor on the interval is a correctness guard, not a nicety: every run +# launches a browser and hits XHS with several requests, so a short interval +# across many creators is the pattern that triggers rate limiting. +MIN_INTERVAL_MINUTES = 30 +MAX_INTERVAL_MINUTES = 7 * 24 * 60 + + +class MonitorTaskCreate(BaseModel): + name: str = Field(min_length=1, max_length=200) + platform: str = "xhs" + # One subprocess handles exactly one crawler type, so a task is either + # creator-driven or note-driven. + mode: Literal["creator", "note"] + # None means "use the value configured on the Settings page", which is what + # makes those defaults meaningful. Bounds still apply when a value is given. + interval_minutes: Optional[int] = Field( + default=None, ge=MIN_INTERVAL_MINUTES, le=MAX_INTERVAL_MINUTES + ) + max_notes_count: Optional[int] = Field(default=None, ge=1, le=500) + enable_comments: bool = True + # Raising this widens the comment window, which is the only lever available + # for noticing new comments -- the API has no time-sort. + max_comments_count: Optional[int] = Field(default=None, ge=1, le=500) + run_timeout_seconds: int = Field(default=3600, ge=60, le=86400) + enabled: bool = True + # Push a WeCom summary for runs that failed or found new works. Opt-in per + # task so a single webhook does not get flooded. + notify_enabled: bool = False + # Raw pasted values: full URLs or bare ids, in either form. + targets: List[str] = Field(min_length=1) + + +class MonitorTaskUpdate(BaseModel): + name: Optional[str] = Field(default=None, min_length=1, max_length=200) + enabled: Optional[bool] = None + interval_minutes: Optional[int] = Field( + default=None, ge=MIN_INTERVAL_MINUTES, le=MAX_INTERVAL_MINUTES + ) + max_notes_count: Optional[int] = Field(default=None, ge=1, le=500) + enable_comments: Optional[bool] = None + max_comments_count: Optional[int] = Field(default=None, ge=1, le=500) + run_timeout_seconds: Optional[int] = Field(default=None, ge=60, le=86400) + notify_enabled: Optional[bool] = None + # When present, replaces the whole target list. + targets: Optional[List[str]] = None + + +class CookiePayload(BaseModel): + cookie: str = Field(min_length=1) + + +class WebhookPayload(BaseModel): + url: str = Field(default="", description="企业微信机器人 Webhook 地址,留空表示停用") + + +class WebhookTestPayload(BaseModel): + url: Optional[str] = Field(default=None, description="不传则使用已保存的地址") diff --git a/api/schemas/settings.py b/api/schemas/settings.py new file mode 100644 index 0000000..148033e --- /dev/null +++ b/api/schemas/settings.py @@ -0,0 +1,37 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/schemas/settings.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Request model for the settings endpoint.""" + +from typing import Any, Dict + +from pydantic import BaseModel, ConfigDict + + +class SettingsUpdatePayload(BaseModel): + """Partial update of arbitrary setting keys. + + Fields are not declared here on purpose: ``api/monitor/app_settings.py`` + owns the registry (key, type, bounds, choices) and validates against it, so + adding a setting does not mean editing a matching schema. + """ + + model_config = ConfigDict(extra="allow") + + def values(self) -> Dict[str, Any]: + return dict(self.model_extra or {}) diff --git a/api/services/crawler_manager.py b/api/services/crawler_manager.py index 83119cc..95138c5 100644 --- a/api/services/crawler_manager.py +++ b/api/services/crawler_manager.py @@ -25,6 +25,7 @@ from datetime import datetime from pathlib import Path from ..schemas import CrawlerStartRequest, LogEntry +from .interpreter import resolve_python_cmd class CrawlerManager: @@ -43,6 +44,11 @@ class CrawlerManager: self._project_root = Path(__file__).parent.parent.parent # Log queue - for pushing to WebSocket self._log_queue: Optional[asyncio.Queue] = None + # Completion signalling for run_and_wait(). Polling `status` is unreliable + # because stop() also resets it to "idle", and `self.process` gets replaced + # by any concurrent start(), so waiters need an explicit event instead. + self._done: asyncio.Event = asyncio.Event() + self.last_exit_code: Optional[int] = None @property def logs(self) -> List[LogEntry]: @@ -54,6 +60,43 @@ class CrawlerManager: self._log_queue = asyncio.Queue() return self._log_queue + def is_busy(self) -> bool: + """Whether a crawler process is currently alive. + + This is the authoritative busy check -- `status` is a lagging indicator + that manual stop() also resets. + """ + return self.process is not None and self.process.poll() is None + + async def run_and_wait( + self, + config: CrawlerStartRequest, + extra_args: Optional[List[str]] = None, + timeout: Optional[float] = None, + ) -> int: + """Start a crawler run and block until it exits, returning the exit code. + + Used by the monitor scheduler. Returns a negative value if the run was + killed by `timeout` or if the process could not be started at all. + """ + started = await self.start(config, extra_args=extra_args) + if not started: + return -1 + + # Capture the process we just launched: a concurrent start() would + # replace self.process, so poll this reference rather than the attribute. + proc = self.process + if proc is None: + return -1 + + try: + await asyncio.wait_for(self._done.wait(), timeout=timeout) + except asyncio.TimeoutError: + await self.stop() + return -1 + + return self.last_exit_code if self.last_exit_code is not None else -1 + def _create_log_entry(self, message: str, level: str = "info") -> LogEntry: """Create log entry""" self._log_id += 1 @@ -90,7 +133,11 @@ class CrawlerManager: return "debug" return "info" - async def start(self, config: CrawlerStartRequest) -> bool: + async def start( + self, + config: CrawlerStartRequest, + extra_args: Optional[List[str]] = None, + ) -> bool: """Start crawler process""" async with self._lock: if self.process and self.process.poll() is None: @@ -99,6 +146,9 @@ class CrawlerManager: # Clear old logs self._logs = [] self._log_id = 0 + # Reset completion signalling for this run + self._done.clear() + self.last_exit_code = None # Clear pending queue (don't replace object to avoid WebSocket broadcast coroutine holding old queue reference) if self._log_queue is None: @@ -111,7 +161,7 @@ class CrawlerManager: pass # Build command line arguments - cmd = self._build_command(config) + cmd = self._build_command(config, extra_args=extra_args) # Log start information entry = self._create_log_entry(f"Starting crawler: {' '.join(cmd)}", "info") @@ -202,9 +252,13 @@ class CrawlerManager: "error_message": None } - def _build_command(self, config: CrawlerStartRequest) -> list: + def _build_command( + self, + config: CrawlerStartRequest, + extra_args: Optional[List[str]] = None, + ) -> list: """Build main.py command line arguments""" - cmd = ["uv", "run", "python", "main.py"] + cmd = [*resolve_python_cmd(), "main.py"] cmd.extend(["--platform", config.platform.value]) cmd.extend(["--lt", config.login_type.value]) @@ -232,22 +286,56 @@ class CrawlerManager: if config.max_comments_count is not None: cmd.extend(["--max_comments_count_singlenotes", str(config.max_comments_count)]) - if config.cookies: + # Each of these is only appended when explicitly set, so manual runs from + # the Crawl tab keep exactly their previous behaviour. + if config.save_data_path: + cmd.extend(["--save_data_path", config.save_data_path]) + if config.enable_cdp_mode is not None: + cmd.extend(["--enable_cdp_mode", "true" if config.enable_cdp_mode else "false"]) + if config.inject_all_cookies is not None: + cmd.extend(["--inject_all_cookies", "true" if config.inject_all_cookies else "false"]) + if config.save_login_state is not None: + cmd.extend(["--save_login_state", "true" if config.save_login_state else "false"]) + if config.max_concurrency_num is not None: + cmd.extend(["--max_concurrency_num", str(config.max_concurrency_num)]) + if config.crawler_max_sleep_sec is not None: + cmd.extend(["--crawler_max_sleep_sec", str(config.crawler_max_sleep_sec)]) + if config.enable_ip_proxy is not None: + cmd.extend(["--enable_ip_proxy", "true" if config.enable_ip_proxy else "false"]) + if config.ip_proxy_pool_count is not None: + cmd.extend(["--ip_proxy_pool_count", str(config.ip_proxy_pool_count)]) + if config.ip_proxy_provider_name: + cmd.extend(["--ip_proxy_provider_name", config.ip_proxy_provider_name]) + if config.static_proxy_url: + cmd.extend(["--static_proxy_url", config.static_proxy_url]) + + # Prefer a cookie file over passing the cookie on the command line, where + # it would be visible in the process list. + if config.cookies_file: + cmd.extend(["--cookies_file", config.cookies_file]) + elif config.cookies: cmd.extend(["--cookies", config.cookies]) cmd.extend(["--headless", "true" if config.headless else "false"]) + if extra_args: + cmd.extend(extra_args) + return cmd async def _read_output(self): """Asynchronously read process output""" loop = asyncio.get_event_loop() + # Capture the process this reader was started for. self.process can be + # replaced by a subsequent start(), which would otherwise make us read + # the exit code of the wrong run. + proc = self.process try: - while self.process and self.process.poll() is None: + while proc and proc.poll() is None: # Read a line in thread pool line = await loop.run_in_executor( - None, self.process.stdout.readline + None, proc.stdout.readline ) if line: line = line.strip() @@ -257,9 +345,9 @@ class CrawlerManager: await self._push_log(entry) # Read remaining output - if self.process and self.process.stdout: + if proc and proc.stdout: remaining = await loop.run_in_executor( - None, self.process.stdout.read + None, proc.stdout.read ) if remaining: for line in remaining.strip().split('\n'): @@ -270,7 +358,7 @@ class CrawlerManager: # Process ended if self.status == "running": - exit_code = self.process.returncode if self.process else -1 + exit_code = proc.returncode if proc else -1 if exit_code == 0: entry = self._create_log_entry("Crawler completed successfully", "success") else: @@ -283,6 +371,11 @@ class CrawlerManager: except Exception as e: entry = self._create_log_entry(f"Error reading output: {str(e)}", "error") await self._push_log(entry) + finally: + # Record the exit code and wake any run_and_wait() waiter. Runs in a + # finally so a cancelled read task still releases the waiter. + self.last_exit_code = proc.returncode if proc else None + self._done.set() # Global singleton diff --git a/api/services/interpreter.py b/api/services/interpreter.py new file mode 100644 index 0000000..e4d569c --- /dev/null +++ b/api/services/interpreter.py @@ -0,0 +1,70 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/api/services/interpreter.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Interpreter resolution for spawning crawler subprocesses. + +Historically both the crawler manager and the environment check hardcoded +``uv run``. ``uv`` is not guaranteed to be installed, so resolve the command +prefix in one place: prefer ``uv`` (matching upstream docs), fall back to a +project-local virtualenv, and finally to the interpreter running the server. +""" + +import shutil +import sys +from pathlib import Path + +# Project root: api/services/interpreter.py -> services -> api -> repo root +PROJECT_ROOT = Path(__file__).parent.parent.parent + + +def venv_python_path(project_root: Path | None = None) -> Path: + """Return the path to the project venv's Python executable.""" + root = project_root if project_root is not None else PROJECT_ROOT + if sys.platform == "win32": + return root / ".venv" / "Scripts" / "python.exe" + return root / ".venv" / "bin" / "python" + + +def resolve_python_cmd(project_root: Path | None = None) -> list[str]: + """Resolve the command prefix used to run ``main.py``. + + Order of preference: + 1. ``uv`` if it is on PATH -- matches the upstream documented workflow. + 2. The project-local ``.venv`` if it exists. + 3. The interpreter currently running the API server. + + Returns a list because the caller appends ``main.py`` and its flags. + """ + if shutil.which("uv"): + return ["uv", "run", "python"] + + venv_python = venv_python_path(project_root) + if venv_python.exists(): + return [str(venv_python)] + + return [sys.executable] + + +def describe_interpreter(project_root: Path | None = None) -> str: + """Human-readable description of what resolve_python_cmd() picks.""" + cmd = resolve_python_cmd(project_root) + if cmd[0] == "uv": + return "uv run python" + if cmd[0] == sys.executable: + return f"current interpreter ({sys.executable})" + return f"project virtualenv ({cmd[0]})" diff --git a/cmd_arg/arg.py b/cmd_arg/arg.py index 8ea133d..cf6fd7f 100644 --- a/cmd_arg/arg.py +++ b/cmd_arg/arg.py @@ -300,6 +300,14 @@ async def parse_cmd(argv: Optional[Sequence[str]] = None): rich_help_panel="Performance Configuration", ), ] = config.MAX_CONCURRENCY_NUM, + crawler_max_sleep_sec: Annotated[ + int, + typer.Option( + "--crawler_max_sleep_sec", + help="Seconds to wait between requests. Higher is slower but far less likely to trip platform rate limiting", + rich_help_panel="Performance Configuration", + ), + ] = config.CRAWLER_MAX_SLEEP_SEC, save_data_path: Annotated[ str, typer.Option( @@ -308,6 +316,41 @@ async def parse_cmd(argv: Optional[Sequence[str]] = None): rich_help_panel="Storage Configuration", ), ] = config.SAVE_DATA_PATH, + enable_cdp_mode: Annotated[ + str, + typer.Option( + "--enable_cdp_mode", + help="Whether to drive the user's local Chrome over CDP instead of launching a browser, supports yes/true/t/y/1 or no/false/f/n/0. Set to false for unattended/server runs", + rich_help_panel="Runtime Configuration", + show_default=True, + ), + ] = str(config.ENABLE_CDP_MODE), + save_login_state: Annotated[ + str, + typer.Option( + "--save_login_state", + help="Whether to persist the browser profile so a previous login can be reused, supports yes/true/t/y/1 or no/false/f/n/0", + rich_help_panel="Runtime Configuration", + show_default=True, + ), + ] = str(config.SAVE_LOGIN_STATE), + inject_all_cookies: Annotated[ + str, + typer.Option( + "--inject_all_cookies", + help="Whether to inject every cookie supplied via --cookies/--cookies_file instead of only web_session, supports yes/true/t/y/1 or no/false/f/n/0", + rich_help_panel="Runtime Configuration", + show_default=True, + ), + ] = str(config.INJECT_ALL_COOKIES), + cookies_file: Annotated[ + str, + typer.Option( + "--cookies_file", + help="Path to a file holding the cookie string. Preferred over --cookies, whose value is visible in the process list", + rich_help_panel="Runtime Configuration", + ), + ] = "", enable_ip_proxy: Annotated[ str, typer.Option( @@ -350,6 +393,18 @@ async def parse_cmd(argv: Optional[Sequence[str]] = None): enable_headless = _to_bool(headless) enable_ip_proxy_value = _to_bool(enable_ip_proxy) init_db_value = init_db.value if init_db else None + enable_cdp_mode_value = _to_bool(enable_cdp_mode) + save_login_state_value = _to_bool(save_login_state) + inject_all_cookies_value = _to_bool(inject_all_cookies) + + # A file is preferred over --cookies: a literal value on the command line + # is visible to any other user on the machine via the process list. + if cookies_file: + try: + with open(cookies_file, "r", encoding="utf-8") as f: + cookies = f.read().strip() + except OSError as e: + raise typer.BadParameter(f"Unable to read --cookies_file: {e}") # Parse specified_id and creator_id into lists specified_id_list = [id.strip() for id in specified_id.split(",") if id.strip()] if specified_id else [] @@ -368,9 +423,13 @@ async def parse_cmd(argv: Optional[Sequence[str]] = None): config.CDP_HEADLESS = enable_headless config.SAVE_DATA_OPTION = save_data_option.value config.COOKIES = cookies + config.ENABLE_CDP_MODE = enable_cdp_mode_value + config.SAVE_LOGIN_STATE = save_login_state_value + config.INJECT_ALL_COOKIES = inject_all_cookies_value config.CRAWLER_MAX_COMMENTS_COUNT_SINGLENOTES = max_comments_count_singlenotes config.CRAWLER_MAX_NOTES_COUNT = crawler_max_notes_count config.MAX_CONCURRENCY_NUM = max_concurrency_num + config.CRAWLER_MAX_SLEEP_SEC = crawler_max_sleep_sec config.SAVE_DATA_PATH = save_data_path config.ENABLE_IP_PROXY = enable_ip_proxy_value config.IP_PROXY_POOL_COUNT = ip_proxy_pool_count diff --git a/config/base_config.py b/config/base_config.py index cf6823c..0c8c06c 100644 --- a/config/base_config.py +++ b/config/base_config.py @@ -52,6 +52,12 @@ HEADLESS = False # Whether to save login status SAVE_LOGIN_STATE = True +# 是否注入完整 cookie(默认 False,保持上游原有行为)。 +# False 时 login_by_cookies 只写入 web_session;a1 / webId 等签名所需 cookie 只能靠 +# browser_data 下的持久化 profile 补齐。无人值守场景(服务器上跑定时监控)应设为 True, +# 否则冷 profile 下 API 签名失败,且表现为「退出码 0 但抓到 0 条」的静默失败。 +INJECT_ALL_COOKIES = False + # ==================== CDP (Chrome DevTools Protocol) 配置 ==================== # 是否启用 CDP 模式 - 使用用户本地的 Chrome/Edge 浏览器进行爬取,具有更好的反检测能力 # 开启后,会自动检测并启动用户的 Chrome/Edge 浏览器,通过 CDP 协议进行控制 diff --git a/docs/监控功能使用说明.md b/docs/监控功能使用说明.md new file mode 100644 index 0000000..dfec3fa --- /dev/null +++ b/docs/监控功能使用说明.md @@ -0,0 +1,476 @@ +# 小红书监控功能使用说明 + +> 定时重复采集一批博主或笔记,与上一轮快照对比,产出**新增作品 / 新增评论 / 点赞收藏评论数涨跌**。 + +本功能是在 MediaCrawler 之上新增的一层,代码集中在 `api/monitor/`,不侵入原有的 +`media_platform/`、`store/` 等目录。 + +--- + +## 一、为什么需要单独一层 + +原项目是**一次性采集**:跑完即退出,没有调度、没有历史、没有差分。直接复用会遇到三个硬伤: + +1. **指标会被覆盖**。`store/xhs/_store_impl.py::XhsDbStoreImplement.update_content()` 对已存在的笔记执行 + `UPDATE ... SET liked_count = ...`,历史值直接丢失。跑第二遍根本看不出"点赞从 100 涨到了 500"。 +2. **单进程串行**。`api/services/crawler_manager.py` 是全局单例,同一时刻只能跑一个 `main.py` 子进程。 +3. **运行输出无法区分**。`AsyncFileWriter` 的文件名只带日期(`creator_contents_2026-10-07.jsonl`), + 同一天多次运行会追加进同一个文件。 + +监控层为此做了对应处理:独立的快照库(保留历史)、调度器与手动采集互斥排队、以及**每轮采集写入独立目录** +(复用早已存在、但 API 层从未转发的 `--save_data_path` 参数)。 + +--- + +## 二、快速开始 + +### 1. 准备环境 + +```bash +# 依赖(若未安装 uv,本项目的解释器探测会自动回退到 .venv) +python -m venv .venv +.venv/Scripts/python -m pip install -r requirements.txt +.venv/Scripts/python -m playwright install chromium # 非 CDP 模式必需 + +# 前端 +cd webui && npm install && npm run build +``` + +### 2. 启动 + +```bash +.venv/Scripts/python -m api.main # 或 uvicorn api.main:app --port 8080 +``` + +打开 ,右上角切换到「监控」。 + +> 解释器探测顺序:`uv`(若在 PATH)→ 项目 `.venv` → 当前解释器。 +> `/api/env/check` 使用同一套逻辑,不会出现"检测失败但其实能跑"的情况。 + +### 3. 配置登录态(**无人值守的前提**) + +在监控页左下角「小红书登录态」粘贴 Cookie。定时监控不能每次都扫码,必须持久化登录态。 + +> **强烈建议先手动扫码登录一次**,以播种 `browser_data/xhs_user_data_dir`, +> 之后再粘贴 Cookie 才可靠。原因见下方「限制」。 + +### 4. 新建监控任务 + +- **类型** + - `博主`:监控其作品,填博主主页链接或纯 ID + - `笔记`:批量监控指定内容,填笔记链接或纯 ID +- **目标**:每行一个。**建议只填纯 ID** —— 链接里的 `xsec_token` 会过期,纯 ID 永久有效。 +- **间隔**:最小 30 分钟。每次运行都要拉起一次浏览器并多次请求平台,间隔过短容易触发风控。 +- **每篇评论抓取条数**:默认 50。这个值直接决定能发现多少新评论,见下方限制。 + +任务创建后立即生效,也可随时点「立即运行」手动触发一轮。 + +--- + +## 二·五、报表 + +「报表」视图是**跨任务**的统计,用来回答"这批账号这段时间表现如何"。 + +**筛选**:勾选参与统计的任务(默认全选),选日期区间(或点「近 7/30/90 天」)。 + +**两类指标,含义不同,所以分列展示**: + +| 列 | 含义 | +|---|---| +| 新增作品 / 新增评论 | 该日**首次发现**的作品数 / 评论条数 | +| 点赞 Δ / 评论 Δ / 收藏 Δ / 分享 Δ | 该日**互动增量**:Σ(当日末值 − 当日之前最后一次采到的值) | + +增量的口径有两个要点: + +- **作品首次出现的那天从 0 起算**,所以新作品的全部点赞都计入其首次发现日。这样做是为了让"新作品带来了多少赞"这件事可见,而不是把它的既有数据丢掉。 +- **某天没采到某篇作品,那天的增量算 0**,不会把跨天的增长平摊到每一天。 + +底部会标明两件事:一是**哪些指标无法解析**(小红书可能返回 `"1.2万"` 这类值,解析失败的不会被当成 0 计入,否则会伪造出一个大的负增长),二是评论数受接口限制只覆盖前 N 条。 + +> 实现上聚合是在 Python 里做的,不是一条大 SQL。原因:按笔记、按天的"上一个基线值"查询是窗口操作,SQLite 表达起来很别扭,而这里的数据量很小,可读性比压榨查询计划更值钱。 + +--- + +## 二·六、企业微信通知 + +在「监控」视图左下角配置 Webhook 地址(企业微信群 → 添加群机器人 → 复制 Webhook 地址)。 + +**两个设计取舍**: + +1. **一轮只发一条汇总**,不是每条事件发一条。一次跑出 20 篇新作品时,你收到的是"新增作品 20 篇"加前 10 条标题,而不是 20 条消息。 +2. **推送失败绝不影响采集**。通知是在数据提交之后、用独立会话发送的,任何网络错误只记日志。爬虫跑成功了不会因为 webhook 挂了而被回滚。 + +**触发时机**(仅这两类): + +- 任务失败 / 疑似登录态失效 +- 发现新增作品 + +指标变化和新增评论**不会**推送(指标变化太频繁,评论量可能很大)。 + +**任务范围**:每个任务在编辑弹窗里有「推送企业微信通知」开关,**默认关闭**。这样一个 webhook 不会被一堆无关任务刷屏。 + +- 配置好地址后可以点「发测试」验证,也可以「保存前先测」。 +- 地址里的 key 等同凭据,**服务端只回传打码形式**,要换只能重新粘贴(和 Cookie 一致)。 +- 任务卡片上的 `last_notified_at`(列表接口会返回)可以回答"为什么这轮没收到推送"。 + +--- + +## 二·七、评论视图与导出 + +「评论」页默认**按作品分组**:每篇作品一个可折叠区块,**默认只展开最新的一组**, +避免打开就是一屏文字。切到「平铺」则是一条流,每条评论下方标注它属于哪篇作品 +(封面缩略图 + 标题 + 跳原文链接)。 + +顶部可按作品筛选,选项里带每篇的评论数: + +``` +全部作品 +烤面筋热量计算 (33) +孜卷热量计算 (3) +``` + +> 评论与作品的关联是后端 JOIN 出来的(`note_title` / `note_cover` / `note_url`), +> 因为评论表本身只存 `note_id`,光看 ID 没有任何可读性。 + +### 导出 + +评论页和报表页都有「导出」按钮,走浏览器下载: + +| 端点 | 内容 | +|---|---| +| `?kind=notes` | 作品表(含互动增量列) | +| `?kind=comments` | 评论(含所属作品标题) | +| `?kind=report` | 报表按天汇总 | + +- 支持 `csv` 与 `xlsx` +- **CSV 带 UTF-8 BOM**(`utf-8-sig`)—— 否则 Excel 打开中文是乱码,这是最常见的投诉 +- 下载是**页面导航**(`window.open`),带不了自定义请求头,所以导出依赖 Cookie 鉴权 —— + 这也是会话必须存在 Cookie 里的原因之一 + +--- + +## 二·八、登录与访问控制 + +面板默认要求登录 —— `/api` 下的所有接口都需要会话,只有 `/api/health`、 +`/api/auth/login`、`/api/auth/logout` 例外。静态资源(页面本身、JS/CSS)不受限制, +否则登录页自己都加载不出来。 + +### 首次启动 + +自动生成一个随机密码并**打印在启动日志里**(只打印一次): + +``` +==================================================================== + WebUI 首次启动,已生成登录密码: + + 94Shn1fMa7dV0jqF + + 请立即登录并修改。 +==================================================================== +``` + +> 刻意**不做**"打开页面让你设置密码"的流程。在局域网监听下,任何能访问到的人 +> 都能抢先设置密码成为管理员;自动生成 + 打印避免了这种抢占,也避免了把自己锁在外面。 + +### 忘记密码 + +设置环境变量 `MC_PASSWORD` 后重启即可: + +```bash +MC_PASSWORD=我的新密码 # Linux/macOS +set MC_PASSWORD=我的新密码 # Windows cmd +``` + +该变量**优先级始终高于**数据库里的密码,且**不会被写入磁盘**。登录后到设置页改成正式密码即可。 + +### 环境变量 + +配置写在项目根目录的 `.env`(已被 gitignore)。 + +| 变量 | 默认 | 说明 | +|---|---|---| +| `MC_HOST` | `127.0.0.1` | 监听地址。**要局域网访问须设为 `0.0.0.0`** | +| `MC_PORT` | `8080` | 端口 | +| `MC_PASSWORD` | 空 | 覆盖数据库密码,忘记密码时的恢复通道 | +| `MC_COOKIE_SECURE` | 关 | **面板走 HTTPS 时才开**。局域网明文下开启会导致浏览器丢弃 Cookie,表现为**登录页反复刷新且无任何报错** | +| `MC_SESSION_TTL_HOURS` | `336` | 登录有效期(14 天) | +| `MC_TRUST_PROXY` | 关 | 仅在受信任的反向代理之后开启,否则 `X-Forwarded-For` 可被伪造以绕过登录节流 | +| `MC_CORS_ORIGINS` | 空 | 附加的允许来源,逗号分隔 | +| `MC_CORS_ORIGIN_REGEX` | 空 | 允许来源的正则,用于局域网里的 Vite 开发服务器 | + +> 本项目的 `.env` 此前**从未被加载过**(代码里没有任何 `load_dotenv` 调用,尽管 +> `python-dotenv` 一直是依赖、`.env.example` 也一直在仓库里)。现已修复。 + +### 安全边界(请务必了解) + +- **局域网是明文 HTTP**,所以 Cookie 没开 `Secure`,`SameSite=Lax`。 + 这意味着**同网段抓包能看到会话令牌**。安全边界是"内网 + 密码",不是传输加密。 +- **超出可信网络之外请走 HTTPS 反向代理**,不要把本服务直接暴露到公网。 +- **登录节流是进程内的**:重启即清零。单用户单 worker 场景足够; + 若日后多 worker,节流会按 worker 各算各的。反向代理下 `request.client.host` 是代理地址, + 需配合 `MC_TRUST_PROXY` 才能正确识别来源。 +- `/docs`、`/redoc`、`/openapi.json` **已关闭** —— 它们默认不鉴权,等于免费公开整个 API 地图。 + +### 会话与登出 + +- 会话存在服务端(`auth_session` 表),库里只存令牌的 **SHA-256**,不存令牌本身 +- 退出登录、**修改密码**都会立即失效(改密码会踢掉所有设备,并给当前设备补发一个新会话) +- 令牌可放在 Cookie(浏览器自动携带,WebSocket 与文件下载都依赖它) + 或 `Authorization: Bearer`(方便脚本调用) + +--- + +## 二·九、设置页 + +原先挤在监控页左下角的 Cookie 与 Webhook 面板已迁到这里,并补齐了采集策略、代理与账号安全。 + +### 分区与生效方式 + +| 分区 | 内容 | 生效时机 | +|---|---|---| +| 登录态 | 小红书 Cookie | 下一轮采集 | +| 通知 | 企业微信 Webhook | 下一条推送 | +| 采集策略 | 新任务默认间隔、默认单轮上限、默认评论条数 | **仅影响新建任务** | +| 采集策略 | 请求间隔、抓二级评论 | 下一轮采集 | +| 采集策略 | 活跃时段 | 定时任务的下一次触发 | +| 代理 | 开关、提供方、池大小、静态地址 | 下一轮采集 | +| 账号安全 | 修改密码 | 立即(其他设备全部掉线) | + +**活跃时段**:只在此时段内触发定时采集,窗口外任务保持到期状态、不会丢失, +窗口一开照常执行。默认 `0–23` 即全天;也支持跨午夜(如 `22–6`)。 + +**"仅影响新建任务"** 的那几项是刻意的:改了默认间隔不应该把已有任务的间隔一起改掉。 + +### 设计要点 + +- **敏感值永不回传**:Cookie 和 Webhook 的 `GET` 只返回「是否已配置」与长度,不返回值。 + 表单不会把没动过的敏感项覆盖掉。 +- **部分更新**:只有请求里出现的 key 会被写入。表单一角改动不会清空其他设置。 +- **设置项由后端声明**:`api/monitor/app_settings.py` 里的注册表(类型、范围、选项、默认值) + 是唯一事实来源,前端**按它生成表单**。加一个设置项不需要改前端字段清单。 +- **越界即拒绝**:超出范围、未知的 key、非法的枚举值都返回 400 而不是静默接受。 + +> 「扫码登录」入口**尚未实现**。它需要跑起爬虫子进程、捕获二维码并实时推流, +> 属于一个独立功能而非设置项,这里不做一个半成品。 + +--- + +## 二·十、平台切换与能力矩阵 + +**右上角的下拉框统一切换平台**,「采集 / 监控 / 报表 / 设置」全部跟着变。选择会记住, +刷新后不会跳回小红书。采集页原来那个平台下拉已移除,避免出现两个事实来源。 + +### 已接通 vs 未接通 + +矩阵里有两个**不同**的概念,混淆会误导: + +| 字段 | 含义 | +|---|---| +| `crawler_modes` / `metrics` / `comment_levels` / `media` | **上游爬虫模块**能做什么 | +| `monitor_wired` | **监控层**是否已接线 | + +**7 个平台的爬虫模块都实现了 search / detail / creator**,真正的差异在指标上: + +| 平台 | 指标 | 评论层级 | 媒体 | 监控接线 | +|---|---|---|---|---| +| 小红书 | 点赞 / 评论 / 收藏 / 分享 | 2 | ✅ | ✅ | +| 抖音 | 点赞 / 评论 / 收藏 / 分享(**无播放量**) | 2 | ✅ | ❌ | +| 快手 | 点赞 / 播放(无评论、分享、收藏) | 1 | ✅ | ❌ | +| B站 | 点赞 / **播放** / **弹幕** / 评论 / 收藏 / 投币 / 分享(最全) | 2 | ✅ | ❌ | +| 微博 | 点赞 / 评论 / 转发(无收藏) | 2 | ✅ | ❌ | +| 贴吧 | 仅回复数 | 2 | ❌ | ❌ | +| 知乎 | 赞同 / 评论 | 2 | ❌ | ❌ | + +> **要更正一个常见误解**:这个代码库里**抖音不存播放量**(只映射点赞/收藏/评论/分享)。 +> 有播放量的是 **B 站**,它还有弹幕。 + +未接通的平台**可以选,但各页会显示明确的说明面板**,并且**创建任务会被直接拒绝**: + +``` +400 抖音的爬虫已支持,但监控层尚未接通,暂时无法创建监控任务。 +``` + +而不是接受任务、然后让它永远跑不出数据 —— 那正是之前"博主主页解析失败被误报成登录失效"的同一种静默故障。 + +### 设置的两层 + +| 位置 | 范围 | 内容 | +|---|---|---| +| 左侧导航「设置」 | **按平台** | 登录 Cookie、采集策略、代理 | +| 右上角「系统设置」 | **全局** | 通知、活跃时段、账号安全 | + +**这不是随便分的**:企业微信只有一个群、调度器只有一套时段规则、密码只有一份 —— +把它们放进"小红书专属"的页面里,会让人以为它们是按平台存的。 + +存储上键名带作用域前缀:`platform.<平台>.<项>` 与 `system.<项>`。 +**旧键会在启动时自动迁移**(`xhs_cookie` → `platform.xhs.cookie`), +且是幂等的:新键已存在时以新键为准,不会覆盖你后来改的值。 + +--- + +## 三、必须知道的限制 + +### 1. 「新增评论」是近似值 —— 最重要的一条 + +小红书评论接口 `/api/sns/web/v2/comment/page` **没有排序参数**,只能拿到平台默认排序(热评优先)的 +前 N 条。因此: + +- 我们只能"每次抓前 N 条做差集",**新发布但沉底的评论不会被发现** +- N 调大能提高发现率,但请求量线性增长,风控风险上升 +- 评论事件区分两种,UI 上也分别标注: + - `new_comment_posted`(新评论):`create_time` 晚于上一轮开始时间,是真·新发布 + - `new_comment_seen`(新出现评论):只是本轮才进入可见窗口的历史评论 + +**这条限制无法通过调参绕过**,是该接口的固有限制。 + +### 2. Cookie 失效是「静默失败」 + +`login_by_cookies()` 只注入 `web_session`,而 API 签名还需要 `a1` / `webId` 等; +更麻烦的是**cookie 登录不做任何校验** —— 坏 Cookie 不会让进程报错退出,而是 +**退出码 0、抓到 0 条**。 + +监控层因此把「退出码 0 且 0 条作品」判定为 `suspected_auth_failure` 并在 UI 上标红, +而不是当成"该博主没发新作品"。这是无人值守场景最容易误报的地方。 + +本实现额外做了两件事: +- 通过 `--inject_all_cookies` 注入**完整** Cookie(默认关闭,保持上游行为不变) +- 通过 `--cookies_file` 传 Cookie,避免明文出现在进程列表里 + +### 3. 作品窗口被截断 + +`每轮最多采集作品数`(默认 20)限定了"该博主的作品"到底指多少条。 +UI 会把该上限显示在作品表旁,避免误以为看到了全部。 + +### 4. 昵称与用户 ID 已被上游脱敏 + +`store/xhs/__init__.py` 落库前调用 `mask_nickname()` 与 `anonymize_user_id()`, +存储的是**打码昵称**与哈希后的 `creator_hash`,没有真实昵称和 user_id。 +这是上游的隐私保护设计,监控层未做改动。 + +### 5. 不发「笔记被删」事件 + +`creator` 模式只取前 N 条,笔记"消失"多半只是掉出窗口;`detail` 模式遇到 +`xsec_token` 过期也会失败。二者与"真被删"无法区分,因此不产生删除事件, +改为在作品表里展示 `last_seen_at`。 + +### 6. 定时任务与手动采集互斥 + +二者共用同一个爬虫子进程。监控任务运行期间点「采集」会被拒绝(返回 400); +反之若有手动采集在跑,到期的监控任务会**保持到期状态排队**,不会丢失,空闲后自动补上。 + +--- + +## 四、数据存放 + +| 内容 | 位置 | +|---|---| +| 监控库(任务/快照/事件/评论/设置/会话) | **MySQL**,库名由 `MYSQL_DB_NAME` 指定 | +| 每轮原始 jsonl | `data/monitor_runs/{task_id}/{run_id}/{platform}/jsonl/` | + +> 爬虫每轮的原始产出**仍然写独立 jsonl 目录**,不进 MySQL。 +> 这是差分机制的基础:每轮写在单独目录里,才能算出"这轮新增了什么"。 +> 多轮数据混在同一批表里的话,这个判断就做不到了。 + +### MySQL 配置与安全边界 + +连接信息写在 `.env`(已被 gitignore,不会进版本库): + +```ini +MYSQL_DB_HOST=<数据库地址> +MYSQL_DB_PORT=3306 +MYSQL_DB_USER=<账号> +MYSQL_DB_PWD=<密码> +MYSQL_DB_NAME=mediacrawler +``` + +> 真实凭据只写在 `.env` 里(已被 gitignore),**不要写进这个文档或任何会提交的文件**。 + +**"只操作这个库"由两层保证,缺一不可**: + +1. **数据库授权(真正的保证)**。账号应只被授予目标库的权限: + + ```sql + REVOKE ALL PRIVILEGES, GRANT OPTION FROM 'MediaCrawler'@'%'; + GRANT ALL PRIVILEGES ON `mediacrawler`.* TO 'MediaCrawler'@'%'; + FLUSH PRIVILEGES; + ``` + + 这样该账号 `SHOW DATABASES` 只能看到目标库,**代码就算写错也碰不到别的库**。 + +2. **启动自检(防配置写错)**。应用启动时会执行 `SELECT DATABASE()`, + 与 `MYSQL_DB_NAME` 不符就**拒绝启动**,而不是往错误的库里写。 + +**字符集**:这台服务的服务端和库默认都是 `latin1`。代码在建表时**逐表强制 +`utf8mb4`**,不依赖库默认值 —— 否则中文会被拒或变成问号。 + +**连接保活**:MySQL 默认 8 小时断开空闲连接,而监控服务是常驻的。 +已配置 `pool_recycle=3600` + `pool_pre_ping`,避免"server has gone away"。 + +**表引擎**:全部 InnoDB(`monitor_run.exit_code` 用 `BIGINT` —— Windows 的退出码是 +无符号 32 位,`0xC0000142` 会溢出有符号 `INT`)。 + +### 从 SQLite 迁移(如有旧数据) + +```bash +python -m api.monitor.migrate_from_sqlite --dry-run # 先看要迁什么 +python -m api.monitor.migrate_from_sqlite # 正式迁移 +``` + +保留原主键(否则 `task_id` 关联会错位);目标库非空时会拒绝执行,除非加 `--force`。 + +监控库中的 Cookie 为明文存储,这是当前版本的已知取舍。 + +--- + +## 五、API + +所有操作都有对应的 HTTP 接口,UI 只是其中一层封装: + +``` +GET /api/monitor/overview 看板汇总 +GET /api/monitor/tasks 任务列表 +POST /api/monitor/tasks 新建任务 +PATCH /api/monitor/tasks/{id} 修改 +DELETE /api/monitor/tasks/{id} 删除 +POST /api/monitor/tasks/{id}/run 立即运行(后台执行,立即返回) +GET /api/monitor/tasks/{id}/runs 运行历史 +GET /api/monitor/notes 作品表(含与上一轮的 Δ) +GET /api/monitor/notes/{id}/series 单篇指标时间序列 +GET /api/monitor/comments 评论流(带所属作品;?note_id= 筛选,?group_by=note 按作品分组) +GET /api/monitor/comment-notes 有评论的作品及其条数(评论筛选下拉用) +GET /api/monitor/export 导出(?kind=notes|comments|report&format=csv|xlsx) +GET /api/monitor/events 变化事件流 +POST /api/monitor/events/read 标记已读 +GET /api/monitor/cookie 登录态健康度(**不返回 Cookie 值**) +POST /api/monitor/cookie 保存 Cookie +DELETE /api/monitor/cookie 清除 Cookie + +GET /api/config/platforms 平台能力矩阵(含 monitor_wired,前端据此渲染切换器) +GET /api/settings 设置 + 表单描述(?platform=,敏感值只回状态) +PUT /api/settings 部分更新(只写请求里出现的 key) + +GET /api/auth/me 身份探测(401 即未登录) +POST /api/auth/login 登录(发 HttpOnly Cookie) +POST /api/auth/logout 退出 +POST /api/auth/password 修改密码(踢掉所有其他设备) + +GET /api/monitor/report 报表(?task_id=1&task_id=2&start_date=&end_date=) +GET /api/monitor/webhook 通知配置状态(**只返回打码地址**) +POST /api/monitor/webhook 保存 Webhook 地址 +DELETE /api/monitor/webhook 删除 Webhook +POST /api/monitor/webhook/test 发送测试消息 +``` + +> `task_id` 用**重复参数**而非逗号拼接(`?task_id=1&task_id=2`);不传表示统计全部任务。 + +--- + +## 六、故障排查 + +| 现象 | 原因 / 处理 | +|---|---| +| 任务一直不运行 | 未配置 Cookie(调度器会跳过并保持任务到期);或全局已有采集在跑 | +| 任务标红「疑似登录态失效」 | Cookie 过期。重新粘贴;若反复失败,先手动扫码登录一次播种浏览器 profile | +| 抓到的作品数长期为 0 | 同上;也可能是该博主确实没有作品 | +| 发现不了新评论 | 评论接口无时间排序所致,调大「每篇评论抓取条数」可缓解但无法根治 | +| 首轮没有任何"新增"事件 | 刻意设计:首轮建立基线,全部数据视为已有,不产生变化事件 | diff --git a/media_platform/xhs/core.py b/media_platform/xhs/core.py index 8110907..080daba 100644 --- a/media_platform/xhs/core.py +++ b/media_platform/xhs/core.py @@ -201,9 +201,22 @@ class XiaoHongShuCrawler(AbstractCrawler): # Parse creator URL to get user_id and security tokens creator_info: CreatorUrlInfo = parse_creator_info_from_url(creator_url) utils.logger.info(f"[XiaoHongShuCrawler.get_creators_and_notes] Parse creator URL info: {creator_info}") - user_id = creator_info.user_id + except ValueError as e: + utils.logger.error(f"[XiaoHongShuCrawler.get_creators_and_notes] Failed to parse creator URL: {e}") + continue - # get creator detail info from web html content + user_id = creator_info.user_id + + # Fetching the profile page is best-effort and must not abort the run. + # It only feeds save_creator(), which is a no-op in this build, while + # the notes themselves come from a completely different endpoint. + # Scraping the profile means parsing window.__INITIAL_STATE__ out of + # HTML, which fails whenever the platform serves a different page -- + # a JSONDecodeError there is especially misleading because it is a + # ValueError subclass, so it used to be reported as "failed to parse + # creator URL" and then skipped the creator entirely, yielding zero + # notes for a perfectly valid target. + try: createor_info: Dict = await self.xhs_client.get_creator_info( user_id=user_id, xsec_token=creator_info.xsec_token, @@ -211,9 +224,6 @@ class XiaoHongShuCrawler(AbstractCrawler): ) if createor_info: await xhs_store.save_creator(user_id, creator=createor_info) - except ValueError as e: - utils.logger.error(f"[XiaoHongShuCrawler.get_creators_and_notes] Failed to parse creator URL: {e}") - continue except (IPBlockError, PlatformAccessError) as e: # Access restricted on the creator homepage, skip this creator instead of crashing the run. utils.logger.error( @@ -221,6 +231,11 @@ class XiaoHongShuCrawler(AbstractCrawler): f"建议降低采集频率、更换 IP 或检查账号状态" ) continue + except Exception as e: + utils.logger.warning( + f"[XiaoHongShuCrawler.get_creators_and_notes] Could not fetch profile for {user_id} " + f"({type(e).__name__}: {e}); continuing to fetch the creator's notes anyway" + ) # Use fixed crawling interval crawl_interval = config.CRAWLER_MAX_SLEEP_SEC diff --git a/media_platform/xhs/login.py b/media_platform/xhs/login.py index e382954..ba6b2a7 100644 --- a/media_platform/xhs/login.py +++ b/media_platform/xhs/login.py @@ -213,8 +213,12 @@ class XiaoHongShuLogin(AbstractLogin): async def login_by_cookies(self): """login xiaohongshu website by cookies""" utils.logger.info("[XiaoHongShuLogin.login_by_cookies] Begin login xiaohongshu by cookie ...") + injected = 0 for key, value in utils.convert_str_cookie_to_dict(self.cookie_str).items(): - if key != "web_session": # Only set web_session cookie attribute + # Default (upstream) behaviour injects only web_session. Unattended runs + # need a1 / webId as well, otherwise signed API calls fail and the run + # exits 0 having fetched nothing -- a silent failure. + if not config.INJECT_ALL_COOKIES and key != "web_session": continue await self.browser_context.add_cookies([{ 'name': key, @@ -222,3 +226,8 @@ class XiaoHongShuLogin(AbstractLogin): 'domain': ".rednote.com" if config.XHS_INTERNATIONAL else ".xiaohongshu.com", 'path': "/" }]) + injected += 1 + utils.logger.info( + f"[XiaoHongShuLogin.login_by_cookies] Injected {injected} cookie(s), " + f"inject_all_cookies={config.INJECT_ALL_COOKIES}" + ) diff --git a/requirements.txt b/requirements.txt index 747cce7..4309a3b 100644 --- a/requirements.txt +++ b/requirements.txt @@ -28,4 +28,8 @@ motor>=3.3.0 openpyxl>=3.1.2 pytest>=7.4.0 pytest-asyncio>=0.21.0 -xhshow>=0.2.0 \ No newline at end of file +xhshow>=0.2.0 +# Required by uvicorn to handle WebSocket upgrades. Declared in pyproject.toml +# but previously missing here, so installing from this file left the live log +# stream silently non-functional (uvicorn answers every upgrade with 404). +websockets>=15.0.1 \ No newline at end of file diff --git a/tests/conftest.py b/tests/conftest.py index dc593a4..2c21f5e 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -96,3 +96,32 @@ def sample_xhs_creator(): "interaction": 50000, "tag_list": '{"profession": "Designer", "interest": "Photography"}' } + + +@pytest.fixture(autouse=True) +def _bypass_auth_for_non_auth_suites(request): + """Skip API authentication for suites that are not about authentication. + + Adding auth to every /api route breaks any test that speaks HTTP, so those + suites override the dependency here. This uses FastAPI's own + ``dependency_overrides`` mechanism rather than a production-visible + "test mode" switch, which could be shipped enabled by accident. + + ``tests/test_auth.py`` is deliberately excluded: it must exercise the real + enforcement path, including the route-enumeration guard that asserts every + other /api route really does return 401. + """ + if request.node.fspath.basename == "test_auth.py": + yield + return + + from api.auth import require_auth, require_ws_auth + from api.main import app + + app.dependency_overrides[require_auth] = lambda: None + app.dependency_overrides[require_ws_auth] = lambda: None + try: + yield + finally: + app.dependency_overrides.pop(require_auth, None) + app.dependency_overrides.pop(require_ws_auth, None) diff --git a/tests/test_auth.py b/tests/test_auth.py new file mode 100644 index 0000000..f980875 --- /dev/null +++ b/tests/test_auth.py @@ -0,0 +1,503 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/tests/test_auth.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Tests for WebUI authentication. + +Deliberately does NOT install ``app.dependency_overrides``: the point of this +file is to exercise the real enforcement path. Other suites override +``require_auth`` so they can keep testing their own concerns. +""" + +import asyncio +import time + +import httpx +import pytest +import pytest_asyncio +from fastapi import WebSocketException +from sqlalchemy import func, select + +from api import auth +from api.main import app +from api.monitor import db as monitor_db +from api.monitor.models import AuthSession + +PASSWORD = "correct-horse-battery" + +# Captured at import, i.e. before the autouse fixture patches the module global, +# so the guard test below checks the value that actually ships. +REAL_PBKDF2_ITERATIONS = auth.PBKDF2_ITERATIONS + +# Every /api route that is allowed to answer without a session. +EXEMPT_PATHS = {"/api/health", "/api/auth/login", "/api/auth/logout"} + + +@pytest.fixture(autouse=True) +def cheap_hashing(monkeypatch): + """600k iterations is right in production and unusable in a test suite. + + hash_password() resolves the count at call time precisely so this works. + """ + monkeypatch.setattr(auth, "PBKDF2_ITERATIONS", 1_000) + monkeypatch.delenv("MC_PASSWORD", raising=False) + auth.reset_throttle_state() + yield + auth.reset_throttle_state() + + +@pytest_asyncio.fixture +async def db(tmp_path): + monitor_db.set_sqlite_path(tmp_path / "monitor.db") + await monitor_db.init_db() + yield monitor_db + await monitor_db.dispose_engine() + + +@pytest_asyncio.fixture +async def client(db): + transport = httpx.ASGITransport(app=app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as http_client: + yield http_client + + +async def _seed_password(password: str = PASSWORD) -> None: + async with monitor_db.get_session() as session: + await auth.set_password(session, password) + + +# -------------------------------------------------------------------------- +# Password hashing +# -------------------------------------------------------------------------- + +class TestPasswordHashing: + def test_iteration_count_has_not_been_lowered(self): + """Guard: someone trimming this for speed would weaken every install.""" + assert REAL_PBKDF2_ITERATIONS >= 600_000 + + def test_round_trip(self): + stored = auth.hash_password(PASSWORD) + assert auth._verify_password_sync(PASSWORD, stored) is True + + def test_wrong_password_rejected(self): + stored = auth.hash_password(PASSWORD) + assert auth._verify_password_sync("wrong", stored) is False + + def test_same_password_hashes_differently(self): + """A fixed salt would let one rainbow table crack every install.""" + assert auth.hash_password(PASSWORD) != auth.hash_password(PASSWORD) + + def test_format_is_self_describing(self): + algo, iterations, salt, digest = auth.hash_password(PASSWORD).split("$") + assert algo == "pbkdf2_sha256" + assert int(iterations) == auth.PBKDF2_ITERATIONS + assert salt and digest + + @pytest.mark.parametrize("stored", ["", "garbage", "md5$1$a$b", "pbkdf2_sha256$x$a$b"]) + def test_malformed_stored_hash_is_rejected_not_raised(self, stored): + assert auth._verify_password_sync(PASSWORD, stored) is False + + +# -------------------------------------------------------------------------- +# Credentials +# -------------------------------------------------------------------------- + +class TestCredentials: + @pytest.mark.asyncio + async def test_check_password_against_stored_hash(self, db): + await _seed_password() + async with monitor_db.get_session() as session: + assert await auth.check_password(session, PASSWORD) is True + assert await auth.check_password(session, "nope") is False + + @pytest.mark.asyncio + async def test_no_password_configured_denies_everything(self, db): + """An unset credential must not mean "open".""" + async with monitor_db.get_session() as session: + assert await auth.check_password(session, "") is False + assert await auth.check_password(session, PASSWORD) is False + + @pytest.mark.asyncio + async def test_env_override_wins_and_is_not_persisted(self, db, monkeypatch): + """The documented way back in after forgetting the password.""" + await _seed_password("stored-password") + monkeypatch.setenv("MC_PASSWORD", "env-password") + + async with monitor_db.get_session() as session: + assert await auth.check_password(session, "env-password") is True + assert await auth.check_password(session, "stored-password") is False + + # Override must never be written to disk. + assert await auth.current_password_hash(session) != "" + assert "env-password" not in (await auth.current_password_hash(session)) + + @pytest.mark.asyncio + async def test_first_run_generates_a_credential(self, db, monkeypatch): + monkeypatch.delenv("MC_PASSWORD", raising=False) + + generated = await auth.ensure_initial_credential() + assert generated + + # Second call is a no-op. + assert await auth.ensure_initial_credential() is None + + async with monitor_db.get_session() as session: + assert await auth.check_password(session, generated) is True + + @pytest.mark.asyncio + async def test_first_run_defers_to_env_password(self, db, monkeypatch): + monkeypatch.setenv("MC_PASSWORD", "env-password") + + assert await auth.ensure_initial_credential() is None + + async with monitor_db.get_session() as session: + assert await auth.current_password_hash(session) == "" + + +# -------------------------------------------------------------------------- +# Sessions +# -------------------------------------------------------------------------- + +class TestSessions: + @pytest.mark.asyncio + async def test_round_trip(self, db): + async with monitor_db.get_session() as session: + token, expires_at = await auth.create_session(session) + + async with monitor_db.get_session() as session: + assert await auth.resolve_session(session, token) is not None + assert expires_at > 0 + + @pytest.mark.asyncio + async def test_only_the_hash_is_stored(self, db): + """A database leak must not hand over live sessions.""" + async with monitor_db.get_session() as session: + token, _ = await auth.create_session(session) + + async with monitor_db.get_session() as session: + stored = (await session.scalars(select(AuthSession.token_hash))).all() + assert token not in stored + assert auth._hash_token(token) in stored + + @pytest.mark.asyncio + async def test_expired_session_is_rejected_and_removed(self, db): + async with monitor_db.get_session() as session: + token, _ = await auth.create_session(session) + row = await session.get(AuthSession, auth._hash_token(token)) + row.expires_at = 1 # long past + + async with monitor_db.get_session() as session: + assert await auth.resolve_session(session, token) is None + + # Fresh session: the identity map in the one above still holds the + # pending-delete object, so it would answer as if the row were present. + async with monitor_db.get_session() as session: + assert await session.get(AuthSession, auth._hash_token(token)) is None + + @pytest.mark.asyncio + async def test_unknown_token_is_rejected(self, db): + async with monitor_db.get_session() as session: + assert await auth.resolve_session(session, "never-issued") is None + assert await auth.resolve_session(session, "") is None + + @pytest.mark.asyncio + async def test_logout_revokes_only_that_session(self, db): + async with monitor_db.get_session() as session: + first, _ = await auth.create_session(session) + second, _ = await auth.create_session(session) + + async with monitor_db.get_session() as session: + await auth.revoke_session(session, first) + + async with monitor_db.get_session() as session: + assert await auth.resolve_session(session, first) is None + assert await auth.resolve_session(session, second) is not None + + @pytest.mark.asyncio + async def test_revoke_all_clears_every_session(self, db): + async with monitor_db.get_session() as session: + await auth.create_session(session) + await auth.create_session(session) + + async with monitor_db.get_session() as session: + removed = await auth.revoke_all_sessions(session) + assert removed == 2 + + async with monitor_db.get_session() as session: + assert await session.scalar(select(func.count()).select_from(AuthSession)) == 0 + + +# -------------------------------------------------------------------------- +# Throttle +# -------------------------------------------------------------------------- + +class TestThrottle: + @pytest.mark.asyncio + async def test_below_threshold_is_not_throttled(self, db): + key = "1.2.3.4" + for _ in range(auth.THROTTLE_THRESHOLD - 1): + await auth.record_failure(key) + assert await auth.retry_after_seconds(key) == 0 + + @pytest.mark.asyncio + async def test_lockout_after_repeated_failures(self, db): + key = "1.2.3.4" + for _ in range(auth.THROTTLE_THRESHOLD): + await auth.record_failure(key) + assert await auth.retry_after_seconds(key) > 0 + + @pytest.mark.asyncio + async def test_success_clears_failures(self, db): + key = "1.2.3.4" + for _ in range(auth.THROTTLE_THRESHOLD): + await auth.record_failure(key) + await auth.clear_failures(key) + assert await auth.retry_after_seconds(key) == 0 + + @pytest.mark.asyncio + async def test_failures_age_out_of_the_window(self, db, monkeypatch): + """Driven by a fake clock rather than sleeping 15 minutes.""" + key = "1.2.3.4" + clock = {"now": 1000.0} + monkeypatch.setattr(auth, "_now", lambda: clock["now"]) + + for _ in range(auth.THROTTLE_THRESHOLD): + await auth.record_failure(key) + assert await auth.retry_after_seconds(key) > 0 + + clock["now"] += auth.THROTTLE_WINDOW_SECONDS + 1 + assert await auth.retry_after_seconds(key) == 0 + + @pytest.mark.asyncio + async def test_keys_are_independent(self, db): + for _ in range(auth.THROTTLE_THRESHOLD): + await auth.record_failure("attacker") + assert await auth.retry_after_seconds("attacker") > 0 + assert await auth.retry_after_seconds("innocent") == 0 + + +# -------------------------------------------------------------------------- +# HTTP enforcement — the acceptance criteria +# -------------------------------------------------------------------------- + +class TestEnforcement: + @pytest.mark.asyncio + async def test_health_is_reachable_without_a_session(self, client): + assert (await client.get("/api/health")).status_code == 200 + + @pytest.mark.asyncio + async def test_protected_endpoint_returns_401_without_a_session(self, client): + response = await client.get("/api/monitor/tasks") + assert response.status_code == 401 + + @pytest.mark.asyncio + async def test_wrong_password_is_401_and_generic(self, client): + await _seed_password() + response = await client.post("/api/auth/login", json={"password": "wrong"}) + assert response.status_code == 401 + # Must not reveal whether a password is even configured. + assert response.json()["detail"] == auth.INVALID_CREDENTIALS + + @pytest.mark.asyncio + async def test_login_unlocks_the_api(self, client): + await _seed_password() + + login = await client.post("/api/auth/login", json={"password": PASSWORD}) + assert login.status_code == 200 + assert auth.SESSION_COOKIE_NAME in client.cookies + + assert (await client.get("/api/monitor/tasks")).status_code == 200 + + @pytest.mark.asyncio + async def test_cookie_is_httponly_and_lax(self, client): + await _seed_password() + login = await client.post("/api/auth/login", json={"password": PASSWORD}) + raw = login.headers["set-cookie"].lower() + assert "httponly" in raw + assert "samesite=lax" in raw + # Secure must be OFF by default: the LAN bind is plain HTTP and a Secure + # cookie is silently dropped there, looping the login page. + assert "secure" not in raw + + @pytest.mark.asyncio + async def test_bearer_token_also_works(self, client): + """Scripts and curl cannot use a cookie jar conveniently.""" + await _seed_password() + login = await client.post("/api/auth/login", json={"password": PASSWORD}) + token = login.cookies[auth.SESSION_COOKIE_NAME] + + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), base_url="http://test" + ) as bare: + bare.headers["Authorization"] = f"Bearer {token}" + assert (await bare.get("/api/monitor/tasks")).status_code == 200 + + @pytest.mark.asyncio + async def test_tampered_token_is_rejected(self, client): + await _seed_password() + await client.post("/api/auth/login", json={"password": PASSWORD}) + client.cookies.set(auth.SESSION_COOKIE_NAME, "not-a-real-token") + assert (await client.get("/api/monitor/tasks")).status_code == 401 + + @pytest.mark.asyncio + async def test_logout_invalidates_the_session(self, client): + await _seed_password() + await client.post("/api/auth/login", json={"password": PASSWORD}) + assert (await client.get("/api/monitor/tasks")).status_code == 200 + + assert (await client.post("/api/auth/logout")).status_code == 200 + assert (await client.get("/api/monitor/tasks")).status_code == 401 + + @pytest.mark.asyncio + async def test_me_reports_401_when_logged_out(self, client): + await _seed_password() + assert (await client.get("/api/auth/me")).status_code == 401 + + await client.post("/api/auth/login", json={"password": PASSWORD}) + me = await client.get("/api/auth/me") + assert me.status_code == 200 + assert me.json()["authenticated"] is True + + @pytest.mark.asyncio + async def test_password_change_evicts_other_devices(self, client): + await _seed_password() + + # A second "device" holds its own session. + login = await client.post("/api/auth/login", json={"password": PASSWORD}) + other_token = login.cookies[auth.SESSION_COOKIE_NAME] + + changed = await client.post( + "/api/auth/password", + json={"current": PASSWORD, "new": "brand-new-password"}, + ) + assert changed.status_code == 200 + + # The old token is dead. + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), base_url="http://test" + ) as other: + other.cookies.set(auth.SESSION_COOKIE_NAME, other_token) + assert (await other.get("/api/monitor/tasks")).status_code == 401 + + # ...and the caller is still logged in. + assert (await client.get("/api/monitor/tasks")).status_code == 200 + + @pytest.mark.asyncio + async def test_password_change_requires_the_current_password(self, client): + await _seed_password() + await client.post("/api/auth/login", json={"password": PASSWORD}) + + response = await client.post( + "/api/auth/password", json={"current": "wrong", "new": "whatever-new"} + ) + assert response.status_code == 401 + + @pytest.mark.asyncio + async def test_repeated_failures_get_throttled(self, client): + await _seed_password() + for _ in range(auth.THROTTLE_THRESHOLD): + await client.post("/api/auth/login", json={"password": "wrong"}) + + blocked = await client.post("/api/auth/login", json={"password": PASSWORD}) + assert blocked.status_code == 429 + assert "retry-after" in {k.lower() for k in blocked.headers} + + @pytest.mark.asyncio + async def test_docs_are_not_exposed(self, client): + for path in ("/docs", "/redoc", "/openapi.json"): + assert (await client.get(path)).status_code == 404 + + +class TestEveryRouteIsGuarded: + @pytest.mark.asyncio + async def test_no_api_route_is_accidentally_open(self, client): + """The guard that stops the next endpoint from shipping unauthenticated.""" + unguarded = [] + + for route in app.routes: + path = getattr(route, "path", "") + methods = getattr(route, "methods", None) + if not path.startswith("/api") or not methods or path in EXEMPT_PATHS: + continue + + # Substitute dummy values for path params so we reach the auth check + # rather than a 404/422 on the parameter itself. + concrete = "/".join( + "1" if segment.startswith("{") else segment for segment in path.split("/") + ) + + for method in methods - {"HEAD", "OPTIONS"}: + response = await client.request(method, concrete, json={}) + if response.status_code != 401: + unguarded.append(f"{method} {path} -> {response.status_code}") + + assert not unguarded, f"以下 /api 路由未受鉴权保护:{unguarded}" + + +# -------------------------------------------------------------------------- +# WebSocket enforcement +# -------------------------------------------------------------------------- + +class _FakeWebSocket: + """Only `.cookies` is read by require_ws_auth.""" + + def __init__(self, cookies): + self.cookies = cookies + + +class TestWebSocketAuth: + """Guarding websockets needs its own mechanism: BaseHTTPMiddleware returns + early for non-http scopes, and HTTP router dependencies never run for them. + Without this the live crawl log stream would be wide open. + """ + + @pytest.mark.asyncio + async def test_missing_cookie_is_rejected(self, db): + with pytest.raises(WebSocketException) as excinfo: + await auth.require_ws_auth(_FakeWebSocket({})) + assert excinfo.value.code == 1008 + + @pytest.mark.asyncio + async def test_valid_cookie_is_accepted(self, db): + async with monitor_db.get_session() as session: + token, _ = await auth.create_session(session) + + # No exception means accepted. + await auth.require_ws_auth(_FakeWebSocket({auth.SESSION_COOKIE_NAME: token})) + + @pytest.mark.asyncio + async def test_unknown_cookie_is_rejected(self, db): + with pytest.raises(WebSocketException): + await auth.require_ws_auth(_FakeWebSocket({auth.SESSION_COOKIE_NAME: "bogus"})) + + def test_every_websocket_route_carries_the_guard(self): + guarded = { + route.path + for route in app.routes + if route.__class__.__name__ == "APIWebSocketRoute" + and any( + getattr(dep.dependency, "__name__", "") == "require_ws_auth" + for dep in (route.dependencies or []) + ) + } + every_ws = { + route.path + for route in app.routes + if route.__class__.__name__ == "APIWebSocketRoute" + } + assert every_ws, "expected at least one websocket route" + assert every_ws == guarded diff --git a/tests/test_cmd_arg_monitor_flags.py b/tests/test_cmd_arg_monitor_flags.py new file mode 100644 index 0000000..27ccc9d --- /dev/null +++ b/tests/test_cmd_arg_monitor_flags.py @@ -0,0 +1,128 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/tests/test_cmd_arg_monitor_flags.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Tests for the CLI flags added to support unattended monitoring runs. + +Each flag must default to the existing config value, so a manual crawl that does +not pass them behaves exactly as before. +""" + +import pytest + +import config +from cmd_arg.arg import parse_cmd + +BASE_ARGS = ["--platform", "xhs", "--type", "creator", "--creator_id", "abc123"] + + +@pytest.fixture(autouse=True) +def _isolate_config(monkeypatch): + monkeypatch.setattr(config, "ENABLE_CDP_MODE", True) + monkeypatch.setattr(config, "INJECT_ALL_COOKIES", False) + monkeypatch.setattr(config, "SAVE_LOGIN_STATE", True) + monkeypatch.setattr(config, "COOKIES", "") + monkeypatch.setattr(config, "SAVE_DATA_PATH", "") + monkeypatch.setattr(config, "CRAWLER_MAX_SLEEP_SEC", 2) + yield + + +class TestEnableCdpMode: + """CDP attaches to the user's desktop Chrome, which cannot work on a server.""" + + @pytest.mark.asyncio + async def test_false_disables_cdp(self): + await parse_cmd([*BASE_ARGS, "--enable_cdp_mode", "false"]) + assert config.ENABLE_CDP_MODE is False + + @pytest.mark.asyncio + async def test_defaults_to_config_value(self): + await parse_cmd(BASE_ARGS) + assert config.ENABLE_CDP_MODE is True + + +class TestCookieFlags: + @pytest.mark.asyncio + async def test_inject_all_cookies_enables_switch(self): + await parse_cmd([*BASE_ARGS, "--inject_all_cookies", "true"]) + assert config.INJECT_ALL_COOKIES is True + + @pytest.mark.asyncio + async def test_inject_all_cookies_defaults_off(self): + await parse_cmd(BASE_ARGS) + assert config.INJECT_ALL_COOKIES is False + + @pytest.mark.asyncio + async def test_cookies_file_is_read_into_config(self, tmp_path): + cookie_file = tmp_path / "cookies.txt" + cookie_file.write_text("web_session=abc; a1=def", encoding="utf-8") + + await parse_cmd([*BASE_ARGS, "--cookies_file", str(cookie_file)]) + + assert config.COOKIES == "web_session=abc; a1=def" + + @pytest.mark.asyncio + async def test_cookies_file_wins_over_inline_cookies(self, tmp_path): + cookie_file = tmp_path / "cookies.txt" + cookie_file.write_text("web_session=fromfile", encoding="utf-8") + + await parse_cmd( + [*BASE_ARGS, "--cookies", "web_session=inline", "--cookies_file", str(cookie_file)] + ) + + assert config.COOKIES == "web_session=fromfile" + + @pytest.mark.asyncio + async def test_missing_cookies_file_is_rejected(self, tmp_path): + missing = tmp_path / "nope.txt" + + with pytest.raises(Exception) as excinfo: + await parse_cmd([*BASE_ARGS, "--cookies_file", str(missing)]) + + # A silently-ignored unreadable cookie file would produce a crawl that + # returns nothing, which is exactly the failure mode this flag exists + # to avoid. + assert "cookies_file" in str(excinfo.value) + + +class TestSaveDataPath: + @pytest.mark.asyncio + async def test_save_data_path_is_applied(self): + await parse_cmd([*BASE_ARGS, "--save_data_path", "data/monitor_runs/1/2"]) + assert config.SAVE_DATA_PATH == "data/monitor_runs/1/2" + + +class TestSaveLoginState: + @pytest.mark.asyncio + async def test_save_login_state_can_be_disabled(self): + await parse_cmd([*BASE_ARGS, "--save_login_state", "false"]) + assert config.SAVE_LOGIN_STATE is False + + +class TestCrawlSleepSec: + """Exposed on the Settings page; previously had no CLI flag at all.""" + + @pytest.mark.asyncio + async def test_value_is_applied(self): + await parse_cmd([*BASE_ARGS, "--crawler_max_sleep_sec", "9"]) + assert config.CRAWLER_MAX_SLEEP_SEC == 9 + + @pytest.mark.asyncio + async def test_defaults_to_config_value(self, monkeypatch): + monkeypatch.setattr(config, "CRAWLER_MAX_SLEEP_SEC", 4) + await parse_cmd(BASE_ARGS) + assert config.CRAWLER_MAX_SLEEP_SEC == 4 diff --git a/tests/test_interpreter.py b/tests/test_interpreter.py new file mode 100644 index 0000000..b335235 --- /dev/null +++ b/tests/test_interpreter.py @@ -0,0 +1,67 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/tests/test_interpreter.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Tests for the subprocess interpreter resolver.""" + +import sys +from pathlib import Path + +from api.services.interpreter import ( + describe_interpreter, + resolve_python_cmd, + venv_python_path, +) + + +def _make_venv(root: Path) -> Path: + """Create a fake venv layout and return the expected python path.""" + exe = venv_python_path(root) + exe.parent.mkdir(parents=True, exist_ok=True) + exe.write_text("", encoding="utf-8") + return exe + + +def test_prefers_uv_when_available(monkeypatch, tmp_path): + monkeypatch.setattr("shutil.which", lambda name: "/usr/bin/uv" if name == "uv" else None) + _make_venv(tmp_path) + + # uv wins even when a venv exists, matching the upstream documented workflow. + assert resolve_python_cmd(tmp_path) == ["uv", "run", "python"] + + +def test_falls_back_to_project_venv(monkeypatch, tmp_path): + monkeypatch.setattr("shutil.which", lambda name: None) + exe = _make_venv(tmp_path) + + assert resolve_python_cmd(tmp_path) == [str(exe)] + + +def test_falls_back_to_current_interpreter(monkeypatch, tmp_path): + monkeypatch.setattr("shutil.which", lambda name: None) + + # No uv, no venv anywhere under the given root. + assert resolve_python_cmd(tmp_path) == [sys.executable] + + +def test_describe_is_human_readable(monkeypatch, tmp_path): + monkeypatch.setattr("shutil.which", lambda name: None) + _make_venv(tmp_path) + assert "virtualenv" in describe_interpreter(tmp_path) + + monkeypatch.setattr("shutil.which", lambda name: "/usr/bin/uv" if name == "uv" else None) + assert describe_interpreter(tmp_path) == "uv run python" diff --git a/tests/test_monitor_api.py b/tests/test_monitor_api.py new file mode 100644 index 0000000..0dee8f3 --- /dev/null +++ b/tests/test_monitor_api.py @@ -0,0 +1,208 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/tests/test_monitor_api.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""API-level tests for the monitoring endpoints. + +Run against an ASGI transport with a temporary database, so no server, network +or login is required. Lifespan is deliberately not exercised: it would start the +scheduler, and these tests only cover routing, validation and persistence. +""" + +import httpx +import pytest +import pytest_asyncio + +from api.main import app +from api.monitor import db as monitor_db +from api.monitor.service import TargetParseError, parse_target_input + +CREATOR_URL = ( + "https://www.xiaohongshu.com/user/profile/5f58bd990000000001003753" + "?xsec_token=ABYVg1evluJZZzpMX-VWzchxQ1qSNVW3r-jOEnKqMcgZw=&xsec_source=pc_search" +) +NOTE_URL = "https://www.xiaohongshu.com/explore/6aa3d827000000002802c5c8?xsec_token=TOKEN&xsec_source=pc_search" + + +@pytest_asyncio.fixture +async def client(tmp_path): + monitor_db.set_sqlite_path(tmp_path / "monitor.db") + await monitor_db.init_db() + + transport = httpx.ASGITransport(app=app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as http_client: + yield http_client + + await monitor_db.dispose_engine() + + +class TestParseTargetInput: + def test_full_url_splits_id_from_token(self): + """The id is the stable key; the token is a refreshable credential.""" + parsed = parse_target_input(CREATOR_URL, "creator") + assert parsed["external_id"] == "5f58bd990000000001003753" + assert parsed["xsec_token"].startswith("ABYVg1evluJZZzpMX") + assert parsed["xsec_source"] == "pc_search" + + def test_bare_id_is_accepted(self): + parsed = parse_target_input("5f58bd990000000001003753", "creator") + assert parsed["external_id"] == "5f58bd990000000001003753" + assert parsed["xsec_token"] == "" + + def test_note_url_without_token_still_parses(self): + parsed = parse_target_input( + "https://www.xiaohongshu.com/explore/6aa3d827000000002802c5c8", "note" + ) + assert parsed["external_id"] == "6aa3d827000000002802c5c8" + assert parsed["xsec_token"] == "" + + def test_creator_url_rejected_in_note_mode(self): + with pytest.raises(TargetParseError): + parse_target_input(CREATOR_URL, "note") + + def test_garbage_is_rejected(self): + with pytest.raises(TargetParseError): + parse_target_input("not a url at all !!", "creator") + + +class TestTaskCrud: + @pytest.mark.asyncio + async def test_create_and_list_task(self, client): + response = await client.post( + "/api/monitor/tasks", + json={ + "name": "网文作者监控", + "mode": "creator", + "interval_minutes": 120, + "targets": [CREATOR_URL, "5f58bd990000000001003754"], + }, + ) + assert response.status_code == 201 + task_id = response.json()["id"] + + listing = await client.get("/api/monitor/tasks") + assert listing.status_code == 200 + tasks = listing.json()["tasks"] + assert len(tasks) == 1 + assert tasks[0]["id"] == task_id + assert tasks[0]["target_count"] == 2 + # next_run_at is persisted so the schedule survives a restart. + assert tasks[0]["next_run_at"] is not None + + @pytest.mark.asyncio + async def test_duplicate_targets_are_deduplicated(self, client): + response = await client.post( + "/api/monitor/tasks", + json={ + "name": "dedup", + "mode": "creator", + "targets": [CREATOR_URL, CREATOR_URL], + }, + ) + assert response.status_code == 201 + + listing = await client.get("/api/monitor/tasks") + assert listing.json()["tasks"][0]["target_count"] == 1 + + @pytest.mark.asyncio + async def test_invalid_target_returns_400(self, client): + response = await client.post( + "/api/monitor/tasks", + json={"name": "bad", "mode": "creator", "targets": ["!!! nonsense !!!"]}, + ) + assert response.status_code == 400 + + @pytest.mark.asyncio + async def test_interval_floor_is_enforced(self, client): + """A tight poll loop is the pattern that triggers platform rate limits.""" + response = await client.post( + "/api/monitor/tasks", + json={"name": "too fast", "mode": "creator", "interval_minutes": 1, "targets": [CREATOR_URL]}, + ) + assert response.status_code == 422 + + @pytest.mark.asyncio + async def test_update_and_delete(self, client): + created = await client.post( + "/api/monitor/tasks", + json={"name": "t", "mode": "note", "targets": [NOTE_URL]}, + ) + task_id = created.json()["id"] + + patched = await client.patch(f"/api/monitor/tasks/{task_id}", json={"enabled": False}) + assert patched.status_code == 200 + listing = await client.get("/api/monitor/tasks") + assert listing.json()["tasks"][0]["enabled"] is False + + deleted = await client.delete(f"/api/monitor/tasks/{task_id}") + assert deleted.status_code == 200 + assert (await client.get("/api/monitor/tasks")).json()["tasks"] == [] + + @pytest.mark.asyncio + async def test_run_now_on_missing_task_is_404(self, client): + response = await client.post("/api/monitor/tasks/9999/run") + assert response.status_code == 404 + + @pytest.mark.asyncio + async def test_run_history_starts_empty(self, client): + created = await client.post( + "/api/monitor/tasks", + json={"name": "t", "mode": "creator", "targets": [CREATOR_URL]}, + ) + task_id = created.json()["id"] + runs = await client.get(f"/api/monitor/tasks/{task_id}/runs") + assert runs.status_code == 200 + assert runs.json()["runs"] == [] + + +class TestCookieEndpoints: + @pytest.mark.asyncio + async def test_cookie_value_is_never_returned(self, client): + """The GET must expose health only, never the credential.""" + secret = "web_session=SUPERSECRETVALUE; a1=abc123" + saved = await client.post("/api/monitor/cookie", json={"cookie": secret}) + assert saved.status_code == 200 + + status_response = await client.get("/api/monitor/cookie") + assert status_response.status_code == 200 + body = status_response.json() + + assert body["present"] is True + assert body["length"] == len(secret) + assert "SUPERSECRETVALUE" not in status_response.text + + @pytest.mark.asyncio + async def test_cookie_initially_absent_and_clearable(self, client): + assert (await client.get("/api/monitor/cookie")).json()["present"] is False + + await client.post("/api/monitor/cookie", json={"cookie": "web_session=x"}) + assert (await client.get("/api/monitor/cookie")).json()["present"] is True + + await client.delete("/api/monitor/cookie") + assert (await client.get("/api/monitor/cookie")).json()["present"] is False + + +class TestDashboardQueries: + @pytest.mark.asyncio + async def test_empty_dashboard_shapes(self, client): + assert (await client.get("/api/monitor/notes")).json()["notes"] == [] + assert (await client.get("/api/monitor/comments")).json()["comments"] == [] + assert (await client.get("/api/monitor/events")).json()["events"] == [] + + overview = (await client.get("/api/monitor/overview")).json() + assert overview["tasks"] == 0 + assert overview["notes"] == 0 diff --git a/tests/test_monitor_comments.py b/tests/test_monitor_comments.py new file mode 100644 index 0000000..6ca7152 --- /dev/null +++ b/tests/test_monitor_comments.py @@ -0,0 +1,232 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/tests/test_monitor_comments.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Comment note-association, grouping, and the export endpoint.""" + +import csv +import io + +import httpx +import pytest +import pytest_asyncio + +from api.main import app +from api.monitor import db as monitor_db +from api.monitor.models import ( + MODE_CREATOR, + MonitorComment, + MonitorNote, + MonitorTask, +) + +TASK_NAME = "评论归属测试" + + +async def _seed(): + """Two works; three comments on the first, one on the second.""" + async with monitor_db.get_session() as session: + task = MonitorTask( + name=TASK_NAME, platform="xhs", mode=MODE_CREATOR, enabled=True, + interval_minutes=60, max_notes_count=20, enable_comments=True, + max_comments_count=50, run_timeout_seconds=3600, + notify_enabled=False, created_at=0, updated_at=0, + ) + session.add(task) + await session.flush() + + for note_id, title in (("note-a", "作品甲"), ("note-b", "作品乙")): + session.add( + MonitorNote( + task_id=task.id, note_id=note_id, title=title, + note_url=f"https://www.xiaohongshu.com/explore/{note_id}", + cover=f"https://img/{note_id}.jpg", creator_hash="h", + source_kind="video", published_at=None, + first_seen_run_id=1, first_seen_at=1_700_000_000_000, + last_seen_run_id=1, last_seen_at=1_700_000_000_000, + ) + ) + + # note-a has three comments, note-b has one. + plan = [ + ("c1", "note-a", 1_700_000_001_000), + ("c2", "note-a", 1_700_000_002_000), + ("c3", "note-a", 1_700_000_003_000), + ("c4", "note-b", 1_700_000_004_000), + ] + for comment_id, note_id, seen_at in plan: + session.add( + MonitorComment( + task_id=task.id, note_id=note_id, comment_id=comment_id, + content=f"内容-{comment_id}", nickname="u***r", creator_hash="h", + create_time=seen_at, like_count=1, sub_comment_count=0, + parent_comment_id="", first_seen_run_id=1, first_seen_at=seen_at, + ) + ) + return task.id + + +@pytest_asyncio.fixture +async def client(tmp_path): + monitor_db.set_sqlite_path(tmp_path / "monitor.db") + await monitor_db.init_db() + await _seed() + + transport = httpx.ASGITransport(app=app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as http_client: + yield http_client + + await monitor_db.dispose_engine() + + +class TestCommentsCarryTheirNote: + @pytest.mark.asyncio + async def test_each_comment_names_its_work(self, client): + """A bare note_id is unreadable -- the title is the whole point.""" + response = await client.get("/api/monitor/comments") + assert response.status_code == 200 + + comments = response.json()["comments"] + assert len(comments) == 4 + + by_id = {c["comment_id"]: c for c in comments} + assert by_id["c1"]["note_title"] == "作品甲" + assert by_id["c1"]["note_url"].endswith("note-a") + assert by_id["c1"]["note_cover"].endswith("note-a.jpg") + assert by_id["c4"]["note_title"] == "作品乙" + + @pytest.mark.asyncio + async def test_note_id_filters_the_stream(self, client): + response = await client.get("/api/monitor/comments", params={"note_id": "note-a"}) + comments = response.json()["comments"] + assert {c["comment_id"] for c in comments} == {"c1", "c2", "c3"} + + +class TestGroupByNote: + @pytest.mark.asyncio + async def test_groups_bucket_by_work(self, client): + response = await client.get("/api/monitor/comments", params={"group_by": "note"}) + body = response.json() + + assert "groups" in body + assert body["total"] == 4 + + groups = {g["note_id"]: g for g in body["groups"]} + assert set(groups) == {"note-a", "note-b"} + assert len(groups["note-a"]["comments"]) == 3 + assert len(groups["note-b"]["comments"]) == 1 + assert groups["note-a"]["note_title"] == "作品甲" + + @pytest.mark.asyncio + async def test_newest_group_comes_first(self, client): + """The UI expands the first group by default, so it must be the newest.""" + response = await client.get("/api/monitor/comments", params={"group_by": "note"}) + groups = response.json()["groups"] + # note-b's only comment is the most recent overall. + assert groups[0]["note_id"] == "note-b" + + @pytest.mark.asyncio + async def test_flat_shape_is_unchanged_without_the_flag(self, client): + body = (await client.get("/api/monitor/comments")).json() + assert "comments" in body and "groups" not in body + + +class TestCommentNoteFilterOptions: + @pytest.mark.asyncio + async def test_options_carry_counts_and_titles(self, client): + response = await client.get("/api/monitor/comment-notes") + assert response.status_code == 200 + + notes = {n["note_id"]: n for n in response.json()["notes"]} + assert notes["note-a"]["comment_count"] == 3 + assert notes["note-b"]["comment_count"] == 1 + assert notes["note-a"]["note_title"] == "作品甲" + + @pytest.mark.asyncio + async def test_scoped_to_a_task(self, client): + tasks = (await client.get("/api/monitor/tasks")).json()["tasks"] + task_id = tasks[0]["id"] + + scoped = await client.get("/api/monitor/comment-notes", params={"task_id": task_id}) + assert len(scoped.json()["notes"]) == 2 + + # A task with no comments yields an empty list, not an error. + other = await client.get("/api/monitor/comment-notes", params={"task_id": 9999}) + assert other.json()["notes"] == [] + + +class TestExport: + @pytest.mark.asyncio + async def test_csv_has_a_bom_so_excel_does_not_mangle_chinese(self, client): + response = await client.get("/api/monitor/export", params={"kind": "comments"}) + assert response.status_code == 200 + assert response.content.startswith(b"\xef\xbb\xbf") + assert "attachment" in response.headers["content-disposition"] + + text = response.content.decode("utf-8-sig") + rows = list(csv.DictReader(io.StringIO(text))) + assert len(rows) == 4 + assert rows[0]["所属作品"] in ("作品甲", "作品乙") + + @pytest.mark.asyncio + async def test_notes_export(self, client): + response = await client.get( + "/api/monitor/export", params={"kind": "notes", "format": "csv"} + ) + rows = list(csv.DictReader(io.StringIO(response.content.decode("utf-8-sig")))) + assert {r["作品ID"] for r in rows} == {"note-a", "note-b"} + + @pytest.mark.asyncio + async def test_xlsx_is_a_readable_workbook(self, client): + from openpyxl import load_workbook + + response = await client.get( + "/api/monitor/export", params={"kind": "comments", "format": "xlsx"} + ) + assert response.status_code == 200 + + workbook = load_workbook(io.BytesIO(response.content)) + sheet = workbook.active + assert sheet.max_row == 5 # header + four comments + assert sheet.cell(row=1, column=1).value == "所属作品" + + @pytest.mark.asyncio + async def test_report_export(self, client): + response = await client.get( + "/api/monitor/export", + params={"kind": "report", "days": 3}, + ) + rows = list(csv.DictReader(io.StringIO(response.content.decode("utf-8-sig")))) + assert len(rows) == 3 + assert "日期" in rows[0] + + @pytest.mark.asyncio + async def test_unknown_kind_and_format_are_rejected(self, client): + assert ( + await client.get("/api/monitor/export", params={"kind": "nope"}) + ).status_code == 400 + assert ( + await client.get("/api/monitor/export", params={"kind": "notes", "format": "pdf"}) + ).status_code == 400 + + @pytest.mark.asyncio + async def test_empty_selection_is_a_404_not_an_empty_file(self, client): + """An empty download looks like a bug; say so instead.""" + response = await client.get( + "/api/monitor/export", params={"kind": "comments", "note_id": "no-such-note"} + ) + assert response.status_code == 404 diff --git a/tests/test_monitor_ingest.py b/tests/test_monitor_ingest.py new file mode 100644 index 0000000..d090025 --- /dev/null +++ b/tests/test_monitor_ingest.py @@ -0,0 +1,527 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/tests/test_monitor_ingest.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Offline tests for the monitoring ingest/diff layer. + +These run without network, browser or login and cover the correctness caveats +that matter most: baseline suppression, count parsing, NULL-vs-zero, the +posted/seen comment split, idempotency, and the silent-cookie-failure signal. +""" + +import json +from pathlib import Path +from typing import Any, Dict, List, Optional + +import pytest +import pytest_asyncio +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine +from sqlalchemy.pool import StaticPool + +from tools.time_util import get_current_timestamp + +from api.monitor.ingest import describe_exit_code, ingest_run, parse_count +from api.monitor.models import ( + EVENT_AUTH_FAILURE, + EVENT_METRIC_DELTA, + EVENT_NEW_COMMENT_POSTED, + EVENT_NEW_COMMENT_SEEN, + EVENT_NEW_NOTE, + EVENT_NO_DATA, + EVENT_RUN_FAILED, + MODE_CREATOR, + MonitorBase, + MonitorEvent, + MonitorNote, + MonitorNoteMetric, + MonitorRun, + MonitorTask, + RUN_FAILED, + RUN_PARTIAL, + RUN_SUCCESS, +) + + +@pytest_asyncio.fixture +async def db(): + """An isolated in-memory monitoring database.""" + engine = create_async_engine("sqlite+aiosqlite://", poolclass=StaticPool) + async with engine.begin() as conn: + await conn.run_sync(MonitorBase.metadata.create_all) + + factory = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False) + async with factory() as db_session: + yield db_session + + await engine.dispose() + + +async def _make_task(db: AsyncSession, **overrides) -> MonitorTask: + defaults = dict( + name="test task", + platform="xhs", + mode=MODE_CREATOR, + enabled=True, + interval_minutes=60, + max_notes_count=20, + enable_comments=True, + max_comments_count=50, + run_timeout_seconds=3600, + created_at=0, + updated_at=0, + ) + defaults.update(overrides) + task = MonitorTask(**defaults) + db.add(task) + await db.flush() + return task + + +async def _make_run( + db: AsyncSession, + task: MonitorTask, + started_at: int, + exit_code: Optional[int] = 0, +) -> MonitorRun: + run = MonitorRun( + task_id=task.id, + trigger="manual", + status=RUN_SUCCESS, + phase=task.mode, + save_data_path="", + queued_at=started_at, + not_before=0, + started_at=started_at, + exit_code=exit_code, + ) + db.add(run) + await db.flush() + return run + + +def _write_run_dir( + root: Path, + notes: List[Dict[str, Any]], + comments: Optional[List[Dict[str, Any]]] = None, +) -> Path: + """Write a run's jsonl output in the crawler's own layout.""" + jsonl_dir = root / "xhs" / "jsonl" + jsonl_dir.mkdir(parents=True, exist_ok=True) + + contents = jsonl_dir / "creator_contents_2026-01-01.jsonl" + contents.write_text( + "\n".join(json.dumps(n, ensure_ascii=False) for n in notes), + encoding="utf-8", + ) + if comments is not None: + comment_file = jsonl_dir / "creator_comments_2026-01-01.jsonl" + comment_file.write_text( + "\n".join(json.dumps(c, ensure_ascii=False) for c in comments), + encoding="utf-8", + ) + return root + + +def _note(note_id: str, liked: Any = "10", **extra) -> Dict[str, Any]: + record = { + "note_id": note_id, + "title": f"title-{note_id}", + "note_url": f"https://www.xiaohongshu.com/explore/{note_id}", + "image_list": "https://img/cover.jpg", + "creator_hash": "hash", + "time": 1700000000000, + "liked_count": liked, + "comment_count": "1", + "collected_count": "1", + "share_count": "1", + } + record.update(extra) + return record + + +def _comment(comment_id: str, note_id: str, create_time: int, **extra) -> Dict[str, Any]: + record = { + "comment_id": comment_id, + "note_id": note_id, + "content": f"content-{comment_id}", + "nickname": "u***r", + "creator_hash": "hash", + "create_time": create_time, + "like_count": "0", + "sub_comment_count": 0, + "parent_comment_id": "", + } + record.update(extra) + return record + + +async def _events(db: AsyncSession, event_type: Optional[str] = None) -> List[MonitorEvent]: + stmt = select(MonitorEvent) + if event_type: + stmt = stmt.where(MonitorEvent.type == event_type) + return list((await db.scalars(stmt)).all()) + + +# -------------------------------------------------------------------------- +# parse_count +# -------------------------------------------------------------------------- + +class TestParseCount: + @pytest.mark.parametrize( + "raw,expected", + [ + ("1234", 1234), + ("1.2万", 12000), + ("1.2w", 12000), + ("3亿", 300000000), + ("1,234", 1234), + (42, 42), + ], + ) + def test_parses_platform_formats(self, raw, expected): + assert parse_count(raw) == expected + + @pytest.mark.parametrize("raw", ["", None, "暂无", "-", "abc", True]) + def test_unparseable_values_return_none(self, raw): + assert parse_count(raw) is None + + +# -------------------------------------------------------------------------- +# Exit codes +# -------------------------------------------------------------------------- + +class TestExitCodeStorage: + """Guards a bug that only showed up when the data moved to MySQL. + + Windows reports process failures as unsigned 32-bit NTSTATUS values + (0xC0000142 = 3221225794). That overflows MySQL's signed INT, while SQLite's + dynamic typing accepted it happily -- so the column silently worked until a + real migration hit it with real data. + """ + + def test_column_is_bigint_not_int(self): + from sqlalchemy import BigInteger + + from api.monitor.models import MonitorRun + + column_type = MonitorRun.__table__.c.exit_code.type + assert isinstance(column_type, BigInteger), ( + f"exit_code must be BigInteger to hold unsigned 32-bit codes, got {column_type!r}" + ) + + @pytest.mark.asyncio + async def test_an_ntstatus_value_round_trips(self, db): + task = await _make_task(db) + run = await _make_run(db, task, started_at=1000, exit_code=3221225794) + await db.commit() + + stored = await db.scalar( + select(MonitorRun.exit_code).where(MonitorRun.id == run.id) + ) + assert stored == 3221225794 + + +class TestDescribeExitCode: + def test_windows_status_code_is_decoded(self): + """3221225794 is 0xC0000142, which is meaningless without decoding.""" + message = describe_exit_code(3221225794) + assert "0xC0000142" in message + assert "DLL_INIT_FAILED" in message + + def test_negative_signed_form_is_also_decoded(self): + # Python may hand back the signed form depending on how it was launched. + assert "0xC0000142" in describe_exit_code(-1073741502) + + def test_unknown_code_degrades_to_the_raw_number(self): + assert describe_exit_code(1) == "Crawler exited with code 1" + + +# -------------------------------------------------------------------------- +# Notes +# -------------------------------------------------------------------------- + +class TestNoteIngest: + @pytest.mark.asyncio + async def test_baseline_run_emits_no_new_note_events(self, db, tmp_path): + task = await _make_task(db) + run = await _make_run(db, task, started_at=1000) + _write_run_dir(tmp_path, [_note("n1"), _note("n2")], comments=[]) + + result = await ingest_run(db, run, task, tmp_path) + + assert result.status == RUN_SUCCESS + assert result.is_baseline is True + assert result.new_notes == 2 + # Everything is "new" on the first run; emitting that would be pure noise. + assert await _events(db, EVENT_NEW_NOTE) == [] + assert len(list((await db.scalars(select(MonitorNote))).all())) == 2 + + @pytest.mark.asyncio + async def test_an_empty_run_does_not_establish_a_baseline(self, db, tmp_path): + """A run that fetched nothing observed nothing, so it is not a baseline. + + Otherwise the first crawl that actually works reports every work as + newly discovered. + """ + task = await _make_task(db) + + empty_run = await _make_run(db, task, started_at=1000) + (tmp_path / "empty").mkdir(parents=True, exist_ok=True) + await ingest_run(db, empty_run, task, tmp_path / "empty") + + real_run = await _make_run(db, task, started_at=2000) + result = await ingest_run( + db, real_run, task, _write_run_dir(tmp_path / "ok", [_note("n1")], comments=[]) + ) + + assert result.is_baseline is True + assert await _events(db, EVENT_NEW_NOTE) == [] + + @pytest.mark.asyncio + async def test_second_run_reports_only_the_added_note(self, db, tmp_path): + task = await _make_task(db) + + first_dir = _write_run_dir(tmp_path / "run1", [_note("n1")], comments=[]) + run1 = await _make_run(db, task, started_at=1000) + await ingest_run(db, run1, task, first_dir) + + second_dir = _write_run_dir(tmp_path / "run2", [_note("n1"), _note("n2")], comments=[]) + run2 = await _make_run(db, task, started_at=2000) + result = await ingest_run(db, run2, task, second_dir) + + assert result.is_baseline is False + assert result.new_notes == 1 + + events = await _events(db, EVENT_NEW_NOTE) + assert len(events) == 1 + assert events[0].target_id == "n2" + assert events[0].run_id == run2.id + + +class TestMetricSnapshots: + @pytest.mark.asyncio + async def test_delta_event_emitted_when_like_count_changes(self, db, tmp_path): + task = await _make_task(db) + + run1 = await _make_run(db, task, started_at=1000) + await ingest_run(db, run1, task, _write_run_dir(tmp_path / "r1", [_note("n1", "100")], comments=[])) + + run2 = await _make_run(db, task, started_at=2000) + await ingest_run(db, run2, task, _write_run_dir(tmp_path / "r2", [_note("n1", "150")], comments=[])) + + events = await _events(db, EVENT_METRIC_DELTA) + assert len(events) == 1 + + payload = json.loads(events[0].payload_json) + assert payload["deltas"]["liked_count"] == {"from": 100, "to": 150, "delta": 50} + + @pytest.mark.asyncio + async def test_no_delta_when_nothing_changed(self, db, tmp_path): + task = await _make_task(db) + run1 = await _make_run(db, task, started_at=1000) + await ingest_run(db, run1, task, _write_run_dir(tmp_path / "r1", [_note("n1", "100")], comments=[])) + run2 = await _make_run(db, task, started_at=2000) + await ingest_run(db, run2, task, _write_run_dir(tmp_path / "r2", [_note("n1", "100")], comments=[])) + + assert await _events(db, EVENT_METRIC_DELTA) == [] + + @pytest.mark.asyncio + async def test_unparseable_count_is_null_not_zero(self, db, tmp_path): + task = await _make_task(db) + run = await _make_run(db, task, started_at=1000) + await ingest_run(db, run, task, _write_run_dir(tmp_path, [_note("n1", "暂无")], comments=[])) + + metric = await db.scalar(select(MonitorNoteMetric).where(MonitorNoteMetric.note_id == "n1")) + # Zero would forge a large negative delta on the next comparison. + assert metric.liked_count is None + assert metric.raw_liked_count == "暂无" + + @pytest.mark.asyncio + async def test_no_delta_when_previous_value_was_unparseable(self, db, tmp_path): + task = await _make_task(db) + run1 = await _make_run(db, task, started_at=1000) + await ingest_run(db, run1, task, _write_run_dir(tmp_path / "r1", [_note("n1", "暂无")], comments=[])) + run2 = await _make_run(db, task, started_at=2000) + await ingest_run(db, run2, task, _write_run_dir(tmp_path / "r2", [_note("n1", "50")], comments=[])) + + assert await _events(db, EVENT_METRIC_DELTA) == [] + + @pytest.mark.asyncio + async def test_metric_snapshot_survives_across_runs(self, db, tmp_path): + """The crawler's own DB store overwrites metrics; ours must not.""" + task = await _make_task(db) + for index, liked in enumerate(["100", "150", "300"]): + run = await _make_run(db, task, started_at=1000 * (index + 1)) + await ingest_run( + db, run, task, _write_run_dir(tmp_path / f"r{index}", [_note("n1", liked)], comments=[]) + ) + + snapshots = list( + ( + await db.scalars( + select(MonitorNoteMetric) + .where(MonitorNoteMetric.note_id == "n1") + .order_by(MonitorNoteMetric.run_id) + ) + ).all() + ) + assert [s.liked_count for s in snapshots] == [100, 150, 300] + + +# -------------------------------------------------------------------------- +# Comments +# -------------------------------------------------------------------------- + +class TestCommentIngest: + @pytest.mark.asyncio + async def test_posted_vs_seen_split_by_create_time(self, db, tmp_path): + task = await _make_task(db) + + # Baseline establishes the seen-set; no events on the first run. + run1 = await _make_run(db, task, started_at=1000) + await ingest_run( + db, run1, task, + _write_run_dir(tmp_path / "r1", [_note("n1")], comments=[_comment("c1", "n1", create_time=500)]), + ) + assert await _events(db, EVENT_NEW_COMMENT_POSTED) == [] + + # c2 was published after run1 started -> genuinely new. + # c3 is old but only just surfaced in the top-N window -> seen, not posted. + run2 = await _make_run(db, task, started_at=2000) + await ingest_run( + db, run2, task, + _write_run_dir( + tmp_path / "r2", + [_note("n1")], + comments=[ + _comment("c1", "n1", create_time=500), + _comment("c2", "n1", create_time=2500), + _comment("c3", "n1", create_time=100), + ], + ), + ) + + posted = await _events(db, EVENT_NEW_COMMENT_POSTED) + seen = await _events(db, EVENT_NEW_COMMENT_SEEN) + assert len(posted) == 1 + assert json.loads(posted[0].payload_json)["comment_id"] == "c2" + assert len(seen) == 1 + assert json.loads(seen[0].payload_json)["comment_id"] == "c3" + + @pytest.mark.asyncio + async def test_comments_not_ingested_when_disabled(self, db, tmp_path): + task = await _make_task(db, enable_comments=False) + run = await _make_run(db, task, started_at=1000) + result = await ingest_run( + db, run, task, + _write_run_dir(tmp_path, [_note("n1")], comments=[_comment("c1", "n1", 500)]), + ) + assert result.new_comments == 0 + + +# -------------------------------------------------------------------------- +# Failure handling +# -------------------------------------------------------------------------- + +class TestFailureHandling: + @pytest.mark.asyncio + async def test_nonzero_exit_is_a_failure(self, db, tmp_path): + task = await _make_task(db) + run = await _make_run(db, task, started_at=1000, exit_code=1) + _write_run_dir(tmp_path, [_note("n1")], comments=[]) + + result = await ingest_run(db, run, task, tmp_path) + + assert result.status == RUN_FAILED + assert len(await _events(db, EVENT_RUN_FAILED)) == 1 + # A crashed run must not touch the seen-set. + assert await db.scalar(select(MonitorNote.id)) is None + + @pytest.mark.asyncio + async def test_zero_notes_with_exit_zero_is_a_suspected_auth_failure(self, db, tmp_path): + """The silent-cookie-failure signature: exit 0 but nothing fetched. + + A real bad-cookie run writes no output file at all, which is why the + exit code has to be checked before the files are. + """ + task = await _make_task(db) + run = await _make_run(db, task, started_at=1000, exit_code=0) + tmp_path.mkdir(parents=True, exist_ok=True) + + result = await ingest_run(db, run, task, tmp_path) + + assert result.status == RUN_PARTIAL + assert len(await _events(db, EVENT_AUTH_FAILURE)) == 1 + assert await _events(db, EVENT_RUN_FAILED) == [] + + @pytest.mark.asyncio + async def test_no_data_is_not_blamed_on_the_cookie_when_a_sibling_succeeded( + self, db, tmp_path + ): + """A task that just worked proves the login is fine; do not cry wolf.""" + healthy = await _make_task(db, name="healthy") + healthy_run = await _make_run(db, healthy, started_at=get_current_timestamp()) + await ingest_run( + db, healthy_run, healthy, + _write_run_dir(tmp_path / "ok", [_note("n1")], comments=[]), + ) + + task = await _make_task(db, name="suspect") + run = await _make_run(db, task, started_at=get_current_timestamp()) + (tmp_path / "empty").mkdir(parents=True, exist_ok=True) + result = await ingest_run(db, run, task, tmp_path / "empty") + + assert result.status == RUN_PARTIAL + assert await _events(db, EVENT_NO_DATA) != [] + assert await _events(db, EVENT_AUTH_FAILURE) == [] + + @pytest.mark.asyncio + async def test_empty_contents_file_is_also_an_auth_failure(self, db, tmp_path): + task = await _make_task(db) + run = await _make_run(db, task, started_at=1000, exit_code=0) + _write_run_dir(tmp_path, [], comments=[]) + + result = await ingest_run(db, run, task, tmp_path) + + assert result.status == RUN_PARTIAL + assert len(await _events(db, EVENT_AUTH_FAILURE)) == 1 + + +# -------------------------------------------------------------------------- +# Idempotency +# -------------------------------------------------------------------------- + +class TestIdempotency: + @pytest.mark.asyncio + async def test_reingesting_the_same_data_adds_nothing(self, db, tmp_path): + task = await _make_task(db) + run_dir = _write_run_dir( + tmp_path, [_note("n1"), _note("n2")], comments=[_comment("c1", "n1", 500)] + ) + + run1 = await _make_run(db, task, started_at=1000) + await ingest_run(db, run1, task, run_dir) + notes_after_first = len(list((await db.scalars(select(MonitorNote))).all())) + + # A retry of the same crawl content must not duplicate rows or events. + run2 = await _make_run(db, task, started_at=2000) + result = await ingest_run(db, run2, task, run_dir) + + assert result.new_notes == 0 + assert result.new_comments == 0 + assert len(list((await db.scalars(select(MonitorNote))).all())) == notes_after_first diff --git a/tests/test_monitor_notify.py b/tests/test_monitor_notify.py new file mode 100644 index 0000000..9db23fe --- /dev/null +++ b/tests/test_monitor_notify.py @@ -0,0 +1,315 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/tests/test_monitor_notify.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Tests for the WeCom notification layer. + +The webhook is stubbed, so nothing here touches the network. +""" + +import json + +import pytest +import pytest_asyncio +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine +from sqlalchemy.pool import StaticPool + +from api.monitor import notify +from api.monitor.models import ( + EVENT_AUTH_FAILURE, + EVENT_METRIC_DELTA, + EVENT_NEW_NOTE, + EVENT_NEW_COMMENT_POSTED, + MODE_CREATOR, + SETTING_WECOM_WEBHOOK, + MonitorBase, + MonitorEvent, + MonitorRun, + MonitorTask, + RUN_SUCCESS, +) +from api.monitor.settings import set_setting + +WEBHOOK = "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=abc123" + + +@pytest_asyncio.fixture +async def db(): + engine = create_async_engine("sqlite+aiosqlite://", poolclass=StaticPool) + async with engine.begin() as conn: + await conn.run_sync(MonitorBase.metadata.create_all) + factory = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False) + async with factory() as session: + yield session + await engine.dispose() + + +async def _seed(db: AsyncSession, notify_enabled: bool = True): + task = MonitorTask( + name="竞品监控", platform="xhs", mode=MODE_CREATOR, enabled=True, + interval_minutes=60, max_notes_count=20, enable_comments=True, + max_comments_count=50, run_timeout_seconds=3600, + notify_enabled=notify_enabled, created_at=0, updated_at=0, + ) + db.add(task) + await db.flush() + + run = MonitorRun( + task_id=task.id, trigger="scheduled", status=RUN_SUCCESS, phase=MODE_CREATOR, + save_data_path="", queued_at=0, not_before=0, max_comments_count=50, + ) + db.add(run) + await db.flush() + return task, run + + +def _add_event(db, task, run, event_type, title, payload=None, severity="info"): + db.add( + MonitorEvent( + task_id=task.id, run_id=run.id, type=event_type, severity=severity, + target_kind="note", target_id="note-1", title=title, + payload_json=json.dumps(payload or {}, ensure_ascii=False), created_at=0, + ) + ) + + +# -------------------------------------------------------------------------- +# Message building +# -------------------------------------------------------------------------- + +class TestBuildRunMessage: + @pytest.mark.asyncio + async def test_no_notifiable_events_means_no_message(self, db): + task, run = await _seed(db) + # Metric deltas are not something anyone wants pushed. + _add_event(db, task, run, EVENT_METRIC_DELTA, "点赞 10→20") + _add_event(db, task, run, EVENT_NEW_COMMENT_POSTED, "新评论") + await db.flush() + + assert await notify.build_run_message(db, task, run) is None + + @pytest.mark.asyncio + async def test_new_notes_are_listed_with_links(self, db): + task, run = await _seed(db) + _add_event( + db, task, run, EVENT_NEW_NOTE, "新作品:标题A", + payload={"note_id": "abc123", "title": "标题A"}, + ) + await db.flush() + + message = await notify.build_run_message(db, task, run) + + assert "竞品监控" in message + assert "新增作品 **1** 篇" in message + assert "标题A" in message + assert "https://www.xiaohongshu.com/explore/abc123" in message + + @pytest.mark.asyncio + async def test_long_note_lists_are_truncated(self, db): + """A first run can find dozens; a wall of text is worse than a count.""" + task, run = await _seed(db) + for index in range(14): + _add_event( + db, task, run, EVENT_NEW_NOTE, f"新作品:{index}", + payload={"note_id": f"n{index}", "title": f"标题{index}"}, + ) + await db.flush() + + message = await notify.build_run_message(db, task, run) + + assert "新增作品 **14** 篇" in message + assert "标题0" in message + assert "标题13" not in message + assert "等共 14 篇" in message + + @pytest.mark.asyncio + async def test_failure_is_reported_as_a_warning(self, db): + task, run = await _seed(db) + _add_event( + db, task, run, EVENT_AUTH_FAILURE, + "疑似登录态失效:本次未抓到任何作品", severity="error", + ) + await db.flush() + + message = await notify.build_run_message(db, task, run) + + assert "异常" in message + assert "登录态失效" in message + assert notify._COLOR_WARNING in message + + @pytest.mark.asyncio + async def test_baseline_runs_say_so(self, db): + task, run = await _seed(db) + run.is_baseline = True + _add_event(db, task, run, EVENT_NEW_NOTE, "新作品", payload={"note_id": "x", "title": "t"}) + await db.flush() + + message = await notify.build_run_message(db, task, run) + + assert "基线" in message + + +# -------------------------------------------------------------------------- +# notify_run gating +# -------------------------------------------------------------------------- + +class TestNotifyRunGating: + @pytest.mark.asyncio + async def test_disabled_task_is_skipped(self, db, monkeypatch): + task, run = await _seed(db, notify_enabled=False) + _add_event(db, task, run, EVENT_NEW_NOTE, "新作品", payload={"note_id": "x", "title": "t"}) + await set_setting(db, SETTING_WECOM_WEBHOOK, WEBHOOK) + await db.flush() + + called = [] + monkeypatch.setattr(notify, "send_wecom", lambda *a, **k: called.append(a) or _ok()) + + assert await notify.notify_run(db, task, run) is None + assert called == [] + + @pytest.mark.asyncio + async def test_missing_webhook_is_skipped(self, db, monkeypatch): + task, run = await _seed(db, notify_enabled=True) + _add_event(db, task, run, EVENT_NEW_NOTE, "新作品", payload={"note_id": "x", "title": "t"}) + await db.flush() + + called = [] + monkeypatch.setattr(notify, "send_wecom", lambda *a, **k: called.append(a) or _ok()) + + assert await notify.notify_run(db, task, run) is None + assert called == [] + + @pytest.mark.asyncio + async def test_successful_push_records_the_timestamp(self, db, monkeypatch): + task, run = await _seed(db, notify_enabled=True) + _add_event(db, task, run, EVENT_NEW_NOTE, "新作品", payload={"note_id": "x", "title": "t"}) + await set_setting(db, SETTING_WECOM_WEBHOOK, WEBHOOK) + await db.flush() + + monkeypatch.setattr(notify, "send_wecom", lambda *a, **k: _ok()) + + message = await notify.notify_run(db, task, run) + + assert message is not None + # Lets the UI answer "why did I not get a push for this run?". + assert task.last_notified_at is not None + + @pytest.mark.asyncio + async def test_push_failure_never_raises(self, db, monkeypatch): + """A broken webhook must not take down the crawl that just succeeded.""" + task, run = await _seed(db, notify_enabled=True) + _add_event(db, task, run, EVENT_NEW_NOTE, "新作品", payload={"note_id": "x", "title": "t"}) + await set_setting(db, SETTING_WECOM_WEBHOOK, WEBHOOK) + await db.flush() + + async def _boom(*args, **kwargs): + raise RuntimeError("network exploded") + + monkeypatch.setattr(notify, "send_wecom", _boom) + + assert await notify.notify_run(db, task, run) is None + + +async def _ok(): + return True, "发送成功" + + +# -------------------------------------------------------------------------- +# send_wecom +# -------------------------------------------------------------------------- + +class _FakeResponse: + def __init__(self, payload): + self._payload = payload + + def raise_for_status(self): + return None + + def json(self): + return self._payload + + +class _FakeClient: + """Captures the request and replays a canned WeCom reply.""" + + last_payload = None + + def __init__(self, reply=None, error=None): + self._reply = reply if reply is not None else {"errcode": 0, "errmsg": "ok"} + self._error = error + + def __call__(self, *args, **kwargs): + return self + + async def __aenter__(self): + return self + + async def __aexit__(self, *exc): + return False + + async def post(self, url, json=None): + if self._error: + raise self._error + type(self).last_payload = json + return _FakeResponse(self._reply) + + +class TestSendWecom: + @pytest.mark.asyncio + async def test_missing_url_is_reported(self): + ok, detail = await notify.send_wecom("", "hi") + assert ok is False + assert "未配置" in detail + + @pytest.mark.asyncio + async def test_success(self, monkeypatch): + monkeypatch.setattr(notify.httpx, "AsyncClient", _FakeClient()) + + ok, detail = await notify.send_wecom(WEBHOOK, "**标题**\n> 内容") + + assert ok is True + assert detail == "发送成功" + # WeCom expects a markdown message envelope. + assert _FakeClient.last_payload["msgtype"] == "markdown" + assert _FakeClient.last_payload["markdown"]["content"] == "**标题**\n> 内容" + + @pytest.mark.asyncio + async def test_nonzero_errcode_is_a_failure(self, monkeypatch): + """WeCom answers HTTP 200 even when it rejects the message.""" + monkeypatch.setattr( + notify.httpx, "AsyncClient", + _FakeClient(reply={"errcode": 93000, "errmsg": "invalid webhook url"}), + ) + + ok, detail = await notify.send_wecom(WEBHOOK, "hi") + + assert ok is False + assert "93000" in detail + + @pytest.mark.asyncio + async def test_network_error_is_returned_not_raised(self, monkeypatch): + import httpx + + monkeypatch.setattr( + notify.httpx, "AsyncClient", + _FakeClient(error=httpx.ConnectError("boom")), + ) + + ok, detail = await notify.send_wecom(WEBHOOK, "hi") + + assert ok is False + assert "请求失败" in detail diff --git a/tests/test_monitor_report.py b/tests/test_monitor_report.py new file mode 100644 index 0000000..f71726b --- /dev/null +++ b/tests/test_monitor_report.py @@ -0,0 +1,304 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/tests/test_monitor_report.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Tests for the cross-task report aggregation. + +The interaction delta is the part that is easy to get subtly wrong, so it is +covered directly against the pure aggregation function. +""" + +from datetime import date, datetime + +import pytest +import pytest_asyncio +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine +from sqlalchemy.pool import StaticPool + +from api.monitor.models import ( + MODE_CREATOR, + MonitorBase, + MonitorComment, + MonitorNote, + MonitorNoteMetric, + MonitorTask, +) +from api.monitor.report import build_report, compute_daily_rows, day_bounds, iter_days + + +def _ms(year: int, month: int, day: int, hour: int = 12) -> int: + return int(datetime(year, month, day, hour).timestamp() * 1000) + + +def _metrics(liked=0, comment=0, collected=0, share=0): + """All four metrics default to parsed values; pass None to simulate a + platform value we could not parse.""" + return { + "liked_count": liked, + "comment_count": comment, + "collected_count": collected, + "share_count": share, + } + + +class TestDayHelpers: + def test_day_bounds_cover_the_whole_local_day(self): + start, end = day_bounds(date(2026, 1, 10)) + assert start < _ms(2026, 1, 10, 0) or start == _ms(2026, 1, 10, 0) + assert end > _ms(2026, 1, 10, 23) + + def test_iter_days_is_inclusive(self): + days = iter_days(date(2026, 1, 10), date(2026, 1, 12)) + assert days == [date(2026, 1, 10), date(2026, 1, 11), date(2026, 1, 12)] + + +class TestInteractionDelta: + def test_note_first_seen_counts_all_of_its_value(self): + """A brand-new note has no earlier baseline, so it starts from zero.""" + day = date(2026, 1, 10) + series = {"n1": [(_ms(2026, 1, 10, 10), _metrics(liked=100, comment=5))]} + + rows = compute_daily_rows(series, {}, {}, [day]) + + assert rows[0]["liked_count_delta"] == 100 + assert rows[0]["comment_count_delta"] == 5 + + def test_growth_is_split_across_days(self): + series = { + "n1": [ + (_ms(2026, 1, 10, 10), _metrics(liked=100)), + (_ms(2026, 1, 11, 10), _metrics(liked=300)), + ] + } + + rows = compute_daily_rows(series, {}, {}, [date(2026, 1, 10), date(2026, 1, 11)]) + + # Day 1: 0 -> 100. Day 2: 100 -> 300. + assert [row["liked_count_delta"] for row in rows] == [100, 200] + + def test_day_without_a_snapshot_reports_no_growth(self): + series = { + "n1": [ + (_ms(2026, 1, 10, 10), _metrics(liked=100)), + (_ms(2026, 1, 12, 10), _metrics(liked=400)), + ] + } + days = [date(2026, 1, 10), date(2026, 1, 11), date(2026, 1, 12)] + + rows = compute_daily_rows(series, {}, {}, days) + + # The note was not crawled on the 11th, so nothing is claimed for it. + assert [row["liked_count_delta"] for row in rows] == [100, 0, 300] + + def test_deltas_aggregate_across_notes(self): + series = { + "n1": [ + (_ms(2026, 1, 10, 10), _metrics(liked=100)), + (_ms(2026, 1, 11, 10), _metrics(liked=150)), + ], + "n2": [ + (_ms(2026, 1, 10, 10), _metrics(liked=10)), + (_ms(2026, 1, 11, 10), _metrics(liked=40)), + ], + } + + rows = compute_daily_rows(series, {}, {}, [date(2026, 1, 10), date(2026, 1, 11)]) + + assert [row["liked_count_delta"] for row in rows] == [110, 80] + + def test_unparseable_metric_names_the_offending_field(self): + """A NULL count makes the delta unknown; it must not be reported as 0.""" + series = { + "n1": [ + (_ms(2026, 1, 10, 10), _metrics(liked=100, comment=None)), + (_ms(2026, 1, 11, 10), _metrics(liked=200, comment=None)), + ] + } + + rows = compute_daily_rows(series, {}, {}, [date(2026, 1, 11)]) + + # Naming the field is actionable; a bare boolean is not. + assert rows[0]["partial_metrics"] == ["comment_count"] + # The parseable metric is still summed correctly. + assert rows[0]["liked_count_delta"] == 100 + + def test_unknown_value_only_taints_the_days_it_touches(self): + series = { + "n1": [ + (_ms(2026, 1, 10, 10), _metrics(liked=None)), + (_ms(2026, 1, 11, 10), _metrics(liked=50)), + (_ms(2026, 1, 12, 10), _metrics(liked=90)), + ] + } + days = [date(2026, 1, 10), date(2026, 1, 11), date(2026, 1, 12)] + + rows = compute_daily_rows(series, {}, {}, days) + + # Day 12 compares two known values, so it is clean. + assert [row["partial_metrics"] for row in rows] == [ + ["liked_count"], + ["liked_count"], + [], + ] + assert rows[2]["liked_count_delta"] == 40 + + def test_new_content_counts_come_from_the_day_maps(self): + rows = compute_daily_rows( + {}, + {date(2026, 1, 10): 3}, + {date(2026, 1, 10): 7}, + [date(2026, 1, 10), date(2026, 1, 11)], + ) + + assert rows[0]["new_notes"] == 3 + assert rows[0]["new_comments"] == 7 + assert rows[1]["new_notes"] == 0 + + +# -------------------------------------------------------------------------- +# DB-backed report + task filtering +# -------------------------------------------------------------------------- + +@pytest_asyncio.fixture +async def db(): + engine = create_async_engine("sqlite+aiosqlite://", poolclass=StaticPool) + async with engine.begin() as conn: + await conn.run_sync(MonitorBase.metadata.create_all) + factory = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False) + async with factory() as session: + yield session + await engine.dispose() + + +async def _seed_task(db: AsyncSession, name: str) -> MonitorTask: + task = MonitorTask( + name=name, platform="xhs", mode=MODE_CREATOR, enabled=True, + interval_minutes=60, max_notes_count=20, enable_comments=True, + max_comments_count=50, run_timeout_seconds=3600, + notify_enabled=False, created_at=0, updated_at=0, + ) + db.add(task) + await db.flush() + return task + + +async def _seed_note_with_metrics( + db: AsyncSession, task: MonitorTask, note_id: str, samples +) -> None: + db.add( + MonitorNote( + task_id=task.id, note_id=note_id, title=note_id, note_url="", + cover="", creator_hash="", source_kind="", published_at=None, + first_seen_run_id=1, first_seen_at=samples[0][0], + last_seen_run_id=len(samples), last_seen_at=samples[-1][0], + ) + ) + for run_id, (ts, liked) in enumerate(samples, start=1): + db.add( + MonitorNoteMetric( + task_id=task.id, note_id=note_id, run_id=run_id, captured_at=ts, + liked_count=liked, comment_count=0, collected_count=0, share_count=0, + raw_liked_count=str(liked), raw_comment_count="0", + raw_collected_count="0", raw_share_count="0", + ) + ) + + +class TestBuildReport: + @pytest.mark.asyncio + async def test_totals_and_rows(self, db): + task = await _seed_task(db, "t1") + await _seed_note_with_metrics( + db, task, "n1", + [(_ms(2026, 1, 10, 10), 100), (_ms(2026, 1, 11, 10), 250)], + ) + await db.commit() + + result = await build_report(db, [task.id], date(2026, 1, 10), date(2026, 1, 11)) + + assert result["totals"]["liked_count_delta"] == 250 + assert len(result["rows"]) == 2 + assert result["note_count"] == 1 + + @pytest.mark.asyncio + async def test_task_selection_isolates_the_report(self, db): + """The whole point: a report for a chosen subset must exclude the rest.""" + kept = await _seed_task(db, "kept") + other = await _seed_task(db, "other") + await _seed_note_with_metrics(db, kept, "n1", [(_ms(2026, 1, 10, 10), 100)]) + await _seed_note_with_metrics(db, other, "n2", [(_ms(2026, 1, 10, 10), 999)]) + await db.commit() + + only_kept = await build_report(db, [kept.id], date(2026, 1, 10), date(2026, 1, 10)) + assert only_kept["totals"]["liked_count_delta"] == 100 + assert only_kept["note_count"] == 1 + + both = await build_report(db, [kept.id, other.id], date(2026, 1, 10), date(2026, 1, 10)) + assert both["totals"]["liked_count_delta"] == 1099 + + @pytest.mark.asyncio + async def test_no_task_filter_covers_everything(self, db): + first = await _seed_task(db, "a") + second = await _seed_task(db, "b") + await _seed_note_with_metrics(db, first, "n1", [(_ms(2026, 1, 10, 10), 10)]) + await _seed_note_with_metrics(db, second, "n2", [(_ms(2026, 1, 10, 10), 20)]) + await db.commit() + + result = await build_report(db, None, date(2026, 1, 10), date(2026, 1, 10)) + + assert result["totals"]["liked_count_delta"] == 30 + assert result["task_ids"] is None + + @pytest.mark.asyncio + async def test_baseline_from_before_the_range_is_used(self, db): + """Growth is measured against the last value before the window opens.""" + task = await _seed_task(db, "t") + await _seed_note_with_metrics( + db, task, "n1", + [(_ms(2026, 1, 5, 10), 1000), (_ms(2026, 1, 10, 10), 1050)], + ) + await db.commit() + + # Report only for the 10th: the delta must be 50, not 1050. + result = await build_report(db, [task.id], date(2026, 1, 10), date(2026, 1, 10)) + + assert result["totals"]["liked_count_delta"] == 50 + + @pytest.mark.asyncio + async def test_empty_range_returns_zeroed_rows(self, db): + result = await build_report(db, None, date(2026, 2, 1), date(2026, 2, 3)) + + assert len(result["rows"]) == 3 + assert result["totals"]["liked_count_delta"] == 0 + assert result["totals"]["new_notes"] == 0 + + @pytest.mark.asyncio + async def test_new_comments_are_counted_by_first_seen_day(self, db): + task = await _seed_task(db, "t") + db.add( + MonitorComment( + task_id=task.id, note_id="n1", comment_id="c1", content="x", + nickname="u", creator_hash="h", create_time=_ms(2026, 1, 9), + like_count=0, sub_comment_count=0, parent_comment_id="", + first_seen_run_id=1, first_seen_at=_ms(2026, 1, 10, 10), + ) + ) + await db.commit() + + result = await build_report(db, [task.id], date(2026, 1, 10), date(2026, 1, 10)) + + assert result["totals"]["new_comments"] == 1 diff --git a/tests/test_monitor_scheduler.py b/tests/test_monitor_scheduler.py new file mode 100644 index 0000000..339a487 --- /dev/null +++ b/tests/test_monitor_scheduler.py @@ -0,0 +1,266 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/tests/test_monitor_scheduler.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Tests for the monitor scheduler's firing, deferral and recovery rules.""" + +import pytest +import pytest_asyncio +from sqlalchemy import select + +from api.monitor import db as monitor_db +from api.monitor import scheduler as scheduler_module +from api.monitor.models import ( + MODE_CREATOR, + MonitorRun, + MonitorTarget, + MonitorTask, + RUN_INTERRUPTED, + RUN_RUNNING, + RUN_SUCCESS, +) +from api.monitor.scheduler import MonitorScheduler +from api.monitor.settings import set_cookie +from tools.time_util import get_current_timestamp + +MS_PER_MINUTE = 60_000 + + +class FakeCrawlerManager: + """Stands in for the global subprocess singleton.""" + + def __init__(self, busy: bool = False) -> None: + self.busy = busy + + def is_busy(self) -> bool: + return self.busy + + +@pytest_asyncio.fixture +async def db(tmp_path): + monitor_db.set_sqlite_path(tmp_path / "monitor.db") + await monitor_db.init_db() + async with monitor_db.get_session() as session: + await set_cookie(session, "web_session=test") + yield monitor_db + await monitor_db.dispose_engine() + + +@pytest_asyncio.fixture +async def executed(monkeypatch): + """Record execute_task calls instead of launching a real crawl.""" + calls: list[tuple[int, str]] = [] + + async def _fake_execute(task_id: int, trigger: str = "manual"): + calls.append((task_id, trigger)) + + monkeypatch.setattr(scheduler_module, "execute_task", _fake_execute) + return calls + + +async def _make_task(next_run_at, enabled: bool = True, interval: int = 60) -> int: + async with monitor_db.get_session() as session: + now = get_current_timestamp() + task = MonitorTask( + name="t", + platform="xhs", + mode=MODE_CREATOR, + enabled=enabled, + interval_minutes=interval, + max_notes_count=20, + enable_comments=True, + max_comments_count=50, + run_timeout_seconds=3600, + next_run_at=next_run_at, + last_status="idle", + created_at=now, + updated_at=now, + ) + session.add(task) + await session.flush() + session.add( + MonitorTarget( + task_id=task.id, + kind=MODE_CREATOR, + external_id="abc123", + xsec_token="", + xsec_source="", + raw_value="abc123", + label="abc123", + enabled=True, + created_at=now, + ) + ) + return task.id + + +async def _get_task(task_id: int) -> MonitorTask: + async with monitor_db.get_session() as session: + return await session.get(MonitorTask, task_id) + + +class TestFiring: + @pytest.mark.asyncio + async def test_due_task_runs_and_advances(self, monkeypatch, db, executed): + monkeypatch.setattr(scheduler_module, "crawler_manager", FakeCrawlerManager(busy=False)) + past = get_current_timestamp() - MS_PER_MINUTE + task_id = await _make_task(past) + + await MonitorScheduler().tick() + + assert executed == [(task_id, "scheduled")] + task = await _get_task(task_id) + # Fixed-delay: the next fire is measured from now, not from the missed slot. + assert task.next_run_at > get_current_timestamp() + + @pytest.mark.asyncio + async def test_future_task_does_not_run(self, monkeypatch, db, executed): + monkeypatch.setattr(scheduler_module, "crawler_manager", FakeCrawlerManager(busy=False)) + await _make_task(get_current_timestamp() + 10 * MS_PER_MINUTE) + + await MonitorScheduler().tick() + + assert executed == [] + + @pytest.mark.asyncio + async def test_disabled_task_does_not_run(self, monkeypatch, db, executed): + monkeypatch.setattr(scheduler_module, "crawler_manager", FakeCrawlerManager(busy=False)) + await _make_task(get_current_timestamp() - MS_PER_MINUTE, enabled=False) + + await MonitorScheduler().tick() + + assert executed == [] + + @pytest.mark.asyncio + async def test_long_outage_coalesces_into_one_run(self, monkeypatch, db, executed): + """A missed schedule fires once, not once per missed interval.""" + monkeypatch.setattr(scheduler_module, "crawler_manager", FakeCrawlerManager(busy=False)) + # Due two days ago on a 1-hour interval. + await _make_task(get_current_timestamp() - 48 * 60 * MS_PER_MINUTE) + + scheduler = MonitorScheduler() + await scheduler.tick() + await scheduler.tick() + + assert len(executed) == 1 + + +class TestDeferral: + @pytest.mark.asyncio + async def test_busy_crawler_defers_without_advancing(self, monkeypatch, db, executed): + """A manual crawl must not consume the monitor task's slot or lose it.""" + monkeypatch.setattr(scheduler_module, "crawler_manager", FakeCrawlerManager(busy=True)) + due_at = get_current_timestamp() - MS_PER_MINUTE + task_id = await _make_task(due_at) + + await MonitorScheduler().tick() + + assert executed == [] + task = await _get_task(task_id) + # Still due, so the next free tick picks it up rather than skipping a cycle. + assert task.next_run_at == due_at + + @pytest.mark.asyncio + async def test_deferred_task_runs_once_crawler_frees_up(self, monkeypatch, db, executed): + fake = FakeCrawlerManager(busy=True) + monkeypatch.setattr(scheduler_module, "crawler_manager", fake) + task_id = await _make_task(get_current_timestamp() - MS_PER_MINUTE) + + scheduler = MonitorScheduler() + await scheduler.tick() + assert executed == [] + + fake.busy = False + await scheduler.tick() + assert executed == [(task_id, "scheduled")] + + +class TestCookieGuard: + @pytest.mark.asyncio + async def test_no_cookie_blocks_run_and_keeps_task_due(self, monkeypatch, db, executed): + """Without a cookie every run would be an auth failure; skip instead.""" + monkeypatch.setattr(scheduler_module, "crawler_manager", FakeCrawlerManager(busy=False)) + async with monitor_db.get_session() as session: + from api.monitor.settings import cookie_key, delete_setting + + await delete_setting(session, cookie_key("xhs")) + + due_at = get_current_timestamp() - MS_PER_MINUTE + task_id = await _make_task(due_at) + + await MonitorScheduler().tick() + + assert executed == [] + task = await _get_task(task_id) + # Left due so it starts working the moment a cookie is pasted. + assert task.next_run_at == due_at + + +class TestRecovery: + @pytest.mark.asyncio + async def test_running_runs_are_marked_interrupted(self, db): + """A run left 'running' cannot be alive -- its process died with the server.""" + async with monitor_db.get_session() as session: + now = get_current_timestamp() + task = MonitorTask( + name="t", platform="xhs", mode=MODE_CREATOR, enabled=True, + interval_minutes=60, max_notes_count=20, enable_comments=True, + max_comments_count=50, run_timeout_seconds=3600, + next_run_at=now, last_status="running", created_at=now, updated_at=now, + ) + session.add(task) + await session.flush() + session.add( + MonitorRun( + task_id=task.id, trigger="scheduled", status=RUN_RUNNING, + phase=MODE_CREATOR, save_data_path="", queued_at=now, not_before=0, + started_at=now, max_comments_count=50, + ) + ) + + await MonitorScheduler().recover() + + async with monitor_db.get_session() as session: + run = await session.scalar(select(MonitorRun)) + assert run.status == RUN_INTERRUPTED + assert run.finished_at is not None + + @pytest.mark.asyncio + async def test_completed_runs_are_left_alone(self, db): + async with monitor_db.get_session() as session: + now = get_current_timestamp() + task = MonitorTask( + name="t", platform="xhs", mode=MODE_CREATOR, enabled=True, + interval_minutes=60, max_notes_count=20, enable_comments=True, + max_comments_count=50, run_timeout_seconds=3600, + next_run_at=now, last_status="success", created_at=now, updated_at=now, + ) + session.add(task) + await session.flush() + session.add( + MonitorRun( + task_id=task.id, trigger="scheduled", status=RUN_SUCCESS, + phase=MODE_CREATOR, save_data_path="", queued_at=now, not_before=0, + started_at=now, finished_at=now, max_comments_count=50, + ) + ) + + await MonitorScheduler().recover() + + async with monitor_db.get_session() as session: + run = await session.scalar(select(MonitorRun)) + assert run.status == RUN_SUCCESS diff --git a/tests/test_platforms.py b/tests/test_platforms.py new file mode 100644 index 0000000..a079ed0 --- /dev/null +++ b/tests/test_platforms.py @@ -0,0 +1,302 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/tests/test_platforms.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Platform capability matrix, platform scoping, and per-platform settings.""" + +import httpx +import pytest +import pytest_asyncio +from sqlalchemy import text + +from api.main import app +from api.monitor import db as monitor_db +from api.monitor import platforms +from api.monitor.models import MonitorTask + +XHS_TARGET = "5f58bd990000000001003753" + + +@pytest_asyncio.fixture +async def client(tmp_path): + monitor_db.set_sqlite_path(tmp_path / "monitor.db") + await monitor_db.init_db() + + transport = httpx.ASGITransport(app=app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as http_client: + yield http_client + + await monitor_db.dispose_engine() + + +class TestCapabilityMatrix: + @pytest.mark.asyncio + async def test_matrix_is_exposed_to_the_ui(self, client): + body = (await client.get("/api/config/platforms")).json() + by_value = {p["value"]: p for p in body["platforms"]} + + assert set(by_value) == {"xhs", "dy", "ks", "bili", "wb", "tieba", "zhihu"} + # Every entry must say whether monitoring is actually wired up -- this is + # what stops the UI offering a platform that can never produce data. + assert all("monitor_wired" in p for p in body["platforms"]) + assert by_value["xhs"]["monitor_wired"] is True + assert by_value["dy"]["monitor_wired"] is False + + @pytest.mark.asyncio + async def test_metrics_are_per_platform_and_labelled(self, client): + body = (await client.get("/api/config/platforms")).json() + by_value = {p["value"]: p for p in body["platforms"]} + + # Bilibili has play count and danmaku; Xiaohongshu has neither. + assert "video_play_count" in by_value["bili"]["metrics"] + assert "video_danmaku" in by_value["bili"]["metrics"] + assert "video_play_count" not in by_value["xhs"]["metrics"] + + # Every metric shown to a user must have a human label. + for capability in body["platforms"]: + for metric in capability["metrics"]: + assert capability["metric_labels"][metric] + + def test_unknown_platform_is_not_monitor_wired(self): + assert platforms.is_known("xhs") is True + assert platforms.is_known("myspace") is False + assert platforms.is_monitor_wired("myspace") is False + + +class TestTaskCreationGuard: + @pytest.mark.asyncio + async def test_unwired_platform_is_rejected_with_an_explanation(self, client): + """Accepting it would create a task that silently never produces data.""" + response = await client.post( + "/api/monitor/tasks", + json={"name": "抖音任务", "mode": "creator", "platform": "dy", "targets": ["x"]}, + ) + assert response.status_code == 400 + detail = response.json()["detail"] + assert "抖音" in detail + assert "尚未接通" in detail + + @pytest.mark.asyncio + async def test_unknown_platform_is_rejected(self, client): + response = await client.post( + "/api/monitor/tasks", + json={"name": "x", "mode": "creator", "platform": "myspace", "targets": ["x"]}, + ) + assert response.status_code == 400 + + @pytest.mark.asyncio + async def test_no_task_row_is_created_when_rejected(self, client): + await client.post( + "/api/monitor/tasks", + json={"name": "抖音任务", "mode": "creator", "platform": "dy", "targets": ["x"]}, + ) + assert (await client.get("/api/monitor/tasks")).json()["tasks"] == [] + + @pytest.mark.asyncio + async def test_xhs_still_works_and_is_the_default(self, client): + explicit = await client.post( + "/api/monitor/tasks", + json={"name": "显式", "mode": "creator", "platform": "xhs", "targets": [XHS_TARGET]}, + ) + assert explicit.status_code == 201 + + defaulted = await client.post( + "/api/monitor/tasks", + json={"name": "默认", "mode": "creator", "targets": [XHS_TARGET]}, + ) + assert defaulted.status_code == 201 + + tasks = (await client.get("/api/monitor/tasks")).json()["tasks"] + assert {t["platform"] for t in tasks} == {"xhs"} + + +class TestPlatformScoping: + async def _seed_two_platforms(self, client): + """One real XHS task plus a Douyin task inserted directly, since the API + refuses to create the latter.""" + await client.post( + "/api/monitor/tasks", + json={"name": "小红书任务", "mode": "creator", "targets": [XHS_TARGET]}, + ) + async with monitor_db.get_session() as session: + session.add( + MonitorTask( + name="抖音任务", platform="dy", mode="creator", enabled=True, + interval_minutes=60, max_notes_count=20, enable_comments=True, + max_comments_count=50, run_timeout_seconds=3600, + notify_enabled=False, created_at=0, updated_at=0, + ) + ) + + @pytest.mark.asyncio + async def test_tasks_are_filtered_by_platform(self, client): + await self._seed_two_platforms(client) + + all_tasks = (await client.get("/api/monitor/tasks")).json()["tasks"] + assert len(all_tasks) == 2 + + xhs_only = (await client.get("/api/monitor/tasks", params={"platform": "xhs"})).json() + assert [t["name"] for t in xhs_only["tasks"]] == ["小红书任务"] + + dy_only = (await client.get("/api/monitor/tasks", params={"platform": "dy"})).json() + assert [t["name"] for t in dy_only["tasks"]] == ["抖音任务"] + + @pytest.mark.asyncio + async def test_overview_is_scoped(self, client): + await self._seed_two_platforms(client) + + assert (await client.get("/api/monitor/overview")).json()["tasks"] == 2 + assert ( + await client.get("/api/monitor/overview", params={"platform": "xhs"}) + ).json()["tasks"] == 1 + + @pytest.mark.asyncio + async def test_a_platform_with_no_tasks_yields_empty_not_everything(self, client): + """An empty task set must not degrade into "no filter".""" + await self._seed_two_platforms(client) + + body = (await client.get("/api/monitor/notes", params={"platform": "bili"})).json() + assert body["notes"] == [] + + report = ( + await client.get("/api/monitor/report", params={"platform": "bili"}) + ).json() + assert report["totals"]["liked_count_delta"] == 0 + assert report["note_count"] == 0 + + +class TestPerPlatformSettings: + @pytest.mark.asyncio + async def test_each_platform_keeps_its_own_values(self, client): + await client.put( + "/api/settings", + params={"platform": "xhs"}, + json={"platform.xhs.crawl_sleep_sec": 3}, + ) + await client.put( + "/api/settings", + params={"platform": "dy"}, + json={"platform.dy.crawl_sleep_sec": 9}, + ) + + xhs = (await client.get("/api/settings", params={"platform": "xhs"})).json() + dy = (await client.get("/api/settings", params={"platform": "dy"})).json() + + assert xhs["values"]["platform.xhs.crawl_sleep_sec"] == 3 + assert dy["values"]["platform.dy.crawl_sleep_sec"] == 9 + + @pytest.mark.asyncio + async def test_system_settings_are_shared_across_platforms(self, client): + await client.put( + "/api/settings", + params={"platform": "xhs"}, + json={"system.active_hours_start": 8}, + ) + + dy = (await client.get("/api/settings", params={"platform": "dy"})).json() + assert dy["values"]["system.active_hours_start"] == 8 + # ...and the system specs are present in every platform's response. + assert "system.active_hours_end" in dy["values"] + + @pytest.mark.asyncio + async def test_a_key_for_another_platform_is_rejected(self, client): + """Writing xhs's key while scoped to dy would land somewhere unexpected.""" + response = await client.put( + "/api/settings", + params={"platform": "dy"}, + json={"platform.xhs.crawl_sleep_sec": 5}, + ) + assert response.status_code == 400 + + @pytest.mark.asyncio + async def test_cookies_are_per_platform(self, client): + await client.post( + "/api/monitor/cookie", + params={"platform": "xhs"}, + json={"cookie": "web_session=xhs-secret"}, + ) + + xhs = (await client.get("/api/monitor/cookie", params={"platform": "xhs"})).json() + dy = (await client.get("/api/monitor/cookie", params={"platform": "dy"})).json() + + assert xhs["present"] is True + assert dy["present"] is False + + # The old endpoint still defaults to Xiaohongshu. + assert (await client.get("/api/monitor/cookie")).json()["present"] is True + + +class TestLegacyKeyMigration: + @pytest.mark.asyncio + async def test_old_flat_keys_are_moved_to_the_new_namespace(self, tmp_path): + """Existing installs must not lose their cookie on upgrade.""" + monitor_db.set_sqlite_path(tmp_path / "monitor.db") + await monitor_db.init_db() + + async with monitor_db.get_engine().begin() as conn: + await conn.execute( + text( + "INSERT INTO monitor_setting (key, value, updated_at) " + "VALUES ('xhs_cookie', 'web_session=legacy', 1)" + ) + ) + await conn.execute( + text( + "INSERT INTO monitor_setting (key, value, updated_at) " + "VALUES ('wecom_webhook', 'https://qyapi.weixin.qq.com/x', 1)" + ) + ) + + # Re-running init performs the rename. + await monitor_db.init_db() + + async with monitor_db.get_engine().begin() as conn: + rows = dict( + (await conn.execute(text("SELECT key, value FROM monitor_setting"))).all() + ) + + assert rows.get("platform.xhs.cookie") == "web_session=legacy" + assert rows.get("system.wecom_webhook") == "https://qyapi.weixin.qq.com/x" + assert "xhs_cookie" not in rows + assert "wecom_webhook" not in rows + + await monitor_db.dispose_engine() + + @pytest.mark.asyncio + async def test_migration_is_idempotent_and_keeps_the_newer_value(self, tmp_path): + monitor_db.set_sqlite_path(tmp_path / "monitor.db") + await monitor_db.init_db() + + async with monitor_db.get_engine().begin() as conn: + await conn.execute( + text( + "INSERT INTO monitor_setting (key, value, updated_at) VALUES " + "('platform.xhs.cookie', 'current', 2), ('xhs_cookie', 'stale', 1)" + ) + ) + + await monitor_db.init_db() + + async with monitor_db.get_engine().begin() as conn: + rows = dict( + (await conn.execute(text("SELECT key, value FROM monitor_setting"))).all() + ) + + assert rows.get("platform.xhs.cookie") == "current" + assert "xhs_cookie" not in rows + + await monitor_db.dispose_engine() diff --git a/tests/test_settings.py b/tests/test_settings.py new file mode 100644 index 0000000..40c3eed --- /dev/null +++ b/tests/test_settings.py @@ -0,0 +1,374 @@ +# -*- coding: utf-8 -*- +# Copyright (c) 2025 relakkes@gmail.com +# +# This file is part of MediaCrawler project. +# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/tests/test_settings.py +# GitHub: https://github.com/NanmiCoder +# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1 +# +# 声明:本代码仅供学习和研究目的使用。使用者应遵守以下原则: +# 1. 不得用于任何商业用途。 +# 2. 使用时应遵守目标平台的使用条款和robots.txt规则。 +# 3. 不得进行大规模爬取或对平台造成运营干扰。 +# 4. 应合理控制请求频率,避免给目标平台带来不必要的负担。 +# 5. 不得用于任何非法或不当的用途。 +# +# 详细许可条款请参阅项目根目录下的LICENSE文件。 +# 使用本代码即表示您同意遵守上述原则和LICENSE中的所有条款。 + +"""Unified settings endpoint, and the effect its values actually have.""" + +from datetime import datetime + +import httpx +import pytest +import pytest_asyncio + +from api.main import app +from api.monitor import app_settings, db as monitor_db +from api.monitor import scheduler as scheduler_module +from api.monitor.scheduler import MonitorScheduler +from api.monitor.settings import get_setting + +SECRET_VALUE = "web_session=SUPERSECRET; a1=abc" + + +@pytest_asyncio.fixture +async def client(tmp_path): + monitor_db.set_sqlite_path(tmp_path / "monitor.db") + await monitor_db.init_db() + + transport = httpx.ASGITransport(app=app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as http_client: + yield http_client + + await monitor_db.dispose_engine() + + +class TestReadSettings: + @pytest.mark.asyncio + async def test_returns_values_secrets_and_the_spec(self, client): + body = (await client.get("/api/settings")).json() + + assert "values" in body and "secrets" in body and "specs" in body + # The spec drives the UI form, so every key must be described. + spec_keys = {spec["key"] for spec in body["specs"]} + assert "platform.xhs.default_interval_minutes" in spec_keys + assert "platform.xhs.enable_ip_proxy" in spec_keys + + @pytest.mark.asyncio + async def test_unset_values_fall_back_to_spec_defaults(self, client): + values = (await client.get("/api/settings")).json()["values"] + assert values["platform.xhs.default_interval_minutes"] == 360 + assert values["platform.xhs.enable_ip_proxy"] is False + + @pytest.mark.asyncio + async def test_secrets_are_masked_never_returned(self, client): + await client.put("/api/settings", json={"platform.xhs.cookie": SECRET_VALUE}) + + response = await client.get("/api/settings") + assert SECRET_VALUE not in response.text + + secret = response.json()["secrets"]["platform.xhs.cookie"] + assert secret["present"] is True + assert secret["length"] == len(SECRET_VALUE) + + +class TestUpdateSettings: + @pytest.mark.asyncio + async def test_partial_update_leaves_other_keys_alone(self, client): + await client.put( + "/api/settings", + json={"platform.xhs.default_interval_minutes": 120, "platform.xhs.cookie": SECRET_VALUE}, + ) + + # A form that only submits the interval must not blank the cookie. + await client.put("/api/settings", json={"platform.xhs.default_interval_minutes": 240}) + + body = (await client.get("/api/settings")).json() + assert body["values"]["platform.xhs.default_interval_minutes"] == 240 + assert body["secrets"]["platform.xhs.cookie"]["present"] is True + + @pytest.mark.asyncio + async def test_empty_string_clears_a_secret(self, client): + await client.put("/api/settings", json={"platform.xhs.cookie": SECRET_VALUE}) + await client.put("/api/settings", json={"platform.xhs.cookie": ""}) + + assert (await client.get("/api/settings")).json()["secrets"]["platform.xhs.cookie"][ + "present" + ] is False + + @pytest.mark.asyncio + async def test_unknown_key_is_rejected(self, client): + response = await client.put("/api/settings", json={"nope.not.a.setting": 1}) + assert response.status_code == 400 + + @pytest.mark.asyncio + async def test_out_of_range_is_rejected(self, client): + response = await client.put( + "/api/settings", json={"platform.xhs.default_interval_minutes": 1} + ) + assert response.status_code == 400 + + @pytest.mark.asyncio + async def test_invalid_choice_is_rejected(self, client): + response = await client.put("/api/settings", json={"platform.xhs.proxy_provider": "nonsense"}) + assert response.status_code == 400 + + @pytest.mark.asyncio + async def test_bools_accept_the_ui_shapes(self, client): + for raw in (True, "true", "1", "yes"): + response = await client.put("/api/settings", json={"platform.xhs.enable_ip_proxy": raw}) + assert response.status_code == 200 + assert (await client.get("/api/settings")).json()["values"][ + "platform.xhs.enable_ip_proxy" + ] is True + + @pytest.mark.asyncio + async def test_password_hash_cannot_be_written_through_this_endpoint(self, client): + """It has its own authenticated endpoint; this must not be a back door.""" + await client.put("/api/settings", json={"auth_password_hash": "pbkdf2_sha256$1$a$b"}) + + async with monitor_db.get_session() as session: + assert await get_setting(session, "auth_password_hash") is None + + +class TestSettingsActuallyTakeEffect: + @pytest.mark.asyncio + async def test_new_tasks_use_the_configured_defaults(self, client): + await client.put( + "/api/settings", + json={ + "platform.xhs.default_interval_minutes": 120, + "platform.xhs.default_max_notes": 7, + "platform.xhs.default_max_comments": 33, + }, + ) + + await client.post( + "/api/monitor/tasks", + json={"name": "用默认值", "mode": "creator", "targets": ["5f58bd990000000001003753"]}, + ) + + task = (await client.get("/api/monitor/tasks")).json()["tasks"][0] + assert task["interval_minutes"] == 120 + assert task["max_notes_count"] == 7 + assert task["max_comments_count"] == 33 + + @pytest.mark.asyncio + async def test_explicit_values_still_win_over_defaults(self, client): + await client.put("/api/settings", json={"platform.xhs.default_interval_minutes": 120}) + + await client.post( + "/api/monitor/tasks", + json={ + "name": "显式值", + "mode": "creator", + "interval_minutes": 720, + "targets": ["5f58bd990000000001003753"], + }, + ) + + task = (await client.get("/api/monitor/tasks")).json()["tasks"][0] + assert task["interval_minutes"] == 720 + + +class TestRunnerAppliesStrategy: + @pytest.mark.asyncio + async def test_strategy_settings_reach_the_command(self, client): + """Stored settings must actually change how the crawler is invoked.""" + from api.services.crawler_manager import CrawlerManager + from api.schemas import CrawlerStartRequest, PlatformEnum, CrawlerTypeEnum + + await client.put( + "/api/settings", + json={ + "platform.xhs.crawl_sleep_sec": 7, + "platform.xhs.enable_sub_comments": True, + "platform.xhs.enable_ip_proxy": True, + "platform.xhs.proxy_provider": "static", + "platform.xhs.proxy_pool_count": 5, + "platform.xhs.static_proxy_url": "http://127.0.0.1:8888", + }, + ) + + async with monitor_db.get_session() as session: + strategy = await scheduler_module.app_settings.get_value( + session, "crawl_sleep_sec", "xhs", 2 + ) + assert strategy == 7 + + # And the flag builder forwards them when present. + command = CrawlerManager()._build_command( + CrawlerStartRequest( + platform=PlatformEnum.XHS, + crawler_type=CrawlerTypeEnum.CREATOR, + creator_ids="abc", + crawler_max_sleep_sec=7, + enable_ip_proxy=True, + ip_proxy_provider_name="static", + ip_proxy_pool_count=5, + static_proxy_url="http://127.0.0.1:8888", + ) + ) + joined = " ".join(command) + assert "--crawler_max_sleep_sec 7" in joined + assert "--enable_ip_proxy true" in joined + assert "--ip_proxy_provider_name static" in joined + assert "--static_proxy_url http://127.0.0.1:8888" in joined + + +def _frozen_clock(hour: int): + """Stand-in for the datetime class whose now() is pinned to a given hour. + + Testing an hour window by sleeping is not an option; patching the class the + scheduler imported is the whole mechanism. + """ + + class _Frozen: + @staticmethod + def now(tz=None): + return datetime(2026, 1, 1, hour) + + return _Frozen + + +class TestManualCrawlCookieFallback: + """The crawl page no longer has its own paste box; it reuses Settings.""" + + @pytest_asyncio.fixture + async def captured(self, monkeypatch): + # api.services re-exports the singleton instance, not the module. + from api.services import crawler_manager + + seen: dict = {} + + async def _fake_start(request, extra_args=None): + seen["cookies"] = request.cookies + return True + + monkeypatch.setattr(crawler_manager, "start", _fake_start) + return seen + + @pytest.mark.asyncio + async def test_falls_back_to_the_stored_cookie(self, client, captured): + await client.put( + "/api/settings", json={"platform.xhs.cookie": "web_session=stored"} + ) + + response = await client.post( + "/api/crawler/start", + json={ + "platform": "xhs", + "login_type": "cookie", + "crawler_type": "creator", + "creator_ids": "abc", + }, + ) + + assert response.status_code == 200 + assert captured["cookies"] == "web_session=stored" + + @pytest.mark.asyncio + async def test_an_explicit_cookie_still_wins(self, client, captured): + await client.put( + "/api/settings", json={"platform.xhs.cookie": "web_session=stored"} + ) + + await client.post( + "/api/crawler/start", + json={ + "platform": "xhs", + "login_type": "cookie", + "crawler_type": "creator", + "creator_ids": "abc", + "cookies": "web_session=explicit", + }, + ) + + assert captured["cookies"] == "web_session=explicit" + + @pytest.mark.asyncio + async def test_missing_cookie_is_a_clear_error_not_a_silent_failure( + self, client, captured + ): + """Better a 400 that names the fix than a run that fetches nothing.""" + response = await client.post( + "/api/crawler/start", + json={ + "platform": "xhs", + "login_type": "cookie", + "crawler_type": "creator", + "creator_ids": "abc", + }, + ) + + assert response.status_code == 400 + assert "设置" in response.json()["detail"] + assert "cookies" not in captured + + @pytest.mark.asyncio + async def test_the_cookie_is_read_per_platform(self, client, captured): + await client.put( + "/api/settings", + params={"platform": "xhs"}, + json={"platform.xhs.cookie": "web_session=xhs-only"}, + ) + + # Douyin has no stored cookie, so it must not borrow Xiaohongshu's. + response = await client.post( + "/api/crawler/start", + json={ + "platform": "dy", + "login_type": "cookie", + "crawler_type": "creator", + "creator_ids": "abc", + }, + ) + + assert response.status_code == 400 + + +class TestActiveHours: + """The window gate lives in the scheduler, not the crawler.""" + + @pytest.mark.asyncio + async def test_inside_a_daytime_window(self, client, monkeypatch): + await client.put( + "/api/settings", + json={"system.active_hours_start": 8, "system.active_hours_end": 22}, + ) + monkeypatch.setattr(scheduler_module, "datetime", _frozen_clock(12)) + + async with monitor_db.get_session() as session: + assert await MonitorScheduler()._within_active_hours(session) is True + + @pytest.mark.asyncio + async def test_outside_a_daytime_window(self, client, monkeypatch): + await client.put( + "/api/settings", + json={"system.active_hours_start": 8, "system.active_hours_end": 22}, + ) + monkeypatch.setattr(scheduler_module, "datetime", _frozen_clock(3)) + + async with monitor_db.get_session() as session: + assert await MonitorScheduler()._within_active_hours(session) is False + + @pytest.mark.asyncio + async def test_window_wrapping_past_midnight(self, client, monkeypatch): + await client.put( + "/api/settings", + json={"system.active_hours_start": 22, "system.active_hours_end": 6}, + ) + + for hour, expected in ((23, True), (3, True), (12, False)): + monkeypatch.setattr(scheduler_module, "datetime", _frozen_clock(hour)) + async with monitor_db.get_session() as session: + assert await MonitorScheduler()._within_active_hours(session) is expected + + @pytest.mark.asyncio + async def test_default_window_covers_the_whole_day(self, client, monkeypatch): + for hour in (0, 12, 23): + monkeypatch.setattr(scheduler_module, "datetime", _frozen_clock(hour)) + async with monitor_db.get_session() as session: + assert await MonitorScheduler()._within_active_hours(session) is True diff --git a/webui/index.html b/webui/index.html index 722c324..4b25344 100644 --- a/webui/index.html +++ b/webui/index.html @@ -4,7 +4,7 @@ - MediaCrawler - Command Center + 综合采集平台 diff --git a/webui/public/logos/bilibili_logo.png b/webui/public/logos/bilibili_logo.png deleted file mode 100644 index 7fce00a..0000000 Binary files a/webui/public/logos/bilibili_logo.png and /dev/null differ diff --git a/webui/public/logos/douyin.png b/webui/public/logos/douyin.png deleted file mode 100644 index 1921765..0000000 Binary files a/webui/public/logos/douyin.png and /dev/null differ diff --git a/webui/public/logos/github.png b/webui/public/logos/github.png deleted file mode 100644 index 8f5cef5..0000000 Binary files a/webui/public/logos/github.png and /dev/null differ diff --git a/webui/public/logos/my_logo.png b/webui/public/logos/my_logo.png deleted file mode 100644 index 8105593..0000000 Binary files a/webui/public/logos/my_logo.png and /dev/null differ diff --git a/webui/public/logos/xiaohongshu_logo.png b/webui/public/logos/xiaohongshu_logo.png deleted file mode 100644 index 5ccfe5d..0000000 Binary files a/webui/public/logos/xiaohongshu_logo.png and /dev/null differ diff --git a/webui/src/App.tsx b/webui/src/App.tsx index 27dd38a..871195f 100644 --- a/webui/src/App.tsx +++ b/webui/src/App.tsx @@ -1,62 +1,100 @@ -import { useState } from 'react' +import { useEffect, useState } from 'react' import { Toaster } from 'sonner' +import { Loader2 } from 'lucide-react' import { Sidebar } from '@/components/layout/Sidebar' import { MainContent } from '@/components/layout/MainContent' -import { AuthorFooter } from '@/components/layout/AuthorFooter' import { CrawlerConfigPanel } from '@/components/config/CrawlerConfigPanel' +import { MonitorDashboard } from '@/components/monitor/MonitorDashboard' +import { ReportView } from '@/components/monitor/ReportView' +import { SettingsView } from '@/components/settings/SettingsView' +import { Login } from '@/components/auth/Login' import { EnvironmentCheck, isEnvChecked } from '@/components/env/EnvironmentCheck' -import { LicenseDisclaimer, isLicenseAccepted } from '@/components/license/LicenseDisclaimer' +import { authApi, setUnauthorizedHandler } from '@/lib/api' + +export type AppView = 'crawler' | 'monitor' | 'report' | 'settings' function App() { - // Initialize by checking localStorage if license has been accepted - const [licenseAccepted, setLicenseAccepted] = useState(() => isLicenseAccepted()) + // null = still probing. Rendering the app while unknown would briefly mount + // the log WebSocket before we know whether the user is authenticated. + const [authed, setAuthed] = useState(null) // Initialize by checking localStorage if env check has passed const [envChecked, setEnvChecked] = useState(() => isEnvChecked()) - // State for showing disclaimer manually - const [showDisclaimer, setShowDisclaimer] = useState(false) + // Which top-level workspace is visible. Only one is mounted at a time: the + // crawler view owns a live log WebSocket and a 2s status poll, which should + // not keep running while the user is looking at the monitor dashboard. + const [view, setView] = useState('crawler') + + useEffect(() => { + let cancelled = false + authApi + .me() + .then(() => !cancelled && setAuthed(true)) + .catch(() => !cancelled && setAuthed(false)) + return () => { + cancelled = true + } + }, []) + + // A 401 from any request means the session expired or was revoked elsewhere + // (e.g. the password was changed on another device), so drop back to login. + useEffect(() => { + setUnauthorizedHandler(() => setAuthed(false)) + return () => setUnauthorizedHandler(null) + }, []) const handleEnvCheckComplete = () => { setEnvChecked(true) } - const handleLicenseAccept = () => { - setLicenseAccepted(true) - setShowDisclaimer(false) + const handleLogout = async () => { + try { + await authApi.logout() + } finally { + // Even if the call fails, the local session is over; unmounting the tree + // closes the log WebSocket via the existing connection-count cleanup. + setAuthed(false) + } } - const handleShowDisclaimer = () => { - setShowDisclaimer(true) + if (authed === null) { + return ( +

+ +
+ ) + } + + if (!authed) { + return setAuthed(true)} /> } return (
- {/* License Disclaimer Modal - Shows first or when triggered */} - {(!licenseAccepted || showDisclaimer) && ( - - )} - - {/* Environment Check Modal - Shows after license accepted */} - {licenseAccepted && !showDisclaimer && !envChecked && ( - - )} + {/* Environment Check Modal - shown until the check passes or is skipped. + A configuration concern, so it lives inside the authenticated tree. */} + {!envChecked && } {/* Header Bar */} - + {/* Main Area */}
- {/* Config Panel - Primary Action Area (Always Expanded) */} -
- -
+ {view === 'crawler' && ( + <> + {/* Config Panel - Primary Action Area (Always Expanded) */} +
+ +
- {/* Console - Collapsible Terminal */} - + {/* Console - Collapsible Terminal */} + + + )} + {view === 'monitor' && } + {view === 'report' && } + {view === 'settings' && }
- {/* Author Footer */} - - {/* Toast notifications - Theme-aware style */} void +} + +/** + * Full-screen login gate. + * + * Never an overlay on top of the app: the authenticated tree owns the log + * WebSocket, and mounting it behind a modal would open the socket before the + * user is authenticated. App.tsx therefore renders this *instead of* the app. + */ +export function Login({ onSuccess }: LoginProps) { + const [password, setPassword] = useState('') + const [error, setError] = useState('') + const [busy, setBusy] = useState(false) + + const handleSubmit = async (event: React.FormEvent) => { + event.preventDefault() + if (!password || busy) return + + setBusy(true) + setError('') + try { + await authApi.login(password) + setPassword('') + onSuccess() + } catch (err: unknown) { + const response = (err as { response?: { status?: number; data?: { detail?: string } } }) + ?.response + if (response?.status === 429) { + // Throttled, not wrong -- say so rather than implying a bad password. + setError(response.data?.detail ?? '尝试过于频繁,请稍后再试') + } else { + setError(response?.data?.detail ?? '登录失败') + } + } finally { + setBusy(false) + } + } + + return ( +
+
+ {/* Corner accents, matching the other full-screen gates */} +
+
+
+
+ +
+ + + 综合采集平台 + +
+

+ 需要登录才能访问控制面板 +

+ +
+ + setPassword(event.target.value)} + placeholder="请输入密码" + className="h-10 font-mono" + /> +
+ + {error && ( +
+ + {error} +
+ )} + + + +

+ 首次启动的密码打印在服务端启动日志里。 + 忘记密码时,设置环境变量 MC_PASSWORD 后重启即可恢复, + 并在登录后于设置页修改。 +

+ +
+ ) +} diff --git a/webui/src/components/config/CrawlerConfigPanel.tsx b/webui/src/components/config/CrawlerConfigPanel.tsx index b09762a..5b9c8c3 100644 --- a/webui/src/components/config/CrawlerConfigPanel.tsx +++ b/webui/src/components/config/CrawlerConfigPanel.tsx @@ -1,5 +1,5 @@ import type { ComponentType, ReactNode, KeyboardEvent } from 'react' -import { useState } from 'react' +import { useEffect, useState } from 'react' import { useTranslation } from 'react-i18next' import { Database, Globe, Image as ImageIcon, KeyRound, MessageSquare, Play, Square, X } from 'lucide-react' import { Label } from '@/components/ui/label' @@ -8,7 +8,9 @@ import { Input } from '@/components/ui/input' import { Checkbox } from '@/components/ui/checkbox' import { Button } from '@/components/ui/button' import { useCrawlerStore } from '@/store/crawlerStore' -import { usePlatforms, useConfigOptions, useStartCrawler, useStopCrawler } from '@/hooks/useCrawler' +import { useConfigOptions, useStartCrawler, useStopCrawler } from '@/hooks/useCrawler' +import { useCurrentPlatform } from '@/hooks/usePlatform' +import { useCookieStatus } from '@/hooks/useMonitor' import { ParsedIdList } from './ParsedIdList' type SectionProps = { @@ -137,9 +139,19 @@ export function CrawlerConfigPanel() { const updateConfig = useCrawlerStore((state) => state.updateConfig) const status = useCrawlerStore((state) => state.status) - const { data: platforms } = usePlatforms() const { data: options } = useConfigOptions() const { mutate: startCrawler, isPending: isStarting } = useStartCrawler() + + // Follow the global platform selection rather than keeping its own copy. + const { platform: hostPlatform, capability: hostCapability } = useCurrentPlatform() + const { data: cookieStatus } = useCookieStatus() + const storedCookieOk = cookieStatus?.present ?? false + + useEffect(() => { + if (config.platform !== hostPlatform) { + updateConfig({ platform: hostPlatform }) + } + }, [hostPlatform, config.platform, updateConfig]) const { mutate: stopCrawler, isPending: isStopping } = useStopCrawler() const isDisabled = status === 'running' || status === 'stopping' @@ -165,22 +177,15 @@ export function CrawlerConfigPanel() { icon={Globe} > - + {/* Driven by the global switcher in the header. A second platform + control here would be a second source of truth for the same + decision, and the two would drift. */} +
+ + {hostCapability?.label ?? hostPlatform} + + 由右上角切换 +
@@ -300,14 +305,23 @@ export function CrawlerConfigPanel() { {config.login_type === 'cookie' ? ( - -