diff --git a/exca/map.py b/exca/map.py index 2b786086..069f9ca4 100644 --- a/exca/map.py +++ b/exca/map.py @@ -11,19 +11,16 @@ import itertools import logging import os -import pickle import typing as tp -import uuid from concurrent import futures from pathlib import Path import numpy as np import pydantic -import submitit -from submitit.core import utils from . import base, slurm from .cachedict import CacheDict +from .utils import ItemQueue MapFunc = tp.Callable[[tp.Sequence[tp.Any]], tp.Iterator[tp.Any]] X = tp.TypeVar("X") @@ -65,45 +62,6 @@ def __call__(self, items: tp.Sequence[tp.Any]) -> tp.Iterator[tp.Any]: return self.infra._method_override(items) -class JobChecker: - """Keeps a record of running jobs in a folder - and enables waiting for them to complete. - """ - - def __init__(self, folder: Path | str) -> None: - basefolder = utils.JobPaths.get_first_id_independent_folder(folder) - self.folder = basefolder / "running-jobs" - - def add(self, jobs: tp.Iterable[tp.Any]) -> None: - """Add jobs to the list of running jobs""" - self.folder.mkdir(exist_ok=True, parents=True) - for job in jobs: - if not job.done(): - job_path = self.folder / (uuid.uuid4().hex[:8] + ".pkl") - with job_path.open("wb") as f: - pickle.dump(job, f) - - def wait(self) -> bool: - """Wait for completion of running jobs""" - waited = False - for fp in self.folder.glob("*.pkl"): - try: # avoid concurrency issues with deleted items - with fp.open("rb") as f: - job: tp.Any = pickle.load(f) - except Exception: # pylint: disable=broad-except - continue - if not job.done(): - msg = "Waiting for completion of pre-existing map job: %s\nin '%s'" - logger.info(msg, job, self.folder) - job.wait() - waited = True - # delete the file as it is not needed anymore - fp.unlink(missing_ok=True) - if waited: - logger.info("Waiting is over") - return waited - - def to_chunks( items: tp.List[X], *, max_chunks: int | None, min_items_per_chunk: int = 1 ) -> tp.Iterator[tp.List[X]]: @@ -317,9 +275,7 @@ def _find_missing(self, items: tp.Dict[str, tp.Any]) -> tp.Dict[str, tp.Any]: if not hasattr(self, "mode"): # compatibility self.mode = "cached" if self.mode == "force": - # remove any item already computed, but not items being computed - # in another process (waited for by JobChecker) - # will not be removed + # remove any item already computed, but not items being recomputed to_remove = set(items) - set(missing) - self._recomputed if to_remove: msg = "Clearing %s items for %s (infra.mode=%s)" @@ -333,13 +289,6 @@ def _find_missing(self, items: tp.Dict[str, tp.Any]) -> tp.Dict[str, tp.Any]: if missing: if self.mode == "read-only": raise RuntimeError(f"{self.mode=} but found {len(missing)} missing items") - executor: submitit.Executor | None = self.executor() - if executor is not None: # wait for items being computed - jcheck = JobChecker(folder=executor.folder) - jcheck.wait() - # update cache dict and recheck as actual checking for keys updates the dict - keys = set(self.cache_dict) # update cache dict - missing = {k: item for k, item in missing.items() if k not in keys} if len(items) == len(missing) == 1 and self.forbid_single_item_computation: key, item = next(iter(missing.items())) raise RuntimeError( @@ -370,151 +319,206 @@ def _method_override(self, *args: tp.Any, **kwargs: tp.Any) -> tp.Iterator[tp.An msg = f"Method {imethod.method} takes parameters {exp}, got {params}" raise NameError(msg) items = next(iter(kwargs.values())) - # specific function for thread and process pool executors - if self.cluster in [None, "threadpool", "processpool"]: - return self._method_override_futures(items) uid_func = imethod.item_uid # we need to keep order for output: uid_items = [(uid_func(item), item) for item in items] missing = list(self._find_missing(dict(uid_items)).items()) + out: dict[str, tp.Any] = {} if missing: + if self.cluster is None: + # Run locally in current thread (no queue needed) + logger.debug("Computing %s missing items locally", len(missing)) + out = self._process_items(missing, use_cache=self.folder is not None) + else: + # Use queue for coordination between workers + out = self._process_with_queue(missing) + folder = self.uid_folder() + if folder is not None: + os.utime(folder) # make sure the modified time is updated + # Return results from cache (or from out if no caching) + try: + cache_dict = self.cache_dict + except ValueError: # no caching + return (out[k] for k, _ in uid_items) + if out: # results not yet in cache (no folder case) + with cache_dict.writer() as writer: + for k, v in out.items(): + writer[k] = v + logger.debug( + "Recovering %s items for %s from %s", len(items), self._factory(), cache_dict + ) + return (cache_dict[k] for k, _ in uid_items) + + def _process_with_queue( + self, missing: tp.List[tp.Tuple[str, tp.Any]] + ) -> dict[str, tp.Any]: + """Add items to queue and spawn workers to process them.""" + # Determine queue folder based on executor type + if self.cluster in ("threadpool", "processpool"): + # For local pools, use a temp folder or the cache folder + queue_folder = self.uid_folder(create=True) + if queue_folder is None: + import tempfile + + queue_folder = Path(tempfile.mkdtemp()) + else: executor = self.executor() if executor is None: raise RuntimeError(f"Executor is None for {self.cluster!r}") - # avoid processing same files at same time if several jobs overlap - np.random.shuffle(missing) - # run on cluster - jobs = [] - chunks = list( - to_chunks( - [ki[1] for ki in missing], - max_chunks=self.max_jobs, - min_items_per_chunk=self.min_samples_per_job, - ) - ) - executor.update_parameters(slurm_array_parallelism=len(chunks)) - with self._work_env(), executor.batch(): # submitit>=1.4.6 - for chunk in chunks: - # select a batch/chunk of samples_per_job items to send to a job - j = executor.submit(self._call_and_store, chunk, use_cache_dict=True) - jobs.append(j) - jcheck = JobChecker(folder=executor.folder) - jcheck.add(jobs) - # pylint: disable=expression-not-assigned - uid = self.uid() - msg = "Sent %s samples for %s into %s jobs on cluster '%s' (eg: %s)" + queue_folder = executor.folder + + item_queue = ItemQueue(queue_folder) + expected_uids = [uid for uid, _ in missing] + + # Add missing items to the queue (workers will claim them) + added = item_queue.add_items(missing) + if added < len(missing): logger.info( - msg, len(missing), uid, len(jobs), executor.cluster, jobs[0].job_id + "Added %s/%s items to queue (others already queued)", + added, + len(missing), + ) + else: + logger.info("Added %s items to queue", added) + + # Calculate number of workers based on pending items + pending = item_queue.pending_count() + if pending > 0: + num_workers = min( + pending if self.max_jobs is None else self.max_jobs, + int(np.ceil(pending / self.min_samples_per_job)), ) - [j.result() for j in jobs] # wait for processing to complete - logger.info("Finished processing %s samples for %s", len(missing), uid) - folder = self.uid_folder() - if folder is not None: - os.utime(folder) # make sure the modified time is updated - msg = "Recovering %s items for %s from %s" - # using factory because uid is too slow for here - logger.debug(msg, len(items), self._factory(), self.cache_dict) - return (self.cache_dict[k] for k, _ in uid_items) - def _method_override_futures(self, items: tp.Sequence[tp.Any]) -> tp.Iterator[tp.Any]: - imethod = self._infra_method - if imethod is None: - raise RuntimeError(f"Infra was not applied: {self!r}") - uid_func = imethod.item_uid # type: ignore - uid_items = [ - (uid_func(item), item) for item in items - ] # we need to keep order for output - missing = list(self._find_missing(dict(uid_items)).items()) - out = {} - if missing: - pool = self.cluster - if len(missing) == 1: - pool = None - # avoid processing same files at same time if several jobs overlap - np.random.shuffle(missing) - if pool is None: - # run locally - msg = "Computing %s missing items" - logger.debug(msg, len(missing)) - cached = self.folder is not None - out = self._call_and_store( - [ki[1] for ki in missing], use_cache_dict=cached - ) - elif pool not in ("processpool", "threadpool"): - raise RuntimeError(f"Unexpected pool {pool!r}") - else: + if self.cluster in ("threadpool", "processpool"): + # Use concurrent.futures for local parallel execution ExecutorCls = ( futures.ThreadPoolExecutor - if pool == "threadpool" + if self.cluster == "threadpool" else futures.ProcessPoolExecutor ) + with ExecutorCls(max_workers=num_workers) as ex: + jobs = [ + ex.submit( + self._process_from_queue, + queue_folder, + self.min_samples_per_job, + ) + for _ in range(num_workers) + ] + logger.info( + "Sent %s workers for %s items into %s", + num_workers, + pending, + self.cluster, + ) + for job in _set_tqdm(futures.as_completed(jobs), total=len(jobs)): + job.result() # raise asap + logger.info("Finished processing %s items for %s", pending, self.uid()) + else: + # Use submitit for slurm/local/debug + executor = self.executor() + if executor is None: + raise RuntimeError(f"Executor is None for {self.cluster!r}") + executor.update_parameters(slurm_array_parallelism=num_workers) jobs = [] - max_workers = self.max_jobs - if max_workers is not None: - max_workers = min(len(missing), max_workers) - with ExecutorCls(max_workers=max_workers) as ex: - # split in a manageable number of chunks - mitems = [ki[1] for ki in missing] - max_workers = ex._max_workers # type: ignore - chunks = to_chunks(mitems, max_chunks=3 * max_workers) # type: ignore - for chunk in chunks: - j = ex.submit( - self._call_and_store, - chunk, - use_cache_dict=self.folder is not None, + with self._work_env(), executor.batch(): + for _ in range(num_workers): + j = executor.submit( + self._process_from_queue, + queue_folder=queue_folder, + batch_size=self.min_samples_per_job, ) jobs.append(j) - uid = self.uid() - msg = "Sent %s items for %s into a %s" - logger.info(msg, len(missing), uid, pool) - iterator = _set_tqdm(futures.as_completed(jobs), total=len(jobs)) - for job in iterator: - out.update(job.result()) # raise asap - logger.info("Finished processing %s items for %s", len(missing), uid) - folder = self.uid_folder() - if folder is not None: - os.utime(folder) # make sure the modified time is updated - try: - cache_dict = self.cache_dict - except ValueError: # no caching - return (out[k] for k, _ in uid_items) - if out: # keep in ram activated but no folder - with cache_dict.writer() as writer: - for x, y in out.items(): - writer[x] = y - msg = "Recovering %s items for %s from %s" - # using factory because uid is too slow for here - logger.debug(msg, len(uid_items), self._factory(), self.cache_dict) - return (cache_dict[k] for k, _ in uid_items) + logger.info( + "Sent %s jobs for %s items on cluster '%s' (eg: %s)", + len(jobs), + pending, + executor.cluster, + jobs[0].job_id, + ) + [j.result() for j in jobs] # wait for completion + logger.info("Finished processing items for %s", self.uid()) - def _call_and_store( - self, items: tp.Sequence[tp.Any], use_cache_dict: bool = True + # Wait for all expected items to be removed from queue + # (handles items processed by concurrent workers, with stale reclaim) + item_queue.wait_for_completion(expected_uids) + return {} + + def _process_items( + self, + uid_items: tp.Sequence[tp.Tuple[str, tp.Any]], + use_cache: bool = True, ) -> dict[str, tp.Any]: - d: dict[str, tp.Any] = self.cache_dict if use_cache_dict else {} # type: ignore - imethod = self._infra_method - if imethod is None: - raise RuntimeError(f"Infra was not applied: {self!r}") - item_uid = imethod.item_uid - if items: # make sure some overlapping job did not already run stuff - keys = set(d) # update cache dict - items = [item for item in items if item_uid(item) not in keys] - if isinstance(self, slurm.SubmititMixin): # dependence to mixin - if self.workdir is not None and self.cluster is not None and items: + """Core processing: run method on items and store results. + + Parameters + ---------- + uid_items: sequence of (uid, item) tuples + Items to process with their cache keys + use_cache: bool + If True, store results in cache_dict; if False, return dict of results + + Returns + ------- + dict + Empty if use_cache=True, otherwise {uid: result} for each item + """ + if not uid_items: + return {} + d: dict[str, tp.Any] = self.cache_dict if use_cache else {} # type: ignore + # Filter out items already in cache + if uid_items: + keys = set(d) if isinstance(d, (dict, CacheDict)) else set() + uid_items = [(uid, item) for uid, item in uid_items if uid not in keys] + if not uid_items: + return {} + if isinstance(self, slurm.SubmititMixin): + if self.workdir is not None and self.cluster is not None: logger.info("Running from working directory: '%s'", os.getcwd()) + # Process items + items = [item for _, item in uid_items] outputs = self._run_method(items) + # Store results sentinel = base.Sentinel() with contextlib.ExitStack() as estack: writer = d if isinstance(d, CacheDict): writer = estack.enter_context(d.writer()) # type: ignore - in_out = itertools.zip_longest(_set_tqdm(items), outputs, fillvalue=sentinel) - for item, output in in_out: - if item is sentinel or output is sentinel: + in_out = itertools.zip_longest( + _set_tqdm(uid_items), outputs, fillvalue=sentinel + ) + for (uid, item), output in in_out: + if (uid, item) is sentinel or output is sentinel: msg = f"Cached function did not yield exactly once per item: {item=!r}, {output=!r}" raise RuntimeError(msg) - writer[item_uid(item)] = output - # don't return the whole cache dict if data is cached - return {} if use_cache_dict else d + writer[uid] = output + return {} if use_cache else d + + def _process_from_queue( + self, queue_folder: Path | str, batch_size: int = 100 + ) -> None: + """Worker method: claim items from queue and process them in batches. + + Workers claim batches of items from the shared queue, process them, + mark them done, and repeat until no pending items remain. + This allows dynamic load balancing across workers. + """ + item_queue = ItemQueue(queue_folder) + total_processed = 0 + while True: + # Claim a batch of pending items from the queue + claimed = item_queue.claim_batch(batch_size) + if not claimed: + break # No pending items + # Process the batch (filtering for cache is done inside _process_items) + self._process_items(claimed, use_cache=True) + # Mark items as done (removes from queue) + item_queue.mark_done([uid for uid, _ in claimed]) + total_processed += len(claimed) + logger.debug( + "Processed batch of %s items (total: %s)", len(claimed), total_processed + ) + logger.info("Worker finished, processed %s items total", total_processed) @dataclasses.dataclass diff --git a/exca/test_utils.py b/exca/test_utils.py index 065df5d1..f2c499b0 100644 --- a/exca/test_utils.py +++ b/exca/test_utils.py @@ -458,3 +458,215 @@ def test_basic_pydantic() -> None: with pytest.raises(RuntimeError) as e: b.infra.clone_obj() assert "discriminated union" in e.value.args[0] + + +# ItemQueue tests + + +def test_item_queue_basic(tmp_path: Path) -> None: + """Test basic add, claim, and mark_done operations.""" + queue = utils.ItemQueue(tmp_path / "queue") + + # Add items + items = [("a", 1), ("b", 2), ("c", 3)] + added = queue.add_items(items) + assert added == 3 + assert len(queue) == 3 + assert queue.pending_count() == 3 + + # Claim all items (marks them as claimed, doesn't remove) + claimed = queue.claim_batch(batch_size=10) + assert len(claimed) == 3 + assert set(uid for uid, _ in claimed) == {"a", "b", "c"} + assert len(queue) == 3 # Still in queue (claimed status) + assert queue.pending_count() == 0 # No pending items + + # Claim again - no pending items + claimed2 = queue.claim_batch() + assert claimed2 == [] + + # Mark items as done + queue.mark_done([uid for uid, _ in claimed]) + assert len(queue) == 0 + + +def test_item_queue_batch_claiming(tmp_path: Path) -> None: + """Test claiming items in batches.""" + queue = utils.ItemQueue(tmp_path / "queue") + + # Add 10 items + items = [(str(i), i) for i in range(10)] + queue.add_items(items) + assert len(queue) == 10 + + # Claim in batches of 3 + batch1 = queue.claim_batch(batch_size=3) + assert len(batch1) == 3 + assert queue.pending_count() == 7 + + batch2 = queue.claim_batch(batch_size=3) + assert len(batch2) == 3 + assert queue.pending_count() == 4 + + # Mark first batch done + queue.mark_done([uid for uid, _ in batch1]) + assert len(queue) == 7 # 3 removed + + batch3 = queue.claim_batch(batch_size=3) + assert len(batch3) == 3 + assert queue.pending_count() == 1 + + batch4 = queue.claim_batch(batch_size=3) + assert len(batch4) == 1 + assert queue.pending_count() == 0 + + # Mark all done + queue.mark_done([uid for uid, _ in batch2 + batch3 + batch4]) + assert len(queue) == 0 + + +def test_item_queue_duplicate_add(tmp_path: Path) -> None: + """Test that duplicate items are skipped.""" + queue = utils.ItemQueue(tmp_path / "queue") + + # Add initial items + added1 = queue.add_items([("a", 1), ("b", 2)]) + assert added1 == 2 + + # Try to add overlapping items + added2 = queue.add_items([("b", 20), ("c", 3)]) # "b" should be skipped + assert added2 == 1 # Only "c" added + assert len(queue) == 3 + + +def test_item_queue_concurrent_queues(tmp_path: Path) -> None: + """Test that multiple queue instances share the same data.""" + folder = tmp_path / "queue" + q1 = utils.ItemQueue(folder) + q2 = utils.ItemQueue(folder) + + # Add items from q1 + q1.add_items([("a", 1), ("b", 2), ("c", 3)]) + + # Both see the same count + assert len(q1) == 3 + assert len(q2) == 3 + + # q1 claims some + claimed1 = q1.claim_batch(batch_size=2) + assert len(claimed1) == 2 + assert q1.pending_count() == 1 + + # q2 sees same state + assert q2.pending_count() == 1 + + # q2 claims the rest + claimed2 = q2.claim_batch(batch_size=10) + assert len(claimed2) == 1 + + # All items accounted for + all_uids = {uid for uid, _ in claimed1} | {uid for uid, _ in claimed2} + assert all_uids == {"a", "b", "c"} + + # Mark all done + q1.mark_done([uid for uid, _ in claimed1]) + q2.mark_done([uid for uid, _ in claimed2]) + assert len(q1) == 0 + + +def test_item_queue_complex_items(tmp_path: Path) -> None: + """Test that complex picklable items work.""" + queue = utils.ItemQueue(tmp_path / "queue") + + # Add complex items + items = [ + ("dict", {"key": "value", "nested": [1, 2, 3]}), + ("list", [1, "two", 3.0]), + ("tuple", (1, 2, 3)), + ] + queue.add_items(items) + + # Claim and verify + claimed = queue.claim_batch() + claimed_dict = {uid: item for uid, item in claimed} + + assert claimed_dict["dict"] == {"key": "value", "nested": [1, 2, 3]} + assert claimed_dict["list"] == [1, "two", 3.0] + assert claimed_dict["tuple"] == (1, 2, 3) + + # Clean up + queue.mark_done([uid for uid, _ in claimed]) + + +def test_item_queue_wait_for_completion(tmp_path: Path) -> None: + """Test wait_for_completion blocks until items are done.""" + import threading + + queue = utils.ItemQueue(tmp_path / "queue") + queue.add_items([("a", 1), ("b", 2)]) + + # Claim items in another thread, then mark done after delay + def worker(): + import time + + claimed = queue.claim_batch() + time.sleep(0.1) + queue.mark_done([uid for uid, _ in claimed]) + + thread = threading.Thread(target=worker) + thread.start() + + # Wait for completion (should block until worker marks done) + queue.wait_for_completion(["a", "b"], poll_interval=0.05) + thread.join() + + assert len(queue) == 0 + + +def test_item_queue_stale_reclaim(tmp_path: Path) -> None: + """Test that stale claimed items are reclaimed based on observed processing time.""" + import time + + # Use a low stale multiplier for faster testing + queue = utils.ItemQueue(tmp_path / "queue", stale_multiplier=2.0) + + # Add two items + queue.add_items([("a", 1), ("b", 2)]) + + # Claim and quickly complete item "a" to establish baseline + claimed_a = queue.claim_batch(batch_size=1) + assert len(claimed_a) == 1 + time.sleep(0.05) # Simulate 50ms processing + queue.mark_done([claimed_a[0][0]]) + + # Now claim item "b" but don't complete it + claimed_b = queue.claim_batch(batch_size=1) + assert len(claimed_b) == 1 + assert queue.pending_count() == 0 + + # Wait for item to become stale (>2x the 50ms baseline = >100ms) + time.sleep(0.15) + + # Reclaim should move it back to pending + reclaimed = queue._reclaim_stale() + assert reclaimed == 1 + assert queue.pending_count() == 1 + + # Can claim again + claimed2 = queue.claim_batch() + assert len(claimed2) == 1 + + +def test_item_queue_no_baseline(tmp_path: Path) -> None: + """Test that without any completed items, stale reclaim is disabled.""" + queue = utils.ItemQueue(tmp_path / "queue") + queue.add_items([("a", 1)]) + + # Claim item but don't complete + claimed = queue.claim_batch() + assert len(claimed) == 1 + + # Reclaim should do nothing (no baseline yet) + reclaimed = queue._reclaim_stale() + assert reclaimed == 0 + assert queue.pending_count() == 0 # Still claimed, not reclaimed diff --git a/exca/utils.py b/exca/utils.py index 0ad8ae03..d707541c 100644 --- a/exca/utils.py +++ b/exca/utils.py @@ -9,7 +9,9 @@ import copy import logging import os +import pickle import shutil +import sqlite3 import sys import typing as tp import uuid @@ -513,3 +515,273 @@ def environment_variables(**kwargs: tp.Any) -> tp.Iterator[None]: for x in kwargs: del os.environ[x] os.environ.update(backup) + + +class ItemQueue: + """SQLite3-based queue for coordinating item processing across concurrent jobs. + + Instead of waiting for other jobs to complete, this allows workers to atomically + claim items from a shared queue. The main process adds items to the queue, + and workers claim them batch by batch. + + Flow: + 1. Main process calls add_items() with all missing (uid, item) pairs + 2. Main process submits worker jobs to cluster + 3. Workers call claim_batch() to get items (marks them as claimed with timestamp) + 4. Workers process items, call mark_done() to remove from queue + (this records processing time for stale detection) + 5. Main process calls wait_for_completion() to block until items are done + (stale items are reclaimed based on observed max processing time) + """ + + # Status constants + PENDING = "pending" + CLAIMED = "claimed" + + def __init__(self, folder: Path | str, stale_multiplier: float = 3.0) -> None: + """ + Parameters + ---------- + folder: Path or str + Directory for the SQLite database + stale_multiplier: float + Multiplier applied to max observed processing time to detect stale items. + An item is stale if claimed_time > max_processing_time * stale_multiplier. + Default: 3.0 (items taking 3x longer than the slowest completed item are stale) + """ + self.folder = Path(folder) + self.folder.mkdir(exist_ok=True, parents=True) + self.db_path = self.folder / "item_queue.db" + self.stale_multiplier = stale_multiplier + self._init_db() + + def _init_db(self) -> None: + """Initialize the SQLite database with the items table and stats.""" + with self._connect() as conn: + conn.execute( + """ + CREATE TABLE IF NOT EXISTS items ( + uid TEXT PRIMARY KEY, + item BLOB NOT NULL, + status TEXT NOT NULL DEFAULT 'pending', + claimed_at REAL + ) + """ + ) + # Stats table to track max processing time + conn.execute( + """ + CREATE TABLE IF NOT EXISTS stats ( + key TEXT PRIMARY KEY, + value REAL NOT NULL + ) + """ + ) + conn.commit() + + @contextlib.contextmanager + def _connect(self) -> tp.Iterator[sqlite3.Connection]: + """Context manager for database connections with proper isolation.""" + conn = sqlite3.connect( + str(self.db_path), timeout=30.0, isolation_level="IMMEDIATE" + ) + try: + yield conn + finally: + conn.close() + + def add_items(self, uid_items: tp.Sequence[tp.Tuple[str, tp.Any]]) -> int: + """Add items to the queue. Called by main process. + + Parameters + ---------- + uid_items: sequence of (uid, item) tuples + Items to add to the queue + + Returns + ------- + int + Number of items actually added (existing items are skipped) + """ + if not uid_items: + return 0 + added = 0 + with self._connect() as conn: + for uid, item in uid_items: + try: + item_blob = pickle.dumps(item) + conn.execute( + "INSERT INTO items (uid, item, status) VALUES (?, ?, ?)", + (uid, item_blob, self.PENDING), + ) + added += 1 + except sqlite3.IntegrityError: + # Already in queue (from concurrent main process), skip + pass + conn.commit() + return added + + def claim_batch(self, batch_size: int = 100) -> tp.List[tp.Tuple[str, tp.Any]]: + """Claim a batch of pending items from the queue. Called by workers. + + Atomically selects pending items and marks them as claimed with timestamp. + + Parameters + ---------- + batch_size: int + Maximum number of items to claim + + Returns + ------- + list of (uid, item) tuples + Items that were claimed (empty if no pending items) + """ + import time + + claimed: tp.List[tp.Tuple[str, tp.Any]] = [] + now = time.time() + with self._connect() as conn: + # Select pending items + cursor = conn.execute( + "SELECT uid, item FROM items WHERE status = ? LIMIT ?", + (self.PENDING, batch_size), + ) + rows = cursor.fetchall() + if not rows: + return claimed + # Mark as claimed with timestamp + uids = [row[0] for row in rows] + placeholders = ",".join("?" * len(uids)) + conn.execute( + f"UPDATE items SET status = ?, claimed_at = ? WHERE uid IN ({placeholders})", + (self.CLAIMED, now, *uids), + ) + conn.commit() + # Deserialize items + for uid, item_blob in rows: + item = pickle.loads(item_blob) + claimed.append((uid, item)) + return claimed + + def mark_done(self, uids: tp.Sequence[str]) -> None: + """Mark items as done by removing them from the queue. + + Called by workers after successfully caching processed items. + Also updates the max processing time for stale detection. + + Parameters + ---------- + uids: sequence of str + UIDs of items to remove from queue + """ + if not uids: + return + import time + + now = time.time() + with self._connect() as conn: + # Get claimed_at times to compute processing durations + placeholders = ",".join("?" * len(uids)) + cursor = conn.execute( + f"SELECT claimed_at FROM items WHERE uid IN ({placeholders}) AND claimed_at IS NOT NULL", + list(uids), + ) + claimed_times = [row[0] for row in cursor.fetchall()] + + # Update max processing time if we have valid times + if claimed_times: + max_duration = max(now - claimed_at for claimed_at in claimed_times) + conn.execute( + """ + INSERT INTO stats (key, value) VALUES ('max_processing_time', ?) + ON CONFLICT(key) DO UPDATE SET value = MAX(value, ?) + """, + (max_duration, max_duration), + ) + + # Delete the items + conn.execute(f"DELETE FROM items WHERE uid IN ({placeholders})", list(uids)) + conn.commit() + + def _reclaim_stale(self) -> int: + """Reclaim items that have been claimed for too long (worker probably crashed). + + Uses the max observed processing time * stale_multiplier as the threshold. + Only reclaims if at least one item has completed (so we have a baseline). + + Returns the number of items reclaimed. + """ + import time + + with self._connect() as conn: + # Get max processing time from stats + cursor = conn.execute( + "SELECT value FROM stats WHERE key = 'max_processing_time'" + ) + row = cursor.fetchone() + if row is None: + # No items completed yet, can't determine stale threshold + return 0 + + max_processing_time = row[0] + stale_threshold = max_processing_time * self.stale_multiplier + cutoff = time.time() - stale_threshold + + cursor = conn.execute( + "UPDATE items SET status = ?, claimed_at = NULL " + "WHERE status = ? AND claimed_at < ?", + (self.PENDING, self.CLAIMED, cutoff), + ) + conn.commit() + return cursor.rowcount + + def wait_for_completion( + self, uids: tp.Sequence[str], poll_interval: float = 1.0 + ) -> None: + """Block until all specified items are no longer in the queue. + + Automatically reclaims stale items (claimed too long ago) so they can + be retried by other workers. + + Parameters + ---------- + uids: sequence of str + UIDs to wait for + poll_interval: float + Seconds to wait between checks + """ + import time + + uids_set = set(uids) + while True: + # Reclaim any stale items first + reclaimed = self._reclaim_stale() + if reclaimed > 0: + logger.warning( + "Reclaimed %s stale items (worker likely crashed)", reclaimed + ) + + with self._connect() as conn: + placeholders = ",".join("?" * len(uids_set)) + cursor = conn.execute( + f"SELECT COUNT(*) FROM items WHERE uid IN ({placeholders})", + list(uids_set), + ) + remaining = cursor.fetchone()[0] + if remaining == 0: + return + time.sleep(poll_interval) + + def __len__(self) -> int: + """Return the number of items remaining in the queue (pending + claimed).""" + with self._connect() as conn: + cursor = conn.execute("SELECT COUNT(*) FROM items") + return cursor.fetchone()[0] + + def pending_count(self) -> int: + """Return the number of pending (unclaimed) items.""" + with self._connect() as conn: + cursor = conn.execute( + "SELECT COUNT(*) FROM items WHERE status = ?", (self.PENDING,) + ) + return cursor.fetchone()[0]