diff --git a/docs/PROVIDERS.md b/docs/PROVIDERS.md index 337598b3b..813da46d8 100644 --- a/docs/PROVIDERS.md +++ b/docs/PROVIDERS.md @@ -255,9 +255,11 @@ processes without a scheduler service or in-memory queue. Waiting async calls are cancellation-safe; exceptions and cancellation release acquired locks; process exit releases kernel locks. Queue time consumes the caller's existing provider deadline, so waiting cannot silently extend a request beyond its -configured timeout. Cogitate holds one permit for its run because the OpenHands -SDK owns its internal multi-turn HTTP calls; this is conservative and avoids an -uncontrolled second path to the same server. +configured timeout. Cogitate holds one parent permit across model turns, but +temporarily yields that permit while the OpenHands `sol` tool runs a nested +`sol` child process. The parent reacquires through the same FIFO admission pool +before any further model request; failure to reacquire is a terminal +`local_queue_timeout`. Every bundled-local attempt appends a content-free JSON record to `health/local-inference/YYYYMMDD.jsonl`. These files follow the configured diff --git a/docs/design/local-permit-yield-on-sol-tool.md b/docs/design/local-permit-yield-on-sol-tool.md new file mode 100644 index 000000000..df3753f78 --- /dev/null +++ b/docs/design/local-permit-yield-on-sol-tool.md @@ -0,0 +1,454 @@ +# Local Permit Yield On Sol Tool + +## Scope + +This note covers local cogitate runs that already hold a governed local +inference permit and then invoke the OpenHands `sol` tool. It changes only the +permit lifecycle around that `sol` tool call. + +It does not change cloud providers, confidential local endpoints with no +governed `parallel_slots`, raw-read tools, bundled-local condenser calls, +provider/model resolution, cogitate policy, command policy, on-disk admission +artifacts, or telemetry schema. + +Ground truth from the installed OpenHands SDK: + +- `conversation.pause()` is not safe from a tool worker thread in the non-ACP + `arun` path because `arun` holds `self._state` across + `await agent.astep(...)` and `pause()` reacquires the same state lock. +- `conversation.interrupt()` is safe from a worker thread while `arun` is + active: it sets a thread-safe cancellation token and schedules task + cancellation with `loop.call_soon_threadsafe`. +- `arun` catches that `CancelledError`, emits synthetic orphaned-action + observations, sets status to `PAUSED`, and returns normally. The next LLM + request is not issued. +- `ParallelToolExecutor._run_safe()` converts every `Exception` from a tool + executor into an `AgentErrorEvent`. A normal exception from `SolExecutor` + cannot terminate the run by itself. Escaping a `BaseException` bypasses SDK + cleanup and is rejected. +- Multiple concurrent `sol` calls are not reachable today: non-native local + tool parsing yields one tool call, agent tool concurrency defaults to `1`, + and undeclared `sol` resources resolve to a tool-wide mutex. +- `_run_command()` runs a real child process and inherits environment, so the + nested `sol` process contends on the same + `health/local-inference-admission/` flock files. + +## Decisions + +### D1. Lease Seam + +Decision: introduce one local admission lease object owned by +`local.run_cogitate()` and pass it into the OpenHands provider as a keyword-only +argument. + +Public shape: + +- `local.run_cogitate()` creates a lease only when it actually acquired a + governed local permit: bundled local, or non-confidential BYO local endpoint + with `parallel_slots`. +- `openhands.run_cogitate(config, on_event=None, *, slot_lease=None)` accepts the + lease. +- `_build_sol_tools(..., slot_lease=None)` passes it into `SolExecutor`. +- Cloud callers pass nothing. Provider registry routing remains unchanged: + cloud providers call `solstone.think.providers.openhands`, while `local` + calls `solstone.think.providers.local`. +- Confidential local endpoints with `parallel_slots is None` create no lease and + keep current behavior. + +Ownership stays with `local.run_cogitate()`. The current `permit.release()` in +its `finally` becomes `slot_lease.close()` when a lease exists. OpenHands and +the executor may yield and reacquire the lease, but they do not own final +release. + +Implementation detail: after creating the lease, clear the old local `permit` +variable or route all telemetry/final release reads through the lease. Do not +leave both `permit.release()` and `lease.close()` active for the same permit. + +### D2. Lease Location And Shape + +Decision: implement the lease in `solstone/think/providers/local_admission.py`, +next to `LocalPermit`, because it is a small lifecycle wrapper around the +existing admission primitive and needs the same cancellation semantics. + +Lease state: + +- `capacity: int` +- `deadline: float`, an absolute `time.monotonic()` deadline captured by + `local.run_cogitate()` as `started + timeout` +- current `LocalPermit | None` +- `threading.Lock` +- `threading.Event` used to cancel pending reacquire +- closed flag +- initial telemetry values copied from the first permit + +Methods: + +- `yield_slot()`: synchronous, worker-thread safe. It removes the current permit + under the lock and releases it outside the lock. If the lease is closed, it is + a no-op after releasing any held permit. If the lease is already yielded, raise + loudly because this means the SDK concurrency contract changed. +- `reacquire()`: synchronous. It computes remaining time from the absolute + deadline and calls `acquire_local_slot(capacity, remaining, + cancel_event=...)`. On success it stores the new permit under the lock. If the + lease was closed or cancelled before storage, it immediately releases the + newly acquired permit and raises the cancellation exception. Deadline expiry + raises `LocalAdmissionTimeout`. +- `cancel_pending_reacquire()`: sets the cancel event without releasing a + currently held permit. This is what `SolExecutor.interrupt()` uses. +- `close()`: sets closed and the cancel event, releases any currently held + permit, and makes any in-flight reacquire drop its ticket and stop. If a + reacquire wins a slot in the race with `close()`, the post-acquire closed + check releases that permit immediately. + +The lease should be single-flight. `SolExecutor` should also hold a private +`threading.Lock` around `yield -> command -> reacquire` so a future SDK bump +cannot run two `sol` actions from the same parent permit at the same time. + +### D3. Failed Reacquire Termination + +Decision: failed reacquire is a terminal local admission failure, but the tool +executor must stop the SDK run with `interrupt()` rather than by raising the +exception directly. + +Flow in `SolExecutor.__call__`: + +1. Resolve and validate the command as today. +2. If there is no lease, run the command exactly as today. +3. With the executor's lease lock held: + - call `slot_lease.yield_slot()` immediately before `_run_command()`; + - call `_run_command()`; + - call `slot_lease.reacquire()` in a `finally`, before returning any + observation or re-raising any command exception. +4. If `reacquire()` raises `LocalAdmissionTimeout`: + - store that exact exception on the executor behind a lock; + - call `conversation.interrupt()` before returning, so the task cancellation + is queued before the worker future completes; + - return an error observation for history consistency. +5. If `reacquire()` raises `LocalAdmissionCancelled`, swallow it and return the + observation already produced by the command, or an internal error observation + if the command did not produce one. `LocalAdmissionCancelled` is not terminal. + +The invariant is: once the parent permit is yielded, control cannot leave +`SolExecutor.__call__` unless the parent permit is held again or a terminal +`LocalAdmissionTimeout` marker has been stored and `conversation.interrupt()` +has been called. Policy-denied and read-budget-exhausted returns happen before +the yield and therefore never enter this path. + +Flow in `openhands.run_cogitate()`: + +1. Keep the `SolExecutor` reference returned by `_build_sol_tools()`. +2. Bind the constructed `Conversation` back onto the executor as a fallback. + The installed SDK passes `conversation` into tool calls, but this pin avoids a + silent behavioral change if the SDK call signature moves. +3. After `arun` returns and before wall-clock, cost, turn, stuck, paused, or + no-output classification, check `sol_executor.take_terminal_error()`. +4. If present, call `conversation.close()` and re-raise the exact stored + `LocalAdmissionTimeout`. + +Exact type preservation: + +- Bundled: `local.run_cogitate()` already treats `LocalAdmissionTimeout` as a + timeout and, because `local_endpoint_reason_copy("local_queue_timeout")` + returns `None`, re-raises the original exception. +- Governed BYO: `classify_byo_cogitate_error(exc)` returns `None` for + `LocalAdmissionTimeout`; `getattr(exc, "reason_code")` supplies + `local_queue_timeout`; `local_endpoint_reason_copy("local_queue_timeout")` + returns `None`; the original exception is re-raised. +- Tests should pin both paths so a future endpoint-copy mapping cannot wrap this + exception into `LocalProviderError`. + +The executor needs the live `Conversation`. The installed SDK supplies it via +`tool(action_event.action, conversation)` inside `_execute_action_event()`. +Implementation should still bind the conversation onto the executor after +construction and use `conversation or self._conversation` in `__call__`. + +### D4. Budget Arithmetic + +Decision: reacquire uses the same absolute budget as the parent local cogitate +run. + +Deadline: + +- `local.run_cogitate()` already captures `started = time.monotonic()` and + `timeout = config["timeout_seconds"] or 600`. +- The lease deadline is `started + timeout`. +- Every reacquire computes `remaining = deadline - time.monotonic()` at the + moment it starts waiting. If `remaining <= 0`, it raises + `LocalAdmissionTimeout` immediately. + +Composition: + +- `_SHELL_TIMEOUT_SECONDS` remains 30 seconds. The child command can consume up + to that amount while the parent permit is yielded. +- `openhands.run_cogitate()` still derives an in-process wall-clock deadline + inside the same budget: `timeout_seconds - 30s`, or half the timeout for very + small budgets. +- If the shell command plus queue wait leaves no budget, the parent terminates + as `LocalAdmissionTimeout`. +- If reacquire blocks past the OpenHands wall-clock deadline, `run_task.cancel()` + cancels `arun`; `_arun_safe()` calls `SolExecutor.interrupt()`; the lease + cancel event wakes or bounds the blocking reacquire; `local.run_cogitate()`'s + `finally` closes the lease. Any late-acquired slot is released by the + post-acquire closed check. + +The parent is allowed to make another model request only after +`slot_lease.reacquire()` has succeeded and stored a current permit. + +### D5. Admission Cancellation + +Decision: extend only the synchronous admission primitive with optional +cancellation, because the new blocking reacquire runs in a worker thread and +uses sync admission. + +Public shape: + +- `acquire_local_slot(capacity, timeout_s, *, exclusive=False, + cancel_event: threading.Event | None = None)` + +Semantics: + +- If `cancel_event` is set before ticket creation, raise + `LocalAdmissionCancelled`. +- If it is set while queued, raise `LocalAdmissionCancelled`. +- If a permit is acquired and the event is set before returning, release the + permit immediately and raise `LocalAdmissionCancelled`. +- Tickets are dropped in the existing `finally` block for timeout, cancellation, + and unexpected errors. +- `LocalAdmissionCancelled` is not a provider readiness reason and should not + be mapped to owner-facing copy. It is a plain internal `Exception`, not a + `LocalAdmissionTimeout` subclass, has no `reason_code`, is not exported from + `local_admission.__all__`, and is never stored as a terminal marker. + +The wait loop can keep the existing 25 ms poll cadence, but should use +`cancel_event.wait(sleep_s)` when a cancel event is present so close wakes the +worker promptly. Async admission does not need this parameter for this design. + +`SolExecutor.interrupt()` should override the base no-op and call +`slot_lease.cancel_pending_reacquire()` when a lease exists. It does not need to +kill the child process in this design; `_run_command()` already has a 30 second +subprocess timeout, and any later reacquire will see the cancelled lease and +exit without taking or leaking a parent permit. + +`slot_lease.cancel_pending_reacquire()` only sets cancellation state; it does +not release a currently held permit. `slot_lease.close()` sets closed state, +cancels any pending reacquire, and releases a currently held permit. The +separation matters because the SDK can interrupt a worker while +`local.run_cogitate()` still owns cleanup. + +### D6. Telemetry + +Decision: keep telemetry schema unchanged and report the original top-level +parent admission values. + +- `queue_wait_ms` becomes `slot_lease.initial_queue_wait_ms`. + Justification: this preserves the field's existing meaning as the wait for + the cogitate run's top-level admission; internal yield/reacquire waits remain + reflected in wall time and terminal reason. +- `admission_slot` becomes `slot_lease.initial_slot_index`. + Justification: one cogitate row represents one parent run, so the stable + original admission slot is less misleading than a last slot after a temporary + yield. + +Do not add fields for yield count, reacquire queue time, or final slot. If a +future observability pass needs that detail, it should be a separate schema +change. + +### D7. Unchanged Paths + +These stay put: + +- Raw-read tools in `read_tools.py` keep the permit held for their whole tool + execution. They do not receive the lease. +- The bundled-local condenser keeps using the parent-held permit. +- Policy classification and read budgets are unchanged. +- `parallel_slots` resolution is unchanged. +- Local endpoint error copy is unchanged. +- No new on-disk artifacts are introduced, so no layer-ownership row is needed. + +## Docs Correction + +Update `docs/PROVIDERS.md` in the "Local admission and bundled inference +telemetry" section. Replace the current statement: + +- `Cogitate holds one permit for its run because the OpenHands SDK owns its internal multi-turn HTTP calls; this is conservative and avoids an uncontrolled second path to the same server.` + +with: + +- `Cogitate holds one parent permit across model turns, but temporarily yields that permit while the OpenHands sol tool runs a nested sol child process. The parent reacquires through the same FIFO admission pool before any further model request; failure to reacquire is a terminal local_queue_timeout.` + +Keep the surrounding claims about governed lanes, confidential bypass, and +bundled-only telemetry. + +## File-By-File Diff Plan + +### `solstone/think/providers/local_admission.py` + +- Add `LocalAdmissionCancelled`. +- Add optional `cancel_event` support to `acquire_local_slot()`. +- Add the lease class beside `LocalPermit`. +- Keep existing `LocalPermit` and exclusive-mode behavior intact. + +### `solstone/think/providers/local.py` + +- In `run_cogitate()`, after initial governed acquire, wrap the permit in the + lease with `capacity.parallel_slots` or `endpoint.parallel_slots` and + `deadline=started + timeout`. +- Pass `slot_lease=lease` to `openhands.run_cogitate()`. +- Replace final `permit.release()` with `lease.close()` for lease-backed runs. +- Route bundled cogitate telemetry through lease initial telemetry properties. +- Preserve exact `LocalAdmissionTimeout` behavior for bundled and governed BYO. + +### `solstone/think/providers/openhands.py` + +- Add keyword-only `slot_lease=None` to `run_cogitate()`. +- Thread the lease through `_build_sol_tools()` and `SolExecutor`. +- Keep and bind the `SolExecutor` reference after creating the `Conversation`. +- In `SolExecutor.__call__`, use the lease around `_run_command()` and terminal + marker plus `conversation.interrupt()` on reacquire timeout. +- Add `SolExecutor.interrupt()` to cancel pending reacquire. +- Add helper methods on `SolExecutor` for storing and retrieving the terminal + exception. +- Check the terminal marker immediately after `arun` returns and before existing + wall-clock/cost/turn/stuck classification. + +### `docs/PROVIDERS.md` + +- Apply the wording correction above. + +### Tests + +- Update local-provider and admission tests for the new lease. +- Expand real SDK shape pins in `tests/test_openhands_sdk_shape.py`. +- Keep fake OpenHands tests only for provider plumbing and event translation; + do not rely on them for SDK lock, interrupt, or worker-thread behavior. + +## Implementation Sequence + +1. Add cancel-aware sync admission and the lease in `local_admission.py`. +2. Refactor `local.run_cogitate()` to create, pass, close, and report through + the lease without changing behavior when no permit is acquired. +3. Thread the keyword-only lease through OpenHands tool construction. +4. Update `SolExecutor` for yield/reacquire, terminal marker, bound + conversation fallback, and interrupt cancellation. +5. Add the marker check in `openhands.run_cogitate()` before existing terminal + classifications. +6. Update `docs/PROVIDERS.md`. +7. Add tests in the order below, starting with lower-level admission tests, then + local provider tests, then real SDK shape pins. + +## Test Plan + +### `tests/test_local.py` + +- Replace `test_run_cogitate_byo_acquires_permit_and_records_no_telemetry`. + The fake `openhands.run_cogitate()` should accept `slot_lease`, call + `slot_lease.yield_slot()`, prove a nested `acquire_local_slot(1, 0.1)` + succeeds, release the nested permit, call `slot_lease.reacquire()`, and return + `"ok"`. On the old tree this fails because the parent never yields and the + nested acquire times out. +- Add bundled parity for the same yield/reacquire behavior, with telemetry still + written once and using the lease's initial queue wait and slot. +- Add a capacity-1 "parent holds again before resume" test. After + `slot_lease.reacquire()` and before the fake OpenHands run returns, a + competing `acquire_local_slot(1, 0.03)` must time out. This proves the parent + has reacquired before any possible next model turn. +- Add sol command failure coverage by monkeypatching `_run_command()` to return + `is_error=True`; assert the lease is reacquired, no lock leaks, and the result + stays an error observation rather than an ungoverned resume. +- Add sol command timeout coverage by monkeypatching `_run_command()` to return + the same shape it returns on `subprocess.TimeoutExpired`; assert reacquire and + no leak. +- Add failed reacquire coverage: yield the lease, queue or hold capacity so + `reacquire()` exceeds the remaining budget, and assert + `local.run_cogitate()` raises the exact `LocalAdmissionTimeout` with + `reason_code == "local_queue_timeout"`. +- Add BYO exact-type coverage for the failed reacquire path. Existing initial + queue timeout coverage proves the pre-OpenHands path; this pins the + post-`sol` reacquire path. + +### `tests/test_local_admission.py` + +- Add `test_sync_queue_cancel_event_drops_ticket`: hold capacity, start + `acquire_local_slot(..., cancel_event=event)` in a thread, wait until its + ticket exists, set the event, assert the thread exits with + `LocalAdmissionCancelled`, no wait tickets remain, and a new acquire succeeds. +- Add `test_cancel_event_after_acquire_releases_permit`: arrange the event to be + set at the acquire-return boundary, then assert the acquired permit is released + and capacity is fully restored. +- Add lease-specific close coverage: start `lease.reacquire()` while capacity is + held elsewhere, wait for the wait ticket, call `lease.close()`, assert no stale + ticket and full capacity restored. +- Add FIFO coverage for yielded parent: with parent yielded, an unrelated waiter + that already has the oldest ticket must acquire before the parent reacquire. + Synchronize on admission ticket files and thread events with bounded waits, + not sleeps. +- Add a cross-process nested acquire test using `subprocess.run([sys.executable, + "-c", ...])` with `SOLSTONE_JOURNAL` pointing at the tmp journal. The parent + lease yields before the child starts; the child performs a real + `acquire_local_slot()` against the same admission directory and exits `0`. +- Add capacity-2 two-parent coverage: create two leases holding capacity, yield + both, run two nested cross-process acquires, then reacquire both parents. + Synchronize on tickets/events with bounded waits and assert full capacity is + restored. + +### `tests/test_openhands_provider.py` + +- Add fake-provider plumbing tests for: + - `openhands.run_cogitate()` passes the lease into `_build_sol_tools()`; + - `SolExecutor` binds the live conversation fallback; + - a stored terminal timeout marker is checked before wall-clock, cost, turn, + stuck, paused, or no-output classification. +- Keep these tests focused on solstone provider logic only. Do not use fakes to + prove real SDK interrupt or `_run_safe` behavior. + +### `tests/test_openhands_sdk_shape.py` + +Use the installed SDK, not `tests/openhands_fakes.py`. + +- Add `test_interrupt_from_tool_worker_stops_before_next_llm_completion`. + Build a real `Conversation` with `openhands.sdk.testing.TestLLM`. First + scripted response calls a custom test tool; the tool runs on the executor + worker, calls `conversation.interrupt()`, and returns an observation. A second + scripted assistant response is available. Assert `arun()` returns paused and + `TestLLM` consumed exactly one completion. +- Add `test_run_safe_exception_becomes_agent_error_and_loop_continues`. + Use a real tool executor that raises `RuntimeError` on the first call and a + `TestLLM` script with a second response. Assert a second LLM completion is + consumed and an `AgentErrorEvent` exists. This proves marker plus interrupt is + necessary. +- Add `test_default_tool_concurrency_and_undeclared_tool_mutex_shape`. + Assert `Agent(...).tool_concurrency_limit == 1`; assert a plain + `ToolDefinition.declared_resources()` returns `declared=False`; assert + `ParallelToolExecutor._resolve_lock_keys()` maps that to `["tool:"]`. + If private helper access is too brittle, assert through two same-tool actions + with blocking executors and a max-workers executor that they serialize. + +These are feasible with the installed SDK because `TestLLM` exists and supports +scripted tool-call messages. The private lock-key assertion is the only brittle +part; prefer a behavior test if construction overhead is acceptable. + +### Raw-Read Tool Coverage + +- Add a local cogitate test where a raw-read tool executes while a competing + acquire attempts capacity 1. Because raw-read tools never receive the lease, + the competing acquire must time out until the tool finishes. This can be a + provider plumbing test with `build_read_tools()` monkeypatched to return a + simple blocking read tool, or a real SDK test if setup remains small. + +## Risks And Open Questions + +- `SolExecutor.interrupt()` will not terminate an already-running child + subprocess. The child is still bounded by `_SHELL_TIMEOUT_SECONDS` and owns its + own admission lifecycle. This is simpler and satisfies the no-leak constraint, + but it means a parent wall-clock timeout can leave the nested child running for + up to 30 seconds. +- The terminal marker relies on `conversation.interrupt()` being queued before + the worker future completes. The executor must call `interrupt()` before + returning the error observation. +- The lease is intentionally single-flight. If a future SDK permits concurrent + same-executor `sol` calls despite the current parser/defaults/mutex, tests + should fail loudly rather than allowing an ungoverned resume. +- Reacquire wait time is not represented separately in telemetry. That is a + deliberate schema-preserving choice. +- I verified the installed SDK exposes `TestLLM`, so the real SDK tests are + feasible. I did not prototype those tests in this design pass. diff --git a/solstone/think/providers/local.py b/solstone/think/providers/local.py index 16bcbabca..ccf651beb 100644 --- a/solstone/think/providers/local.py +++ b/solstone/think/providers/local.py @@ -888,6 +888,7 @@ async def run_cogitate( from solstone.think.providers import local_server, openhands from solstone.think.providers.local_admission import ( LocalAdmissionTimeout, + LocalSlotLease, acquire_local_slot_async, record_local_inference, ) @@ -899,7 +900,7 @@ async def run_cogitate( timeout = float(config.get("timeout_seconds", 600) or 600) server = None capacity = None - permit = None + slot_lease = None outcome = "success" reason_code: str | None = None try: @@ -910,12 +911,26 @@ async def run_cogitate( capacity.parallel_slots, _remaining_timeout(started, timeout), ) + slot_lease = LocalSlotLease( + capacity=capacity.parallel_slots, + deadline=started + timeout, + permit=permit, + ) elif endpoint.parallel_slots is not None: permit = await acquire_local_slot_async( endpoint.parallel_slots, _remaining_timeout(started, timeout), ) - return await openhands.run_cogitate(config, on_event=on_event) + slot_lease = LocalSlotLease( + capacity=endpoint.parallel_slots, + deadline=started + timeout, + permit=permit, + ) + return await openhands.run_cogitate( + config, + on_event=on_event, + slot_lease=slot_lease, + ) except asyncio.CancelledError: outcome = "cancelled" reason_code = "cancelled" @@ -968,8 +983,8 @@ async def run_cogitate( raise wrapped from exc raise finally: - if permit is not None: - permit.release() + if slot_lease is not None: + slot_lease.close() if server is not None and capacity is not None: record_local_inference( _telemetry_record( @@ -981,11 +996,15 @@ async def run_cogitate( capacity_source=capacity.source, started=started, queue_wait_ms=( - permit.queue_wait_ms - if permit is not None + slot_lease.initial_queue_wait_ms + if slot_lease is not None else (time.monotonic() - started) * 1000.0 ), - admission_slot=permit.slot_index if permit is not None else None, + admission_slot=( + slot_lease.initial_slot_index + if slot_lease is not None + else None + ), retry_index=None, outcome=outcome, finish_reason="stop" if outcome == "success" else None, diff --git a/solstone/think/providers/local_admission.py b/solstone/think/providers/local_admission.py index 26c6457ec..9705fb8c5 100644 --- a/solstone/think/providers/local_admission.py +++ b/solstone/think/providers/local_admission.py @@ -10,6 +10,7 @@ import errno import fcntl import logging import os +import threading import time import uuid from dataclasses import dataclass @@ -30,6 +31,10 @@ class LocalAdmissionTimeout(TimeoutError): reason_code = "local_queue_timeout" +class LocalAdmissionCancelled(Exception): + """Internal cancellation of a synchronous local admission wait.""" + + @dataclass class LocalPermit: """One flock-backed serving-capacity permit.""" @@ -189,17 +194,25 @@ def _deadline(started: float, timeout_s: float | None) -> float | None: def acquire_local_slot( - capacity: int, timeout_s: float | None, *, exclusive: bool = False + capacity: int, + timeout_s: float | None, + *, + exclusive: bool = False, + cancel_event: threading.Event | None = None, ) -> LocalPermit: """Wait synchronously for governed local serving capacity.""" if capacity < 1: raise ValueError("local inference capacity must be at least one") + if cancel_event is not None and cancel_event.is_set(): + raise LocalAdmissionCancelled("local inference admission was cancelled") started = time.monotonic() deadline = _deadline(started, timeout_s) root = _admission_dir() ticket = _create_ticket(root) try: while True: + if cancel_event is not None and cancel_event.is_set(): + raise LocalAdmissionCancelled("local inference admission was cancelled") if _ticket_has_turn(root, ticket): permit = ( _try_acquire_exclusive(capacity, started, root) @@ -207,6 +220,11 @@ def acquire_local_slot( else _try_acquire(capacity, started, root) ) if permit is not None: + if cancel_event is not None and cancel_event.is_set(): + permit.release() + raise LocalAdmissionCancelled( + "local inference admission was cancelled" + ) return permit now = time.monotonic() if deadline is not None and now >= deadline: @@ -216,7 +234,13 @@ def acquire_local_slot( sleep_s = _POLL_INTERVAL_S if deadline is not None: sleep_s = min(sleep_s, max(0.0, deadline - now)) - time.sleep(sleep_s) + if cancel_event is not None: + if cancel_event.wait(sleep_s): + raise LocalAdmissionCancelled( + "local inference admission was cancelled" + ) + else: + time.sleep(sleep_s) finally: _drop_ticket(ticket) @@ -254,6 +278,90 @@ async def acquire_local_slot_async( _drop_ticket(ticket) +class LocalSlotLease: + """Thread-safe owner for a local inference permit that can yield and reacquire.""" + + def __init__( + self, + *, + capacity: int, + deadline: float | None, + permit: LocalPermit, + ) -> None: + if capacity < 1: + raise ValueError("local inference capacity must be at least one") + self.capacity = capacity + self.deadline = deadline + self.initial_queue_wait_ms = permit.queue_wait_ms + self.initial_slot_index = permit.slot_index + self._permit: LocalPermit | None = permit + self._lock = threading.Lock() + self._cancel_event = threading.Event() + self._closed = False + self._reacquiring = False + + def yield_slot(self) -> None: + """Release the currently held permit without closing the lease.""" + with self._lock: + if self._closed: + raise LocalAdmissionCancelled("local inference lease is closed") + permit = self._permit + if permit is None: + raise RuntimeError("local inference lease has no held permit to yield") + self._permit = None + permit.release() + + def reacquire(self) -> LocalPermit: + """Reacquire the permit through the FIFO admission queue.""" + with self._lock: + if self._closed or self._cancel_event.is_set(): + raise LocalAdmissionCancelled("local inference lease is cancelled") + if self._permit is not None: + return self._permit + if self._reacquiring: + raise RuntimeError("local inference lease reacquire already in flight") + self._reacquiring = True + + try: + timeout_s = self._remaining_timeout() + permit = acquire_local_slot( + self.capacity, + timeout_s, + cancel_event=self._cancel_event, + ) + except BaseException: + with self._lock: + self._reacquiring = False + raise + + with self._lock: + self._reacquiring = False + if self._closed or self._cancel_event.is_set(): + permit.release() + raise LocalAdmissionCancelled("local inference lease is cancelled") + self._permit = permit + return permit + + def cancel_pending_reacquire(self) -> None: + """Cancel a pending or future reacquire without releasing a held permit.""" + self._cancel_event.set() + + def close(self) -> None: + """Close the lease, cancel pending reacquire, and release a held permit.""" + with self._lock: + self._closed = True + self._cancel_event.set() + permit = self._permit + self._permit = None + if permit is not None: + permit.release() + + def _remaining_timeout(self) -> float | None: + if self.deadline is None: + return None + return max(0.0, self.deadline - time.monotonic()) + + def record_local_inference(record: dict[str, Any]) -> None: """Durably append one prompt/output-free local inference record.""" try: @@ -271,6 +379,7 @@ def record_local_inference(record: dict[str, Any]) -> None: __all__ = [ "LocalAdmissionTimeout", "LocalPermit", + "LocalSlotLease", "acquire_local_slot", "acquire_local_slot_async", "record_local_inference", diff --git a/solstone/think/providers/openhands.py b/solstone/think/providers/openhands.py index f4228d202..5b3cde7ce 100644 --- a/solstone/think/providers/openhands.py +++ b/solstone/think/providers/openhands.py @@ -17,6 +17,7 @@ import math import os import shutil import sys +import threading import traceback import uuid from collections.abc import Callable @@ -46,6 +47,11 @@ from solstone.think.cogitate_policy import ( resolve_read_scope, ) from solstone.think.providers.cli import QuotaExhaustedError, assemble_prompt +from solstone.think.providers.local_admission import ( + LocalAdmissionCancelled, + LocalAdmissionTimeout, + LocalSlotLease, +) from solstone.think.providers.local_server import LOCAL_MIN_CONTEXT_TOKENS from solstone.think.providers.shared import ( USAGE_KEYS, @@ -259,16 +265,33 @@ def _ensure_sol_types() -> dict[str, Any]: policy: CogitatePolicy, callback: JSONEventCallback, read_call_budget: int, + slot_lease: LocalSlotLease | None = None, ) -> None: self.policy = policy self.callback = callback self.read_call_budget = read_call_budget + self.slot_lease = slot_lease self.read_call_count = 0 self._budget_exhausted_emitted = False + self._conversation: Any | None = None + self._terminal_error: LocalAdmissionTimeout | None = None + self._terminal_error_lock = threading.Lock() + self._slot_cycle_lock = threading.Lock() - def __call__(self, action: Any, conversation: Any = None) -> Any: - del conversation + def bind_conversation(self, conversation: Any) -> None: + self._conversation = conversation + + def take_terminal_error(self) -> LocalAdmissionTimeout | None: + with self._terminal_error_lock: + error = self._terminal_error + self._terminal_error = None + return error + def interrupt(self) -> None: + if self.slot_lease is not None: + self.slot_lease.cancel_pending_reacquire() + + def __call__(self, action: Any, conversation: Any = None) -> Any: command = str(action.command) decision = self.policy.classify_command(command) if not decision.allowed: @@ -293,9 +316,48 @@ def _ensure_sol_types() -> dict[str, Any]: ) assert decision.argv is not None - result = _run_command(decision.argv) + if self.slot_lease is None: + result = _run_command(decision.argv) + return SolObservation.from_text( + result["text"], is_error=result["is_error"] + ) + + with self._slot_cycle_lock: + self.slot_lease.yield_slot() + result: dict[str, Any] | None = None + command_error: Exception | None = None + try: + result = _run_command(decision.argv) + except Exception as exc: + command_error = exc + finally: + try: + self.slot_lease.reacquire() + except LocalAdmissionTimeout as exc: + self._store_terminal_error(exc) + live_conversation = conversation or self._conversation + if live_conversation is not None: + live_conversation.interrupt() + return SolObservation.from_text(str(exc), is_error=True) + except LocalAdmissionCancelled: + if result is not None: + return SolObservation.from_text( + result["text"], is_error=result["is_error"] + ) + return SolObservation.from_text( + "local_admission_cancelled: cogitate run interrupted " + "before reacquiring local inference", + is_error=True, + ) + if command_error is not None: + raise command_error + assert result is not None return SolObservation.from_text(result["text"], is_error=result["is_error"]) + def _store_terminal_error(self, error: LocalAdmissionTimeout) -> None: + with self._terminal_error_lock: + self._terminal_error = error + class SolTool(ToolDefinition[SolAction, SolObservation]): name = "sol" @@ -330,6 +392,7 @@ def _build_sol_tools( policy: CogitatePolicy, callback: JSONEventCallback, read_call_budget: int, + slot_lease: LocalSlotLease | None = None, ) -> tuple[list[Any], Any]: types = _ensure_sol_types() sol_action = types["SolAction"] @@ -342,6 +405,7 @@ def _build_sol_tools( policy=policy, callback=callback, read_call_budget=read_call_budget, + slot_lease=slot_lease, ) tool = sol_tool_cls( description=( @@ -1038,6 +1102,8 @@ def _wall_clock_deadline_s(timeout_seconds: float) -> float: async def run_cogitate( config: dict[str, Any], on_event: Callable[[dict], None] | None = None, + *, + slot_lease: LocalSlotLease | None = None, ) -> str | None: """Run a cogitate prompt through OpenHands SDK.""" callback = JSONEventCallback(on_event) @@ -1080,11 +1146,13 @@ async def run_cogitate( llm = _build_llm(provider, model) usage_start = _usage_snapshot(llm) tool_specs = [] + sol_executor = None if caps.sol: - sol_tools, _executor = _build_sol_tools( + sol_tools, sol_executor = _build_sol_tools( policy=policy, callback=callback, read_call_budget=read_call_budget, + slot_lease=slot_lease, ) # openhands-sdk v1.23 resolves Agent.tools by spec name via the # registry; passing ToolDefinition instances directly fails pydantic @@ -1146,6 +1214,8 @@ async def run_cogitate( visualizer=None, ) translator.conversation = conversation + if sol_executor is not None: + sol_executor.bind_conversation(conversation) conversation.send_message(prompt_body) timeout_seconds = float(config.get("timeout_seconds", 600) or 600) wall_clock_s = _wall_clock_deadline_s(timeout_seconds) @@ -1172,6 +1242,12 @@ async def run_cogitate( # generic except-Exception classification path unchanged. run_task.result() + if sol_executor is not None: + terminal_error = sol_executor.take_terminal_error() + if terminal_error is not None: + conversation.close() + raise terminal_error + result = translator.result() usage = _usage_delta(usage_start, llm) if wall_clock_exceeded: diff --git a/tests/openhands_fakes.py b/tests/openhands_fakes.py index 712018241..f3b00e402 100644 --- a/tests/openhands_fakes.py +++ b/tests/openhands_fakes.py @@ -141,6 +141,7 @@ class Conversation(FakeModel): super().__init__(**kwargs) self.messages: list[str] = [] self.paused = False + self.interrupted = False self.closed = False self.state = SimpleNamespace(execution_status=None) type(self).instances.append(self) @@ -151,6 +152,10 @@ class Conversation(FakeModel): def pause(self) -> None: self.paused = True + def interrupt(self) -> None: + self.interrupted = True + self.state.execution_status = "paused" + def close(self) -> None: self.closed = True diff --git a/tests/test_local.py b/tests/test_local.py index 87e734623..c915d851e 100644 --- a/tests/test_local.py +++ b/tests/test_local.py @@ -1817,7 +1817,12 @@ def test_run_cogitate_byo_acquires_permit_and_records_no_telemetry(monkeypatch): records = [] monkeypatch.setattr(local_admission, "record_local_inference", records.append) - async def fake_cogitate(*_args, **_kwargs): + async def fake_cogitate(*_args, slot_lease=None, **_kwargs): + assert slot_lease is not None + slot_lease.yield_slot() + with local_admission.acquire_local_slot(1, 0.1) as nested: + assert nested.slot_index == 0 + slot_lease.reacquire() with pytest.raises(local_admission.LocalAdmissionTimeout): local_admission.acquire_local_slot(1, 0.03) return "ok" @@ -1837,6 +1842,82 @@ def test_run_cogitate_byo_acquires_permit_and_records_no_telemetry(monkeypatch): assert permit.slot_index == 0 +def test_run_cogitate_byo_keeps_permit_for_non_sol_work(monkeypatch): + provider = _provider() + monkeypatch.setattr( + provider, + "resolve_local_endpoint", + lambda: _byo_endpoint(parallel_slots=1), + ) + monkeypatch.setattr( + "solstone.think.providers.local_server.connect", + lambda: (_ for _ in ()).throw(AssertionError("connect not expected")), + ) + + from solstone.think.providers import local_admission + + async def fake_cogitate(*_args, slot_lease=None, **_kwargs): + assert slot_lease is not None + with pytest.raises(local_admission.LocalAdmissionTimeout): + local_admission.acquire_local_slot(1, 0.03) + return "ok" + + monkeypatch.setattr( + "solstone.think.providers.openhands.run_cogitate", + fake_cogitate, + ) + + assert ( + asyncio.run(provider.run_cogitate({"model": LOCAL_MODEL, "timeout_seconds": 1})) + == "ok" + ) + + +@pytest.mark.parametrize("bundled", [False, True]) +def test_run_cogitate_reacquire_timeout_preserves_exact_type( + monkeypatch, + bundled, +): + provider = _provider() + if bundled: + _patch_bundled_server(monkeypatch) + else: + monkeypatch.setattr( + provider, + "resolve_local_endpoint", + lambda: _byo_endpoint(parallel_slots=1), + ) + monkeypatch.setattr( + "solstone.think.providers.local_server.connect", + lambda: (_ for _ in ()).throw(AssertionError("connect not expected")), + ) + + from solstone.think.providers import local_admission + + async def fake_cogitate(*_args, slot_lease=None, **_kwargs): + assert slot_lease is not None + slot_lease.yield_slot() + holder = local_admission.acquire_local_slot(1, 0.1) + try: + slot_lease.reacquire() + finally: + holder.release() + + monkeypatch.setattr( + "solstone.think.providers.openhands.run_cogitate", + fake_cogitate, + ) + + with pytest.raises(local_admission.LocalAdmissionTimeout) as exc: + asyncio.run( + provider.run_cogitate({"model": LOCAL_MODEL, "timeout_seconds": 0.03}) + ) + + assert exc.type is local_admission.LocalAdmissionTimeout + assert exc.value.reason_code == "local_queue_timeout" + assert not list(Path(local_admission._admission_dir()).glob("wait-*.ticket")) + + def test_run_cogitate_byo_queue_timeout_preserves_exact_type(monkeypatch): provider = _provider() monkeypatch.setattr( diff --git a/tests/test_local_admission.py b/tests/test_local_admission.py index 311f75967..fb97545be 100644 --- a/tests/test_local_admission.py +++ b/tests/test_local_admission.py @@ -6,6 +6,9 @@ from __future__ import annotations import asyncio import errno import json +import os +import subprocess +import sys import threading import time @@ -13,6 +16,7 @@ import pytest from solstone.think.providers.local_admission import ( LocalAdmissionTimeout, + LocalSlotLease, acquire_local_slot, acquire_local_slot_async, record_local_inference, @@ -26,6 +30,23 @@ def _isolated_journal(monkeypatch, tmp_path): think_utils._journal_path_cache = None +def _admission_root(tmp_path): + return tmp_path / "health" / "local-inference-admission" + + +def _wait_for(predicate, timeout_s: float = 1.0) -> None: + deadline = time.monotonic() + timeout_s + while not predicate(): + if time.monotonic() >= deadline: + raise AssertionError("condition was not met before timeout") + time.sleep(0.005) + + +def _wait_for_ticket_count(tmp_path, count: int) -> None: + root = _admission_root(tmp_path) + _wait_for(lambda: len(list(root.glob("wait-*.ticket"))) >= count) + + def test_cross_thread_admission_never_exceeds_capacity(monkeypatch, tmp_path): _isolated_journal(monkeypatch, tmp_path) active = 0 @@ -179,9 +200,104 @@ def test_async_queued_cancellation_does_not_leak(monkeypatch, tmp_path): asyncio.run(exercise()) +def test_sync_cancel_event_drops_ticket_and_releases_queue(monkeypatch, tmp_path): + _isolated_journal(monkeypatch, tmp_path) + + from solstone.think.providers import local_admission + + holder = acquire_local_slot(1, 0.1) + cancel_event = threading.Event() + errors: list[BaseException] = [] + + def wait_for_slot() -> None: + try: + acquire_local_slot(1, 2.0, cancel_event=cancel_event) + except BaseException as exc: + errors.append(exc) + + thread = threading.Thread(target=wait_for_slot) + thread.start() + _wait_for_ticket_count(tmp_path, 1) + + cancel_event.set() + thread.join(timeout=1.0) + holder.release() + + assert not thread.is_alive() + assert len(errors) == 1 + assert isinstance(errors[0], local_admission.LocalAdmissionCancelled) + assert not list(_admission_root(tmp_path).glob("wait-*.ticket")) + with acquire_local_slot(1, 0.1) as permit: + assert permit.slot_index == 0 + + +def test_sync_cancel_event_after_acquire_releases_permit(monkeypatch, tmp_path): + _isolated_journal(monkeypatch, tmp_path) + + from solstone.think.providers import local_admission + + cancel_event = threading.Event() + real_try_acquire = local_admission._try_acquire + + def cancel_after_acquire(capacity, started, root): + permit = real_try_acquire(capacity, started, root) + if permit is not None: + cancel_event.set() + return permit + + monkeypatch.setattr(local_admission, "_try_acquire", cancel_after_acquire) + + with pytest.raises(local_admission.LocalAdmissionCancelled): + acquire_local_slot(1, 0.1, cancel_event=cancel_event) + + monkeypatch.setattr(local_admission, "_try_acquire", real_try_acquire) + with acquire_local_slot(1, 0.1) as permit: + assert permit.slot_index == 0 + + +def test_lease_close_cancels_pending_reacquire_without_leaking( + monkeypatch, + tmp_path, +): + _isolated_journal(monkeypatch, tmp_path) + + from solstone.think.providers import local_admission + + initial = acquire_local_slot(1, 0.1) + lease = LocalSlotLease( + capacity=1, + deadline=time.monotonic() + 2.0, + permit=initial, + ) + lease.yield_slot() + holder = acquire_local_slot(1, 0.1) + errors: list[BaseException] = [] + + def reacquire() -> None: + try: + lease.reacquire() + except BaseException as exc: + errors.append(exc) + + thread = threading.Thread(target=reacquire) + thread.start() + _wait_for_ticket_count(tmp_path, 1) + + lease.close() + thread.join(timeout=1.0) + holder.release() + + assert not thread.is_alive() + assert len(errors) == 1 + assert isinstance(errors[0], local_admission.LocalAdmissionCancelled) + assert not list(_admission_root(tmp_path).glob("wait-*.ticket")) + with acquire_local_slot(1, 0.1) as permit: + assert permit.slot_index == 0 + + def test_waiters_are_admitted_in_ticket_order(monkeypatch, tmp_path): _isolated_journal(monkeypatch, tmp_path) - root = tmp_path / "health" / "local-inference-admission" + root = _admission_root(tmp_path) first = acquire_local_slot(1, 1) order: list[int] = [] @@ -207,6 +323,163 @@ def test_waiters_are_admitted_in_ticket_order(monkeypatch, tmp_path): assert order == list(range(5)) +def test_yielded_parent_reacquire_respects_existing_fifo_waiter( + monkeypatch, + tmp_path, +): + _isolated_journal(monkeypatch, tmp_path) + + initial = acquire_local_slot(1, 0.1) + lease = LocalSlotLease( + capacity=1, + deadline=time.monotonic() + 2.0, + permit=initial, + ) + order: list[str] = [] + waiter_entered = threading.Event() + release_waiter = threading.Event() + parent_reacquired = threading.Event() + + def waiter() -> None: + with acquire_local_slot(1, 2.0): + order.append("waiter") + waiter_entered.set() + assert release_waiter.wait(1.0) + + waiter_thread = threading.Thread(target=waiter) + waiter_thread.start() + _wait_for_ticket_count(tmp_path, 1) + + lease.yield_slot() + assert waiter_entered.wait(1.0) + + def parent() -> None: + lease.reacquire() + order.append("parent") + parent_reacquired.set() + + parent_thread = threading.Thread(target=parent) + parent_thread.start() + _wait_for_ticket_count(tmp_path, 1) + + assert order == ["waiter"] + release_waiter.set() + assert parent_reacquired.wait(1.0) + parent_thread.join(timeout=1.0) + waiter_thread.join(timeout=1.0) + lease.close() + + assert not parent_thread.is_alive() + assert not waiter_thread.is_alive() + assert order == ["waiter", "parent"] + + +def test_cross_process_nested_acquire_succeeds_while_parent_yielded( + monkeypatch, + tmp_path, +): + _isolated_journal(monkeypatch, tmp_path) + + initial = acquire_local_slot(1, 0.1) + lease = LocalSlotLease( + capacity=1, + deadline=time.monotonic() + 2.0, + permit=initial, + ) + env = {**os.environ, "SOLSTONE_JOURNAL": str(tmp_path)} + script = """ +from solstone.think.providers.local_admission import acquire_local_slot +with acquire_local_slot(1, 1.0) as permit: + assert permit.slot_index == 0 +""" + + try: + lease.yield_slot() + completed = subprocess.run( + [sys.executable, "-c", script], + env=env, + text=True, + capture_output=True, + timeout=2.0, + check=False, + ) + assert completed.returncode == 0, completed.stderr + lease.reacquire() + finally: + lease.close() + + with acquire_local_slot(1, 0.1) as permit: + assert permit.slot_index == 0 + + +def test_capacity_two_two_parents_run_cross_process_nested_acquires( + monkeypatch, + tmp_path, +): + _isolated_journal(monkeypatch, tmp_path) + + env = {**os.environ, "SOLSTONE_JOURNAL": str(tmp_path)} + script = """ +from solstone.think.providers.local_admission import acquire_local_slot +with acquire_local_slot(2, 1.0): + pass +""" + barrier = threading.Barrier(2) + results: list[int] = [] + errors: list[BaseException] = [] + lock = threading.Lock() + + def parent() -> None: + lease: LocalSlotLease | None = None + try: + permit = acquire_local_slot(2, 1.0) + lease = LocalSlotLease( + capacity=2, + deadline=time.monotonic() + 3.0, + permit=permit, + ) + barrier.wait(timeout=1.0) + lease.yield_slot() + completed = subprocess.run( + [sys.executable, "-c", script], + env=env, + text=True, + capture_output=True, + timeout=2.0, + check=False, + ) + with lock: + results.append(completed.returncode) + if completed.returncode != 0: + raise AssertionError(completed.stderr) + lease.reacquire() + except BaseException as exc: + with lock: + errors.append(exc) + finally: + if lease is not None: + lease.close() + + threads = [threading.Thread(target=parent) for _ in range(2)] + for thread in threads: + thread.start() + for thread in threads: + thread.join(timeout=4.0) + + assert all(not thread.is_alive() for thread in threads) + assert errors == [] + assert results == [0, 0] + first = acquire_local_slot(2, 0.1) + try: + second = acquire_local_slot(2, 0.1) + try: + assert {first.slot_index, second.slot_index} == {0, 1} + finally: + second.release() + finally: + first.release() + + def test_stale_ticket_from_exited_owner_is_pruned(monkeypatch, tmp_path): _isolated_journal(monkeypatch, tmp_path) root = tmp_path / "health" / "local-inference-admission" diff --git a/tests/test_openhands_provider.py b/tests/test_openhands_provider.py index d350a46e3..0d4f74a33 100644 --- a/tests/test_openhands_provider.py +++ b/tests/test_openhands_provider.py @@ -18,6 +18,7 @@ from solstone.think.cogitate_policy import ( MAX_TURNS_HEADROOM, ) from solstone.think.providers import openhands +from solstone.think.providers.local_admission import LocalAdmissionTimeout from solstone.think.providers.shared import USAGE_KEYS, JSONEventCallback from solstone.think.talent import get_talent, get_talent_configs from tests.openhands_fakes import _REGISTERED_TOOLS, install_fake_openhands @@ -732,6 +733,49 @@ def test_run_cogitate_keeps_finish_branch_without_output_path( assert [event for event in events if event["event"] == "error"] == [] +def test_run_cogitate_threads_slot_lease_to_sol_executor_and_binds_conversation( + fake_openhands, + monkeypatch, + tmp_path, +): + config = _run_config(monkeypatch, tmp_path) + events: list[dict] = [] + slot_lease = object() + + asyncio.run(openhands.run_cogitate(config, events.append, slot_lease=slot_lease)) + + conversation = fake_openhands.Conversation.instances[0] + executor = _REGISTERED_TOOLS["sol"].executor + assert executor.slot_lease is slot_lease + assert executor._conversation is conversation + + +def test_run_cogitate_terminal_sol_reacquire_error_preempts_stuck_classification( + fake_openhands, + monkeypatch, + tmp_path, +): + terminal = LocalAdmissionTimeout("busy") + + async def mark_terminal_and_stuck(conversation): + _REGISTERED_TOOLS["sol"].executor._store_terminal_error(terminal) + conversation.state.execution_status = "stuck" + + fake_openhands.Conversation.arun_impl = mark_terminal_and_stuck + config = _run_config(monkeypatch, tmp_path) + events: list[dict] = [] + + with pytest.raises(LocalAdmissionTimeout) as exc: + asyncio.run(openhands.run_cogitate(config, events.append)) + + assert exc.value is terminal + conversation = fake_openhands.Conversation.instances[0] + assert conversation.closed is True + assert [ + event for event in events if event.get("reason_code") == "agent_stuck" + ] == [] + + def test_run_cogitate_uses_emit_final_branch_for_daily_no_output( fake_openhands, monkeypatch, diff --git a/tests/test_openhands_sdk_shape.py b/tests/test_openhands_sdk_shape.py index 9f2ed3498..322a45b6e 100644 --- a/tests/test_openhands_sdk_shape.py +++ b/tests/test_openhands_sdk_shape.py @@ -3,10 +3,65 @@ from __future__ import annotations +import asyncio import inspect +import os +import threading + +from pydantic import Field from tests._logging_isolation import preserve_global_logging +os.environ.setdefault("OPENHANDS_SUPPRESS_BANNER", "1") + +from openhands.sdk.tool import ( # noqa: E402 + ToolAnnotations, + ToolDefinition, + ToolExecutor, +) +from openhands.sdk.tool.schema import Action, Observation # noqa: E402 + + +class _ShapeAction(Action): + value: str = Field(default="") + + +class _ShapeObservation(Observation): + pass + + +class _ShapeTool(ToolDefinition[_ShapeAction, _ShapeObservation]): + name = "sdk_shape" + + @classmethod + def create(cls, *args, **kwargs): + del args, kwargs + return [] + + +class _InterruptingExecutor(ToolExecutor): + def __init__(self) -> None: + self.call_count = 0 + self.worker_thread_id: int | None = None + + def __call__(self, action, conversation=None): + del action + self.call_count += 1 + self.worker_thread_id = threading.get_ident() + assert conversation is not None + conversation.interrupt() + return _ShapeObservation.from_text("interrupted", is_error=True) + + +class _RaisingExecutor(ToolExecutor): + def __init__(self) -> None: + self.call_count = 0 + + def __call__(self, action, conversation=None): + del action, conversation + self.call_count += 1 + raise RuntimeError("executor boom") + def test_local_conversation_methods_match_provider_await_sites(monkeypatch): monkeypatch.setenv("OPENHANDS_SUPPRESS_BANNER", "1") @@ -17,3 +72,138 @@ def test_local_conversation_methods_match_provider_await_sites(monkeypatch): assert inspect.iscoroutinefunction(LocalConversation.arun) is True assert inspect.iscoroutinefunction(LocalConversation.send_message) is False + + +def _shape_tool(executor): + return _ShapeTool( + description="SDK shape test tool", + action_type=_ShapeAction, + observation_type=_ShapeObservation, + executor=executor, + annotations=ToolAnnotations(title="sdk_shape"), + ) + + +def _tool_call_message(call_id: str = "call-1"): + from openhands.sdk.llm import Message, MessageToolCall + + return Message( + role="assistant", + content=[], + tool_calls=[ + MessageToolCall( + id=call_id, + name="sdk_shape", + arguments='{"value":"ok"}', + origin="completion", + ) + ], + ) + + +def _assistant_message(text: str): + from openhands.sdk.llm import Message, TextContent + + return Message(role="assistant", content=[TextContent(text=text)]) + + +def _conversation(llm, tool, tmp_path, callbacks=None): + from openhands.sdk import Agent, Conversation + from openhands.sdk.tool.registry import register_tool + from openhands.sdk.tool.spec import Tool + + register_tool("sdk_shape", tool) + agent = Agent( + llm=llm, + tools=[Tool(name="sdk_shape")], + include_default_tools=[], + system_prompt="system", + ) + return Conversation( + agent=agent, + workspace=str(tmp_path), + persistence_dir=str(tmp_path / "history"), + callbacks=callbacks or [], + visualizer=None, + stuck_detection=False, + ) + + +def test_interrupt_from_tool_worker_stops_before_next_llm_completion(tmp_path): + from openhands.sdk.testing import TestLLM + + executor = _InterruptingExecutor() + llm = TestLLM.from_messages( + [_tool_call_message(), _assistant_message("should not be consumed")] + ) + conversation = _conversation(llm, _shape_tool(executor), tmp_path) + conversation.send_message("go") + + asyncio.run(conversation.arun()) + + assert executor.call_count == 1 + assert executor.worker_thread_id != threading.get_ident() + assert llm.call_count == 1 + conversation.close() + + +def test_executor_exception_becomes_agent_error_and_loop_continues(tmp_path): + from openhands.sdk.event.llm_convertible import AgentErrorEvent + from openhands.sdk.testing import TestLLM + + executor = _RaisingExecutor() + events = [] + llm = TestLLM.from_messages([_tool_call_message(), _assistant_message("done")]) + conversation = _conversation( + llm, + _shape_tool(executor), + tmp_path, + callbacks=[events.append], + ) + conversation.send_message("go") + + asyncio.run(conversation.arun()) + + assert executor.call_count == 1 + assert llm.call_count == 2 + assert any(isinstance(event, AgentErrorEvent) for event in events) + conversation.close() + + +def test_agent_default_concurrency_and_undeclared_tool_mutex(): + from openhands.sdk import Agent + from openhands.sdk.agent.parallel_executor import ParallelToolExecutor + from openhands.sdk.event.llm_convertible import ActionEvent + from openhands.sdk.llm import MessageToolCall + from openhands.sdk.testing import TestLLM + + tool = _shape_tool(_InterruptingExecutor()) + agent = Agent( + llm=TestLLM.from_messages([]), + tools=[], + include_default_tools=[], + system_prompt="system", + ) + action_event = ActionEvent( + thought=[], + tool_name="sdk_shape", + tool_call_id="call-1", + tool_call=MessageToolCall( + id="call-1", + name="sdk_shape", + arguments='{"value":"ok"}', + origin="completion", + ), + llm_response_id="response-1", + action=_ShapeAction(value="ok"), + ) + + assert agent.tool_concurrency_limit == 1 + # Private helper pin: undeclared resources resolve to the same tool-wide + # mutex key that serializes same-tool calls when concurrency is raised. + resources = ParallelToolExecutor._extract_declared_resources(action_event, tool) + assert resources is not None + assert resources.declared is False + assert ParallelToolExecutor._resolve_lock_keys(resources, tool) == [ + "tool:sdk_shape" + ] diff --git a/tests/test_openhands_sol_tool.py b/tests/test_openhands_sol_tool.py index 40ba1265e..ad4e3bd7f 100644 --- a/tests/test_openhands_sol_tool.py +++ b/tests/test_openhands_sol_tool.py @@ -3,10 +3,13 @@ from __future__ import annotations +import time + import pytest from solstone.think.cogitate_policy import CogitatePolicy -from solstone.think.providers import openhands +from solstone.think.providers import local_admission, openhands +from solstone.think.providers.local_admission import LocalSlotLease from solstone.think.providers.shared import JSONEventCallback from tests.openhands_fakes import install_fake_openhands @@ -26,18 +29,34 @@ def _sol_tool_and_executor( tmp_path, events: list[dict], read_call_budget: int = 200, + slot_lease=None, ): policy = CogitatePolicy(allowed_roots=[tmp_path], access_tier="normal") tools, executor = openhands._build_sol_tools( policy=policy, callback=JSONEventCallback(events.append), read_call_budget=read_call_budget, + slot_lease=slot_lease, ) assert len(tools) == 1 assert tools[0].name == "sol" return tools[0], executor +def _isolated_lease(monkeypatch, tmp_path, *, timeout_s: float = 1.0): + monkeypatch.setattr( + local_admission, + "_admission_dir", + lambda: tmp_path / "local-inference-admission", + ) + permit = local_admission.acquire_local_slot(1, 0.1) + return LocalSlotLease( + capacity=1, + deadline=time.monotonic() + timeout_s, + permit=permit, + ) + + def test_read_only_allowed_sol_call_returns_non_error_observation( fake_openhands, fixed_time, @@ -137,3 +156,142 @@ def test_read_call_budget_overflow_emits_once_and_denies_recoverably( "ts": 123456, } ] + + +def test_unexpected_command_exception_reacquires_before_escaping( + fake_openhands, + fixed_time, + tmp_path, + monkeypatch, +): + events: list[dict] = [] + lease = _isolated_lease(monkeypatch, tmp_path) + tool, _executor = _sol_tool_and_executor( + tmp_path=tmp_path, + events=events, + slot_lease=lease, + ) + + def fail_command(_argv: list[str]): + raise RuntimeError("boom") + + monkeypatch.setattr(openhands, "_run_command", fail_command) + + with pytest.raises(RuntimeError, match="boom"): + tool(tool.action_from_arguments({"command": "sol call journal search x"})) + + with pytest.raises(local_admission.LocalAdmissionTimeout): + local_admission.acquire_local_slot(1, 0.03) + lease.close() + with local_admission.acquire_local_slot(1, 0.1) as permit: + assert permit.slot_index == 0 + + +def test_reacquire_timeout_sets_terminal_marker_and_interrupts_conversation( + fake_openhands, + fixed_time, + tmp_path, + monkeypatch, +): + events: list[dict] = [] + lease = _isolated_lease(monkeypatch, tmp_path, timeout_s=0.03) + tool, executor = _sol_tool_and_executor( + tmp_path=tmp_path, + events=events, + slot_lease=lease, + ) + holder = None + + def run_and_hold(_argv: list[str]): + nonlocal holder + holder = local_admission.acquire_local_slot(1, 0.1) + return {"text": "ran", "is_error": False} + + conversation = fake_openhands.Conversation() + monkeypatch.setattr(openhands, "_run_command", run_and_hold) + + observation = tool( + tool.action_from_arguments({"command": "sol call journal search x"}), + conversation, + ) + + try: + terminal = executor.take_terminal_error() + assert observation.is_error is True + assert isinstance(terminal, local_admission.LocalAdmissionTimeout) + assert terminal.reason_code == "local_queue_timeout" + assert conversation.interrupted is True + finally: + if holder is not None: + holder.release() + lease.close() + + with local_admission.acquire_local_slot(1, 0.1) as permit: + assert permit.slot_index == 0 + + +def test_reacquire_cancel_is_recoverable_observation_not_terminal( + fake_openhands, + fixed_time, + tmp_path, + monkeypatch, +): + events: list[dict] = [] + lease = _isolated_lease(monkeypatch, tmp_path) + tool, executor = _sol_tool_and_executor( + tmp_path=tmp_path, + events=events, + slot_lease=lease, + ) + + def run_and_cancel(_argv: list[str]): + lease.cancel_pending_reacquire() + return {"text": "ran before cancel", "is_error": False} + + monkeypatch.setattr(openhands, "_run_command", run_and_cancel) + + observation = tool( + tool.action_from_arguments({"command": "sol call journal search x"}) + ) + + assert observation.text == "ran before cancel" + assert observation.is_error is False + assert executor.take_terminal_error() is None + lease.close() + with local_admission.acquire_local_slot(1, 0.1) as permit: + assert permit.slot_index == 0 + + +@pytest.mark.parametrize( + "command_result", + [ + {"text": "exit_code: 2", "is_error": True}, + {"text": "timeout: command exceeded 30s", "is_error": True}, + ], +) +def test_command_error_results_reacquire_before_returning( + fake_openhands, + fixed_time, + tmp_path, + monkeypatch, + command_result, +): + events: list[dict] = [] + lease = _isolated_lease(monkeypatch, tmp_path) + tool, executor = _sol_tool_and_executor( + tmp_path=tmp_path, + events=events, + slot_lease=lease, + ) + monkeypatch.setattr(openhands, "_run_command", lambda _argv: command_result) + + observation = tool( + tool.action_from_arguments({"command": "sol call journal search x"}) + ) + + assert observation.text == command_result["text"] + assert observation.is_error is True + assert executor.take_terminal_error() is None + with pytest.raises(local_admission.LocalAdmissionTimeout): + local_admission.acquire_local_slot(1, 0.03) + lease.close()