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
244 changes: 191 additions & 53 deletions src/campaign/README.md

Large diffs are not rendered by default.

137 changes: 107 additions & 30 deletions src/campaign/campaign_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,9 +14,10 @@
Pass 2 — fill remaining capacity up to ``max_replicas`` (highest priority
first).

A group becomes *eligible* when each dependency group has either explicitly
called ``signal_ready()`` (workflow-driven) or has ``dep_threshold`` or more
finished replicas (count-based fallback, default 1).
A group becomes *eligible* either when a parent workflow calls
``_trigger_dependent()`` (explicit activation) or when each dependency group
has ``dep_threshold`` or more finished replicas / has called ``_signal_done()``
(count-based fallback, default 1).

Usage
-----
Expand Down Expand Up @@ -122,14 +123,16 @@ class BaseWorkflow:
def __init__(
self,
config: Optional[dict] = None,
on_ready: Optional[object] = None,
_cm: Optional[object] = None,
_group_name: Optional[str] = None,
asyncflow: Optional[object] = None,
policies: Optional[list] = None,
engine_dragon: Optional[object] = None,
) -> None:
self.config = config
# Async callable injected by AsyncCampaignManager; invoke via _signal_ready().
self._on_ready = on_ready
# AsyncCampaignManager reference injected at construction.
self._cm = _cm
self._group_name = _group_name
# Shared WorkflowEngine injected by AsyncCampaignManager (optional).
self.asyncflow = asyncflow
# One Dragon Policy per assigned GPU, injected by AsyncCampaignManager.
Expand All @@ -138,20 +141,34 @@ def __init__(
# Dragon backend handle for routing tasks to specific backends.
self.engine_dragon: Optional[object] = engine_dragon

async def _signal_ready(self) -> None:
async def _trigger_dependent(
self,
name: str,
replicas: int = 1,
**kwargs,
) -> None:
"""
Tell the CM to activate a dependent workflow group with *replicas* replicas.

The group must already be registered (via config or register_group) with
``replicas=0``. Calling this mid-run is the canonical way to start a
workflow that depends on data produced by the current workflow.
No-op when no CM was injected.
"""
Signal the CM that this workflow has produced enough data for its
dependents to start.
if self._cm is not None:
await self._cm.trigger_dependent(name, replicas=replicas, **kwargs)

Calls the ``on_ready`` coroutine injected by the CM at construction.
No-op if no callback was provided.
async def _signal_done(self) -> None:
"""
import asyncio
Signal the CM that this workflow has finished producing data for
its dependents (count-based dependency fallback).

if self._on_ready is not None:
result = self._on_ready()
if asyncio.iscoroutine(result):
await result
Marks this group ``ready=True`` in the CM, which unblocks any groups
that list this one as a dependency with ``dep_threshold > finished_replicas``.
No-op when no CM was injected.
"""
if self._cm is not None and self._group_name is not None:
await self._cm.signal_done(self._group_name)

def run(self, replica_id: str) -> None:
"""
Expand Down Expand Up @@ -387,10 +404,14 @@ def from_config(
cm._log.warning(f"from_config: no class registered for {name!r} — skipping")
continue

# Groups with dependencies default to replicas=0 — they remain
# inactive until a parent workflow calls _trigger_dependent().
has_deps = bool(wf_cfg.get("dependencies", []))
default_replicas = 0 if has_deps else 1
cm.register_group(
name=name,
workflow_class=wf_class,
replicas=int(wf_cfg.get("replicas", 1)),
replicas=int(wf_cfg.get("replicas", default_replicas)),
dependencies=list(wf_cfg.get("dependencies", [])),
dep_threshold=int(wf_cfg.get("dependency_threshold", 1)),
priority=int(wf_cfg.get("priority", 0)),
Expand Down Expand Up @@ -637,22 +658,65 @@ async def close(self) -> None:
# Monitoring
# ------------------------------------------------------------------

async def signal_ready(self, group_name: str) -> None:
async def signal_done(self, group_name: str) -> None:
"""
Mark *group_name* as having produced enough data for its dependents.
Signal that *group_name* has produced output and downstream work should run.

Called by a workflow via ``self._signal_ready()`` (injected at construction).
Once set, the flag overrides the ``dep_threshold`` replica-count check so
dependent groups are unblocked immediately.
Called by a workflow via ``self._signal_done()`` — can be called multiple
times per replica (e.g. once per iteration). Each call queues **1 more
replica** in every group that lists *group_name* in its ``dependencies``.
The CM routes the signal to the right downstream groups automatically, so
the calling workflow does not need to know their names.

Idempotent — subsequent calls for the same group are no-ops.
Also sets ``group.ready = True`` to satisfy count-based dep checks.
"""
async with self._lock:
group = self._groups.get(group_name)
if group is None or group.ready:
if group is None:
return
group.ready = True
self._log.info(f"Group {group_name!r} signaled ready — unblocking dependents")
dependents = [g for g in self._groups.values() if group_name in g.dependencies]
for dep in dependents:
dep.replicas += 1
dep.configured_replicas += 1
if dep.status == "done":
dep.status = "pending"
if dependents:
self._log.info(
f"{group_name!r} signaled done → +1 replica for {[d.name for d in dependents]}"
)
await self._schedule()

async def trigger_dependent(
self,
name: str,
replicas: int = 1,
config: Optional[dict] = None,
) -> None:
"""
Queue *replicas* more runs of the dependent group *name*.

Called by a parent workflow (via ``self._trigger_dependent()``) each
time its execution logic decides to launch downstream work. May be
called multiple times — each call adds *replicas* to the group's
total and re-opens the group for scheduling if it had already finished.

The group must be pre-registered (via config or ``register_group``).
"""
async with self._lock:
group = self._groups.get(name)
if group is None:
self._log.warning(f"trigger_dependent: group {name!r} not registered — ignoring")
return
group.replicas += replicas
group.configured_replicas += replicas
if group.status == "done":
group.status = "pending"
if config:
group.group_config = {**(group.group_config or {}), **config}
self._log.info(
f"trigger_dependent: {name!r} +{replicas} replicas (total={group.replicas})"
)
await self._schedule()

async def add_replicas(self, group_name: str, n: int = 1) -> None:
Expand Down Expand Up @@ -739,7 +803,11 @@ def _can_start_locked(self, group: _GroupInfo) -> bool:
return False
if group.started_count >= group.replicas:
return False
if group.running_count >= group.max_replicas:
# max_replicas == 0 means "no explicit cap — use replicas count".
# This handles dependent groups registered with replicas=0 that later
# receive replicas via trigger_dependent() or signal_done().
effective_max = group.max_replicas if group.max_replicas > 0 else group.replicas
if group.running_count >= effective_max:
return False
if not self._deps_satisfied_locked(group):
return False
Expand Down Expand Up @@ -816,7 +884,7 @@ def _schedule_locked(self) -> list[tuple["_GroupInfo", int]]:
for g in eligible:
if (
g.started_count < g.replicas
and g.running_count < g.max_replicas
and g.running_count < (g.max_replicas if g.max_replicas > 0 else g.replicas)
and self._deps_satisfied_locked(g)
and not self._resources.can_fit(g.required_cpus, g.required_gpus)
):
Expand Down Expand Up @@ -901,7 +969,8 @@ async def _run_replica(self, group: _GroupInfo, replica_idx: int) -> None:
}
wf = group.workflow_class(
config=replica_config,
on_ready=lambda: self.signal_ready(group.name),
_cm=self,
_group_name=group.name,
asyncflow=self._asyncflow,
policies=policies,
engine_dragon=self._engine_dragon,
Expand Down Expand Up @@ -1005,7 +1074,13 @@ async def _on_replica_finished(self, group: _GroupInfo, replica_id: str) -> None
await self._schedule()

async with self._lock:
all_done = bool(self._groups) and all(g.status == "done" for g in self._groups.values())
# Groups registered with replicas=0 (triggered dependents not yet activated)
# are excluded from the completion check — they only count once triggered.
all_done = (
bool(self._groups)
and all(g.status == "done" or g.replicas == 0 for g in self._groups.values())
and any(g.replicas > 0 for g in self._groups.values())
)

if all_done:
self._all_done.set()
Expand Down Expand Up @@ -1131,10 +1206,12 @@ def from_config(
wf_class = workflow_registry.get(name)
if wf_class is None:
continue
has_deps = bool(wf_cfg.get("dependencies", []))
default_replicas = 0 if has_deps else 1
cm.register_group(
name=name,
workflow_class=wf_class,
replicas=int(wf_cfg.get("replicas", 1)),
replicas=int(wf_cfg.get("replicas", default_replicas)),
dependencies=list(wf_cfg.get("dependencies", [])),
dep_threshold=int(wf_cfg.get("dependency_threshold", 1)),
priority=int(wf_cfg.get("priority", 0)),
Expand Down
85 changes: 61 additions & 24 deletions tests/test_campaign_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,16 +55,27 @@ async def run(self, replica_id: str) -> None:
RecordingWorkflow.ran.append(replica_id)


class SignalWorkflow(BaseWorkflow):
"""Fires _signal_ready() immediately then finishes after a tiny sleep."""
class SignalDoneWorkflow(BaseWorkflow):
"""Fires _signal_done() immediately then finishes after a tiny sleep."""

workflow_id = "signal"
workflow_id = "signal_done"

async def run(self, replica_id: str) -> None:
await self._signal_ready()
await self._signal_done()
await asyncio.sleep(0.01)


class TriggerWorkflow(BaseWorkflow):
"""Triggers a dependent group named 'downstream' then finishes."""

workflow_id = "trigger"
dependent_name: str = "downstream"
dependent_replicas: int = 1

async def run(self, replica_id: str) -> None:
await self._trigger_dependent(self.dependent_name, replicas=self.dependent_replicas)


class HookWorkflow(BaseWorkflow):
"""Records (replica_id, final_state) tuples in on_replica_done."""

Expand Down Expand Up @@ -149,25 +160,25 @@ async def _fake_init():


class TestBaseWorkflow:
async def test_signal_ready_no_callback_is_noop(self):
async def test_signal_done_no_cm_is_noop(self):
wf = NullWorkflow()
await wf._signal_ready() # must not raise
await wf._signal_done() # must not raise

async def test_signal_ready_calls_sync_callback(self):
called = []
wf = NullWorkflow(on_ready=lambda: called.append(1))
await wf._signal_ready()
assert called == [1]

async def test_signal_ready_awaits_async_callback(self):
called = []
async def test_trigger_dependent_no_cm_is_noop(self):
wf = NullWorkflow()
await wf._trigger_dependent("some_group", replicas=2) # must not raise

async def cb():
called.append(1)
async def test_signal_done_calls_cm(self):
cm_mock = AsyncMock()
wf = NullWorkflow(_cm=cm_mock, _group_name="mygroup")
await wf._signal_done()
cm_mock.signal_done.assert_awaited_once_with("mygroup")

wf = NullWorkflow(on_ready=cb)
await wf._signal_ready()
assert called == [1]
async def test_trigger_dependent_calls_cm(self):
cm_mock = AsyncMock()
wf = NullWorkflow(_cm=cm_mock)
await wf._trigger_dependent("dep", replicas=3)
cm_mock.trigger_dependent.assert_awaited_once_with("dep", replicas=3)

def test_base_run_raises_not_implemented(self):
wf = BaseWorkflow()
Expand Down Expand Up @@ -258,9 +269,9 @@ async def run(self, replica_id: str) -> None:
b_idx = next(i for i, (wf, _) in enumerate(order) if wf == "B")
assert all(wf == "A" for wf, _ in order[:b_idx])

async def test_dependency_via_signal_ready(self, acm):
"""_signal_ready() unblocks B even before all of A's replicas finish."""
acm.register_group("a", SignalWorkflow, replicas=1)
async def test_dependency_via_signal_done(self, acm):
"""_signal_done() unblocks B even before all of A's replicas finish."""
acm.register_group("a", SignalDoneWorkflow, replicas=1)
acm.register_group(
"b",
NullWorkflow,
Expand All @@ -275,6 +286,25 @@ async def test_dependency_via_signal_ready(self, acm):
assert s["groups"]["a"]["ready"] is True
assert s["groups"]["b"]["status"] == "done"

async def test_trigger_dependent_activates_group(self, acm):
"""Parent workflow calls _trigger_dependent to start a replicas=0 group."""
acm.register_group("upstream", TriggerWorkflow, replicas=1)
acm.register_group("downstream", RecordingWorkflow, replicas=0)
await acm.start()
assert await acm.wait(timeout=3.0)

s = acm.status()
assert s["groups"]["downstream"]["status"] == "done"
assert "downstream_0" in RecordingWorkflow.ran

async def test_untriggered_group_does_not_block_completion(self, acm):
"""A replicas=0 group that is never triggered must not prevent _all_done."""
acm.register_group("a", NullWorkflow, replicas=1)
acm.register_group("never_triggered", NullWorkflow, replicas=0)
await acm.start()
assert await acm.wait(timeout=3.0)
assert acm.status()["groups"]["a"]["status"] == "done"

async def test_on_replica_done_hook_called(self, acm):
acm.register_group("a", HookWorkflow, replicas=2)
await acm.start()
Expand Down Expand Up @@ -328,14 +358,16 @@ async def test_from_config_registers_groups(self):
config = {
"workflows": {
"x": {"replicas": 2, "max_replicas": 1, "priority": 3},
"y": {"replicas": 1, "dependencies": ["x"], "dependency_threshold": 2},
# y has dependencies → replicas defaults to 0 (triggered group)
"y": {"dependencies": ["x"], "dependency_threshold": 2},
}
}
cm = AsyncCampaignManager.from_config(config, {"x": NullWorkflow, "y": NullWorkflow})
s = cm.status()["groups"]
assert s["x"]["replicas_total"] == 2
assert s["x"]["max_replicas"] == 1
assert s["x"]["priority"] == 3
assert s["y"]["replicas_total"] == 0 # triggered group: not yet activated
assert s["y"]["dependencies"] == ["x"]
assert s["y"]["dep_threshold"] == 2

Expand Down Expand Up @@ -443,15 +475,20 @@ def test_from_config_registers_groups(self):
config = {
"workflows": {
"alpha": {"replicas": 3, "max_replicas": 2, "priority": 7},
# beta has dependencies → replicas defaults to 0 (triggered group)
"beta": {"dependencies": ["alpha"]},
}
}
cm = CampaignManager.from_config(config, {"alpha": SyncRecordingWorkflow})
cm = CampaignManager.from_config(
config, {"alpha": SyncRecordingWorkflow, "beta": SyncRecordingWorkflow}
)
s = cm.status()["groups"]
cm.close()
assert "alpha" in s
assert s["alpha"]["replicas_total"] == 3
assert s["alpha"]["max_replicas"] == 2
assert s["alpha"]["priority"] == 7
assert s["beta"]["replicas_total"] == 0 # triggered group: not yet activated

def test_unknown_group_skipped_in_from_config(self):
config = {"workflows": {"ghost": {"replicas": 1}}}
Expand Down
Loading
Loading