implement cancellation handling in CaptureService and Database, update job states accordingly

This commit is contained in:
kstyagi23
2026-09-05 06:54:10 +05:30
parent 9c1a301c1a
commit cab7a68b8b
7 changed files with 293 additions and 38 deletions
+30 -5
View File
@@ -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,
+21 -9
View File
@@ -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<url>\S+)")
META_REFRESH_RE = re.compile(r"^(?P<delay>\s*\d+(?:\.\d+)?\s*;?\s*url\s*=\s*)(?P<url>.+)$", 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:
+60 -8
View File
@@ -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:
+14 -14
View File
@@ -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"
+13
View File
@@ -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);
}
}