Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
69 changes: 69 additions & 0 deletions src/malwar/core/concurrency.py
Original file line number Diff line number Diff line change
@@ -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
45 changes: 19 additions & 26 deletions src/malwar/monitor/snapshot.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand All @@ -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:
Expand All @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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(
Expand Down
125 changes: 125 additions & 0 deletions tests/unit/test_concurrency.py
Original file line number Diff line number Diff line change
@@ -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"
Loading