diff --git a/sdk/agentserver/azure-ai-agentserver-core/CHANGELOG.md b/sdk/agentserver/azure-ai-agentserver-core/CHANGELOG.md index 177ca386dd6f..f09e9bde7d48 100644 --- a/sdk/agentserver/azure-ai-agentserver-core/CHANGELOG.md +++ b/sdk/agentserver/azure-ai-agentserver-core/CHANGELOG.md @@ -1,5 +1,13 @@ # Release History +## 2.2.1 (Unreleased) + +### Bugs Fixed + +- Failed steering queue appends no longer retain an unreturned acknowledgment + future. Per-slot acknowledgment IDs prevent a committed append with a lost + response from routing the next accepted input's result incorrectly. + ## 2.2.0 (2026-09-23) ### Other Changes diff --git a/sdk/agentserver/azure-ai-agentserver-core/azure/ai/agentserver/core/_version.py b/sdk/agentserver/azure-ai-agentserver-core/azure/ai/agentserver/core/_version.py index 4df7819b5eb2..2d76447c9689 100644 --- a/sdk/agentserver/azure-ai-agentserver-core/azure/ai/agentserver/core/_version.py +++ b/sdk/agentserver/azure-ai-agentserver-core/azure/ai/agentserver/core/_version.py @@ -2,4 +2,4 @@ # Copyright (c) Microsoft Corporation. All rights reserved. # --------------------------------------------------------- -VERSION = "2.2.0" +VERSION = "2.2.1" diff --git a/sdk/agentserver/azure-ai-agentserver-core/azure/ai/agentserver/core/tasks/_attachments.py b/sdk/agentserver/azure-ai-agentserver-core/azure/ai/agentserver/core/tasks/_attachments.py index 8aef73bc7647..c268721e26a1 100644 --- a/sdk/agentserver/azure-ai-agentserver-core/azure/ai/agentserver/core/tasks/_attachments.py +++ b/sdk/agentserver/azure-ai-agentserver-core/azure/ai/agentserver/core/tasks/_attachments.py @@ -99,6 +99,29 @@ _STEERING_QUEUE_CAP = 9 +def _steering_pending_ack_ids(steering: dict[str, Any], pending_count: int) -> list[str | None]: + """Align internal acknowledgment IDs with queued inputs from older records. + + :param steering: Persisted steering state. + :type steering: dict[str, Any] + :param pending_count: Number of queued inputs. + :type pending_count: int + :return: IDs aligned with pending inputs, including untagged legacy slots. + :rtype: list[str | None] + :raises ValueError: If stored IDs cannot align with the pending queue. + """ + if pending_count == 0: + return [] + ids = steering.get("pending_ack_ids") + if ids is None: + return [None] * pending_count + if not isinstance(ids, list) or len(ids) > pending_count or any( + value is not None and not isinstance(value, str) for value in ids + ): + raise ValueError("Invalid steering pending_ack_ids for pending queue") + return ids + [None] * (pending_count - len(ids)) + + # --------------------------------------------------------------------------- # # Hash helper # --------------------------------------------------------------------------- # diff --git a/sdk/agentserver/azure-ai-agentserver-core/azure/ai/agentserver/core/tasks/_decorator.py b/sdk/agentserver/azure-ai-agentserver-core/azure/ai/agentserver/core/tasks/_decorator.py index a88112320a2d..6823085c6ae9 100644 --- a/sdk/agentserver/azure-ai-agentserver-core/azure/ai/agentserver/core/tasks/_decorator.py +++ b/sdk/agentserver/azure-ai-agentserver-core/azure/ai/agentserver/core/tasks/_decorator.py @@ -798,6 +798,7 @@ async def _append_steering_input( # pylint: disable=protected-access,too-many-l task_id: str, input_val: Any, existing: Any, + ack_id: str, input_id: str | None = None, if_last_input_id: str | None = None, ) -> None: @@ -812,6 +813,8 @@ async def _append_steering_input( # pylint: disable=protected-access,too-many-l :keyword existing: The previously-fetched task record (used for the first etag attempt; later attempts re-fetch internally). :paramtype existing: Any + :keyword ack_id: Internal identifier binding this queue slot to its acknowledgment. + :paramtype ack_id: str :keyword input_id: When set, the new input's identity. Used to advance ``payload["last_input_id"]`` atomically with the queue append. @@ -867,8 +870,10 @@ async def _append_steering_input( # pylint: disable=protected-access,too-many-l _STEERING_INPUT_KEY_PREFIX, _STEERING_THRESHOLD_BYTES, _resolve_input_storage, + _steering_pending_ack_ids, ) + pending_ack_ids = _steering_pending_ack_ids(steering, len(pending)) next_seq = int(steering.get("next_input_seq", 0)) steering_key = f"{_STEERING_INPUT_KEY_PREFIX}{next_seq}" store_mode, queue_entry = _resolve_input_storage( @@ -883,7 +888,9 @@ async def _append_steering_input( # pylint: disable=protected-access,too-many-l steering["next_input_seq"] = next_seq + 1 pending.append(queue_entry) + pending_ack_ids.append(ack_id) steering["pending_inputs"] = pending + steering["pending_ack_ids"] = pending_ack_ids steering["cancel_requested"] = True # SOT: the # internal _steering["generation"] payload field is removed @@ -949,8 +956,8 @@ def _create_steering_ack_run( manager: Any, task_id: str, future: Any, + ack_id: str, input_id: str | None = None, - input_val: Any = None, ) -> TaskRun[Output]: """Create a TaskRun for a queued steering input. @@ -960,11 +967,10 @@ def _create_steering_ack_run( :type task_id: str :param future: Future that will resolve with the next-turn outcome. :type future: Any + :param ack_id: Internal identifier of the queued steering slot. + :type ack_id: str :param input_id: The input_id stamped on the queued input (if any). :type input_id: str | None - :param input_val: The raw queued input value (used to identify the - slot when ``cancel()`` is invoked on the returned handle). - :type input_val: Any :return: A :class:`TaskRun` whose result resolves with the queued turn. :rtype: TaskRun[Output] """ @@ -973,8 +979,7 @@ async def _queued_cancel_cb() -> None: await manager._cancel_queued_steering_input( # pylint: disable=protected-access task_id=task_id, future=future, - input_id=input_id, - input_val=input_val, + ack_id=ack_id, ) return TaskRun( @@ -1306,21 +1311,31 @@ async def _lifecycle_start_inner( # pylint: disable=too-many-locals,too-many-st if self._opts.steerable: # Steering path: append input to queue, signal cancel, return ack # pylint: disable=protected-access - ack_future = manager._register_steering_future(task_id) - await self._append_steering_input( - manager, - task_id=task_id, - input_val=input, - existing=existing, - input_id=input_id, - if_last_input_id=if_last_input_id, - ) + # Keep drain from binding a future before its append is accepted. + async with manager._get_task_write_lock(task_id): + ack_id = _generate_input_id() + ack_future = manager._register_steering_future(task_id, ack_id) + appended = False + try: + await self._append_steering_input( + manager, + task_id=task_id, + input_val=input, + existing=existing, + ack_id=ack_id, + input_id=input_id, + if_last_input_id=if_last_input_id, + ) + appended = True + finally: + if not appended: + manager._unregister_steering_future(task_id, ack_id, ack_future) # Set cancel on in-memory context if task runs in this process active = manager._active_tasks.get(task_id) # pylint: enable=protected-access if active: active.context.cancel.set() - return self._create_steering_ack_run(manager, task_id, ack_future, input_id=input_id, input_val=input) + return self._create_steering_ack_run(manager, task_id, ack_future, ack_id, input_id=input_id) raise TaskConflictError(task_id, "in_progress") # completed (or any other terminal status) @@ -1748,8 +1763,8 @@ async def delete(self, task_id: str) -> None: exec_task.cancel() # 2. Resolve all queued steerer futures with TaskCancelled. - pending = getattr(mgr, "_pending_steering_futures", {}).pop(task_id, []) - for queued_fut in pending: + pending = getattr(mgr, "_pending_steering_futures", {}).pop(task_id, {}) + for queued_fut in pending.values(): if not queued_fut.done(): queued_fut.set_exception(TaskCancelled()) diff --git a/sdk/agentserver/azure-ai-agentserver-core/azure/ai/agentserver/core/tasks/_manager.py b/sdk/agentserver/azure-ai-agentserver-core/azure/ai/agentserver/core/tasks/_manager.py index 0600ec43b670..fceb4662e151 100644 --- a/sdk/agentserver/azure-ai-agentserver-core/azure/ai/agentserver/core/tasks/_manager.py +++ b/sdk/agentserver/azure-ai-agentserver-core/azure/ai/agentserver/core/tasks/_manager.py @@ -33,6 +33,7 @@ _read_input_value, _ref_key, _resolve_input_storage, + _steering_pending_ack_ids, ) from ._decorator import TaskOptions, _deserialize_input, _resolve_effective_timeout, _serialize_input from ._exceptions import ( @@ -186,7 +187,7 @@ def _parse_turn_started_at(value: Any) -> float | None: def _resolve_queued_steerers_on_terminal( - pending_steering_futures: dict[str, list["asyncio.Future[Any]"]], + pending_steering_futures: dict[str, dict[str, "asyncio.Future[Any]"]], task_id: str, *, current_status: str, @@ -203,9 +204,9 @@ def _resolve_queued_steerers_on_terminal( Pops every queued steerer future for ``task_id`` and resolves each with ``TaskConflictError(current_status=current_status)``. - :param pending_steering_futures: Per-task list of pending steerer + :param pending_steering_futures: Per-task map of pending steerer futures (mutated in-place — emptied for the given ``task_id``). - :type pending_steering_futures: dict[str, list[asyncio.Future[Any]]] + :type pending_steering_futures: dict[str, dict[str, asyncio.Future[Any]]] :param task_id: The task whose queued steerers should be resolved. :type task_id: str :keyword current_status: Status string to carry on @@ -214,8 +215,8 @@ def _resolve_queued_steerers_on_terminal( """ # TaskConflictError is already imported at module top-level (line 24). - queued = pending_steering_futures.pop(task_id, []) - for fut in queued: + queued = pending_steering_futures.pop(task_id, {}) + for fut in queued.values(): if not fut.done(): fut.set_exception(TaskConflictError(task_id, current_status)) @@ -431,7 +432,7 @@ def __init__( self._shutdown_event = shutdown_event or asyncio.Event() self._shutdown_grace_seconds = shutdown_grace_seconds self._active_generation_future: dict[str, asyncio.Future[Any]] = {} - self._pending_steering_futures: dict[str, list[asyncio.Future[Any]]] = {} + self._pending_steering_futures: dict[str, dict[str, asyncio.Future[Any]]] = {} # Layer 2: periodic recovery scan task. Created # at startup() time; cancelled at shutdown(). self._periodic_recovery_task: asyncio.Task[None] | None = None @@ -615,7 +616,7 @@ async def list_tasks( raise RuntimeError("Task list did not converge after retryable conflict") from exc raise translated from exc - def _register_steering_future(self, task_id: str) -> asyncio.Future[Any]: + def _register_steering_future(self, task_id: str, ack_id: str) -> asyncio.Future[Any]: """Create and register a future for a queued steering input. Must be called BEFORE ``_append_steering_input()`` to avoid a race @@ -623,23 +624,54 @@ def _register_steering_future(self, task_id: str) -> asyncio.Future[Any]: :param task_id: The task identifier. :type task_id: str + :param ack_id: Unique internal identifier of the queue slot. + :type ack_id: str :return: The registered future. :rtype: asyncio.Future[Any] """ + pending = self._pending_steering_futures.setdefault(task_id, {}) + if ack_id in pending: + raise ValueError(f"Steering acknowledgment ID already registered for task {task_id!r}") loop = asyncio.get_running_loop() future: asyncio.Future[Any] = loop.create_future() - if task_id not in self._pending_steering_futures: - self._pending_steering_futures[task_id] = [] - self._pending_steering_futures[task_id].append(future) + pending[ack_id] = future return future - async def _cancel_queued_steering_input( # pylint: disable=unused-argument + def _remove_steering_future(self, task_id: str, ack_id: str, future: asyncio.Future[Any]) -> None: + """Remove a registered future only if this request owns its ID. + + :param task_id: The task identifier. + :type task_id: str + :param ack_id: Internal identifier of the queued slot. + :type ack_id: str + :param future: The request's registered future. + :type future: asyncio.Future[Any] + """ + pending = self._pending_steering_futures.get(task_id) + if pending is not None and pending.get(ack_id) is future: + del pending[ack_id] + if not pending: + self._pending_steering_futures.pop(task_id, None) + + def _unregister_steering_future(self, task_id: str, ack_id: str, future: asyncio.Future[Any]) -> None: + """Discard only the future owned by a failed steering append. + + :param task_id: The task identifier. + :type task_id: str + :param ack_id: Internal identifier of the attempted queue slot. + :type ack_id: str + :param future: The failed append's registered future. + :type future: asyncio.Future[Any] + """ + self._remove_steering_future(task_id, ack_id, future) + future.cancel() + + async def _cancel_queued_steering_input( self, *, task_id: str, future: asyncio.Future[Any], - input_id: str | None, - input_val: Any, + ack_id: str, ) -> None: """Remove a queued steering input from the chain's pending queue. @@ -652,11 +684,8 @@ async def _cancel_queued_steering_input( # pylint: disable=unused-argument :keyword task_id: The chain task identifier. :keyword future: The queued steerer's result_future. - :keyword input_id: The input_id of the queued slot (used for the - future-list cleanup; the queue entry itself is identified by - ``input_val``). - :keyword input_val: The raw queued value used to identify which - ``pending_inputs`` entry to remove. + :keyword ack_id: Internal ID identifying the exact queued slot, even + when multiple inputs have identical values or public input IDs. """ from ._attachments import _is_ref, _ref_key # pylint: disable=import-outside-toplevel from ._exceptions import TaskCancelled # pylint: disable=import-outside-toplevel @@ -668,37 +697,28 @@ async def _cancel_queued_steering_input( # pylint: disable=unused-argument task_info = None if task_info is None or not task_info.payload: # Chain already gone — just resolve the future. + self._remove_steering_future(task_id, ack_id, future) if not future.done(): future.set_exception(TaskCancelled()) return steering = dict(task_info.payload.get("steering") or {}) pending = list(steering.get("pending_inputs") or []) + pending_ack_ids = _steering_pending_ack_ids(steering, len(pending)) attachments_patch: dict[str, Any] = {} - # Drop the first queue entry whose raw value matches ``input_val``. - removed = False - new_pending: list[Any] = [] - for entry in pending: - if not removed: - raw = entry - if _is_ref(entry): - # For ref-shaped entries, resolve via attachment to - # compare against input_val. If the attachment is - # missing, fall back to ref identity (unlikely). - key = _ref_key(entry) - raw = (task_info.attachments or {}).get(key, entry) - if raw == input_val: - removed = True - if _is_ref(entry): - attachments_patch[_ref_key(entry)] = None - continue - new_pending.append(entry) - if not removed: + if ack_id not in pending_ack_ids: # Queue entry already drained or never landed; just resolve. + self._remove_steering_future(task_id, ack_id, future) if not future.done(): future.set_exception(TaskCancelled()) return - steering["pending_inputs"] = new_pending - steering["cancel_requested"] = len(new_pending) > 0 + index = pending_ack_ids.index(ack_id) + entry = pending.pop(index) + pending_ack_ids.pop(index) + if _is_ref(entry): + attachments_patch[_ref_key(entry)] = None + steering["pending_inputs"] = pending + steering["pending_ack_ids"] = pending_ack_ids + steering["cancel_requested"] = len(pending) > 0 payload_patch: dict[str, Any] = {"steering": steering} try: # Spec 031 / FR-005a+b: the outer lock is already held, so use @@ -721,10 +741,8 @@ async def _cancel_queued_steering_input( # pylint: disable=unused-argument task_id, exc_info=True, ) - # Remove the future from the registered pending list and resolve it. - pending_list = self._pending_steering_futures.get(task_id) or [] - if future in pending_list: - pending_list.remove(future) + # Remove this request's registration and resolve its returned handle. + self._remove_steering_future(task_id, ack_id, future) if not future.done(): future.set_exception(TaskCancelled()) @@ -2351,6 +2369,8 @@ async def _try_drain_steering( # pylint: disable=too-many-branches,too-many-sta if not pending: return None + pending_ack_ids = _steering_pending_ack_ids(steering, len(pending)) + next_ack_id = pending_ack_ids.pop(0) # Pop the next input from the queue.: the entry may be # either a raw inline value (≤ 20 KiB at append) or a ref slot # pointing into ``task_info.attachments``. Resolve uniformly via @@ -2367,6 +2387,7 @@ async def _try_drain_steering( # pylint: disable=too-many-branches,too-many-sta # state need to survive a crash mid-drain.) steering["active_input"] = next_input_raw steering["pending_inputs"] = pending + steering["pending_ack_ids"] = pending_ack_ids # SOT: internal # _steering["generation"] writes removed. The drain transition # IS the generation advance — no separate counter needed. @@ -2442,11 +2463,15 @@ async def _try_drain_steering( # pylint: disable=too-many-branches,too-many-sta _conflict_attempt=_conflict_attempt + 1, ) - # Pop and bind the next pending steering future (if any) + # Bind only the future registered for this durable queue slot. new_future: asyncio.Future[Any] | None = None - steering_futures = self._pending_steering_futures.get(task_id, []) - if steering_futures: - new_future = steering_futures.pop(0) + steering_futures = self._pending_steering_futures.get(task_id) + if next_ack_id is not None and steering_futures is not None: + new_future = steering_futures.pop(next_ack_id, None) + if not steering_futures: + self._pending_steering_futures.pop(task_id, None) + if next_ack_id is not None and new_future is None: + logger.debug("Steering input %s for task %s has no local acknowledgment future", next_ack_id, task_id) # Resolve the queued steerer's future binding for the new turn. # / (Subscriber): the OLD result_future is NOT diff --git a/sdk/agentserver/azure-ai-agentserver-core/docs/task-and-streaming-spec.md b/sdk/agentserver/azure-ai-agentserver-core/docs/task-and-streaming-spec.md index 306a6bf76d07..5b1030826b77 100644 --- a/sdk/agentserver/azure-ai-agentserver-core/docs/task-and-streaming-spec.md +++ b/sdk/agentserver/azure-ai-agentserver-core/docs/task-and-streaming-spec.md @@ -1042,11 +1042,17 @@ chain record itself does not carry the per-turn diagnostic. | Sub-key | Type | Meaning | |---|---|---| | `pending_inputs` | array of input values OR refs (§23) | FIFO of queued steering inputs. | +| `pending_ack_ids` | array of opaque strings or nulls | Internal acknowledgment IDs aligned with `pending_inputs`. Each new append adds a unique ID in the same PATCH; drain and queued cancel remove the corresponding ID. Null marks an untagged slot from a pre-upgrade task record. These are independent of the caller's `input_id`. | | `next_input_seq` | integer | Monotonic counter for promoted-attachment key allocation (NEVER reused). | | `cancel_requested` | boolean | Resilient cancel signal; set on steering append; cleared after drain when pending is empty. | | `drain_in_progress` | boolean | True between the start of a drain PATCH and the next turn-start; protects against partial drain on crash. | | `active_input` | any JSON value OR ref | The single input being drained (mirror copy used by the race-recovery contract). Cleared at suspend / terminal. | +Older task records may omit `pending_ack_ids`; those queued slots are treated +as untagged. Workers handling the same chain concurrently must all understand +this field, since older workers do not keep it aligned when draining or +cancelling inputs. + Implementers in other languages MUST use these exact key names. A process built in language X must be able to recover a task created by language Y. @@ -2996,6 +3002,8 @@ executes this PATCH as a single round-trip: pending.append(input) # raw inline attachments_patch = None 7. steering['pending_inputs'] = pending + steering['pending_ack_ids'] = existing IDs (null-padded for untagged + legacy slots) + [new unique ack ID] steering['cancel_requested'] = True 8. payload_patch = {'steering': steering} if input_id provided: payload_patch['last_input_id'] = input_id @@ -3027,6 +3035,8 @@ Phase 1 — "Drain start" PATCH (atomic across payload + attachments): 4. If pending is empty: return None (no drain happens; caller proceeds to suspend/complete normally). 5. next_entry = pending.pop(0) + next_ack_id = steering['pending_ack_ids'].pop(0) if present, + or null for an untagged legacy entry 6. attachments_patch = {} 7. If next_entry is a ref (§23.3): attachments_patch[ref_key(next_entry)] = None # delete attachment @@ -3035,6 +3045,7 @@ Phase 1 — "Drain start" PATCH (atomic across payload + attachments): active_input_value = next_entry 8. steering['active_input'] = active_input_value 9. steering['pending_inputs'] = pending + steering['pending_ack_ids'] = remaining IDs 10. steering['drain_in_progress'] = True 11. steering['cancel_requested'] = len(pending) > 0 # more pending => keep advisory 12. payload['steering'] = steering diff --git a/sdk/agentserver/azure-ai-agentserver-core/tests/tasks/test_steering.py b/sdk/agentserver/azure-ai-agentserver-core/tests/tasks/test_steering.py index dde0b9884c25..110220c2fb66 100644 --- a/sdk/agentserver/azure-ai-agentserver-core/tests/tasks/test_steering.py +++ b/sdk/agentserver/azure-ai-agentserver-core/tests/tasks/test_steering.py @@ -15,6 +15,7 @@ task, EntryMode, SteeringQueueFull, + TaskCancelled, TaskConflictError, multi_turn_task, ) @@ -123,40 +124,265 @@ async def regular(ctx: TaskContext[dict]) -> dict: @pytest.mark.asyncio async def test_steering_queue_full(self, tmp_path): - """start() raises SteeringQueueFull when queue is at capacity. - - : the per-task ``max_pending`` knob was - demoted; the framework-wide default - ``_DEFAULT_MAX_PENDING_STEERING`` (10) applies. This test fills the - queue at that default to verify the exception still surfaces. - """ + """Rejected steers do not leak futures or steal accepted acknowledgments.""" from azure.ai.agentserver.core.tasks._decorator import _DEFAULT_MAX_PENDING_STEERING manager, mgr_mod = await self._setup_manager(tmp_path) + gate = asyncio.Event() try: - gate = asyncio.Event() - @multi_turn_task(name="chat", steerable=True) async def chat(ctx: TaskContext[dict]) -> dict: await gate.wait() - return {"msg": "done"} + return {"msg": ctx.input["msg"]} run1 = await chat.start(task_id="t1", input={"msg": "A"}) - # Fill the queue to the framework default + accepted = [] for i in range(_DEFAULT_MAX_PENDING_STEERING): - await chat.start(task_id="t1", input={"msg": f"fill-{i}"}) + accepted.append(await chat.start(task_id="t1", input={"msg": f"fill-{i}"})) + registered = dict(manager._pending_steering_futures["t1"]) + assert len(registered) == _DEFAULT_MAX_PENDING_STEERING + + for i in range(3): + with pytest.raises(SteeringQueueFull): + await chat.start(task_id="t1", input={"msg": f"overflow-{i}"}) + assert manager._pending_steering_futures["t1"] == registered + + info = await manager.provider.get("t1") + assert info is not None + assert [entry["msg"] for entry in info.payload["steering"]["pending_inputs"]] == [ + f"fill-{i}" for i in range(_DEFAULT_MAX_PENDING_STEERING) + ] + + gate.set() + assert await asyncio.wait_for(run1.result(), timeout=5.0) == {"msg": "A"} + for i, run in enumerate(accepted): + assert await asyncio.wait_for(run.result(), timeout=5.0) == {"msg": f"fill-{i}"} + assert not manager._pending_steering_futures.get("t1") + finally: + gate.set() + await self._teardown_manager(manager, mgr_mod) - # Queue is full — should raise - with pytest.raises(SteeringQueueFull): - await chat.start(task_id="t1", input={"msg": "overflow"}) + @pytest.mark.asyncio + async def test_steering_storage_failure_unregisters_only_failed_future(self, tmp_path, monkeypatch): + manager, mgr_mod = await self._setup_manager(tmp_path) + gate = asyncio.Event() + try: - #: SteeringQueueFull is bare exception (no max_pending) + @multi_turn_task(name="chat", steerable=True) + async def chat(ctx: TaskContext[dict]) -> dict: + await gate.wait() + return {"msg": ctx.input["msg"]} + + first = await chat.start(task_id="t1", input={"msg": "first"}) + accepted = await chat.start(task_id="t1", input={"msg": "accepted"}) + registered = dict(manager._pending_steering_futures["t1"]) + accepted_future = next(iter(registered.values())) + rejected_future: asyncio.Future[Any] | None = None + + async def reject_update(task_id: str, _patch: Any) -> Any: + nonlocal rejected_future + rejected_future = list(manager._pending_steering_futures[task_id].values())[-1] + raise OSError("storage unavailable") + + with monkeypatch.context() as patcher: + patcher.setattr(manager.provider, "update", reject_update) + with pytest.raises(OSError, match="storage unavailable"): + await chat.start(task_id="t1", input={"msg": "rejected"}) + + assert rejected_future is not None + assert rejected_future.cancelled() + assert manager._pending_steering_futures["t1"] == registered + assert not accepted_future.done() + info = await manager.provider.get("t1") + assert info is not None + assert info.payload["steering"]["pending_inputs"] == [{"msg": "accepted"}] gate.set() - await asyncio.wait_for(run1.result(), timeout=5.0) + assert await asyncio.wait_for(first.result(), timeout=5.0) == {"msg": "first"} + assert await asyncio.wait_for(accepted.result(), timeout=5.0) == {"msg": "accepted"} + finally: + gate.set() + await self._teardown_manager(manager, mgr_mod) + + @pytest.mark.asyncio + async def test_steering_cancelled_append_unregisters_future(self, tmp_path, monkeypatch): + manager, mgr_mod = await self._setup_manager(tmp_path) + gate = asyncio.Event() + append_started = asyncio.Event() + try: + + @multi_turn_task(name="chat", steerable=True) + async def chat(ctx: TaskContext[dict]) -> dict: + await gate.wait() + return {"msg": ctx.input["msg"]} + + first = await chat.start(task_id="t1", input={"msg": "first"}) + aborted_future: asyncio.Future[Any] | None = None + + async def blocked_update(task_id: str, _patch: Any) -> Any: + nonlocal aborted_future + aborted_future = list(manager._pending_steering_futures[task_id].values())[-1] + append_started.set() + await asyncio.Event().wait() + + with monkeypatch.context() as patcher: + patcher.setattr(manager.provider, "update", blocked_update) + aborted = asyncio.create_task(chat.start(task_id="t1", input={"msg": "aborted"})) + try: + await asyncio.wait_for(append_started.wait(), timeout=5.0) + assert manager._get_task_write_lock("t1").locked() + finally: + aborted.cancel() + with pytest.raises(asyncio.CancelledError): + await aborted + + assert aborted_future is not None + assert aborted_future.cancelled() + assert "t1" not in manager._pending_steering_futures + info = await manager.provider.get("t1") + assert info is not None + assert info.payload.get("steering", {}).get("pending_inputs", []) == [] + + accepted = await chat.start(task_id="t1", input={"msg": "accepted"}) + gate.set() + assert await asyncio.wait_for(first.result(), timeout=5.0) == {"msg": "first"} + assert await asyncio.wait_for(accepted.result(), timeout=5.0) == {"msg": "accepted"} + finally: + gate.set() + await self._teardown_manager(manager, mgr_mod) + + @pytest.mark.asyncio + @pytest.mark.parametrize("cancel_request", [False, True]) + async def test_committed_append_without_response_preserves_later_ack(self, tmp_path, monkeypatch, cancel_request): + manager, mgr_mod = await self._setup_manager(tmp_path) + gate = asyncio.Event() + persisted = asyncio.Event() + seen: list[str] = [] + try: + + @multi_turn_task(name="chat", steerable=True) + async def chat(ctx: TaskContext[dict]) -> dict: + await gate.wait() + seen.append(ctx.input["msg"]) + return {"msg": ctx.input["msg"], "turn": len(seen)} + + first = await chat.start(task_id="t1", input={"msg": "first"}) + earlier = await chat.start(task_id="t1", input={"msg": "earlier"}) + original_update = manager.provider.update + + async def commit_without_response(task_id: str, patch: Any) -> Any: + await original_update(task_id, patch) + persisted.set() + if cancel_request: + await asyncio.Event().wait() + raise OSError("response lost after commit") + + with monkeypatch.context() as patcher: + patcher.setattr(manager.provider, "update", commit_without_response) + if cancel_request: + interrupted = asyncio.create_task( + chat.start(task_id="t1", input={"msg": "duplicate"}, input_id="same-id") + ) + try: + await asyncio.wait_for(persisted.wait(), timeout=5.0) + finally: + interrupted.cancel() + with pytest.raises(asyncio.CancelledError): + await interrupted + else: + with pytest.raises(OSError, match="response lost after commit"): + await chat.start(task_id="t1", input={"msg": "duplicate"}, input_id="same-id") + + later = await chat.start(task_id="t1", input={"msg": "duplicate"}, input_id="same-id") + info = await manager.provider.get("t1") + assert info is not None + assert [entry["msg"] for entry in info.payload["steering"]["pending_inputs"]] == [ + "earlier", + "duplicate", + "duplicate", + ] + ack_ids = info.payload["steering"]["pending_ack_ids"] + assert len(ack_ids) == len(set(ack_ids)) == 3 + + gate.set() + assert await asyncio.wait_for(first.result(), timeout=5.0) == {"msg": "first", "turn": 1} + assert await asyncio.wait_for(earlier.result(), timeout=5.0) == {"msg": "earlier", "turn": 2} + assert await asyncio.wait_for(later.result(), timeout=5.0) == {"msg": "duplicate", "turn": 4} + assert seen == ["first", "earlier", "duplicate", "duplicate"] + finally: + gate.set() + await self._teardown_manager(manager, mgr_mod) + + @pytest.mark.asyncio + async def test_legacy_backlog_does_not_claim_new_steering_ack(self, tmp_path): + from azure.ai.agentserver.core.tasks._models import TaskPatchRequest + + manager, mgr_mod = await self._setup_manager(tmp_path) + gate = asyncio.Event() + seen: list[str] = [] + try: + + @multi_turn_task(name="chat", steerable=True) + async def chat(ctx: TaskContext[dict]) -> dict: + await gate.wait() + seen.append(ctx.input["msg"]) + return {"msg": ctx.input["msg"], "turn": len(seen)} + + first = await chat.start(task_id="t1", input={"msg": "first"}) + existing = await manager.provider.get("t1") + assert existing is not None + steering = dict(existing.payload.get("steering") or {}) + steering["pending_inputs"] = [{"msg": "remote"}] + steering["cancel_requested"] = True + await manager.provider.update("t1", TaskPatchRequest(payload={"steering": steering}, if_match=existing.etag)) + + local = await chat.start(task_id="t1", input={"msg": "local"}) + info = await manager.provider.get("t1") + assert info is not None + ack_ids = info.payload["steering"]["pending_ack_ids"] + assert ack_ids[0] is None + assert isinstance(ack_ids[1], str) + + gate.set() + assert await asyncio.wait_for(first.result(), timeout=5.0) == {"msg": "first", "turn": 1} + assert await asyncio.wait_for(local.result(), timeout=5.0) == {"msg": "local", "turn": 3} + assert seen == ["first", "remote", "local"] + finally: + gate.set() + await self._teardown_manager(manager, mgr_mod) + + @pytest.mark.asyncio + async def test_queued_cancel_keeps_identical_input_owned_by_other_handle(self, tmp_path): + manager, mgr_mod = await self._setup_manager(tmp_path) + gate = asyncio.Event() + try: + @multi_turn_task(name="chat", steerable=True) + async def chat(ctx: TaskContext[dict]) -> dict: + await gate.wait() + return {"msg": ctx.input["msg"]} + + first = await chat.start(task_id="t1", input={"msg": "first"}) + earlier = await chat.start(task_id="t1", input={"msg": "same"}, input_id="shared-id") + later = await chat.start(task_id="t1", input={"msg": "same"}, input_id="shared-id") + info = await manager.provider.get("t1") + assert info is not None + earlier_ack_id = info.payload["steering"]["pending_ack_ids"][0] + + await later.cancel() + with pytest.raises(TaskCancelled): + await later.result() + info = await manager.provider.get("t1") + assert info is not None + assert info.payload["steering"]["pending_inputs"] == [{"msg": "same"}] + assert info.payload["steering"]["pending_ack_ids"] == [earlier_ack_id] + + gate.set() + assert await asyncio.wait_for(first.result(), timeout=5.0) == {"msg": "first"} + assert await asyncio.wait_for(earlier.result(), timeout=5.0) == {"msg": "same"} finally: + gate.set() await self._teardown_manager(manager, mgr_mod) @pytest.mark.asyncio diff --git a/sdk/agentserver/azure-ai-agentserver-core/tests/tasks/test_steering_attachment_queue.py b/sdk/agentserver/azure-ai-agentserver-core/tests/tasks/test_steering_attachment_queue.py index 2589372752d3..c0289091961c 100644 --- a/sdk/agentserver/azure-ai-agentserver-core/tests/tasks/test_steering_attachment_queue.py +++ b/sdk/agentserver/azure-ai-agentserver-core/tests/tasks/test_steering_attachment_queue.py @@ -194,6 +194,8 @@ async def runner(ctx: TaskContext[dict]) -> dict: assert len(pending_pre) == 2 assert _ref_key(pending_pre[0]) == "steering_input_0" assert _ref_key(pending_pre[1]) == "steering_input_1" + ack_ids_pre = info_pre.payload["steering"]["pending_ack_ids"] + assert len(ack_ids_pre) == len(set(ack_ids_pre)) == 2 assert info_pre.attachments["steering_input_0"] == a_value assert info_pre.attachments["steering_input_1"] == b_value @@ -207,6 +209,7 @@ async def runner(ctx: TaskContext[dict]) -> dict: pending_mid = info_mid.payload["steering"]["pending_inputs"] # Only B left in the queue. assert len(pending_mid) == 1 + assert info_mid.payload["steering"]["pending_ack_ids"] == [ack_ids_pre[1]] # B's attachment key MUST still be steering_input_1 (not renamed to _0). assert _ref_key(pending_mid[0]) == "steering_input_1" # A's attachment is gone; B's is unchanged.