implement cancellation handling in CaptureService and Database, update job states accordingly
This commit is contained in:
+30
-5
@@ -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
@@ -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:
|
||||
|
||||
@@ -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
@@ -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"
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user