diff --git a/README.md b/README.md index 5131642..e594767 100644 --- a/README.md +++ b/README.md @@ -61,6 +61,18 @@ Or with `uv`: uv sync ``` +## Smoke Test + +A smoke test is a quick "does it start and answer basic requests?" check. It does +not replace full tests, but it catches broken imports, routes, and JSON responses. + +```bash +python scripts/smoke_test.py +``` + +The smoke test starts the web server on `127.0.0.1:18080` and disables automatic +continuous scraping for that process. + ## Running ### Web UI @@ -90,6 +102,8 @@ The current web panel includes: - scrape/export/media actions - shared `scrape_media` toggle - local message viewer powered by SQLite + media files +- API docs at `/swagger` and `/openapi.json` +- health checks at `/health` and `/health/continuous` ### CLI diff --git a/app_state.py b/app_state.py new file mode 100644 index 0000000..e89fc10 --- /dev/null +++ b/app_state.py @@ -0,0 +1,84 @@ +import json +import threading +from pathlib import Path +from typing import Any, Callable, Dict + + +DEFAULT_STATE: Dict[str, Any] = { + "api_id": None, + "api_hash": None, + "channels": {}, + "channel_names": {}, + "scrape_media": True, + "forwarding_rules": [], + "continuous_scraping": { + "enabled": True, + "interval_minutes": 1, + "channels": [], + "run_all_tracked": True, + }, +} + + +class StateStore: + def __init__(self, path: Path): + self.path = path + self.lock = threading.RLock() + + def load(self) -> Dict[str, Any]: + with self.lock: + if not self.path.exists(): + return self._default_state() + try: + with self.path.open("r", encoding="utf-8") as handle: + state = json.load(handle) + except (json.JSONDecodeError, OSError): + return self._default_state() + return self._merge_defaults(state) + + def save(self, state: Dict[str, Any]) -> None: + with self.lock: + self.path.parent.mkdir(parents=True, exist_ok=True) + tmp_path = self.path.with_suffix(self.path.suffix + ".tmp") + with tmp_path.open("w", encoding="utf-8") as handle: + json.dump(self._merge_defaults(state), handle, ensure_ascii=False, indent=2) + handle.write("\n") + tmp_path.replace(self.path) + + def update(self, mutator: Callable[[Dict[str, Any]], None]) -> Dict[str, Any]: + with self.lock: + state = self.load() + mutator(state) + self.save(state) + return state + + def continuous_config(self) -> Dict[str, Any]: + state = self.load() + return dict(state.get("continuous_scraping") or DEFAULT_STATE["continuous_scraping"]) + + def save_continuous_config(self, config: Dict[str, Any]) -> Dict[str, Any]: + def mutate(state: Dict[str, Any]) -> None: + state["continuous_scraping"] = { + "enabled": bool(config.get("enabled", True)), + "interval_minutes": max(1, int(config.get("interval_minutes", 1) or 1)), + "channels": [ + str(item).strip() + for item in config.get("channels", []) + if str(item).strip() + ], + "run_all_tracked": bool(config.get("run_all_tracked", True)), + } + + return self.update(mutate)["continuous_scraping"] + + def _default_state(self) -> Dict[str, Any]: + return json.loads(json.dumps(DEFAULT_STATE)) + + def _merge_defaults(self, state: Dict[str, Any]) -> Dict[str, Any]: + merged = self._default_state() + for key, value in state.items(): + if key == "continuous_scraping" and isinstance(value, dict): + merged[key].update(value) + else: + merged[key] = value + return merged diff --git a/health.py b/health.py new file mode 100644 index 0000000..f7a1ba8 --- /dev/null +++ b/health.py @@ -0,0 +1,73 @@ +import sqlite3 +from pathlib import Path +from typing import Any, Dict + +from app_state import StateStore + + +def health_payload( + data_dir: Path, + session_dir: Path, + state_store: StateStore, + continuous_snapshot: Dict[str, Any], + job_queue_size: int, +) -> Dict[str, Any]: + checks = { + "data_dir": _dir_check(data_dir, writable=True), + "session_dir": _dir_check(session_dir, writable=True), + "state_file": _state_check(state_store), + "sqlite": _sqlite_check(), + "continuous": _continuous_check(continuous_snapshot), + "job_queue": {"ok": True, "size": job_queue_size}, + } + ok = all(item.get("ok", False) for item in checks.values()) + return {"ok": ok, "status": "ok" if ok else "degraded", "checks": checks} + + +def _dir_check(path: Path, writable: bool = False) -> Dict[str, Any]: + path.mkdir(parents=True, exist_ok=True) + ok = path.exists() and path.is_dir() + payload: Dict[str, Any] = {"ok": ok, "path": str(path)} + if writable: + probe = path / ".healthcheck" + try: + probe.write_text("ok", encoding="utf-8") + probe.unlink(missing_ok=True) + payload["writable"] = True + except OSError as exc: + payload.update({"ok": False, "writable": False, "error": str(exc)}) + return payload + + +def _state_check(state_store: StateStore) -> Dict[str, Any]: + try: + state = state_store.load() + return { + "ok": True, + "path": str(state_store.path), + "has_api_credentials": bool(state.get("api_id") and state.get("api_hash")), + "tracked_channels": len(state.get("channels", {})), + } + except Exception as exc: + return {"ok": False, "path": str(state_store.path), "error": str(exc)} + + +def _sqlite_check() -> Dict[str, Any]: + try: + conn = sqlite3.connect(":memory:") + conn.execute("SELECT 1") + conn.close() + return {"ok": True} + except sqlite3.Error as exc: + return {"ok": False, "error": str(exc)} + + +def _continuous_check(snapshot: Dict[str, Any]) -> Dict[str, Any]: + status = snapshot.get("status", {}) + config = snapshot.get("config", {}) + return { + "ok": True, + "enabled": bool(config.get("enabled")), + "running": bool(status.get("running")), + "last_error": status.get("last_error"), + } diff --git a/pyproject.toml b/pyproject.toml index 2f9d252..55d89ab 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -8,7 +8,6 @@ dependencies = [ "aiohappyeyeballs==2.6.1", "aiohttp==3.12.14", "aiosignal==1.4.0", - "asyncio==3.4.3", "attrs==25.3.0", "frozenlist==1.7.0", "idna==3.10", @@ -18,6 +17,6 @@ dependencies = [ "pyasn1==0.6.1", "qrcode==8.0", "rsa==4.9.1", - "telethon==1.40.0", + "Telethon==1.40.0", "yarl==1.20.1", ] diff --git a/requirements.txt b/requirements.txt index ee961d7..c91dc6c 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,8 +1,7 @@ aiohappyeyeballs==2.6.1 aiohttp==3.12.14 aiosignal==1.4.0 -asyncio==3.4.3 -attrs==25.3.0 +attrs==25.3.0 frozenlist==1.7.0 idna==3.10 multidict==6.6.3 diff --git a/scraper_jobs.py b/scraper_jobs.py new file mode 100644 index 0000000..1ebd546 --- /dev/null +++ b/scraper_jobs.py @@ -0,0 +1,70 @@ +import asyncio +from typing import Any, Dict, List + +from app_state import StateStore + + +class ScraperJobService: + def __init__(self, state_store: StateStore): + self.state_store = state_store + + def run(self, job_type: str, payload: Dict[str, Any]) -> None: + if job_type == "set_scrape_media": + value = bool(payload["value"]) + + def mutate(state: Dict[str, Any]) -> None: + state["scrape_media"] = value + + self.state_store.update(mutate) + print(f"Media scraping set to {value}") + return + + asyncio.run(self._run_async(job_type, payload)) + + async def _run_async(self, job_type: str, payload: Dict[str, Any]) -> None: + ScraperClass = self._import_scraper_class() + scraper = ScraperClass() + initialized = await scraper.initialize_client(interactive=False) + if not initialized: + raise RuntimeError( + "Telegram client is not ready. Check credentials, session, and write access to /app/session." + ) + try: + if job_type == "scrape_channel": + await self._scrape_channels(scraper, [payload["channel_id"]]) + elif job_type == "scrape_all": + await self._scrape_channels( + scraper, list(scraper.state.get("channels", {}).keys()) + ) + elif job_type == "scrape_selected": + await self._scrape_channels( + scraper, [str(channel_id) for channel_id in payload.get("channels", [])] + ) + elif job_type == "export_all": + await scraper.export_data() + elif job_type == "export_channel": + channel_id = payload["channel_id"] + scraper.export_to_csv(channel_id) + scraper.export_to_json(channel_id) + elif job_type == "rescrape_media": + await scraper.rescrape_media(payload["channel_id"]) + elif job_type == "fix_missing_media": + await scraper.fix_missing_media(payload["channel_id"]) + elif job_type == "refresh_dialogs": + await scraper.list_channels() + else: + raise RuntimeError(f"Unsupported job type: {job_type}") + finally: + scraper.close_db_connections() + if scraper.client: + await scraper.client.disconnect() + + async def _scrape_channels(self, scraper, channels: List[str]) -> None: + for channel_id in channels: + offset = int(scraper.state.get("channels", {}).get(channel_id, 0) or 0) + await scraper.scrape_channel(channel_id, offset) + + def _import_scraper_class(self): + from telegram_scraper_with_forwarding import OptimizedTelegramScraper + + return OptimizedTelegramScraper diff --git a/scripts/smoke_test.py b/scripts/smoke_test.py new file mode 100644 index 0000000..33e2333 --- /dev/null +++ b/scripts/smoke_test.py @@ -0,0 +1,81 @@ +import json +import os +import subprocess +import sys +import time +import urllib.request + + +PORT = os.environ.get("TELEGRAM_SCRAPER_SMOKE_PORT", "18080") +BASE_URL = f"http://127.0.0.1:{PORT}" +ENDPOINTS = [ + "/", + "/viewer", + "/swagger", + "/openapi.json", + "/health", + "/health/continuous", + "/api/dashboard", + "/api/continuous", + "/api/channels", +] + + +def fetch(path: str) -> tuple[int, bytes]: + with urllib.request.urlopen(BASE_URL + path, timeout=3) as response: + return response.status, response.read() + + +def wait_for_server(process: subprocess.Popen) -> None: + deadline = time.time() + 15 + last_error = None + while time.time() < deadline: + if process.poll() is not None: + raise RuntimeError(f"server exited with code {process.returncode}") + try: + fetch("/health") + return + except Exception as exc: + last_error = exc + time.sleep(0.3) + raise RuntimeError(f"server did not become ready: {last_error}") + + +def main() -> int: + env = { + **os.environ, + "TELEGRAM_SCRAPER_PORT": PORT, + "TELEGRAM_SCRAPER_START_CONTINUOUS": "0", + } + process = subprocess.Popen( + [sys.executable, "main.py"], + env=env, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + ) + try: + wait_for_server(process) + for endpoint in ENDPOINTS: + status, body = fetch(endpoint) + if status != 200: + raise RuntimeError(f"{endpoint} returned HTTP {status}") + if endpoint.endswith(".json") or endpoint.startswith("/api") or endpoint.startswith("/health"): + json.loads(body.decode("utf-8")) + print(f"ok {endpoint}") + return 0 + except Exception as exc: + print(f"smoke test failed: {exc}", file=sys.stderr) + if process.stdout: + print(process.stdout.read(), file=sys.stderr) + return 1 + finally: + process.terminate() + try: + process.wait(timeout=5) + except subprocess.TimeoutExpired: + process.kill() + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/telegram_scraper_with_forwarding.py b/telegram_scraper_with_forwarding.py index dbf7240..9754594 100644 --- a/telegram_scraper_with_forwarding.py +++ b/telegram_scraper_with_forwarding.py @@ -15,6 +15,7 @@ from telethon import TelegramClient, events from telethon.tl.types import MessageMediaPhoto, MessageMediaDocument, MessageMediaWebPage, User, PeerChannel, Channel, Chat from telethon.errors import FloodWaitError, SessionPasswordNeededError import qrcode +from app_state import StateStore warnings.filterwarnings("ignore", message="Using async sessions support is an experimental feature") @@ -66,6 +67,7 @@ class OptimizedTelegramScraper: self.DATA_DIR.mkdir(exist_ok=True) self.SESSION_DIR.mkdir(exist_ok=True) self.STATE_FILE = str(self.DATA_DIR / 'state.json') + self.state_store = StateStore(self.DATA_DIR / 'state.json') self.state = self.load_state() self.client = None self.continuous_scraping_active = False @@ -77,25 +79,11 @@ class OptimizedTelegramScraper: self.forwarding_handler = None def load_state(self) -> Dict[str, Any]: - if os.path.exists(self.STATE_FILE): - try: - with open(self.STATE_FILE, 'r') as f: - return json.load(f) - except: - pass - return { - 'api_id': None, - 'api_hash': None, - 'channels': {}, - 'channel_names': {}, - 'scrape_media': True, - 'forwarding_rules': [], - } + return self.state_store.load() def save_state(self): try: - with open(self.STATE_FILE, 'w') as f: - json.dump(self.state, f, indent=2) + self.state_store.save(self.state) except Exception as e: print(f"Failed to save state: {e}") diff --git a/webui/app.js b/webui/app.js index b18f67b..8af9b4d 100644 --- a/webui/app.js +++ b/webui/app.js @@ -54,6 +54,10 @@ function setBusy(button, busy) { button.disabled = busy; } +function confirmAction(message) { + return window.confirm(message); +} + function renderAuth(auth) { document.getElementById("auth-status").textContent = auth.status; document.getElementById("auth-details").textContent = auth.details || ""; @@ -133,6 +137,7 @@ function renderChannels(channels) { }); node.querySelector(".remove-btn").addEventListener("click", async () => { + if (!confirmAction(`Remove ${channel.name} from tracked channels?`)) return; await api("/api/channels/remove", { method: "POST", body: JSON.stringify({ channel_id: channel.channel_id }), @@ -321,6 +326,7 @@ async function main() { }); document.getElementById("scrape-all-btn").addEventListener("click", async (event) => { + if (!confirmAction("Queue scraping for all tracked channels?")) return; setBusy(event.currentTarget, true); try { await api("/api/jobs/scrape", { method: "POST", body: JSON.stringify({}) }); @@ -331,6 +337,7 @@ async function main() { }); document.getElementById("export-all-btn").addEventListener("click", async (event) => { + if (!confirmAction("Queue export for all tracked channels?")) return; setBusy(event.currentTarget, true); try { await api("/api/jobs/export", { method: "POST", body: JSON.stringify({}) }); @@ -439,6 +446,9 @@ async function main() { const channels = Array.from(document.querySelectorAll(".continuous-channel-checkbox:checked")).map( (item) => item.value, ); + if (!runAllTracked && channels.length === 0 && enabled) { + if (!confirmAction("Continuous scraping is enabled with no selected channels. Save anyway?")) return; + } await api("/api/continuous", { method: "POST", body: JSON.stringify({ diff --git a/webui/index.html b/webui/index.html index 9385b14..f5a2f49 100644 --- a/webui/index.html +++ b/webui/index.html @@ -19,6 +19,7 @@ Dashboard Message Viewer API Docs + Health
diff --git a/webui/swagger.html b/webui/swagger.html index 558c0fe..9c6c3cd 100644 --- a/webui/swagger.html +++ b/webui/swagger.html @@ -18,6 +18,7 @@ diff --git a/webui_server.py b/webui_server.py index 3266148..fae401b 100644 --- a/webui_server.py +++ b/webui_server.py @@ -21,6 +21,9 @@ from typing import Any, Dict, List, Optional import qrcode import qrcode.image.svg +from app_state import StateStore +from health import health_payload +from scraper_jobs import ScraperJobService from telethon import TelegramClient from telethon.errors import SessionPasswordNeededError @@ -31,8 +34,11 @@ WEBUI_DIR = BASE_DIR / "webui" STATE_FILE = DATA_DIR / "state.json" DEFAULT_HOST = os.environ.get("TELEGRAM_SCRAPER_HOST", "0.0.0.0") DEFAULT_PORT = int(os.environ.get("TELEGRAM_SCRAPER_PORT", "8080")) +START_CONTINUOUS = os.environ.get("TELEGRAM_SCRAPER_START_CONTINUOUS", "1") != "0" SESSION_DIR = BASE_DIR / "session" SESSION_DIR.mkdir(exist_ok=True) +STATE_STORE = StateStore(STATE_FILE) +SCRAPER_JOBS = ScraperJobService(STATE_STORE) def utc_now_iso() -> str: @@ -40,23 +46,11 @@ def utc_now_iso() -> str: def load_state() -> Dict[str, Any]: - if STATE_FILE.exists(): - with STATE_FILE.open("r", encoding="utf-8") as handle: - return json.load(handle) - return { - "api_id": None, - "api_hash": None, - "channels": {}, - "channel_names": {}, - "scrape_media": True, - "forwarding_rules": [], - } + return STATE_STORE.load() def save_state(state: Dict[str, Any]) -> None: - DATA_DIR.mkdir(parents=True, exist_ok=True) - with STATE_FILE.open("w", encoding="utf-8") as handle: - json.dump(state, handle, ensure_ascii=False, indent=2) + STATE_STORE.save(state) def channel_db_path(channel_id: str) -> Path: @@ -496,12 +490,7 @@ class ContinuousScrapeManager: self.lock = threading.RLock() self.thread: Optional[threading.Thread] = None self.stop_event = threading.Event() - self.config: Dict[str, Any] = { - "enabled": True, - "interval_minutes": 1, - "channels": [], - "run_all_tracked": True, - } + self.config: Dict[str, Any] = STATE_STORE.continuous_config() self.status: Dict[str, Any] = { "running": False, "last_started_at": None, @@ -546,12 +535,14 @@ class ContinuousScrapeManager: ) -> Dict[str, Any]: interval_minutes = max(1, int(interval_minutes)) with self.lock: - self.config = { - "enabled": bool(enabled), - "interval_minutes": interval_minutes, - "channels": channels, - "run_all_tracked": bool(run_all_tracked), - } + self.config = STATE_STORE.save_continuous_config( + { + "enabled": bool(enabled), + "interval_minutes": interval_minutes, + "channels": channels, + "run_all_tracked": bool(run_all_tracked), + } + ) if enabled: self.start() else: @@ -641,63 +632,7 @@ def import_scraper_class(): def run_job(job_type: str, payload: Dict[str, Any]) -> None: - if job_type == "set_scrape_media": - state = load_state() - state["scrape_media"] = bool(payload["value"]) - save_state(state) - print(f"Media scraping set to {state['scrape_media']}") - return - - async def _async_job() -> None: - ScraperClass = import_scraper_class() - scraper = ScraperClass() - initialized = await scraper.initialize_client(interactive=False) - if not initialized: - raise RuntimeError( - "Telegram client is not ready. Check credentials, session, and write access to /app/session." - ) - try: - if job_type == "scrape_channel": - channel_id = payload["channel_id"] - state = scraper.load_state() - offset = int(state.get("channels", {}).get(channel_id, 0) or 0) - await scraper.scrape_channel(channel_id, offset) - elif job_type == "scrape_all": - channels = list(scraper.state.get("channels", {}).keys()) - for channel_id in channels: - offset = int( - scraper.state.get("channels", {}).get(channel_id, 0) or 0 - ) - await scraper.scrape_channel(channel_id, offset) - elif job_type == "scrape_selected": - channels = [ - str(channel_id) for channel_id in payload.get("channels", []) - ] - for channel_id in channels: - offset = int( - scraper.state.get("channels", {}).get(channel_id, 0) or 0 - ) - await scraper.scrape_channel(channel_id, offset) - elif job_type == "export_all": - await scraper.export_data() - elif job_type == "export_channel": - channel_id = payload["channel_id"] - scraper.export_to_csv(channel_id) - scraper.export_to_json(channel_id) - elif job_type == "rescrape_media": - await scraper.rescrape_media(payload["channel_id"]) - elif job_type == "fix_missing_media": - await scraper.fix_missing_media(payload["channel_id"]) - elif job_type == "refresh_dialogs": - await scraper.list_channels() - else: - raise RuntimeError(f"Unsupported job type: {job_type}") - finally: - scraper.close_db_connections() - if scraper.client: - await scraper.client.disconnect() - - asyncio.run(_async_job()) + SCRAPER_JOBS.run(job_type, payload) def auth_status() -> Dict[str, Any]: @@ -882,6 +817,13 @@ def openapi_payload() -> Dict[str, Any]: "responses": json_response, } }, + "/health": { + "get": { + "summary": "Application health", + "description": "Checks data/session write access, state loading, SQLite, continuous status, and job queue size.", + "responses": json_response, + } + }, "/api/channels": { "get": { "summary": "List tracked channels", @@ -1066,6 +1008,16 @@ class TelegramScraperRequestHandler(BaseHTTPRequestHandler): return self.send_json(self.app.continuous_manager.snapshot()) if path == "/health/continuous": return self.send_json(self.app.continuous_manager.snapshot()) + if path == "/health": + return self.send_json( + health_payload( + DATA_DIR, + SESSION_DIR, + STATE_STORE, + self.app.continuous_manager.snapshot(), + self.app.job_runner.queue.qsize(), + ) + ) if path == "/api/jobs": return self.send_json(self.app.job_runner.recent_jobs()) if path.startswith("/api/jobs/"): @@ -1131,6 +1083,7 @@ class TelegramScraperRequestHandler(BaseHTTPRequestHandler): "/api/continuous", "/api/jobs", "/api/channels", + "/health", "/health/continuous", "/openapi.json", } @@ -1397,7 +1350,8 @@ class TelegramScraperWebServer(ThreadingHTTPServer): self.job_runner = JobRunner() self.auth_manager = TelegramAuthManager() self.continuous_manager = ContinuousScrapeManager() - self.continuous_manager.start() + if START_CONTINUOUS and self.continuous_manager.snapshot()["config"].get("enabled", True): + self.continuous_manager.start() def run_server(host: str = DEFAULT_HOST, port: int = DEFAULT_PORT) -> None: