From 0d3f5081e332c74dc9aeb65456628e5d1acf05bb Mon Sep 17 00:00:00 2001 From: Claude Date: Fri, 11 Sep 2026 16:40:24 +0000 Subject: [PATCH] fix(monitor): bound pending tasks, not just concurrent scans (#57) build_snapshot admitted one asyncio Task per skill to gather() and bounded only how many ran, via a semaphore inside the worker. Live Tasks therefore scaled with the registry rather than with the concurrency limit. Reproduced independently at the reporter's parameters and at the sweep's real ceiling of --max-scans 12000: n=512 gather-all tasks_alive=513 peak_active=8 707 KiB n=512 worker-pool tasks_alive=9 peak_active=8 26 KiB n=12000 gather-all tasks_alive=12001 peak_active=8 16,579 KiB n=12000 worker-pool tasks_alive=9 peak_active=8 477 KiB So ~16 MiB of Tasks to keep 8 of them busy. Not a correctness or security problem -- the concurrency limit worked as documented, and no extra load ever reached the registry -- but it is waste that grows with the registry, and the registry is the thing that grows. run_bounded inverts the structure: a fixed pool pulls from a queue, so live Tasks equal the limit. Applied to all three sites (scan, escalate, enrich), not just the one reported; the escalation and enrichment phases had the same shape. On failure the pool is cancelled rather than left running, so a sweep that is already failing stops hitting the registry. The tests were checked against the old implementation and fail there. That matters here: peak concurrency was already correct before this change, so a test asserting only the limit would have passed against the bug. Reported by @Nievesjyl. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01DNoTXU8k3pfSBzR7aJubqL --- src/malwar/core/concurrency.py | 69 ++++++++++++++++++ src/malwar/monitor/snapshot.py | 45 +++++------- tests/unit/test_concurrency.py | 125 +++++++++++++++++++++++++++++++++ 3 files changed, 213 insertions(+), 26 deletions(-) create mode 100644 src/malwar/core/concurrency.py create mode 100644 tests/unit/test_concurrency.py diff --git a/src/malwar/core/concurrency.py b/src/malwar/core/concurrency.py new file mode 100644 index 0000000..b550f05 --- /dev/null +++ b/src/malwar/core/concurrency.py @@ -0,0 +1,69 @@ +"""Bounded concurrent execution over an iterable of work items. + +``asyncio.gather(*(worker(x) for x in items))`` wraps every coroutine in a Task +immediately, so a semaphore *inside* the worker bounds how many run but not how +many exist. Over a 12,000-skill sweep that is 12,000 live Tasks holding ~16 MiB +to keep 8 of them busy. + +:func:`run_bounded` inverts it: a fixed pool of workers pulls from a queue, so +the number of live Tasks is the concurrency limit rather than the input size. + +Reported by @Nievesjyl in Ap6pack/malwar#57, and reproduced independently: + + n=12000 gather-all tasks_alive=12001 peak_active=8 16,579 KiB + n=12000 worker-pool tasks_alive=9 peak_active=8 477 KiB +""" + +from __future__ import annotations + +import asyncio +from collections.abc import Awaitable, Callable, Iterable +from typing import TypeVar + +T = TypeVar("T") + + +async def run_bounded( + items: Iterable[T], + worker: Callable[[T], Awaitable[None]], + *, + concurrency: int, +) -> None: + """Await ``worker(item)`` for every item, at most ``concurrency`` at a time. + + Unlike ``gather`` with an internal semaphore, only ``concurrency`` Tasks + exist at any moment, so memory is flat in the size of the input. + + Exception semantics match ``asyncio.gather`` without ``return_exceptions``: + the first failure propagates. Items not yet started are dropped, and the + other in-flight workers are cancelled rather than left running. Callers in + this codebase catch inside their own worker so one bad skill cannot abort a + sweep; this only covers a genuinely unexpected failure. + + Order is not preserved. Nothing here depends on it -- each worker writes to + a dict keyed by slug -- and requiring it would mean holding results. + """ + queue: asyncio.Queue[T] = asyncio.Queue() + for item in items: + queue.put_nowait(item) + if queue.empty(): + return + + async def _pump() -> None: + while True: + try: + item = queue.get_nowait() + except asyncio.QueueEmpty: + return + await worker(item) + + pool = [asyncio.create_task(_pump()) for _ in range(max(1, concurrency))] + try: + await asyncio.gather(*pool) + except BaseException: + # Stop the rest of the pool instead of leaving workers running against + # a sweep that is already failing. + for task in pool: + task.cancel() + await asyncio.gather(*pool, return_exceptions=True) + raise diff --git a/src/malwar/monitor/snapshot.py b/src/malwar/monitor/snapshot.py index 476c1ef..ffbae85 100644 --- a/src/malwar/monitor/snapshot.py +++ b/src/malwar/monitor/snapshot.py @@ -27,6 +27,7 @@ from pathlib import Path from typing import NamedTuple +from malwar.core.concurrency import run_bounded from malwar.crawl.client import ClawHubClient from malwar.monitor.escalation import ( EscalationBackend, @@ -479,18 +480,16 @@ async def build_snapshot( contents: dict[str, str] = {} # --- Phase 1: rule-scan every skill (fast, free, one request each) --- - semaphore = asyncio.Semaphore(max(1, concurrency)) done = len(snapshot.skills) lock = asyncio.Lock() async def _worker(meta: SkillMeta) -> None: nonlocal done - async with semaphore: - try: - record, content = await _scan_slug(client, meta) - except Exception as exc: # last-resort guard; never abort the sweep - record = SkillRecord(slug=meta.slug, verdict="UNKNOWN", error=f"worker: {exc}") - content = None + try: + record, content = await _scan_slug(client, meta) + except Exception as exc: # last-resort guard; never abort the sweep + record = SkillRecord(slug=meta.slug, verdict="UNKNOWN", error=f"worker: {exc}") + content = None async with lock: snapshot.skills[meta.slug] = record if record.error: @@ -501,7 +500,7 @@ async def _worker(meta: SkillMeta) -> None: if on_progress is not None: on_progress(done, total, meta.slug) - await asyncio.gather(*(_worker(meta) for meta in to_scan)) + await run_bounded(to_scan, _worker, concurrency=concurrency) # --- Phase 2: targeted second opinion on the ambiguous band --- if keep_content and contents: @@ -516,15 +515,12 @@ async def _worker(meta: SkillMeta) -> None: backend.name, ) - esc_sem = asyncio.Semaphore(max(1, concurrency)) - async def _escalate(slug: str) -> None: - async with esc_sem: - try: - res = await backend.assess(contents[slug], file_name=f"{slug}/SKILL.md") - except Exception as exc: # a bad escalation must not abort the sweep - logger.warning("escalation failed for %s: %s", slug, exc) - return + try: + res = await backend.assess(contents[slug], file_name=f"{slug}/SKILL.md") + except Exception as exc: # a bad escalation must not abort the sweep + logger.warning("escalation failed for %s: %s", slug, exc) + return rec = snapshot.skills[slug] rec.escalation_backend = res.backend rec.escalation_verdict = res.verdict @@ -537,7 +533,7 @@ async def _escalate(slug: str) -> None: if res.score is not None: rec.risk_score = round(res.score * 100) - await asyncio.gather(*(_escalate(slug) for slug in candidates)) + await run_bounded(candidates, _escalate, concurrency=concurrency) # --- Fail-safe: downgrade unverified fragile-MALICIOUS verdicts --- # A MALICIOUS verdict resting on a single high-false-positive rule (see @@ -610,15 +606,12 @@ async def _escalate(slug: str) -> None: backfill.sort(key=lambda item: -verdict_rank(item[1].verdict)) flagged.extend(slug for slug, _ in backfill[:enrich_backfill_limit]) if flagged: - enrich_sem = asyncio.Semaphore(max(1, concurrency)) - async def _enrich(slug: str) -> None: - async with enrich_sem: - try: - detail = await client.get_skill(slug) - except Exception as exc: # never fatal; attribution is a bonus - logger.debug("detail fetch failed for %s: %s", slug, exc) - return + try: + detail = await client.get_skill(slug) + except Exception as exc: # never fatal; attribution is a bonus + logger.debug("detail fetch failed for %s: %s", slug, exc) + return rec = snapshot.skills[slug] rec.detail_fetched = True if detail.owner is not None: @@ -654,7 +647,7 @@ async def _enrich(slug: str) -> None: rec.moderation_engine = "" rec.moderation_scanned_at = None - await asyncio.gather(*(_enrich(slug) for slug in flagged)) + await run_bounded(flagged, _enrich, concurrency=concurrency) enriched = sum(1 for s in flagged if snapshot.skills[s].moderation_checked) attributed = sum(1 for s in flagged if snapshot.skills[s].publisher) unblocked = sum( diff --git a/tests/unit/test_concurrency.py b/tests/unit/test_concurrency.py new file mode 100644 index 0000000..2bb8fca --- /dev/null +++ b/tests/unit/test_concurrency.py @@ -0,0 +1,125 @@ +"""Tests for bounded concurrent execution. + +Ap6pack/malwar#57 reported that ``build_snapshot`` admitted one Task per skill +to ``asyncio.gather`` and bounded only how many *ran*, so live Tasks scaled +with the registry rather than with the concurrency limit. Reproduced +independently at the reporter's parameters (512 items, limit 8) and at the +sweep's real ceiling (12,000). + +The load-bearing assertion is the second one in each pair: peak concurrency was +already correct under the old code, so a test that only checked it would have +passed against the bug. +""" + +from __future__ import annotations + +import asyncio + +import pytest + +from malwar.core.concurrency import run_bounded + + +async def _tracking_worker(limit_box: dict[str, int], gate: asyncio.Event): + """Worker that records peak concurrency and blocks until released.""" + state = {"active": 0} + + async def worker(_item: int) -> None: + state["active"] += 1 + limit_box["peak"] = max(limit_box["peak"], state["active"]) + await gate.wait() + state["active"] -= 1 + + return worker + + +class TestTaskCountIsBoundedByConcurrency: + @pytest.mark.parametrize(("n", "concurrency"), [(512, 8), (12_000, 8), (100, 4)]) + async def test_live_tasks_track_the_limit_not_the_input(self, n, concurrency): + box = {"peak": 0} + gate = asyncio.Event() + worker = await _tracking_worker(box, gate) + + before = len(asyncio.all_tasks()) + runner = asyncio.create_task( + run_bounded(range(n), worker, concurrency=concurrency) + ) + await asyncio.sleep(0.05) + live = len(asyncio.all_tasks()) - before + + assert box["peak"] == concurrency, "concurrency limit not respected" + # The regression itself: under gather-over-all this was n + 1. + assert live <= concurrency + 2, ( + f"{live} live tasks for {n} items at concurrency {concurrency}; " + "task count must not scale with input size" + ) + + gate.set() + await runner + + +class TestEveryItemRuns: + async def test_all_items_are_processed_exactly_once(self): + seen: list[int] = [] + + async def worker(item: int) -> None: + await asyncio.sleep(0) + seen.append(item) + + await run_bounded(range(250), worker, concurrency=7) + assert sorted(seen) == list(range(250)) + assert len(seen) == len(set(seen)), "an item ran more than once" + + async def test_empty_input_is_a_no_op(self): + async def worker(_item: int) -> None: # pragma: no cover - must not run + raise AssertionError("worker ran for an empty input") + + await run_bounded([], worker, concurrency=4) + + async def test_fewer_items_than_workers(self): + seen: list[int] = [] + + async def worker(item: int) -> None: + seen.append(item) + + await run_bounded([1, 2], worker, concurrency=16) + assert sorted(seen) == [1, 2] + + async def test_zero_concurrency_still_makes_progress(self): + # A misconfigured limit must not deadlock a sweep. + seen: list[int] = [] + + async def worker(item: int) -> None: + seen.append(item) + + await run_bounded([1, 2, 3], worker, concurrency=0) + assert sorted(seen) == [1, 2, 3] + + +class TestFailureSemantics: + async def test_first_exception_propagates(self): + async def worker(item: int) -> None: + if item == 3: + raise ValueError("boom") + await asyncio.sleep(0) + + with pytest.raises(ValueError, match="boom"): + await run_bounded(range(20), worker, concurrency=4) + + async def test_pool_does_not_outlive_a_failure(self): + # A worker still running after the call returned would keep hitting the + # registry for a sweep that has already failed. + running = {"count": 0} + + async def worker(item: int) -> None: + if item == 0: + raise ValueError("boom") + running["count"] += 1 + await asyncio.sleep(5) + running["count"] -= 1 + + before = len(asyncio.all_tasks()) + with pytest.raises(ValueError): + await run_bounded(range(50), worker, concurrency=8) + await asyncio.sleep(0.05) + assert len(asyncio.all_tasks()) <= before, "workers left running after failure"