From 716c76b318e71a69452a458ccc6d8819ab6fdc18 Mon Sep 17 00:00:00 2001 From: Zijian Zhang <35801754+futrime@users.noreply.github.com> Date: Fri, 25 Sep 2026 12:36:28 +0000 Subject: [PATCH] fix(flowing): close sessions a flow can no longer reach The flow API has no `Session.close`, and the engine closed a session only as the flow call that opened it ended, holding every view strongly until then. A loop opening a fresh session a round -- ralph_loop, flame_chase, rlar's reviewer, lane runtimes, for days -- kept every one of them open, and a stream harness keeps a CLI process per session. - The engine now holds a SessionView only weakly, through an `Opened` record (a `weakref.ref` subclass) the call's cleanup and the agent's hooks keep instead. A view let go of has its session closed on the run's loop -- marshalled there from whichever thread let go of it or a collection found it -- in an eagerly started task the call's and the run's cleanup wait for; the next spawn drains what was let go of first, so a flow that never yields holds few open too. A session is closed exactly once whichever of that and its call's end comes first, and one handed up to a caller still closes as the call that opened it ends. `Call.res` is keyed by id, so a session closed early leaves it. - A hook heard for a session let go of gets a closed stand-in; a fork keeps the session it was cut from open until its own first turn succeeds; a hook failing between turns no longer ties its session into a reference cycle through its own frame; `steer` refuses a session that is over. - A turn's hard deadline is kept when a sooner graceful one is what its driver is told: `Call.limits` also answers the hard deadline, and the engine interrupts the turn there and raises `DurationExceeded`; a turn hard for its cost or tokens is told the hard deadline, not the graceful one. A graceful deadline's cancel is sent from the loop, so a flow that returns as its last turn ends no longer leaks a bare CancelledError to its caller. - Fakes: `FakeAgentDriver.live` and `.peak` count open sessions; a fake turn held open by its reply is cut at a hard deadline, as a real one is; `FakeEnvDriver` lets go of temporary-copy holds as the run that took them closes (or as the root fake closes), so a run resumed on the same fake takes its copy again instead of raising TempCloneBusy. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/features/budgets.md | 3 +- docs/reference/flows.md | 11 +- docs/weaver/testing-flows.md | 13 +- specs/runtime/flowing.md | 10 + src/hmz/flows/agents.py | 7 +- src/hmz/runtime/flowing/engine.py | 131 ++++-- src/hmz/runtime/flowing/fakes.py | 124 ++++- src/hmz/runtime/flowing/viewing.py | 263 ++++++++++- tests/unit/flows/test_engine_budget.py | 176 ++++++++ tests/unit/flows/test_engine_resume.py | 33 ++ tests/unit/flows/test_engine_scale.py | 40 ++ tests/unit/flows/test_engine_sessions.py | 549 +++++++++++++++++++++++ tests/unit/flows/test_fakes.py | 24 + 13 files changed, 1304 insertions(+), 80 deletions(-) create mode 100644 tests/unit/flows/test_engine_sessions.py diff --git a/docs/features/budgets.md b/docs/features/budgets.md index bb190e74..8c1c9bb5 100644 --- a/docs/features/budgets.md +++ b/docs/features/budgets.md @@ -69,7 +69,8 @@ that does not count as it goes. A graceful per-turn budget is a budget for the run's accounting rather than a way to shorten an answer, so a turn meant to be cut off says `graceful=False`. Where several budgets are over one turn — its own, its flow's, the run's — each limit is the least any of them leaves, and it is -hard wherever the budget that sets it is. +hard wherever the budget that sets it is. A hard deadline holds under a sooner graceful one too: +the turn runs on past the graceful deadline, and is cut off at the hard one. **A turn cut off did what it did.** Its edits are on disk and its conversation is open to the next turn, so the round after a short round carries the same session on rather than starting diff --git a/docs/reference/flows.md b/docs/reference/flows.md index 6b508c7d..35a369f1 100644 --- a/docs/reference/flows.md +++ b/docs/reference/flows.md @@ -389,9 +389,14 @@ Antigravity and dsh do not fork, and raise `UnsupportedOperation`. A fork is cut its first turn, so it is refused then if the session it came from has taken a turn since. See [Branching a conversation](/weaver/branching). -**Every session a flow call opened is closed when that call ends**, and every call it started -has. A turn that is cancelled — a `TaskGroup` sibling failing, a deadline, ctrl+c — interrupts -the CLI rather than leaving it running. +**A session is closed when the flow call that opened it ends** — and every call it started has — +**or as soon as nothing holds it any more**, whichever comes first. There is no `close`: a loop +that opens a fresh session a round holds one or two open however many rounds it runs, and a +session kept in a variable, a list or a dict stays open for as long as it is kept. One handed +back to a caller is closed all the same as the call that opened it ends. A fork keeps the +session it was forked from open until its own first turn, which is where it is cut. A turn that +is cancelled — a `TaskGroup` sibling failing, a deadline, ctrl+c — interrupts the CLI rather +than leaving it running. ## Where each agent works diff --git a/docs/weaver/testing-flows.md b/docs/weaver/testing-flows.md index 385719db..bd18114b 100644 --- a/docs/weaver/testing-flows.md +++ b/docs/weaver/testing-flows.md @@ -171,7 +171,9 @@ Afterwards the driver says what happened: `.prompts` is every prompt any of its given, hooks' additions included, in order; `.sessions` is every session it opened, each a `FakeSession` with its own `.prompts`, the `.requests` it was asked for with their limits, what it was `.steered` with, the `.tools` its answers reached for, whether it is `.closed`, and -the session it was `.forked_from`. +the session it was `.forked_from`. `.live` is how many of its sessions are open now and `.peak` +the most that were open at once, which is how a test sees a loop let go of the sessions it is +done with. `FakeAgentDriver("codex")` is Codex: it serves exactly what Codex serves, so a flow that asks for more is refused the way it would be with the real one: @@ -199,6 +201,9 @@ async def test_the_budget_stops_it() -> None: local=FakeEnvDriver(run=GREEN)) ``` +A turn held open by its reply — one waiting on `until_steered()` — is cut off at a hard +deadline, as a real one is, and raises `DurationExceeded`. + ## Script an environment `FakeEnvDriver(files)` is a working directory held in a dictionary: what is in it to start @@ -227,8 +232,10 @@ Afterwards it says what happened: `.files` is what is under the workdir now, `.t file as text, `.commands` every command run there, `.machine` every file on the fake machine — copies and worktrees included — and `.clones` and `.scratches` the temporary copies and scratch directories still there, which is how a test sees them [cleaned -up](/weaver/worktrees#how-long-they-last). `refs=` is the git refs `derive_worktree` knows, and -`repo=False` a workdir that is not a repository. +up](/weaver/worktrees#how-long-they-last). A temporary copy is held by whoever made it until +the run that made it is over, and a run resumed on the same fake takes it again as it was left. +`refs=` is the git refs `derive_worktree` knows, and `repo=False` a workdir that is not a +repository. ## Script the person diff --git a/specs/runtime/flowing.md b/specs/runtime/flowing.md index f5186055..d325ef2b 100644 --- a/specs/runtime/flowing.md +++ b/specs/runtime/flowing.md @@ -373,6 +373,16 @@ def under() -> Path: ... made removed, when the call and every call it started are over; a resumable run that keeps a journal MUST keep its directories and write them down instead. Removing MUST be shielded from cancellation and limited in time. `run_flow` MUST close whatever it opened. +- A session MUST also be closed as soon as nothing can reach it -- where only a reference cycle + holds it, as soon as a collection finds it -- however long its call goes on: the engine MUST + hold a `SessionView` only weakly, so that a flow opening a fresh session a round holds a + bounded number open however many rounds it runs, on fakes that never let the loop go on too. + A session let go of MUST be closed on the run's loop, whichever thread let go of it, and + exactly once however that races its call's end; the call's cleanup MUST wait for a close + still under way. A hook arriving for a session let go of and not yet closed -- its + `SESSION_END` among them -- MUST be handed a stand-in that is over, never the view that went. + A fork MUST keep the session it was forked from open until its own first turn, which is + where a harness cuts it. - A call whose caller has ended MUST raise `FlowCancelled` at its next operation. - `running` MUST answer with every call going now, and nothing of a call once it has ended. diff --git a/src/hmz/flows/agents.py b/src/hmz/flows/agents.py index 757d720e..2fb8bf2e 100644 --- a/src/hmz/flows/agents.py +++ b/src/hmz/flows/agents.py @@ -275,7 +275,12 @@ class Usage(pydantic.BaseModel): class Session(Protocol): - """One conversation of one agent, held in one environment.""" + """One conversation of one agent, held in one environment. + + A flow does not close one: it is closed when the flow call that opened it ends, or as + soon as nothing holds it any more, whichever comes first. Keep it in a variable for as + long as there are turns to take in it, and let go of it when there are not. + """ @property def agent(self) -> Agent: diff --git a/src/hmz/runtime/flowing/engine.py b/src/hmz/runtime/flowing/engine.py index d006a19b..1ded67fc 100644 --- a/src/hmz/runtime/flowing/engine.py +++ b/src/hmz/runtime/flowing/engine.py @@ -85,7 +85,7 @@ ) if TYPE_CHECKING: - from collections.abc import Callable, Coroutine + from collections.abc import Callable, Coroutine, Reversible from hmz.flows import ( AgentCollection, @@ -104,7 +104,7 @@ SessionHandle, Skill, ) - from .viewing import Releasable + from .viewing import Opened, Releasable __all__ = [ "DEPTH", @@ -739,7 +739,7 @@ def __init__( self.inflight = 0 self.expired = False self.fired = False - self.handle: asyncio.TimerHandle | None = asyncio.get_running_loop().call_later( + self.handle: asyncio.Handle | None = asyncio.get_running_loop().call_later( max(delay, 0.0), self.expire ) @@ -755,10 +755,19 @@ def expire(self) -> None: def turned(self, delta: int) -> None: """A turn under the call started or ended.""" self.inflight += delta - if self.expired and self.inflight == 0 and not self.fired: - self._fire() + if ( + self.expired + and self.inflight == 0 + and not self.fired + and self.handle is None + ): + # From the loop rather than from here: the turn that ended may be the call's + # own, in its task, and a flow returning without awaiting again would carry a + # cancel sent now out to whoever called it. + self.handle = asyncio.get_running_loop().call_soon(self._fire) def _fire(self) -> None: + self.handle = None if not self.node.ended: self.fired = True self.task.cancel(f"{self.node.ref}: its budget's duration is spent") @@ -846,7 +855,7 @@ def __init__( self.resumed = False self.state: FlowStateImpl | None = None self.clock: _Clock | None = None - self.res: list[Releasable] | None = None + self.res: dict[int, Releasable] | None = None self.seqs: dict[str, int] | None = None self._record: LiveCall | None = None @@ -923,9 +932,14 @@ def check(self) -> None: ) node = node.parent - def limits(self, budget: Budget | None) -> Limits: + def limits(self, budget: Budget | None) -> tuple[Limits, float | None]: """What a turn of this call may spend now, every budget over it taken together. + Returns: + The limits its driver is handed, and the deadline the engine holds the turn to + itself -- a hard one due after the graceful deadline the limits carry, which lets + the turn run on past it -- or None where the driver's limits do. + Raises: FlowCancelled: If the call's caller, or the run, is over. BudgetExceeded: The leaf for a budget over it that is spent. @@ -933,7 +947,7 @@ def limits(self, budget: Budget | None) -> Limits: # Each limit is the least any budget over the turn leaves, and the turn is hard -- # stopped mid-turn when a limit is reached -- where a budget that sets one of them # is: a hard budget stays hard under a graceful one, and the other way round. - cost = tokens = until = _INF + cost = tokens = until = cut = _INF cost_hard = tokens_hard = until_hard = False node: Call | None = self while node is not None: @@ -954,6 +968,8 @@ def limits(self, budget: Budget | None) -> Limits: ends = node.since + own.duration.total_seconds() if ends < until or (hard and ends == until): until, until_hard = ends, hard + if hard and ends < cut: + cut = ends node = node.parent now = time.monotonic() deadline = self.deadline @@ -984,16 +1000,23 @@ def limits(self, budget: Budget | None) -> Limits: ends, hard or (until_hard and ends == deadline), ) + if hard and ends < cut: + cut = ends + # `Limits` carry one deadline, and one `graceful` for every limit. A graceful soonest + # deadline lets the turn run on past it, which a driver told a hard limit would not: + # one held to a hard cost or token limit is told the hard deadline instead, if any, + # and one held to none is held to the hard deadline by the engine. + bites = (cost_hard and cost != _INF) or (tokens_hard and tokens != _INF) + if deadline != _INF and not until_hard: + if bites: + return limits_of(cost, tokens, cut, graceful=False), None + return ( + limits_of(cost, tokens, deadline, graceful=True), + None if cut == _INF else cut, + ) return limits_of( - cost, - tokens, - deadline, - graceful=not ( - (cost_hard and cost != _INF) - or (tokens_hard and tokens != _INF) - or (until_hard and deadline != _INF) - ), - ) + cost, tokens, deadline, graceful=not (bites or deadline != _INF) + ), None def turning(self, delta: int) -> None: """A turn of this call started (+1) or ended (-1), which a deadline waits on.""" @@ -1005,13 +1028,17 @@ def turning(self, delta: int) -> None: node = node.parent def hold(self, resource: Releasable) -> None: - """Keeps something the call made, to release when the call and its callees end.""" + """Keeps something the call made, to release when the call and its callees end. + + Kept by `id`, in the order it was made, so that a session closed before then -- its + view let go of -- is let go of here too, however many the call opens. + """ res = self.res if res is None: - self.res = [resource] + self.res = {id(resource): resource} self.run.holding.add(self) else: - res.append(resource) + res[id(resource)] = resource def made(self, view: EnvView, kind: str, name: str, derived: EnvView) -> None: """A temporary copy or scratch directory was made through one of this call's views. @@ -1034,7 +1061,7 @@ def made(self, view: EnvView, kind: str, name: str, derived: EnvView) -> None: } ) return - for one in self.res or (): + for one in (self.res or {}).values(): if ( type(one) is Made and one.driver is driver @@ -1047,16 +1074,16 @@ def made(self, view: EnvView, kind: str, name: str, derived: EnvView) -> None: def unmade(self, driver: EnvDriver, kind: str, name: str) -> None: """A temporary copy or scratch directory was removed by the flow itself.""" if self.res: - self.res = [ - one - for one in self.res + self.res = { + key: one + for key, one in self.res.items() if not ( type(one) is Made and one.driver is driver and one.kind == kind and one.id == name ) - ] + } def arm(self, own: Budget) -> None: """Starts the call's own deadline, where it is sooner than the one above it.""" @@ -1115,7 +1142,7 @@ async def _cascade(node: Call) -> None: if res is not None: node.res = None node.run.holding.discard(node) - await _released(res) + await _released(res.values()) parent = node.parent if parent is None: return @@ -1125,7 +1152,8 @@ async def _cascade(node: Call) -> None: node = parent -async def _released(res: list[Releasable]) -> None: +async def _released(res: Reversible[Releasable]) -> None: + """Releases what a call made, the last made first.""" for one in reversed(res): try: await one.release() @@ -1187,7 +1215,10 @@ class Run: __slots__ = ( "bringing", + "closing", "derived", + "dropped", + "due", "fetched", "here_chain", "here_driver", @@ -1199,6 +1230,7 @@ class Run: "loads", "local", "lock", + "loop", "opened", "past", "person", @@ -1207,6 +1239,7 @@ class Run: "recorder", "skills", "specs", + "thread", ) def __init__( @@ -1216,6 +1249,13 @@ def __init__( local: EnvDriver | None, recorder: Recorder | None, ) -> None: + """A run about to start, on the loop running now. + + Raises: + RuntimeError: If no loop is. + """ + self.loop = asyncio.get_running_loop() + self.thread = threading.get_ident() self.lock = threading.Lock() self.live: dict[Call, None] = {} self.person = person @@ -1237,6 +1277,9 @@ def __init__( self.holding: set[Call] = set() self.fetched: dict[tuple[str, str | None], asyncio.Future[Path]] = {} self.reapers: set[asyncio.Future[None]] = set() + self.dropped: list[Opened] = [] + self.due = False + self.closing: set[asyncio.Task[None]] = set() self.derived: dict[int, EnvDriver] = {} self.specs: dict[tuple[AgentDriver, Grant], str] = {} @@ -1274,6 +1317,24 @@ def spec_of(self, view: AgentView | OutworlderView) -> str: ) return said + def drain(self) -> None: + """Starts closing every session whose view went since this was last done.""" + dropped = self.dropped + if dropped: + self.dropped = [] + for opened in dropped: + opened.drop() + + def drain_soon(self) -> None: + """Drains as soon as the loop gets to it, asked once however many views go first.""" + if not self.due: + self.due = True + self.loop.call_soon(self._drained) + + def _drained(self) -> None: + self.due = False + self.drain() + def spawned( self, node: Call, role: str, handle: SessionHandle, driver: AgentDriver ) -> None: @@ -1297,19 +1358,27 @@ async def close(self) -> None: from .loading import unpin try: + self.drain() if self.reapers: await asyncio.wait(set(self.reapers), timeout=REAP) # What calls still going made -- ones a flow started and never waited for -- # goes with the run, which is over whether they are or not. left: list[Releasable] = [] for node in list(self.holding): - left.extend(node.res or ()) + left.extend((node.res or {}).values()) node.res = None self.holding.clear() + # With them, the sessions let go of whose close is still under way, which the + # environments they are in outlive. + waiting = set(self.closing) if left: - try: - await asyncio.wait_for(asyncio.shield(_released(left)), REAP) - except TimeoutError: + releasing = asyncio.ensure_future(_released(left)) + self.reapers.add(releasing) + releasing.add_done_callback(self.reapers.discard) + waiting.add(releasing) + if waiting: + _, late = await asyncio.wait(waiting, timeout=REAP) + if late: log.warning("cleaning up after a run took longer than %ss", REAP) closing = list(reversed(self.derived.values())) if self.opened and self.local is not None: diff --git a/src/hmz/runtime/flowing/fakes.py b/src/hmz/runtime/flowing/fakes.py index 20b4293d..19df074d 100644 --- a/src/hmz/runtime/flowing/fakes.py +++ b/src/hmz/runtime/flowing/fakes.py @@ -20,11 +20,15 @@ async def test_fix_stops_when_the_tests_pass(): others a scripted answer reaches for through its session -- :meth:`FakeSession.tool`, :meth:`~FakeSession.ask`, :meth:`~FakeSession.notify`, :meth:`~FakeSession.subagent` -- so that a flow's hooks are exercised too. A `STOP` hook that blocks keeps the turn going - with its reason as the next prompt, as a real one does. + with its reason as the next prompt, as a real one does, and a hard deadline cuts off a turn + its answer holds open -- one waiting on :meth:`FakeSession.until_steered` -- as a real + driver's does. - :class:`FakeEnvDriver` is a dictionary of files under a workdir, with worktrees, temporary copies and scratch directories as copies of it, and `exec` answered by a function or a table, with a few commands -- `true`, `false`, `echo`, `cat`, `ls`, `sleep` - and `git rev-parse` -- answered by default. + and `git rev-parse` -- answered by default. A temporary copy is its holder's until the + run that took it closes, and a run resumed on the same fake takes it again as it was + left. - :class:`FakeOutworlder` answers as the person outside the run would, or is away. - :func:`run_fake` runs a flow with a fake for every role nobody gave one for. @@ -268,6 +272,7 @@ def __init__( self._turning = False self._started = False self._interrupted = threading.Event() + self._expired = False self._steers: deque[str] = deque() self._loop: asyncio.AbstractEventLoop | None = None self._wake: asyncio.Event | None = None @@ -295,11 +300,21 @@ async def turn(self, request: TurnRequest, sink: UsageSink) -> Any: if self._turning: raise SessionError(f"{self._id} is taking a turn already") driver = self.driver - self._loop = asyncio.get_running_loop() + loop = self._loop = asyncio.get_running_loop() self._wake = asyncio.Event() self._interrupted.clear() + self._expired = False self._steers.clear() prompt = request.prompt + limits = request.limits + # A hard deadline stops the turn where it has got to, as a real driver's does. + timer = ( + None + if limits.graceful or limits.deadline is None + else loop.call_later( + max(limits.deadline - time.monotonic(), 0.0), self._expire + ) + ) self._turning = True try: if not self._started: @@ -320,10 +335,21 @@ async def turn(self, request: TurnRequest, sink: UsageSink) -> Any: again = 0 while True: self.prompts.append(prompt) - said = await driver.script.next( - prompt, request.output_schema, driver.harness, session=self - ) + try: + said = await driver.script.next( + prompt, request.output_schema, driver.harness, session=self + ) + except SessionError as error: + if self._expired: + raise DurationExceeded( + f"{self._id} reached the deadline of its turn" + ) from error + raise if self._interrupted.is_set(): + if self._expired: + raise DurationExceeded( + f"{self._id} reached the deadline of its turn" + ) raise SessionError(f"{self._id} was interrupted") self._spend(request, sink) text = said if isinstance(said, str) else _as(said, None, "") @@ -336,6 +362,13 @@ async def turn(self, request: TurnRequest, sink: UsageSink) -> Any: prompt = stopped.reason finally: self._turning = False + if timer is not None: + timer.cancel() + + def _expire(self) -> None: + """The turn's hard deadline has come: it stops.""" + self._expired = True + self.interrupt() def _spend(self, request: TurnRequest, sink: UsageSink) -> None: """Spends one answer's worth, or what the limits leave where they are hard.""" @@ -383,6 +416,7 @@ async def close(self) -> None: if self.closed: return self.closed = True + self.driver.live -= 1 self.interrupt() await self._fire(HookKind.SESSION_END) @@ -483,6 +517,8 @@ class FakeAgentDriver: Attributes: sessions: Every session it opened, in order. + live: How many of them are open now. + peak: The most of them that were open at once. closed: How many times it was closed. """ @@ -515,6 +551,8 @@ def __init__( self.seconds = seconds self.forks = forks self.sessions: list[FakeSession] = [] + self.live = 0 + self.peak = 0 self.closed = 0 def __repr__(self) -> str: @@ -545,6 +583,8 @@ async def open( forked = fork_of session = FakeSession(self, placement, permission, skills, hooks, forked) self.sessions.append(session) + self.live += 1 + self.peak = max(self.peak, self.live) return session async def close(self) -> None: @@ -557,12 +597,20 @@ async def close(self) -> None: class _Disk: - """The files of one fake machine, shared by every environment on it.""" + """The files of one fake machine, shared by every environment on it. + + Its temporary copies are the machine's too, by the workdir copied and the id, as a real + machine's are by where they are made -- so that an environment derived again, as a + resumed run derives it, finds the copies made from it. + """ def __init__(self) -> None: self.files: dict[PurePosixPath, bytes] = {} self.commands: list[tuple[PurePosixPath, Command]] = [] self.numbers = itertools.count(1) + self.clones: dict[tuple[PurePosixPath, str], FakeEnvDriver] = {} + #: Who holds each copy, until the run that took it closes. + self.holders: dict[tuple[PurePosixPath, str], object] = {} class FakeEnvDriver: @@ -584,6 +632,10 @@ class FakeEnvDriver: refs: The git refs `derive_worktree` knows. repo: Whether the workdir is a git repository, as `git rev-parse` answers. + Closing one lets go of the holds on the temporary copies taken through it, leaving the + copies where they are, as a real driver does: a copy's own as the run that derived it + closes it, and every copy's on the machine as the fake they were made from is closed. + Attributes: closed: How many times it was closed. """ @@ -604,6 +656,7 @@ def __init__( refs: Iterable[str] = ("HEAD", "main"), repo: bool = True, _disk: _Disk | None = None, + _clone: tuple[PurePosixPath, str] | None = None, ) -> None: self.workdir = PurePosixPath(workdir) self.backend = EnvBackendKind(backend) @@ -620,8 +673,9 @@ def __init__( self.repo = repo self.available = True self.closed = 0 + self._root = _disk is None self._disk = _Disk() if _disk is None else _disk - self._held: dict[str, tuple[object, FakeEnvDriver]] = {} + self._clone = _clone self._scratch: dict[str, FakeEnvDriver] = {} for path, data in (files or {}).items(): self._disk.files[self._at(path)] = ( @@ -664,7 +718,13 @@ def placement(self) -> Placement: def _at(self, path: str | PurePosixPath) -> PurePosixPath: return self.workdir / path - def _there(self, workdir: PurePosixPath, *, copy: bool) -> FakeEnvDriver: + def _there( + self, + workdir: PurePosixPath, + *, + copy: bool, + clone: tuple[PurePosixPath, str] | None = None, + ) -> FakeEnvDriver: if copy: files = self._disk.files for path, data in list(files.items()): @@ -685,6 +745,7 @@ def _there(self, workdir: PurePosixPath, *, copy: bool) -> FakeEnvDriver: refs=self.refs, repo=self.repo, _disk=self._disk, + _clone=clone, ) def _gone(self, workdir: PurePosixPath) -> None: @@ -800,21 +861,33 @@ async def derive_temp_clone( *, holder: object, ) -> FakeEnvDriver: - held = self._held.get(id) - if held is not None: - if held[0] != holder: + disk = self._disk + key = (self.workdir, id) + copy = disk.clones.get(key) + if copy is None: + copy = self._there( + PurePosixPath(f"/clones/{next(disk.numbers)}-{id}"), + copy=True, + clone=key, + ) + elif key in disk.holders: + if disk.holders[key] != holder: raise TempCloneBusy(f"the copy {id!r} is somebody else's") - return held[1] - copy = self._there( - PurePosixPath(f"/clones/{next(self._disk.numbers)}-{id}"), copy=True - ) - self._held[id] = (holder, copy) + return copy + else: + # Let go of as the run that took it closed, and taken again as it was left -- + # which is how a resumed run finds its copies. + copy = self._there(copy.workdir, copy=False, clone=key) + disk.clones[key] = copy + disk.holders[key] = holder return copy async def destroy_temp_clone(self, id: str) -> None: # noqa: A002 -- the driver's - held = self._held.pop(id, None) - if held is not None: - self._gone(held[1].workdir) + key = (self.workdir, id) + copy = self._disk.clones.pop(key, None) + self._disk.holders.pop(key, None) + if copy is not None: + self._gone(copy.workdir) async def derive_scratch(self, id: str) -> FakeEnvDriver: # noqa: A002 -- the driver's scratch = self._scratch.get(id) @@ -832,7 +905,7 @@ async def destroy_scratch(self, id: str) -> None: # noqa: A002 -- the driver's @property def clones(self) -> list[str]: """The ids of the temporary copies made here and not yet removed.""" - return sorted(self._held) + return sorted(name for at, name in self._disk.clones if at == self.workdir) @property def scratches(self) -> list[str]: @@ -841,6 +914,15 @@ def scratches(self) -> list[str]: async def close(self) -> None: self.closed += 1 + # Holds on copies are let go of, and the copies left where they are, as a real + # driver's are -- so that a run resumed on this machine takes them again. A copy is + # closed by the run that derived it as it ends; the driver it was derived from, by + # whoever made it. + disk = self._disk + if self._root: + disk.holders.clear() + elif self._clone is not None and disk.clones.get(self._clone) is self: + disk.holders.pop(self._clone, None) # ---------------------------------------------------------------------------- outworlders diff --git a/src/hmz/runtime/flowing/viewing.py b/src/hmz/runtime/flowing/viewing.py index 2d164b6f..e2580179 100644 --- a/src/hmz/runtime/flowing/viewing.py +++ b/src/hmz/runtime/flowing/viewing.py @@ -15,6 +15,17 @@ Coder(Agent, GoalCommandAgentMixin)` is a type for a type checker, and a view deriving from it would carry the protocols' stub bodies, lose its `__slots__`, and make every `isinstance` slower for nothing. A view is a handful of pointers, made per flow call. + +A session has no `close` in the flow API, so the engine closes it: when the flow call that +opened it ends -- whoever holds it by then, the caller it was handed up to included -- or as +soon as nothing can reach its view any more, whichever comes first. The engine holds a +:class:`SessionView` only weakly, through the :class:`Opened` it keeps of each session, so a +flow opening a fresh session a round holds a few open however many rounds it runs. Letting go +of a view -- wherever its last reference goes, or on whichever thread a collection finds it +-- has the session closed on the run's loop, once, in a task the call's cleanup waits for. A +hook arriving for a session whose view is gone, as it closes, is handed a stand-in: a view of +it that is over, which nothing can take a turn in. A fork holds the session it was forked from +until its first turn, which is where a harness cuts it. """ from __future__ import annotations @@ -26,6 +37,9 @@ import asyncio import contextvars import logging +import threading +import time +import weakref from typing import TYPE_CHECKING, Any, ClassVar, Self, overload import pydantic @@ -35,6 +49,7 @@ AskUserHookParams, BashEnvMixin, CapabilityNotGranted, + DurationExceeded, EnvBackendKind, FilesEnvMixin, GitWorktreeEnvMixin, @@ -73,6 +88,7 @@ if TYPE_CHECKING: from collections.abc import Awaitable, Callable from pathlib import PurePosixPath + from types import FrameType from hmz.flows import ( AskUserHookResult, @@ -121,13 +137,17 @@ class _Line: - """What an agent and every agent derived from it share: its hooks and its sessions.""" + """What an agent and every agent derived from it share: its hooks and its sessions. + + Its sessions are the ones not closed yet, by the `id` of their handles, which is what a + hook arrives with. + """ __slots__ = ("hooks", "sessions") def __init__(self) -> None: self.hooks = HookTable() - self.sessions: dict[int, SessionView] = {} + self.sessions: dict[int, Opened] = {} class _Sink: @@ -149,6 +169,39 @@ def add(self, *, cost: float, output_tokens: int, duration: float) -> None: node = node.parent +class _Cut: + """A hard deadline a turn is held to by the engine, its driver being told a sooner one. + + `Limits` carry one deadline, and a turn is told the soonest of those over it. Where that + is a graceful one -- which lets the turn run on past it -- and a hard one is due later, + the turn is still to stop at the hard one, and it is this that stops it there: it + interrupts the session, and the turn raises `DurationExceeded`. + """ + + __slots__ = ("fired", "timer") + + def __init__( + self, loop: asyncio.AbstractEventLoop, handle: SessionHandle, deadline: float + ) -> None: + self.fired = False + self.timer = loop.call_later( + max(deadline - time.monotonic(), 0.0), self._fire, handle + ) + + def _fire(self, handle: SessionHandle) -> None: + self.fired = True + handle.interrupt() + + def stop(self) -> None: + """The turn is over: the timer goes.""" + self.timer.cancel() + + @staticmethod + def exceeded(role: str) -> DurationExceeded: + """What the turn it stopped raises.""" + return DurationExceeded(f"{role}: the turn ran to a hard deadline over it") + + def _command(prompt: str, word: str) -> bool: """Whether a prompt is one of the harness's own commands, as `/goal` or `/loop`.""" return prompt.startswith(word) and ( @@ -276,6 +329,8 @@ async def _opened(self, env: Env, fork_of: SessionView | None) -> SessionView: node.check() line = self._lined() run = node.run + if run.dropped: + run.drain() handle = await self._driver.open( env._driver.placement(), permission=self._grant.permission, @@ -284,8 +339,12 @@ async def _opened(self, env: Env, fork_of: SessionView | None) -> SessionView: fork_of=None if fork_of is None else fork_of._handle, ) session = SessionView(self, env, handle) - line.sessions[id(handle)] = session - node.hold(session) + # A harness cuts a fork as the fork's first turn goes, from the session it was + # forked from, which is kept open until then however soon the flow lets go of it. + session._parent = fork_of + opened = Opened(session, self, env, handle) + line.sessions[id(handle)] = opened + node.hold(opened) run.spawned(node, self._role, handle, self._driver) return session @@ -329,7 +388,7 @@ async def run( f"{self._role}: /loop needs LoopCommandAgentMixin on the role" ) node = self._node - limits = node.limits(budget) + limits, hard = node.limits(budget) if taken._busy: raise SessionError(f"{self._role}: a turn of this session is under way") if taken._closed: @@ -337,6 +396,7 @@ async def run( handle = taken._handle taken._busy = True node.turning(1) + cut = None if hard is None else _Cut(node.run.loop, handle, hard) try: said = await handle.turn( TurnRequest(prompt, output_schema, limits), _Sink(node) @@ -344,15 +404,23 @@ async def run( except asyncio.CancelledError: handle.interrupt() raise - except BaseException: + except BaseException as error: failed = taken._error if failed is not None: taken._error = None raise failed from None + # What an interrupted turn raises; a turn that answered, or failed of itself, + # before the deadline's interrupt reached it did so in time. + if cut is not None and cut.fired and isinstance(error, SessionError): + raise cut.exceeded(self._role) from error raise finally: taken._busy = False node.turning(-1) + if cut is not None: + cut.stop() + # A fork is cut by now, and the session it was cut from may go. + taken._parent = None failed = taken._error if failed is not None: taken._error = None @@ -376,6 +444,8 @@ async def steer( ) taken = self._own(session) self._node.check() + if taken._closed: + raise SessionError(f"{self._role}: the session is over") await taken._handle.steer(prompt, queued=queued) def _brought(self) -> tuple[Skill, ...]: @@ -427,15 +497,30 @@ def _hang( sessions = line.sessions async def bound(handle: SessionHandle, fields: dict[str, Any]) -> HookResult: - session = sessions.get(id(handle)) + opened = sessions.get(id(handle)) + session = None + if opened is not None: + session = opened() + if session is None: + session = opened.standin() token = CALLING.set(node) + hooked: FrameType | None = None try: - return await fn(params(ctx=node, session=session, **fields)) # pyright: ignore[reportArgumentType] + awaited = fn(params(ctx=node, session=session, **fields)) # pyright: ignore[reportArgumentType] + hooked = getattr(awaited, "cr_frame", None) + return await awaited except Exception as error: # The flow's to raise, not the driver's: the turn it arrived in stops, and # raises it where the flow is waiting on that turn. if session is not None and not session._closed: + # Kept on the session till then, where the frames it came out of -- + # this one's `session`, the hook's own `params`, over now -- would hold + # the session in a cycle through it: one let go of meanwhile would stay + # open until a collection found it. + if hooked is not None: + hooked.clear() session._failed(error) + del session else: log.exception("a %s hook raised", kind) return default_result(kind) @@ -510,9 +595,23 @@ def on_ask_user( class SessionView: - """One session, as the flow that opened it holds it.""" + """One session, as the flow that opened it holds it. - __slots__ = ("_agent", "_busy", "_closed", "_env", "_error", "_handle", "_line") + The engine keeps it only weakly -- see :class:`Opened` -- so the session closes as soon + as the flow lets go of it, if the call that opened it has not ended first. + """ + + __slots__ = ( + "__weakref__", + "_agent", + "_busy", + "_closed", + "_env", + "_error", + "_handle", + "_line", + "_parent", + ) def __init__( self, @@ -528,6 +627,7 @@ def __init__( self._busy = False self._closed = False self._error: Exception | None = None + self._parent: SessionView | None = None def __repr__(self) -> str: return f"" @@ -555,15 +655,138 @@ def _failed(self, error: Exception) -> None: self._error = error self._handle.interrupt() + +class Opened(weakref.ref["SessionView"]): + """A session an agent view opened, as the engine keeps it: what closes it, and when. + + A weak reference to the session's view, since the view is the flow's to hold -- the + engine holding it would keep every session a flow ever opened open until its call ends + -- and what the call's cleanup and the agent's hooks hold instead, which is why it holds + nothing that holds the view. The view going has the session closed on the run's loop, + in a task the call's cleanup waits for (:meth:`drop`); the call ending first closes it + there and then (:meth:`release`). Whichever comes first closes it, once. + + Attributes: + agent: The agent view that opened it, whose call it belongs to. + env: The environment view it was opened in. + handle: The driver's session. + closed: Whether its closing has started. + task: The close the view going started, while it is under way. + """ + + __slots__ = ("agent", "closed", "env", "handle", "task") + + def __new__( + cls, + view: SessionView, + agent: AgentView, + env: EnvView, + handle: SessionHandle, + ) -> Self: + """A record of a session, watching its view go.""" + del agent, env, handle + return super().__new__(cls, view, _dropped) + + def __init__( + self, + view: SessionView, + agent: AgentView, + env: EnvView, + handle: SessionHandle, + ) -> None: + """A record of a session, watching its view go.""" + # Not `weakref.ref.__init__`, which only checks the arguments `__new__` took. + del view + self.agent = agent + self.env = env + self.handle = handle + self.closed = False + self.task: asyncio.Task[None] | None = None + + def __repr__(self) -> str: + return f"" + + def standin(self) -> SessionView: + """The session as a hook arriving after its view went is handed it: one that is over.""" + view = SessionView(self.agent, self.env, self.handle) + view._closed = True + return view + + def drop(self) -> None: + """Closes a session whose view went, on the run's loop, in a task of its own. + + The task is started eagerly: a close that need not wait -- a fake's, or that of a + session whose CLI never started -- is over before this returns, and costs no turn of + the loop. + """ + if self.closed: + return + self.closed = True + run = self.agent._node.run + task = asyncio.Task(self._shut(), loop=run.loop, eager_start=True) + if task.done(): + self._settled(task) + return + self.task = task + run.closing.add(task) + task.add_done_callback(self._settled) + async def release(self) -> None: - """Closes the session, which the flow call that opened it does as it ends.""" - if not self._closed: - self._closed = True - try: - await self._handle.close() - finally: - if self._line is not None: - self._line.sessions.pop(id(self._handle), None) + """Closes the session as the call that opened it ends, or waits out its close.""" + if not self.closed: + self.closed = True + await self._shut() + elif self.task is not None: + await asyncio.shield(self.task) + + async def _shut(self) -> None: + view = self() + if view is not None: + view._closed = True + handle = self.handle + try: + await handle.close() + except Exception: + log.exception("closing %r failed", self) + finally: + line = self.agent._line + if line is not None: + line.sessions.pop(id(handle), None) + + def _settled(self, task: asyncio.Task[None]) -> None: + """Its close is over, and the call has nothing of it left to release.""" + self.task = None + node = self.agent._node + node.run.closing.discard(task) + res = node.res + if res is not None: + res.pop(id(self), None) + if not res: + node.res = None + node.run.holding.discard(node) + + +def _dropped(opened: Opened) -> None: + """A session's view went, so nothing can reach the session: it is to be closed. + + Called by the interpreter as the view goes -- in the middle of whatever statement let go + of it, or on whichever thread a collection found it -- so all it does is put the session + where the run's loop closes it: among the run's `dropped`, which the loop drains as soon + as it gets to it, and the next session opened in the run drains first -- so that a flow + that never lets the loop go on, as one on fakes may not, holds few open all the same. + """ + if opened.closed: + return + run = opened.agent._node.run + try: + if threading.get_ident() != run.thread: + run.loop.call_soon_threadsafe(opened.drop) + return + run.dropped.append(opened) + run.drain_soon() + except RuntimeError: + # The loop is closed: the run ended with it, closing what it held as it did. + return # ------------------------------------------------------------------------- environments @@ -985,7 +1208,7 @@ async def run( ) node = self._node if node is not None: - node.limits(budget) + node.limits(budget) # refuses a turn under a spent budget if source.made: hook = source.hook if hook is None: @@ -1063,7 +1286,7 @@ def on_session_end( #: Every resource a flow call made that goes when it does, by what makes one: a session #: closes, a temporary copy or scratch directory is removed. -type Releasable = SessionView | Made +type Releasable = Opened | Made class Made: diff --git a/tests/unit/flows/test_engine_budget.py b/tests/unit/flows/test_engine_budget.py index 6195fa3a..83e5ea80 100644 --- a/tests/unit/flows/test_engine_budget.py +++ b/tests/unit/flows/test_engine_budget.py @@ -542,3 +542,179 @@ async def calling( driver = FakeAgentDriver() await run_fake(calling, agents={"agent": driver}, budget=run) assert driver.sessions[0].requests[0].limits.graceful is graceful + + +# ------------------------------------------------------- a hard deadline under a graceful one + + +async def _held(prompt: str, *, session: Any, output_schema: Any) -> str: + """A turn that goes on until it is steered or stopped, which nobody steers it.""" + del prompt, output_schema + return await session.until_steered() + + +def _seconds(value: float) -> datetime.timedelta: + return datetime.timedelta(seconds=value) + + +class Deadlines(FlowParams): + #: The turn's own budget's duration, and whether it is hard; none for 0. + turn: float = 0.0 + turn_hard: bool = True + #: The budget of the flow the turn is taken in: its duration, and whether it is hard. + own: float = 60.0 + own_hard: bool = False + + +@flow(agents=Solo, envs=Place, params=Deadlines) +async def held_turn( + task: str, *, agents: Solo, envs: Place, params: Deadlines, ctx: FlowContext +) -> float: + """Takes one turn that nothing ends but a deadline, and says when it was ended.""" + agent = agents["agent"] + session = await agent.spawn(env=envs["env"]) + started = time.monotonic() + turn = ( + Budget(duration=_seconds(params.turn), graceful=not params.turn_hard) + if params.turn + else None + ) + try: + await agent.run("hold on", session=session, budget=turn) + except DurationExceeded: + return time.monotonic() - started + return -1.0 + + +@flow(agents=Solo, envs=Place, params=Deadlines) +async def under( + task: str, *, agents: Solo, envs: Place, params: Deadlines, ctx: FlowContext +) -> float: + """Calls `held_turn` under a budget of its own.""" + own = Budget(duration=_seconds(params.own), graceful=not params.own_hard) + return await held_turn(task, agents=agents, envs=envs, params=params, budget=own) + + +@pytest.mark.parametrize( + ("params", "run", "cut"), + [ + (Deadlines(turn=0.3), Budget(duration=_seconds(0.05)), 0.3), + (Deadlines(turn=0.3, own=0.05), None, 0.3), + (Deadlines(turn=0.1, own=1.0), None, 0.1), + (Deadlines(own=0.3, own_hard=True), Budget(duration=_seconds(0.05)), 0.3), + ], + ids=[ + "the turn's own hard, the run's graceful sooner", + "the turn's own hard, its flow's graceful sooner", + "the turn's own hard, its flow's graceful later", + "its flow's hard, the run's graceful sooner", + ], +) +async def test_a_hard_deadline_stops_a_turn_whatever_graceful_one_comes_first( + params: Deadlines, run: Budget | None, cut: float +) -> None: + driver = FakeAgentDriver(reply=_held) + started = time.monotonic() + try: + ended = await asyncio.wait_for( + run_fake( + under, + agents={"agent": driver}, + params=params, + budget=run or Budget(cost=math.inf), + ), + 5, + ) + except DurationExceeded: + # The graceful deadline over the turn has passed too, so the flow is stopped as + # the turn ends -- which is when the hard deadline came. + ended = time.monotonic() - started + assert cut - 0.02 <= ended < cut + 1, f"ended after {ended:.3f}s, not {cut}s" + assert driver.sessions[0].closed + + +async def test_a_turn_hard_for_its_cost_is_told_the_hard_deadline_not_a_graceful_one() -> ( + None +): + @flow(agents=Solo, envs=Place, params=Depth) + async def one_turn( + task: str, *, agents: Solo, envs: Place, params: Depth, ctx: FlowContext + ) -> None: + agent = agents["agent"] + session = await agent.spawn(env=envs["env"]) + await agent.run("x", session=session, budget=Budget(cost=1, graceful=False)) + + @flow(agents=Solo, envs=Place, params=Depth) + async def calling( + task: str, *, agents: Solo, envs: Place, params: Depth, ctx: FlowContext + ) -> None: + own = Budget(duration=_seconds(30)) + await one_turn(task, agents=agents, envs=envs, params=params, budget=own) + + driver = FakeAgentDriver() + started = time.monotonic() + await run_fake( + calling, + agents={"agent": driver}, + budget=Budget(duration=_seconds(60), graceful=False), + ) + limits = driver.sessions[0].requests[0].limits + assert not limits.graceful + assert limits.deadline is not None + assert 45 < limits.deadline - started < 61, "told the graceful one as hard" + + +async def test_a_flow_returning_as_its_graceful_deadline_lets_go_returns() -> None: + @flow(agents=Solo, envs=Place, params=Depth) + async def finishing( + task: str, *, agents: Solo, envs: Place, params: Depth, ctx: FlowContext + ) -> str: + agent = agents["agent"] + session = await agent.spawn(env=envs["env"]) + return await agent.run("slow", session=session) + + @flow(agents=Solo, envs=Place, params=Depth) + async def calling( + task: str, *, agents: Solo, envs: Place, params: Depth, ctx: FlowContext + ) -> list[str]: + own = Budget(duration=_seconds(0.05)) + said = await finishing( + task, agents=agents, envs=envs, params=params, budget=own + ) + await asyncio.sleep(0.01) + return [said, "and on"] + + async def slow(prompt: str, **_: Any) -> str: + await asyncio.sleep(0.15) + return prompt + + driver = FakeAgentDriver(reply=slow) + assert await run_fake(calling, agents={"agent": driver}) == ["slow", "and on"] + + +async def test_the_timer_a_turn_is_held_to_goes_with_the_turn( + monkeypatch: pytest.MonkeyPatch, +) -> None: + loop = asyncio.get_running_loop() + armed: list[asyncio.TimerHandle] = [] + real = loop.call_later + + def arming(delay: float, callback: Any, *args: Any, **kwargs: Any) -> Any: + handle = real(delay, callback, *args, **kwargs) + armed.append(handle) + return handle + + @flow(agents=Solo, envs=Place, params=Depth) + async def many( + task: str, *, agents: Solo, envs: Place, params: Depth, ctx: FlowContext + ) -> None: + agent = agents["agent"] + session = await agent.spawn(env=envs["env"]) + hard = Budget(duration=_seconds(120), graceful=False) + for _ in range(params.turns): + await agent.run("x", session=session, budget=hard) + + monkeypatch.setattr(loop, "call_later", arming) + await run_fake(many, params={"turns": 10_000}, budget=Budget(duration=_seconds(60))) + assert len(armed) >= 10_000, "no turn was held to its hard deadline by the engine" + assert all(one.cancelled() for one in armed), "a timer outlived its turn" diff --git a/tests/unit/flows/test_engine_resume.py b/tests/unit/flows/test_engine_resume.py index 7da996de..a228770f 100644 --- a/tests/unit/flows/test_engine_resume.py +++ b/tests/unit/flows/test_engine_resume.py @@ -22,6 +22,7 @@ Budget, Env, EnvCollection, + FilesEnvMixin, FlowContext, FlowParams, StateNotSerializable, @@ -247,6 +248,38 @@ async def cloning( assert (before["resumed"], after["resumed"]) == (False, True) +class Copies(Env, TemporaryClonedDirEnvMixin, FilesEnvMixin): ... + + +class CopyPlace(EnvCollection): + env: Copies + + +async def test_a_run_resumed_on_the_same_fake_takes_its_copy_again( + tmp_path: Path, +) -> None: + @flow(agents=Solo, envs=CopyPlace, params=Step, resumable=True) + async def copying( + task: str, *, agents: Solo, envs: CopyPlace, params: Step, ctx: FlowContext + ) -> tuple[str, bytes]: + clone = await envs["env"].derive_temp_clone("work") + if params.fail_at >= 0: + await clone.write("notes.txt", b"left off here") + raise CrashError(str(clone.workdir)) + return str(clone.workdir), await clone.read("notes.txt") + + env = FakeEnvDriver({"notes.txt": "fresh"}) + journal = tmp_path / "run.jsonl" + with pytest.raises(CrashError) as crashed: + await _run( + copying, journal, resume=False, envs={"env": env}, params={"fail_at": 0} + ) + assert env.clones == ["work"] + workdir, notes = await _run(copying, journal, resume=True, envs={"env": env}) + assert (workdir, notes) == (str(crashed.value), b"left off here") + assert env.clones == ["work"] + + # ---------------------------------------------------------------------------- the file diff --git a/tests/unit/flows/test_engine_scale.py b/tests/unit/flows/test_engine_scale.py index a264ab8c..483cf91f 100644 --- a/tests/unit/flows/test_engine_scale.py +++ b/tests/unit/flows/test_engine_scale.py @@ -61,6 +61,7 @@ class Shape(FlowParams): level: int = 0 gather: bool = False calls: int = 0 + keep: bool = False #: The params each level of a tree is called with, made once so that neither the engine nor @@ -172,6 +173,20 @@ async def looping_resumable( return (time.thread_time() - started) / max(params.calls, 1) +@flow(agents=Pair, envs=Envs, params=Shape) +async def spawning( + task: str, *, agents: Pair, envs: Envs, params: Shape, ctx: FlowContext +) -> None: + """Opens a session a round and takes a turn in it, keeping every one or letting go.""" + agent, env = agents["a"], envs["repo"] + kept: list[Any] = [] + for _ in range(params.calls): + session = await agent.spawn(env=env) + await agent.run(task, session=session) + if params.keep: + kept.append(session) + + async def _fake(flow_: Any, **said: Any) -> Any: return await run_fake( flow_, @@ -337,6 +352,31 @@ async def measured() -> float: assert best < 25e-6 * 3, f"{_us(best)} per journaled call" +async def test_a_session_let_go_of_a_round_costs_little_more_than_one_kept() -> None: + """What closing a session as its view goes adds to opening and using it. + + Kept, every session closes as the call ends, which is what closing cost before a + session could be let go of; let go of, each closes as the next is opened. Both are timed + whole -- the closes included -- and the second is held to a small multiple of the first. + """ + rounds = 5_000 + + async def per(*, keep: bool) -> float: + async def measured() -> float: + started = time.thread_time() + await _fake(spawning, params=Shape(calls=rounds, keep=keep)) + return (time.thread_time() - started) / rounds + + return await _best(measured) + + kept, let_go = await per(keep=True), await per(keep=False) + TIMINGS["session a round: spawn + turn + close, let go of / kept"] = ( + f"{_us(let_go)} / {_us(kept)} ({let_go / kept:.2f}x)" + ) + assert let_go / kept <= 1.6, f"{let_go / kept:.2f}x a session kept to the end" + assert let_go < 30e-6 * 3, f"{_us(let_go)} a round" + + async def test_calls_scale_linearly() -> None: async def per(calls: int) -> float: async def measured() -> float: diff --git a/tests/unit/flows/test_engine_sessions.py b/tests/unit/flows/test_engine_sessions.py new file mode 100644 index 00000000..e8d8d11e --- /dev/null +++ b/tests/unit/flows/test_engine_sessions.py @@ -0,0 +1,549 @@ +"""How long a session lives: until its flow call ends, or until the flow lets go of it. + +The flow API gives a flow no way to close a session, so the engine closes it: when the flow +call that opened it ends, whoever holds it by then, or as soon as nothing can reach it any +more -- whichever comes first. A flow opening a fresh session a round holds a few open however +many rounds it runs, and every session is closed once, on the run's loop, whichever thread let +go of it. +""" + +from __future__ import annotations + +import asyncio +import gc +import math +import threading +from typing import TYPE_CHECKING, Any, cast + +import pytest + +from hmz.flows import ( + Agent, + AgentCollection, + Budget, + Env, + EnvCollection, + FlowContext, + FlowParams, + NotificationHookParams, + NotificationHookResult, + SessionEndHookParams, + SessionEndHookResult, + SessionError, + UserPromptSubmitHookParams, + UserPromptSubmitHookResult, + flow, + load, +) +from hmz.runtime.flowing.engine import run_flow +from hmz.runtime.flowing.fakes import FakeAgentDriver, FakeEnvDriver, run_fake +from tests.flows.kit import flowverse + +if TYPE_CHECKING: + from pathlib import Path + + +class Solo(AgentCollection): + agent: Agent + + +class Place(EnvCollection): + env: Env + + +class Rounds(FlowParams): + rounds: int = 0 + fanout: int = 1 + + +def _driver(agent: object) -> FakeAgentDriver: + """The fake under an agent view.""" + return cast("Any", agent).driver + + +def _handle(session: object) -> Any: + """The fake session under a session view.""" + return cast("Any", session)._handle + + +def _echo(prompt: str, **_: Any) -> str: + return prompt + + +async def _yielding(prompt: str, **_: Any) -> str: + """Answers after letting the loop go on, as every real turn does.""" + await asyncio.sleep(0) + return prompt + + +class Counting(FakeAgentDriver): + """A fake whose sessions write down every close asked of them, and where it was asked. + + Args: + delay: How long a close takes, after it is asked. + """ + + def __init__(self, *, delay: float = 0.0, **kwargs: Any) -> None: + super().__init__(**kwargs) + self.delay = delay + self.closes: list[tuple[str | None, int]] = [] + + async def open(self, *args: Any, **kwargs: Any) -> Any: + session = await super().open(*args, **kwargs) + real = session.close + + async def close() -> None: + self.closes.append((session.id, threading.get_ident())) + if self.delay: + await asyncio.sleep(self.delay) + await real() + + session.close = close + return session + + +# ------------------------------------------------------------------------ a round apiece + + +@flow(agents=Solo, envs=Place, params=Rounds) +async def rounds( + task: str, *, agents: Solo, envs: Place, params: Rounds, ctx: FlowContext +) -> tuple[int, int, int]: + """A fresh session a round, a turn in it, and nothing kept.""" + agent, env = agents["agent"], envs["env"] + driver = _driver(agent) + live = kept = hooked = 0 + for _ in range(params.rounds): + session = await agent.spawn(env=env) + await agent.run(task, session=session) + live = max(live, driver.live) + kept = max(kept, len(cast("Any", ctx).res or ())) + hooked = max(hooked, len(cast("Any", agent)._line.sessions)) + return live, kept, hooked + + +@pytest.mark.parametrize("reply", [None, _yielding], ids=["instant", "yielding"]) +async def test_ten_thousand_rounds_hold_two_sessions_open_at_most(reply: Any) -> None: + driver = FakeAgentDriver(reply=reply) + live, kept, hooked = await run_fake( + rounds, "go", agents={"agent": driver}, params={"rounds": 10_000} + ) + assert len(driver.sessions) == 10_000 + assert max(live, driver.peak) <= 2, "sessions let go of were left open" + assert kept <= 2, f"the call kept {kept} sessions it had closed" + assert hooked <= 2, f"the agent kept {hooked} sessions it had closed" + assert driver.live == 0 + assert all(one.closed for one in driver.sessions) + + +@flow(agents=Solo, envs=Place, params=Rounds) +async def fanned( + task: str, *, agents: Solo, envs: Place, params: Rounds, ctx: FlowContext +) -> int: + """Rounds in batches of `fanout` gathered at once, each in a session of its own.""" + agent, env = agents["agent"], envs["env"] + driver = _driver(agent) + live = 0 + + async def one() -> str: + session = await agent.spawn(env=env) + return await agent.run(task, session=session) + + for _ in range(params.rounds // params.fanout): + said = await asyncio.gather(*(one() for _ in range(params.fanout))) + assert said == [task] * params.fanout + live = max(live, driver.live) + return live + + +@flow(agents=Solo, envs=Place, params=Rounds) +async def workers( + task: str, *, agents: Solo, envs: Place, params: Rounds, ctx: FlowContext +) -> None: + """`fanout` loops at once, each opening a fresh session a round.""" + agent, env = agents["agent"], envs["env"] + + async def worker(rounds: int) -> None: + for _ in range(rounds): + session = await agent.spawn(env=env) + await agent.run(task, session=session) + + await asyncio.gather( + *(worker(params.rounds // params.fanout) for _ in range(params.fanout)) + ) + + +async def test_gathered_rounds_hold_a_batch_open_at_most() -> None: + driver = FakeAgentDriver(reply=_yielding) + live = await run_fake( + fanned, + "go", + agents={"agent": driver}, + params={"rounds": 10_000, "fanout": 100}, + ) + assert len(driver.sessions) == 10_000 + assert live <= 100 + assert driver.peak <= 100 + 2, f"{driver.peak} open at once, in batches of 100" + assert driver.live == 0 + assert all(one.closed for one in driver.sessions) + + +async def test_loops_gathered_hold_two_sessions_apiece_at_most() -> None: + driver = FakeAgentDriver(reply=_yielding) + await run_fake( + workers, "go", agents={"agent": driver}, params={"rounds": 10_000, "fanout": 10} + ) + assert len(driver.sessions) == 10_000 + assert driver.peak <= 2 * 10, f"{driver.peak} open at once, by 10 loops" + assert driver.live == 0 + + +async def test_an_on_disk_flow_run_over_drivers_holds_few_open(tmp_path: Path) -> None: + flows = flowverse( + tmp_path, + { + "churn": """ + from hmz.flows import Agent, AgentCollection, Env, EnvCollection + from hmz.flows import FlowContext, FlowParams, flow + + class Solo(AgentCollection): + agent: Agent + + class Place(EnvCollection): + env: Env + + class Rounds(FlowParams): + rounds: int = 0 + + @flow(agents=Solo, envs=Place, params=Rounds) + async def churn(task, *, agents, envs, params, ctx: FlowContext) -> int: + agent = agents["agent"] + for _ in range(params.rounds): + session = await agent.spawn(env=envs["env"]) + await agent.run(task, session=session) + return params.rounds + """ + }, + ) + driver = FakeAgentDriver() + said = await run_flow( + load(str(flows / "churn")), + "go", + agents={"agent": driver}, + envs={"env": FakeEnvDriver()}, + params={"rounds": "1000"}, + budget=Budget(cost=math.inf), + ) + assert said == 1_000 + assert len(driver.sessions) == 1_000 + assert driver.peak <= 2 + assert driver.live == 0 + + +# ------------------------------------------------------------------------- what is kept + + +async def test_a_session_the_flow_still_holds_stays_open() -> None: + @flow(agents=Solo, envs=Place, params=Rounds) + async def holding( + task: str, *, agents: Solo, envs: Place, params: Rounds, ctx: FlowContext + ) -> tuple[list[bool], list[bool]]: + agent, env = agents["agent"], envs["env"] + kept = await agent.spawn(env=env) + await agent.run("one", session=kept) + listed = [await agent.spawn(env=env)] + mapped = {"s": await agent.spawn(env=env)} + handles = [_handle(one) for one in (kept, listed[0], mapped["s"])] + for _ in range(50): + await agent.run("churn", session=await agent.spawn(env=env)) + gc.collect() + await asyncio.sleep(0) + for one in (kept, listed[0], mapped["s"]): + await agent.run("still here", session=one) + del one + before = [one.closed for one in handles] + listed.clear() + del mapped["s"] + await asyncio.sleep(0) + return before, [one.closed for one in handles] + + before, after = await run_fake(holding) + assert before == [False, False, False] + assert after == [False, True, True] + + +async def test_a_fork_outlives_the_session_it_was_cut_from() -> None: + @flow(agents=Solo, envs=Place, params=Rounds) + async def forking( + task: str, *, agents: Solo, envs: Place, params: Rounds, ctx: FlowContext + ) -> tuple[list[bool], str, str, bool]: + agent, env = agents["agent"], envs["env"] + parent = await agent.spawn(env=env) + await agent.run("a", session=parent) + child = await agent.fork(parent, env=env) + cut = _handle(parent) + del parent + await asyncio.sleep(0) + # A harness cuts the fork as its first turn goes, so the parent is kept till then. + closed = [cut.closed] + said = await agent.run("b", session=child) + await asyncio.sleep(0) + closed.append(cut.closed) + orphan = await agent.fork(await agent.spawn(env=env), env=env) + await asyncio.sleep(0) + again = await agent.run("c", session=orphan) + return closed, said, again, _handle(child).closed + + driver = FakeAgentDriver(reply=_echo) + assert await run_fake(forking, agents={"agent": driver}) == ( + [False, True], + "b", + "c", + False, + ) + assert driver.sessions[1].prompts == ["a", "b"] + assert [one.closed for one in driver.sessions] == [True] * 4 + + +async def test_a_fork_whose_first_turn_failed_still_holds_its_parent() -> None: + refused: list[bool] = [True] + + async def refusing( + params: UserPromptSubmitHookParams, + ) -> UserPromptSubmitHookResult: + del params + return UserPromptSubmitHookResult(block=refused.pop(), reason="not yet") + + @flow(agents=Solo, envs=Place, params=Rounds) + async def forking( + task: str, *, agents: Solo, envs: Place, params: Rounds, ctx: FlowContext + ) -> list[bool]: + agent, env = agents["agent"], envs["env"] + parent = await agent.spawn(env=env) + await agent.run("a", session=parent) + child = await agent.fork(parent, env=env) + cut = _handle(parent) + del parent + agent.on_user_prompt_submit(refusing) + with pytest.raises(SessionError): + await agent.run("b", session=child) + await asyncio.sleep(0) + closed = [cut.closed] + agent.on_user_prompt_submit(None) + await agent.run("b", session=child) + await asyncio.sleep(0) + return [*closed, cut.closed] + + assert await run_fake(forking) == [False, True] + + +async def test_a_chain_of_forks_holds_few_open() -> None: + @flow(agents=Solo, envs=Place, params=Rounds) + async def chaining( + task: str, *, agents: Solo, envs: Place, params: Rounds, ctx: FlowContext + ) -> int: + agent, env = agents["agent"], envs["env"] + driver = _driver(agent) + live = 0 + session = await agent.spawn(env=env) + for _ in range(params.rounds): + await agent.run(task, session=session) + session = await agent.fork(session, env=env) + live = max(live, driver.live) + return live + + driver = FakeAgentDriver() + live = await run_fake( + chaining, "go", agents={"agent": driver}, params={"rounds": 1_000} + ) + assert len(driver.sessions) == 1_001 + assert max(live, driver.peak) <= 3 + assert driver.live == 0 + + +async def test_a_session_handed_up_closes_with_the_call_that_opened_it() -> None: + driver = Counting() + + @flow(agents=Solo, envs=Place, params=Rounds) + async def opening( + task: str, *, agents: Solo, envs: Place, params: Rounds, ctx: FlowContext + ) -> Any: + session = await agents["agent"].spawn(env=envs["env"]) + await agents["agent"].run("go", session=session) + return session + + @flow(agents=Solo, envs=Place, params=Rounds) + async def caller( + task: str, *, agents: Solo, envs: Place, params: Rounds, ctx: FlowContext + ) -> tuple[bool, int, int]: + session = await opening(task, agents=agents, envs=envs, params=params) + closed = _handle(session).closed + with pytest.raises(SessionError): + await agents["agent"].run("more", session=session) + closes = len(driver.closes) + del session + await asyncio.sleep(0) + return closed, closes, len(driver.closes) + + assert await run_fake(caller, agents={"agent": driver}) == (True, 1, 1) + assert len(driver.closes) == 1 + + +# ---------------------------------------------------------------------- closed, once + + +async def test_a_session_let_go_of_as_its_call_ends_is_closed_once() -> None: + driver = Counting() + + @flow(agents=Solo, envs=Place, params=Rounds) + async def leaving( + task: str, *, agents: Solo, envs: Place, params: Rounds, ctx: FlowContext + ) -> None: + session = await agents["agent"].spawn(env=envs["env"]) + await agents["agent"].run("go", session=session) + + @flow(agents=Solo, envs=Place, params=Rounds) + async def caller( + task: str, *, agents: Solo, envs: Place, params: Rounds, ctx: FlowContext + ) -> list[int]: + counted: list[int] = [] + for _ in range(20): + await leaving(task, agents=agents, envs=envs, params=params) + counted.append(len(driver.closes)) + await asyncio.sleep(0) + return counted + + assert await run_fake(caller, agents={"agent": driver}) == list(range(1, 21)) + assert sorted(set(driver.closes)) == sorted(driver.closes) + + +async def test_a_close_under_way_as_its_call_ends_is_waited_for() -> None: + driver = Counting(delay=0.05) + + @flow(agents=Solo, envs=Place, params=Rounds) + async def dropping( + task: str, *, agents: Solo, envs: Place, params: Rounds, ctx: FlowContext + ) -> int: + session = await agents["agent"].spawn(env=envs["env"]) + await agents["agent"].run("go", session=session) + del session + await asyncio.sleep(0) + return len(driver.closes) + + @flow(agents=Solo, envs=Place, params=Rounds) + async def caller( + task: str, *, agents: Solo, envs: Place, params: Rounds, ctx: FlowContext + ) -> tuple[int, bool]: + started = await dropping(task, agents=agents, envs=envs, params=params) + return started, driver.sessions[0].closed + + assert await run_fake(caller, agents={"agent": driver}) == (1, True) + assert len(driver.closes) == 1 + + +@pytest.mark.parametrize("how", ["released", "collected"]) +async def test_a_session_let_go_of_on_another_thread_is_closed_on_the_loop( + how: str, +) -> None: + driver = Counting() + loop = asyncio.get_running_loop() + debug = loop.get_debug() + + @flow(agents=Solo, envs=Place, params=Rounds) + async def elsewhere( + task: str, *, agents: Solo, envs: Place, params: Rounds, ctx: FlowContext + ) -> tuple[list[tuple[str | None, int]], str | None]: + session = await agents["agent"].spawn(env=envs["env"]) + await agents["agent"].run("go", session=session) + name = _handle(session).id + knot: list[Any] = [session] + del session + if how == "released": + await asyncio.to_thread(knot.clear) + else: + knot.append(knot) + del knot + await asyncio.to_thread(gc.collect) + for _ in range(10): + if driver.closes: + break + await asyncio.sleep(0) + return list(driver.closes), name + + # A loop in debug mode refuses a call from another thread that is not thread-safe. + loop.set_debug(True) + gc.disable() + try: + closes, name = await run_fake(elsewhere, agents={"agent": driver}) + finally: + gc.enable() + loop.set_debug(debug) + assert closes == [(name, threading.get_ident())] + assert len(driver.closes) == 1 + + +# ------------------------------------------------------------------------------- hooks + + +async def test_a_hook_heard_once_a_session_is_let_go_of_gets_one_that_is_over() -> None: + heard: list[tuple[Any, Any, type[BaseException] | None]] = [] + + async def ending(params: SessionEndHookParams) -> SessionEndHookResult: + session = params.session + refused: type[BaseException] | None = None + try: + await cast("Any", session).agent.run("more", session=session) + except SessionError as error: + refused = type(error) + heard.append((session, cast("Any", session).id, refused)) + return SessionEndHookResult() + + @flow(agents=Solo, envs=Place, params=Rounds) + async def dropping( + task: str, *, agents: Solo, envs: Place, params: Rounds, ctx: FlowContext + ) -> tuple[Any, str]: + agent = agents["agent"] + agent.on_session_end(ending) + session = await agent.spawn(env=envs["env"]) + await agent.run("go", session=session) + name = _handle(session).id + del session + await asyncio.sleep(0) + return agent, name + + driver = FakeAgentDriver() + agent, name = await run_fake(dropping, agents={"agent": driver}) + assert len(heard) == 1 + session, said, refused = heard[0] + assert (said, refused) == (name, SessionError) + assert session.agent is agent + assert driver.sessions[0].closed + + +class BoomError(Exception): + pass + + +async def test_a_hook_failing_between_turns_keeps_no_session_open() -> None: + async def notified(params: NotificationHookParams) -> NotificationHookResult: + raise BoomError(params.message) + + @flow(agents=Solo, envs=Place, params=Rounds) + async def failing( + task: str, *, agents: Solo, envs: Place, params: Rounds, ctx: FlowContext + ) -> bool: + agent = agents["agent"] + agent.on_notification(notified) + session = await agent.spawn(env=envs["env"]) + handle = _handle(session) + await handle.notify("between turns") + del session + await asyncio.sleep(0) + return handle.closed + + # Nothing collected: the session is to close for being let go of, not for a collection. + gc.disable() + try: + assert await run_fake(failing) + finally: + gc.enable() diff --git a/tests/unit/flows/test_fakes.py b/tests/unit/flows/test_fakes.py index 33076430..1d78fd25 100644 --- a/tests/unit/flows/test_fakes.py +++ b/tests/unit/flows/test_fakes.py @@ -31,6 +31,7 @@ Outworlder, Permission, SessionError, + TempCloneBusy, UnsupportedOperation, WorktreeError, flow, @@ -249,6 +250,29 @@ async def test_a_fake_env_copies_and_forgets() -> None: ) +async def test_a_fake_env_lets_go_of_a_copy_as_it_is_closed_and_keeps_it() -> None: + env = FakeEnvDriver({"a.txt": "A"}) + clone = await env.derive_temp_clone("c", holder="first") + await clone.write("b.txt", b"B") + with pytest.raises(TempCloneBusy): + await env.derive_temp_clone("c", holder="second") + await clone.close() + again = await env.derive_temp_clone("c", holder="second") + assert again.workdir == clone.workdir + assert again.files == {"a.txt": b"A", "b.txt": b"B"} + with pytest.raises(TempCloneBusy): + await env.derive_temp_clone("c", holder="first") + await clone.close() + with pytest.raises(TempCloneBusy): + await env.derive_temp_clone("c", holder="first") + await env.close() + third = await env.derive_temp_clone("c", holder="third") + assert third.workdir == clone.workdir + assert env.clones == ["c"] + await env.destroy_temp_clone("c") + assert (env.clones, list(env.machine)) == ([], ["/work/a.txt"]) + + async def test_a_fake_env_refuses_a_worktree_where_one_is() -> None: env = FakeEnvDriver({"a.txt": "A"}) await env.derive_worktree(ref="main", dir="/elsewhere")