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.db import Database
|
||||||
from app.url_policy import NetworkRules, UrlPolicyError, resolve_and_validate
|
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:
|
def _remove_tree(path: Path, allowed_root: Path) -> None:
|
||||||
"""Delete only a child path beneath the configured work root."""
|
"""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"
|
report_path = self.settings.reports_dir / f"{job_id}-{attempt_id}.json"
|
||||||
heartbeat_stop = threading.Event()
|
heartbeat_stop = threading.Event()
|
||||||
lease_lost = threading.Event()
|
lease_lost = threading.Event()
|
||||||
|
cancellation_requested = threading.Event()
|
||||||
|
cancellation_lock = threading.Lock()
|
||||||
heartbeat_thread: threading.Thread | None = None
|
heartbeat_thread: threading.Thread | None = None
|
||||||
capture_deadline = time.monotonic() + int(config["max_duration_seconds"])
|
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:
|
def capture_cancelled() -> bool:
|
||||||
return (
|
nonlocal last_cancellation_check
|
||||||
lease_lost.is_set()
|
if lease_lost.is_set() or cancellation_requested.is_set():
|
||||||
or self.database.is_cancel_requested(job_id)
|
return True
|
||||||
or not self.database.owns_lease(job_id, self.worker_id)
|
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:
|
def renew_lease() -> None:
|
||||||
interval = min(20, max(3, self.settings.lease_seconds // 3))
|
interval = min(20, max(3, self.settings.lease_seconds // 3))
|
||||||
@@ -170,6 +186,15 @@ class CaptureService:
|
|||||||
def update(
|
def update(
|
||||||
phase: str, message: str, stats: dict[str, int], warnings: list[str]
|
phase: str, message: str, stats: dict[str, int], warnings: list[str]
|
||||||
) -> None:
|
) -> 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(
|
renewed = self.database.update_progress(
|
||||||
job_id,
|
job_id,
|
||||||
self.worker_id,
|
self.worker_id,
|
||||||
|
|||||||
+20
-8
@@ -9,7 +9,7 @@ import threading
|
|||||||
import time
|
import time
|
||||||
from collections import deque
|
from collections import deque
|
||||||
from collections.abc import Callable
|
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 dataclasses import asdict, dataclass, field
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Literal
|
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)
|
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_WARNINGS = 100
|
||||||
MAX_RECORDED_ERRORS = 250
|
MAX_RECORDED_ERRORS = 250
|
||||||
|
MAX_CONNECTION_TIMEOUT_SECONDS = 5.0
|
||||||
|
MAX_READ_TIMEOUT_SECONDS = 5.0
|
||||||
WINDOWS_RESERVED_NAMES = {
|
WINDOWS_RESERVED_NAMES = {
|
||||||
"CON",
|
"CON",
|
||||||
"PRN",
|
"PRN",
|
||||||
@@ -294,7 +296,8 @@ class SiteCrawler:
|
|||||||
if self._proxy_url:
|
if self._proxy_url:
|
||||||
self._pin_proxy()
|
self._pin_proxy()
|
||||||
self._enqueue(CrawlRequest(self.source_url, "document", 0))
|
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:
|
while self.pending and not self._limit_reached:
|
||||||
self._ensure_not_cancelled()
|
self._ensure_not_cancelled()
|
||||||
if self._duration_exceeded():
|
if self._duration_exceeded():
|
||||||
@@ -304,12 +307,21 @@ class SiteCrawler:
|
|||||||
self.pending.popleft()
|
self.pending.popleft()
|
||||||
for _ in range(min(len(self.pending), self.parallel_connections))
|
for _ in range(min(len(self.pending), self.parallel_connections))
|
||||||
]
|
]
|
||||||
futures = [executor.submit(self._fetch, request) for request in batch]
|
pending_futures = {executor.submit(self._fetch, request) for request in batch}
|
||||||
for future in as_completed(futures):
|
while pending_futures:
|
||||||
|
completed, pending_futures = wait(
|
||||||
|
pending_futures,
|
||||||
|
timeout=0.1,
|
||||||
|
return_when=FIRST_COMPLETED,
|
||||||
|
)
|
||||||
self._ensure_not_cancelled()
|
self._ensure_not_cancelled()
|
||||||
|
for future in completed:
|
||||||
self._handle_fetch_result(future.result())
|
self._handle_fetch_result(future.result())
|
||||||
|
|
||||||
self._emit_progress("Capturing", "Discovering pages and assets")
|
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:
|
if not self._root_record:
|
||||||
raise CaptureError(self._root_error or "The starting page could not be captured.")
|
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:
|
def _request_timeout(self) -> httpx.Timeout:
|
||||||
remaining = max(0.05, self._remaining_seconds())
|
remaining = max(0.05, self._remaining_seconds())
|
||||||
return httpx.Timeout(
|
return httpx.Timeout(
|
||||||
connect=min(10.0, remaining),
|
connect=min(MAX_CONNECTION_TIMEOUT_SECONDS, remaining),
|
||||||
read=min(20.0, remaining),
|
read=min(MAX_READ_TIMEOUT_SECONDS, remaining),
|
||||||
write=min(10.0, remaining),
|
write=min(MAX_CONNECTION_TIMEOUT_SECONDS, remaining),
|
||||||
pool=min(10.0, remaining),
|
pool=min(MAX_CONNECTION_TIMEOUT_SECONDS, remaining),
|
||||||
)
|
)
|
||||||
|
|
||||||
def _pin_target(self, url: str) -> None:
|
def _pin_target(self, url: str) -> None:
|
||||||
|
|||||||
@@ -27,11 +27,11 @@ class Database:
|
|||||||
self.path = path
|
self.path = path
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def connection(self) -> Iterator[sqlite3.Connection]:
|
def connection(self, *, timeout: float = 15.0) -> Iterator[sqlite3.Connection]:
|
||||||
connection = sqlite3.connect(self.path, timeout=15, isolation_level=None)
|
connection = sqlite3.connect(self.path, timeout=timeout, isolation_level=None)
|
||||||
connection.row_factory = sqlite3.Row
|
connection.row_factory = sqlite3.Row
|
||||||
connection.execute("PRAGMA foreign_keys = ON")
|
connection.execute("PRAGMA foreign_keys = ON")
|
||||||
connection.execute("PRAGMA busy_timeout = 15000")
|
connection.execute(f"PRAGMA busy_timeout = {int(timeout * 1_000)}")
|
||||||
try:
|
try:
|
||||||
yield connection
|
yield connection
|
||||||
finally:
|
finally:
|
||||||
@@ -106,6 +106,32 @@ class Database:
|
|||||||
"CREATE INDEX IF NOT EXISTS jobs_owner_created_idx "
|
"CREATE INDEX IF NOT EXISTS jobs_owner_created_idx "
|
||||||
"ON jobs(owner_token, created_at DESC)"
|
"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(
|
def create_job(
|
||||||
self,
|
self,
|
||||||
@@ -429,6 +455,30 @@ class Database:
|
|||||||
and row["lease_expires_at"] >= now
|
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:
|
def request_cancel(self, job_id: str) -> dict[str, Any] | None:
|
||||||
now = timestamp()
|
now = timestamp()
|
||||||
with self.connection() as connection:
|
with self.connection() as connection:
|
||||||
@@ -463,14 +513,16 @@ class Database:
|
|||||||
cursor = connection.execute(
|
cursor = connection.execute(
|
||||||
"""
|
"""
|
||||||
UPDATE jobs
|
UPDATE jobs
|
||||||
SET cancel_requested = 1, message = 'Cancellation requested', updated_at = ?
|
SET cancel_requested = 1, state = 'cancelled', phase = 'Cancelled',
|
||||||
WHERE id = ? AND state = 'running' AND cancel_requested = 0
|
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 = {
|
payload = {
|
||||||
"state": "running",
|
"state": "cancelled",
|
||||||
"phase": "Cancelling",
|
"phase": "Cancelled",
|
||||||
"message": "Cancellation requested",
|
"message": "Cancellation requested",
|
||||||
}
|
}
|
||||||
if cursor.rowcount:
|
if cursor.rowcount:
|
||||||
|
|||||||
+14
-14
@@ -66,7 +66,7 @@ async def secure_headers(_: Request, call_next: Any) -> Response:
|
|||||||
|
|
||||||
|
|
||||||
@app.get("/", response_class=HTMLResponse)
|
@app.get("/", response_class=HTMLResponse)
|
||||||
async def index(request: Request) -> Response:
|
def index(request: Request) -> Response:
|
||||||
return templates.TemplateResponse(
|
return templates.TemplateResponse(
|
||||||
request=request,
|
request=request,
|
||||||
name="index.html",
|
name="index.html",
|
||||||
@@ -75,7 +75,7 @@ async def index(request: Request) -> Response:
|
|||||||
|
|
||||||
|
|
||||||
@app.get("/api/health")
|
@app.get("/api/health")
|
||||||
async def health() -> dict[str, Any]:
|
def health() -> dict[str, Any]:
|
||||||
return {
|
return {
|
||||||
"status": "ok",
|
"status": "ok",
|
||||||
"worker_enabled": settings.worker_enabled,
|
"worker_enabled": settings.worker_enabled,
|
||||||
@@ -85,7 +85,7 @@ async def health() -> dict[str, Any]:
|
|||||||
|
|
||||||
|
|
||||||
@app.get("/api/capabilities")
|
@app.get("/api/capabilities")
|
||||||
async def capabilities() -> dict[str, Any]:
|
def capabilities() -> dict[str, Any]:
|
||||||
return {
|
return {
|
||||||
"defaults": {
|
"defaults": {
|
||||||
"max_pages": min(150, settings.max_pages_cap),
|
"max_pages": min(150, settings.max_pages_cap),
|
||||||
@@ -113,12 +113,12 @@ async def capabilities() -> dict[str, Any]:
|
|||||||
|
|
||||||
|
|
||||||
@app.get("/api/stats")
|
@app.get("/api/stats")
|
||||||
async def public_stats() -> dict[str, int]:
|
def public_stats() -> dict[str, int]:
|
||||||
return database.public_stats()
|
return database.public_stats()
|
||||||
|
|
||||||
|
|
||||||
@app.post("/api/captures", status_code=status.HTTP_202_ACCEPTED)
|
@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)
|
owner_token = _owner_token(request)
|
||||||
config = _validated_config(payload)
|
config = _validated_config(payload)
|
||||||
rules = NetworkRules(
|
rules = NetworkRules(
|
||||||
@@ -135,18 +135,18 @@ async def create_capture(payload: CaptureRequest, request: Request) -> dict[str,
|
|||||||
|
|
||||||
|
|
||||||
@app.get("/api/captures")
|
@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)
|
owner_token = _owner_token(request)
|
||||||
return {"items": [_public_job(job) for job in database.list_owned_jobs(owner_token)]}
|
return {"items": [_public_job(job) for job in database.list_owned_jobs(owner_token)]}
|
||||||
|
|
||||||
|
|
||||||
@app.get("/api/captures/{job_id}")
|
@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))
|
return _require_owned_job(job_id, _owner_token(request))
|
||||||
|
|
||||||
|
|
||||||
@app.post("/api/captures/{job_id}/cancel")
|
@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))
|
_require_owned_job(job_id, _owner_token(request))
|
||||||
job = database.request_cancel(job_id)
|
job = database.request_cancel(job_id)
|
||||||
if not job:
|
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}")
|
@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))
|
job = _owned_job(job_id, _owner_token(request))
|
||||||
if job["state"] in {"queued", "running"}:
|
if job["state"] in {"queued", "running"}:
|
||||||
cancelled = database.request_cancel(job_id)
|
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")
|
@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))
|
job = _owned_job(job_id, _owner_token(request))
|
||||||
if job["state"] != "ready" or not job.get("report_path"):
|
if job["state"] != "ready" or not job.get("report_path"):
|
||||||
raise HTTPException(status_code=409, detail="A capture report is not available yet.")
|
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")
|
@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))
|
job = _owned_job(job_id, _owner_token(request))
|
||||||
if job["state"] != "ready" or not job.get("artifact_path"):
|
if job["state"] != "ready" or not job.get("artifact_path"):
|
||||||
raise HTTPException(status_code=409, detail="The archive is not available yet.")
|
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")
|
@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))
|
_owned_job(job_id, _owner_token(request, allow_query=True))
|
||||||
after_header = request.headers.get("last-event-id", "0")
|
after_header = request.headers.get("last-event-id", "0")
|
||||||
try:
|
try:
|
||||||
@@ -221,13 +221,13 @@ async def _event_stream(request: Request, job_id: str, after_id: int) -> AsyncIt
|
|||||||
while True:
|
while True:
|
||||||
if await request.is_disconnected():
|
if await request.is_disconnected():
|
||||||
return
|
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:
|
for event in events:
|
||||||
last_id = int(event["id"])
|
last_id = int(event["id"])
|
||||||
payload = json.dumps(event["payload"], separators=(",", ":"))
|
payload = json.dumps(event["payload"], separators=(",", ":"))
|
||||||
yield f"id: {last_id}\nevent: {event['kind']}\ndata: {payload}\n\n"
|
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:
|
if not job or job["state"] in terminal_states:
|
||||||
terminal = json.dumps({"state": job["state"] if job else "missing"})
|
terminal = json.dumps({"state": job["state"] if job else "missing"})
|
||||||
yield f"event: terminal\ndata: {terminal}\n\n"
|
yield f"event: terminal\ndata: {terminal}\n\n"
|
||||||
|
|||||||
@@ -469,10 +469,23 @@ async function loadPublicStats() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
async function cancelJob(jobId) {
|
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 {
|
try {
|
||||||
upsertJob(await request(`/api/captures/${encodeURIComponent(jobId)}/cancel`, { method: "POST" }));
|
upsertJob(await request(`/api/captures/${encodeURIComponent(jobId)}/cancel`, { method: "POST" }));
|
||||||
void loadPublicStats();
|
void loadPublicStats();
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
|
if (current) upsertJob(current);
|
||||||
setMessage(error.message);
|
setMessage(error.message);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -37,6 +37,38 @@ class MinimalSiteHandler(BaseHTTPRequestHandler):
|
|||||||
return
|
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"<html><body><img src='/slow.bin'></body></html>"
|
||||||
|
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:
|
def test_capture_service_creates_private_archive(tmp_path) -> None:
|
||||||
server = ThreadingHTTPServer(("127.0.0.1", 0), MinimalSiteHandler)
|
server = ThreadingHTTPServer(("127.0.0.1", 0), MinimalSiteHandler)
|
||||||
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
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()
|
names = archive.namelist()
|
||||||
assert any(name.endswith("siteharbor-manifest.json") for name in names)
|
assert any(name.endswith("siteharbor-manifest.json") for name in names)
|
||||||
assert any(name.endswith("site.css") 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"
|
||||||
|
|||||||
+62
-2
@@ -97,13 +97,73 @@ def test_worker_cannot_cancel_without_a_user_cancellation_request(tmp_path) -> N
|
|||||||
assert running is not None
|
assert running is not None
|
||||||
assert running["state"] == "running"
|
assert running["state"] == "running"
|
||||||
|
|
||||||
database.request_cancel(job["id"])
|
cancellation = database.request_cancel(job["id"])
|
||||||
assert database.mark_cancelled(job["id"], "worker-a", "Cancelled by user")
|
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"])
|
cancelled = database.get_job(job["id"])
|
||||||
assert cancelled is not None
|
assert cancelled is not None
|
||||||
assert cancelled["state"] == "cancelled"
|
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:
|
def test_expired_cancel_requested_job_is_terminalized_before_reclaim(tmp_path) -> None:
|
||||||
database = Database(tmp_path / "siteharbor.sqlite3")
|
database = Database(tmp_path / "siteharbor.sqlite3")
|
||||||
database.initialize()
|
database.initialize()
|
||||||
|
|||||||
Reference in New Issue
Block a user