From 729b1abf06b1925358d130e1bc48f3779a28b699 Mon Sep 17 00:00:00 2001 From: Josh Stevenson Date: Sat, 3 Oct 2026 06:54:54 -0700 Subject: [PATCH 1/2] Fix PR117 native continuity, owned resource retirement and partial outcomes --- agent/codex_runtime.py | 14 +- agent/context_engine.py | 10 + agent/memory_manager.py | 7 +- agent/transports/codex_app_server_session.py | 194 ++++++-- cli.py | 27 +- gateway/platforms/api_server.py | 115 ++++- gateway/run.py | 39 +- run_agent.py | 239 ++++++---- .../test_codex_app_server_session.py | 166 ++++++- .../test_codex_native_wire_hardening.py | 69 ++- .../gateway/test_native_partial_consumers.py | 402 ++++++++++++++++ .../test_codex_native_runtime_hardening.py | 201 ++++++++ .../test_bot_capability_refresh.py | 3 + .../tui_gateway/test_bot_local_retirement.py | 447 ++++++++++++++++++ .../tui_gateway/test_bot_native_retirement.py | 1 - .../test_prompt_recovery_contract.py | 61 +++ tui_gateway/server.py | 31 +- 17 files changed, 1836 insertions(+), 190 deletions(-) create mode 100644 tests/gateway/test_native_partial_consumers.py create mode 100644 tests/tui_gateway/test_bot_local_retirement.py diff --git a/agent/codex_runtime.py b/agent/codex_runtime.py index 3e876b4152cf3..e89e3def27e7d 100644 --- a/agent/codex_runtime.py +++ b/agent/codex_runtime.py @@ -727,8 +727,18 @@ def run_codex_app_server_turn( model=getattr(agent, "model", None), provider=getattr(agent, "provider", None), subscription_only_trial=subscription_only_trial, ): - existing.close() - agent._codex_session = None + if not existing.update_model( + model=getattr(agent, "model", None), provider=getattr(agent, "provider", None), + subscription_only_trial=subscription_only_trial, + ): + error = ( + "Codex native route change cannot preserve this thread. " + "Reset the conversation explicitly before changing provider or runtime mode." + ) + return { + "final_response": error, "messages": messages, "api_calls": 0, + "completed": False, "partial": True, "error": error, + } # Lazy session: one CodexAppServerSession per AIAgent instance. # Spawned on first turn, reused across turns, closed at AIAgent diff --git a/agent/context_engine.py b/agent/context_engine.py index 43ec54aa90875..35f14f8de3113 100644 --- a/agent/context_engine.py +++ b/agent/context_engine.py @@ -403,6 +403,16 @@ def on_session_end(self, session_id: str, messages: List[Dict[str, Any]]) -> Non NOT called per-turn — only when the session truly ends. """ + def shutdown(self) -> None: + """Release only handles owned by this engine instance. + + The logical session may continue in another engine instance with the + same session ID. Do not end it, delete durable context, or close shared + host resources here. Resource-bearing engines should override this + independently of on_session_end; resource-free engines need no work. + Tolerate handles already released by on_session_end at a real boundary. + """ + def on_session_reset(self) -> None: """Called on /new or /reset. Reset per-session state. diff --git a/agent/memory_manager.py b/agent/memory_manager.py index 818320ca723aa..7a3df22722dfb 100644 --- a/agent/memory_manager.py +++ b/agent/memory_manager.py @@ -1293,7 +1293,7 @@ def on_delegation(self, task: str, result: str, *, provider.name, e, ) - def shutdown_all(self) -> None: + def shutdown_all(self, *, preserve_providers=()) -> None: """Shut down all providers (reverse order for clean teardown). Drains the background sync/prefetch executor first (bounded by @@ -1301,9 +1301,14 @@ def shutdown_all(self) -> None: land before providers are torn down. The worker threads are daemon, so anything still wedged past the drain window dies with the interpreter rather than blocking exit. + + A same-session replacement may borrow provider instances. Drain this + manager's executor, but leave those exact live-owner instances open. """ self._drain_sync_executor() for provider in reversed(self._providers): + if any(provider is preserved for preserved in preserve_providers): + continue try: provider.shutdown() except Exception as e: diff --git a/agent/transports/codex_app_server_session.py b/agent/transports/codex_app_server_session.py index 22d73a2442ee9..70c24de8f032f 100644 --- a/agent/transports/codex_app_server_session.py +++ b/agent/transports/codex_app_server_session.py @@ -137,6 +137,7 @@ def _notification_belongs_to_turn( *, thread_id: Optional[str], turn_id: Optional[str], + require_explicit_scope: bool = False, ) -> bool: """Return whether a multiplexed notification belongs to this turn. @@ -150,6 +151,19 @@ def _notification_belongs_to_turn( return False observed_thread_id, observed_turn_id = _notification_scope_ids(note) + if require_explicit_scope: + if observed_thread_id is None or observed_turn_id is None: + return False + params = note.get("params") or {} + for name, scope in (("params", params), ("turn", params.get("turn")), ("item", params.get("item"))): + if not isinstance(scope, dict): + continue + if any(scope.get(key) is not None and str(scope[key]) != str(thread_id) + for key in ("threadId", "thread_id")): + return False + keys = ("id", "turnId", "turn_id") if name == "turn" else ("turnId", "turn_id") + if any(scope.get(key) is not None and str(scope[key]) != str(turn_id) for key in keys): + return False if ( thread_id is not None @@ -315,6 +329,8 @@ def __init__( self._interrupt_event = threading.Event() self._active_turn_id: Optional[str] = None self._active_turn_lock = threading.Lock() + # One caller owns a lifecycle boundary, including request/ACK latency. + self._operation_lock = threading.Lock() # Exclusions for prior turns belong to this session, alongside its # canonical thread identity; they confer no native qualification. self._known_turn_ids: set[str] = set() @@ -450,6 +466,36 @@ def matches_route( and self._subscription_only_trial == subscription_only_trial ) + def update_model( + self, model: Optional[str], provider: Optional[str], *, + subscription_only_trial: bool = False, + ) -> bool: + """Select the next turn's model without replacing authoritative history. + + turn/start.model is a sticky protocol override; sending the selected + model every turn also restores the primary model after a --once turn. + Provider and qualification-mode changes require an explicit reset. + """ + if not self._operation_lock.acquire(blocking=False): + return False + try: + if not self.matches_route( + self._model, provider, subscription_only_trial=subscription_only_trial, + ): + return False + selected = str(model or "").strip() + for prefix in ("openai/", "openai-codex/"): + if selected.startswith(prefix): + selected = selected[len(prefix):] + break + # An omitted override cannot restore an unknown thread default. + if not selected: + return False + self._model = selected + return True + finally: + self._operation_lock.release() + def close(self) -> None: if self._closed: return @@ -553,6 +599,25 @@ def run_turn( turn_timeout: float = 600.0, notification_poll_timeout: float = 0.25, post_tool_quiet_timeout: float = 90.0, + ) -> TurnResult: + if not self._operation_lock.acquire(blocking=False): + return TurnResult(thread_id=self._thread_id, error="codex session already has an active operation") + try: + return self._run_turn( + user_input, turn_timeout=turn_timeout, + notification_poll_timeout=notification_poll_timeout, + post_tool_quiet_timeout=post_tool_quiet_timeout, + ) + finally: + self._operation_lock.release() + + def _run_turn( + self, + user_input: Any, + *, + turn_timeout: float = 600.0, + notification_poll_timeout: float = 0.25, + post_tool_quiet_timeout: float = 90.0, ) -> TurnResult: """Send a user message and block until turn/completed, while forwarding server-initiated approval requests and projecting items @@ -597,15 +662,14 @@ def run_turn( # Send turn/start with the user input. Text-only for now (codex # supports rich content but Hermes' text path is the common case). + turn_params: dict[str, Any] = { + "threadId": self._thread_id, + "input": [{"type": "text", "text": user_input_text}], + } + if self._model: + turn_params["model"] = self._model try: - ts = self._client.request( - "turn/start", - { - "threadId": self._thread_id, - "input": [{"type": "text", "text": user_input_text}], - }, - timeout=10, - ) + ts = self._client.request("turn/start", turn_params, timeout=10) except CodexAppServerError as exc: # Classify auth/refresh failures so the user gets a clear # `codex login` pointer instead of a raw RPC error string. @@ -879,13 +943,30 @@ def compact_thread( *, turn_timeout: float = 600.0, notification_poll_timeout: float = 0.25, + ) -> TurnResult: + if not self._operation_lock.acquire(blocking=False): + return TurnResult(thread_id=self._thread_id, error="codex session already has an active operation") + try: + return self._compact_thread( + turn_timeout=turn_timeout, + notification_poll_timeout=notification_poll_timeout, + ) + finally: + self._operation_lock.release() + + def _compact_thread( + self, + *, + turn_timeout: float = 600.0, + notification_poll_timeout: float = 0.25, ) -> TurnResult: """Trigger Codex-native history compaction for the current thread. - Success requires an acknowledgement binding the requested operation - to a turn ID, followed by explicitly scoped started/completed events. - The current native protocol returns no ID, so it cannot establish - that binding and is conservatively retired with a correlation error. + The supported ACK is empty. Under the session's single-writer lock, + read prior history and drain pre-request lifecycle, then bind a fresh + same-thread turn/started. Only its completed contextCompaction item + and successful terminal certify completion, including events queued + while the request is waiting for its ACK. """ result = TurnResult() try: @@ -906,8 +987,44 @@ def compact_thread( return result projector = CodexEventProjector() + # thread/read is a supported non-mutating history boundary. The + # FIFO client reader has queued notifications preceding its response; + # persisted turn IDs also exclude delayed lifecycle for prior turns. + # Unsupported/malformed boundaries refuse before launching compaction + # and preserve the canonical thread for ordinary conversation. + try: + snapshot = self._client.request( + "thread/read", {"threadId": self._thread_id, "includeTurns": True}, timeout=10, + ) + except CodexAppServerError as exc: + result.error = self._format_error_with_stderr("cannot establish compaction history boundary", exc) + self._interrupt_event.clear() + return result + except (RuntimeError, TimeoutError, OSError) as exc: + result.error = self._format_error_with_stderr("compaction history boundary transport failed", exc) + result.should_retire = True + self.close() + return result + if self._interrupt_event.is_set(): + result.interrupted = True + self._interrupt_event.clear() + return result + thread = snapshot.get("thread") if isinstance(snapshot, dict) else None + turns = thread.get("turns") if isinstance(thread, dict) else None + status = thread.get("status") if isinstance(thread, dict) else None + if ( + not isinstance(thread, dict) or thread.get("id") != self._thread_id + or not isinstance(status, dict) or status.get("type") != "idle" + or not isinstance(turns, list) + or any(not isinstance(turn, dict) or not isinstance(turn.get("id"), str) + or not turn["id"] or turn.get("status") == "inProgress" for turn in turns) + ): + result.error = "cannot establish compaction history boundary: expected an idle thread with prior turn IDs" + return result + self._known_turn_ids.update(turn["id"] for turn in turns) + # Reject lifecycle already queued before the new request. Bound the - # drain so an endlessly streaming peer cannot prevent retirement. + # drain so an endlessly streaming peer cannot prevent refusal. for _ in range(1024): note = self._client.take_notification(timeout=0) if note is None: @@ -917,10 +1034,12 @@ def compact_thread( self._known_turn_ids.add(str(turn_id)) else: result.error = "cannot correlate compaction: notification backlog exceeds boundary limit" - result.should_retire = True - self.close() return result + if self._interrupt_event.is_set(): + result.interrupted = True + self._interrupt_event.clear() + return result try: acknowledgment = self._client.request( "thread/compact/start", @@ -956,16 +1075,14 @@ def compact_thread( self.close() return result - acknowledged_turn = acknowledgment.get("turn") if isinstance(acknowledgment, dict) else None - acknowledged_id = acknowledged_turn.get("id") if isinstance(acknowledged_turn, dict) else None - if not isinstance(acknowledged_id, str) or not acknowledged_id or acknowledged_id in self._known_turn_ids: - result.error = "cannot correlate compaction: acknowledgement has no fresh bound turn id" + if not isinstance(acknowledgment, dict): + result.error = "invalid thread/compact/start acknowledgement" result.should_retire = True self.close() return result - result.turn_id = acknowledged_id - self._known_turn_ids.add(acknowledged_id) turn_started = False + compaction_item_id: Optional[str] = None + compaction_completed = False deadline = time.monotonic() + turn_timeout turn_complete = False @@ -1012,33 +1129,46 @@ def compact_thread( method = note.get("method", "") observed_thread_id, observed_turn_id = _notification_scope_ids(note) + # Validate all explicit aliases before binding a new turn; a + # conflicting nested identity is not a legitimate start boundary. + if not _notification_belongs_to_turn( + note, thread_id=self._thread_id, + turn_id=result.turn_id if turn_started else observed_turn_id, + require_explicit_scope=True, + ): + continue if not turn_started: if ( method != "turn/started" or observed_thread_id is None or observed_turn_id is None or str(observed_thread_id) != str(self._thread_id) - or str(observed_turn_id) != acknowledged_id + or not isinstance(observed_turn_id, str) or not observed_turn_id + or observed_turn_id in self._known_turn_ids ): continue + result.turn_id = observed_turn_id + self._known_turn_ids.add(observed_turn_id) turn_started = True - if not _notification_belongs_to_turn( - note, - thread_id=self._thread_id, - turn_id=result.turn_id, - ): - logger.debug( - "ignoring foreign codex notification: method=%s", method - ) - continue - with self._active_turn_lock: self._active_turn_id = result.turn_id + params = note.get("params") or {} + item = params.get("item") or {} + if isinstance(item, dict) and item.get("type") == "contextCompaction": + item_id = item.get("id") + if method == "item/started" and isinstance(item_id, str) and item_id: + compaction_item_id = item_id + elif method == "item/completed" and item_id == compaction_item_id and item_id: + compaction_completed = True if method == "turn/completed": turn_complete = self._accept_terminal(note, result) if not turn_complete: continue + if result.completed and not compaction_completed: + result.completed = False + result.error = "compaction terminal lacked a completed contextCompaction item" + result.should_retire = True if self._on_event is not None: try: diff --git a/cli.py b/cli.py index 10ba49606f20f..29bf2d90f43dc 100644 --- a/cli.py +++ b/cli.py @@ -16968,6 +16968,13 @@ def run_agent(): # Get the final response response = result.get("final_response", "") if result else "" + # Native incomplete turns keep their draft separate from final_response. + # Project it for display without changing the completion contract. + _partial_draft = result.get("partial_response", "") if result and not response else "" + _partial_notice = "" + if _partial_draft: + _partial_notice = f"[Partial response — {result.get('error') or 'processing incomplete'}]" + response = f"{_partial_draft}\n\n{_partial_notice}" # Session titling now runs at TURN START (agent/turn_context.py) # from the user's message alone, so it is already done — or in @@ -17071,7 +17078,16 @@ def run_agent(): is_error_response = result and (result.get("failed") or result.get("partial")) already_streamed = self._stream_started and self._stream_box_opened and not is_error_response - if use_streaming_tts and _streaming_box_opened and not is_error_response: + if _partial_draft and ( + self._stream_started and self._stream_box_opened + or use_streaming_tts and _streaming_box_opened + ): + # The draft is already visible; label its outcome once. + if use_streaming_tts and _streaming_box_opened and not self._stream_box_opened: + w = self._scrollback_box_width() + _cprint(f"\n{_ACCENT}╰{'─' * (w - 2)}╯{_RST}") + _cprint(f"\n{_DIM}{_partial_notice}{_RST}") + elif use_streaming_tts and _streaming_box_opened and not is_error_response: # Text was already printed sentence-by-sentence; just close the box w = self._scrollback_box_width() _cprint(f"\n{_ACCENT}╰{'─' * (w - 2)}╯{_RST}") @@ -21700,6 +21716,9 @@ def _signal_handler_q(signum, frame): ): cli.session_id = cli.agent.session_id response = result.get("final_response", "") if isinstance(result, dict) else str(result) + if isinstance(result, dict) and not response and result.get("partial_response"): + print(result["partial_response"]) + print(f"Partial response: {result.get('error') or 'processing incomplete'}", file=sys.stderr) # Surface backend errors that produced no visible output # (e.g. invalid model slug → provider 4xx). Mirrors the # interactive CLI path. Write to stderr so piped stdout @@ -21707,6 +21726,7 @@ def _signal_handler_q(signum, frame): if ( not response and isinstance(result, dict) + and not result.get("partial_response") and result.get("error") and (result.get("failed") or result.get("partial")) ): @@ -21743,7 +21763,10 @@ def _signal_handler_q(signum, frame): # permanently block the card. Non-kanban runs keep the # plain 0/1 contract automation wrappers expect. _exit_code = 0 - if isinstance(result, dict) and result.get("failed"): + if isinstance(result, dict) and ( + result.get("failed") or result.get("partial") + or result.get("completed") is False or result.get("interrupted") + ): _exit_code = 1 if os.environ.get("HERMES_KANBAN_TASK") and result.get( "failure_reason" diff --git a/gateway/platforms/api_server.py b/gateway/platforms/api_server.py index c52b2efa8857b..17d2809496e75 100644 --- a/gateway/platforms/api_server.py +++ b/gateway/platforms/api_server.py @@ -4683,7 +4683,7 @@ async def _handle_session_chat(self, request: "web.Request") -> "web.Response": **agent_overrides, ) effective_session_id = result.get("session_id") if isinstance(result, dict) else session_id - final_response = _resolve_media_to_data_urls(result.get("final_response", "") if isinstance(result, dict) else "") + final_response = _resolve_media_to_data_urls((result.get("final_response") or result.get("partial_response") or "") if isinstance(result, dict) else "") headers = {"X-Hermes-Session-Id": effective_session_id or session_id} if gateway_session_key: headers["X-Hermes-Session-Key"] = gateway_session_key @@ -4709,6 +4709,10 @@ async def _handle_session_chat(self, request: "web.Request") -> "web.Response": "object": "hermes.session.chat.completion", "session_id": effective_session_id or session_id, "message": {"role": "assistant", "content": final_response}, + "completed": bool(result.get("completed", True)), + "partial": bool(result.get("partial")), + "interrupted": bool(result.get("interrupted")), + "error": _redact_api_error_text(result["error"]) if result.get("error") else None, "usage": usage, "runtime": runtime, }, @@ -4855,7 +4859,13 @@ async def _run_and_signal() -> None: confirmed_runtime_lock=lock_active, **agent_overrides, ) - final_response = _resolve_media_to_data_urls(result.get("final_response", "") if isinstance(result, dict) else "") + final_response = _resolve_media_to_data_urls((result.get("final_response") or result.get("partial_response") or "") if isinstance(result, dict) else "") + outcome = { + "completed": bool(result.get("completed", True)), + "partial": bool(result.get("partial")), + "interrupted": bool(result.get("interrupted")), + "error": _redact_api_error_text(result["error"]) if result.get("error") else None, + } effective_session_id = result.get("session_id", session_id) if isinstance(result, dict) else session_id turn_messages = self._turn_transcript_messages(history, user_message, result) if isinstance(result, dict) else [] effective_runtime = {} @@ -4879,9 +4889,7 @@ async def _run_and_signal() -> None: "session_id": effective_session_id, "message_id": message_id, "content": final_response, - "completed": True, - "partial": False, - "interrupted": False, + **outcome, "runtime": effective_runtime, })) # A steer accepted after the final assistant response is drained @@ -4892,20 +4900,26 @@ async def _run_and_signal() -> None: completed_payload = { "session_id": effective_session_id, "message_id": message_id, - "completed": True, + **outcome, "messages": turn_messages, "usage": usage, "runtime": effective_runtime, } if pending_steer: completed_payload["pending_steer"] = pending_steer - await queue.put(_event_payload("run.completed", completed_payload)) + terminal_status = ( + "completed" if outcome["completed"] + else "cancelled" if outcome["interrupted"] else "failed" + ) + terminal_event = f"run.{terminal_status}" + await queue.put(_event_payload(terminal_event, completed_payload)) self._set_run_status( run_id, - "completed", + terminal_status, session_id=effective_session_id, usage=usage, - last_event="run.completed", + last_event=terminal_event, + **outcome, **({"pending_steer": pending_steer} if pending_steer else {}), ) except asyncio.CancelledError: @@ -5319,7 +5333,7 @@ async def _compute_completion(): status=500, ) - final_response = _resolve_media_to_data_urls(result.get("final_response") or "") + final_response = _resolve_media_to_data_urls(result.get("final_response") or result.get("partial_response") or "") is_partial = bool(result.get("partial")) is_failed = bool(result.get("failed")) completed = bool(result.get("completed", True)) @@ -5331,7 +5345,7 @@ async def _compute_completion(): # codes. See issue #22496. if is_partial and err_msg and "truncat" in err_msg.lower(): finish_reason = "length" - elif is_failed or (not completed and err_msg): + elif is_failed or is_partial or not completed or result.get("interrupted"): finish_reason = "error" else: finish_reason = "stop" @@ -5388,6 +5402,7 @@ async def _compute_completion(): response_data["hermes"] = { "completed": completed, "partial": is_partial, + "interrupted": bool(result.get("interrupted")), "failed": is_failed, "error": err_msg, "error_code": "output_truncated" if finish_reason == "length" else "agent_error", @@ -5439,6 +5454,7 @@ async def _write_sse_chat_completion( } await response.write(_sse_frame(role_chunk)) last_activity = time.monotonic() + content_emitted = False # Helper — route a queue item to the correct SSE event. async def _emit(item): @@ -5451,9 +5467,11 @@ async def _emit(item): conversation history. See #6972 for the original event, #16588 for the ``toolCallId``/``status`` lifecycle fields. """ + nonlocal content_emitted if isinstance(item, tuple) and len(item) == 2 and item[0] == "__tool_progress__": await response.write(_sse_frame(item[1], event="hermes.tool.progress")) else: + content_emitted = content_emitted or bool(item) content_chunk = { "id": completion_id, "object": "chat.completion.chunk", "created": created, "model": model, @@ -5516,6 +5534,10 @@ async def _emit(item): is_failed = bool(result.get("failed")) if isinstance(result, dict) else False completed = bool(result.get("completed", True)) if isinstance(result, dict) else True err_msg = result.get("error") if isinstance(result, dict) else None + if isinstance(result, dict) and not content_emitted: + draft = result.get("partial_response") + if draft: + await _emit(_resolve_media_to_data_urls(draft)) if agent_error is not None: is_failed = True err_msg = err_msg or str(agent_error) @@ -5524,7 +5546,7 @@ async def _emit(item): # for truncation, "error" for failure, "stop" for normal completion. if is_partial and err_msg and "truncat" in err_msg.lower(): finish_reason = "length" - elif agent_error is not None or is_failed or (not completed and err_msg): + elif agent_error is not None or is_failed or is_partial or not completed or (isinstance(result, dict) and result.get("interrupted")): finish_reason = "error" else: finish_reason = "stop" @@ -5550,6 +5572,7 @@ async def _emit(item): finish_chunk["hermes"] = { "completed": completed, "partial": is_partial, + "interrupted": bool(result.get("interrupted")) if isinstance(result, dict) else False, "failed": is_failed, "error": err_msg, "error_code": "output_truncated" if finish_reason == "length" else "agent_error", @@ -5695,6 +5718,7 @@ def _envelope(status: str) -> Dict[str, Any]: return env final_response_text = "" + result = None agent_error: Optional[str] = None usage: Dict[str, int] = {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0} terminal_snapshot_persisted = False @@ -6002,12 +6026,12 @@ async def _flush_batch() -> None: # deltas were streamed (e.g. some providers only emit # the full response at the end), emit a single fallback # delta so Responses clients still receive a live text part. - agent_final = result.get("final_response", "") if isinstance(result, dict) else "" + agent_final = (result.get("final_response") or result.get("partial_response") or "") if isinstance(result, dict) else "" if agent_final and not final_text_parts: await _emit_text_delta(agent_final) if agent_final and not final_response_text: final_response_text = agent_final - if isinstance(result, dict) and result.get("error") and not final_response_text: + if isinstance(result, dict) and result.get("error"): agent_error = _redact_api_error_text(result["error"]) except Exception as e: # noqa: BLE001 logger.error("Error running agent for streaming responses: %s", e, exc_info=True) @@ -6015,6 +6039,11 @@ async def _flush_batch() -> None: # Close the message item if it was opened final_response_text = "".join(final_text_parts) or final_response_text + incomplete_result = bool(isinstance(result, dict) and ( + result.get("partial") or result.get("failed") + or result.get("completed") is False or result.get("interrupted") + )) + message_status = "incomplete" if incomplete_result or agent_error else "completed" if message_opened: await _write_event("response.output_text.done", { "type": "response.output_text.done", @@ -6027,7 +6056,7 @@ async def _flush_batch() -> None: msg_done_item = { "id": message_item_id, "type": "message", - "status": "completed", + "status": message_status, "role": "assistant", "content": [ {"type": "output_text", "text": final_response_text} @@ -6071,16 +6100,26 @@ async def _flush_batch() -> None: final_items.append({ "type": "message", + "status": message_status, "role": "assistant", "content": [ {"type": "output_text", "text": final_response_text or (_redact_api_error_text(agent_error) if agent_error else "")} ], }) - if agent_error: - failed_env = _envelope("failed") + if agent_error or incomplete_result: + terminal_status = "incomplete" if isinstance(result, dict) and result.get("partial") and not result.get("failed") else "failed" + failed_env = _envelope(terminal_status) failed_env["output"] = final_items - failed_env["error"] = {"message": _redact_api_error_text(agent_error), "type": "server_error"} + if agent_error: + failed_env["error"] = {"message": _redact_api_error_text(agent_error), "type": "server_error"} + if isinstance(result, dict): + failed_env["hermes"] = { + "completed": bool(result.get("completed", True)), + "partial": bool(result.get("partial")), + "interrupted": bool(result.get("interrupted")), + "error": agent_error, + } failed_env["usage"] = { "input_tokens": usage.get("input_tokens", 0), "output_tokens": usage.get("output_tokens", 0), @@ -6098,8 +6137,8 @@ async def _flush_batch() -> None: conversation_history_snapshot=_failed_history, ) terminal_snapshot_persisted = True - await _write_event("response.failed", { - "type": "response.failed", + await _write_event(f"response.{terminal_status}", { + "type": f"response.{terminal_status}", "response": failed_env, }) else: @@ -6453,7 +6492,7 @@ async def _compute_response(): status=500, ) - final_response = _resolve_media_to_data_urls(result.get("final_response", "")) + final_response = _resolve_media_to_data_urls(result.get("final_response") or result.get("partial_response") or "") if not final_response: final_response = _redact_api_error_text(result.get("error", "(No response generated)")) @@ -6503,6 +6542,17 @@ async def _compute_response(): }, } + if result.get("partial") or result.get("failed") or result.get("completed") is False or result.get("interrupted"): + response_data["status"] = "failed" if result.get("failed") else "incomplete" + response_data["hermes"] = { + "completed": bool(result.get("completed", True)), + "partial": bool(result.get("partial")), + "interrupted": bool(result.get("interrupted")), + "error": _redact_api_error_text(result["error"]) if result.get("error") else None, + } + if result.get("error"): + response_data["error"] = {"message": _redact_api_error_text(result["error"]), "type": "server_error"} + # Store the complete response object for future chaining / GET retrieval if store: self._response_store.put(response_id, { @@ -7108,12 +7158,13 @@ def _extract_output_items(result: Dict[str, Any], start_index: int = 0) -> List[ }) # Final assistant message - final = result.get("final_response", "") + final = result.get("final_response") or result.get("partial_response") or "" if not final: final = _redact_api_error_text(result.get("error", "(No response generated)")) items.append({ "type": "message", + "status": "incomplete" if (result.get("partial") or result.get("failed") or result.get("completed") is False or result.get("interrupted")) else "completed", "role": "assistant", "content": [ { @@ -7861,19 +7912,31 @@ def _run_sync(): # Check for structured failure (non-retryable client errors like # 401/400 return failed=True instead of raising, so the except # block below never fires — issue #15561). - elif isinstance(result, dict) and result.get("failed"): + elif isinstance(result, dict) and ( + result.get("failed") or result.get("partial") + or result.get("completed") is False or result.get("interrupted") + ): error_msg = _redact_api_error_text(result.get("error") or "agent run failed") + status = "cancelled" if result.get("interrupted") else "failed" + outcome = { + "output": result.get("final_response") or result.get("partial_response") or "", + "completed": False, + "partial": bool(result.get("partial")), + "interrupted": bool(result.get("interrupted")), + } _put_event_if_active({ - "event": "run.failed", + "event": f"run.{status}", "run_id": run_id, "timestamp": time.time(), "error": error_msg, + **outcome, }) self._set_run_status( run_id, - "failed", + status, error=error_msg, - last_event="run.failed", + last_event=f"run.{status}", + **outcome, ) else: final_response = result.get("final_response", "") if isinstance(result, dict) else "" diff --git a/gateway/run.py b/gateway/run.py index dbcd55ebf2c7b..2fac415fd011f 100644 --- a/gateway/run.py +++ b/gateway/run.py @@ -4030,6 +4030,14 @@ def _normalize_empty_agent_response( if response: return response + partial_response = agent_result.get("partial_response") + if partial_response: + notice = f"⚠️ Partial response: {agent_result.get('error') or 'processing incomplete'}" + # Streaming already delivered the draft. Send only its outcome label. + if agent_result.get("partial_response_previewed"): + return notice + return f"{partial_response}\n\n{notice}" + if agent_result.get("failed"): # None-safe: the gateway result dict is built with # ``'error': holder.get('error')`` and can carry an EXPLICIT None, @@ -4136,6 +4144,8 @@ def _is_gateway_hidden_reasoning_incomplete_turn(agent_result: dict) -> bool: return False if not agent_result.get("partial"): return False + if agent_result.get("partial_response"): + return False error_text = str(agent_result.get("error", "") or "").strip() if "remained incomplete after" not in error_text.lower(): return False @@ -6725,14 +6735,18 @@ def _approval_notify_sync(approval_data: dict) -> None: ) if not final_response: - final_response = _normalize_empty_agent_response( - result, final_response or "", history_len=len(agent_history), - ) - final_response = _sanitize_gateway_final_response(ctx.source.platform, final_response) - if not final_response: - final_response = f"⚠️ {result['error']}" if result.get("error") else "" + # Preserve the producer's empty final for native incomplete turns. + # Delivery normalization below renders the separately labeled draft. + if not result.get("partial_response"): + final_response = _normalize_empty_agent_response( + result, final_response or "", history_len=len(agent_history), + ) + final_response = _sanitize_gateway_final_response(ctx.source.platform, final_response) + if not final_response: + final_response = f"⚠️ {result['error']}" if result.get("error") else "" return { "final_response": final_response, + "partial_response": result.get("partial_response", ""), "messages": result.get("messages", []), "api_calls": result.get("api_calls", 0), "failed": result.get("failed", False), @@ -6809,6 +6823,7 @@ def _approval_notify_sync(approval_data: dict) -> None: return { "final_response": final_response, + "partial_response": result.get("partial_response", ""), "last_reasoning": result.get("last_reasoning"), "messages": ctx.result_holder[0].get("messages", []) if ctx.result_holder[0] else [], "api_calls": ctx.result_holder[0].get("api_calls", 0) if ctx.result_holder[0] else 0, @@ -30294,7 +30309,12 @@ def _run_sync_with_timeout_lifecycle(): # _run_agent_task; sending the raw copy bypasses those steps. _delivery_result = response if isinstance(response, dict) else (result or {}) _previewed = bool(_delivery_result.get("response_previewed")) - first_response = _delivery_result.get("final_response", "") + _partial_draft = _delivery_result.get("partial_response", "") + if _partial_draft and _stream_confirmed_final_delivery(_sc, _partial_draft): + _delivery_result["partial_response_previewed"] = True + first_response = _normalize_empty_agent_response( + _delivery_result, _delivery_result.get("final_response", ""), + ) _already_streamed = _stream_confirmed_final_delivery( _sc, first_response, @@ -30561,6 +30581,11 @@ def _run_sync_with_timeout_lifecycle(): # final answer. Suppressing delivery here leaves the user staring # at silence. (#10xxx — "agent stops after web search") _sc = stream_consumer_holder[0] + if isinstance(response, dict) and response.get("partial_response"): + # A successful stream seal preserved the accumulator. Keep the + # subsequent result-only delivery to a notice, without repeating it. + if _stream_confirmed_final_delivery(_sc, response["partial_response"]): + response["partial_response_previewed"] = True if isinstance(response, dict) and not response.get("failed"): _final = response.get("final_response") or "" _is_empty_sentinel = not _final or _final == "(empty)" diff --git a/run_agent.py b/run_agent.py index 58bba015ed537..28efc34ca2384 100644 --- a/run_agent.py +++ b/run_agent.py @@ -551,6 +551,7 @@ def __init__( checkpoint_max_file_size_mb: int = 10, pass_session_id: bool = False, requested_provider: str = None, + _resource_preserve_agent=None, ): """Forwarder — see ``agent.agent_init.init_agent``.""" if tool_delay is not None: @@ -561,87 +562,97 @@ def __init__( stacklevel=2, ) from agent.agent_init import init_agent - init_agent( - self, - base_url=base_url, - api_key=api_key, - provider=provider, - requested_provider=requested_provider, - api_mode=api_mode, - acp_command=acp_command, - acp_args=acp_args, - command=command, - args=args, - model=model, - max_iterations=max_iterations, - enabled_toolsets=enabled_toolsets, - disabled_toolsets=disabled_toolsets, - save_trajectories=save_trajectories, - verbose_logging=verbose_logging, - quiet_mode=quiet_mode, - tool_progress_mode=tool_progress_mode, - ephemeral_system_prompt=ephemeral_system_prompt, - log_prefix_chars=log_prefix_chars, - log_prefix=log_prefix, - providers_allowed=providers_allowed, - providers_ignored=providers_ignored, - providers_order=providers_order, - provider_sort=provider_sort, - provider_require_parameters=provider_require_parameters, - provider_data_collection=provider_data_collection, - openrouter_min_coding_score=openrouter_min_coding_score, - session_id=session_id, - tool_progress_callback=tool_progress_callback, - tool_start_callback=tool_start_callback, - tool_complete_callback=tool_complete_callback, - thinking_callback=thinking_callback, - reasoning_callback=reasoning_callback, - clarify_callback=clarify_callback, - read_terminal_callback=read_terminal_callback, - read_preview_callback=read_preview_callback, - drive_preview_callback=drive_preview_callback, - read_window_below_callback=read_window_below_callback, - setup_mcp_callback=setup_mcp_callback, - tour_callback=tour_callback, - step_callback=step_callback, - stream_delta_callback=stream_delta_callback, - interim_assistant_callback=interim_assistant_callback, - tool_gen_callback=tool_gen_callback, - status_callback=status_callback, - notice_callback=notice_callback, - notice_clear_callback=notice_clear_callback, - event_callback=event_callback, - reaction_callback=reaction_callback, - max_tokens=max_tokens, - reasoning_config=reasoning_config, - service_tier=service_tier, - request_overrides=request_overrides, - prefill_messages=prefill_messages, - platform=platform, - user_id=user_id, - user_id_alt=user_id_alt, - user_name=user_name, - chat_id=chat_id, - chat_name=chat_name, - chat_type=chat_type, - thread_id=thread_id, - gateway_session_key=gateway_session_key, - skip_context_files=skip_context_files, - load_soul_identity=load_soul_identity, - skip_memory=skip_memory, - skip_background_review=skip_background_review, - session_db=session_db, - parent_session_id=parent_session_id, - iteration_budget=iteration_budget, - run_budget_seconds=run_budget_seconds, - fallback_model=fallback_model, - credential_pool=credential_pool, - checkpoints_enabled=checkpoints_enabled, - checkpoint_max_snapshots=checkpoint_max_snapshots, - checkpoint_max_total_size_mb=checkpoint_max_total_size_mb, - checkpoint_max_file_size_mb=checkpoint_max_file_size_mb, - pass_session_id=pass_session_id, - ) + try: + init_agent( + self, + base_url=base_url, + api_key=api_key, + provider=provider, + requested_provider=requested_provider, + api_mode=api_mode, + acp_command=acp_command, + acp_args=acp_args, + command=command, + args=args, + model=model, + max_iterations=max_iterations, + enabled_toolsets=enabled_toolsets, + disabled_toolsets=disabled_toolsets, + save_trajectories=save_trajectories, + verbose_logging=verbose_logging, + quiet_mode=quiet_mode, + tool_progress_mode=tool_progress_mode, + ephemeral_system_prompt=ephemeral_system_prompt, + log_prefix_chars=log_prefix_chars, + log_prefix=log_prefix, + providers_allowed=providers_allowed, + providers_ignored=providers_ignored, + providers_order=providers_order, + provider_sort=provider_sort, + provider_require_parameters=provider_require_parameters, + provider_data_collection=provider_data_collection, + openrouter_min_coding_score=openrouter_min_coding_score, + session_id=session_id, + tool_progress_callback=tool_progress_callback, + tool_start_callback=tool_start_callback, + tool_complete_callback=tool_complete_callback, + thinking_callback=thinking_callback, + reasoning_callback=reasoning_callback, + clarify_callback=clarify_callback, + read_terminal_callback=read_terminal_callback, + read_preview_callback=read_preview_callback, + drive_preview_callback=drive_preview_callback, + read_window_below_callback=read_window_below_callback, + setup_mcp_callback=setup_mcp_callback, + tour_callback=tour_callback, + step_callback=step_callback, + stream_delta_callback=stream_delta_callback, + interim_assistant_callback=interim_assistant_callback, + tool_gen_callback=tool_gen_callback, + status_callback=status_callback, + notice_callback=notice_callback, + notice_clear_callback=notice_clear_callback, + event_callback=event_callback, + reaction_callback=reaction_callback, + max_tokens=max_tokens, + reasoning_config=reasoning_config, + service_tier=service_tier, + request_overrides=request_overrides, + prefill_messages=prefill_messages, + platform=platform, + user_id=user_id, + user_id_alt=user_id_alt, + user_name=user_name, + chat_id=chat_id, + chat_name=chat_name, + chat_type=chat_type, + thread_id=thread_id, + gateway_session_key=gateway_session_key, + skip_context_files=skip_context_files, + load_soul_identity=load_soul_identity, + skip_memory=skip_memory, + skip_background_review=skip_background_review, + session_db=session_db, + parent_session_id=parent_session_id, + iteration_budget=iteration_budget, + run_budget_seconds=run_budget_seconds, + fallback_model=fallback_model, + credential_pool=credential_pool, + checkpoints_enabled=checkpoints_enabled, + checkpoint_max_snapshots=checkpoint_max_snapshots, + checkpoint_max_total_size_mb=checkpoint_max_total_size_mb, + checkpoint_max_file_size_mb=checkpoint_max_file_size_mb, + pass_session_id=pass_session_id, + ) + except Exception: + # The constructor may fail after creating local clients/providers. + # No caller receives this object, so retire them here without + # ending its shared session or masking the construction failure. + try: + self.retire_local_resources(preserve_agent=_resource_preserve_agent) + except Exception: + logger.debug("Partial agent initialization cleanup failed", exc_info=True) + raise def _get_session_db_for_recall(self): """Return a SessionDB for recall, lazily creating it if an entrypoint forgot. @@ -4486,26 +4497,40 @@ def get_activity_summary(self) -> dict: }, ) - def shutdown_memory_provider(self, messages: list = None) -> None: - """Shut down the memory provider and context engine at session end. + def shutdown_memory_provider( + self, messages: list = None, *, end_session: bool = True, preserve_agent=None, + ) -> None: + """Shut down this instance's memory provider. Idempotent: gateway cleanup and AIAgent.close() may share this - ownership boundary. + ownership boundary. Local replacement passes end_session=False: + provider/engine shutdown releases instance resources, but logical + session-end hooks must wait for a real session boundary. + preserve_agent identifies the live owner of any borrowed instances. """ if getattr(self, "_memory_provider_shutdown", False): return self._memory_provider_shutdown = True - if self._memory_manager: - try: - self._memory_manager.on_session_end(messages or []) - except Exception as e: - logger.warning("Memory provider on_session_end failed during shutdown: %s", e, exc_info=True) + memory_manager = getattr(self, "_memory_manager", None) + preserved_manager = getattr(preserve_agent, "_memory_manager", None) + if memory_manager and memory_manager is not preserved_manager: + if end_session: + try: + memory_manager.on_session_end(messages or []) + except Exception as e: + logger.warning("Memory provider on_session_end failed during shutdown: %s", e, exc_info=True) try: - self._memory_manager.shutdown_all() + if preserved_manager is None: + memory_manager.shutdown_all() + else: + memory_manager.shutdown_all( + preserve_providers=preserved_manager.providers, + ) except Exception: pass + # Notify context engine of session end (flush DAG, close DBs, etc.) - if hasattr(self, "context_compressor") and self.context_compressor: + if end_session and hasattr(self, "context_compressor") and self.context_compressor: try: self.context_compressor.on_session_end( self.session_id or "", @@ -4514,6 +4539,14 @@ def shutdown_memory_provider(self, messages: list = None) -> None: except Exception: pass + engine = getattr(self, "context_compressor", None) + engine_shutdown = getattr(engine, "shutdown", None) + if engine is not getattr(preserve_agent, "context_compressor", None) and callable(engine_shutdown): + try: + engine_shutdown() + except Exception: + logger.debug("Context engine local shutdown failed", exc_info=True) + def commit_memory_session(self, messages: list = None) -> None: """Trigger end-of-session extraction without tearing providers down. Called when session_id rotates (e.g. /new, context compression); @@ -4669,6 +4702,30 @@ def release_clients(self) -> None: except Exception: pass + def retire_local_resources(self, *, preserve_agent=None) -> None: + """Retire this instance while its logical session continues elsewhere. + + Capability refresh constructs separate providers for the same session + ID. Drain their work without logical session-end hooks, ending the + SQLite row, closing its shared DB, or destroying task-scoped tools. + Context engines release only instance handles through shutdown(), + reserving on_session_end for real session boundaries. Native transport + retirement has a separate ownership boundary in the caller. + Borrowed provider/engine instances belonging to preserve_agent survive. + """ + if getattr(self, "_local_resources_retired", False): + return + self._local_resources_retired = True + try: + self.release_clients() + finally: + session_messages = getattr(self, "_session_messages", None) + self.shutdown_memory_provider( + session_messages if isinstance(session_messages, list) else None, + end_session=False, + preserve_agent=preserve_agent, + ) + def close(self) -> None: """Release all resources held by this agent instance. diff --git a/tests/agent/transports/test_codex_app_server_session.py b/tests/agent/transports/test_codex_app_server_session.py index ecd0530c732b3..4fbc9ff560fb0 100644 --- a/tests/agent/transports/test_codex_app_server_session.py +++ b/tests/agent/transports/test_codex_app_server_session.py @@ -53,6 +53,9 @@ def request(self, method: str, params: Optional[dict] = None, timeout: float = 3 if method == "thread/start": return {"thread": {"id": "thread-fake-001"}, "activePermissionProfile": {"id": "workspace-write"}} + if method == "thread/read": + return {"thread": {"id": "thread-fake-001", "status": {"type": "idle"}, + "turns": getattr(self, "_history", [])}} if method == "turn/start": return {"turn": {"id": "turn-fake-001"}} if method == "turn/interrupt": @@ -126,21 +129,30 @@ def make_session(client: FakeClient, **kwargs) -> CodexAppServerSession: ) -def emit_compaction_on_request(client, *, turn_id="compact-turn-1", bind_turn=True): - """Model new events from a fake request, optionally with an explicit ID. - - The current native compact protocol has no such ID; bound success tests - exercise the adapter contract, not native protocol qualification. - """ +def emit_compaction_on_request(client, *, turn_id="compact-turn-1", complete_item=True): + """Emit scoped lifecycle before returning the protocol's empty ACK.""" notes = list(client._notifications) client._notifications.clear() + if complete_item: + # Supply the documented single contextCompaction item pair. Existing + # transcript/terminal fixtures still control start and terminal order. + index = next((i + 1 for i, note in enumerate(notes) + if note["method"] == "turn/started" + and note["params"].get("threadId") == "thread-fake-001" + and note["params"].get("turn", {}).get("id") == turn_id), len(notes)) + item = {"type": "contextCompaction", "id": "fixture-compaction"} + notes[index:index] = [ + {"method": method, "params": {"threadId": "thread-fake-001", "turnId": turn_id, + "item": dict(item)}} + for method in ("item/started", "item/completed") + ] original_request = client.request def request(method, params=None, timeout=30): response = original_request(method, params, timeout) if method == "thread/compact/start": client._notifications.extend(notes) - return {"turn": {"id": turn_id}} if bind_turn else {} + return {} return response client.request = request @@ -516,9 +528,11 @@ def test_compact_thread_ignores_foreign_child_completion(self): assert result.error is None assert result.turn_id == "compact-turn-1" assert result.final_text == "parent compacted" - assert result.projected_messages == [ - {"role": "assistant", "content": "parent compacted"} - ] + parent_message = {"role": "assistant", "content": "parent compacted"} + assert result.projected_messages.count(parent_message) == 1 + assert result.projected_messages[-1] == parent_message + assert not any("child" in message.get("content", "") for message in result.projected_messages) + assert result.compacted @@ -1216,7 +1230,7 @@ def request(method, params): class TestNativeCompactionTerminalProof: @pytest.mark.parametrize("preexisting", [False, True]) - def test_no_id_ack_cannot_certify_even_typed_compact_lifecycle(self, preexisting): + def test_empty_ack_preserves_thread_only_for_fresh_lifecycle(self, preexisting): client = FakeClient() client.queue_notification("turn/started", threadId="t", turn={"id": "old"}) client.queue_notification("item/completed", threadId="t", turnId="old", @@ -1226,11 +1240,17 @@ def test_no_id_ack_cannot_certify_even_typed_compact_lifecycle(self, preexisting client.queue_notification("turn/completed", threadId="t", turn={"id": "old", "status": "completed"}) if not preexisting: - emit_compaction_on_request(client, bind_turn=False) - result = make_session(client).compact_thread(turn_timeout=0.01) - assert not result.completed and result.final_text == "" - assert result.should_retire and client._closed - assert "correlat" in result.error + emit_compaction_on_request(client, turn_id="old") + session = make_session(client) + result = session.compact_thread(turn_timeout=0.01) + assert result.completed is (not preexisting) + assert result.should_retire is preexisting + assert client._closed is preexisting + if not preexisting: + assert session._thread_id == "thread-fake-001" + assert result.turn_id == "old" and result.final_text == "old summary" + else: + assert result.final_text == "" and "timed out" in result.error def test_compaction_ack_cannot_reuse_known_previous_turn(self): client = FakeClient() @@ -1255,7 +1275,7 @@ def test_compaction_ack_cannot_accept_prequeued_matching_lifecycle(self): def request(method, params=None, timeout=30): response = original_request(method, params, timeout) - return {"turn": {"id": "prior"}} if method == "thread/compact/start" else response + return {} if method == "thread/compact/start" else response client.request = request result = make_session(client).compact_thread(turn_timeout=0.01) @@ -1303,6 +1323,8 @@ def test_compaction_start_uncertainty_closes_child(self, error): def request(method, params): if method == "thread/start": return {"thread": {"id": "thread-fake-001"}} + if method == "thread/read": + return {"thread": {"id": "thread-fake-001", "status": {"type": "idle"}, "turns": []}} if method == "thread/compact/start": raise error return {} @@ -1380,3 +1402,113 @@ def initialize(**kwargs): assert not any(method == "thread/compact/start" for method, _ in client.requests) assert session._active_turn_id is None assert not session._interrupt_event.is_set() + + +class TestNativeCompactionBoundary: + def test_late_prior_lifecycle_is_excluded_by_history_read(self): + client = FakeClient() + client._history = [{"id": "delayed-prior", "status": "completed"}] + for turn_id in ("delayed-prior", "fresh"): + client.queue_notification("turn/started", threadId="t", turn={"id": turn_id}) + client.queue_notification("item/started", threadId="t", turnId=turn_id, + item={"type": "contextCompaction", "id": turn_id + "-item"}) + client.queue_notification("item/completed", threadId="t", turnId=turn_id, + item={"type": "contextCompaction", "id": turn_id + "-item"}) + client.queue_notification("turn/completed", threadId="t", + turn={"id": turn_id, "status": "completed"}) + emit_compaction_on_request(client, turn_id="fresh", complete_item=False) + session = make_session(client) + result = session.compact_thread(turn_timeout=0.1) + assert result.completed and result.compacted and result.turn_id == "fresh" + assert not session._closed and session._thread_id == "thread-fake-001" + assert client.requests[1] == ("thread/read", {"threadId": "thread-fake-001", "includeTurns": True}) + + @pytest.mark.parametrize("boundary", [ + {}, {"thread": {"id": "foreign", "status": {"type": "idle"}, "turns": []}}, + {"thread": {"id": "thread-fake-001", "status": {"type": "active"}, "turns": []}}, + {"thread": {"id": "thread-fake-001", "status": {"type": "idle"}, + "turns": [{"id": "inflight", "status": "inProgress"}]}}, + ]) + def test_unavailable_boundary_refuses_before_compact_without_closing(self, boundary): + client = FakeClient() + original = client.request + client.request = lambda method, params=None, timeout=30: ( + boundary if method == "thread/read" else original(method, params, timeout)) + session = make_session(client) + result = session.compact_thread(turn_timeout=0.01) + assert result.error and not result.should_retire and not session._closed + assert session._thread_id == "thread-fake-001" + assert not any(method == "thread/compact/start" for method, _ in client.requests) + + def test_regular_turn_cannot_certify_compaction_without_item_pair(self): + client = FakeClient() + client.queue_notification("turn/started", threadId="t", turn={"id": "regular"}) + client.queue_notification("turn/completed", threadId="t", + turn={"id": "regular", "status": "completed"}) + emit_compaction_on_request(client, turn_id="regular", complete_item=False) + result = make_session(client).compact_thread(turn_timeout=0.01) + assert not result.completed and result.should_retire + assert "contextCompaction" in result.error + + @pytest.mark.parametrize("outer", ["run_turn", "compact_thread"]) + def test_concurrent_operation_is_refused_during_request_ack(self, outer): + client = FakeClient() + session = make_session(client) + refused = [] + client.queue_notification("turn/started", threadId="t", turn={"id": "compact-turn-1"}) + client.queue_notification("turn/completed", threadId="t", + turn={"id": "compact-turn-1", "status": "completed"}) + if outer == "compact_thread": + emit_compaction_on_request(client) + else: + client._notifications.clear() + client.queue_notification("turn/completed", threadId="t", + turn={"id": "turn-fake-001", "status": "completed"}) + original = client.request + launch = "turn/start" if outer == "run_turn" else "thread/compact/start" + def request(method, params=None, timeout=30): + if method == launch: + refused.append( + session.compact_thread(turn_timeout=0.02, notification_poll_timeout=0.001) + if outer == "run_turn" else + session.run_turn("overlap", turn_timeout=0.02, notification_poll_timeout=0.001) + ) + return original(method, params, timeout) + client.request = request + result = session.run_turn("outer", turn_timeout=0.1) if outer == "run_turn" else session.compact_thread(turn_timeout=0.1) + assert result.completed and not session._closed + assert len(refused) == 1 and refused[0].error and not refused[0].should_retire + assert sum(method == launch for method, _ in client.requests) == 1 + + +@pytest.mark.parametrize("conflict", [ + {"threadId": "thread-fake-001", "turnId": "fresh", "turn": {"id": "prior"}}, + {"threadId": "thread-fake-001", "turn": {"id": "fresh", "threadId": "foreign"}}, +]) +def test_compaction_rejects_conflicting_start_scope(conflict): + client = FakeClient() + client._notifications.append({"method": "turn/started", "params": conflict}) + client.queue_notification("item/started", threadId="t", turnId="fresh", + item={"type": "contextCompaction", "id": "compact-item"}) + client.queue_notification("item/completed", threadId="t", turnId="fresh", + item={"type": "contextCompaction", "id": "compact-item"}) + client.queue_notification("turn/completed", threadId="t", turn={"id": "fresh", "status": "completed"}) + emit_compaction_on_request(client, turn_id="fresh", complete_item=False) + result = make_session(client).compact_thread(turn_timeout=0.01) + assert not result.completed and result.should_retire and result.turn_id is None + + +def test_compaction_interrupt_at_history_boundary_preserves_idle_thread(): + client = FakeClient() + session = make_session(client) + original = client.request + def request(method, params=None, timeout=30): + response = original(method, params, timeout) + if method == "thread/read": + session.request_interrupt() + return response + client.request = request + result = session.compact_thread(turn_timeout=0.01) + assert result.interrupted and not result.should_retire and not session._closed + assert not session._interrupt_event.is_set() + assert not any(method == "thread/compact/start" for method, _ in client.requests) diff --git a/tests/agent/transports/test_codex_native_wire_hardening.py b/tests/agent/transports/test_codex_native_wire_hardening.py index f3d18293a11b8..172bb42e1006d 100644 --- a/tests/agent/transports/test_codex_native_wire_hardening.py +++ b/tests/agent/transports/test_codex_native_wire_hardening.py @@ -188,7 +188,14 @@ def send(obj, **kwargs): return original_send(obj, **kwargs) client._send = send - client.request = lambda *args, **kwargs: {"turn": {"id": "new-turn"}} + def request(method, *args, **kwargs): + if method == "thread/read": + return {"thread": {"id": "dummy-thread", "status": {"type": "idle"}, "turns": []}} + if method == "thread/compact/start": + return {} + return {"turn": {"id": "new-turn"}} + + client.request = request client._server_requests.put({"id": 42, "method": "item/permissions/requestApproval", "params": {"threadId": "dummy-thread", "turnId": "new-turn"}}) session = CodexAppServerSession() @@ -231,9 +238,11 @@ def run(): session.close() -def test_no_id_compaction_ack_retires_dummy_before_server_reply(dummy_codex): +def test_empty_compaction_ack_without_lifecycle_times_out_and_retires(dummy_codex): client = cas.CodexAppServerClient(codex_bin=dummy_codex) - client.request = lambda *args, **kwargs: {} + client.request = lambda method, *args, **kwargs: ( + {"thread": {"id": "dummy-thread", "status": {"type": "idle"}, "turns": []}} + if method == "thread/read" else {}) client._server_requests.put({"id": 42, "method": "item/permissions/requestApproval"}) session = CodexAppServerSession() session._client = client @@ -241,13 +250,13 @@ def test_no_id_compaction_ack_retires_dummy_before_server_reply(dummy_codex): try: result = session.compact_thread(turn_timeout=0.01) assert not result.completed and result.final_text == "" - assert result.should_retire and "correlat" in result.error + assert result.should_retire and "timed out" in result.error assert client._closed and session._closed client._proc.wait(timeout=2) assert not client.is_alive() assert session._active_turn_id is None assert not session._interrupt_event.is_set() - assert client._server_requests.qsize() == 1 + assert client._server_requests.qsize() == 0 finally: client.close(timeout=0) session.close() @@ -473,3 +482,53 @@ def send(index): assert len({frame["id"] for frame in frames if frame.get("method") == "ping"}) == 5 assert sorted((frame.get("params") or frame.get("result"))["tag"] for frame in frames[:-1]) == list(range(8)) assert all((frame.get("params") or frame.get("result"))["text"] == payload for frame in frames[:-1]) + + +@pytest.mark.parametrize("before_ack", [False, True]) +def test_supported_empty_ack_compaction_preserves_thread_over_real_wire(tmp_path, before_ack): + """The inert child speaks protocol lifecycle, including delayed prior events.""" + dummy = tmp_path / "compact-protocol-child" + dummy.write_text( + f"#!{sys.executable}\n" + "import json, sys\n" + "def emit(value): print(json.dumps(value), flush=True)\n" + "def note(method, turn, item=None):\n" + " params = {'threadId': 'canonical', 'turnId': turn}\n" + " if method.startswith('turn/'): params['turn'] = {'id': turn, 'status': 'completed' if method == 'turn/completed' else 'inProgress'}\n" + " if item: params['item'] = item\n" + " emit({'method': method, 'params': params})\n" + "for line in sys.stdin:\n" + " msg = json.loads(line)\n" + " if 'id' not in msg or 'method' not in msg: continue\n" + " method = msg['method']\n" + " if method == 'thread/read':\n" + " emit({'id': msg['id'], 'result': {'thread': {'id': 'canonical', 'status': {'type': 'idle'}, 'turns': [{'id': 'prior', 'status': 'completed'}]}}})\n" + " elif method == 'thread/compact/start':\n" + f" if not {before_ack!r}: emit({{'id': msg['id'], 'result': {{}}}})\n" + " note('turn/started', 'prior')\n" + " note('turn/completed', 'prior')\n" + " emit({'method': 'turn/started', 'params': {'threadId': 'foreign', 'turn': {'id': 'foreign-turn'}}})\n" + " note('turn/completed', 'compact')\n" + " note('turn/started', 'compact')\n" + " note('item/started', 'compact', {'type': 'contextCompaction', 'id': 'compact-item'})\n" + " note('item/completed', 'compact', {'type': 'contextCompaction', 'id': 'compact-item'})\n" + " note('turn/completed', 'compact')\n" + f" if {before_ack!r}: emit({{'id': msg['id'], 'result': {{}}}})\n" + " elif method == 'turn/start':\n" + " emit({'id': msg['id'], 'result': {'turn': {'id': 'next'}}})\n" + " note('turn/completed', 'next')\n" + " else: emit({'id': msg['id'], 'result': {}})\n", + encoding="utf-8", + ) + dummy.chmod(0o700) + client = cas.CodexAppServerClient(codex_bin=str(dummy)) + session = CodexAppServerSession() + session._client, session._thread_id = client, "canonical" + try: + compact = session.compact_thread(turn_timeout=1, notification_poll_timeout=0.001) + assert compact.completed and compact.compacted and compact.turn_id == "compact" + assert session._thread_id == "canonical" and not session._closed and client.is_alive() + next_turn = session.run_turn("continue", turn_timeout=1, notification_poll_timeout=0.001) + assert next_turn.completed and next_turn.thread_id == "canonical" + finally: + session.close() diff --git a/tests/gateway/test_native_partial_consumers.py b/tests/gateway/test_native_partial_consumers.py new file mode 100644 index 0000000000000..db8e6115c62f2 --- /dev/null +++ b/tests/gateway/test_native_partial_consumers.py @@ -0,0 +1,402 @@ +"""Native producer results reach real consumers offline, without becoming finals.""" + +import asyncio +import json +import sys +import types +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +import pytest +import pytest_asyncio + +from agent import codex_runtime +from agent.transports import codex_app_server_session as sessions +from gateway.config import Platform, PlatformConfig, StreamingConfig +from gateway.platforms import api_server +from gateway.run import _normalize_empty_agent_response, _should_clear_resume_pending_after_turn +from gateway.session import SessionSource +from hermes_state import SessionDB + + +DRAFT = "A native draft that must remain accessible." +ERROR = "native terminal incomplete" + + +@pytest_asyncio.fixture +async def inert_api_executor(monkeypatch): + """Exercise real API execution/result code while excluding worker scheduling. + + All collaborators are inert and SQLite is a temporary local store. Return + completed Futures so these consumer tests do not depend on executor wakeups; + gateway tests above separately retain the real threaded stream boundary. + """ + loop = asyncio.get_running_loop() + + def run_in_executor(executor, func, *args): + assert executor is None, "consumer test unexpectedly requested a custom executor" + future = loop.create_future() + try: + future.set_result(func(*args)) + except Exception as exc: + future.set_exception(exc) + return future + + monkeypatch.setattr(loop, "run_in_executor", run_in_executor) + yield + + +@pytest.fixture +def native(monkeypatch, tmp_path): + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + turn = sessions.TurnResult( + final_text=DRAFT, partial_text=DRAFT, completed=False, + error=ERROR, thread_id="native-thread", turn_id="native-turn", + projected_messages=[{"role": "assistant", "content": DRAFT}], + ) + memory, review, results = [], [], [] + + class NativeSession: + def __init__(self, **kwargs): + pass + + def matches_route(self, **kwargs): + return True + + def run_turn(self, **kwargs): + return turn + + def close(self): + pass + + monkeypatch.setattr(sessions, "CodexAppServerSession", NativeSession) + monkeypatch.setattr(codex_runtime, "_record_codex_app_server_usage", lambda *args: {}) + monkeypatch.setattr(codex_runtime, "_record_codex_app_server_compaction", lambda *args: None) + + class NativeAgent: + def __init__(self, **kwargs): + self.session_id = kwargs.get("session_id", "partial-session") + self.model = "openai-codex/gpt-test" + self.provider = "openai-codex" + self.api_mode = "codex_app_server" + self.session_cwd = str(tmp_path) + self.context_rebase_enabled = False + self._skill_nudge_interval = 1 + self._iters_since_skill = 0 + self.valid_tool_names = {"skill_manage"} + self.tools = [] + self.max_iterations = 500 + self._active_children = [] + self.stream_delta_callback = kwargs.get("stream_delta_callback") + self.emit_deltas = False + + def clear_interrupt(self): + self._interrupt_requested = False + + def _sync_external_memory_for_turn(self, **kwargs): + memory.append(kwargs) + + def _spawn_background_review(self, **kwargs): + review.append(kwargs) + + def run_conversation(self, user_message=None, conversation_history=None, **kwargs): + self._interrupt_requested = turn.interrupted + self._interrupt_message = None + if self.emit_deltas and self.stream_delta_callback: + self.stream_delta_callback(DRAFT) + messages = list(conversation_history or []) + messages.append({"role": "user", "content": user_message}) + result = codex_runtime.run_codex_app_server_turn( + self, user_message=user_message, original_user_message=user_message, + messages=messages, effective_task_id=self.session_id, + should_review_memory=True, + ) + results.append(result) + return result + + return SimpleNamespace(Agent=NativeAgent, turn=turn, results=results, memory=memory, review=review) + + +def _assert_incomplete(native): + assert native.results + for result in native.results: + assert result["final_response"] == "" + assert result["partial_response"] == DRAFT + assert result["completed"] is False + assert result["partial"] is True + assert result["error"] == ERROR + assert result["interrupted"] is native.turn.interrupted + assert native.memory == [] + assert native.review == [] + + +@pytest.mark.parametrize("interrupted", [False, True]) +@pytest.mark.parametrize("streaming", [False, True]) +def test_cli_chat_renders_labeled_native_draft(native, interrupted, streaming): + from tests.cli.test_cli_interrupt_ack_race import _make_cli + import cli as cli_module + + native.turn.interrupted = interrupted + cli = _make_cli() + cli.agent = native.Agent(session_id=cli.session_id) + cli.agent.emit_deltas = streaming + cli.agent.stream_delta_callback = cli._stream_delta + panels = [] + printed = [] + with patch.object(cli, "_ensure_runtime_credentials", return_value=True), \ + patch.object(cli, "_resolve_turn_agent_config", return_value={ + "signature": cli._active_agent_route_signature, + "model": None, "runtime": None, "request_overrides": None, + }), \ + patch.object(cli, "_init_agent", return_value=True), \ + patch.object(cli_module, "ChatConsole") as console, \ + patch.object(cli_module, "_cprint", side_effect=printed.append): + console.return_value.print.side_effect = panels.append + response = cli.chat("continue") + + assert DRAFT in response + assert "Partial response" in response + assert ERROR in response + if streaming: + assert sum(DRAFT in str(text) for text in printed) == 1 + assert sum("Partial response" in str(text) for text in printed) == 1 + assert not any(DRAFT in str(getattr(panel, "renderable", "")) for panel in panels) + else: + assert any(DRAFT in str(getattr(panel, "renderable", "")) for panel in panels) + assert cli._last_turn_interrupted is interrupted + _assert_incomplete(native) + + +@pytest.mark.parametrize("interrupted", [False, True]) +def test_quiet_single_query_main_exposes_draft_and_failure(native, monkeypatch, capsys, interrupted): + import cli as cli_module + import signal + + native.turn.interrupted = interrupted + cli = SimpleNamespace( + agent=native.Agent(), session_id="partial-session", conversation_history=[], + _claim_active_session=lambda *args, **kwargs: True, + _ensure_runtime_credentials=lambda: True, + _init_agent=lambda **kwargs: True, + _active_agent_route_signature=("inert",), + _resolve_turn_agent_config=lambda *args: { + "signature": ("inert",), "model": None, "runtime": None, + "request_overrides": None, + }, + ) + monkeypatch.setattr(cli_module, "HermesCLI", lambda **kwargs: cli) + monkeypatch.setattr(cli_module, "CLI_CONFIG", {"worktree": False}) + monkeypatch.setattr(cli_module.atexit, "register", lambda *args: None) + monkeypatch.setattr(signal, "signal", lambda *args: None) + monkeypatch.setattr(cli_module, "_finalize_single_query", lambda *args: None) + monkeypatch.setenv("HERMES_INTERACTIVE", "") + monkeypatch.setenv("HERMES_SINGLE_QUERY_SESSION", "") + monkeypatch.setenv("HERMES_KANBAN_GOAL_MODE", "0") + with pytest.raises(SystemExit) as exit_info: + cli_module.main(query="continue", quiet=True, toolsets="terminal") + stdout, stderr = capsys.readouterr() + assert stdout.strip() == DRAFT + assert "Partial response" in stderr and ERROR in stderr + assert exit_info.value.code == 1 + _assert_incomplete(native) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("interrupted", [False, True]) +async def test_gateway_actual_turn_preserves_native_partial(native, monkeypatch, tmp_path, streaming, interrupted): + from tests.gateway.test_stale_finalize_suppression import FinalizeCaptureAdapter, _make_runner + import gateway.run as gateway_run + + native.turn.interrupted = interrupted + adapter = FinalizeCaptureAdapter() + runner = _make_runner(adapter) + runner.config.streaming = StreamingConfig(enabled=streaming, edit_interval=0.01, buffer_threshold=1) + (tmp_path / "config.yaml").write_text(json.dumps({ + "display": {"tool_progress": "off", "interim_assistant_messages": False}, + "streaming": {"enabled": streaming, "edit_interval": 0.01, "buffer_threshold": 1}, + })) + monkeypatch.setattr(gateway_run, "_hermes_home", tmp_path) + monkeypatch.setattr(gateway_run, "_resolve_runtime_agent_kwargs", lambda: {}) + fake_agent_module = types.ModuleType("run_agent") + + class StreamingNativeAgent(native.Agent): + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.emit_deltas = streaming + + fake_agent_module.AIAgent = StreamingNativeAgent + monkeypatch.setitem(sys.modules, "run_agent", fake_agent_module) + source = SessionSource(platform=Platform.TELEGRAM, chat_id="partial-chat", chat_type="dm") + result = await asyncio.wait_for(runner._run_agent( + message="continue", context_prompt="", history=[], source=source, + session_id="partial-session", session_key="agent:main:telegram:dm:partial-chat", + ), timeout=5) + + assert result["final_response"] == "" + assert result["partial_response"] == DRAFT + assert result["completed"] is False + assert result["interrupted"] is interrupted + assert not _should_clear_resume_pending_after_turn(result) + delivery = _normalize_empty_agent_response(result, result["final_response"]) + assert "Partial response" in delivery and ERROR in delivery + if streaming: + assert result["partial_response_previewed"] is True + assert DRAFT not in delivery + payloads = [msg["content"] for msg in adapter.sent] + [msg["content"] for msg in adapter.edits] + assert any(DRAFT in payload for payload in payloads) + assert sum(DRAFT in msg["content"] for msg in adapter.sent) == 1 + else: + assert DRAFT in delivery + assert adapter.sent == [] + _assert_incomplete(native) + + +class CaptureStream: + def __init__(self, *, status=200, headers=None): + self.status = status + self.headers = headers or {} + self.chunks = [] + + async def prepare(self, request): + pass + + async def write(self, data): + self.chunks.append(data) + + async def write_eof(self): + pass + + +def _events(stream): + events = [] + for frame in b"".join(stream.chunks).decode().split("\n\n"): + lines = frame.splitlines() + data = next((line[6:] for line in lines if line.startswith("data: ")), None) + if data and data != "[DONE]": + name = next((line[7:] for line in lines if line.startswith("event: ")), None) + events.append((name, json.loads(data))) + return events + + +@pytest.mark.asyncio +@pytest.mark.parametrize("endpoint", ["chat", "responses", "session"]) +@pytest.mark.parametrize("interrupted", [False, True]) +@pytest.mark.parametrize("mode", ["batch", "stream_result_only", "stream_deltas"]) +async def test_api_real_handlers_preserve_native_partial(native, monkeypatch, tmp_path, inert_api_executor, endpoint, interrupted, mode): + native.turn.interrupted = interrupted + adapter = api_server.APIServerAdapter(PlatformConfig(enabled=True)) + db = SessionDB(tmp_path / "consumer-state.db") + db.create_session("partial-session", "api_server") + adapter._session_db = db + + def create_agent(**kwargs): + agent = native.Agent(**kwargs) + agent.emit_deltas = mode == "stream_deltas" + return agent + + monkeypatch.setattr(adapter, "_create_agent", create_agent) + monkeypatch.setattr(api_server, "_publish_turn_process_ownership", lambda *args: None) + monkeypatch.setattr(api_server, "_clear_turn_process_ownership", lambda *args: None) + monkeypatch.setattr(api_server.web, "StreamResponse", CaptureStream) + streaming = mode != "batch" + body = {"stream": streaming, "model": "hermes-agent"} + if endpoint == "chat": + body["messages"] = [{"role": "user", "content": "continue"}] + handler = adapter._handle_chat_completions + elif endpoint == "responses": + body["input"] = "continue" + handler = adapter._handle_responses + else: + body["message"] = "continue" + handler = adapter._handle_session_chat_stream if streaming else adapter._handle_session_chat + request = SimpleNamespace( + headers={}, match_info={"session_id": "partial-session"}, + json=AsyncMock(return_value=body), query={}, + ) + try: + response = await asyncio.wait_for(handler(request), timeout=5) + assert response.status == 200 + if not streaming: + payload = json.loads(response.text) + if endpoint == "chat": + assert payload["choices"][0]["message"]["content"] == DRAFT + assert payload["choices"][0]["finish_reason"] == "error" + outcome = payload["hermes"] + elif endpoint == "responses": + assert payload["status"] == "incomplete" + assert payload["output"][-1]["content"][0]["text"] == DRAFT + assert payload["output"][-1]["status"] == "incomplete" + outcome = payload["hermes"] + assert adapter._response_store.get(payload["id"])["response"]["status"] == "incomplete" + else: + assert payload["message"]["content"] == DRAFT + outcome = payload + else: + events = _events(response) + if endpoint == "chat": + content = "".join(event["choices"][0]["delta"].get("content", "") for _, event in events) + assert content == DRAFT + assert events[-1][1]["choices"][0]["finish_reason"] == "error" + outcome = events[-1][1]["hermes"] + elif endpoint == "responses": + deltas = [event["delta"] for name, event in events if name == "response.output_text.delta"] + assert "".join(deltas) == DRAFT + terminal = events[-1][1]["response"] + assert terminal["status"] == "incomplete" + assert terminal["output"][-1]["content"][0]["text"] == DRAFT + assert terminal["output"][-1]["status"] == "incomplete" + assert not any(name == "response.completed" for name, _ in events) + outcome = terminal["hermes"] + else: + terminal = next(event for name, event in events if name == "assistant.completed") + assert terminal["content"] == DRAFT + outcome = terminal + terminal_name = "run.cancelled" if interrupted else "run.failed" + run = next(event for name, event in events if name == terminal_name) + assert run["completed"] is False + assert adapter._run_statuses[run["run_id"]]["status"] != "completed" + assert adapter._run_statuses[run["run_id"]]["last_event"] == terminal_name + assert not any(name == "run.completed" for name, _ in events) + assert outcome["completed"] is False + assert outcome["partial"] is True + assert outcome["interrupted"] is interrupted + assert outcome["error"] == ERROR + _assert_incomplete(native) + finally: + db.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("interrupted", [False, True]) +async def test_runs_actual_execution_preserves_draft_without_completed_status(native, monkeypatch, inert_api_executor, interrupted): + native.turn.interrupted = interrupted + adapter = api_server.APIServerAdapter(PlatformConfig(enabled=True)) + monkeypatch.setattr(adapter, "_create_agent", lambda **kwargs: native.Agent(**kwargs)) + monkeypatch.setattr(api_server, "_publish_turn_process_ownership", lambda *args: None) + monkeypatch.setattr(api_server, "_clear_turn_process_ownership", lambda *args: None) + request = SimpleNamespace(headers={}, json=AsyncMock(return_value={ + "input": "continue", "session_id": "partial-session", + })) + response = await asyncio.wait_for(adapter._handle_runs(request), timeout=5) + assert response.status == 202 + run_id = json.loads(response.text)["run_id"] + task = adapter._active_run_tasks.get(run_id) + if task is not None: + await asyncio.wait_for(task, timeout=5) + outcome = adapter._run_statuses[run_id] + assert outcome["status"] == ("cancelled" if interrupted else "failed") + assert outcome["output"] == DRAFT + assert outcome["completed"] is False + assert outcome["partial"] is True + assert outcome["interrupted"] is interrupted + assert outcome["error"] == ERROR + events = [] + queue = adapter._run_streams[run_id] + while not queue.empty(): + event = queue.get_nowait() + if event is not None: + events.append(event) + assert not any(event.get("event") == "run.completed" for event in events) + assert any(event.get("output") == DRAFT for event in events) + _assert_incomplete(native) diff --git a/tests/run_agent/test_codex_native_runtime_hardening.py b/tests/run_agent/test_codex_native_runtime_hardening.py index 0d85566994209..65d4c6e90f240 100644 --- a/tests/run_agent/test_codex_native_runtime_hardening.py +++ b/tests/run_agent/test_codex_native_runtime_hardening.py @@ -97,3 +97,204 @@ def restore(): sanitize_surrogates=noop, summarize_user_message_for_log=noop, set_session_context=noop, set_current_write_origin=noop, ra=SimpleNamespace()) assert restored == [] + + +@pytest.fixture +def native_continuity(monkeypatch): + """Drive the actual runtime and session with an inert authoritative thread.""" + transports = [] + + class NativeTransport: + def __init__(self, **kwargs): + self.requests, self.notes, self.history = [], [], [] + self.closed = False + self.thread_id = f"native-{len(transports) + 1}" + self.reject_model = None + transports.append(self) + + def initialize(self, **kwargs): + return {} + + def request(self, method, params=None, timeout=30): + self.requests.append((method, dict(params or {}))) + if method == "thread/start": + return {"thread": {"id": self.thread_id}, "model": params["model"], + "modelProvider": params["modelProvider"]} + if method == "turn/start": + assert params["threadId"] == self.thread_id + if params.get("model") == self.reject_model: + raise sessions.CodexAppServerError(code=-32602, message="unsupported turn model") + self.history.append((params["model"], params["input"][0]["text"])) + turn_id = f"turn-{len(self.history)}" + # The response depends on retained native history, not Hermes' projection. + text = " | ".join(value for _, value in self.history) + self.notes.extend([ + {"method": "item/completed", "params": { + "threadId": self.thread_id, "turnId": turn_id, + "item": {"type": "agentMessage", "id": f"item-{turn_id}", "text": text}}}, + {"method": "turn/completed", "params": { + "threadId": self.thread_id, "turn": {"id": turn_id, "status": "completed"}}}, + ]) + return {"turn": {"id": turn_id}} + raise AssertionError(f"unexpected RPC {method}") + + def take_notification(self, timeout=0): + return self.notes.pop(0) if self.notes else None + + def take_server_request(self, timeout=0): + return None + + def stderr_tail(self, count=20): + return [] + + def is_alive(self): + return not self.closed + + def close(self): + self.closed = True + + monkeypatch.setattr(sessions, "CodexAppServerClient", NativeTransport) + monkeypatch.setattr(codex_runtime, "make_codex_app_server_event_bridge", lambda agent: None) + monkeypatch.setattr(codex_runtime, "_record_codex_app_server_usage", lambda *args: {}) + monkeypatch.setattr(codex_runtime, "_record_codex_app_server_compaction", lambda *args: None) + agent = SimpleNamespace(model="openai-codex/gpt-primary", provider="openai-codex", + api_mode="codex_app_server", session_cwd="/tmp", context_rebase_enabled=False, + _skill_nudge_interval=0, _iters_since_skill=0, valid_tool_names=set()) + messages = [] + + def run(text): + messages.append({"role": "user", "content": text}) + return codex_runtime.run_codex_app_server_turn( + agent, user_message=text, original_user_message=text, messages=messages, + effective_task_id="offline-native-continuity") + + return SimpleNamespace(agent=agent, transports=transports, run=run) + + +def test_native_model_switch_and_once_rollback_keep_authoritative_history(native_continuity): + state = native_continuity + first = state.run("remember this") + canonical = state.agent._codex_session + state.agent.model = "openai/gpt-once" + second = state.run("use another model once") + # This is the caller's --once rollback of the selected route. The protocol + # setting is sticky, so the next turn must explicitly restore primary. + state.agent.model = "openai-codex/gpt-primary" + third = state.run("return to primary") + assert first["completed"] and second["completed"] and third["completed"] + assert state.agent._codex_session is canonical and not canonical._closed + assert len(state.transports) == 1 + client = state.transports[0] + assert not client.closed + assert sum(method == "thread/start" for method, _ in client.requests) == 1 + starts = [params for method, params in client.requests if method == "turn/start"] + assert [params["model"] for params in starts] == ["gpt-primary", "gpt-once", "gpt-primary"] + assert {params["threadId"] for params in starts} == {first["codex_thread_id"]} + assert third["final_response"] == "remember this | use another model once | return to primary" + canonical.close() + + +def test_native_provider_change_refuses_before_closing_or_dispatch(native_continuity): + state = native_continuity + state.run("remember this") + canonical = state.agent._codex_session + count = len(state.transports[0].requests) + state.agent.provider = "custom" + refused = state.run("change provider") + assert not refused["completed"] and "Reset" in refused["error"] + assert state.agent._codex_session is canonical and not canonical._closed + assert not state.transports[0].closed and len(state.transports[0].requests) == count + canonical.close() + + +def test_unsupported_native_turn_model_is_visible_without_history_retirement(native_continuity): + state = native_continuity + state.run("remember this") + canonical = state.agent._codex_session + state.transports[0].reject_model = "unsupported" + state.agent.model = "unsupported" + refused = state.run("must be refused") + assert not refused["completed"] and "unsupported turn model" in refused["error"] + assert state.agent._codex_session is canonical and not canonical._closed + assert state.transports[0].history == [("gpt-primary", "remember this")] + state.agent.model = "gpt-primary" + next_turn = state.run("try primary again") + assert next_turn["completed"] and next_turn["final_response"] == "remember this | try primary again" + canonical.close() + + +def test_explicit_native_reset_creates_a_distinct_thread(native_continuity): + state = native_continuity + first = state.run("old history") + previous = state.agent._codex_session + # Explicit reset ownership already closes and clears the native session. + previous.close() + state.agent._codex_session = None + state.agent.model = "gpt-new" + next_turn = state.run("new conversation") + assert previous._closed and state.transports[0].closed + assert next_turn["completed"] and next_turn["codex_thread_id"] != first["codex_thread_id"] + assert len(state.transports) == 2 and next_turn["final_response"] == "new conversation" + state.agent._codex_session.close() + + +def test_tui_once_restore_owner_preserves_native_thread_and_history(native_continuity, monkeypatch): + """Real TUI snapshot/restore and AIAgent.switch_model; clients stay inert.""" + from types import MethodType + from agent import agent_runtime_helpers as runtime_helpers + from run_agent import AIAgent + from tui_gateway import server + from hermes_cli import config + + state = native_continuity + agent = state.agent + built = [] + agent.switch_model = MethodType(AIAgent.switch_model, agent) + agent.api_key = "synthetic-native-key" + agent.base_url = "https://offline.invalid/v1" + agent.requested_provider = agent.provider + agent.client = SimpleNamespace() + agent._client_kwargs = {} + agent._credential_pool = object() + agent.context_compressor = None + agent._fallback_chain = [] + agent._primary_runtime = None + agent.quiet_mode = True + agent._read_reasoning_echo_from_config = lambda: False + agent._apply_client_headers_for_base_url = lambda *args: None + agent._ensure_lmstudio_runtime_loaded = lambda *args: None + agent._lmstudio_load_was_unverified = lambda *args: False + agent._effective_lmstudio_context_length = lambda *args: None + agent._anthropic_prompt_cache_policy = lambda **kwargs: (False, False) + + def create_client(kwargs, **unused): + client = SimpleNamespace(kwargs=dict(kwargs)) + built.append(client) + return client + + agent._create_openai_client = create_client + monkeypatch.setattr(runtime_helpers, "get_provider_request_timeout", lambda *args: None) + monkeypatch.setattr(runtime_helpers, "sync_credential_pool_entry_id", lambda *args: None) + monkeypatch.setattr(config, "load_config_readonly", lambda *args, **kwargs: {}) + monkeypatch.setattr(config, "load_config", lambda *args, **kwargs: {}) + # No SDK constructor, credential read, metadata probe or external provider. + first = state.run("retain native history") + canonical = agent._codex_session + snapshot = server._snapshot_agent_model_runtime(agent) + agent.switch_model( + new_model="gpt-once", new_provider="openai-codex", api_key=agent.api_key, + base_url=agent.base_url, api_mode="codex_app_server", + ) + temporary = state.run("temporary selection") + server._restore_agent_model_runtime(agent, snapshot) + assert agent.model == "openai-codex/gpt-primary" + assert agent.api_mode == "codex_app_server" and agent._codex_session is canonical + restored = state.run("after actual TUI restore") + assert first["completed"] and temporary["completed"] and restored["completed"] + assert len(built) == 2 and len(state.transports) == 1 + starts = [params for method, params in state.transports[0].requests if method == "turn/start"] + assert [params["model"] for params in starts] == ["gpt-primary", "gpt-once", "gpt-primary"] + assert {params["threadId"] for params in starts} == {first["codex_thread_id"]} + assert restored["final_response"] == "retain native history | temporary selection | after actual TUI restore" + assert not canonical._closed and not state.transports[0].closed + canonical.close() diff --git a/tests/tui_gateway/test_bot_capability_refresh.py b/tests/tui_gateway/test_bot_capability_refresh.py index e5f09b1c56b5f..a336363c31508 100644 --- a/tests/tui_gateway/test_bot_capability_refresh.py +++ b/tests/tui_gateway/test_bot_capability_refresh.py @@ -32,6 +32,9 @@ def release_clients(self): self.client_closes.append(self.client) self.client = None + def retire_local_resources(self, *, preserve_agent=None): + self.release_clients() + def clear_interrupt(self): pass diff --git a/tests/tui_gateway/test_bot_local_retirement.py b/tests/tui_gateway/test_bot_local_retirement.py new file mode 100644 index 0000000000000..304dd6e318a1d --- /dev/null +++ b/tests/tui_gateway/test_bot_local_retirement.py @@ -0,0 +1,447 @@ +"""Instance disposal through real capability refresh and memory manager. + +No provider, terminal, browser, or child process is started. SQLite ownership +and the refresh path are real; counted resources implement the local lifecycle. +""" +from types import SimpleNamespace + +import pytest + +from agent.context_engine import ContextEngine +from agent.memory_manager import MemoryManager +from agent.memory_provider import MemoryProvider +from run_agent import AIAgent +from tests.tui_gateway.test_bot_capability_refresh import FakeAgent, KEY, SID, env # noqa: F401 +from tests.tui_gateway.test_bot_native_retirement import native_session +from tests.tui_gateway.test_prompt_recovery_contract import turn_env # noqa: F401 +from tui_gateway import server + + +class CountedProvider(MemoryProvider): + name = "counted-local" + + def __init__(self): + self.shutdowns = 0 + self.ends = [] + self.lifecycle = [] + self.live = False + + def is_available(self): + return True + + def initialize(self, session_id, **kwargs): + self.session_id = session_id + self.live = True + + def get_tool_schemas(self): + return [] + + def on_session_end(self, messages): + self.lifecycle.append("end") + self.ends.append(messages) + + def shutdown(self): + self.lifecycle.append("shutdown") + self.shutdowns += 1 + self.live = False + + def prefetch(self, query, *, session_id=""): + assert self.live + return "survivor memory" + + +class CountedEngine(ContextEngine): + name = "counted-local" + + def __init__(self): + self.shutdowns = 0 + self.ends = [] + self.lifecycle = [] + self.live = True + + def update_from_response(self, usage): + pass + + def should_compress(self, prompt_tokens=None): + assert self.live + return False + + def compress(self, messages, **kwargs): + pytest.fail("retirement attempted compression") + + def on_session_end(self, session_id, messages): + # A logical end could mutate canonical state for this same session ID. + self.lifecycle.append("end") + self.ends.append((session_id, messages)) + + def shutdown(self): + self.lifecycle.append("shutdown") + self.shutdowns += 1 + self.live = False + + +class ResourceAgent(FakeAgent): + shutdown_memory_provider = AIAgent.shutdown_memory_provider + + def retire_local_resources(self, *, preserve_agent=None): + return AIAgent.retire_local_resources(self, preserve_agent=preserve_agent) + + def attach_resources(self): + self.provider_resource = CountedProvider() + self._memory_manager = MemoryManager() + self._memory_manager.add_provider(self.provider_resource) + self._memory_manager.initialize_all(self.session_id) + self.context_compressor = CountedEngine() + + +@pytest.fixture +def resources(env, monkeypatch): + env.old.__class__ = ResourceAgent + env.old.attach_resources() + task_state = SimpleNamespace(terminal=object(), browser=object(), processes=object()) + env.old._session_messages = [{"role": "assistant", "content": "prior turn"}] + env.old.task_state = task_state + + def construct(**kwargs): + new = ResourceAgent(**kwargs) + new.attach_resources() + new.task_state = task_state + env.built.append(new) + return new + + monkeypatch.setattr("run_agent.AIAgent", construct) + env.resource_construct = construct + + def forbidden(*args, **kwargs): + pytest.fail("local retirement destroyed logical session/task resources") + + with monkeypatch.context() as guard: + guard.setattr(AIAgent, "close", forbidden) + guard.setattr(env.target, "end_session", forbidden) + guard.setattr(env.target, "close", forbidden) + guard.setattr(env.launch, "close", forbidden) + guard.setattr("run_agent.cleanup_vm", forbidden) + guard.setattr("run_agent.cleanup_browser", forbidden) + guard.setattr("tools.process_registry.process_registry.kill_all", forbidden) + yield env + + +def assert_local_state(agent, shutdowns): + assert agent.provider_resource.shutdowns == shutdowns + assert agent.context_compressor.shutdowns == shutdowns + assert agent.provider_resource.live is (shutdowns == 0) + assert agent.context_compressor.live is (shutdowns == 0) + assert agent.provider_resource.ends == [] + assert agent.context_compressor.ends == [] + + +def assert_canonical_store(owner): + current = owner.session["agent"] + assert current._session_db is owner.target + assert current.session_id == KEY + assert current._owns_session_db + assert owner.target.get_session(KEY)["ended_at"] is None + assert owner.launch.get_session(KEY) is None + assert owner.launch.get_messages("launch-sentinel")[0]["content"] == "untouched" + assert current.task_state is owner.old.task_state + + +def test_success_and_repeated_refresh_dispose_only_superseded_instances(resources): + owner = resources + server._sync_bot_capabilities(SID, owner.session) + first = owner.session["agent"] + assert first is not owner.old + assert_local_state(owner.old, 1) + assert_local_state(first, 0) + assert len(owner.old.client_closes) == 1 + owner.old.retire_local_resources() + owner.old.shutdown_memory_provider() + server._sync_bot_capabilities(SID, owner.session) + assert owner.session["agent"] is first + assert_local_state(owner.old, 1) + assert len(owner.built) == 1 + + owner.monkeypatch.setattr("tools.bot_mode_probe.capability_fingerprint", lambda home: "later-caps") + server._sync_bot_capabilities(SID, owner.session) + assert_local_state(first, 1) + assert_local_state(owner.session["agent"], 0) + assert len(first.client_closes) == 1 + assert_canonical_store(owner) + + +@pytest.mark.parametrize("rejection", ["runtime", "transfer", "cancel"]) +def test_rejected_candidate_is_disposed_once_without_ending_predecessor(resources, rejection): + owner = resources + + def rejected(**kwargs): + new = owner.resource_construct(**kwargs) + if rejection == "runtime": + new.model = "wrong-model" + elif rejection == "cancel": + owner.session["_turn_cancel_requested"] = True + return new + + owner.monkeypatch.setattr("run_agent.AIAgent", rejected) + if rejection == "transfer": + owner.monkeypatch.setattr(server, "_transfer_db_to_agent", lambda *args: False) + if rejection == "cancel": + with pytest.raises(RuntimeError, match="BOT_CAPABILITY_OWNER_CHANGED"): + server._sync_bot_capabilities(SID, owner.session) + else: + server._sync_bot_capabilities(SID, owner.session) + candidate = owner.built[0] + assert owner.session["agent"] is owner.old + assert_local_state(owner.old, 0) + assert_local_state(candidate, 1) + candidate.retire_local_resources() + assert_local_state(candidate, 1) + assert len(candidate.client_closes) == 1 + assert not candidate._owns_session_db + assert owner.session["bot_caps_seen"] == "old-caps" + assert_canonical_store(owner) + + +def test_native_retirement_failure_still_disposes_other_instance_resources(resources): + owner = resources + native, client = native_session() + owner.old._codex_session = native + + def fail(): + raise RuntimeError("inert native retirement failure") + + owner.monkeypatch.setattr(native, "close", fail) + server._sync_bot_capabilities(SID, owner.session) + assert owner.session["agent"] is not owner.old + assert owner.session["_bot_native_retirement"]["status"] == "failed" + assert client.closes == 0 + assert_local_state(owner.old, 1) + assert_local_state(owner.session["agent"], 0) + assert_canonical_store(owner) + + +def test_constructor_failure_preserves_predecessor_and_allows_retry(resources): + owner = resources + + def fail(**kwargs): + raise RuntimeError("inert constructor failure before allocation") + + owner.monkeypatch.setattr("run_agent.AIAgent", fail) + server._sync_bot_capabilities(SID, owner.session) + assert owner.built == [] + assert owner.session["agent"] is owner.old + assert_local_state(owner.old, 0) + assert_canonical_store(owner) + owner.monkeypatch.setattr("run_agent.AIAgent", owner.resource_construct) + server._sync_bot_capabilities(SID, owner.session) + assert_local_state(owner.old, 1) + assert_local_state(owner.session["agent"], 0) + + +def test_real_session_boundary_delivers_end_before_local_shutdown(resources): + current = resources.old + messages = current._session_messages + current.shutdown_memory_provider(messages) + assert current.provider_resource.ends == [messages] + assert current.context_compressor.ends == [(KEY, messages)] + assert current.provider_resource.shutdowns == current.context_compressor.shutdowns == 1 + assert current.provider_resource.lifecycle == ["end", "shutdown"] + assert current.context_compressor.lifecycle == ["end", "shutdown"] + current.shutdown_memory_provider(messages) + current.retire_local_resources() + assert current.provider_resource.ends == [messages] + assert current.context_compressor.ends == [(KEY, messages)] + assert current.provider_resource.shutdowns == current.context_compressor.shutdowns == 1 + + +def test_client_release_failure_does_not_skip_local_provider_or_engine_shutdown(resources): + current = resources.old + + def fail(): + raise RuntimeError("inert client release failure") + + current.release_clients = fail + with pytest.raises(RuntimeError, match="client release failure"): + current.retire_local_resources() + assert_local_state(current, 1) + current.retire_local_resources() + assert_local_state(current, 1) + + +@pytest.mark.parametrize("allocation", ["none", "clients", "providers"]) +def test_actual_constructor_failure_reclaims_partial_local_resources(resources, allocation): + owner = resources + allocated = [] + client_retirements = [] + failure = RuntimeError("inert initializer failure") + + def initialize(agent, **kwargs): + allocated.append(agent) + agent.session_id = KEY + agent._session_db = owner.target + agent._owns_session_db = False + if allocation in {"clients", "providers"}: + agent.client = object() + if allocation == "providers": + ResourceAgent.attach_resources(agent) + raise failure + + def retire_client(agent, client, *, reason): + client_retirements.append(client) + + owner.monkeypatch.setattr("agent.agent_init.init_agent", initialize) + owner.monkeypatch.setattr(AIAgent, "_retire_shared_openai_client", retire_client) + # The actual wrapper and local disposal execute; provider initialization + # alone is inert, with no generic provider resolution or transport spawn. + with pytest.raises(RuntimeError) as raised: + AIAgent(session_db=owner.target, session_id=KEY) + assert raised.value is failure + failed = allocated[0] + assert failed._local_resources_retired + assert failed._memory_provider_shutdown + assert len(client_retirements) == (0 if allocation == "none" else 1) + if allocation == "providers": + assert_local_state(failed, 1) + failed.retire_local_resources() + assert len(client_retirements) == (0 if allocation == "none" else 1) + assert_canonical_store(owner) + + +def test_constructor_cleanup_failure_preserves_original_exception(resources): + owner = resources + failure = ValueError("original initializer failure") + + def initialize(agent, **kwargs): + raise failure + + def cleanup_fail(agent, *, preserve_agent=None): + raise RuntimeError("cleanup failure") + + owner.monkeypatch.setattr("agent.agent_init.init_agent", initialize) + owner.monkeypatch.setattr(AIAgent, "retire_local_resources", cleanup_fail) + with pytest.raises(ValueError) as raised: + AIAgent(session_db=owner.target, session_id=KEY) + assert raised.value is failure + assert_local_state(owner.old, 0) + assert_canonical_store(owner) + + +@pytest.mark.parametrize("sharing", ["provider", "engine", "manager"]) +@pytest.mark.parametrize("rejected", [False, True]) +def test_shared_instances_remain_usable_for_the_live_survivor(resources, sharing, rejected): + owner = resources + old = owner.old + + def construct(**kwargs): + new = ResourceAgent(**kwargs) + if sharing == "manager": + new._memory_manager = old._memory_manager + new.provider_resource = old.provider_resource + else: + new.provider_resource = old.provider_resource if sharing == "provider" else CountedProvider() + new._memory_manager = MemoryManager() + new._memory_manager.add_provider(new.provider_resource) + if sharing != "provider": + new._memory_manager.initialize_all(new.session_id) + new.context_compressor = old.context_compressor if sharing == "engine" else CountedEngine() + new.task_state = old.task_state + if rejected: + new.model = "wrong-model" + owner.built.append(new) + return new + + owner.monkeypatch.setattr("run_agent.AIAgent", construct) + server._sync_bot_capabilities(SID, owner.session) + new = owner.built[0] + live, retired = (old, new) if rejected else (new, old) + assert owner.session["agent"] is live + assert live.provider_resource.live + assert live.context_compressor.live + assert live.provider_resource.shutdowns == live.context_compressor.shutdowns == 0 + assert live.provider_resource.prefetch("next turn", session_id=KEY) == "survivor memory" + assert live.context_compressor.should_compress() is False + assert not live._memory_manager._shutting_down + if sharing == "engine": + assert retired.context_compressor is live.context_compressor + assert retired.provider_resource.shutdowns == 1 + else: + assert retired.provider_resource is live.provider_resource + assert retired.context_compressor.shutdowns == 1 + if sharing == "manager": + assert retired._memory_manager is live._memory_manager + else: + # A separate retired manager drains even when one provider is borrowed. + assert retired._memory_manager._shutting_down + assert retired._memory_manager.shutdown_drain_state["status"] == "drained" + assert live.provider_resource.ends == [] + assert live.context_compressor.ends == [] + retired.retire_local_resources(preserve_agent=live) + assert live.provider_resource.shutdowns == live.context_compressor.shutdowns == 0 + assert_canonical_store(owner) + + +@pytest.mark.parametrize("sharing", ["provider", "engine", "manager"]) +@pytest.mark.parametrize("path", ["sync", "wrapper"]) +def test_failed_constructor_preserves_predecessor_borrowed_instances(resources, sharing, path): + owner = resources + old = owner.old + allocated, client_retirements = [], [] + failure = RuntimeError("inert borrowed-resource initializer failure") + + def initialize(agent, **kwargs): + # The private cleanup context belongs to the wrapper, not init_agent. + assert "_resource_preserve_agent" not in kwargs + allocated.append(agent) + agent.session_id = kwargs["session_id"] + agent._session_db = kwargs["session_db"] + agent._owns_session_db = False + agent.client = object() + if sharing == "manager": + agent._memory_manager = old._memory_manager + agent.provider_resource = old.provider_resource + else: + agent.provider_resource = old.provider_resource if sharing == "provider" else CountedProvider() + agent._memory_manager = MemoryManager() + agent._memory_manager.add_provider(agent.provider_resource) + if sharing != "provider": + agent._memory_manager.initialize_all(agent.session_id) + agent.context_compressor = old.context_compressor if sharing == "engine" else CountedEngine() + raise failure + + def retire_client(agent, client, *, reason): + client_retirements.append(client) + + owner.monkeypatch.setattr("run_agent.AIAgent", AIAgent) + owner.monkeypatch.setattr("agent.agent_init.init_agent", initialize) + owner.monkeypatch.setattr(AIAgent, "_retire_shared_openai_client", retire_client) + if path == "sync": + # Real refresh -> real factory -> actual constructor wrapper. + server._sync_bot_capabilities(SID, owner.session) + else: + with pytest.raises(RuntimeError) as raised: + AIAgent(session_db=owner.target, session_id=KEY, + _resource_preserve_agent=old) + assert raised.value is failure + assert len(allocated) == 1 + failed = allocated[0] + assert failed._local_resources_retired and failed._memory_provider_shutdown + assert client_retirements and len(client_retirements) == 1 + assert owner.session["agent"] is old + assert owner.session["bot_caps_seen"] == "old-caps" + assert_local_state(old, 0) + assert old.provider_resource.prefetch("next turn", session_id=KEY) == "survivor memory" + assert old.context_compressor.should_compress() is False + assert not old._memory_manager._shutting_down + if sharing == "engine": + assert failed.context_compressor is old.context_compressor + assert failed.provider_resource.shutdowns == 1 + else: + assert failed.provider_resource is old.provider_resource + assert failed.context_compressor.shutdowns == 1 + if sharing == "manager": + assert failed._memory_manager is old._memory_manager + else: + assert failed._memory_manager.shutdown_drain_state["status"] == "drained" + failed.retire_local_resources(preserve_agent=old) + assert len(client_retirements) == 1 + assert_local_state(old, 0) + assert_canonical_store(owner) diff --git a/tests/tui_gateway/test_bot_native_retirement.py b/tests/tui_gateway/test_bot_native_retirement.py index ac9c1e1451d3b..f872902a67e6d 100644 --- a/tests/tui_gateway/test_bot_native_retirement.py +++ b/tests/tui_gateway/test_bot_native_retirement.py @@ -83,7 +83,6 @@ def forbidden(*a, **k): monkeypatch.setattr('tools.bot_mode_probe.capability_fingerprint', lambda home: 'next-caps') with monkeypatch.context() as guard: guard.setattr(REAL_AGENT, 'close', forbidden) - guard.setattr(REAL_AGENT, 'shutdown_memory_provider', forbidden) guard.setattr(env.target, 'close', forbidden) guard.setattr(env.target, 'end_session', forbidden) guard.setattr(env.launch, 'close', forbidden) diff --git a/tests/tui_gateway/test_prompt_recovery_contract.py b/tests/tui_gateway/test_prompt_recovery_contract.py index 5eb13213e878f..317940e1c9d6f 100644 --- a/tests/tui_gateway/test_prompt_recovery_contract.py +++ b/tests/tui_gateway/test_prompt_recovery_contract.py @@ -178,6 +178,67 @@ def test_returned_provider_error_is_terminal_and_replayable(emits, turn_env): assert session["running"] is False +@pytest.mark.parametrize("streamed", [False, True]) +@pytest.mark.parametrize("interrupted", [False, True]) +def test_native_partial_result_remains_visible_without_success( + monkeypatch, emits, turn_env, streamed, interrupted, +): + from agent import codex_runtime + from agent.transports.codex_app_server_session import TurnResult + + draft = "A draft with unfinished reasoning" + memory, reviews, results = [], [], [] + turn = TurnResult(partial_text=draft, error="native turn stopped", + interrupted=interrupted, thread_id="thread", turn_id="turn") + native_session = types.SimpleNamespace( + matches_route=lambda **kwargs: True, + run_turn=lambda **kwargs: turn, + ) + agent = types.SimpleNamespace( + session_id="prompt-recovery-session", model="gpt-test", provider="openai-codex", + api_mode="codex_app_server", context_rebase_enabled=False, + _codex_session=native_session, _interrupt_requested=interrupted, + _skill_nudge_interval=1, _iters_since_skill=0, valid_tool_names={"skill_manage"}, + clear_interrupt=lambda: None, + _sync_external_memory_for_turn=lambda **kwargs: memory.append(kwargs), + _spawn_background_review=lambda **kwargs: reviews.append(kwargs), + ) + monkeypatch.setattr(codex_runtime, "_record_codex_app_server_usage", lambda *args: {}) + monkeypatch.setattr(codex_runtime, "_record_codex_app_server_compaction", lambda *args: None) + + def run(message, stream_callback=None, **kwargs): + if streamed and stream_callback: + stream_callback(draft) + result = codex_runtime.run_codex_app_server_turn( + agent, user_message=message, original_user_message=message, + messages=[], effective_task_id="offline", should_review_memory=True, + ) + results.append(result) + return result + + agent.run_conversation = run + session = _session(agent=agent, running=True) + server._start_inflight_turn(session, "do the thing") + server._run_prompt_submit("rid", "sid", session, "do the thing") + + completes = _events(emits, "message.complete") + assert len(completes) == 1 + assert completes[0]["text"] == draft + assert completes[0]["partial"] is True + assert completes[0]["completed"] is False + assert completes[0]["error"] == "native turn stopped" + assert completes[0]["status"] == ("interrupted" if interrupted else "error") + assert results[0]["final_response"] == "" + assert results[0]["completed"] is False + assert memory == reviews == [] + assert not session.get("model_verified_for") + if not interrupted: + assert completes[0]["error"] == "native turn stopped" + snapshot = server._inflight_snapshot(session) + assert snapshot["assistant"] == draft + assert snapshot["error"] == "native turn stopped" + + def test_exception_restores_agent_transcript_and_retains_partial(emits, turn_env): def _boom(message, stream_callback=None, **kwargs): if stream_callback is not None: diff --git a/tui_gateway/server.py b/tui_gateway/server.py index 21ba8fa242f6e..e06205106602b 100644 --- a/tui_gateway/server.py +++ b/tui_gateway/server.py @@ -7175,6 +7175,7 @@ def require_owner(): _preserve_default_reasoning=reasoning is None, service_tier_override=tier, platform_override=_session_source(session), + _resource_preserve_agent=agent, ) finally: _clear_session_context(tokens) @@ -7232,15 +7233,15 @@ def require_owner(): finally: retired = agent if published else new_agent if retired is not None and (published or retired is not agent): - # Soft retirement preserves same-session tools and the SQLite row. - # close() would end that row and destroy task-ID-scoped resources. + # Retire only this instance's clients and provider/engine handles. + # close() would end the row and destroy task-ID-scoped resources. if not published: retired._owns_session_db = False retired._end_session_on_close = False try: - retired.release_clients() + retired.retire_local_resources(preserve_agent=new_agent if published else agent) except Exception: - logger.debug("Bot capability client retirement failed", exc_info=True) + logger.debug("Bot capability local retirement failed", exc_info=True) def _sync_agent_model_with_config(sid: str, session: dict) -> None: @@ -9517,6 +9518,7 @@ def _make_agent( service_tier_override: str | None = None, platform_override: str | None = None, _preserve_default_reasoning: bool = False, + _resource_preserve_agent=None, ): # AC-4 test seam: dead unless explicitly armed by the isolated certify # harness. Both inline and compute-host paths construct through _make_agent, @@ -9687,6 +9689,7 @@ def _make_agent( platform=_resolve_agent_platform(platform_override), session_id=session_id or key, session_db=session_db if session_db is not None else _get_db(), + _resource_preserve_agent=_resource_preserve_agent, ephemeral_system_prompt=system_prompt or None, checkpoints_enabled=is_truthy_value(os.environ.get("HERMES_TUI_CHECKPOINTS")), pass_session_id=is_truthy_value(os.environ.get("HERMES_TUI_PASS_SESSION_ID")), @@ -14089,8 +14092,10 @@ def _interim_assistant_cb(text: str, *, already_streamed: bool = False) -> None: status = ( "interrupted" if result.get("interrupted") - else "error" if result.get("error") else "complete" + else "error" if result.get("error") or result.get("partial") else "complete" ) + if not raw and result.get("partial"): + raw = result.get("partial_response") or "" # When the backend produced no visible response AND reported a # real error (e.g. invalid model slug → provider 4xx), surface # that error as the visible text instead of shipping an empty @@ -14120,6 +14125,13 @@ def _interim_assistant_cb(text: str, *, already_streamed: bool = False) -> None: status = "complete" payload = {"text": raw, "usage": _get_usage(agent), "status": status} + if isinstance(result, dict) and result.get("partial"): + payload["partial"] = True + payload["completed"] = False + if result.get("error"): + payload["error"] = str(result["error"]) + if result.get("interrupted"): + payload["interrupted"] = True if last_reasoning: payload["reasoning"] = last_reasoning if status_note: @@ -14157,13 +14169,20 @@ def _interim_assistant_cb(text: str, *, already_streamed: bool = False) -> None: _error_surface = None with session["history_lock"]: if status == "error": + # Native result-only turns may have no streamed deltas. + # Keep the draft in the resume snapshot without appending + # it again after a turn that already streamed it. + if result.get("partial") and result.get("partial_response"): + inflight = session.get("inflight_turn") + if isinstance(inflight, dict): + inflight["assistant"] = result["partial_response"] # Returned-error result (provider 4xx, budget, etc.): retain # the failed turn for resume replay instead of clearing it. # If this terminal frame is lost to a disconnect, resume's # inflight payload is the only carrier of the failure. _fail_inflight_turn( session, - result.get("error") if isinstance(result, dict) else raw, + (result.get("error") or "Turn incomplete") if isinstance(result, dict) else raw, error_surface=_error_surface, ) turn_error_retained = True From ea91ebc7674bdde757f1b0d44e073a9bb4e7252c Mon Sep 17 00:00:00 2001 From: Josh Stevenson Date: Sat, 3 Oct 2026 07:48:27 -0700 Subject: [PATCH 2/2] fix(catalog): isolate target credentials and session catalog ownership Carry profile home and secret scope through catalog workers and cache refresh; exclude AWS SDK catalog prefetch. Retain reviewed native/group/provider/dial fixes. Timer fixture binds the existing runtime seam without altering assertions. --- agent/secret_scope.py | 9 +- .../gateway-event/session-info.ts | 19 +- .../session-info-side-effects.test.tsx | 54 ++- .../tests/hide-bot-chats.runtime.test.ts | 2 + hermes_cli/auth.py | 12 +- hermes_cli/model_switch.py | 74 ++-- hermes_cli/models.py | 50 +-- hermes_cli/web_server.py | 4 +- .../test_catalog_review_boundaries.py | 333 ++++++++++++++++++ tui_gateway/methods_complete.py | 14 +- 10 files changed, 512 insertions(+), 59 deletions(-) create mode 100644 tests/hermes_cli/test_catalog_review_boundaries.py diff --git a/agent/secret_scope.py b/agent/secret_scope.py index 1a58ccb9006a9..81b0652d0ed95 100644 --- a/agent/secret_scope.py +++ b/agent/secret_scope.py @@ -135,6 +135,7 @@ class ProfileSecretScope(Mapping[str, str]): generation: str source_status: str digest: str + allow_environment_fallback: bool = True def __getitem__(self, key: str) -> str: return self.data[key] @@ -204,7 +205,10 @@ def _immutable_scope( profile_home: Path | None, source_status: str, external_generation: int = 0, + allow_environment_fallback: bool = True, ) -> ProfileSecretScope: + if not allow_environment_fallback: + source_status += ";ambient:excluded" copied = {str(key): str(value) for key, value in values.items()} generation, digest = _scope_generation( profile_home, @@ -219,6 +223,7 @@ def _immutable_scope( generation=generation, source_status=source_status, digest=digest, + allow_environment_fallback=allow_environment_fallback, ) @@ -352,7 +357,7 @@ def get_secret(name: str, default: Optional[str] = None) -> Optional[str]: val = scope.get(name) if val is not None: return val - if _MULTIPLEX_ACTIVE: + if _MULTIPLEX_ACTIVE or not scope.allow_environment_fallback: return default # Multiplex off: the scope is an overlay over the process environment, # not an isolation boundary — there is no other profile to leak from. @@ -604,6 +609,7 @@ def build_profile_secret_scope( hermes_home: Path, *, fail_closed_external: bool = False, + allow_environment_fallback: bool = True, ) -> ProfileSecretScope: """Build a profile's secret mapping from its ``.env`` and optional ``.op.env``. @@ -664,6 +670,7 @@ def build_profile_secret_scope( f"external:{external_snapshot.status}" ), external_generation=int(external_snapshot.generation), + allow_environment_fallback=allow_environment_fallback, ) diff --git a/apps/desktop/src/app/session/hooks/use-message-stream/gateway-event/session-info.ts b/apps/desktop/src/app/session/hooks/use-message-stream/gateway-event/session-info.ts index 9215d3cbe21f1..68664bff03d24 100644 --- a/apps/desktop/src/app/session/hooks/use-message-stream/gateway-event/session-info.ts +++ b/apps/desktop/src/app/session/hooks/use-message-stream/gateway-event/session-info.ts @@ -1,3 +1,4 @@ +import { getApiRequestConnection } from '@/api/client' import { normalizePersonalityValue } from '@/lib/chat-runtime' import { modelOptionsQueryKey } from '@/lib/model-options' import { reconcileApprovalModeForProfile } from '@/store/approval-mode' @@ -24,6 +25,7 @@ import { setWorkspaceCwdOwner, setYoloActive } from '@/store/session' +import { knownOwnerForSession } from '@/store/session-states' import { reportInstallMethodWarning } from '@/store/updates' import { finalizeInterruptedMessages } from '../../use-prompt-actions/rewind' @@ -426,8 +428,23 @@ export function handleSessionInfoEvent(ctx: GatewayEventContext): boolean { } if (modelValueChanged || providerValueChanged) { + const knownOwner = knownOwnerForSession(sessionId) + const owner = typeof knownOwner === 'string' ? { connectionId: 'local', profile: knownOwner } : knownOwner + const eventConnection = event.connectionId?.trim() || undefined + const eventProfile = event.profile?.trim() || undefined + + const matchesOwner = owner && (!eventConnection || eventConnection === owner.connectionId) && + (!eventProfile || eventProfile === owner.profile || eventProfile === owner.targetProfile) + + const completeEventOwner = eventConnection && eventProfile + const legacyUnstamped = !eventConnection && !eventProfile && !owner + const catalogProfile = matchesOwner ? owner.targetProfile || owner.profile : eventProfile || activeGatewayProfile + const catalogConnection = matchesOwner ? owner.connectionId : eventConnection || getApiRequestConnection() + // An incomplete stamp that conflicts with a known owner cannot name + // one exact cache. Invalidate the catalog family without mixing owners. + const exactOwner = matchesOwner || completeEventOwner || legacyUnstamped void queryClient.invalidateQueries({ - queryKey: explicitSid && sessionId ? modelOptionsQueryKey(activeGatewayProfile, sessionId) : ['model-options'] + queryKey: explicitSid && sessionId && exactOwner ? modelOptionsQueryKey(catalogProfile, sessionId, catalogConnection) : ['model-options'] }) } diff --git a/apps/desktop/src/app/session/hooks/use-message-stream/session-info-side-effects.test.tsx b/apps/desktop/src/app/session/hooks/use-message-stream/session-info-side-effects.test.tsx index 2d1738ac0f58f..dd50197b64ec4 100644 --- a/apps/desktop/src/app/session/hooks/use-message-stream/session-info-side-effects.test.tsx +++ b/apps/desktop/src/app/session/hooks/use-message-stream/session-info-side-effects.test.tsx @@ -2,11 +2,12 @@ import { QueryClient } from '@tanstack/react-query' import { act, cleanup } from '@testing-library/react' import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' +import { setApiRequestConnection } from '@/api/client' import { isTargetSessionBusy } from '@/app/session/hooks/use-prompt-actions/utils' import type { ClientSessionState } from '@/app/types' import { createClientSessionState } from '@/lib/chat-runtime' import { modelOptionsQueryKey } from '@/lib/model-options' -import { setCurrentModel, setCurrentProvider } from '@/store/session' +import { _resetSessionOwnerHintsForTests, setCurrentModel, setCurrentProvider, setSessionOwnerHint } from '@/store/session' import { type MessageStreamHarness, renderMessageStream } from './test-harness' import { PRE_TURN_LIVE_SETTLE_GRACE_MS } from './utils' @@ -55,6 +56,8 @@ afterEach(() => { setCurrentProvider('') vi.useRealTimers() vi.restoreAllMocks() + setApiRequestConnection(null) + _resetSessionOwnerHintsForTests() }) describe('session.info config refetch gating', () => { @@ -92,6 +95,55 @@ describe('session.info config refetch gating', () => { }) describe('session.info model-options invalidation gating', () => { + it('does not mix a stale owner hint with an incompatible connection-only event stamp', () => { + setApiRequestConnection('connection-a') + setSessionOwnerHint('session-background', { + connectionId: 'connection-a', profile: 'a-profile', targetProfile: 'a-target' + }) + mountStream() + const invalidate = vi.spyOn(queryClient, 'invalidateQueries') + act(() => stream.handleEvent({ + connectionId: 'connection-b', session_id: 'session-background', + type: 'session.info', payload: { model: 'b-new', provider: 'b-provider' } + })) + expect(invalidate).toHaveBeenCalledWith({ queryKey: ['model-options'] }) + expect(invalidate).not.toHaveBeenCalledWith({ + queryKey: modelOptionsQueryKey('a-target', 'session-background', 'connection-b') + }) + }) + + it('invalidates the background event owner catalog and backend target while another connection is foregrounded', () => { + setApiRequestConnection('connection-a') + setSessionOwnerHint('session-background', { + connectionId: 'connection-b', profile: 'tile-profile', targetProfile: 'backend-target' + }) + mountStream() + const invalidate = vi.spyOn(queryClient, 'invalidateQueries') + act(() => stream.handleEvent({ + connectionId: 'connection-b', profile: 'tile-profile', session_id: 'session-background', + type: 'session.info', payload: { model: 'b-new', provider: 'b-provider' } + })) + expect(invalidate).toHaveBeenCalledWith({ + queryKey: modelOptionsQueryKey('backend-target', 'session-background', 'connection-b') + }) + expect(invalidate).not.toHaveBeenCalledWith({ + queryKey: modelOptionsQueryKey(ACTIVE_PROFILE, 'session-background', 'connection-a') + }) + }) + + it('uses a stamped background event connection when no session owner hint is known', () => { + setApiRequestConnection('connection-a') + mountStream() + const invalidate = vi.spyOn(queryClient, 'invalidateQueries') + act(() => stream.handleEvent({ + connectionId: 'connection-b', profile: 'background-profile', session_id: 'session-background', + type: 'session.info', payload: { model: 'b-new', provider: 'b-provider' } + })) + expect(invalidate).toHaveBeenCalledWith({ + queryKey: modelOptionsQueryKey('background-profile', 'session-background', 'connection-b') + }) + }) + it('skips invalidation when model/provider merely restate the known values', () => { mountStream() const invalidate = vi.spyOn(queryClient, 'invalidateQueries') diff --git a/apps/desktop/src/plugins/hermes-bots/tests/hide-bot-chats.runtime.test.ts b/apps/desktop/src/plugins/hermes-bots/tests/hide-bot-chats.runtime.test.ts index da3f06f2bc8c9..f966a08130e02 100644 --- a/apps/desktop/src/plugins/hermes-bots/tests/hide-bot-chats.runtime.test.ts +++ b/apps/desktop/src/plugins/hermes-bots/tests/hide-bot-chats.runtime.test.ts @@ -53,11 +53,13 @@ afterEach(() => { gatewayState.set('closed') vi.clearAllMocks() vi.useRealTimers() + plugin.groupTurnRuntime.bindGroupTurnPorts(globalThis) }) describe('Bot Mode hidden-session reconciliation lifecycle', () => { it('uses persisted REST on load/reconnect and stops with plugin disposal', async () => { vi.useFakeTimers() + plugin.groupTurnRuntime.bindGroupTurnPorts(globalThis) const disposers: Array<() => void> = [] plugin.register(createPluginContext(plugin.id, dispose => disposers.push(dispose))) diff --git a/hermes_cli/auth.py b/hermes_cli/auth.py index 4ed3a72febfa3..b7c1b749713c8 100644 --- a/hermes_cli/auth.py +++ b/hermes_cli/auth.py @@ -1952,6 +1952,8 @@ def is_provider_explicitly_configured(provider_id: str) -> bool: without the user's explicit choice. See PR #4210 for the same pattern applied to the setup wizard gate. """ + from agent.secret_scope import get_secret + normalized = (provider_id or "").strip().lower() # 1. Check auth.json active_provider @@ -2021,7 +2023,7 @@ def _slot_matches_provider(slot): for env_var in pconfig.api_key_env_vars: if env_var in _IMPLICIT_ENV_VARS: continue - if has_usable_secret(os.getenv(env_var, "")): + if has_usable_secret(get_secret(env_var, "")): return True # AWS SDK providers (Bedrock) have auth_type="aws_sdk" and empty @@ -2033,11 +2035,11 @@ def _slot_matches_provider(slot): # Only check explicit env credentials here (NOT boto3's full chain): # ambient sources like EC2 IMDS / SSO profiles must not auto-surface. if pconfig and pconfig.auth_type == "aws_sdk": - if has_usable_secret(os.getenv("AWS_BEARER_TOKEN_BEDROCK", "")): + if has_usable_secret(get_secret("AWS_BEARER_TOKEN_BEDROCK", "")): return True if ( - has_usable_secret(os.getenv("AWS_ACCESS_KEY_ID", "")) - and has_usable_secret(os.getenv("AWS_SECRET_ACCESS_KEY", "")) + has_usable_secret(get_secret("AWS_ACCESS_KEY_ID", "")) + and has_usable_secret(get_secret("AWS_SECRET_ACCESS_KEY", "")) ): return True @@ -2057,7 +2059,7 @@ def _slot_matches_provider(slot): # the user deletes the env var (#55790) — only count it when # the referenced var still resolves to a usable secret NOW. env_var = entry.get("source", "").split(":", 1)[1].strip() - if env_var and has_usable_secret(os.getenv(env_var, "")): + if env_var and has_usable_secret(get_secret(env_var, "")): return True continue if ( diff --git a/hermes_cli/model_switch.py b/hermes_cli/model_switch.py index 78cdeb6549ffe..2900c14a29a81 100644 --- a/hermes_cli/model_switch.py +++ b/hermes_cli/model_switch.py @@ -2426,7 +2426,13 @@ def _fetch_one(slug: str) -> None: max_workers=min(_PARALLEL_PREFETCH_WORKERS, len(stale_slugs)), thread_name_prefix="model-cache-prefetch", ) as executor: - list(executor.map(_fetch_one, stale_slugs)) + # Each task needs its own copy: one Context cannot be entered by + # concurrent workers. Carry both target home and credential scope. + from contextvars import copy_context + context = copy_context() + futures = [executor.submit(context.copy().run, _fetch_one, slug) for slug in stale_slugs] + for future in futures: + future.result() def _collect_authed_provider_slugs( @@ -2560,13 +2566,17 @@ def _collect_authed_provider_slugs( seen.add(pid.lower()) seen.add(hermes_slug.lower()) - # --- Section 2b: Canonical providers cross-check --- + # --- Section 2b: Canonical providers cross-check --- for _cp in CANONICAL_PROVIDERS: if _cp.slug.lower() in seen: continue if _cp.slug.lower() in _excluded_set: continue _cp_config = PROVIDER_REGISTRY.get(_cp.slug) + if _cp_config and getattr(_cp_config, "auth_type", "") == "aws_sdk": + # A saved auth record or pool does not make the SDK's ambient + # credential chain safe for a target-profile catalog prefetch. + continue _cp_has_creds = False if _cp_config and _cp_config.api_key_env_vars: _cp_has_creds = any(_scoped_key_env(ev) for ev in _cp_config.api_key_env_vars) @@ -2584,8 +2594,6 @@ def _collect_authed_provider_slugs( _cp_has_creds = True except Exception: pass - if not _cp_has_creds and _cp_config and getattr(_cp_config, "auth_type", "") == "aws_sdk": - continue # skip AWS SDK in prefetch if _cp_has_creds: slugs.append(_cp.slug) seen.add(_cp.slug.lower()) @@ -2657,8 +2665,17 @@ def list_authenticated_providers( clear_provider_models_cache, get_curated_nous_model_ids, ) + def _catalog_provider_model_ids(slug): + from agent.secret_scope import current_secret_scope + scope = current_secret_scope() + if slug == "bedrock" and scope is not None and not scope.allow_environment_fallback: + # The SDK discovers through its process-global credential chain. + # A cross-profile catalog can use curated IDs without that probe. + return list(_PROVIDER_MODELS.get(slug, [])) + return cached_provider_model_ids(slug) + # Explicit refresh: drop every provider's cached model-id list so the - # cached_provider_model_ids() calls below all re-fetch live. Without this + # _catalog_provider_model_ids() calls below all re-fetch live. Without this # a stale 1h cache can fall back to the curated static list when its live # fetch later fails, silently dropping live-only models (e.g. OpenCode # Zen's free tier) the user had seen before. @@ -2706,7 +2723,7 @@ def _record_builtin_endpoint(slug: str) -> None: return url = "" if getattr(pcfg, "base_url_env_var", ""): - url = os.environ.get(pcfg.base_url_env_var, "") or "" + url = _scoped_key_env(pcfg.base_url_env_var) or "" if not url: url = getattr(pcfg, "inference_base_url", "") or "" normed = _norm_url(url) @@ -2721,15 +2738,15 @@ def _has_fast_aws_sdk_signal() -> bool: botocore may otherwise probe EC2 IMDS (169.254.169.254) on local machines before returning no credentials. """ - if os.environ.get("AWS_BEARER_TOKEN_BEDROCK", "").strip(): + if _scoped_key_env("AWS_BEARER_TOKEN_BEDROCK").strip(): return True if ( - os.environ.get("AWS_ACCESS_KEY_ID", "").strip() - and os.environ.get("AWS_SECRET_ACCESS_KEY", "").strip() + _scoped_key_env("AWS_ACCESS_KEY_ID").strip() + and _scoped_key_env("AWS_SECRET_ACCESS_KEY").strip() ): return True return any( - os.environ.get(name, "").strip() + _scoped_key_env(name).strip() for name in ( "AWS_PROFILE", "AWS_CONTAINER_CREDENTIALS_RELATIVE_URI", @@ -2744,6 +2761,11 @@ def _has_aws_sdk_creds_for_listing(slug: str) -> bool: current_norm = str(current_provider or "").strip().lower() if _has_fast_aws_sdk_signal(): return True + from agent.secret_scope import current_secret_scope + scope = current_secret_scope() + if scope is not None and not scope.allow_environment_fallback: + # boto3's ambient chain is process-owned, not this target profile. + return False if slug_norm != current_norm: return False try: @@ -2774,19 +2796,19 @@ def _has_aws_sdk_creds_for_listing(slug: str) -> bool: # On auth rejection or unreachable server, fall back to the caller-supplied # current model so the picker still shows something when offline / mis-keyed. if "lmstudio" not in curated and ( - os.environ.get("LM_API_KEY") or os.environ.get("LM_BASE_URL") or current_provider.strip().lower() == "lmstudio" + _scoped_key_env("LM_API_KEY") or _scoped_key_env("LM_BASE_URL") or current_provider.strip().lower() == "lmstudio" ): from hermes_cli.models import fetch_lmstudio_models from hermes_cli.auth import AuthError is_current_lmstudio = current_provider.strip().lower() == "lmstudio" lm_base = ( - os.environ.get("LM_BASE_URL") + _scoped_key_env("LM_BASE_URL") or (current_base_url if is_current_lmstudio and current_base_url else None) or "http://127.0.0.1:1234/v1" ) try: live = fetch_lmstudio_models( - api_key=os.environ.get("LM_API_KEY", ""), + api_key=_scoped_key_env("LM_API_KEY"), base_url=lm_base, timeout=1.5, # Smaller timeout for picker ) @@ -2798,7 +2820,7 @@ def _has_aws_sdk_creds_for_listing(slug: str) -> bool: # --- Parallel cache prefetch --------------------------------------------- # The serial loops below (sections 1, 2, 2b) each call - # cached_provider_model_ids(slug) which blocks on a live /v1/models HTTP + # _catalog_provider_model_ids(slug) which blocks on a live /v1/models HTTP # round-trip when the disk cache is stale or missing. With many authed # providers those serial round-trips stack to 15-30s on a cold/expired # cache. Pre-scanning which providers have credentials (without fetching @@ -2891,7 +2913,7 @@ def _has_aws_sdk_creds_for_listing(slug: str) -> bool: continue # Check if any env var is set - has_creds = any(os.environ.get(ev) for ev in env_vars) + has_creds = any(_scoped_key_env(ev) for ev in env_vars) if not has_creds: try: from hermes_cli.auth import _load_auth_store @@ -2908,11 +2930,11 @@ def _has_aws_sdk_creds_for_listing(slug: str) -> bool: if not has_creds: continue - # Unified pathway: route through cached_provider_model_ids() so the + # Unified pathway: route through _catalog_provider_model_ids() so the # /model picker sees the SAME list `hermes model` would build, with # disk caching to keep the picker open snappy. Falls back to the # curated static list when the live fetcher returns nothing. - model_ids = cached_provider_model_ids(hermes_id) + model_ids = _catalog_provider_model_ids(hermes_id) if not model_ids: model_ids = curated.get(hermes_id, []) if hermes_id in _MODELS_DEV_PREFERRED: @@ -2990,13 +3012,13 @@ def _has_aws_sdk_creds_for_listing(slug: str) -> bool: except Exception as exc: logger.debug("Vertex credential check failed: %s", exc) elif overlay.extra_env_vars: - has_creds = any(os.environ.get(ev) for ev in overlay.extra_env_vars) + has_creds = any(_scoped_key_env(ev) for ev in overlay.extra_env_vars) # Also check api_key_env_vars from PROVIDER_REGISTRY for api_key auth_type if not has_creds and overlay.auth_type == "api_key": for _key in (pid, hermes_slug): pcfg = _auth_registry.get(_key) if pcfg and pcfg.api_key_env_vars: - if any(os.environ.get(ev) for ev in pcfg.api_key_env_vars): + if any(_scoped_key_env(ev) for ev in pcfg.api_key_env_vars): has_creds = True break # Check auth store and credential pool for non-env-var credentials. @@ -3064,15 +3086,15 @@ def _has_aws_sdk_creds_for_listing(slug: str) -> bool: # matches what the user's authenticated Codex/Copilot backend # actually serves — including ChatGPT-Pro-only Codex slugs # (e.g. gpt-5.3-codex-spark) that aren't in the static curated - # catalog. ``cached_provider_model_ids()`` falls back to the + # catalog. ``_catalog_provider_model_ids()`` falls back to the # curated list when the live endpoint is unreachable, so this # is safe for unauthenticated and offline cases too. - model_ids = cached_provider_model_ids(hermes_slug) + model_ids = _catalog_provider_model_ids(hermes_slug) # For aws_sdk providers (bedrock), use live discovery so the list # reflects the active region (eu.*, ap.*) not the static us.* list. elif overlay.auth_type == "aws_sdk": try: - _ids = cached_provider_model_ids(hermes_slug) + _ids = _catalog_provider_model_ids(hermes_slug) model_ids = _ids if _ids else (curated.get(hermes_slug, []) or curated.get(pid, [])) except Exception: model_ids = curated.get(hermes_slug, []) or curated.get(pid, []) @@ -3117,7 +3139,7 @@ def _has_aws_sdk_creds_for_listing(slug: str) -> bool: # Unified pathway — see Section 1 rationale. Fall back to the # curated dict (with models.dev merge for preferred providers) # when the live fetcher comes up empty. - model_ids = cached_provider_model_ids(hermes_slug) + model_ids = _catalog_provider_model_ids(hermes_slug) if not model_ids: model_ids = curated.get(hermes_slug, []) or curated.get(pid, []) if hermes_slug in _MODELS_DEV_PREFERRED: @@ -3160,7 +3182,7 @@ def _has_aws_sdk_creds_for_listing(slug: str) -> bool: _cp_config = _auth_registry.get(_cp.slug) _cp_has_creds = False if _cp_config and _cp_config.api_key_env_vars: - _cp_has_creds = any(os.environ.get(ev) for ev in _cp_config.api_key_env_vars) + _cp_has_creds = any(_scoped_key_env(ev) for ev in _cp_config.api_key_env_vars) # Also check auth store and credential pool if not _cp_has_creds: try: @@ -3191,13 +3213,13 @@ def _has_aws_sdk_creds_for_listing(slug: str) -> bool: # region (eu.*, us.*, ap.*) instead of the hardcoded us.* static list. if _cp_config and getattr(_cp_config, "auth_type", "") == "aws_sdk": try: - _ids = cached_provider_model_ids(_cp.slug) + _ids = _catalog_provider_model_ids(_cp.slug) _cp_model_ids = _ids if _ids else curated.get(_cp.slug, []) except Exception: _cp_model_ids = curated.get(_cp.slug, []) else: # Unified pathway — same as sections 1 and 2. - _cp_model_ids = cached_provider_model_ids(_cp.slug) + _cp_model_ids = _catalog_provider_model_ids(_cp.slug) if not _cp_model_ids: _cp_model_ids = curated.get(_cp.slug, []) _cp_total = len(_cp_model_ids) diff --git a/hermes_cli/models.py b/hermes_cli/models.py index 7bc228bed8ef6..f6df2b0856619 100644 --- a/hermes_cli/models.py +++ b/hermes_cli/models.py @@ -18,6 +18,7 @@ import urllib.request import urllib.error import time +from contextvars import copy_context from difflib import get_close_matches from pathlib import Path from typing import Any, NamedTuple, Optional, TYPE_CHECKING @@ -25,6 +26,7 @@ if TYPE_CHECKING: from typing import TypeGuard +from agent.secret_scope import get_secret as _get_secret from hermes_cli import __version__ as _HERMES_VERSION from hermes_cli.urllib_security import open_credentialed_url, url_origin from utils import atomic_json_write, base_url_host_matches @@ -1731,8 +1733,9 @@ def _warm_reasoning_caps_async(refresh) -> None: """ if os.environ.get("PYTEST_CURRENT_TEST"): return + context = copy_context() threading.Thread( - target=refresh, name="reasoning-caps-warm", daemon=True + target=lambda: context.run(refresh), name="reasoning-caps-warm", daemon=True ).start() @@ -2512,7 +2515,7 @@ def fetch_ai_gateway_pricing( def _resolve_openrouter_api_key() -> str: """Best-effort OpenRouter API key for pricing fetch.""" - return os.getenv("OPENROUTER_API_KEY", "").strip() + return _get_secret("OPENROUTER_API_KEY", "").strip() _DEFAULT_NOUS_INFERENCE_BASE = "https://inference-api.nousresearch.com" @@ -2662,11 +2665,11 @@ def _fetch_novita_pricing( matching the pattern used by ``fetch_ai_gateway_pricing`` — without this, every menu render or pricing lookup re-hits the network. """ - api_key = os.getenv("NOVITA_API_KEY", "").strip() + api_key = _get_secret("NOVITA_API_KEY", "").strip() if not api_key: return {} - base_url = os.getenv("NOVITA_BASE_URL", "").strip() or "https://api.novita.ai/openai/v1" + base_url = _get_secret("NOVITA_BASE_URL", "").strip() or "https://api.novita.ai/openai/v1" cache_key = base_url.rstrip("/") if not force_refresh: cached = _cached_catalog(cache_key) @@ -2765,7 +2768,7 @@ def list_available_providers() -> list[dict[str, str]]: custom_base_url = _get_custom_base_url() or "" has_creds = bool(custom_base_url.strip()) elif pid == "openrouter": - has_creds = has_usable_secret(os.getenv("OPENROUTER_API_KEY", "")) + has_creds = has_usable_secret(_get_secret("OPENROUTER_API_KEY", "")) else: status = get_auth_status(pid) has_creds = bool(status.get("logged_in") or status.get("configured")) @@ -2911,7 +2914,7 @@ def _get_ollama_base_url() -> str: except (OSError, RuntimeError, TypeError, ValueError): pass - env_host = os.getenv("OLLAMA_HOST", "").strip() + env_host = _get_secret("OLLAMA_HOST", "").strip() if env_host: if env_host.startswith(":") and not env_host.startswith("::"): env_host = "127.0.0.1" + env_host @@ -2958,7 +2961,7 @@ def _get_ollama_request_headers() -> dict[str, str]: key_env = str( entry.get("key_env") or entry.get("api_key_env") or "" ).strip() - api_key = os.getenv(key_env, "").strip() if key_env else "" + api_key = _get_secret(key_env, "").strip() if key_env else "" if api_key: if not any(key.lower() == "authorization" for key in result): result["Authorization"] = f"Bearer {api_key}" @@ -3850,7 +3853,7 @@ def _openai_discovery_base_url(provider: str) -> str: config-set data-residency host (``us.api.openai.com``) was ignored and the catalog kept coming from ``api.openai.com``. """ - env_raw = os.getenv("OPENAI_BASE_URL", "").strip().rstrip("/") + env_raw = _get_secret("OPENAI_BASE_URL", "").strip().rstrip("/") if env_raw: return env_raw try: @@ -3904,7 +3907,7 @@ def provider_model_ids(provider: Optional[str], *, force_refresh: bool = False) fallback_key = str(config.get("api_key") or "").strip() if not fallback_key: key_env = str(config.get("key_env") or "").strip() - fallback_key = os.getenv(key_env, "").strip() if key_env else "" + fallback_key = _get_secret(key_env, "").strip() if key_env else "" fallback_base = _normalize_openai_base_url( config.get("base_url") or base_url ) @@ -4020,7 +4023,7 @@ def provider_model_ids(provider: Optional[str], *, force_refresh: bool = False) if live: return live if normalized in ("openai", "openai-api"): - api_key = os.getenv("OPENAI_API_KEY", "").strip() + api_key = _get_secret("OPENAI_API_KEY", "").strip() if api_key: base = _openai_discovery_base_url(normalized) # Custom OpenAI-compatible endpoints (proxies, gateways, self-hosted) @@ -4073,9 +4076,9 @@ def provider_model_ids(provider: Optional[str], *, force_refresh: bool = False) # Try common API key env vars for custom endpoints api_key = ( str(model_cfg.get("api_key", "") or "").strip() - or os.getenv("CUSTOM_API_KEY", "") - or os.getenv("OPENAI_API_KEY", "") - or os.getenv("OPENROUTER_API_KEY", "") + or _get_secret("CUSTOM_API_KEY", "") + or _get_secret("OPENAI_API_KEY", "") + or _get_secret("OPENROUTER_API_KEY", "") ) api_mode = "anthropic_messages" if _base_url_looks_like_anthropic_messages(base_url) else None live = fetch_api_models(api_key, base_url, api_mode=api_mode) @@ -4318,8 +4321,9 @@ def _refresh() -> None: with _swr_refresh_lock: _swr_refresh_inflight.discard(cache_key) + context = copy_context() threading.Thread( - target=_refresh, daemon=True, name=f"model-cache-swr-{cache_key}" + target=lambda: context.run(_refresh), daemon=True, name=f"model-cache-swr-{cache_key}" ).start() @@ -4352,10 +4356,10 @@ def _credential_fingerprint(provider: str) -> str: pcfg = PROVIDER_REGISTRY.get(provider) if pcfg is not None: for ev in getattr(pcfg, "api_key_env_vars", ()) or (): - parts.append(f"{ev}={_os.environ.get(ev, '')}") + parts.append(f"{ev}={_get_secret(ev, '')}") bev = getattr(pcfg, "base_url_env_var", "") or "" if bev: - parts.append(f"{bev}={_os.environ.get(bev, '')}") + parts.append(f"{bev}={_get_secret(bev, '')}") except Exception: pass @@ -4371,7 +4375,7 @@ def _credential_fingerprint(provider: str) -> str: pass if provider == "ollama": - parts.append(f"OLLAMA_HOST={_os.environ.get('OLLAMA_HOST', '')}") + parts.append(f"OLLAMA_HOST={_get_secret('OLLAMA_HOST', '')}") provider_cfg = _get_provider_config_dict("ollama") parts.append( "providers.ollama.base_url=" @@ -4381,7 +4385,7 @@ def _credential_fingerprint(provider: str) -> str: key_env = provider_cfg.get("key_env") or provider_cfg.get("api_key_env") or "" parts.append(f"providers.ollama.key_env={key_env}") if key_env: - parts.append(f"{key_env}={_os.environ.get(str(key_env), '')}") + parts.append(f"{key_env}={_get_secret(str(key_env), '')}") model_cfg = _get_model_config_dict() parts.append( "model.provider=" @@ -5899,7 +5903,7 @@ def probe_api_models( def _deepinfra_catalog_url() -> tuple[str, str]: """Return ``(cache_key, full_url)`` for the DeepInfra catalog endpoint.""" - base = os.getenv("DEEPINFRA_BASE_URL", "").strip() or _DEEPINFRA_DEFAULT_BASE_URL + base = _get_secret("DEEPINFRA_BASE_URL", "").strip() or _DEEPINFRA_DEFAULT_BASE_URL cache_key = base.rstrip("/") return cache_key, f"{cache_key}/models?{_DEEPINFRA_MODELS_QUERY}" @@ -5924,7 +5928,7 @@ def _fetch_deepinfra_catalog( return None headers: dict[str, str] = {"User-Agent": _HERMES_USER_AGENT} - api_key = os.getenv("DEEPINFRA_API_KEY", "").strip() + api_key = _get_secret("DEEPINFRA_API_KEY", "").strip() if api_key: headers["Authorization"] = f"Bearer {api_key}" @@ -6037,7 +6041,7 @@ def deepinfra_base_url(section: Optional[dict] = None) -> str: to re-code (with subtly divergent normalization). """ candidate = section.get("base_url") if isinstance(section, dict) else None - value = candidate or os.getenv("DEEPINFRA_BASE_URL") or _DEEPINFRA_DEFAULT_BASE_URL + value = candidate or _get_secret("DEEPINFRA_BASE_URL") or _DEEPINFRA_DEFAULT_BASE_URL return str(value).strip().rstrip("/") @@ -6084,10 +6088,10 @@ def _fetch_deepinfra_pricing( def _fetch_ai_gateway_models(timeout: float = 5.0) -> Optional[list[str]]: """Fetch available language models with tool-use from AI Gateway.""" - api_key = os.getenv("AI_GATEWAY_API_KEY", "").strip() + api_key = _get_secret("AI_GATEWAY_API_KEY", "").strip() if not api_key: return None - base_url = os.getenv("AI_GATEWAY_BASE_URL", "").strip() + base_url = _get_secret("AI_GATEWAY_BASE_URL", "").strip() if not base_url: from hermes_constants import AI_GATEWAY_BASE_URL base_url = AI_GATEWAY_BASE_URL diff --git a/hermes_cli/web_server.py b/hermes_cli/web_server.py index e5b423b04e6da..0e48a523b2ece 100644 --- a/hermes_cli/web_server.py +++ b/hermes_cli/web_server.py @@ -1958,7 +1958,9 @@ def _validate_model_assignment_provider( # Check raw entries as well as the compatibility view: disabled entries # are omitted from that view but must not fall through to a built-in alias. for key, entry in providers_cfg.items(): - matches_builtin = not requested.startswith("custom:") and canonical == _model_assignment_provider_category(str(key)) + # Canonical declarations own their native aliases. An alias-named + # independent endpoint owns only its declared key/name identities. + matches_builtin = not requested.startswith("custom:") and canonical == str(key).strip().lower() if matches_builtin or requested in custom_provider_aliases(str(entry.get("name") or key) if isinstance(entry, dict) else str(key), str(key)): if not isinstance(entry, dict) or not is_provider_enabled(entry): raise HTTPException(status_code=400, detail=f"Provider '{provider}' is disabled or invalid in this profile") diff --git a/tests/hermes_cli/test_catalog_review_boundaries.py b/tests/hermes_cli/test_catalog_review_boundaries.py new file mode 100644 index 0000000000000..e02127f3519f0 --- /dev/null +++ b/tests/hermes_cli/test_catalog_review_boundaries.py @@ -0,0 +1,333 @@ +"""PR120 identity and credential catalog boundaries, using inert inventory I/O.""" +from pathlib import Path + +import pytest +import yaml + + +@pytest.mark.parametrize("requested", ["openrouter", "openai-api", "anthropic"]) +def test_disabled_openai_declaration_does_not_disable_distinct_builtin(requested): + from hermes_cli.web_server import _validate_model_assignment_provider + _validate_model_assignment_provider({"providers": {"openai": {"enabled": False}}}, requested, "") + + +@pytest.mark.parametrize("requested", ["openai", "custom:openai"]) +def test_disabled_exact_declaration_still_rejects(requested): + from fastapi import HTTPException + from hermes_cli.web_server import _validate_model_assignment_provider + with pytest.raises(HTTPException) as caught: + _validate_model_assignment_provider({"providers": {"openai": {"enabled": False}}}, requested, "") + assert caught.value.status_code == 400 + + +@pytest.mark.parametrize("alias,canonical", [("claude", "anthropic"), ("google", "gemini"), ("fw", "fireworks")]) +def test_disabled_independent_alias_endpoint_does_not_disable_canonical_builtin(alias, canonical): + from hermes_cli.web_server import _validate_model_assignment_provider + cfg = {"providers": {alias: {"name": "Independent fixture", "base_url": "http://fixture.invalid/v1", "enabled": False}}} + _validate_model_assignment_provider(cfg, canonical) + + +@pytest.mark.parametrize("alias,canonical", [("claude", "anthropic"), ("google", "gemini"), ("fw", "fireworks")]) +def test_disabled_canonical_declaration_still_owns_native_aliases(alias, canonical): + from fastapi import HTTPException + from hermes_cli.web_server import _validate_model_assignment_provider + with pytest.raises(HTTPException) as caught: + _validate_model_assignment_provider({"providers": {canonical: {"enabled": False}}}, alias) + assert caught.value.status_code == 400 + + +@pytest.mark.parametrize("multiplex", [False, True]) +@pytest.mark.parametrize("target_key", [None, "inert-target-anthropic"]) +def test_actual_rpc_inventory_uses_only_target_credentials_and_restores_scope(tmp_path, monkeypatch, multiplex, target_key): + import agent.secret_scope as secrets + import agent.models_dev as mdev + import hermes_cli.auth as auth + import hermes_cli.inventory as inventory + import hermes_cli.models as models + import hermes_cli.model_switch as switch + import hermes_cli.profiles as profiles + import hermes_cli.providers as providers + import tui_gateway.server as srv + from hermes_constants import get_hermes_home + + launch, target = tmp_path / "launch", tmp_path / "target" + for home in [launch, target]: + home.mkdir() + (home / "config.yaml").write_text(yaml.safe_dump({"model": {"provider": "auto", "default": "fixture"}})) + (launch / ".env").write_text("ANTHROPIC_API_KEY=inert-launch-anthropic\n") + if target_key: + (target / ".env").write_text(f"ANTHROPIC_API_KEY={target_key}\n") + monkeypatch.setattr(Path, "home", lambda: tmp_path) + monkeypatch.setenv("HOME", str(tmp_path)) + monkeypatch.setenv("HERMES_HOME", str(launch)) + monkeypatch.setenv("ANTHROPIC_API_KEY", "inert-launch-anthropic") + monkeypatch.setattr(secrets, "_MULTIPLEX_ACTIVE", multiplex) + monkeypatch.setattr(srv, "_hermes_home", launch) + monkeypatch.setattr(srv, "_sessions", {}) + monkeypatch.setattr(profiles, "get_profile_dir", lambda name: target if name == "target" else launch) + monkeypatch.setattr(mdev, "PROVIDER_TO_MODELS_DEV", {"anthropic": "anthropic"}) + monkeypatch.setattr(mdev, "fetch_models_dev", lambda: {"anthropic": {"env": ["ANTHROPIC_API_KEY"]}}) + monkeypatch.setattr(auth, "PROVIDER_REGISTRY", {"anthropic": auth.PROVIDER_REGISTRY["anthropic"]}) + monkeypatch.setattr(auth, "_load_auth_store", lambda: {}) + monkeypatch.setattr(auth, "read_credential_pool", lambda *a, **k: []) + monkeypatch.setattr(auth, "is_runtime_provider_routable", lambda slug: True) + monkeypatch.setattr(providers, "HERMES_OVERLAYS", {}) + monkeypatch.setattr(models, "CANONICAL_PROVIDERS", []) + monkeypatch.setattr(models, "_PROVIDER_MODELS", {"anthropic": ["fixture-model"]}) + monkeypatch.setattr(models, "get_curated_nous_model_ids", lambda: []) + monkeypatch.setattr(models, "fetch_ollama_cloud_models", lambda: []) + monkeypatch.setattr(switch, "_credential_pool_is_usable", lambda *a, **k: False) + fetched = [] + def inert_models(slug): + fetched.append((slug, Path(get_hermes_home()), secrets.get_secret("ANTHROPIC_API_KEY"))) + return ["fixture-model"] + monkeypatch.setattr(models, "cached_provider_model_ids", inert_models) + seen = [] + def actual_inventory(ctx, **kwargs): + rows = switch.list_authenticated_providers(probe_custom_providers=False) + seen.append((secrets.current_secret_scope(), auth.is_provider_explicitly_configured("anthropic"))) + return {"providers": rows} + monkeypatch.setattr(inventory, "build_model_options_payload", actual_inventory) + outer = secrets.set_secret_scope({"ANTHROPIC_API_KEY": "inert-outer"}) + old_scope = secrets.current_secret_scope() + try: + reply = srv.handle_request({"id": "catalog", "method": "model.options", "params": {"profile": "target", "refresh": True}}) + assert "error" not in reply, reply + slugs = {r["slug"] for r in reply["result"]["providers"]} + assert ("anthropic" in slugs) == bool(target_key) + assert seen[0][0] is not old_scope + assert seen[0][1] == bool(target_key) + assert fetched == ([('anthropic', target, target_key)] if target_key else []) + assert secrets.current_secret_scope() is old_scope + assert Path(get_hermes_home()) == launch + finally: + secrets.reset_secret_scope(outer) + + +def test_prefetch_real_worker_carries_both_profile_home_and_secret_scope(tmp_path, monkeypatch): + import agent.secret_scope as secrets + import hermes_cli.models as models + import hermes_cli.model_switch as switch + from hermes_constants import get_hermes_home, set_hermes_home_override, reset_hermes_home_override + monkeypatch.setattr(secrets, "_MULTIPLEX_ACTIVE", True) + monkeypatch.setattr(switch, "_PARALLEL_PREFETCH_WORKERS", 1) + monkeypatch.setattr(models, "_load_provider_models_cache", lambda: {}) + monkeypatch.setattr(models, "_credential_fingerprint", lambda slug: "inert-fingerprint") + monkeypatch.setattr(models, "update_provider_cache_entry", lambda *a: None) + fetched = [] + def inert_fetch(slug, **kwargs): + fetched.append((slug, Path(get_hermes_home()), secrets.get_secret("ANTHROPIC_API_KEY"))) + return ["fixture-model"] + monkeypatch.setattr(models, "cached_provider_model_ids", inert_fetch) + home_token = set_hermes_home_override(str(tmp_path)) + secret_token = secrets.set_secret_scope({"ANTHROPIC_API_KEY": "inert-prefetch-target"}) + try: + switch._prefetch_provider_models_parallel(["anthropic", "gemini", "fireworks", "deepseek"]) + assert fetched == [(slug, tmp_path, "inert-prefetch-target") for slug in ["anthropic", "gemini", "fireworks", "deepseek"]] + finally: + secrets.reset_secret_scope(secret_token) + reset_hermes_home_override(home_token) + + +def test_catalog_authority_is_opt_in_and_scope_miss_preserves_legacy_env_injection(tmp_path, monkeypatch): + import agent.secret_scope as secrets + monkeypatch.setattr(secrets, "_MULTIPLEX_ACTIVE", False) + monkeypatch.setenv("ANTHROPIC_API_KEY", "inert-environment") + for allow_fallback, expected in [(True, "inert-environment"), (False, None)]: + scope = secrets.build_profile_secret_scope(tmp_path, allow_environment_fallback=allow_fallback) + token = secrets.set_secret_scope(scope) + try: + assert secrets.get_secret("ANTHROPIC_API_KEY") == expected + assert secrets.get_secret("PATH") is not None + finally: + secrets.reset_secret_scope(token) + + +def test_rpc_exception_restores_previous_secret_and_home_context(tmp_path, monkeypatch): + import agent.secret_scope as secrets + import hermes_cli.inventory as inventory + import hermes_cli.profiles as profiles + import tui_gateway.server as srv + from hermes_constants import get_hermes_home + launch, target = tmp_path / "launch", tmp_path / "target" + launch.mkdir(); target.mkdir() + monkeypatch.setenv("HERMES_HOME", str(launch)) + monkeypatch.setattr(srv, "_hermes_home", launch) + monkeypatch.setattr(srv, "_sessions", {}) + monkeypatch.setattr(profiles, "get_profile_dir", lambda name: target) + def fail(ctx, **kwargs): + assert Path(get_hermes_home()) == target + assert secrets.current_secret_scope().profile_home == target + assert not secrets.current_secret_scope().allow_environment_fallback + raise RuntimeError("inert inventory failure") + monkeypatch.setattr(inventory, "build_model_options_payload", fail) + token = secrets.set_secret_scope({"ANTHROPIC_API_KEY": "inert-outer"}) + previous = secrets.current_secret_scope() + try: + reply = srv.handle_request({"id": "catalog", "method": "model.options", "params": {"profile": "target"}}) + assert reply["error"]["code"] == 5033 + assert "inert inventory failure" in reply["error"]["message"] + assert secrets.current_secret_scope() is previous + assert Path(get_hermes_home()) == launch + finally: + secrets.reset_secret_scope(token) + + +@pytest.mark.parametrize("auth_source", ["saved_record", "credential_pool"]) +def test_canonical_sdk_prefetch_skips_saved_auth_and_pools(tmp_path, monkeypatch, auth_source): + from types import SimpleNamespace + import agent.secret_scope as secrets + import agent.models_dev as mdev + import hermes_cli.auth as auth + import hermes_cli.models as models + import hermes_cli.model_switch as switch + import hermes_cli.providers as providers + slugs = ["anthropic", "gemini", "fireworks", "deepseek"] + monkeypatch.setattr(mdev, "PROVIDER_TO_MODELS_DEV", {}) + monkeypatch.setattr(providers, "HERMES_OVERLAYS", {}) + monkeypatch.setattr(models, "CANONICAL_PROVIDERS", [SimpleNamespace(slug=slug) for slug in [*slugs, "bedrock"]]) + store = {"providers": {slug: {} for slug in slugs}} + if auth_source == "saved_record": + store["providers"]["bedrock"] = {"fixture": True} + monkeypatch.setattr(auth, "_load_auth_store", lambda: store) + monkeypatch.setattr(switch, "_credential_pool_is_usable", lambda slug, **kwargs: slug == "bedrock" and auth_source == "credential_pool") + monkeypatch.setattr(switch, "_PARALLEL_PREFETCH_WORKERS", 1) + monkeypatch.setattr(models, "_load_provider_models_cache", lambda: {}) + monkeypatch.setattr(models, "_credential_fingerprint", lambda slug: "inert-fingerprint") + monkeypatch.setattr(models, "update_provider_cache_entry", lambda *args: None) + fetched = [] + def inert_fetch(slug, **kwargs): + assert slug != "bedrock", "SDK discovery must never enter catalog prefetch" + fetched.append(slug) + return ["fixture-model"] + monkeypatch.setattr(models, "cached_provider_model_ids", inert_fetch) + token = secrets.set_secret_scope(secrets.build_profile_secret_scope(tmp_path, allow_environment_fallback=False)) + try: + selected = switch._collect_authed_provider_slugs({}, {}, []) + assert selected == slugs + switch._prefetch_provider_models_parallel(selected) + assert fetched == slugs + finally: + secrets.reset_secret_scope(token) + + +@pytest.mark.parametrize("multiplex", [False, True]) +@pytest.mark.parametrize("params", [{}, {"profile": "default"}]) +def test_bare_and_explicit_launch_catalog_bind_launch_scope(tmp_path, monkeypatch, multiplex, params): + import agent.secret_scope as secrets + import hermes_cli.inventory as inventory + import hermes_cli.profiles as profiles + import tui_gateway.server as srv + tmp_path.joinpath('.env').write_text('ANTHROPIC_API_KEY=inert-launch\n') + monkeypatch.setenv('HERMES_HOME', str(tmp_path)) + monkeypatch.setattr(secrets, '_MULTIPLEX_ACTIVE', multiplex) + monkeypatch.setattr(srv, '_hermes_home', tmp_path) + monkeypatch.setattr(srv, '_sessions', {}) + monkeypatch.setattr(profiles, 'get_profile_dir', lambda name: tmp_path) + seen = [] + def inventory_launch(ctx, **kwargs): + seen.append(secrets.get_secret('ANTHROPIC_API_KEY')) + return {'providers': []} + monkeypatch.setattr(inventory, 'build_model_options_payload', inventory_launch) + previous = secrets.current_secret_scope() + reply = srv.handle_request({'id': 'catalog', 'method': 'model.options', 'params': params}) + assert 'error' not in reply, reply + assert seen == ['inert-launch'] + assert secrets.current_secret_scope() is previous + + +def test_model_catalog_key_and_fingerprint_use_the_bound_secret_scope(monkeypatch): + import agent.secret_scope as secrets + import hermes_cli.models as models + monkeypatch.setattr(secrets, '_MULTIPLEX_ACTIVE', True) + monkeypatch.setenv('OPENROUTER_API_KEY', 'inert-launch-a') + token = secrets.set_secret_scope({'OPENROUTER_API_KEY': 'inert-target'}) + try: + assert models._resolve_openrouter_api_key() == 'inert-target' + before = models._credential_fingerprint('openrouter') + monkeypatch.setenv('OPENROUTER_API_KEY', 'inert-launch-b') + assert models._credential_fingerprint('openrouter') == before + finally: + secrets.reset_secret_scope(token) + + +def test_swr_real_thread_preserves_target_home_and_scope(tmp_path, monkeypatch): + import threading + import agent.secret_scope as secrets + import hermes_cli.models as models + from hermes_constants import get_hermes_home, set_hermes_home_override, reset_hermes_home_override + monkeypatch.setattr(secrets, '_MULTIPLEX_ACTIVE', True) + done = threading.Event() + seen = [] + def inert_refresh(): + try: + seen.append((Path(get_hermes_home()), secrets.get_secret('OPENROUTER_API_KEY'))) + finally: + done.set() + return None + home_token = set_hermes_home_override(str(tmp_path)) + secret_token = secrets.set_secret_scope({'OPENROUTER_API_KEY': 'inert-target'}) + try: + models._spawn_swr_refresh('pr120-inert-unique', inert_refresh) + assert done.wait(3), 'bounded inert refresh must settle' + assert seen == [(tmp_path, 'inert-target')] + finally: + secrets.reset_secret_scope(secret_token) + reset_hermes_home_override(home_token) + + +def test_actual_openai_catalog_transport_receives_target_key_and_endpoint(monkeypatch): + import agent.secret_scope as secrets + import hermes_cli.models as models + monkeypatch.setattr(secrets, '_MULTIPLEX_ACTIVE', True) + monkeypatch.setenv('OPENAI_API_KEY', 'inert-launch') + monkeypatch.setenv('OPENAI_BASE_URL', 'http://launch.invalid/v1') + monkeypatch.setattr(models, '_get_model_config_dict', lambda: {}) + seen = [] + def inert_transport(api_key, base_url, **kwargs): + seen.append((api_key, base_url)) + return ['inert-target-model'] + monkeypatch.setattr(models, 'fetch_api_models', inert_transport) + token = secrets.set_secret_scope({'OPENAI_API_KEY': 'inert-target', 'OPENAI_BASE_URL': 'http://target.invalid/v1'}) + try: + assert models.provider_model_ids('openai-api') == ['inert-target-model'] + assert seen == [('inert-target', 'http://target.invalid/v1')] + finally: + secrets.reset_secret_scope(token) + + +@pytest.mark.parametrize('target_credential', [None, 'inert-target-aws']) +def test_strict_cross_profile_aws_catalog_never_uses_ambient_sdk_chain(tmp_path, monkeypatch, target_credential): + import agent.secret_scope as secrets + import agent.bedrock_adapter as bedrock + import agent.models_dev as mdev + import hermes_cli.auth as auth + import hermes_cli.models as models + import hermes_cli.model_switch as switch + import hermes_cli.providers as providers + monkeypatch.setattr(secrets, '_MULTIPLEX_ACTIVE', False) + monkeypatch.setenv('AWS_BEARER_TOKEN_BEDROCK', 'inert-launch-aws') + monkeypatch.setattr(auth, 'PROVIDER_REGISTRY', {'bedrock': auth.PROVIDER_REGISTRY['bedrock']}) + monkeypatch.setattr(auth, '_load_auth_store', lambda: {}) + monkeypatch.setattr(switch, '_credential_pool_is_usable', lambda *a, **k: False) + monkeypatch.setattr(mdev, 'PROVIDER_TO_MODELS_DEV', {}) + monkeypatch.setattr(mdev, 'fetch_models_dev', lambda: {}) + monkeypatch.setattr(providers, 'HERMES_OVERLAYS', {'bedrock': providers.HERMES_OVERLAYS['bedrock']}) + monkeypatch.setattr(models, 'CANONICAL_PROVIDERS', []) + monkeypatch.setattr(models, 'get_curated_nous_model_ids', lambda: []) + monkeypatch.setattr(models, 'fetch_ollama_cloud_models', lambda: []) + def denied(*a, **k): + pytest.fail('ambient SDK/provider discovery must not run for a strict catalog') + monkeypatch.setattr(bedrock, 'has_aws_credentials', denied) + monkeypatch.setattr(models, 'cached_provider_model_ids', denied) + if target_credential: + tmp_path.joinpath('.env').write_text(f'AWS_BEARER_TOKEN_BEDROCK={target_credential}\n') + token = secrets.set_secret_scope(secrets.build_profile_secret_scope(tmp_path, allow_environment_fallback=False)) + try: + rows = switch.list_authenticated_providers(current_provider='bedrock', probe_custom_providers=False) + assert any(row['slug'] == 'bedrock' for row in rows) == bool(target_credential) + if target_credential: + row = next(row for row in rows if row['slug'] == 'bedrock') + assert row['models'] == models._PROVIDER_MODELS['bedrock'] + finally: + secrets.reset_secret_scope(token) diff --git a/tui_gateway/methods_complete.py b/tui_gateway/methods_complete.py index d0ec9febf16fd..a163d3eb5a0c6 100644 --- a/tui_gateway/methods_complete.py +++ b/tui_gateway/methods_complete.py @@ -469,6 +469,7 @@ def _(rid, params: dict) -> dict: @method("model.options") def _(rid, params: dict) -> dict: token = None + secret_token = None try: from hermes_cli.inventory import build_model_options_payload from hermes_cli.profiles import get_profile_dir, normalize_profile_name, validate_profile_name @@ -496,11 +497,19 @@ def _(rid, params: dict) -> dict: session_home = Path(session.get("profile_home") or _hermes_home).resolve() if session else None if requested_home is not None and session_home is not None and requested_home != session_home: return _err(rid, 4033, "Model catalog profile does not match the session owner") - home = session_home or requested_home + from hermes_constants import get_hermes_home + home = session_home or requested_home or Path(get_hermes_home()).resolve() if home is not None: if not home.is_dir(): return _err(rid, 4033, "Model catalog session owner does not exist") token = set_hermes_home_override(str(home)) + from agent.secret_scope import build_profile_secret_scope, set_secret_scope + secret_token = set_secret_scope(build_profile_secret_scope( + home, + # Cross-profile catalogs must not borrow launch credentials. + # The launch catalog retains normal environment injection. + allow_environment_fallback=home == Path(_hermes_home).resolve(), + )) agent = session.get("agent") if session else None # Layer agent-session state on top of disk config — once an agent # is spawned, IT owns the live provider/model/base_url. Empty @@ -517,6 +526,9 @@ def _(rid, params: dict) -> dict: except Exception as e: return _err(rid, 5033, str(e)) finally: + if secret_token is not None: + from agent.secret_scope import reset_secret_scope + reset_secret_scope(secret_token) if token is not None: reset_hermes_home_override(token)