diff --git a/app/capture.py b/app/capture.py index c519242..c495b88 100644 --- a/app/capture.py +++ b/app/capture.py @@ -22,6 +22,9 @@ from app.crawler import ( from app.db import Database from app.url_policy import NetworkRules, UrlPolicyError, resolve_and_validate +CANCELLATION_POLL_SECONDS = 0.25 +PROGRESS_WRITE_INTERVAL_SECONDS = 0.5 + def _remove_tree(path: Path, allowed_root: Path) -> None: """Delete only a child path beneath the configured work root.""" @@ -104,15 +107,28 @@ class CaptureService: report_path = self.settings.reports_dir / f"{job_id}-{attempt_id}.json" heartbeat_stop = threading.Event() lease_lost = threading.Event() + cancellation_requested = threading.Event() + cancellation_lock = threading.Lock() heartbeat_thread: threading.Thread | None = None capture_deadline = time.monotonic() + int(config["max_duration_seconds"]) + last_cancellation_check = 0.0 + last_progress_at = 0.0 + last_progress_phase: str | None = None def capture_cancelled() -> bool: - return ( - lease_lost.is_set() - or self.database.is_cancel_requested(job_id) - or not self.database.owns_lease(job_id, self.worker_id) - ) + nonlocal last_cancellation_check + if lease_lost.is_set() or cancellation_requested.is_set(): + return True + now = time.monotonic() + with cancellation_lock: + if now - last_cancellation_check < CANCELLATION_POLL_SECONDS: + return False + last_cancellation_check = now + # Fetch workers call this predicate for every response chunk. Poll SQLite sparingly. + if self.database.worker_should_stop(job_id, self.worker_id): + cancellation_requested.set() + return True + return False def renew_lease() -> None: interval = min(20, max(3, self.settings.lease_seconds // 3)) @@ -170,6 +186,15 @@ class CaptureService: def update( phase: str, message: str, stats: dict[str, int], warnings: list[str] ) -> None: + nonlocal last_progress_at, last_progress_phase + now = time.monotonic() + if ( + phase == last_progress_phase + and now - last_progress_at < PROGRESS_WRITE_INTERVAL_SECONDS + ): + return + last_progress_at = now + last_progress_phase = phase renewed = self.database.update_progress( job_id, self.worker_id, diff --git a/app/crawler.py b/app/crawler.py index c9f14e8..9f6df22 100644 --- a/app/crawler.py +++ b/app/crawler.py @@ -9,7 +9,7 @@ import threading import time from collections import deque from collections.abc import Callable -from concurrent.futures import ThreadPoolExecutor, as_completed +from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait from dataclasses import asdict, dataclass, field from pathlib import Path from typing import Literal @@ -60,6 +60,8 @@ JS_SOURCE_MAP_RE = re.compile(r"(?://[#@]\s*sourceMappingURL=)(?P\S+)") META_REFRESH_RE = re.compile(r"^(?P\s*\d+(?:\.\d+)?\s*;?\s*url\s*=\s*)(?P.+)$", re.I) MAX_RECORDED_WARNINGS = 100 MAX_RECORDED_ERRORS = 250 +MAX_CONNECTION_TIMEOUT_SECONDS = 5.0 +MAX_READ_TIMEOUT_SECONDS = 5.0 WINDOWS_RESERVED_NAMES = { "CON", "PRN", @@ -294,7 +296,8 @@ class SiteCrawler: if self._proxy_url: self._pin_proxy() self._enqueue(CrawlRequest(self.source_url, "document", 0)) - with ThreadPoolExecutor(max_workers=self.parallel_connections) as executor: + executor = ThreadPoolExecutor(max_workers=self.parallel_connections) + try: while self.pending and not self._limit_reached: self._ensure_not_cancelled() if self._duration_exceeded(): @@ -304,12 +307,21 @@ class SiteCrawler: self.pending.popleft() for _ in range(min(len(self.pending), self.parallel_connections)) ] - futures = [executor.submit(self._fetch, request) for request in batch] - for future in as_completed(futures): + pending_futures = {executor.submit(self._fetch, request) for request in batch} + while pending_futures: + completed, pending_futures = wait( + pending_futures, + timeout=0.1, + return_when=FIRST_COMPLETED, + ) self._ensure_not_cancelled() - self._handle_fetch_result(future.result()) + for future in completed: + self._handle_fetch_result(future.result()) self._emit_progress("Capturing", "Discovering pages and assets") + finally: + # Do not wait for a stalled remote server after a user cancellation. + executor.shutdown(wait=not self.cancelled(), cancel_futures=True) if not self._root_record: raise CaptureError(self._root_error or "The starting page could not be captured.") @@ -969,10 +981,10 @@ class SiteCrawler: def _request_timeout(self) -> httpx.Timeout: remaining = max(0.05, self._remaining_seconds()) return httpx.Timeout( - connect=min(10.0, remaining), - read=min(20.0, remaining), - write=min(10.0, remaining), - pool=min(10.0, remaining), + connect=min(MAX_CONNECTION_TIMEOUT_SECONDS, remaining), + read=min(MAX_READ_TIMEOUT_SECONDS, remaining), + write=min(MAX_CONNECTION_TIMEOUT_SECONDS, remaining), + pool=min(MAX_CONNECTION_TIMEOUT_SECONDS, remaining), ) def _pin_target(self, url: str) -> None: diff --git a/app/db.py b/app/db.py index 41972e3..634bfc8 100644 --- a/app/db.py +++ b/app/db.py @@ -27,11 +27,11 @@ class Database: self.path = path @contextmanager - def connection(self) -> Iterator[sqlite3.Connection]: - connection = sqlite3.connect(self.path, timeout=15, isolation_level=None) + def connection(self, *, timeout: float = 15.0) -> Iterator[sqlite3.Connection]: + connection = sqlite3.connect(self.path, timeout=timeout, isolation_level=None) connection.row_factory = sqlite3.Row connection.execute("PRAGMA foreign_keys = ON") - connection.execute("PRAGMA busy_timeout = 15000") + connection.execute(f"PRAGMA busy_timeout = {int(timeout * 1_000)}") try: yield connection finally: @@ -106,6 +106,32 @@ class Database: "CREATE INDEX IF NOT EXISTS jobs_owner_created_idx " "ON jobs(owner_token, created_at DESC)" ) + cancelled_rows = connection.execute( + "SELECT id FROM jobs WHERE state = 'running' AND cancel_requested = 1" + ).fetchall() + if cancelled_rows: + now = timestamp() + connection.execute( + """ + UPDATE jobs + SET state = 'cancelled', phase = 'Cancelled', + message = 'Cancellation requested', + worker_id = NULL, lease_expires_at = NULL, completed_at = ?, updated_at = ? + WHERE state = 'running' AND cancel_requested = 1 + """, + (now, now), + ) + for row in cancelled_rows: + self._append_event( + connection, + str(row["id"]), + "state", + { + "state": "cancelled", + "phase": "Cancelled", + "message": "Cancellation requested", + }, + ) def create_job( self, @@ -429,6 +455,30 @@ class Database: and row["lease_expires_at"] >= now ) + def worker_should_stop(self, job_id: str, worker_id: str) -> bool: + """Check cancellation and lease ownership without stalling a crawl on a busy database.""" + now = timestamp() + try: + with self.connection(timeout=0.5) as connection: + row = connection.execute( + """ + SELECT state, cancel_requested, worker_id, lease_expires_at + FROM jobs + WHERE id = ? + """, + (job_id,), + ).fetchone() + except sqlite3.OperationalError: + return False + return not bool( + row + and not row["cancel_requested"] + and row["state"] == "running" + and row["worker_id"] == worker_id + and row["lease_expires_at"] + and row["lease_expires_at"] >= now + ) + def request_cancel(self, job_id: str) -> dict[str, Any] | None: now = timestamp() with self.connection() as connection: @@ -463,14 +513,16 @@ class Database: cursor = connection.execute( """ UPDATE jobs - SET cancel_requested = 1, message = 'Cancellation requested', updated_at = ? - WHERE id = ? AND state = 'running' AND cancel_requested = 0 + SET cancel_requested = 1, state = 'cancelled', phase = 'Cancelled', + message = 'Cancellation requested', worker_id = NULL, + lease_expires_at = NULL, completed_at = ?, updated_at = ? + WHERE id = ? AND state = 'running' """, - (now, job_id), + (now, now, job_id), ) payload = { - "state": "running", - "phase": "Cancelling", + "state": "cancelled", + "phase": "Cancelled", "message": "Cancellation requested", } if cursor.rowcount: diff --git a/app/main.py b/app/main.py index f27487b..d4d88f8 100644 --- a/app/main.py +++ b/app/main.py @@ -66,7 +66,7 @@ async def secure_headers(_: Request, call_next: Any) -> Response: @app.get("/", response_class=HTMLResponse) -async def index(request: Request) -> Response: +def index(request: Request) -> Response: return templates.TemplateResponse( request=request, name="index.html", @@ -75,7 +75,7 @@ async def index(request: Request) -> Response: @app.get("/api/health") -async def health() -> dict[str, Any]: +def health() -> dict[str, Any]: return { "status": "ok", "worker_enabled": settings.worker_enabled, @@ -85,7 +85,7 @@ async def health() -> dict[str, Any]: @app.get("/api/capabilities") -async def capabilities() -> dict[str, Any]: +def capabilities() -> dict[str, Any]: return { "defaults": { "max_pages": min(150, settings.max_pages_cap), @@ -113,12 +113,12 @@ async def capabilities() -> dict[str, Any]: @app.get("/api/stats") -async def public_stats() -> dict[str, int]: +def public_stats() -> dict[str, int]: return database.public_stats() @app.post("/api/captures", status_code=status.HTTP_202_ACCEPTED) -async def create_capture(payload: CaptureRequest, request: Request) -> dict[str, Any]: +def create_capture(payload: CaptureRequest, request: Request) -> dict[str, Any]: owner_token = _owner_token(request) config = _validated_config(payload) rules = NetworkRules( @@ -135,18 +135,18 @@ async def create_capture(payload: CaptureRequest, request: Request) -> dict[str, @app.get("/api/captures") -async def list_captures(request: Request) -> dict[str, list[dict[str, Any]]]: +def list_captures(request: Request) -> dict[str, list[dict[str, Any]]]: owner_token = _owner_token(request) return {"items": [_public_job(job) for job in database.list_owned_jobs(owner_token)]} @app.get("/api/captures/{job_id}") -async def get_capture(job_id: str, request: Request) -> dict[str, Any]: +def get_capture(job_id: str, request: Request) -> dict[str, Any]: return _require_owned_job(job_id, _owner_token(request)) @app.post("/api/captures/{job_id}/cancel") -async def cancel_capture(job_id: str, request: Request) -> dict[str, Any]: +def cancel_capture(job_id: str, request: Request) -> dict[str, Any]: _require_owned_job(job_id, _owner_token(request)) job = database.request_cancel(job_id) if not job: @@ -155,7 +155,7 @@ async def cancel_capture(job_id: str, request: Request) -> dict[str, Any]: @app.delete("/api/captures/{job_id}") -async def delete_capture(job_id: str, request: Request) -> dict[str, Any]: +def delete_capture(job_id: str, request: Request) -> dict[str, Any]: job = _owned_job(job_id, _owner_token(request)) if job["state"] in {"queued", "running"}: cancelled = database.request_cancel(job_id) @@ -174,7 +174,7 @@ async def delete_capture(job_id: str, request: Request) -> dict[str, Any]: @app.get("/api/captures/{job_id}/report") -async def capture_report(job_id: str, request: Request) -> JSONResponse: +def capture_report(job_id: str, request: Request) -> JSONResponse: job = _owned_job(job_id, _owner_token(request)) if job["state"] != "ready" or not job.get("report_path"): raise HTTPException(status_code=409, detail="A capture report is not available yet.") @@ -185,7 +185,7 @@ async def capture_report(job_id: str, request: Request) -> JSONResponse: @app.get("/api/captures/{job_id}/download") -async def download_capture(job_id: str, request: Request) -> FileResponse: +def download_capture(job_id: str, request: Request) -> FileResponse: job = _owned_job(job_id, _owner_token(request)) if job["state"] != "ready" or not job.get("artifact_path"): raise HTTPException(status_code=409, detail="The archive is not available yet.") @@ -201,7 +201,7 @@ async def download_capture(job_id: str, request: Request) -> FileResponse: @app.get("/api/captures/{job_id}/events") -async def capture_events(request: Request, job_id: str) -> StreamingResponse: +def capture_events(request: Request, job_id: str) -> StreamingResponse: _owned_job(job_id, _owner_token(request, allow_query=True)) after_header = request.headers.get("last-event-id", "0") try: @@ -221,13 +221,13 @@ async def _event_stream(request: Request, job_id: str, after_id: int) -> AsyncIt while True: if await request.is_disconnected(): return - events = database.get_events(job_id, last_id) + events = await asyncio.to_thread(database.get_events, job_id, last_id) for event in events: last_id = int(event["id"]) payload = json.dumps(event["payload"], separators=(",", ":")) yield f"id: {last_id}\nevent: {event['kind']}\ndata: {payload}\n\n" - job = database.get_job(job_id) + job = await asyncio.to_thread(database.get_job, job_id) if not job or job["state"] in terminal_states: terminal = json.dumps({"state": job["state"] if job else "missing"}) yield f"event: terminal\ndata: {terminal}\n\n" diff --git a/app/static/app.js b/app/static/app.js index 9c0fe83..17f8e38 100644 --- a/app/static/app.js +++ b/app/static/app.js @@ -469,10 +469,23 @@ async function loadPublicStats() { } async function cancelJob(jobId) { + const current = jobs.get(jobId); + if (current && !current.cancel_requested) { + upsertJob( + { + ...current, + cancel_requested: true, + phase: "Cancelling", + message: "Cancellation requested", + }, + { persist: false, connect: false }, + ); + } try { upsertJob(await request(`/api/captures/${encodeURIComponent(jobId)}/cancel`, { method: "POST" })); void loadPublicStats(); } catch (error) { + if (current) upsertJob(current); setMessage(error.message); } } diff --git a/tests/test_capture_service.py b/tests/test_capture_service.py index 26ac78d..f79802d 100644 --- a/tests/test_capture_service.py +++ b/tests/test_capture_service.py @@ -37,6 +37,38 @@ class MinimalSiteHandler(BaseHTTPRequestHandler): return +class SlowResponseHandler(BaseHTTPRequestHandler): + slow_request_started = threading.Event() + release_slow_request = threading.Event() + + def do_GET(self) -> None: # noqa: N802 + if self.path == "/": + body = b"" + self.send_response(200) + self.send_header("Content-Type", "text/html") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + return + if self.path == "/slow.bin": + self.send_response(200) + self.send_header("Content-Type", "application/octet-stream") + self.send_header("Content-Length", "1") + self.end_headers() + self.slow_request_started.set() + self.release_slow_request.wait(timeout=30) + try: + self.wfile.write(b"x") + except BrokenPipeError: + pass + return + self.send_response(404) + self.end_headers() + + def log_message(self, format: str, *args) -> None: # type: ignore[override] + return + + def test_capture_service_creates_private_archive(tmp_path) -> None: server = ThreadingHTTPServer(("127.0.0.1", 0), MinimalSiteHandler) thread = threading.Thread(target=server.serve_forever, daemon=True) @@ -89,3 +121,64 @@ def test_capture_service_creates_private_archive(tmp_path) -> None: names = archive.namelist() assert any(name.endswith("siteharbor-manifest.json") for name in names) assert any(name.endswith("site.css") for name in names) + + +def test_capture_cancels_while_a_resource_is_stalled(tmp_path) -> None: + SlowResponseHandler.slow_request_started.clear() + SlowResponseHandler.release_slow_request.clear() + server = ThreadingHTTPServer(("127.0.0.1", 0), SlowResponseHandler) + server_thread = threading.Thread(target=server.serve_forever, daemon=True) + server_thread.start() + port = server.server_address[1] + + test_settings = replace( + settings, + data_dir=tmp_path, + database_path=tmp_path / "siteharbor.sqlite3", + work_dir=tmp_path / "work", + artifacts_dir=tmp_path / "artifacts", + reports_dir=tmp_path / "reports", + allow_private_networks=True, + allow_nonstandard_ports=True, + respect_robots=False, + ) + test_settings.ensure_directories() + database = Database(test_settings.database_path) + database.initialize() + source_url = f"http://127.0.0.1:{port}/" + job = database.create_job( + source_url, + source_url, + { + "include_external_assets": False, + "max_pages": 10, + "max_depth": 2, + "max_bytes": 2_000_000, + "max_duration_seconds": 60, + "retention_days": 1, + }, + ) + claimed = database.claim_next("test-worker", lease_seconds=120) + assert claimed is not None + capture_thread = threading.Thread( + target=CaptureService(database, test_settings, "test-worker").run, + args=(claimed,), + daemon=True, + ) + capture_thread.start() + + try: + assert SlowResponseHandler.slow_request_started.wait(timeout=2) + cancelled = database.request_cancel(job["id"]) + assert cancelled is not None + capture_thread.join(timeout=3) + assert not capture_thread.is_alive() + finally: + SlowResponseHandler.release_slow_request.set() + capture_thread.join(timeout=2) + server.shutdown() + server.server_close() + + cancelled_job = database.get_job(job["id"]) + assert cancelled_job is not None + assert cancelled_job["state"] == "cancelled" diff --git a/tests/test_database.py b/tests/test_database.py index 9f45d63..b5d5902 100644 --- a/tests/test_database.py +++ b/tests/test_database.py @@ -97,13 +97,73 @@ def test_worker_cannot_cancel_without_a_user_cancellation_request(tmp_path) -> N assert running is not None assert running["state"] == "running" - database.request_cancel(job["id"]) - assert database.mark_cancelled(job["id"], "worker-a", "Cancelled by user") + cancellation = database.request_cancel(job["id"]) + assert cancellation is not None + assert cancellation["state"] == "cancelled" + assert not database.mark_cancelled(job["id"], "worker-a", "Cancelled by user") cancelled = database.get_job(job["id"]) assert cancelled is not None assert cancelled["state"] == "cancelled" +def test_worker_stop_check_detects_a_cancellation_request(tmp_path) -> None: + database = Database(tmp_path / "siteharbor.sqlite3") + database.initialize() + job = database.create_job( + "https://example.com", + "https://example.com/", + capture_config(), + ) + assert database.claim_next("worker-a", lease_seconds=30) + + assert not database.worker_should_stop(job["id"], "worker-a") + assert database.request_cancel(job["id"]) + assert database.worker_should_stop(job["id"], "worker-a") + + +def test_cancel_terminalizes_an_existing_cancellation_request(tmp_path) -> None: + database = Database(tmp_path / "siteharbor.sqlite3") + database.initialize() + job = database.create_job( + "https://example.com", + "https://example.com/", + capture_config(), + ) + assert database.claim_next("worker-a", lease_seconds=30) + with database.connection() as connection: + connection.execute( + "UPDATE jobs SET cancel_requested = 1 WHERE id = ?", + (job["id"],), + ) + + cancelled = database.request_cancel(job["id"]) + + assert cancelled is not None + assert cancelled["state"] == "cancelled" + + +def test_initialize_recovers_a_legacy_cancellation_request(tmp_path) -> None: + database = Database(tmp_path / "siteharbor.sqlite3") + database.initialize() + job = database.create_job( + "https://example.com", + "https://example.com/", + capture_config(), + ) + assert database.claim_next("worker-a", lease_seconds=30) + with database.connection() as connection: + connection.execute( + "UPDATE jobs SET cancel_requested = 1 WHERE id = ?", + (job["id"],), + ) + + database.initialize() + + recovered = database.get_job(job["id"]) + assert recovered is not None + assert recovered["state"] == "cancelled" + + def test_expired_cancel_requested_job_is_terminalized_before_reclaim(tmp_path) -> None: database = Database(tmp_path / "siteharbor.sqlite3") database.initialize()